卷积神经网络的引入6 —— 为什么 Transformer 也能“看懂”图像?CNN vs ViT 的直观对比实验
卷积神经网络的引入6 —— 为什么 Transformer 也能“看懂”图像?CNN vs ViT 的直观对比实验
在上一篇文章里,我们做了一个非常“狠”的实验:
当图像的空间结构被彻底打乱时,CNN 的优势会明显下降,而 MLP 反而没有那么吃亏。
这个结论其实非常重要。
因为它说明:
- CNN 的强大,并不是“凭空而来”;
- CNN 之所以能处理图像,是因为它天然相信局部结构、邻近像素关系和空间连续性;
- 一旦这种结构假设被摧毁,卷积的意义就会迅速减弱。
但这又会引出一个新的问题:
如果 Transformer 根本没有卷积核,它为什么也能“看懂”图像?
换句话说:没有局部卷积窗口,模型凭什么完成图像分类?
这正是 Vision Transformer(ViT)最值得理解的地方。
一、问题本质:CNN 在“看局部”,Transformer 在“看关系”
CNN 的核心机制是:
- 用卷积核在局部窗口内滑动;
- 逐层提取边缘、纹理、形状;
- 通过池化和层级堆叠形成更大感受野。
所以 CNN 天然带着一种很强的 空间归纳偏置(inductive bias):
- 邻近像素比远处像素更相关;
- 同样的局部模式在不同位置可以复用;
- 图像理解应该从局部到整体逐步建立。
而 Transformer 完全不是这一套思路。
它不预设“哪个像素必须先看、哪个区域必须后看”,而是把输入切成一块一块 token,再通过 Self-Attention 去学习:
- 哪些 patch 之间应该互相关注;
- 哪些区域组合起来更能决定类别;
- 全局依赖关系如何建立。
所以从直觉上说:
- CNN 像是在显微镜下一层层扫描图像;
- Transformer 像是在会议室里同时比较整张图上的多个区域之间的关系。
它不是先天“懂图像”,而是靠数据驱动地学会“哪些 patch 该彼此联系”。
二、ViT 到底做了什么?
Vision Transformer 的核心流程其实非常直接:
1. 把图像切成 patch
例如一张 32×32 的 CIFAR-10 图片,如果 patch size 取 4×4,那么每张图就会被切成:
8 × 8 = 64个 patch
每个 patch 本质上就是一小块局部区域。
2. 把每个 patch 映射成向量(embedding)
每个 patch 会被拉平成一段向量,再通过线性层映射到统一维度,比如 128 维。
于是原始图像就从:
- 一张二维/三维网格图像
变成了:
- 一串 patch token 序列
3. 加入位置编码(positional embedding)
因为 Transformer 本身并不知道“第几个 token 在左上角、哪个 token 在中间”,所以必须额外告诉它位置信息。
这一步非常关键。
如果没有位置编码,Transformer 只会看到一堆 patch 向量,却不知道它们在空间上如何排列。
4. 通过多层 Self-Attention 建立全局关系
每个 patch 都可以和其他 patch 计算相关性。
于是模型会逐渐学到:
- 猫耳朵区域应该和猫脸区域一起看;
- 轮胎区域应该和车身区域一起看;
- 飞机机翼和机头之间有强依赖关系。
这就是 ViT “看懂图像” 的关键。
它不是靠卷积去扫描局部,而是靠注意力去学习全局依赖。
三、这是不是说明 Transformer 比 CNN 更强?
不能这么简单下结论。
因为 CNN 和 ViT 的优势来源不同:
CNN 的优势
- 自带局部归纳偏置;
- 在小数据集上更容易收敛;
- 参数利用效率高;
- 对图像任务很“对口”。
ViT 的优势
- 更容易建模长距离依赖;
- 架构统一,方便与 NLP 中的 Transformer 思想融合;
- 在大规模数据上通常更容易做大做强。
所以真正的区别不是:
谁绝对更厉害?
而是:
谁更依赖先验结构,谁更依赖数据规模。
CNN 把“图像应该局部相关”这个先验写进了结构里;
ViT 则尽量少写先验,把更多自由度交给数据本身去学习。
四、实验目标:在同一数据集上直观看 CNN 和 ViT 的差异
这一篇我们不追求大模型,不追求 SOTA。
我们只做一个 直观、可复现、足够说明问题的 mini 对比实验:
实验问题
- 在同样的 CIFAR-10 数据集上,CNN 与 Tiny ViT 谁收敛更快?
- 在小样本、低分辨率场景下,ViT 能否直接打赢 CNN?
- 如果 patch 切得太大,ViT 会发生什么?
- Transformer 到底是“真的懂图像”,还是“靠足够多的数据硬学出来”?
控制变量
- 数据集:CIFAR-10
- 输入尺寸:
32×32 - 训练轮数:统一
- 优化器:统一 AdamW
- 数据增强:统一 RandomCrop + RandomHorizontalFlip
- 对比模型:SimpleCNN vs TinyViT
五、理论预期
在真正跑实验之前,我们先做一个理论预测:
| 模型 | 理论预期 |
|---|---|
| SimpleCNN | 收敛更快,小数据集下更稳 |
| TinyViT(小 patch) | 可以学会分类,但前期收敛通常慢于 CNN |
| TinyViT(大 patch) | token 太少,局部细节损失大,效果通常更差 |
原因并不复杂:
- CNN 先验更强,所以更容易学;
- ViT 更灵活,但也更“挑数据”;
- patch 越大,局部信息损失越明显;
- patch 越小,token 越多,注意力计算开销也越大。
六、实验代码:SimpleCNN vs TinyViT
下面给出一个可以直接复现的 PyTorch mini 实验框架。
import math
import time
import random
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
# ============================
# 1. 全局配置
# ============================
SEED = 42
BATCH_SIZE = 128
EPOCHS = 15
LR = 3e-4
WEIGHT_DECAY = 1e-4
PATCH_SIZE = 4 # 可改成 8 观察效果变化
EMBED_DIM = 128
NUM_HEADS = 4
DEPTH = 4
MLP_RATIO = 4
NUM_CLASSES = 10
def seed_everything(seed=SEED):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
seed_everything(SEED)
if torch.backends.mps.is_available():
device = torch.device("mps")
elif torch.cuda.is_available():
device = torch.device("cuda")
else:
device = torch.device("cpu")
print("Using device:", device)
# ============================
# 2. 数据集
# ============================
mean = (0.4914, 0.4822, 0.4465)
std = (0.2470, 0.2435, 0.2616)
transform_train = transforms.Compose([
transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean, std),
])
transform_test = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize(mean, std),
])
trainset = datasets.CIFAR10(root="./data", train=True, download=True, transform=transform_train)
testset = datasets.CIFAR10(root="./data", train=False, download=True, transform=transform_test)
trainloader = DataLoader(trainset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)
testloader = DataLoader(testset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)
# ============================
# 3. CNN
# ============================
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, 3, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, 3, padding=1),
nn.ReLU(inplace=True),
nn.AdaptiveAvgPool2d((1, 1))
)
self.classifier = nn.Linear(128, num_classes)
def forward(self, x):
x = self.features(x)
x = x.flatten(1)
return self.classifier(x)
# ============================
# 4. ViT 组件
# ============================
class PatchEmbedding(nn.Module):
def __init__(self, img_size=32, patch_size=4, in_chans=3, embed_dim=128):
super().__init__()
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = (img_size // patch_size) ** 2
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x):
x = self.proj(x) # [B, C, H/P, W/P]
x = x.flatten(2).transpose(1, 2) # [B, N, C]
return x
class MLPBlock(nn.Module):
def __init__(self, dim, hidden_dim):
super().__init__()
self.net = nn.Sequential(
nn.Linear(dim, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, dim)
)
def forward(self, x):
return self.net(x)
class TransformerEncoderBlock(nn.Module):
def __init__(self, dim, num_heads, mlp_ratio=4.0):
super().__init__()
self.norm1 = nn.LayerNorm(dim)
self.attn = nn.MultiheadAttention(embed_dim=dim, num_heads=num_heads, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
self.mlp = MLPBlock(dim, int(dim * mlp_ratio))
def forward(self, x):
x_norm = self.norm1(x)
attn_out, _ = self.attn(x_norm, x_norm, x_norm)
x = x + attn_out
x = x + self.mlp(self.norm2(x))
return x
class TinyViT(nn.Module):
def __init__(self, img_size=32, patch_size=4, embed_dim=128, depth=4, num_heads=4, num_classes=10):
super().__init__()
self.patch_embed = PatchEmbedding(img_size, patch_size, 3, embed_dim)
num_patches = self.patch_embed.num_patches
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
self.blocks = nn.Sequential(*[
TransformerEncoderBlock(embed_dim, num_heads, mlp_ratio=MLP_RATIO)
for _ in range(depth)
])
self.norm = nn.LayerNorm(embed_dim)
self.head = nn.Linear(embed_dim, num_classes)
def forward(self, x):
x = self.patch_embed(x) # [B, N, D]
B = x.size(0)
cls_token = self.cls_token.expand(B, -1, -1)
x = torch.cat([cls_token, x], dim=1) # [B, N+1, D]
x = x + self.pos_embed
x = self.blocks(x)
x = self.norm(x)
cls = x[:, 0]
return self.head(cls)
# ============================
# 5. 训练与评估
# ============================
@torch.no_grad()
def evaluate(model, loader):
model.eval()
correct = 0
total = 0
for x, y in loader:
x, y = x.to(device), y.to(device)
logits = model(x)
pred = logits.argmax(dim=1)
correct += (pred == y).sum().item()
total += y.size(0)
return correct / total
def train_model(model, trainloader, testloader, name="model"):
model = model.to(device)
optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)
criterion = nn.CrossEntropyLoss()
history = {"train_acc": [], "test_acc": []}
for epoch in range(EPOCHS):
model.train()
correct = 0
total = 0
running_loss = 0.0
t0 = time.perf_counter()
for x, y in trainloader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
logits = model(x)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
running_loss += loss.item() * x.size(0)
pred = logits.argmax(dim=1)
correct += (pred == y).sum().item()
total += y.size(0)
train_acc = correct / total
test_acc = evaluate(model, testloader)
history["train_acc"].append(train_acc)
history["test_acc"].append(test_acc)
dt = time.perf_counter() - t0
print(f"[{name}] Epoch {epoch+1:02d}/{EPOCHS} | train={train_acc:.4f} | test={test_acc:.4f} | time={dt:.1f}s")
return history
cnn = SimpleCNN()
vit = TinyViT(patch_size=PATCH_SIZE, embed_dim=EMBED_DIM, depth=DEPTH, num_heads=NUM_HEADS)
cnn_hist = train_model(cnn, trainloader, testloader, name="CNN")
vit_hist = train_model(vit, trainloader, testloader, name="TinyViT")
七、如何理解这个实验?
这个实验最想说明的,不是“谁能多拿几个点”,而是:
CNN 与 Transformer 是通过完全不同的机制理解图像的。
对 CNN 来说
它的强项是:
- 先天知道“局部”很重要;
- 先天知道“同一个模式在不同位置可复用”;
- 所以在小数据集、低分辨率任务上往往学得更快。
对 ViT 来说
它的强项是:
- 不把局部结构写死;
- 直接学习 patch 与 patch 的全局依赖;
- 在更大数据、更大模型、更大训练预算下,往往潜力更强。
也就是说:
- CNN 更像“有经验的图像工程师”;
- ViT 更像“没有先验、但很擅长从数据中总结关系的分析师”。
八、你通常会观察到什么现象?
如果你真的跑这组 mini 实验,通常会观察到以下现象:
1. CNN 在前几个 epoch 往往收敛更快
这很正常。
因为 CIFAR-10 不算大,图像分辨率也只有 32×32,局部纹理和边缘信息本来就很重要。
CNN 的归纳偏置在这种任务里天然占优。
2. TinyViT 不是学不会,而是通常更“慢热”
ViT 并不是不能做图像分类。
问题在于:
- 它没有卷积这种强结构先验;
- 它要先学“哪些 patch 有关”,再学“这些关系如何对应类别”。
所以它通常更依赖:
- 合理的 patch 设计;
- 足够训练轮数;
- 更强的数据增强;
- 更大的训练数据规模。
3. patch 太大时,ViT 会明显吃亏
比如把 patch_size=4 改成 8,你会发现 token 数量会迅速减少:
4×4 patch:共64个 token8×8 patch:共16个 token
这意味着什么?
意味着很多局部细节在 patch embedding 那一步就被粗暴压缩掉了。
对于 CIFAR-10 这种本来分辨率就不高的数据,patch 过大几乎等于“信息先丢一轮,再开始学习”。
九、这件事和上一篇有什么联系?
联系非常强。
上一篇我们讨论的是:
当空间结构被打乱时,CNN 为什么会受伤?
这一篇我们讨论的是:
如果不靠卷积,Transformer 为什么仍能做图像理解?
两篇文章合起来,其实在回答同一件事:
图像模型到底依赖什么结构先验?
CNN 的答案是:
- 我相信局部;
- 我相信邻域;
- 我相信平移复用;
- 所以我在图像任务上天然高效。
Transformer 的答案是:
- 我不强行写死局部规则;
- 我把图像切成 token;
- 我通过注意力自己学 patch 之间的关系;
- 只要数据足够,我同样可以建立图像理解能力。
所以,真正重要的不是“谁彻底替代谁”,而是:
- CNN 代表的是 强归纳偏置;
- ViT 代表的是 弱归纳偏置 + 更强数据驱动能力。
十、结论
这一篇可以先得出 4 个核心结论:
1. Transformer 并不是天然“懂图像”
它之所以能处理图像,是因为图像被切成了 patch token,并结合位置编码后,Transformer 可以学习 token 之间的全局关系。
2. CNN 和 ViT 的核心差异,不在于有没有卷积,而在于“先验写入程度”
- CNN:把局部结构先验直接写进架构;
- ViT:尽量少写先验,把更多自由度交给数据。
3. 在小数据、低分辨率任务上,CNN 往往更高效
因为它更符合这类任务的结构特征。
4. ViT 的真正优势,往往出现在更大规模数据和训练资源下
当数据足够大时,弱先验不再是劣势,反而会变成更强的表达自由度。

浙公网安备 33010602011771号