3 行 for 循环解锁生成模型端到端训练:拆解 Explorative Modeling 并跑通官方代码

深度学习有一条铁律:端到端训练永远打得过手工拆分流水线。AlexNet 证明了它,之后图像分类、检测、分割全被这句话统治。但有一个领域是例外——生成模型。今天最强的自回归和扩散模型,训练时只学「预测一小步」,推理时却要把这一步展开几百上千次,训练和推理用的根本不是同一套采样方式,每一步误差还会喂进下一步,这就是臭名昭著的 暴露偏差(exposure bias)。

UIUC 与哈佛最近放出的论文 Explorative Modeling(XM) 试图补上这块拼图,而它的全部内核,竟然只是一个 3~5 行的 for 循环。代码已开源(PyTorch,含 ImageNet 256×256 / 视频世界模型 / 语言模型全套训练脚本),今天这篇文章带你拆透它,然后直接跑起来。

为什么生成模型不能端到端?根源是「取平均」

分类任务每个输入基本只有一个正确答案,学个确定映射就行。但生成不一样——你让模型「画一只狗」,合法答案无穷多,这些答案就是数据分布里的一个个 mode(独立的峰)。

麻烦在于主流生成模型用的是重构损失(如 MSE)。当一个输入被随机配上一堆合法目标时,重构损失的最优解是这些目标的平均值。而对绝大多数数据,平均值根本不在数据流形上——它落在几个 mode 中间,谁也不像。这就是 mode 模糊(mode blurring):三堆散点被预测成正中间一个点,一张狗的照片糊成一团,一句话退化成重复「the」。

现有模型怎么绕?答案是把生成拆碎:自回归一次只预测一个元素,扩散一次只去一点噪声,每一步目标被切到几乎只剩单一 mode,重构损失就不会取平均。拆生成这条路保住了质量,却杀死了端到端。

作者的追问很直接:生成模型只有两样东西可拆——怎么生成、怎么训练。既然拆生成会毁掉端到端,那为什么不拆训练?

核心机制:best-of-K 探索

XM 拆的就是训练循环本身:每个训练步,模型不只生成一个样本硬凑目标,而是生成 K 个候选,只挑离真实数据最近的那一个回传梯度。官方 README 的伪代码长这样:

# 探索之前:生成一个,直接回传
y = model(sample_latent())          # 从噪声/掩码生成一个输出
loss = recon_loss(y, x)             # 和真实目标 x 打分
loss.backward()

# 探索之后(Forward XM):探索 K 个候选,只保留最好的
losses = []
for _ in range(K):                  # 探索 K 个候选输出
    y = model(sample_latent())      # 生成一个候选
    losses.append(recon_loss(y, x)) # 每个都和 x 打分
min(losses).backward()              # 只对最近的那个候选回传梯度

为什么这能解决 mode 模糊?打个比方:猜飞镖落点,只让你猜一次,最优策略是猜平均位置——那往往是没几支飞镖扎中的地方;但允许猜 K 次、只按最接近的一次算分,最优策略立刻变成把猜测散开,让每次覆盖一簇不同的落点。模型同理:不同输入噪声各自「认领」一个 mode,而不是全挤到中间取平均。探索多少次,就能稳稳抓住多少个 mode。

作者把这个被长期忽视的能力命名为 生成表达力(generative expressivity),并指出它由训练目标决定,参数和数据堆多大都不会自己涨——这也解释了为什么最强模型都那么依赖 guidance:无分类器引导本质是把预测「推离」模糊平均值,模型自己不糊,就不用推。

加到扩散/流模型上:只改几行

XM 和现有生成模型的结合极其简单——探索潜在变量(扩散里就是噪声),只回传最佳候选的梯度。官方示例:

t = sample_timestep()               # 采样一个噪声等级
losses = []
for _ in range(K):                  # 探索 K 个候选噪声
    z = randn_like(x)               # 一个候选噪声
    x_t = add_noise(x, z, t)        # 把数据加噪到 t 级
    losses.append(diffusion_loss(model(x_t, t), x, z))
min(losses).backward()              # 只回传最近候选的梯度

在官方仓库里,探索就是一个命令行开关:--xm_best_of_k K,K=1 就是无探索的 baseline 模型。两个探索方向可以互补组合:

  • Forward XM:固定真实目标,在自己的生成里搜最近的一个 → 偏向「查全」,覆盖所有 mode
  • Reverse XM:固定一个生成,在真实数据里搜最近的一个 → 偏向「查准」,几乎不增加算力开销(代价是可能塌缩到少数 mode)

跑通官方代码

仓库结构很干净:model/ 下是 DiT、flow matching、Jumpy、语言模型实现,job_scripts/<modality>/ 按模态分好训练脚本,slurm_executor.sh 帮你把脚本提交到 Slurm 集群。上手:

git clone https://github.com/alexiglad/XM.git
cd XM
conda create -n xm python=3.12
conda activate xm
pip install -r requirements.txt

# 关键配置
export HF_HOME=/path/to/cache     # 数据集/模型缓存
export HF_TOKEN=...               # ImageNet 需要接受 license
wandb login                        # 日志走 W&B

# 直接跑 ImageNet 256×256 类条件训练(XDiffusion = DiT + 探索)
bash job_scripts/img/pretrain_class_conditional/xdit.sh

# 或者提交到 Slurm(把 example_h100.slurm 里的 TODO 填成你的集群配置)
bash slurm_executor.sh example_h100 job_scripts/img/pretrain_class_conditional/xdit.sh

视频数据(SSv2 / Kinetics-400)需要 ffprobe,先看 data/vid/README.md。FID/FVD 通过 --run_online_evaluation 在训练间隙在线评测,不用等训练完。

实践建议

  • 先理解 K 是第三条 scaling 轴:论文里探索增益随规模增长——数据放大时从 7% 爬到 36%,模型放大时从 13% 爬到 23%,算力翻三倍效率增益翻一倍多。小规模时参数和数据是瓶颈,探索收益不明显;规模越大越值得开。
  • 从 K=2 起步:Reverse XM 几乎不增加算力开销,适合先验证;想要覆盖全部 mode 再加 Forward。别一上来就 K 很大,探索次数 × 前向成本是实打实的。
  • 关注端到端场景:XM 当独立端到端模型用时,行为克隆只需一次前向就能追平需要 100 次前向的 Diffusion Policy;世界建模比 Diffuser 少 16~256 倍推理算力。推理延迟敏感的控制/具身场景是它最值钱的地方。
  • 复现数字做个对照:论文报的 4.1× FLOP 效率、6.2× 样本效率、47% 参数效率、ImageNet 无引导 1.43 FID,建议用 --xm_best_of_k 1 和 --xm_best_of_k 5 各跑一组,自己验证收益曲线。

资源链接

十几年来我们习惯了两个旋钮调生成模型:做大、喂更多数据。这篇论文给了第三个旋钮,朴素到有点可疑:多猜几次,只留最好的那次。但它的所有实验都指向同一件事——前两个旋钮迟早拧到头,第三个才刚刚开始转。

滚动至顶部