从零到一:用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_2245M224x224极轻量,适合移动端或对速度要求极高的原型验证。
vit_small_patch16_22422M224x224平衡型,在速度和精度间取得较好权衡,是快速实验的常用选择。
vit_base_patch16_22486M224x224最常用的基准模型,在多数任务上能提供优秀的性能,资源消耗相对合理。
vit_large_patch16_224307M224x224大型模型,精度高,但需要更多显存和计算资源,适合对精度有极致要求且资源充足的场景。
vit_base_patch16_38486M384x384更高输入分辨率,能捕捉更细粒度特征,通常比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)

这一行代码做了两件事:

  1. 加载了 vit_base_patch16_224 在ImageNet上预训练的权重。
  2. 将模型最后的全连接分类层(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库,这一切都变得异常简洁和高效,让你能更专注于任务本身和算法创新,而不是陷于工程实现的泥潭。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐