第 11 章 2020—2021:视觉新范式——ViT、Swin 与扩散模型
学习目标
- 理解视觉 Transformer(ViT,2020)如何把图片变成 token 序列
- 理解 Swin Transformer(2021)的移动窗口注意力
- 理解扩散模型(DDPM,2020)的前向加噪与反向去噪
- 分别用 PyTorch 实现三种模型
11.1 时代背景:Transformer 反攻视觉
第 9 章之后,Transformer 统治了 NLP。2020 年,Google 的 ViT(Vision Transformer) 证明:把图片切成 patch,当成「视觉词」,直接喂给标准 Transformer 编码器,在大规模数据上可以超越 CNN。2021 年,微软的 Swin Transformer 引入移动窗口注意力,把 Transformer 的效率带到高分辨率图像。同年,DDPM 扩散模型在图像生成上超越 GAN。
视觉世界完成了从 CNN 到「注意力 + 扩散」的范式转移。
11.2 ViT:图片即序列
ViT 的思路非常直接:
- 把 224×224 图片切成 16×16 的 patch(14×14 = 196 个);
- 每个 patch 线性投影成一个向量(相当于一个「词」);
- 前面加一个
[CLS]记号(分类标记),加上位置编码; - 送入标准 Transformer 编码器,取
[CLS]的输出接分类头。
examples/ch11_vit.py 在 Fashion-MNIST 上实现 ViT(4×4 patch,49 个视觉 token):
class ViT(nn.Module):
def __init__(self):
super().__init__()
self.patchify = nn.Conv2d(1, D_MODEL, kernel_size=PATCH, stride=PATCH)
self.cls = nn.Parameter(torch.zeros(1, 1, D_MODEL))
self.pos = nn.Parameter(torch.randn(1, N_PATCHES + 1, D_MODEL) * 0.02)
encoder_layer = nn.TransformerEncoderLayer(
D_MODEL, NHEAD, dim_feedforward=D_MODEL * 4,
batch_first=True, activation="gelu", norm_first=True,
)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=2)
self.head = nn.Linear(D_MODEL, 10)
def forward(self, x):
patches = self.patchify(x).flatten(2).transpose(1, 2) # (B, 49, D)
tokens = torch.cat([self.cls.expand(x.shape[0], -1, -1), patches], dim=1)
tokens = tokens + self.pos
out = self.encoder(tokens)
return self.head(out[:, 0])运行结果:
ViT 参数量:104,970
epoch 1:loss = 0.9528,训练准确率 = 78.98%
epoch 2:loss = 0.5372,训练准确率 = 82.97%
epoch 3:loss = 0.4613,训练准确率 = 84.31%
测试准确率 = 82.94%10 万参数、3 个 epoch 达到 82.9%。ViT 的代价是数据饥饿:它没有 CNN 的归纳偏置(局部性、平移不变性),需要大量数据才能学出来。在 ImageNet-21k 上预训练后,ViT 才全面超过 CNN——这也是「预训练 + 微调」在视觉领域的开端。
11.3 Swin:窗口里的高效注意力
自注意力的计算量是序列长度的平方。图片的 patch 数随分辨率平方增长,直接全局注意力在高分辨率下会爆炸。Swin Transformer 的解法:只在局部窗口内做注意力。
- 窗口划分:把特征图切成互不重叠的小窗口(如 4×4 个 patch 一个窗口);
- 窗口内自注意力:每个窗口内部做多头注意力;
- 移动窗口:下一层把窗口平移半个窗口再划分,让信息跨窗口流动——新窗口会包含原图中不同窗口的内容,注意力就能建立跨窗口联系。
examples/ch11_swin.py 在 8×8 特征图上演示完整机制:
第一次划分:窗口数 = 4 ,每个窗口 token 数 = 16
窗口 0 的像素: [0.0, 1.0, 2.0, 3.0, 8.0, 9.0, 10.0, 11.0, 16.0, 17.0, 18.0, 19.0, 24.0, 25.0, 26.0, 27.0]
窗口注意力输出形状: (4, 16, 1) → 逆变换回 (1, 1, 8, 8)
移动窗口后,窗口 0 的像素: [18.0, 19.0, 20.0, 21.0, 26.0, 27.0, 28.0, 29.0, 34.0, 35.0, 36.0, 37.0, 42.0, 43.0, 44.0, 45.0]
→ 现在它同时包含原图中 4 个不同窗口的像素,信息可以跨窗口流动第一次划分后,窗口 0 只含原图左上角;移动窗口后,窗口 0 的像素来自原图 4 个不同区域——跨窗口信息由此建立。窗口注意力的复杂度与图像尺寸线性相关,这让 Swin 能处理高分辨率图像,成为 2021 年后检测、分割任务的主流骨干。
11.4 扩散模型:从噪声中逐步雕琢图像
2020 年,Ho 等人提出 DDPM(Denoising Diffusion Probabilistic Models),思路非常直觉:
前向过程(加噪):把一张干净图片逐步加高斯噪声,T 步之后变成纯噪声。
xₜ = √ᾱₜ·x₀ + √(1-ᾱₜ)·ε反向过程(去噪):训练一个 U-Net,给它「带噪图片 + 时间步 t」,预测噪声 ε。训练损失就是预测噪声与真实噪声的 MSE。
采样:从纯噪声出发,用训练好的网络逐步去噪 T 步,得到新图片。
examples/ch11_diffusion.py 在 16×16 的 Fashion-MNIST 上训练最小 DDPM:
U-Net 参数量:184,065
epoch 1:loss = 0.2475
epoch 2:loss = 0.1336
采样 8 张图,像素统计:
每张均值: 0.12 0.05 0.19 0.05 0.20 -0.02 0.24 0.14
每张最大值: 0.82 1.34 0.58 0.49 0.76 0.61 0.76 0.732 个 epoch 后,从纯噪声采样出的图片已经有亮部(最大值 0.5~1.3),均值接近真实分布——扩散模型开始「学会」衣服的像素结构。扩散模型 2021-2022 年在生成质量上全面超越 GAN,成为 Stable Diffusion、Midjourney、DALL·E 的底层引擎(第 12 章)。
11.5 三条新范式的共同点
ViT、Swin、扩散模型看似不同,共享同一个底座:注意力与自监督的规模化。ViT 把视觉变成 token,扩散模型把生成变成逐步去噪,都建立在第 9 章 Transformer 和自监督训练之上——这也是「预训练大模型」时代的开端。
完整代码
本章用到的完整示例代码:
examples/ch11_vit.py
"""第 11 章示例(2020):视觉 Transformer(ViT)在 Fashion-MNIST 上训练。
28×28 图片切成 4×4 的 patch(7×7=49 个),每个 patch 线性投影成向量,
交给 2 层 Transformer 编码器,再接分类头。
运行方式:
uv run python examples/ch11_vit.py
"""
import warnings
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
warnings.filterwarnings("ignore", message=".*nested_tensor.*")
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.2860,), (0.3530,)),
])
train = datasets.FashionMNIST("data", train=True, download=True, transform=transform)
test = datasets.FashionMNIST("data", train=False, download=True, transform=transform)
train_loader = DataLoader(train, batch_size=256, shuffle=True)
test_loader = DataLoader(test, batch_size=1024)
PATCH, D_MODEL, NHEAD = 4, 64, 4
N_PATCHES = (28 // PATCH) ** 2
class ViT(nn.Module):
def __init__(self):
super().__init__()
self.patchify = nn.Conv2d(1, D_MODEL, kernel_size=PATCH, stride=PATCH)
self.cls = nn.Parameter(torch.zeros(1, 1, D_MODEL))
self.pos = nn.Parameter(torch.randn(1, N_PATCHES + 1, D_MODEL) * 0.02)
encoder_layer = nn.TransformerEncoderLayer(
D_MODEL, NHEAD, dim_feedforward=D_MODEL * 4,
batch_first=True, activation="gelu", norm_first=True,
)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=2)
self.head = nn.Linear(D_MODEL, 10)
def forward(self, x):
patches = self.patchify(x).flatten(2).transpose(1, 2) # (B, 49, D)
tokens = torch.cat([self.cls.expand(x.shape[0], -1, -1), patches], dim=1)
tokens = tokens + self.pos
out = self.encoder(tokens)
return self.head(out[:, 0])
model = ViT().to(device)
opt = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-2)
criterion = nn.CrossEntropyLoss()
print(f"ViT 参数量:{sum(p.numel() for p in model.parameters()):,}")
def evaluate(loader):
model.eval()
correct = total = 0
with torch.no_grad():
for xb, yb in loader:
pred = model(xb.to(device)).argmax(1)
correct += (pred == yb.to(device)).sum().item()
total += len(yb)
return correct / total
for epoch in range(3):
model.train()
total_loss = 0
for xb, yb in train_loader:
xb, yb = xb.to(device), yb.to(device)
opt.zero_grad()
loss = criterion(model(xb), yb)
loss.backward()
opt.step()
total_loss += loss.item() * len(xb)
print(f"epoch {epoch + 1}:loss = {total_loss / len(train):.4f},"
f"训练准确率 = {evaluate(train_loader) * 100:.2f}%")
print(f"测试准确率 = {evaluate(test_loader) * 100:.2f}%")examples/ch11_swin.py
"""第 11 章示例(2021):Swin Transformer 的移动窗口注意力。
窗口注意力只在局部窗口内计算,复杂度与图像尺寸线性相关;
通过移动窗口(shifted window),信息可以跨窗口传播。
本示例演示"划分窗口→窗口内注意力→移动窗口→再次划分"的完整机制。
运行方式:
uv run python examples/ch11_swin.py
"""
import torch
from torch import nn
# 8×8 特征图,像素值 0~63
x = torch.arange(64, dtype=torch.float32).view(1, 1, 8, 8)
WS = 4
def window_partition(x, ws):
"""把 (B, C, H, W) 切成 (n_windows, ws*ws, C)。"""
B, C, H, W = x.shape
x = x.view(B, C, H // ws, ws, W // ws, ws)
return x.permute(0, 2, 4, 3, 5, 1).reshape(-1, ws * ws, C)
def window_reverse(wins, ws, H, W):
"""划分的逆操作:拼回特征图。"""
C = wins.shape[-1]
n = H // ws
return wins.view(1, n, n, ws, ws, C).permute(0, 5, 1, 3, 2, 4).reshape(1, C, H, W)
wins = window_partition(x, WS)
print("第一次划分:窗口数 =", wins.shape[0], ",每个窗口 token 数 =", wins.shape[1])
print("窗口 0 的像素:", wins[0].squeeze().tolist())
# 窗口内多头自注意力
mha = nn.MultiheadAttention(1, 1, batch_first=True)
out, _ = mha(wins, wins, wins)
print("窗口注意力输出形状:", tuple(out.shape), "→ 逆变换回", tuple(window_reverse(out, WS, 8, 8).shape))
# 移动窗口:特征图沿 (H, W) 各平移 -2
shifted = torch.roll(x, shifts=(-2, -2), dims=(2, 3))
wins_shifted = window_partition(shifted, WS)
print("移动窗口后,窗口 0 的像素:", wins_shifted[0].squeeze().tolist())
print("→ 现在它同时包含原图中 4 个不同窗口的像素,信息可以跨窗口流动")examples/ch11_diffusion.py
"""第 11 章示例:最小扩散模型(DDPM)在 16×16 Fashion-MNIST 上训练。
前向过程逐步加噪,反向过程用一个小 U-Net 预测噪声;采样时从纯噪声逐步去噪。
运行方式:
uv run python examples/ch11_diffusion.py
"""
import torch
from torch import nn
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
transform = transforms.Compose([
transforms.Resize(16),
transforms.ToTensor(),
])
train = datasets.FashionMNIST("data", train=True, download=True, transform=transform)
loader = DataLoader(Subset(train, range(8000)), batch_size=128, shuffle=True)
T = 100
beta = torch.linspace(1e-4, 0.02, T, device=device)
alpha = 1 - beta
alpha_bar = torch.cumprod(alpha, dim=0)
class TBlock(nn.Module):
def __init__(self, in_ch, out_ch, t_dim=32, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(in_ch, out_ch, 3, stride, padding=1)
self.gn1 = nn.GroupNorm(4, out_ch)
self.t_lin = nn.Linear(t_dim, out_ch)
self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)
self.gn2 = nn.GroupNorm(4, out_ch)
self.silu = nn.SiLU()
def forward(self, x, t):
h = self.silu(self.gn1(self.conv1(x)))
h = h + self.t_lin(t)[:, :, None, None]
return self.silu(self.gn2(self.conv2(h)))
class MiniUNet(nn.Module):
def __init__(self, in_ch=1, base=32, t_dim=32):
super().__init__()
self.t_emb = nn.Sequential(nn.Linear(1, t_dim), nn.SiLU(), nn.Linear(t_dim, t_dim))
self.down1 = TBlock(in_ch, base, t_dim) # 16×16
self.down2 = TBlock(base, base * 2, t_dim, stride=2) # 8×8
self.mid = TBlock(base * 2, base * 2, t_dim)
self.up1 = nn.Sequential(
nn.Upsample(scale_factor=2, mode="nearest"),
nn.Conv2d(base * 2, base, 3, padding=1),
)
self.up1_t = TBlock(base, base, t_dim)
self.out = nn.Conv2d(base, in_ch, 1)
def forward(self, x, t):
te = self.t_emb(t)
h1 = self.down1(x, te)
h2 = self.down2(h1, te)
h = self.mid(h2, te)
h = self.up1(h)
h = self.up1_t(h + h1, te)
return self.out(h)
model = MiniUNet().to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
print(f"U-Net 参数量:{sum(p.numel() for p in model.parameters()):,}")
for epoch in range(2):
total_loss = 0
for xb, _ in loader:
xb = xb.to(device)
t = torch.randint(1, T + 1, (xb.shape[0],), device=device)
noise = torch.randn_like(xb)
ab = alpha_bar[t - 1].view(-1, 1, 1, 1)
x_noisy = ab.sqrt() * xb + (1 - ab).sqrt() * noise
t_in = (t / T).float().view(-1, 1)
opt.zero_grad()
loss = nn.functional.mse_loss(model(x_noisy, t_in), noise)
loss.backward()
opt.step()
total_loss += loss.item() * len(xb)
print(f"epoch {epoch + 1}:loss = {total_loss / 8000:.4f}")
# 采样:从纯噪声出发,逐步去噪
model.eval()
x = torch.randn(8, 1, 16, 16, device=device)
with torch.no_grad():
for t in range(T, 0, -1):
t_in = torch.full((8, 1), t / T, device=device)
pred_noise = model(x, t_in)
ab_t = alpha_bar[t - 1]
x = (x - (1 - alpha[t - 1]) / (1 - ab_t).sqrt() * pred_noise) / alpha[t - 1].sqrt()
if t > 1:
x = x + beta[t - 1].sqrt() * torch.randn_like(x)
print("采样 8 张图,像素统计:")
print(" 每张均值:", " ".join(f"{v:.2f}" for v in x.mean(dim=(1, 2, 3)).tolist()))
print(" 每张最大值:", " ".join(f"{v:.2f}" for v in x.max(dim=1)[0].max(dim=1)[0].max(dim=1)[0].tolist()))动手实践
- 运行
examples/ch11_vit.py,把 patch 从 4×4 改成 7×7(4 个 patch),对比参数量与准确率。 - 运行
examples/ch11_swin.py,把窗口大小从 4 改成 2,观察窗口数与移动窗口后的像素组成。 - 修改
examples/ch11_diffusion.py的 T(时间步)从 100 改成 50,观察训练 loss 与采样效果。 - 给 ViT 去掉
[CLS]token,改用平均池化做分类,对比准确率。
常见错误
| 错误写法 | 现象 | 原因 |
|---|---|---|
| ViT 的 patch 数算错 | 位置编码维度不匹配 | patch 数 = (H/P)×(W/P),加 CLS 后 +1 |
[CLS] 复制时忘了 batch 维 | 形状错误 | 用 cls.expand(B, -1, -1) |
| 窗口划分的 reshape 顺序错 | 窗口内容错乱 | 要先 view 再 permute,注意维度顺序 |
| 移动窗口后忘了「平移回去」 | 特征错位 | 真实 Swin 采样前要把 shift 还原并加 mask |
| 扩散模型采样时不用 α 公式 | 生成全噪声 | 去噪步必须按 x = (x - β/√(1-ᾱ)·ε)/√α 更新 |
| DDPM 的 beta 调度与时间步不匹配 | loss 发散 | t 与 ᾱ 的索引要一一对应 |
章末练习
基础
- ViT 如何把图片变成 token 序列?
- Swin 为什么用窗口注意力?
- 扩散模型的前向与反向过程分别是什么?
提高
- 计算:8×8 特征图、4×4 窗口时,窗口内注意力的计算量 vs 全局注意力(序列长度 64)。
- 在 ViT 中把 2 层编码器改成 4 层,观察准确率与训练时间。
- 修改 DDPM 示例:用余弦 beta 调度替换线性调度,对比 loss 曲线。
挑战
- 实现带移动窗口的完整 Swin 块(含 cyclic shift 与 mask),在 32×32 输入上验证窗口间信息传播。
- 实现 DDPM 的
DDIM采样(确定性、少步数),对比 100 步与 20 步的采样质量。
章末自测
- ViT 中图片被切成?
- A. 像素
- B. patch
- C. 通道
- D. 窗口
- ViT 分类用哪个 token 的输出?
- A. 第一个 patch
- B. [CLS]
- C. 平均池化
- D. 最后一个 patch
- Swin 的移动窗口解决了?
- A. 数据不足
- B. 跨窗口信息流动
- C. 过拟合
- D. 梯度消失
- 窗口注意力的复杂度与图像尺寸?
- A. 平方关系
- B. 线性关系
- C. 无关
- D. 指数关系
- DDPM 的反向过程学习预测?
- A. 图片
- B. 噪声 ε
- C. 类别
- D. 时间步
- 扩散模型采样从什么开始?
- A. 一张真图
- B. 纯噪声
- C. 随机类别
- D. 一个 patch
- ViT 的缺点之一是?
- A. 无法训练
- B. 数据饥饿(需要大规模预训练)
- C. 参数太少
- D. 只能分类
- DDPM 的 U-Net 输入包括?
- A. 图片和时间步
- B. 图片和标签
- C. 噪声和时间步
- D. 只有图片
- Stable Diffusion 等产品使用的生成技术是?
- A. VAE
- B. 扩散模型
- C. GAN
- D. RNN
- Swin Transformer 提出于?
- A. 2019
- B. 2021
- C. 2020
- D. 2023
