Hugging Face Transformers timm 集成
Hugging Face 已经引入了 TimmWrapper,这是一个工具,使得 PyTorch Image Models (timm) 库中的任何模型都可以直接在 🤗 transformers 生态系统中使用。此集成使用户能够在利用 transformers 高级 API 进行推理、量化和微调的同时,发挥 timm 在计算机视觉模型方面的丰富收藏。
通过 TimmWrapper 实现无缝集成
TimmWrapper 在 timm 库和 transformers 库之间架起了桥梁,使得 timm 模型与标准的 Hugging Face 工作流兼容。此集成提供了几个关键的技术优势:
- Pipeline API 支持:
timm模型可以插入到高级的transformerspipeline 中,以实现流畅的推理。 - Auto Class 兼容性:可以使用
AutoModelForImageClassification和AutoImageProcessor加载模型,抽象了模型和处理器加载的复杂性。 - Trainer API 集成:用户可以使用标准的
TrainerAPI 微调timm模型,在不同模型架构之间保持一致的工作流。 - 往返兼容性:在
transformers生态系统中微调的模型可以使用timm.create_model('hf-hub:my_org/my_fine_tuned_model', pretrained=True)加载回timm。
优化的推理和部署
此集成使得可以在 timm 模型上使用 transformers 生态系统中的高级优化技术:
使用 bitsandbytes 进行量化
用户可以使用 BitsAndBytesConfig 对任何 timm 模型进行量化以实现高效推理。在使用 ViT 基础模型的提供的示例中,8 位量化将模型大小从 346.27 MB 减少到 88.20 MB(减少 74.53%),同时保持了几乎相同的准确率(特定标签为 0.33% 对比 0.35%)。
使用 torch.compile 进行加速
timm 集成与 torch.compile(在 PyTorch 2.0 中引入)完全兼容,允许用户通过仅一行代码编译模型来实现更快的推理时间。
灵活的微调选项
TimmWrapper 支持标准和参数高效微调(PEFT)两种方法:
标准微调
timm 模型可以使用 Trainer 类在自定义数据集上进行微调,该类管理训练循环、日志和评估。这与原生 transformers 模型使用的工作流完全相同。
LoRA(低秩适应)
通过 PEFT 库,用户可以将 LoRA 应用到 timm 模型上,仅训练极少数参数。在一个示例中,ViT 模型仅有 0.77% 的参数可训练(在 86,543,818 个总参数中,可训练参数为 667,493),这使得在消费级硬件上进行高效训练成为可能。
实际实施示例
- 图像分类:可以使用
pipelineAPI 加载如mobilenetv4_conv_medium(该模型没有原生transformers实现)的模型以进行即时推理。 - 交互式演示:该集成可与 Gradio 配合使用,使开发者能够使用微调的
timmViT 模型构建食物分类器网页应用。 - 模型加载:可以使用
AutoImageProcessor和AutoModelForImageClassification直接从 Hugging Face Hub 加载timm检查点。