游雁
2024-01-13 835369d6315e96c1820326ed11ea4b999793720f
funasr/download/download_from_hub.py
@@ -7,17 +7,18 @@
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))
   kwargs["model_path"] = model_or_path
   
   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")):
@@ -36,6 +37,8 @@
      kwargs["model"] = cfg["model"]
      if os.path.exists(os.path.join(model_or_path, "am.mvn")):
         kwargs["frontend_conf"]["cmvn_file"] = os.path.join(model_or_path, "am.mvn")
      if os.path.exists(os.path.join(model_or_path, "jieba_usr_dict")):
         kwargs["jieba_usr_dict"] = os.path.join(model_or_path, "jieba_usr_dict")
   else:# configuration.json
      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:
@@ -49,9 +52,10 @@
   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.
@@ -67,7 +71,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: