很多刚接触生成模型的朋友第一次跑通自回归或者 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 12GBPython 3.10PyTorch 2.1.0 CUDA 12.1torchvision 0.16.0numpy 1.24.3matplotlib 3.7.1einops 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 的三种解法这里必须单独拿出来说因为这是复现路上第一个拦路虎。你用下面的写法想直接下载 MNISTtrainset torchvision.datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform)大概率会碰到类似这样的报错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/目录下文件名必须保持原样然后在代码里使用downloadFalse。如果 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里设置shuffleTrue并且开num_workers4预取数据批次大小设为 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_min0.002, sigma_max80.0, rho7.0): t t.clamp(0, 1) return (sigma_min ** (1 / rho) t * (sigma_max ** (1 / rho) - sigma_min ** (1 / rho))) ** rho这个函数接收的时间 t 范围是 [0, 1]t0 时输出接近 sigma_mint1 时输出 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_samples64, steps100, devicecuda): model.eval() ts torch.linspace(1.0, 0.0, steps 1, devicedevice) sigmas sigma_schedule(ts) # 从大到小排列 x torch.randn(n_samples, 1, 28, 28, devicedevice) * 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[i1] - 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 通道 → 下采样1ResBlock(64→64) 下采样到14×14 → 下采样2ResBlock(64→128) 下采样到7×7 → 中间层ResBlock(128→128) Self-Attention → 上采样1ResBlock(128→128) 上采样到14×14并与对应下采样特征拼接 → 上采样2ResBlock(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, devicet.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 Annealingwarmup 500 步批次大小128训练轮数30Weight Decay1e-5EMA 衰减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 附近因为标准高斯噪声的每个分量方差是 1MSE 预测如果输出全零初始 loss 大约就是 1。经过 10 个 epoch 左右loss 会降到 0.2 到 0.3 之间后面下降速度会放缓。到了第 20 个 epochloss 基本稳定在一个平台这时候再往后训练主要是让细节更干净。不过我要提醒一句不要只看 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_norm1.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 步采样对比同一组“初始噪声”生成的图像马上就能看出步数带来的细节变化。这个小习惯帮我省下了大量来回对比的时间也推荐给你用。
