downloadModel.py 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125
  1. #!/usr/bin/env python3
  2. """
  3. 从 Hugging Face 下载模型到本地 model 文件夹
  4. 支持镜像源、自动按模型名创建子文件夹
  5. """
  6. import argparse
  7. import os
  8. from pathlib import Path
  9. # 项目根目录
  10. PROJECT_ROOT = Path(__file__).parent.parent.resolve()
  11. # 默认模型保存目录:项目根目录/model/weights
  12. DEFAULT_MODEL_DIR = PROJECT_ROOT / "model" / "weights"
  13. def download_model(model_id: str, output_dir: str = None, use_mirror: bool = True):
  14. """
  15. 从 Hugging Face 下载模型
  16. Args:
  17. model_id: 模型 ID,格式为 "username/model_name"
  18. output_dir: 输出根目录,默认为项目根目录/model/weights
  19. use_mirror: 是否使用镜像源,默认 True
  20. """
  21. if output_dir is None:
  22. output_dir = DEFAULT_MODEL_DIR
  23. try:
  24. from huggingface_hub import snapshot_download
  25. except ImportError:
  26. print("错误:未安装 huggingface_hub 库")
  27. print("请运行:pip install huggingface_hub")
  28. return
  29. # 提取模型名作为子文件夹名
  30. model_name = model_id.split("/")[-1]
  31. output_path = Path(output_dir).resolve() / model_name
  32. output_path.mkdir(parents=True, exist_ok=True)
  33. print(f"开始下载模型:{model_id}")
  34. print(f"保存路径:{output_path}")
  35. if use_mirror:
  36. print("使用镜像源:hf-mirror.com")
  37. print("-" * 50)
  38. try:
  39. # 设置镜像源环境变量
  40. if use_mirror:
  41. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  42. # 下载整个模型仓库
  43. downloaded_path = snapshot_download(
  44. repo_id=model_id,
  45. local_dir=str(output_path),
  46. local_dir_use_symlinks=False, # Windows 兼容性
  47. )
  48. print("-" * 50)
  49. print(f"✓ 模型下载完成!")
  50. print(f"模型位置:{downloaded_path}")
  51. # 列出下载的文件
  52. files = list(output_path.iterdir())
  53. print(f"\n下载的文件 ({len(files)} 个):")
  54. for f in files:
  55. if f.is_file():
  56. size = f.stat().st_size
  57. if size > 1024 * 1024 * 1024:
  58. size_str = f"{size / (1024*1024*1024):.2f} GB"
  59. elif size > 1024 * 1024:
  60. size_str = f"{size / (1024*1024):.2f} MB"
  61. elif size > 1024:
  62. size_str = f"{size / 1024:.2f} KB"
  63. else:
  64. size_str = f"{size} B"
  65. print(f" - {f.name} ({size_str})")
  66. except Exception as e:
  67. print(f"下载失败:{e}")
  68. raise
  69. def main():
  70. parser = argparse.ArgumentParser(
  71. description="从 Hugging Face 下载模型到本地",
  72. formatter_class=argparse.RawDescriptionHelpFormatter,
  73. epilog=f"""
  74. 示例:
  75. python downloadModel.py hf-internal-testing/tiny-random-BertModel
  76. python downloadModel.py bert-base-chinese --no-mirror
  77. python downloadModel.py Qwen/Qwen2.5-7B-Instruct -o ./my_models
  78. 默认保存目录:{DEFAULT_MODEL_DIR}
  79. """
  80. )
  81. parser.add_argument(
  82. "model_id",
  83. type=str,
  84. help="Hugging Face 模型 ID (格式:username/model_name)"
  85. )
  86. parser.add_argument(
  87. "-o", "--output",
  88. type=str,
  89. default=None,
  90. help=f"模型保存根目录 (默认:{DEFAULT_MODEL_DIR})"
  91. )
  92. parser.add_argument(
  93. "--no-mirror",
  94. action="store_true",
  95. help="不使用镜像源,直接使用 Hugging Face 官方源"
  96. )
  97. args = parser.parse_args()
  98. download_model(args.model_id, args.output, use_mirror=not args.no_mirror)
  99. if __name__ == "__main__":
  100. main()