From 0cf5dfec2c8313fc2ed2aab8d10bf3dc4b9c283f Mon Sep 17 00:00:00 2001
From: 雾聪 <wucong.lyb@alibaba-inc.com>
Date: 星期四, 14 三月 2024 14:41:49 +0800
Subject: [PATCH] update cmakelist

---
 funasr/download/download_from_hub.py |   25 ++++++++++++++++---------
 1 files changed, 16 insertions(+), 9 deletions(-)

diff --git a/funasr/download/download_from_hub.py b/funasr/download/download_from_hub.py
index 4f0ae74..ef2832f 100644
--- a/funasr/download/download_from_hub.py
+++ b/funasr/download/download_from_hub.py
@@ -13,10 +13,16 @@
         pass
     elif hub == "openai":
         model_or_path = kwargs.get("model")
-        if model_or_path in name_maps_openai:
-            model_or_path = name_maps_openai[model_or_path]
-        kwargs["model_path"] = model_or_path
-    
+        if os.path.exists(model_or_path):
+            # local path
+            kwargs["model_path"] = model_or_path
+            kwargs["model"] = "WhisperWarp"
+        else:
+            # model name
+            if model_or_path in name_maps_openai:
+                model_or_path = name_maps_openai[model_or_path]
+            kwargs["model_path"] = model_or_path
+   
     return kwargs
 
 def download_from_ms(**kwargs):
@@ -24,7 +30,7 @@
     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):
+    if not os.path.exists(model_or_path) and "model_path" not in kwargs:
         try:
             model_or_path = get_or_download_model_dir(model_or_path, model_revision,
                                                       is_training=kwargs.get("is_training"),
@@ -32,7 +38,7 @@
         except Exception as e:
             print(f"Download: {model_or_path} failed!: {e}")
     
-    kwargs["model_path"] = model_or_path
+    kwargs["model_path"] = model_or_path if "model_path" not in kwargs else kwargs["model_path"]
     
     if 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:
@@ -42,9 +48,10 @@
             if "file_path_metas" in conf_json:
                 add_file_root_path(model_or_path, conf_json["file_path_metas"], cfg)
             cfg.update(kwargs)
-            config = OmegaConf.load(cfg["config"])
-            kwargs = OmegaConf.merge(config, cfg)
-        kwargs["model"] = config["model"]
+            if "config" in cfg:
+                config = OmegaConf.load(cfg["config"])
+                kwargs = OmegaConf.merge(config, cfg)
+                kwargs["model"] = config["model"]
     elif os.path.exists(os.path.join(model_or_path, "config.yaml")) and os.path.exists(os.path.join(model_or_path, "model.pt")):
         config = OmegaConf.load(os.path.join(model_or_path, "config.yaml"))
         kwargs = OmegaConf.merge(config, kwargs)

--
Gitblit v1.9.1