从零理解CNN与Transformer:核心原理对比与实战入门指南

1次阅读
没有评论

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

image.webp

背景痛点:为什么需要 Transformer

传统卷积神经网络(CNN)在计算机视觉(CV)领域取得了巨大成功,但其核心设计存在一个根本限制:局部感受野。CNN 通过滑动窗口的方式逐步提取特征,这意味着:

  • 浅层卷积只能捕捉像素级的局部模式(如边缘、纹理)
  • 深层网络通过堆叠卷积层才能建立远距离特征关系
  • 这种间接的远程建模方式效率低下,且难以处理长距离依赖

而 Transformer 最初在自然语言处理(NLP)中展现出强大的 全局上下文建模能力,其核心是通过 Self-Attention 机制直接计算任意两个位置的关系。2020 年《An Image is Worth 16×16 Words》论文首次将纯 Transformer 架构(Vision Transformer/ViT)成功应用于图像分类,引发了 CV 领域的架构革命。

核心原理对比

CNN 的局部特征提取

典型 CNN 层的计算可表示为:

$$\text{Output}(x,y) = \sum_{i=-k}^{k}\sum_{j=-k}^{k} \text{Input}(x+i,y+j) \cdot \text{Kernel}(i,j)$$

其中 kernel size $k$ 通常为 3 或 5,这种设计带来两个特性:

  • 平移不变性:相同模式在不同位置会被同样检测
  • 局部性:每个输出仅依赖附近 $k\times k$ 区域的输入

从零理解 CNN 与 Transformer:核心原理对比与实战入门指南
(示意图:CNN 逐层扩大感受野)

Transformer 的全局注意力

Transformer 的核心是 Multi-Head Self-Attention(多头自注意力),其关键计算步骤:

  1. 将输入投影到 Query(Q)、Key(K)、Value(V)空间
  2. 计算注意力权重:
    $$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$
  3. 多头结果拼接后通过线性层融合

(示意图:Transformer 可直接建模任意两像素关系)

实战代码对比

CNN 图像分类器实现

import torch
import torch.nn as nn

class SimpleCNN(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        # 卷积层 1:3 输入通道,16 输出通道,3x3 卷积核
        self.conv1 = nn.Conv2d(3, 16, kernel_size=3, stride=1, padding=1)
        # 最大池化,2x2 窗口,步长 2
        self.pool = nn.MaxPool2d(2, 2)
        # 卷积层 2:16->32 通道
        self.conv2 = nn.Conv2d(16, 32, 3, 1, 1)
        # 全连接层
        self.fc = nn.Linear(32 * 8 * 8, num_classes)  # 假设输入为 32x32 图像

    def forward(self, x):
        x = self.pool(torch.relu(self.conv1(x)))  # 输出 16x16x16
        x = self.pool(torch.relu(self.conv2(x)))  # 输出 8x8x32
        x = x.view(-1, 32 * 8 * 8)  # 展平
        return self.fc(x)

ViT 关键模块实现

class PatchEmbedding(nn.Module):
    """将图像分割为 patches 并嵌入"""
    def __init__(self, img_size=32, patch_size=4, in_chans=3, embed_dim=64):
        super().__init__()
        self.num_patches = (img_size // patch_size) ** 2
        # 使用卷积实现分块和投影
        self.proj = nn.Conv2d(in_chans, embed_dim, 
                             kernel_size=patch_size, 
                             stride=patch_size)
        # 可学习的分类 token
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))

    def forward(self, x):
        # x 形状: [B, C, H, W]
        x = self.proj(x)  # [B, E, H/P, W/P]
        x = x.flatten(2).transpose(1, 2)  # [B, num_patches, E]
        # 添加 class token
        cls_tokens = self.cls_token.expand(x.shape[0], -1, -1)
        x = torch.cat((cls_tokens, x), dim=1)
        return x

三大常见陷阱与解决方案

  1. 显存溢出问题
  2. 错误:直接对高分辨率图像(如 224×224)计算全局注意力
  3. 解决:必须先将图像分块(如 16×16 patches),ViT 中典型设置:

    patch_size = 16  # 224x224 图像→14x14 个 patch

  4. 位置信息丢失

  5. 错误:未添加位置编码(Positional Encoding),导致模型无法理解空间结构
  6. 解决:标准 ViT 使用可学习的 1D 位置编码:

    self.pos_embed = nn.Parameter(torch.randn(1, num_patches+1, embed_dim))

  7. 训练不稳定性

  8. 错误:直接使用 CNN 常用的学习率(如 0.1)导致梯度爆炸
  9. 解决:Transformer 通常需要更小的学习率和预热(warmup):
    optimizer = AdamW(model.parameters(), lr=3e-5)
    scheduler = get_cosine_schedule_with_warmup(optimizer, warmup_steps=500)

性能对比实验(CIFAR-10)

指标 ResNet-18 ViT-Tiny
参数量 11M 5M
训练时间 /epoch 45s 68s
显存占用 1.2GB 2.1GB
测试准确率 93.5% 91.2%

注:ViT 在小数据集上需要更强的数据增强和正则化

融合创新方向

现代架构正尝试结合两者的优势:

  • Swin Transformer:引入局部窗口注意力 + 跨窗口连接,兼具 CNN 的局部性和 Transformer 的全局性
  • ConvNeXt:用 CNN 结构实现类似 Transformer 的特性
  • 混合架构:前端使用 CNN 提取低级特征,后端用 Transformer 建模全局关系

建议初学者先从 PyTorch 官方 Vision Transformer 教程入手,逐步理解模块设计背后的动机。记住:没有绝对优越的架构,只有适合特定任务的解决方案。

正文完
 0
评论(没有评论)