视觉 Transformer 与生成模型 视觉 Transformer(Vision Transformer)把自注意力用到了图像块上,用数据驱动的空间学习挑战了 CNN 的统治地位。本文件涵盖 ViT、DeiT、Swin Transformer,以及用 GAN(StyleGAN)、VAE 和扩散模型(DDPM、Stable Diffusion)做图像生成,再加上超分辨率和神经风格迁移。 CNN(文件 02)内置了很强的空间归纳偏置:局部连接、权重共享、平移等变性。视觉 Transformer 提出了一个很尖锐的问题:如果我们完全丢掉这些归纳偏置,只用第 6 章的注意力机制,让模型从数据中自己学空间结构,会怎么样?
视觉 Transformer(Vision Transformer)把自注意力用到了图像块上,用数据驱动的空间学习挑战了 CNN 的统治地位。本文件涵盖 ViT、DeiT、Swin Transformer,以及用 GAN(StyleGAN)、VAE 和扩散模型(DDPM、Stable Diffusion)做图像生成,再加上超分辨率和神经风格迁移。
CNN(文件 02)内置了很强的空间归纳偏置:局部连接、权重共享、平移等变性。视觉 Transformer 提出了一个很尖锐的问题:如果我们完全丢掉这些归纳偏置,只用第 6 章的注意力机制,让模型从数据中自己学空间结构,会怎么样?
视觉 Transformer(Vision Transformer,ViT)(Dosovitskiy 等,2021)直接把标准 Transformer 编码器用到图像上。核心想法是把一张图像当作一个图像块序列来处理,就像 NLP 把文本当作 token 序列一样。
流程如下:
**图像块嵌入(patch embedding)**等价于一次核大小为 P、步长为 P(不重叠)的卷积。ViT 字面意义上就是把二维图像转换成一维序列,然后用和语言一样的架构来处理它。
ViT 的归纳偏置比 CNN 少:它不强求局部连接或平移等变性。这意味着它需要更多训练数据才能从头学到空间结构。在小数据集上,CNN 表现优于 ViT。但当在超大数据集(JFT-300M,3 亿张图像)上训练时,ViT 能匹敌甚至超过最好的 CNN,这说明 CNN 的归纳偏置有助于数据效率,但对最终性能来说并不是必需的。
ViT 的自注意力在图像块数量上是 O(N^2) 的。对 224x224 的图像、16x16 的图像块,N = 196,还能接受。但对更高分辨率或更小块,二次方的代价就变得不可接受。
DeiT(Data-efficient Image Transformer,Touvron 等,2021)表明,仅靠 ImageNet(而不需要庞大的 JFT 数据集)也能有效训练 ViT,方法是用强数据增强、正则化(随机深度、标签平滑、dropout)以及知识蒸馏(knowledge distillation):一个预训练的 CNN 老师提供软标签,ViT 学生去学习匹配它。DeiT 在 [CLS] token 之外加了一个蒸馏 token,专门训练来预测老师的输出。
Swin Transformer(Liu 等,2021)针对 ViT 的两个主要局限:对图像尺寸的二次方代价,以及缺少层次化特征图(检测和分割需要的)。
Swin 引入了移位窗口(shifted windows):与其在所有图像块上做全局自注意力,不如在局部窗口内(比如 7x7 个图像块)做注意力。这让代价对图像尺寸是线性的:O(N) 而不是 O(N^2)。但光靠局部窗口会阻碍区域之间的信息流动。
**窗口移位(window shifting)**解决了这个问题:在交替的层里,窗口划分平移半个窗口大小。这创造了跨窗口的连接,让信息无需全局注意力的代价就能在各层之间流遍图像的所有部分。
Swin 还通过在各阶段之间合并图像块来构建层次化表示。每个阶段结束后,相邻的 2x2 图像块被拼接并投影,通道维度翻倍、空间分辨率减半。这产生了类似于 CNN 和 FPN(文件 03)的多尺度特征图,让 Swin 能直接对接 Faster R-CNN 这样的检测头和 U-Net 这样的分割头。
PVT(Pyramid Vision Transformer)采用类似的层次化思路,配合空间降维注意力:在每个阶段,先对键和值做空间下采样,再算注意力,在保持全局感受野的同时降低二次方代价。
**自监督视觉学习(self-supervised visual learning)**从无标注图像中训练表示。标签很贵,但图像很丰富。目标是学到能很好地迁移到下游任务、且完全不需要人工标注的特征。
**对比学习(contrastive learning)**训练模型识别出同一张图像的两个增强视图("正样本对")应当有相似的表示,而不同图像的视图("负样本对")应当有不同的表示。
SimCLR(Chen 等,2020)为批量里的每张图像创建两个增强视图,用共享的主干 + 投影头对两者编码,再用 NT-Xent 损失(归一化的温度缩放交叉熵):
其中 \text{sim} 是余弦相似度(第 1 章),\tau 是温度参数。分子把正样本对拉近;分母把负样本对推开。SimCLR 需要很大的批量(4096 以上)来提供足够多的负样本。
MoCo(Momentum Contrast,He 等,2020)通过维护一个动量更新的负样本嵌入队列解决了大批量的要求。查询编码器用梯度下降更新;键编码器作为查询编码器的指数移动平均(EMA,第 4 章)来更新:\theta_k \leftarrow m \theta_k + (1 - m) \theta_q,其中 m = 0.999。队列存储近期的键嵌入,提供一个又大又一致的负样本集合,而不需要超大批量。
BYOL(Bootstrap Your Own Latent,Grill 等,2020)彻底去掉了负样本对。它用两个网络:一个"在线"网络和一个"目标"网络(在线网络的 EMA)。在线网络预测目标网络对不同增强视图的表示。在没有负样本的情况下,BYOL 通过预测器头的不对称性和 EMA 目标避免了坍缩问题(即模型对所有输入都输出同一个向量)。
DINO(Self-Distillation with No Labels,Caron 等,2021)把自蒸馏用到 ViT 上。学生网络预测教师网络(学生的 EMA)在不同增强视图上的输出。教师用更大的裁剪;学生用更小的裁剪。DINO 学到的特征里显式地包含了场景布局信息:DINO 训练的 ViT 的自注意力图在没有任何分割监督的情况下就能自然地把物体分割出来。
**掩码图像建模(masked image modelling)**是 BERT 掩码语言建模(第 7 章)的视觉对应物。把输入图像块的一大半掩掉,让模型学着去重建它们。
MAE(Masked Autoencoders,He 等,2022)掩掉 75% 的图像块,训练一个 ViT 编码器-解码器来重建缺失的像素值。只有未掩码的图像块会被编码器处理(预训练时省下 4 倍计算量),轻量级的解码器再从编码后的可见图像块加上可学习的掩码 token 来重建整张图像。
BEiT(BERT Pre-training of Image Transformers,Bao 等,2022)掩掉图像块,预测离散的视觉 token(由一个预训练的 dVAE 分词器得到),而不是原始像素。这和 BERT 预测离散词 token 一脉相承,也避开了像素重建里的低层细节。
**图像生成(image generation)**旨在产生训练集里不存在的、逼真的新图像。核心挑战是建模自然图像这个高维概率分布。
生成对抗网络(Generative Adversarial Networks,GAN)(Goodfellow 等,2014)用两个相互竞争的网络:一个生成器(generator) G 从随机噪声生成假图像,一个判别器(discriminator) D 试图区分真图和假图。它们对抗地训练:G 想骗过 D,D 想抓住 G。
生成器取一个随机潜在向量 z(从高斯之类的简单分布采样),通过一系列转置卷积把它映射成一张图像。判别器是一个标准的 CNN 分类器。在平衡点上,G 生成的图像和真实数据无法区分,D 对所有输入都输出 0.5。
**模式坍缩(mode collapse)**是 GAN 的主要失败模式:生成器学会只生成少数几种能骗过判别器的图像,忽略了训练数据的多样性。生成器找到了一小撮"安全"的输出,而不是覆盖整个分布。
稳定 GAN 训练的技巧包括:谱归一化(约束判别器的 Lipschitz 常数)、渐进式增长(先在低分辨率上训练,再逐步提高分辨率)、特征匹配(匹配判别器中间特征的统计量而不是最终输出),以及用 Wasserstein 距离替代原始的 JS 散度目标。
StyleGAN(Karras 等,2019)是高质量图像合成领域最有影响力的 GAN 架构。它的关键创新是基于风格的生成器(style-based generator):与其把潜在向量 z 直接喂给生成器,不如先用一个**映射网络(mapping network,一个 8 层 MLP)把它映射成一个风格向量 w。这个风格向量通过自适应实例归一化(adaptive instance normalisation,AdaIN)**注入到生成器的每一层,调制特征图的统计量:
其中 y_s 和 y_b 是从 w 推导出的缩放和偏置。不同层控制不同方面:浅层控制粗粒度特征(姿态、脸型),中层控制中等粒度特征(发型、眼睛),深层控制细节(雀斑、发丝纹理)。StyleGAN 能生成 1024x1024 分辨率的光照级真实人脸。
变分自编码器(Variational Autoencoders,VAE)(第 6 章)提供了另一种生成方法。和 GAN 不同,VAE 有原则性的概率框架和明确的训练目标(ELBO)。它们生成的图像往往比 GAN 模糊,但潜在空间更平滑、更有结构。VAE 也是潜在扩散模型中用于在图像和潜在空间之间压缩的编码器-解码器对。
**扩散模型(diffusion models)**已成为图像生成的主流范式,在质量和多样性上都超过了 GAN。它的概念很简单:逐步给数据加噪声,直到它变成纯高斯噪声(前向过程),然后学着一步一步地反转这个过程(反向过程)。
前向过程在 T 个时间步上添加高斯噪声:
DDPM(Denoising Diffusion Probabilistic Models,Ho 等,2020)确立了这个框架。采样需要遍历全部 T 步(通常是 1000),很慢。DDIM(Denoising Diffusion Implicit Models,Song 等,2021)把采样过程重新表述成一个确定性映射,允许大步跳跃(比如 50 步代替 1000 步),质量损失极小。
基于分数的模型(score-based models)(Song 和 Ermon,2019)提供了一个替代视角。模型不预测噪声 \epsilon,而是估计分数函数(score function) \nabla_{x_t} \log p(x_t),即对数概率对带噪声图像的梯度。这个梯度指向数据分布中更高概率(更干净)的区域。采样时按这个梯度用 Langevin 动力学前进。基于分数的模型和 DDPM 在**随机微分方程(SDE)**的框架下被统一:前向过程是一个加噪声的 SDE,反向过程是时间反转的 SDE。
无分类器引导(classifier-free guidance)(Ho 和 Salimans,2022)控制样本质量和多样性之间的权衡。模型同时被条件式(带文本提示或类别标签)和无条件式(随机丢弃条件)训练。采样时,预测是两者的加权组合:
其中 c 是条件,\varnothing 是空条件,s > 1 是引导尺度。s 越大,图像越强地匹配条件,但多样性越低。s = 1 给出无引导的模型;s = 7.5 是常用的默认值。
潜在扩散(latent diffusion)(Rombach 等,2022;Stable Diffusion)把扩散过程从像素空间搬到一个学习到的潜在空间。一个预训练的 VAE 编码器把图像压缩成更低维的潜在表示(通常是 4 倍或 8 倍空间下采样),扩散在这个压缩空间里进行,VAE 解码器再从去噪后的潜在表示重建像素。这高效得多:在像素空间扩散一张 512x512 图像意味着处理一个 512 \times 512 \times 3 的张量,而在潜在空间只需处理一个 64 \times 64 \times 4 的张量。
潜在扩散里的去噪 U-Net 接收带噪声的潜在表示、时间步(编码成正弦嵌入,类似于 Transformer 里的位置编码)以及一个条件信号(来自一个冻结的 CLIP 或 T5 文本编码器的文本嵌入)。文本条件通过 U-Net 内部的交叉注意力层进入:文本嵌入作为键和值,图像特征作为查询。这让模型能在每个空间位置上注意到文本提示的相关部分。
**流匹配(flow matching)**是扩散的一种新兴替代方案,它学习的是噪声和数据之间的一条直接传输路径,而不是 DDPM 那种迭代去噪。
一个**连续归一化流(continuous normalising flow,CNF)**定义了一个时变速度场 v_\theta(x, t),沿平滑轨迹把样本从简单分布 p_0(噪声)推到数据分布 p_1。这个变换遵循一个常微分方程(ODE):
从 x_0 \sim \mathcal{N}(0, I) 出发,把 ODE 向前积分到 t = 1,就得到一个来自数据分布的样本。速度场由神经网络参数化,训练去匹配一个目标条件流。
最优传输(Optimal Transport,OT)流匹配(Lipman 等,2023)用噪声和数据之间的直线作为目标流:从噪声样本 x_0 到数据样本 x_1 的条件路径就是 x_t = (1 - t) x_0 + t x_1,目标速度是 v = x_1 - x_0。训练损失变成:
整流流(rectified flows)(Liu 等,2022)迭代地把学到的流路径拉直。第一轮训练后,用模型通过模拟 ODE 生成(噪声,数据)对。这些对之间比随机配对更对齐,用它们重新训练模型。重复这个过程会让路径越来越直,从而可以用更少的 ODE 步数(甚至一步)走完,实现极快的生成。
流匹配相比扩散有几个优势:训练目标更简单(直接的速度回归,不需要噪声调度),采样 ODE 更平滑(需要更少的积分步数),而且和最优传输的联系提供了理论支撑。Stable Diffusion 3 和 Flux 用的是流匹配,而不是传统的 DDPM。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def create_patch_embedding(image, patch_size, d_model, params): """把图像转换成图像块嵌入序列。""" H, W, C = image.shape n_patches_h = H // patch_size n_patches_w = W // patch_size n_patches = n_patches_h * n_patches_w # 提取图像块 patches = [] for i in range(n_patches_h): for j in range(n_patches_w): patch = image[i*patch_size:(i+1)*patch_size, j*patch_size:(j+1)*patch_size, :] patches.append(patch.ravel()) patches = jnp.stack(patches) # (N, P*P*C) # 线性投影到 d_model embeddings = patches @ params['proj_w'] + params['proj_b'] # (N, d_model) # 前置 CLS token cls_token = params['cls_token'] # (1, d_model) embeddings = jnp.concatenate([cls_token, embeddings], axis=0) # (N+1, d_model) # 加位置嵌入 embeddings = embeddings + params['pos_embed'] # (N+1, d_model) return embeddings, patches # 配置 H, W, C = 32, 32, 3 patch_size = 8 d_model = 64 n_patches = (H // patch_size) * (W // patch_size) # 16 key = jax.random.PRNGKey(42) keys = jax.random.split(key, 5) # 构造一张四个象限颜色不同的合成图 image = jnp.zeros((H, W, C)) image = image.at[:16, :16, 0].set(1.0) # 左上红 image = image.at[:16, 16:, 1].set(1.0) # 右上绿 image = image.at[16:, :16, 2].set(1.0) # 左下蓝 image = image.at[16:, 16:, :2].set(1.0) # 右下黄 params = { 'proj_w': jax.random.normal(keys[0], (patch_size**2 * C, d_model)) * 0.02, 'proj_b': jnp.zeros(d_model), 'cls_token': jax.random.normal(keys[1], (1, d_model)) * 0.02, 'pos_embed': jax.random.normal(keys[2], (n_patches + 1, d_model)) * 0.02, } embeddings, patches = create_patch_embedding(image, patch_size, d_model, params) print(f"Image shape: {image.shape}") print(f"Patch size: {patch_size}x{patch_size}") print(f"Number of patches: {n_patches}") print(f"Patch vector length: {patch_size**2 * C}") print(f"Embedding shape: {embeddings.shape} (CLS + {n_patches} patches)") # 可视化图像块 fig, axes = plt.subplots(2, 5, figsize=(14, 6)) axes[0, 0].imshow(image); axes[0, 0].set_title('Full Image'); axes[0, 0].axis('off') for idx in range(min(9, n_patches)): ax = axes[(idx+1) // 5, (idx+1) % 5] patch_img = patches[idx].reshape(patch_size, patch_size, C) ax.imshow(patch_img); ax.set_title(f'Patch {idx}'); ax.axis('off') plt.suptitle('ViT Patch Decomposition') plt.tight_layout(); plt.show()
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def generator(z, params): h = jnp.tanh(z @ params['g_w1'] + params['g_b1']) h = jnp.tanh(h @ params['g_w2'] + params['g_b2']) return h @ params['g_w3'] + params['g_b3'] def discriminator(x, params): h = jax.nn.leaky_relu(x @ params['d_w1'] + params['d_b1'], 0.2) h = jax.nn.leaky_relu(h @ params['d_w2'] + params['d_b2'], 0.2) return jax.nn.sigmoid(h @ params['d_w3'] + params['d_b3']) def init_params(key): keys = jax.random.split(key, 6) z_dim, h_dim, data_dim = 2, 32, 2 scale = 0.1 return { 'g_w1': jax.random.normal(keys[0], (z_dim, h_dim)) * scale, 'g_b1': jnp.zeros(h_dim), 'g_w2': jax.random.normal(keys[1], (h_dim, h_dim)) * scale, 'g_b2': jnp.zeros(h_dim), 'g_w3': jax.random.normal(keys[2], (h_dim, data_dim)) * scale, 'g_b3': jnp.zeros(data_dim), 'd_w1': jax.random.normal(keys[3], (data_dim, h_dim)) * scale, 'd_b1': jnp.zeros(h_dim), 'd_w2': jax.random.normal(keys[4], (h_dim, h_dim)) * scale, 'd_b2': jnp.zeros(h_dim), 'd_w3': jax.random.normal(keys[5], (h_dim, 1)) * scale, 'd_b3': jnp.zeros(1), } def d_loss(params, real_data, fake_data): real_score = discriminator(real_data, params) fake_score = discriminator(fake_data, params) return -jnp.mean(jnp.log(real_score + 1e-7) + jnp.log(1 - fake_score + 1e-7)) def g_loss(params, fake_data): fake_score = discriminator(fake_data, params) return -jnp.mean(jnp.log(fake_score + 1e-7)) # 真实数据:环形分布 key = jax.random.PRNGKey(42) theta = jax.random.uniform(key, (512,)) * 2 * jnp.pi real_data = jnp.stack([jnp.cos(theta), jnp.sin(theta)], axis=1) real_data = real_data + jax.random.normal(key, real_data.shape) * 0.05 params = init_params(jax.random.PRNGKey(0)) d_grad = jax.grad(d_loss) g_grad = jax.grad(g_loss) lr = 0.001 snapshots = [] for step in range(3000): key, k1 = jax.random.split(key) z = jax.random.normal(k1, (512, 2)) fake_data = generator(z, params) # 更新判别器 grads = d_grad(params, real_data, fake_data) for k in ['d_w1', 'd_b1', 'd_w2', 'd_b2', 'd_w3', 'd_b3']: params[k] = params[k] - lr * grads[k] # 更新生成器 fake_data = generator(z, params) grads = g_grad(params, fake_data) for k in ['g_w1', 'g_b1', 'g_w2', 'g_b2', 'g_w3', 'g_b3']: params[k] = params[k] - lr * grads[k] if step in [0, 500, 1500, 2999]: snapshots.append((step, fake_data.copy())) fig, axes = plt.subplots(1, 4, figsize=(16, 4)) for ax, (step, fake) in zip(axes, snapshots): ax.scatter(real_data[:, 0], real_data[:, 1], s=5, alpha=0.3, c='#3498db', label='Real') ax.scatter(fake[:, 0], fake[:, 1], s=5, alpha=0.3, c='#e74c3c', label='Generated') ax.set_title(f'Step {step}'); ax.set_xlim(-2, 2); ax.set_ylim(-2, 2) ax.set_aspect('equal'); ax.legend(markerscale=3) plt.suptitle('GAN Training: Generator Learns the Ring Distribution') plt.tight_layout(); plt.show()
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def noise_schedule(T, beta_start=0.0001, beta_end=0.02): """线性噪声调度。""" betas = jnp.linspace(beta_start, beta_end, T) alphas = 1.0 - betas alpha_bars = jnp.cumprod(alphas) return betas, alphas, alpha_bars def forward_diffusion(x0, t, alpha_bars, key): """在时间步 t 给 x0 加噪声。""" alpha_bar_t = alpha_bars[t] noise = jax.random.normal(key, x0.shape) xt = jnp.sqrt(alpha_bar_t) * x0 + jnp.sqrt(1 - alpha_bar_t) * noise return xt, noise # 构造一张简单的二维"图像"(棋盘格) img = jnp.zeros((32, 32)) for i in range(4): for j in range(4): if (i + j) % 2 == 0: img = img.at[i*8:(i+1)*8, j*8:(j+1)*8].set(1.0) T = 1000 betas, alphas, alpha_bars = noise_schedule(T) # 可视化前向过程 timesteps = [0, 50, 200, 500, 999] key = jax.random.PRNGKey(42) fig, axes = plt.subplots(1, len(timesteps), figsize=(16, 3.5)) for ax, t in zip(axes, timesteps): key, subkey = jax.random.split(key) xt, noise = forward_diffusion(img, t, alpha_bars, subkey) ax.imshow(xt, cmap='gray', vmin=-2, vmax=2) ax.set_title(f't={t}\n$\\bar{{\\alpha}}$={alpha_bars[t]:.3f}') ax.axis('off') plt.suptitle('Diffusion Forward Process: Progressive Noise Addition') plt.tight_layout(); plt.show() # 简单去噪:训练一个微型网络来预测 t=200 处的噪声 t_denoise = 200 key, k1 = jax.random.split(key) xt, true_noise = forward_diffusion(img, t_denoise, alpha_bars, k1) # 微型"去噪器":只学一个常数的噪声估计(用于演示) noise_estimate = jnp.zeros_like(img) lr = 0.01 for step in range(100): residual = noise_estimate - true_noise noise_estimate = noise_estimate - lr * residual # 反向一步 alpha_bar_t = alpha_bars[t_denoise] x_denoised = (xt - jnp.sqrt(1 - alpha_bar_t) * noise_estimate) / jnp.sqrt(alpha_bar_t) fig, axes = plt.subplots(1, 3, figsize=(12, 4)) axes[0].imshow(img, cmap='gray'); axes[0].set_title('Original $x_0$'); axes[0].axis('off') axes[1].imshow(xt, cmap='gray', vmin=-2, vmax=2) axes[1].set_title(f'Noisy $x_{{200}}$'); axes[1].axis('off') axes[2].imshow(x_denoised, cmap='gray') axes[2].set_title('Denoised (one step)'); axes[2].axis('off') plt.tight_layout(); plt.show() mse = jnp.mean((x_denoised - img)**2) print(f"Denoising MSE: {mse:.4f}")