十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

如何用 create_model 加载 Hugging Face Hub 上自定义训练或微调过的 timm 模型权重?

如何用 create_model 加载 Hugging Face Hub 上自定义训练或微调过的 timm 模型权重? 如何用 create_model 加载 Hugging Face Hub 上自定义训练或微调过的 timm 模型权重【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models你训练或微调了一个timm模型想让它脱离本地环境在任何装有timm的机器上一行代码就能把权重拉回来。timm内置了 Hugging Face Hub 集成create_model支持带来源前缀的模型名可以直接从 Hub 仓库加载你推送上去的权重和模型配置。本文给出这条加载路径的完整操作环境准备、模型名写法、加载代码、结果验证以及常见报错的对应处理。准备条件安装timm见 安装文档pip install timm安装后可以用文档给出的命令确认timm可用输出为文档示例列出前 5 个带预训练权重的注册模型名python -c from timm import list_models; print(list_models(pretrainedTrue)[:5])安装huggingface_hub这是 Hub 加载的硬依赖来源pip install huggingface_hub如果 Hub 仓库里放的是.safetensors权重新文件model.safetensors还需要safetensors包缺失时源码中的报错会提示pip install safetensors见 timm/models/_hub.py。权重如何到达 Hub 的timm提供timm.models.push_to_hf_hub会把模型权重和config.json一起推送到 Hub仓库名为your-username/repo_name。官方文档中的示例流程文档示例 import timm model timm.create_model(resnet18, pretrainedTrue, num_classes4) # 此处省略训练 / 微调 model_cfg dict(label_names[a, b, c, d]) timm.models.push_to_hf_hub(model, resnet18-random, model_configmodel_cfg)执行后模型位于你的 Hub 用户名下的resnet18-random仓库。推送私有仓库或首次使用时先按 hf_hub 文档完成认证终端执行huggingface-cli loginnotebook 里用huggingface_hub.notebook_login。模型名怎么写create_model的model_name支持 URI 风格的来源前缀hf-hub:表示从 Hub 加载timm/models/_factory.pyhf-hub:owner/repo_name # 从 Hub 仓库加载 hf-hub:owner/repo_namerevision # 指定 revision分支或 commitrevision是可选的用紧跟在仓库名后面例如hf-hub:timm/resnet18.a1_in1kmain该形式在 tests/test_factory.py 中有测试覆盖一个模型名里只允许出现一个来标识 revision。hf_hub:下划线是已弃用的旧前缀仍可用会被映射为hf-hub前缀大小写不敏感测试中HF-HUB:也成立。新代码建议统一写hf-hub:。前缀之后的部分被当作不透明标识符原样透传不做 URL 解析。用 create_model 加载 Hub 上的模型文档给出的加载方式hfdocs/source/hf_hub.mdx仓库 id 替换为你自己的 model_reloaded timm.create_model(hf_hub:nateraw/resnet18-random, pretrainedTrue)写成当前推荐的hf-hub:前缀等价于model_reloaded timm.create_model(hf-hub:nateraw/resnet18-random, pretrainedTrue)其中nateraw/resnet18-random就是owner/repo_name形式的 Hub 仓库 id改成你推送的仓库即可。pretrainedTrue在这里表示“按config.json下载并加载该仓库的预训练权重”。如果不想把权重下到默认缓存目录可以传cache_dir覆盖 Hugging Face Hub 与 Torch checkpoint 的缓存位置timm/models/_factory.py 文档中的用法示例 model create_model(hf-hub:timm/vit_small_patch14_dinov2.lvd142m, pretrainedTrue, cache_dir/data/my-models) # 数据将存放到 /data/my-models/models--timm--vit_small_patch14_dinov2.lvd142m/注意使用hf-hub:前缀时不要再传pretrained_cfg参数源码里有一条断言明确禁止pretrained_cfg should not be set when sourcing model from Hugging Face Hub.。create_model 从 Hub 仓库读取了什么了解读取逻辑有助于核对结果timm/models/_hub.py先从仓库下载config.json并解析architecture字段决定实例化哪个timm注册模型顶层的num_classes、label_names、label_descriptions会写入该模型的pretrained_cfgmodel_args里的键值会作为默认参数传入模型构建。然后下载权重默认文件名pytorch_model.bin若已安装safetensors会优先检查同名.safetensors替代文件如model.safetensors存在则改用 safetensors 加载。因此加载出来的模型结构、分类数、标签名都来自推送时生成的config.json而不是本地代码里的配置。验证加载结果加载完成后做三项检查。文档示例中推送的模型是num_classes4、label_names[a, b, c, d]的 resnet18所以重新加载后以下输出为文档示例推导出的预期形态# 1. 分类数应与推送时一致 print(model_reloaded.num_classes) # 文档示例中的模型为 4 # 2. pretrained_cfg 应体现来源包含 hf_hub_id 和 source: hf-hub # 以及推送时写入的 label_names print(model_reloaded.pretrained_cfg) # 3. 跑一次前向确认权重真正生效 import torch x torch.randn(1, 3, 224, 224) y model_reloaded(x) print(y.shape)create_model返回的模型默认处于 train 模式做推理前先调用.eval()quickstart 文档 中的提示。推理时建议按模型自身的数据配置构建 transform避免用错预处理 model_reloaded model_reloaded.eval() transform timm.data.create_transform(**timm.data.resolve_data_config(model_reloaded.pretrained_cfg))常见报错与处理报错含义与处理RuntimeError: Hugging Face hub model specified but package not installed. Runpip install huggingface_hub.未安装huggingface_hub执行pip install huggingface_hub后重试ValueError: Model name xxx has no source prefix but looks like a Hugging Face Hub repo id or a local path. Use hf-hub:xxx to load from the Hub or local-dir:xxx to load from a local folder.模型名含/等路径字符却没有来源前缀。timm会拒绝把owner/repo静默当成同名注册模型补上hf-hub:前缀即可ValueError: Unknown model source xxx ...前缀不是合法来源合法来源只有hf-hub与local-dir两种断言失败pretrained_cfg should not be set when sourcing model from Hugging Face Hub.用前缀加载时同时传了pretrained_cfg去掉该参数断言失败hf_hub id should only contain one character to identify revision.只用于标识 revision一个模型名里最多出现一次另外两条边界如果权重文件是.safetensors且未安装safetensors包加载会直接断言失败提示pip install safetensors。前缀只认hf-hub和local-dir两种local-dir:用于从本地目录读path/config.json加权重文件是 Hub 之外的另一条加载路径本文不展开。【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表