diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index 0b8864f..7dfc798 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -54,7 +54,24 @@ impl FunAsrNanoGenerateModel { let model_list = find_type_files(path, "pt")?; let mut dict_to_hashmap = HashMap::new(); for m in model_list { - let dict = read_all_with_key(m, Some("state_dict"))?; + let dict = match read_all_with_key(m.clone(), Some("state_dict")) { + Ok(dict) => dict, + Err(e) => { + println!( + "model read_all_with_key {} get state_dict err: {}, use None try again", + &m, e + ); + match read_all_with_key(m.clone(), None) { + Ok(dict) => dict, + Err(e) => { + return Err(anyhow!(format!( + "model read_all_with_key({}, None): e: {}", + &m, e + ))); + } + } + } + }; for (k, v) in dict { dict_to_hashmap.insert(k, v); } diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 4fabe06..a70aa56 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -164,7 +164,8 @@ fn fun_asr_nano_weight() -> Result<()> { let mut dict_to_hashmap = HashMap::new(); // let mut dtype = candle_core::DType::F32; for m in model_list { - let dict = read_all_with_key(m, Some("state_dict"))?; + // let dict = read_all_with_key(m, Some("state_dict"))?; + let dict = read_all_with_key(m, None)?; // dtype = dict[0].1.dtype(); for (k, v) in dict { if k.contains("model") {