Timm库实战:如何用VIT预训练模型快速搭建图像分类任务(附代码示例)
从零到一:用Timm库与Vision Transformer构建你的首个图像分类器
如果你已经厌倦了在PyTorch里手动搭建模型、处理权重加载的繁琐流程,或者对Transformer在视觉领域的应用跃跃欲试却不知从何下手,那么这篇文章正是为你准备的。我们将绕过那些冗长的理论铺垫,直接进入实战环节,手把手教你如何利用timm这个“模型动物园”和强大的Vision Transformer预训练模型,在短短几十分钟内,搭建并运行一个属于自己的图像分类任务。无论你是想快速验证一个想法,还是需要一个可靠的原型作为项目起点,这里提供的思路和代码都能让你事半功倍。
1. 为什么是Timm与Vision Transformer?
在深度学习项目,尤其是计算机视觉任务中,模型的选择和快速实验能力至关重要。过去,我们可能需要从零开始编写模型定义,或者从某个GitHub仓库克隆代码,再手动下载和加载预训练权重。这个过程不仅耗时,还容易因版本、路径等问题出错。
timm库(全称PyTorch Image Models)的出现,极大地简化了这一流程。它由Ross Wightman创建并维护,集成了数百个经过精心实现的、支持预训练权重的计算机视觉模型,涵盖了从经典的ResNet、EfficientNet到前沿的Vision Transformer、Swin Transformer等几乎所有主流架构。其设计哲学是提供一个统一、简洁的接口,让开发者能够用一两行代码就调用最先进的模型。
而Vision Transformer,无疑是近年来计算机视觉领域最具颠覆性的架构之一。它将自然语言处理中取得巨大成功的Transformer模型引入图像领域,通过将图像分割成一个个图像块(patches),并将其视为序列进行处理,完全摒弃了传统的卷积操作。VIT模型在ImageNet等大型数据集上预训练后,展现出了惊人的性能,尤其在数据量充足时,其表现常常超越传统的CNN模型。
将timm与VIT结合,意味着你可以用最少的代码,立即获得一个在超大规模数据集上预训练好的、性能卓越的视觉骨干网络,并轻松地将其适配到你的特定任务上。
1.1 环境准备与Timm安装
在开始之前,确保你的Python环境(建议3.8及以上)已经安装了PyTorch。timm库的安装非常简单:
pip install timm
或者,如果你想安装最新的开发版本:
pip install git+https://github.com/rwightman/pytorch-image-models.git
安装完成后,可以通过以下命令快速验证,并查看已安装的版本和可用的模型数量:
import timm
print(f"Timm version: {timm.__version__}")
print(f"Number of available models: {len(timm.list_models('*'))}")
你会看到一个庞大的模型列表。为了聚焦于VIT,我们可以列出所有包含‘vit’的模型:
vit_models = timm.list_models('*vit*')
print(f"Available VIT models: {len(vit_models)}")
# 打印前10个看看
for model_name in vit_models[:10]:
print(model_name)
这通常会输出像 vit_base_patch16_224、vit_small_patch16_224、vit_large_patch16_224 这样的模型名称。这些命名通常遵循 {架构}_{规模}_{patch大小}_{输入分辨率} 的格式,非常直观。
2. 探索与加载VIT预训练模型
timm.create_model() 函数是通往所有模型的万能钥匙。让我们从一个最简单的例子开始,加载一个标准的VIT-Base模型,并窥探其内部结构。
import torch
import timm
# 使用默认配置加载预训练的VIT-Base模型
model = timm.create_model('vit_base_patch16_224', pretrained=True)
print(type(model))
print(model)
执行这段代码,timm会自动从云端下载对应的预训练权重文件(通常存储在 ~/.cache/torch/hub/checkpoints/ 目录下)。你会看到模型的结构打印出来,包括Patch Embedding层、Transformer Encoder堆叠层和最后的分类头。
注意:首次运行下载权重可能需要一些时间,取决于你的网络速度。下载完成后,权重会被缓存,后续加载会非常快。
每个模型都有一个 default_cfg 属性,这是一个字典,包含了模型的所有默认配置信息,对于理解模型输入输出规格至关重要。
default_cfg = model.default_cfg
print("Model default configuration:")
for key, value in default_cfg.items():
print(f" {key}: {value}")
你需要特别关注以下几个键:
input_size: 模型期望的输入图像尺寸,通常是(3, 224, 224),即3通道,高宽均为224像素。num_classes: 原始预训练任务的类别数,对于ImageNet预训练模型,通常是1000。pool_size: 对于VIT,这通常是None,因为分类token是直接使用的。
2.2 模型选择与权衡
timm提供了多种VIT变体,如何选择?这里有一个简单的对比表格,帮助你根据需求决策:
| 模型名称示例 | 参数量 (约) | 输入分辨率 | 特点与适用场景 |
|---|---|---|---|
vit_tiny_patch16_224 | 5M | 224x224 | 极轻量,适合移动端或对速度要求极高的原型验证。 |
vit_small_patch16_224 | 22M | 224x224 | 平衡型,在速度和精度间取得较好权衡,是快速实验的常用选择。 |
vit_base_patch16_224 | 86M | 224x224 | 最常用的基准模型,在多数任务上能提供优秀的性能,资源消耗相对合理。 |
vit_large_patch16_224 | 307M | 224x224 | 大型模型,精度高,但需要更多显存和计算资源,适合对精度有极致要求且资源充足的场景。 |
vit_base_patch16_384 | 86M | 384x384 | 更高输入分辨率,能捕捉更细粒度特征,通常比224分辨率版本精度更高,但计算量更大。 |
选择建议:
- 快速原型/学习:从
vit_small_patch16_224或vit_base_patch16_224开始。 - 追求最佳精度:在资源允许下尝试
vit_large_patch16_224或更高分辨率的版本。 - 资源受限:考虑
vit_tiny或vit_small。
3. 定制模型以适应你的分类任务
预训练模型是在ImageNet(1000类)上训练的。我们的任务很可能具有不同的类别数。timm提供了极其灵活的方式来修改模型的分类头。
3.1 修改分类头(输出类别数)
最直接的方式是在创建模型时指定 num_classes 参数:
# 假设我们的新任务是猫狗二分类
num_our_classes = 2
model_custom_head = timm.create_model('vit_base_patch16_224',
pretrained=True,
num_classes=num_our_classes)
这一行代码做了两件事:
- 加载了
vit_base_patch16_224在ImageNet上预训练的权重。 - 将模型最后的全连接分类层(head)替换为一个新的、输出维度为
num_our_classes的线性层。注意,新分类头的权重是随机初始化的,而其他所有层的权重都保留了预训练的值。
我们可以验证一下:
print(f"Original model class count: {model.num_classes}") # 输出 1000
print(f"Custom model class count: {model_custom_head.num_classes}") # 输出 2
3.2 更精细的模型定制
有时,我们可能想冻结一部分预训练层(通常是为了防止在小数据集上过拟合,或进行特征提取),只训练新添加的层或最后几层。timm模型支持通过参数名进行筛选。
首先,看看模型有哪些层:
for name, param in model_custom_head.named_parameters():
print(name)
对于VIT,你可能会看到类似 patch_embed.proj.weight, blocks.0.norm1.weight, head.weight 这样的名称。通常,head 就是我们要训练的新分类头。
冻结除分类头外的所有参数:
for name, param in model_custom_head.named_parameters():
if 'head' not in name: # 冻结所有不包含‘head’的参数
param.requires_grad = False
# 检查一下,只有head层的参数需要梯度
trainable_params = sum(p.numel() for p in model_custom_head.parameters() if p.requires_grad)
total_params = sum(p.numel() for p in model_custom_head.parameters())
print(f"Trainable params: {trainable_params:,}")
print(f"Total params: {total_params:,}")
print(f"Trainable %: {100 * trainable_params / total_params:.2f}%")
3.3 数据预处理与增强
预训练模型有其特定的数据预处理流程。timm为每个模型提供了对应的数据变换函数,可以通过 timm.data.create_transform() 来获取。
from timm.data import create_transform
from PIL import Image
import torchvision.transforms as T
# 获取模型对应的默认预处理变换(用于验证/推理)
data_config = timm.data.resolve_model_data_config(model_custom_head)
transform_eval = create_transform(**data_config, is_training=False)
# 对于训练,通常需要更强的数据增强
transform_train = create_transform(**data_config, is_training=True)
# 让我们看看训练变换具体做了什么
print("训练变换包含:")
for t in transform_train.transforms:
print(f" - {t.__class__.__name__}")
# 使用示例
img = Image.open('path/to/your/image.jpg').convert('RGB')
img_tensor_train = transform_train(img) # 应用于训练
img_tensor_eval = transform_eval(img) # 应用于验证/测试
print(f"Transformed image shape: {img_tensor_train.shape}") # 应为 torch.Size([3, 224, 224])
提示:
create_transform会根据模型的default_cfg自动设置正确的 resize 尺寸、裁剪尺寸、归一化均值与标准差(通常是ImageNet的统计值)。这确保了输入数据与模型预训练时的分布一致,对迁移学习的性能至关重要。
4. 构建完整的训练与评估流程
现在,我们将模型、数据、损失函数和优化器组合起来,构建一个简单的训练循环。这里我们使用PyTorch Lightning来简化代码结构,它能让训练逻辑更清晰,并方便地集成日志、检查点等功能。
首先,安装PyTorch Lightning(如果尚未安装):
pip install pytorch-lightning
4.1 定义Lightning数据模块
假设我们有一个简单的图像文件夹数据集,结构如下:
dataset/
├── train/
│ ├── cat/
│ │ ├── cat001.jpg
│ │ └── ...
│ └── dog/
│ ├── dog001.jpg
│ └── ...
└── val/
├── cat/
└── dog/
import pytorch_lightning as pl
from torch.utils.data import DataLoader, Dataset
from torchvision.datasets import ImageFolder
import torch.nn.functional as F
class CatDogDataModule(pl.LightningDataModule):
def __init__(self, data_dir='./dataset', batch_size=32):
super().__init__()
self.data_dir = data_dir
self.batch_size = batch_size
self.transform_train = None
self.transform_val = None
def setup(self, stage=None):
# 在这里创建变换
model = timm.create_model('vit_base_patch16_224', pretrained=False) # 仅用于获取配置
data_config = timm.data.resolve_model_data_config(model)
self.transform_train = create_transform(**data_config, is_training=True)
self.transform_val = create_transform(**data_config, is_training=False)
# 创建数据集
self.train_dataset = ImageFolder(root=f'{self.data_dir}/train', transform=self.transform_train)
self.val_dataset = ImageFolder(root=f'{self.data_dir}/val', transform=self.transform_val)
# 打印类别信息
self.classes = self.train_dataset.classes
print(f"Found {len(self.classes)} classes: {self.classes}")
def train_dataloader(self):
return DataLoader(self.train_dataset, batch_size=self.batch_size, shuffle=True, num_workers=4)
def val_dataloader(self):
return DataLoader(self.val_dataset, batch_size=self.batch_size, shuffle=False, num_workers=4)
4.2 定义Lightning模型模块
class ViTClassifier(pl.LightningModule):
def __init__(self, model_name='vit_base_patch16_224', num_classes=2, learning_rate=1e-4):
super().__init__()
self.save_hyperparameters() # 保存超参数,便于日志记录
self.model = timm.create_model(model_name,
pretrained=True,
num_classes=num_classes)
self.lr = learning_rate
self.criterion = torch.nn.CrossEntropyLoss()
def forward(self, x):
return self.model(x)
def training_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = self.criterion(logits, y)
preds = torch.argmax(logits, dim=1)
acc = (preds == y).float().mean()
# 记录日志
self.log('train_loss', loss, on_step=True, on_epoch=True, prog_bar=True)
self.log('train_acc', acc, on_step=True, on_epoch=True, prog_bar=True)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = self.criterion(logits, y)
preds = torch.argmax(logits, dim=1)
acc = (preds == y).float().mean()
self.log('val_loss', loss, on_epoch=True, prog_bar=True)
self.log('val_acc', acc, on_epoch=True, prog_bar=True)
return {'val_loss': loss, 'val_acc': acc}
def configure_optimizers(self):
# 区分需要梯度更新的参数(我们之前冻结了非head参数)
optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, self.parameters()),
lr=self.lr,
weight_decay=0.01)
# 添加学习率调度器
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer,
T_max=self.trainer.max_epochs,
eta_min=1e-6)
return [optimizer], [scheduler]
4.3 启动训练
def main():
# 初始化数据模块和模型
dm = CatDogDataModule(data_dir='./dataset', batch_size=16)
model = ViTClassifier(model_name='vit_base_patch16_224',
num_classes=len(dm.classes), # 从数据中获取类别数
learning_rate=2e-4)
# 设置训练器
trainer = pl.Trainer(
max_epochs=10, # 训练轮数
accelerator='auto', # 自动选择GPU/CPU
devices=1, # 使用设备数量
log_every_n_steps=10, # 每10步记录一次日志
check_val_every_n_epoch=1, # 每1轮验证一次
enable_progress_bar=True,
)
# 开始训练!
trainer.fit(model, datamodule=dm)
if __name__ == '__main__':
main()
运行这段代码,你将看到一个标准的训练进度条,并实时观察到训练损失、准确率以及验证集上的表现。PyTorch Lightning会自动处理训练循环、验证、日志记录等繁琐细节。
4.4 模型推理与预测
训练完成后,使用模型进行单张图片预测非常简单:
def predict_single_image(image_path, model, class_names, transform):
"""对单张图片进行预测"""
model.eval() # 切换到评估模式
img = Image.open(image_path).convert('RGB')
img_tensor = transform(img).unsqueeze(0) # 增加batch维度
with torch.no_grad():
logits = model(img_tensor)
probabilities = F.softmax(logits, dim=1)
predicted_class_idx = torch.argmax(probabilities, dim=1).item()
confidence = probabilities[0, predicted_class_idx].item()
predicted_class = class_names[predicted_class_idx]
return predicted_class, confidence
# 使用示例
# 假设我们已经从训练好的checkpoint加载了模型 `trained_model`
# 并且有 `class_names` 列表 (例如 ['cat', 'dog'])
# 以及验证用的变换 `transform_eval`
result_class, result_conf = predict_single_image('path/to/test_image.jpg',
trained_model,
dm.classes,
dm.transform_val)
print(f"预测结果: {result_class}, 置信度: {result_conf:.4f}")
5. 进阶技巧与性能优化
掌握了基础流程后,我们可以探索一些进阶技巧来提升效果或效率。
5.1 学习率策略与优化器选择
迁移学习中,学习率的设置尤为关键。一个常见的策略是使用差分学习率,即对预训练的主干网络(backbone)和新添加的分类头(head)使用不同的学习率。分类头通常需要更大的学习率,因为它从随机初始化开始。
from itertools import chain
# 在ViTClassifier的configure_optimizers方法中,可以这样实现:
def configure_optimizers_advanced(self):
# 将参数分为两组:主干网络和分类头
backbone_params = []
head_params = []
for name, param in self.model.named_parameters():
if 'head' in name:
head_params.append(param)
else:
backbone_params.append(param)
# 为两组参数设置不同的学习率
optimizer = torch.optim.AdamW([
{'params': backbone_params, 'lr': self.lr * 0.1}, # 主干网络学习率小
{'params': head_params, 'lr': self.lr} # 分类头学习率大
], weight_decay=0.01)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)
return [optimizer], [scheduler]
5.2 使用混合精度训练加速
混合精度训练(AMP)可以显著减少显存占用并加快训练速度,尤其对于VIT这类大模型。PyTorch Lightning原生支持,只需在Trainer中设置precision=16。
trainer = pl.Trainer(
max_epochs=10,
accelerator='auto',
devices=1,
precision=16, # 启用混合精度训练
# ... 其他参数
)
5.3 模型集成与测试时增强(TTA)
timm库内置了对测试时增强(Test Time Augmentation, TTA)的支持,这是一种通过应用多种数据增强(如水平翻转、多尺度裁剪)到同一张测试图像,并对所有增强版本的预测结果进行平均,来提升模型鲁棒性和准确率的技巧。
# 启用TTA进行预测
from timm.data import resolve_data_config
from timm.data.transforms_factory import create_transform
model_tta = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=2).eval()
config = resolve_data_config(model_tta.default_cfg)
transform = create_transform(**config, is_training=False)
# 假设我们有一个批次的数据
batch = torch.randn(4, 3, 224, 224)
# 普通推理
with torch.no_grad():
output_normal = model_tta(batch)
# 使用TTA推理
with torch.no_grad():
output_tta = timm.utils.tta.tta_model(model_tta, batch, aug_fn=timm.utils.tta.tta.ClassificationTTAWrapper)
print(f"Normal output shape: {output_normal.shape}")
print(f"TTA output shape: {output_tta.shape}") # 应该和普通输出一致
在实际项目中,我发现对于数据量有限或类别间差异细微的任务,即使只进行简单的水平翻转TTA,也能带来0.5%到2%的稳定精度提升,而计算开销增加有限,性价比很高。
5.4 特征提取与可视化
除了直接用于分类,VIT的中间层特征也极具价值,可以用于可视化、相似性检索或其他下游任务。timm模型可以方便地获取中间特征。
# 创建一个不包含分类头的模型,用于特征提取
model_features = timm.create_model('vit_base_patch16_224',
pretrained=True,
num_classes=0) # num_classes=0 移除分类头
model_features.eval()
# 获取全局特征(通常是[CLS] token对应的输出)
with torch.no_grad():
features = model_features(batch) # 输出形状: (batch_size, feature_dim)
print(f"Extracted features shape: {features.shape}")
# 如果你想获取特定Transformer block的输出,可以注册钩子(hook)
features = {}
def get_features(name):
def hook(model, input, output):
features[name] = output.detach()
return hook
# 注册到第6个block之后
model_features.blocks[5].register_forward_hook(get_features('block_5'))
with torch.no_grad():
_ = model_features(batch)
print(f"Features from block 5 shape: {features['block_5'].shape}")
通过这种方式,你可以深入探索模型“看到”了什么,为模型解释性和进一步的应用开发提供基础。整个流程从环境搭建、模型探索、定制化、训练到进阶优化,形成了一套完整的实战指南。最关键的是,借助timm库,这一切都变得异常简洁和高效,让你能更专注于任务本身和算法创新,而不是陷于工程实现的泥潭。
更多推荐
所有评论(0)