Hugging Face Transformers timm 集成

Hugging Face 已经引入了 TimmWrapper,这是一个工具,使得 PyTorch Image Models (timm) 库中的任何模型都可以直接在 🤗 transformers 生态系统中使用。此集成使用户能够在利用 transformers 高级 API 进行推理、量化和微调的同时,发挥 timm 在计算机视觉模型方面的丰富收藏。

通过 TimmWrapper 实现无缝集成

TimmWrappertimm 库和 transformers 库之间架起了桥梁,使得 timm 模型与标准的 Hugging Face 工作流兼容。此集成提供了几个关键的技术优势:

  • Pipeline API 支持timm 模型可以插入到高级的 transformers pipeline 中,以实现流畅的推理。
  • Auto Class 兼容性:可以使用 AutoModelForImageClassificationAutoImageProcessor 加载模型,抽象了模型和处理器加载的复杂性。
  • Trainer API 集成:用户可以使用标准的 Trainer API 微调 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),这使得在消费级硬件上进行高效训练成为可能。

实际实施示例

  • 图像分类:可以使用 pipeline API 加载如 mobilenetv4_conv_medium (该模型没有原生 transformers 实现)的模型以进行即时推理。
  • 交互式演示:该集成可与 Gradio 配合使用,使开发者能够使用微调的 timm ViT 模型构建食物分类器网页应用。
  • 模型加载:可以使用 AutoImageProcessorAutoModelForImageClassification 直接从 Hugging Face Hub 加载 timm 检查点。

Sources