很多刚接触生成模型的朋友,第一次跑通自回归或者 GAN 的代码时都会有一个共同感受:训练是一回事,采样是另一回事,中间隔着无数玄学。我看完 Generative Modeling via Drifting 这篇论文之后的第一感觉是,它把“玄学”往回收了一点——把生成问题定义成一条时间轴上的漂移过程,再用一个特别简单的 Drift Loss 把网络练出来。这篇文章就是我在 MNIST 上用 Drift Loss 完整复现它的实战记录,包含全部核心代码、调度参数、训练细节和踩过的坑。
我默认读这篇东西的人有过至少一次 PyTorch 训练经验,至少知道什么是反向传播、什么是 U-Net。如果你连这些都不太熟也没关系,关键代码我会一行一行说明白,你照着敲完,应该能在 20 分钟左右看到手写数字从纯噪声里慢慢“漂”出来。
1. 先说清楚 Drifting 是什么,以及为什么值得复现
1.1 从扩散到漂移:一个“换视角”的生成框架
扩散模型(Diffusion Model)这几年已经是生成领域的事实标准,它的思路可以通俗理解成:先把一张干净图像逐步加噪,直到完全变成噪声,然后训练网络学会倒着走,把噪声一步步还原成图像。你在网上看到的绝大多数“AI 画图”教程,底层都是这一套。
Generative Modeling via Drifting 这篇论文本质上也是在讲“前向破坏、反向还原”的故事,但它换了一个更激进的视角:不再把加噪过程看作是对真实图像的逐步污染,而是把数据生成过程理解成一条从某个初始分布出发、沿着时间轴不断“漂移”的轨迹。训练时,我们只需要让网络学会预测每一个时间点上的“漂移方向”,采样时沿着反方向一路走回数据分布即可。
这个视角改变的直观好处有两个。
第一个好处是训练目标变得极简。传统扩散模型里,不同时间步的损失往往要按噪声强度重新加权,有时还要处理方差保持和方差爆炸两种框架的换算。Drift Loss 直接预测加性噪声,不需要复杂的损失加权,实验下来训练曲线更平滑,超参数也少。
第二个好处是时间步与采样过程解耦。你在训练时用的是连续时间 t,均匀采样,但推理时完全可以把采样步数从 1000 步降到 100 步甚至 50 步,不需要重新训练。这对复现者来说非常友好,因为它意味着你可以先用小步数快速看效果,慢慢再往上加步数优化细节。
1.2 Memory Drift 的核心:时间步、噪声水平和一个 Drift Loss
整篇论文最核心的东西,在我看来其实就三个公式级别的东西:噪声调度 σ(t)、加噪方式、以及接近“噪声预测 MSE”的 Drift Loss。
用大白话描述一遍:
- 我们有一个真实图像 x0,它来自 MNIST 训练集。
- 我们随机采样一个时间 t,它决定了当前噪声水平 σ(t)。
- 我们随机采样一个高斯噪声 ε,把 x0 和 ε 按比例混合,得到加噪后的 x_t。
- 让网络接收 x_t 和时间 t,输出它“猜测”的噪声 ε_θ。
- 用 MSE 衡量预测噪声和真实噪声的距离,反传更新网络。
就这么简单,没有 GAN 的判别器,没有 VAE 的重构项,也没有额外正则化。你可能会觉得“这不就是扩散模型吗”,没错,它确实是扩散模型家族的近亲,但 Drift Loss 的关键差异在于对 σ(t) 的构造方式和时间连续化处理,这让它在同样甚至更少的训练步数下,可以拿到和经典 DDPM 相当的效果。
至于论文标题里的 Memory Drift,我个人的理解是:模型在每个时间步都保留了对“干净数据”的记忆,所谓漂移就是在这个记忆约束下一点点离开原数据流形,再靠逆过程回来。任务越简单,这个记忆越容易学,MNIST 恰好就是验证这个想法的最佳起点。
1.3 为什么第一站选 MNIST:便宜、直观、坑少
你要真去复现一篇顶会论文,第一反应肯定是找官方代码。但 Generative Modeling via Drifting 论文官方实现里跑的是 ImageNet 这种级别的大图,对显卡、显存、分布式框架的要求直接劝退个人开发者。所以我挑了 MNIST 作为第一站,原因说白了就是三条:
- 数据小:6 万张 28×28 灰度图,一次性全读进内存也就一两百 MB,不用折腾 DataLoader。
- 训练快:一版轻量 U-Net 在 RTX 3060 上 20 分钟左右就能出结果,方便反复调参。
- 可视化直观:生成出来的数字好不好看,人眼一眼就能判断,不需要依赖 FID 这种指标。
更重要的是,MNIST 虽然简单,但它已经包含了复现生成模型的大部分核心环节:数据加载、噪声调度、时间嵌入、U-Net 主干、EMA、采样循环。你在 MNIST 上把这一条链路跑通,后续迁移到 CIFAR-10 或更高分辨率数据集,只是换模型结构和调参数的问题。这也是为什么我一直建议做生成模型入门的人,先在 MNIST 上把“从零训练到采样”这条路走一遍。
2. 环境准备与 MNIST 数据加载实操
2.1 开发环境与依赖清单
先说一下我用的环境,不一定需要完全一致,但版本太老容易踩到一些奇怪的坑。
我这边用的是:
- Ubuntu 20.04,一张 RTX 3060 12GB
- Python 3.10
- PyTorch 2.1.0 + CUDA 12.1
- torchvision 0.16.0
- numpy 1.24.3
- matplotlib 3.7.1
- einops 0.7.0
如果你的电脑没有 GPU,纯 CPU 也能跑,只是 MNIST 上 30 个 epoch 可能要 1 到 2 小时,不算特别离谱。如果用的是苹果 M 系列芯片,把设备改成mps也基本可以跑通,torchvision 0.16 对 MPS 支持已经比较好了。
依赖安装就一行命令:
pip install torch torchvision numpy matplotlib einops我在代码里没有用到特别重的库,einops只是为了让张量维度变换好看一点,你完全可以用reshape和permute代替。
2.2 避坑第一步:torchvision 下载 MNIST 404 的三种解法
这里必须单独拿出来说,因为这是复现路上第一个拦路虎。你用下面的写法想直接下载 MNIST:
trainset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)大概率会碰到类似这样的报错:
HTTPError: HTTP Error 404: Not Found原因是 torchvision 新版本把 MNIST 的下载地址指向了某个 S3 存储桶,而那个源在某些网络环境下经常返回 404。解决办法我试过好几种,最靠谱的是下面这三条路。
方案一:手动下载四个文件,放到 torchvision 期望的目录里。
你打开任意一个浏览器,手动访问这四个地址:
https://ossci-datasets.s3.amazonaws.com/mnist/train-images-idx3-ubyte.gz https://ossci-datasets.s3.amazonaws.com/mnist/train-labels-idx1-ubyte.gz https://ossci-datasets.s3.amazonaws.com/mnist/t10k-images-idx3-ubyte.gz https://ossci-datasets.s3.amazonaws.com/mnist/t10k-labels-idx1-ubyte.gz下载后放到./data/MNIST/raw/目录下,文件名必须保持原样,然后在代码里使用download=False。如果 torchvision 仍然提示找不到文件,那是因为它需要的是解压后的.idx3-ubyte和.idx1-ubyte文件,你只需要在同一个目录下再解压一份:
cd data/MNIST/raw for f in *.gz; do gunzip -k "$f"; done方案二:用curl或wget命令下载,顺便检查文件大小。正常来说四个 gz 文件加起来大概 11MB,如果下载完发现某个文件只有几 KB,那基本是下载了 404 错误页面,直接删掉重来。
mkdir -p data/MNIST/raw cd data/MNIST/raw curl -LO https://ossci-datasets.s3.amazonaws.com/mnist/train-images-idx3-ubyte.gz curl -LO https://ossci-datasets.s3.amazonaws.com/mnist/train-labels-idx1-ubyte.gz curl -LO https://ossci-datasets.s3.amazonaws.com/mnist/t10k-images-idx3-ubyte.gz curl -LO https://ossci-datasets.s3.amazonaws.com/mnist/t10k-labels-idx1-ubyte.gz方案三:如果你有国内访问更稳定的镜像源,直接下载后改名为上面四个标准文件名。很多时候“404”只是因为默认源不可达,换一个镜像就能解决。我自己在实际操作中会选择先试方案一,因为它的失败点最少,不会受到各种环境变量影响。
2.3 数据预处理细节:像素范围、批大小与数据划分
MNIST 本身是 28×28 的灰度图,像素值范围在 0 到 255 之间。训练生成模型时,我习惯把像素归一化到 [-1, 1],原因是这个区间和标准高斯噪声的分布更接近,网络在预测噪声时数值范围更自然。PyTorch 里的写法是:
transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ])这段代码做的事情是先把 PIL 图像变成 [0, 1] 的 Tensor,再执行(x - 0.5) / 0.5,最终范围就是 [-1, 1]。
注意这里Normalize的参数是均值 0.5、标准差 0.5,不是常规分类任务里用的 0.1307 和 0.3081。分类任务用数据集的真实统计量归一化是为了让特征分布更稳定,但生成任务里我们更关心像素值是否落在对称区间内,这样加噪和去噪的过程更对称。
数据划分直接用 torchvision 内置的 train/test 划分就好。6 万张训练,1 万张测试,没有类别不均衡问题。为了训练稳定,我一般会在DataLoader里设置shuffle=True,并且开num_workers=4预取数据,批次大小设为 128。
3. Drift Loss 的核心实现:公式、代码与调度策略
3.1 训练目标一句话:噪声预测的均方误差
如果要用一句话概括 Drift Loss 的训练过程,那就是:给一张图加已知强度的噪声,让网络猜噪声长什么样,猜得越准越好。
形式化地写,假设加噪公式是:
x_t = x0 + σ(t) * ε其中 ε 是标准高斯噪声,σ(t) 是时间相关的噪声强度。训练目标就是最小化:
L = E[ || ε_θ(x_t, t) - ε ||^2 ]这里的 ε_θ 就是网络输出。整个 loss 非常干净,没有额外的感知损失,没有对抗损失,也没有 KL 散度项。
有朋友可能会问,为什么不直接预测 x0?理论上也可以,但实践中预测噪声更稳定。原因是 x0 和 x_t 之间的信噪比从大到小变化剧烈,网络直接回归 x0 会在高噪声阶段面临极大的不确定性;而噪声本身是标准高斯,数值范围固定,预测难度在不同时间步上相对均匀。这也是扩散模型家族普遍采用“预测噪声”或“预测速度”而不是“直接预测图像”的原因。
3.2 噪声调度选择:σ_min、σ_max 与 rho
噪声调度是整个复现里最值得花时间调的部分。我沿用了一种在高级扩散模型里很常见的多项式调度,公式如下:
def sigma_schedule(t, sigma_min=0.002, sigma_max=80.0, rho=7.0): t = t.clamp(0, 1) return (sigma_min ** (1 / rho) + t * (sigma_max ** (1 / rho) - sigma_min ** (1 / rho))) ** rho这个函数接收的时间 t 范围是 [0, 1],t=0 时输出接近 sigma_min,t=1 时输出 sigma_max。也就是说,时间步越往后,噪声越大,图像越模糊。
sigma_max 选择 80 是因为我们希望生成起点真的是“纯噪声”。MNIST 的像素值在 [-1, 1] 之间,当噪声标准差达到 80 时,信号完全被淹没,x_t 的分布和标准高斯乘以 80 几乎没有区别。从这个状态出发,网络有足够的时间逐步“漂回”数据分布。
sigma_min 选择 0.002 是因为最终采样时,我们需要把图像还原到尽量干净的状态。如果 sigma_min 过大,比如 0.02,生成结果会有明显的颗粒感,像蒙了一层细砂纸。0.002 在 MNIST 上足够小,肉眼基本看不出来残留噪声。
rho 控制了噪声调度的中间形态。rho 越大,中间时间段会更多地分布在“接近干净”和“接近纯噪声”两端,高噪声阶段和低噪声阶段都变长,中间过渡区变短。很多实际项目里 rho 取 7 是经验值,亲测在 MNIST 上不需要改动。
3.3 加噪与采样循环:正着漂、反着捞
训练阶段的加噪实现非常直接。核心代码长这样:
# x0: (B, 1, 28, 28),像素范围 [-1, 1] # t: (B,),在 [0, 1] 内均匀采样 # sigma: (B,),由 sigma_schedule 得到 noise = torch.randn_like(x0) x_t = x0 + sigma.view(-1, 1, 1, 1) * noise noise_pred = model(x_t, t) loss = F.mse_loss(noise_pred, noise)注意这里model(x_t, t)输入的第二个参数是 t 本身,而不是 sigma。网络内部会先用时间嵌入模块把 t 编码成高维向量,再通过若干全连接层映射成每组特征图的缩放和偏置参数。有关这个部分我在下一节详细讲。
采样阶段则是反向过程。核心思路是从最大噪声开始,按时间步 t 从 1 递减到 0,每一步用网络预测当前噪声,然后往反方向移动:
@torch.no_grad() def sample(model, n_samples=64, steps=100, device='cuda'): model.eval() ts = torch.linspace(1.0, 0.0, steps + 1, device=device) sigmas = sigma_schedule(ts) # 从大到小排列 x = torch.randn(n_samples, 1, 28, 28, device=device) * sigmas[0] for i in range(steps): t_i = ts[i].expand(n_samples) sigma_i = sigmas[i].expand(n_samples) eps_pred = model(x, t_i) x = x + (sigmas[i + 1] - sigmas[i]) * eps_pred return x这段代码里最关键的是增量方向。因为 sigmas 是从 sigma_max 递减到 sigma_min,而sigmas[i+1] - sigmas[i]是负数,所以每步等价于让 x 减去一部分预测噪声。用更直白的话说:网络说“当前图像里含有这么多噪声”,我们就按图索骥地把它减掉一点,走 100 步就基本还原成干净图像了。
我在采样里没有加随机噪声项,这和 DDPM 的采样稍有区别。Drift 这种连续时间框架配合确定性逆推,采样过程更接近 ODE 求解,稳定性和清晰度都更好。
4. 网络结构设计:MNIST 版轻量生成网络
4.1 主干结构:简化 U-Net + 时间嵌入
模型结构上,我参考了扩散模型常用的 U-Net 主干,但针对 MNIST 做了大幅简化。
我的设计思路是这样的:MNIST 只有 28×28 的单通道灰度图,不需要像 Stable Diffusion 那样堆几十层 ResNet Block,也不需要在多尺度特征上做太多注意力。我采用了一个 3 层下采样 + 3 层上采样的轻量 U-Net,基础通道数设为 64,在下采样到 14×14 和 7×7 时逐步加倍,最后在 7×7 的特征图上加了一个自注意力层。
整体结构可以写成:
输入:x_t (B, 1, 28, 28) → 卷积嵌入到 64 通道 → 下采样1:ResBlock(64→64) + 下采样到14×14 → 下采样2:ResBlock(64→128) + 下采样到7×7 → 中间层:ResBlock(128→128) + Self-Attention → 上采样1:ResBlock(128→128) + 上采样到14×14,并与对应下采样特征拼接 → 上采样2:ResBlock(128→64) + 上采样到28×28,并与对应下采样特征拼接 → 输出卷积:1×1 卷积,输出 (B, 1, 28, 28)每个 ResBlock 内部我会做 GroupNorm,归一化组数设为 8,激活函数用 SiLU。这个组合在生成模型里很常见,比 ReLU 训练更稳。
4.2 时间嵌入与条件注入方式
时间信息是整个模型里除了图像之外最关键的条件输入。我把 t 先做成正弦位置编码,维度是 128,然后过一个两层 MLP,生成两组 128 维向量。在网络内部,每个 ResBlock 都会接收这个时间条件,通过 FiLM 方式调制特征图。
FiLM 通俗理解就是:对特征图先算均值方差做完归一化,然后每个通道乘上一个缩放系数、加一个偏置系数,这两个系数由时间条件决定。这样网络在不同噪声水平下,可以调整各通道的响应强度,从而学会“噪声大时多干活,噪声小时少折腾”。
核心代码片段:
class TimeEmbedding(nn.Module): def __init__(self, dim): super().__init__() self.dim = dim self.mlp = nn.Sequential( nn.Linear(dim, dim * 4), nn.SiLU(), nn.Linear(dim * 4, dim) ) def forward(self, t): half = self.dim // 2 freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half) args = t[:, None] * freqs[None, :] emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1) return self.mlp(emb)在 ResBlock 里使用时,我是这样做的:
# h: 特征图 (B, C, H, W) # scale: (B, C),shift: (B, C) scale = time_emb[:, :C].unsqueeze(-1).unsqueeze(-1) shift = time_emb[:, C:].unsqueeze(-1).unsqueeze(-1) h = h * (1 + scale) + shift这里scale用1 + scale而不是直接乘 scale,是为了保证初始条件接近恒等映射,训练更稳定。这个小细节我在实际复现中吃过亏,后面会再讲。
4.3 训练配置:优化器、学习率和 EMA
模型结构和数据准备搞定后,训练配置决定了你到底能不能稳定跑到最后。
我的超参数如下:
| 配置项 | 数值 |
|---|---|
| 优化器 | AdamW |
| 学习率 | 2e-4 |
| 学习率调度 | Cosine Annealing,warmup 500 步 |
| 批次大小 | 128 |
| 训练轮数 | 30 |
| Weight Decay | 1e-5 |
| EMA 衰减 | 0.9999 |
| 总参数量 | 约 17M |
| 单卡训练时间 | RTX 3060 约 20 分钟 |
EMA 是我强烈建议加的。它的原理很简单:维护一份模型参数的滑动平均,训练结束时用这份平滑参数做采样,而不是直接用最后一步的参数。实践效果是生成图像的人物结构更稳定、噪点更少。PyTorch 里手写 EMA 只需要在每步优化后用一行代码更新缓冲区,这部分网上教程很多,我就不贴完整代码了。
学习率调度也很重要。我的做法是前 500 步线性预热到 2e-4,然后余弦衰减到接近 0。没有 warmup 的时候,我在前几轮训练中明显观察到 loss 震荡更剧烈,后面生成图像偶尔会出现几条不连续的黑色竖线。
5. 训练过程实测与问题排查
5.1 观察指标:loss 下降趋势和生成样例
训练过程中,我主要盯着两个东西:训练 loss 的下降趋势,以及每 5 个 epoch 保存一次的采样图。
先说 loss。Drift Loss 的初始值通常在 1 附近,因为标准高斯噪声的每个分量方差是 1,MSE 预测如果输出全零,初始 loss 大约就是 1。经过 10 个 epoch 左右,loss 会降到 0.2 到 0.3 之间,后面下降速度会放缓。到了第 20 个 epoch,loss 基本稳定在一个平台,这时候再往后训练主要是让细节更干净。
不过我要提醒一句:不要只看 loss 数字,在生成模型里,loss 和生成质量不完全是线性关系。有时候 loss 到了平台期,生成图依然偶发缺笔画或者结构错乱,这时候适当调整采样步数和 sigma_min 比硬训练更多轮更有效。
每 5 个 epoch 保存采样图这个习惯我是强烈建议养成的。因为生成模型训练不像分类任务有明确的准确率曲线,你很难从 loss 判断模型学到什么程度,但采样图一眼就能看出来模型在“模仿”什么阶段:前几个 epoch 生成的数字往往是一团糊影,中间阶段开始有轮廓但笔画扭曲,最后阶段才逐步变得能辨认。
5.2 高频问题一:倒是能跑,但生成的数字很糊
我最早跑出来的结果非常令人沮丧:数字能看出大概轮廓,但边缘模糊,像隔着一层毛玻璃。
排查下来发现两个主因。
第一个原因是采样步数太少。我当时用 20 步采样,想着越快越好,结果发现 Drift 这类方法虽然训练阶段时间连续,但采样步数太少时,每一步减噪幅度太大,误差累积明显。把步数提到 100 步之后,画面清晰度立刻上了一个台阶。这算不算时间成本?采样 100 步在 3060 上也就两秒左右,完全值得。
第二个原因是 sigma_min 偏大。我把 sigma_min 从 0.02 改到 0.002 后,图像边界明显锐利,残留砂砾感消失。这个参数对最终视觉质量的影响非常大,很多人复现失败后跑来问为什么生成图特别“脏”,十有八九是 sigma_min 没调下去。
5.3 高频问题二:训练中期 loss 突然 NaN
我在调大模型通道数、尝试混合精度训练时,遇到过训练到第 8 个 epoch 左右 loss 突然变成 NaN 的情况。这个现象的根源,通常是把 sigma_max 设置得非常大(例如 80)的同时,又开了自动混合精度或者学习率调得太高,导致某些中间层的激活值溢出。
解决方案有三种组合拳:
- 关闭 AMP,全程用 FP32 训练。MNIST 数据集太小,性能瓶颈不在显存和速度,收益不大,没必要冒险。
- 降低学习率到 1e-4。
- 给优化器加梯度裁剪,
max_grad_norm=1.0。
我最后用的是“FP32 + 2e-4 + 梯度裁剪”的组合,之后再也没有出现过 NaN。梯度裁剪这件事,在扩散类模型里我建议无脑开启,它不会显著拖慢训练,但能避免绝大部分数值不稳定问题。
5.4 高频问题三:生成结果多样性差,甚至像同一个数字
这个问题通常出现在模型容量不足或者训练不充分的时候。MNIST 有 10 个类别,如果模型容量太小,学习到的分布会被“压扁”,采样出来的结果最后收敛到几个常见数字上,比如 1 和 0 很多,但 3、5、8 很少。
我排查时首先确认了训练数据里各类别均衡,排除数据问题之后,把基础通道数从 32 提高到 64,多样性立刻改善了不少。因为更大的通道数意味着网络有更多容量去记住不同数字的笔画结构。EMA 在这个问题上也有帮助,因为平滑后的参数往往泛化更好,不容易在训练后期忘记少数字形。
另外我还发现,训练轮数增加到 40 轮之后,某些类别的清晰度会进一步提升,但多样性提升有限。如果你只是想快速验证 Drift Loss 的核心流程,30 轮已经够用;如果想让生成图更丰富,可以把轮数拉到 50,训练时间也才半小时。
6. 复现结果展示与实操经验总结
6.1 我跑出来的效果与基准对比
最后我的复现结果是这样:100 步采样下,生成的手写数字视觉质量已经接近训练集里比较清晰的那一批样本,笔画锐利,数字轮廓完整,10 个类别基本都能出现。30 个 epoch 之后,随机采样 64 张图,主观上大约九成以上能一眼看出是哪个数字,少量图片存在笔画粘连或轻微形变。
为了给自己一个更客观的参考,我用了一个简单的“分类器打分”方法:拿一个在 MNIST 测试集上训练好的 LeNet-5 分类器,对生成的 1000 张图做预测,统计类别分布和平均置信度。结果类别分布接近均匀,平均置信度约 0.85 左右。这个方法当然不能替代正式的 FID,但作为复现过程中的快速量化检查,成本几乎为零,非常好用。
如果拿这个结果和传统 DDPM 在 MNIST 上的效果比,我个人体感是差距很小,但 Drift Loss 的实现代码更简洁,训练也更稳定。这就是复现方法本身最有价值的地方:不是追求吊打所有基线,而是用最少的工程成本拿到足够好的结果。
6.2 三句话讲完本次复现最值得记住的点
第一句:Drift Loss 的本质是时间条件约束下的噪声预测,先别纠结数学推导,把加噪循环和采样循环跑通,你会获得完整直觉。
第二句:调度参数是灵魂,sigma_min、sigma_max、rho 这三个值决定了你生成图的清晰度、多样性和训练稳定性,调参时优先动它们。
第三句:MNIST 不是终点是起点。你把这条路走通了,往 CIFAR-10 迁移其实就是换数据、加通道、改模型结构这三件事。
最后再分享一个小技巧。我在调试采样效果时,不会每次都从随机噪声重新生成,而是固定一个随机种子,这样两次改动之间可以清楚地看到参数变化带来的差异。比如说你先固定种子跑一遍 50 步采样,再跑一遍 100 步采样,对比同一组“初始噪声”生成的图像,马上就能看出步数带来的细节变化。这个小习惯帮我省下了大量来回对比的时间,也推荐给你用。