Skip to content

第 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 的思路非常直接:

  1. 把 224×224 图片切成 16×16 的 patch(14×14 = 196 个);
  2. 每个 patch 线性投影成一个向量(相当于一个「词」);
  3. 前面加一个 [CLS] 记号(分类标记),加上位置编码;
  4. 送入标准 Transformer 编码器,取 [CLS] 的输出接分类头。

examples/ch11_vit.py 在 Fashion-MNIST 上实现 ViT(4×4 patch,49 个视觉 token):

python
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.73

2 个 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 ​

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 ​

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 ​

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()))

动手实践 ​

  1. 运行 examples/ch11_vit.py,把 patch 从 4×4 改成 7×7(4 个 patch),对比参数量与准确率。
  2. 运行 examples/ch11_swin.py,把窗口大小从 4 改成 2,观察窗口数与移动窗口后的像素组成。
  3. 修改 examples/ch11_diffusion.py 的 T(时间步)从 100 改成 50,观察训练 loss 与采样效果。
  4. 给 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 与 ᾱ 的索引要一一对应

章末练习 ​

基础

  1. ViT 如何把图片变成 token 序列?
  2. Swin 为什么用窗口注意力?
  3. 扩散模型的前向与反向过程分别是什么?

提高

  1. 计算:8×8 特征图、4×4 窗口时,窗口内注意力的计算量 vs 全局注意力(序列长度 64)。
  2. 在 ViT 中把 2 层编码器改成 4 层,观察准确率与训练时间。
  3. 修改 DDPM 示例:用余弦 beta 调度替换线性调度,对比 loss 曲线。

挑战

  1. 实现带移动窗口的完整 Swin 块(含 cyclic shift 与 mask),在 32×32 输入上验证窗口间信息传播。
  2. 实现 DDPM 的 DDIM 采样(确定性、少步数),对比 100 步与 20 步的采样质量。

章末自测 ​

  1. ViT 中图片被切成?
    • A. 像素
    • B. patch
    • C. 通道
    • D. 窗口
  2. ViT 分类用哪个 token 的输出?
    • A. 第一个 patch
    • B. [CLS]
    • C. 平均池化
    • D. 最后一个 patch
  3. Swin 的移动窗口解决了?
    • A. 数据不足
    • B. 跨窗口信息流动
    • C. 过拟合
    • D. 梯度消失
  4. 窗口注意力的复杂度与图像尺寸?
    • A. 平方关系
    • B. 线性关系
    • C. 无关
    • D. 指数关系
  5. DDPM 的反向过程学习预测?
    • A. 图片
    • B. 噪声 ε
    • C. 类别
    • D. 时间步
  6. 扩散模型采样从什么开始?
    • A. 一张真图
    • B. 纯噪声
    • C. 随机类别
    • D. 一个 patch
  7. ViT 的缺点之一是?
    • A. 无法训练
    • B. 数据饥饿(需要大规模预训练)
    • C. 参数太少
    • D. 只能分类
  8. DDPM 的 U-Net 输入包括?
    • A. 图片和时间步
    • B. 图片和标签
    • C. 噪声和时间步
    • D. 只有图片
  9. Stable Diffusion 等产品使用的生成技术是?
    • A. VAE
    • B. 扩散模型
    • C. GAN
    • D. RNN
  10. Swin Transformer 提出于?
    • A. 2019
    • B. 2021
    • C. 2020
    • D. 2023