From e30a17cf4e715b3d139fa1e0ba01cda1bcf0f884 Mon Sep 17 00:00:00 2001
From: shixian.shi <shixian.shi@alibaba-inc.com>
Date: 星期三, 10 一月 2024 11:23:41 +0800
Subject: [PATCH] update funasr-onnx

---
 funasr/download/download_from_hub.py |   21 ++++++++++-----------
 1 files changed, 10 insertions(+), 11 deletions(-)

diff --git a/funasr/download/download_from_hub.py b/funasr/download/download_from_hub.py
index abf3ba0..73578f2 100644
--- a/funasr/download/download_from_hub.py
+++ b/funasr/download/download_from_hub.py
@@ -7,22 +7,20 @@
 def download_model(**kwargs):
 	model_hub = kwargs.get("model_hub", "ms")
 	if model_hub == "ms":
-		kwargs = download_fr_ms(**kwargs)
+		kwargs = download_from_ms(**kwargs)
 	
 	return kwargs
 
-def download_fr_ms(**kwargs):
+def download_from_ms(**kwargs):
 	model_or_path = kwargs.get("model")
 	if model_or_path in name_maps_ms:
 		model_or_path = name_maps_ms[model_or_path]
 	model_revision = kwargs.get("model_revision")
 	if not os.path.exists(model_or_path):
-		model_or_path = get_or_download_model_dir(model_or_path, model_revision, is_training=kwargs.get("is_training"))
+		model_or_path = get_or_download_model_dir(model_or_path, model_revision, is_training=kwargs.get("is_training"), check_latest=kwargs.get("kwargs", True))
 	
 	config = os.path.join(model_or_path, "config.yaml")
 	if os.path.exists(config) and os.path.exists(os.path.join(model_or_path, "model.pb")):
-		# config = os.path.join(model_or_path, "config.yaml")
-		# assert os.path.exists(config), "{} is not exist!".format(config)
 		cfg = OmegaConf.load(config)
 		kwargs = OmegaConf.merge(cfg, kwargs)
 		init_param = os.path.join(model_or_path, "model.pb")
@@ -42,18 +40,19 @@
 		assert os.path.exists(os.path.join(model_or_path, "configuration.json"))
 		with open(os.path.join(model_or_path, "configuration.json"), 'r', encoding='utf-8') as f:
 			conf_json = json.load(f)
-			config = os.path.join(model_or_path, conf_json["model"]["model_config"])
+			config = os.path.join(model_or_path, conf_json["model_config"])
 			cfg = OmegaConf.load(config)
 			kwargs = OmegaConf.merge(cfg, kwargs)
-			init_param = os.path.join(model_or_path, conf_json["model"]["model_name"])
+			init_param = os.path.join(model_or_path, conf_json["model_file"])
 			kwargs["init_param"] = init_param
 		kwargs["model"] = cfg["model"]
 	return OmegaConf.to_container(kwargs, resolve=True)
 
 def get_or_download_model_dir(
-                              model,
-                              model_revision=None,
-							  is_training=False,
+		model,
+		model_revision=None,
+		is_training=False,
+		check_latest=True,
 	):
 	""" Get local model directory or download model if necessary.
 
@@ -69,7 +68,7 @@
 	
 	key = Invoke.LOCAL_TRAINER if is_training else Invoke.PIPELINE
 	
-	if os.path.exists(model):
+	if os.path.exists(model) and check_latest:
 		model_cache_dir = model if os.path.isdir(
 			model) else os.path.dirname(model)
 		try:

--
Gitblit v1.9.1