CNN与Transformer模型结构图解析:如何选择与优化视觉任务架构

1次阅读
没有评论

共计 3906 个字符,预计需要花费 10 分钟才能阅读完成。

image.webp

背景痛点分析

在计算机视觉任务中,CNN(卷积神经网络)凭借局部感受野和权重共享的特性,一直是主流架构。但随着任务复杂度的提升,其局限性逐渐显现:

CNN 与 Transformer 模型结构图解析:如何选择与优化视觉任务架构

  • 长距离依赖建模不足 :传统 CNN 的卷积核大小有限(通常 3×3 或 5×5),难以捕获图像中远距离像素间的关联。例如在医学图像分割中,病灶区域可能分散在多个不连续区域。
  • 全局信息整合效率低 :需通过堆叠多层卷积或池化操作逐步扩大感受野,导致深层网络训练困难。

Transformer 架构虽然在自然语言处理中表现优异,但直接应用于视觉任务时面临:

  • 计算复杂度高 :标准自注意力机制的计算成本与输入序列长度呈平方关系(O(n²))。对于 224×224 图像,ViT 需处理 196 个 16×16 图像块,显存占用激增。
  • 数据依赖性 :纯 Transformer 模型通常需要大规模预训练(如 JFT-300M 数据集)才能达到与 CNN 相当的性能。

技术对比:ResNet50 与 ViT 结构图解析

ResNet50 核心结构(CNN 代表)

# PyTorch 风格的模块定义
class Bottleneck(nn.Module):
    def __init__(self, in_channels, out_channels, stride=1):
        super().__init__()
        self.conv1 = nn.Conv2d(in_channels, out_channels//4, kernel_size=1, bias=False)
        self.bn1 = nn.BatchNorm2d(out_channels//4)
        self.conv2 = nn.Conv2d(out_channels//4, out_channels//4, kernel_size=3, 
                              stride=stride, padding=1, bias=False)  # 关键计算节点
        self.bn2 = nn.BatchNorm2d(out_channels//4)
        self.conv3 = nn.Conv2d(out_channels//4, out_channels, kernel_size=1, bias=False)
        self.bn3 = nn.BatchNorm2d(out_channels)

计算复杂度分析

  • 主要开销来自 3×3 卷积层(占整体 FLOPs 的~70%)
  • 通过残差连接缓解梯度消失,但感受野扩展依赖网络深度

Vision Transformer(ViT)结构

class PatchEmbed(nn.Module):
    def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
        super().__init__()
        num_patches = (img_size // patch_size) ** 2  # 196 for 224×224
        self.proj = nn.Conv2d(in_chans, embed_dim, 
                             kernel_size=patch_size, stride=patch_size)  # 分块嵌入
        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))

关键差异点

  1. 输入处理:ViT 通过线性投影将图像转为序列,CNN 保持空间结构
  2. 特征交互:ViT 使用自注意力实现全局交互,CNN 依赖局部卷积
  3. 位置信息:ViT 需显式添加位置编码,CNN 通过卷积隐含位置信息

混合架构实现示例

CNN-Transformer 混合模型(PyTorch 实现)

class HybridModel(nn.Module):
    def __init__(self):
        super().__init__()
        # CNN 特征提取(带通道注意力)self.cnn_backbone = nn.Sequential(nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3),
            ChannelAttention(64),  # 通道注意力模块
            nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
            ResNetBlock(64, 256, stride=1)
        )

        # Transformer 编码层
        self.transformer = nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model=256, nhead=8),
            num_layers=6
        )

    def forward(self, x):
        # CNN 提取特征 [B,3,224,224] -> [B,256,14,14]
        cnn_feat = self.cnn_backbone(x)

        # 展平为序列 [B,256,14,14] -> [B,196,256]
        patches = cnn_feat.flatten(2).transpose(1, 2)

        # Transformer 处理
        output = self.transformer(patches)
        return output

关键组件说明

  1. 通道注意力模块

    class ChannelAttention(nn.Module):
        def __init__(self, channels, reduction=16):
            super().__init__()
            self.avg_pool = nn.AdaptiveAvgPool2d(1)
            self.fc = nn.Sequential(nn.Linear(channels, channels // reduction),
                nn.ReLU(),
                nn.Linear(channels // reduction, channels)
            )
        def forward(self, x):
            b, c, _, _ = x.size()
            y = self.avg_pool(x).view(b, c)
            y = self.fc(y).view(b, c, 1, 1)
            return x * y.sigmoid()  # 特征图通道权重 

  2. 位置编码可视化

    # 正弦位置编码公式
    pos_enc = torch.zeros(1, num_patches, d_model)
    position = torch.arange(0, num_patches).unsqueeze(1)
    div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
    pos_enc[0, :, 0::2] = torch.sin(position * div_term)
    pos_enc[0, :, 1::2] = torch.cos(position * div_term)

性能考量与优化建议

CIFAR-10 测试结果

模型类型 FLOPs (G) mAP (%) 显存占用 (MB)
ResNet50 4.1 95.2 1200
ViT-Tiny 1.2 93.8 850
Hybrid (Ours) 2.7 96.1 1100

显存优化技巧

  • 使用混合精度训练(AMP)

    scaler = torch.cuda.amp.GradScaler()
    with torch.cuda.amp.autocast():
        outputs = model(inputs)
        loss = criterion(outputs, labels)
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

  • 梯度检查点技术

    from torch.utils.checkpoint import checkpoint
    
    def custom_forward(module, x):
        return module(x)
    
    # 在 forward 中替换原始调用
    x = checkpoint(custom_forward, self.transformer_layer, x)

避坑指南

小数据场景过拟合解决方案

  1. 数据增强策略:

    train_transform = transforms.Compose([transforms.RandomHorizontalFlip(),
        transforms.RandomRotation(15),
        transforms.ColorJitter(brightness=0.2, contrast=0.2),
        transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),
        transforms.ToTensor()])

  2. 正则化配置:

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, 
                                weight_decay=0.05)  # 较大的 weight_decay

边缘设备部署量化建议

  1. 注意力层量化难点:
  2. Query/Key 的点积操作范围动态变化
  3. Softmax 输出需要高精度表示

  4. 可行方案:

    # 使用 PyTorch 量化 API
    model.qconfig = torch.quantization.get_default_qat_qconfig('fbgemm')
    quantized_model = torch.quantization.prepare_qat(model.train())
    quantized_model = torch.quantization.convert(quantized_model.eval())

开放性问题讨论

在边缘设备部署时,如何平衡注意力层的计算精度与推理速度?可能的探索方向包括:

  • 采用稀疏注意力模式(如轴向注意力)
  • 对注意力权重进行 8 -bit 量化
  • 使用低秩近似重构注意力矩阵
正文完
 0
评论(没有评论)