图像识别与分类
ResNet / EfficientNet / ViT / ConvNeXt 对比、图像分类实战、数据增强 Albumentations
图像分类任务概述
图像分类是计算机视觉中最基础也是最核心的任务之一,其目标是将输入图像分配到预定义的类别标签中。给定一张图像,模型需要输出该图像所属类别的概率分布。尽管这一任务对人类而言几乎毫无难度,但对机器来说却充满了挑战。
图像分类面临的主要挑战包括:类内差异大(同一类别的物体在不同角度、光照、姿态下外观差异显著)、类间相似度高(不同类别的物体可能具有相似的外观)、背景干扰(复杂背景会分散模型的注意力)、遮挡与截断、以及尺度多样性(同一物体在图像中的大小可能相差很大)。
ImageNet 是目前图像分类领域最具影响力的基准数据集,包含约 128 万张训练图像和 5 万张验证图像,涵盖 1000 个类别。自 2012 年 AlexNet 在 ImageNet 上取得突破性进展以来,该数据集上的 Top-1 和 Top-5 准确率已成为衡量图像分类模型性能的行业标准。
- Top-1 准确率:模型预测的最高概率类别与真实标签一致的样本占比。
- Top-5 准确率:模型预测的前五个最高概率类别中包含真实标签的样本占比。ImageNet 原始评测中主要使用 Top-5 准确率,因为部分类别之间存在语义模糊性。
经典 CNN 架构
ResNet
ResNet(Residual Network)由何恺明等人于 2015 年提出,是深度学习历史上最具影响力的架构之一。在 ResNet 之前,人们普遍认为更深的网络应该能取得更好的性能,但实验发现网络深度增加到一定程度后,训练误差反而上升,这就是所谓的退化问题(Degradation Problem)——它并非过拟合,而是深层网络难以优化。
ResNet 的核心创新是残差学习(Residual Learning)。传统网络学习的是从输入到输出的直接映射 H(x),而 ResNet 改为学习残差映射 F(x) = H(x) - x,然后通过跳跃连接(Skip Connection)将输入 x 直接加到残差映射上,即输出为 F(x) + x。这样,即使残差映射 F(x) 学习到零,输出至少能复制输入,从而保证深层网络不会比浅层网络更差。从梯度传播的角度看,跳跃连接为梯度提供了高速公路,有效缓解了梯度消失问题。
ResNet 的另一重要设计是 Bottleneck 模块,主要用于深层 ResNet(如 ResNet-50 及更深版本)。Bottleneck 模块采用 1x1 -> 3x3 -> 1x1 的卷积层级结构:第一个 1x1 卷积降维,3x3 卷积进行空间特征提取,最后一个 1x1 卷积恢复维度。这种设计在控制计算量的同时引入了更多非线性变换。
| 模型变体 | 层数 | 参数量 | ImageNet Top-1 准确率 |
|---|---|---|---|
| ResNet-18 | 18 | 11.7M | 69.8% |
| ResNet-34 | 34 | 21.8M | 73.3% |
| ResNet-50 | 50 | 25.6M | 76.1% |
| ResNet-101 | 101 | 44.5M | 77.4% |
| ResNet-152 | 152 | 60.2M | 78.6% |
EfficientNet
EfficientNet 由 Google 团队在 2019 年提出,其核心贡献是发现了**复合缩放(Compound Scaling)**策略。以往的模型缩放通常只调整一个维度——要么增加深度(层数),要么增加宽度(通道数),要么增加输入分辨率。EfficientNet 的研究表明,这三个维度之间存在相互依赖关系:当分辨率提高时,网络需要更深的层来捕捉更大范围内的特征,同时需要更宽的通道来捕捉更细粒度的模式。
复合缩放使用一个复合系数 φ 来同时缩放深度(d = α^φ)、宽度(w = β^φ)和分辨率(r = γ^φ),其中 α、β、γ 是通过网格搜索从基准网络(EfficientNet-B0)上确定的最佳比例。约束条件为 α · β^2 · γ^2 ≈ 2,确保总计算量 (FLOPS) 大致随 2^φ 增长。
EfficientNet 的基础构建模块是 MBConv(Mobile Inverted Bottleneck Convolution),源自 MobileNetV2。MBConv 采用倒置残差结构:先通过 1x1 卷积扩展通道数(通常是 4 倍或 6 倍),再使用深度可分离卷积(Depthwise Separable Convolution)进行空间特征提取,最后通过 1x1 卷积压缩通道数。此外,MBConv 中还引入了 Squeeze-and-Excitation(SE)注意力机制,通过学习通道间的依赖关系来增强重要特征的响应。
EfficientNet-B0 到 B7 系列在 ImageNet 上以远小于 ResNet 和 ResNeXt 的参数量和 FLOPs 达到了同级别甚至更高的准确率。例如,EfficientNet-B7 在 600M 参数量下达到 84.4% 的 Top-1 准确率,大幅刷新了当时的 SOTA。
Vision Transformer(ViT)
Vision Transformer(ViT)由 Google Research 在 2020 年提出,首次将 Transformer 架构成功应用于图像分类任务,并证明在足够大的数据集(如 JFT-300M)上预训练后,纯 Transformer 可以超越最先进的 CNN 架构。
ViT 的核心思路是将图像处理为序列数据,使其适配 Transformer 的标准输入格式。具体步骤如下:
- 图像切块(Patch Embedding):将输入图像(如 224x224x3)切分为固定大小的 patch(如 16x16),得到 (224/16)^2 = 196 个 patch。每个 patch 通过线性投影展平为向量,得到 patch embedding。所有 patch embedding 组成一个序列,类似于 NLP 中的 token 序列。
- 位置编码(Position Embedding):由于 Transformer 本身不具备序列位置感知能力,需要为每个 patch embedding 添加可学习的位置编码。
- [CLS] Token:在序列头部插入一个特殊可学习的 [CLS] token,其对应输出经过 MLP Head 后用于分类预测。这一设计借鉴了 BERT 的做法。
- 标准 Transformer Encoder:由多层 Multi-Head Self-Attention(MSA)和 MLP 组成,每层后接 Layer Normalization 和残差连接。
ViT 与 CNN 的核心差异在于感受野。CNN 通过堆叠小卷积核逐步扩大感受野,本质上是局部优先的——每个卷积核只关注局部邻域。而 ViT 的自注意力机制从第一层开始就能建立 全局 的 patch 间依赖关系,每个 patch 都可以直接关注到图像中任意其他 patch。
这种全局建模能力使 ViT 能更好地捕捉长距离空间依赖,但也带来了新的问题。CNN 天然具有两种归纳偏置(Inductive Bias):平移等变性(一个特征不论出现在图像哪个位置都能被检测到)和局部性(相邻像素更相关)。ViT 没有这些归纳偏置,所有空间关系都必须从数据中学习。这意味着 ViT 在小规模数据集上难以训练——如果不经过大规模预训练,其性能通常不如同等规模的 CNN。然而,当预训练数据量足够大时,ViT 能够超越 CNN 的归纳偏置限制,学习到更通用、更强大的视觉表征。
ConvNeXt
ConvNeXt 由 Facebook AI Research(现 Meta AI)于 2022 年提出,旨在探索一个问题:在 Transformer 时代,纯 CNN 架构是否仍具竞争力? ConvNeXt 的答案是肯定的——通过系统性地借鉴 Swin Transformer 的设计理念对标准 ResNet 进行现代化改造,ConvNeXt 在 ImageNet 上取得了与 Swin Transformer 相当甚至更优的性能。
ConvNeXt 从 ResNet-50 出发,逐步引入以下改进策略:
- 训练策略现代化:使用 AdamW 优化器、数据增强(MixUp、CutMix、RandAugment)、正则化(Stochastic Depth、Label Smoothing)等 ViT 常用的训练技巧,仅此一项就将 ResNet-50 的 ImageNet Top-1 准确率从 76.1% 提升至 78.8%。
- 宏观架构调整:将 ResNet 的 stage(4 个阶段)计算比例从 (3:4:6:3) 调整为 (3:3:9:3),模仿 Swin-T 的设计,增加 Stage 3 的深度。同时将下采样从 conv3x3 stride 2 改为独立的 Patchify Stem(4x4, stride 4 卷积)。
- 引入深度可分离卷积(Depthwise Conv):将 3x3 卷积替换为 depthwise conv,与 Transformer 的 MSA 中逐点运算一致。depthwise conv 仅在空间维度操作,参数量和计算量远小于标准卷积。
- 倒置瓶颈(Inverted Bottleneck):借鉴 MobileNetV2 和 Transformer MLP 的设计,将 ConvNeXt Block 中的通道扩展放在中间层(即 h_dim -> 4h_dim -> h_dim),而非 ResNet 的压缩瓶颈(h_dim -> h_dim/4 -> h_dim)。这与 ViT MLP 的扩展比例一致。
- 大卷积核:将 depthwise conv 的核大小从 3x3 增加到 7x7,仿照 Swin Transformer 的 7x7 窗口大小,以扩大感受野。
- Layer Normalization:将 BatchNorm 替换为 LayerNorm,少用 BN,仅保留一个 LN 在 depthwise conv 之后。
ConvNeXt Block 结构:LN -> 7x7 Depthwise Conv -> GELU -> 1x1 Conv(4x 扩展)-> GELU -> 1x1 Conv(压缩),每个操作后均有残差连接。
在 ImageNet-1K 上,ConvNeXt-Base(89M 参数,15.0G FLOPs)达到 84.1% Top-1 准确率,ConvNeXt-Large(198M 参数,34.4G FLOPs)达到 84.3%,ConvNeXt-XLarge(350M 参数,59.0G FLOPs)达到 84.8%,与同量级的 Swin Transformer 和 ViT 变体不相上下。ConvNeXt 证明了经过现代化设计的纯 CNN 依然具备与 Transformer 同等的建模能力。
架构对比表
下表汇总了四种代表性架构在 ImageNet-1K 上的性能对比(均为 Base 级别模型,输入分辨率 224x224):
| 模型 | 参数量 | FLOPs | Top-1 准确率 | 推理速度(img/s, V100) |
|---|---|---|---|---|
| ResNet-50 | 25.6M | 4.1G | 76.1% | ~1230 |
| EfficientNet-B4 | 19.3M | 4.2G | 82.9% | ~350 |
| ViT-B/16 | 86.6M | 17.6G | 77.9%(ImageNet-1K only) | ~850 |
| ViT-B/16(JFT-300M pretrain) | 86.6M | 17.6G | 84.2% | ~850 |
| ConvNeXt-Base | 88.6M | 15.4G | 84.1% | ~650 |
说明:ViT-B/16 在仅 ImageNet-1K 上训练时准确率较低,需要大规模预训练才能发挥优势。EfficientNet-B4 在相近 FLOPs 下准确率较高,但深度可分离卷积的硬件利用率较低,实际推理速度偏慢。ConvNeXt 在参数量与 ViT 相当的情况下达到了接近的性能,且推理速度具有优势。
数据增强 Albumentations
数据增强是图像分类训练中不可或缺的一环,尤其在数据量有限时,合理的数据增强可以显著提升模型的泛化能力和鲁棒性。
Albumentations 是一个高性能图像增强库,基于 NumPy 和 OpenCV 实现,其核心优势在于:丰富的增强操作、优化的执行速度(支持多线程)、以及统一简洁的 API。
常用增强操作
| 增强操作 | 说明 | 适用场景 |
|---|---|---|
| RandomCrop | 随机裁剪图像,保留局部特征 | 通用,缓解过拟合 |
| HorizontalFlip | 水平翻转,概率执行 | 通用,对称性数据 |
| ColorJitter | 随机调整亮度、对比度、饱和度、色调 | 提高光照和色彩鲁棒性 |
| ShiftScaleRotate | 平移、缩放、旋转的组合变换 | 提高几何不变性 |
| GaussNoise | 添加高斯噪声 | 提高抗噪能力 |
| Blur | 高斯模糊或均值模糊 | 模拟失焦/运动模糊 |
| CoarseDropout | 随机遮挡矩形区域 | 提高对遮挡的鲁棒性 |
| Normalize | 按均值和标准差归一化 | 标准化输入分布 |
CutMix 与 MixUp
- MixUp:以随机比例混合两张图像及其标签,即
x' = λ * x_i + (1-λ) * x_j,标签同理。这鼓励模型学习线性插值行为,提高泛化能力。 - CutMix:从一张图像裁剪矩形区域粘贴到另一张图像上,标签按面积比例混合。相比 MixUp,CutMix 保留了局部图像内容,更适合视觉任务。
Albumentations vs torchvision
| 对比维度 | Albumentations | torchvision.transforms |
|---|---|---|
| 操作种类 | 70+,涵盖像素级、空间级、噪声类 | 约30+,基础操作齐全 |
| 执行速度 | 较快(基于 OpenCV,支持并行) | 较慢(基于 PIL,单线程) |
| 统一接口 | 统一调用 transform(image=img, mask=mask) | 不同操作需单独调用 |
| 与 PyTorch 集成 | 需额外包装,配合 Dataset 使用 | 原生集成,可直接 compose |
| 边框/关键点支持 | 完善,支持同步变换 | 不直接支持 |
在实际项目中,通常将 Albumentations 用于训练阶段的在线增强(数据量大时对速度要求高),而将 torchvision 用于验证/测试阶段的简单预处理(Resize + Normalize)。
代码示例:使用 PyTorch + timm 加载预训练模型进行分类推理
以下代码演示如何使用 timm(PyTorch Image Models)加载预训练模型,对单张图像进行图像分类推理。
import torch
import timm
from PIL import Image
from torchvision import transforms
# 选择模型并加载预训练权重
# 支持 'resnet50', 'efficientnet_b4', 'vit_base_patch16_224', 'convnext_base' 等
model_name = 'convnext_base'
model = timm.create_model(model_name, pretrained=True)
model.eval()
# 加载 ImageNet 类别标签映射
# 可从 https://storage.googleapis.com/download.tensorflow.org/data/ImageNetLabels.txt 下载
import json
import urllib.request
url = "https://raw.githubusercontent.com/anishathalye/imagenet-simple-labels/master/imagenet-simple-labels.json"
with urllib.request.urlopen(url) as f:
labels = json.load(f)
# 图像预处理
# 不同模型的预处理要求可能不同,timm 提供了 data_config 接口获取
data_config = timm.data.resolve_model_data_config(model)
transforms_list = timm.data.create_transform(**data_config, is_training=False)
def classify_image(image_path: str) -> tuple[str, float]:
"""
对单张图像进行分类推理。
Args:
image_path: 图像文件路径
Returns:
(类别名称, 置信度) 元组
"""
image = Image.open(image_path).convert("RGB")
input_tensor = transforms_list(image).unsqueeze(0) # 添加 batch 维度
with torch.no_grad():
output = model(input_tensor)
probabilities = torch.nn.functional.softmax(output[0], dim=0)
top_prob, top_idx = torch.topk(probabilities, k=1)
confidence = top_prob.item()
predicted_label = labels[top_idx.item()]
return predicted_label, confidence
# 使用示例
if __name__ == "__main__":
label, confidence = classify_image("example.jpg")
print(f"预测类别: {label}")
print(f"置信度: {confidence:.4f}")代码说明:
timm.create_model提供了数百种预训练模型的统一加载接口,支持自动下载权重。timm.data.resolve_model_data_config和timm.data.create_transform会根据模型自动获取其训练时使用的预处理参数(输入尺寸、均值、标准差、插值方式等),避免了手动配置的繁琐和错误。- 上述代码同时支持 CPU 和 GPU(若检测到 CUDA 可用,
model会自动使用 GPU;如需显式指定,可调用model.cuda()并将输入张量移到对应设备)。
如果要进行批量推理或集成到训练流程中,可进一步使用 timm.data.ImageDataset 和 PyTorch 的 DataLoader 实现高效数据加载与流水线处理。