Back to blog

Diffusion Forcing: Next-token Prediction Meets Full-Sequence Diffusion

通过为每个 token 分配独立噪声级别,统一自回归预测与全序列扩散,在视频生成、决策规划和机器人控制中实现稳定长序列生成和灵活引导

Diffusion Forcing: Next-token Prediction Meets Full-Sequence Diffusion

一、论文概述

项目内容
标题Diffusion Forcing: Next-token Prediction Meets Full-Sequence Diffusion
作者Boyuan Chen, Diego Marti Monso, Yilun Du, Max Simchowitz, Russ Tedrake, Vincent Sitzmann
机构MIT CSAIL, Technical University of Munich
论文https://arxiv.org/abs/2407.01392
项目https://boyuan.space/diffusion-forcing
发布2024-07-01 (v4: 2024-12-10)
许可CC BY 4.0

二、核心思想

问题定义

序列建模存在两种主流范式,各有优劣:

范式优势劣势
自回归 Next-Token Prediction变长生成、树搜索、在线反馈控制无法引导采样、连续数据易发散
全序列扩散 (Full-Seq Diffusion)可引导采样到高奖励轨迹、擅长连续信号固定长度、非因果架构、帧间不连续

核心矛盾:将 next-token 模型直接用于全序列扩散会导致生成质量差——因为模型无法建模”早期 token 的小不确定性必然导致后期 token 的高不确定性”这一因果关系。

解决方案概述

Diffusion Forcing (DF) 是一种新的训练和采样范式:为序列中每个 token 分配独立的噪声级别,通过共享的因果 next-token 预测模型对任意长度的噪声序列进行去噪。

核心洞察:加噪是一种部分遮蔽(partial masking)——零噪声意味着 token 未被遮蔽,纯噪声意味着完全遮蔽。DF 迫使模型学习对任意噪声级别的 token 集合进行”解遮蔽”。

Figure 1: Diffusion Forcing 能力对比——结合自回归和全序列扩散的优势

三、技术架构

方法总览

Figure 2: 方法总览——DF 训练因果序列网络对每帧独立噪声级别的灵活长度序列去噪

DF 对比两种传统范式:

  • Teacher Forcing:从 ground-truth 序列预测单个 next token
  • 全序列扩散:非因果架构对所有帧使用相同噪声级别去噪
  • Diffusion Forcing:交错序列时间轴和扩散噪声轴,统一两者优势

核心公式

前向扩散过程(每个 token 独立加噪):

q(xk∣xk−1)=N(xk;1−βkxk−1,βkI)q(\mathbf{x}^k | \mathbf{x}^{k-1}) = \mathcal{N}(\mathbf{x}^k; \sqrt{1 - \beta_k} \mathbf{x}^{k-1}, \beta_k \mathbf{I})

反向去噪过程(参数化模型):

pθ(xk−1∣xk)=N(xk−1;μ(xk,k),γkI)p_\theta(\mathbf{x}^{k-1} | \mathbf{x}^k) = \mathcal{N}(\mathbf{x}^{k-1}; \boldsymbol{\mu}(\mathbf{x}^k, k), \gamma_k \mathbf{I})

DF 训练目标(核心损失函数):

\mathcal{L}(\theta) = \mathbb{E}_{\mathbf{z}_t \sim p_\theta(\mathbf{z}_t | \mathbf{z}_{t-1}, \mathbf{x}_t^{k_t}, k_t)} \sum_{t=1}^{T} \left[ \| \boldsymbol{\epsilon}_t - \boldsymbol{\epsilon}_\theta(\mathbf{z}_{t-1}, \mathbf{x}_t^{k_t}, k_t) \| ^2 \right] \tag{3.1}

其中:

  • k1:Tk_{1:T} 从 [K]T[K]^T 均匀采样(每个 token 独立噪声级别)
  • x1:T\mathbf{x}_{1:T} 从训练数据采样
  • ϵt∼N(0,σkt2I)\epsilon_t \sim \mathcal{N}(0, \sigma_{k_t}^2 \mathbf{I})
  • zt\mathbf{z}_t 是 RNN 隐状态,捕获过去 token 的影响

理论保证(Theorem 3.1):DF 训练过程优化的是所有噪声级别序列的期望对数似然的 ELBO 的重加权版本。在适当条件下,优化 (3.1) 同时最大化所有噪声级别序列的似然下界。

采样过程

DF 采样由 M×TM \times T 噪声调度矩阵 K\mathcal{K} 定义:

  • 列对应时间步 tt
  • 行(索引 mm)决定噪声级别
  • Km,t\mathcal{K}_{m,t} 表示第 mm 行第 tt 列 token 的目标噪声级别

采样流程:

  1. 初始化所有 token 为白噪声(k=Kk = K)
  2. 逐行从上到下迭代,每行内从左到右去噪
  3. 最后一行 m=0m = 0 时所有 token 干净(k=0k = 0)

训练算法(Algorithm 1)

1: loop
2:   采样轨迹 (x_1, ..., x_T)
3:   for t = 1,...,T do
4:     采样独立噪声级别 k_t ∈ {0,1,...,K}
5:     x_t^{k_t} = ForwardDiffuse(x_t, k_t)
6:     更新隐状态 z_t ~ p_θ(z_t | z_{t-1}, x_t^{k_t}, k_t)
7:     计算噪声 ε_t = ε_θ(z_{t-1}, x_t^{k_t}, k_t)
8:   end for
9:   L = MSE Loss([ε_1,...,ε_T], [ε_1,...,ε_T])
10:  反向传播更新 θ
11: end loop

采样算法(Algorithm 2)

1: 输入: 模型 θ, 调度矩阵 K, 初始隐状态 z_0, 引导代价 c(·)
2: 初始化 x_1,...,x_T ~ N(0, σ_K² I)
3: for m = M-1,...,0 do        # 逐行去噪
4:   for t = 1,...,T do        # 行内从左到右
5:     z_t^new ~ p_θ(z_t | z_{t-1}, x_t, K_{m+1,t})
6:     k ← K_{m,t}, w ~ N(0, I)
7:     x_t^new ← 1/√α_k (x_t - (1-α_k)/√(1-α̅_k) ε_θ(z_t^new, x_t, k)) + σ_k w
8:     更新 z_t ← z_t^new
9:   end for
10:  x_{1:H} ← AddGuidance(x_{1:H}, ∇_x log c(x_{1:H}))
11: end for
12: 返回 x_{1:T}

模型实现

  • 架构:卷积 RNN(实验中使用),理论上也可用 masked Transformer
  • 隐状态:zt\mathbf{z}_t 通过循环层演化,捕获过去 token 的信息
  • 输入:前一隐状态 zt−1\mathbf{z}_{t-1} + 当前噪声 token xtkt\mathbf{x}_t^{k_t} + 噪声级别 ktk_t
  • 输出:预测噪声 ϵθ\epsilon_\theta,通过仿射重参数化得到去噪结果

四、核心创新

创新 1:稳定自回归生成

传统自回归模型在连续数据(如视频)上生成超过训练长度时会发散。DF 通过以下方式稳定生成:

  • 用略带噪声的 token(0<k≪K0 < k \ll K)更新隐状态
  • 避免了”完全确定的历史 + 完全未知的未来”的极端情况
  • 实验表明可稳定生成超过训练长度 2-5 倍的序列

创新 2:因果不确定性建模

Figure 5: 噪声级别控制——不同 k 值实现不同效果

DF 通过噪声级别编码因果不确定性:

  • 近未来:低噪声(更确定)
  • 远未来:高噪声(更不确定)

“之字形”采样方案:[x10,x2K/2,x3K]⊤→[x10,x20,x3K/2]⊤→[x10,x20,x30]⊤[\mathbf{x}_1^0, \mathbf{x}_2^{K/2}, \mathbf{x}_3^K]^\top \to [\mathbf{x}_1^0, \mathbf{x}_2^0, \mathbf{x}_3^{K/2}]^\top \to [\mathbf{x}_1^0, \mathbf{x}_2^0, \mathbf{x}_3^0]^\top

创新 3:蒙特卡洛引导 (MCG)

DF 允许通过多次采样未来轨迹并平均引导梯度来影响当前 token 的生成:

  • 对同一未来进行多次采样
  • 平均各次采样的引导梯度
  • 效果类似于 MPPI(Model Predictive Path Integral)控制
  • 结合因果不确定性建模效果更佳

创新 4:灵活序列决策框架

DF 同时作为策略(policy)和规划器(planner):

功能实现方式
策略短前瞻窗口 HH,低延迟在线决策
规划器长前瞻窗口 + 引导,离线长程规划
灵活切换无需重训或修改架构

Token 定义:xt=[at,rt,ot+1]⊤\mathbf{x}_t = [\mathbf{a}_t, \mathbf{r}_t, \mathbf{o}_{t+1}]^\top(动作 + 奖励 + 观测)

五、实验结果

视频生成

Figure 3: 视频生成对比——DF 生成的时间一致性最好,不会发散

在 Minecraft 和 DMLab 数据集上:

  • DF:稳定生成超过训练长度(1000+ 帧),时间一致性好
  • Teacher Forcing:快速发散
  • 因果全序列扩散:帧间不连续,视频跳跃

规划任务(D4RL Maze2D)

环境MPPICQLIQLDiffuser*Diffuser (执行动作)Ours w/o MCGOurs
U-Maze33.25.747.4113.9±3.16.3±2.1110.1±3.9116.7±2.0
Medium10.25.034.9121.5±2.713.5±2.3136.1±10.2149.4±7.5
Large5.112.558.6123.0±6.46.3±2.1142.8±5.6159.0±2.7
单任务平均16.27.747.0119.58.7129.67141.7
Multi U-Maze41.2-24.8128.9±1.832.8±1.7107.7±4.9119.1±4.0
Multi Medium15.4-12.1127.2±3.422.0±2.7145.6±6.5152.3±9.9
Multi Large8.0-13.9132.1±5.86.9±1.7129.8±1.5167.1±2.7
多任务平均21.5-16.9129.420.6127.7146.2

关键发现:

  • Diffuser 直接执行生成的动作时性能急剧下降(需手写 PD 控制器)
  • DF 的原始动作生成自洽,甚至优于 Diffuser 状态预测 + PD 控制器
  • MCG 引导带来显著性能提升

机器人任务

Figure 4: 真实机器人任务——水果位置互换

任务:机械臂交换两个水果的位置(需要记忆初始配置)

方法成功率
Diffusion Policy(无记忆)0%(失败)
DF(正常观测)80%
DF(视觉干扰/遮挡)76%
Next-frame diffusion(扰动观测)48%

DF 通过隐状态自然融入记忆,且对噪声/缺失观测鲁棒。

时间序列预测

在 Electricity 数据集上的预测区间:

Figure 6: 电力数据集预测区间

DF 在多变量时间序列预测上与 prior diffusion 和 transformer 方法具有竞争力。

六、核心创新总结

创新点说明关键优势
独立噪声级别每个 token 独立噪声,统一自回归与扩散结合两者优势
因果不确定性近未来低噪声、远未来高噪声有效长程引导
蒙特卡洛引导多次采样未来平均引导梯度显著提升规划性能
灵活决策框架同时作为策略和规划器无需重训切换
稳定长序列生成超越训练长度 2-5 倍不发散视频/连续信号生成
理论保证优化所有子序列似然的 ELBO数学基础扎实

七、技术影响

对序列建模的指导

  • 统一范式:DF 证明自回归和扩散可以统一,而非二选一
  • 噪声即遮蔽:为理解扩散模型提供了新视角——加噪是部分遮蔽的一种形式
  • 因果+扩散:因果架构与扩散的结合可产生新能力(如 MCG)

应用前景

  • 视频生成:稳定长视频生成,避免滑动窗口
  • 机器人控制:带记忆的模仿学习,对噪声鲁棒
  • 决策规划:灵活的在线/离线规划切换
  • 时间序列:通用序列建模

局限性

  1. 架构限制:当前基于 RNN,高分辨率视频需要更大的 Transformer
  2. 规模验证:未在互联网规模数据集上验证 scaling 行为
  3. 计算成本:多次采样未来用于 MCG 增加推理开销
  4. 训练复杂度:独立噪声级别采样增加了训练的随机性

八、参考资源

论文

相关工作

  • Diffuser (Janner et al., ICML 2022): 全序列扩散规划
  • Diffusion Policy (Chi et al., 2024): 扩散策略模仿学习
  • AR-Diffusion (Wu et al., NeurIPS 2023): 自回归文本扩散
  • DDPM (Ho et al., NeurIPS 2020): 去噪扩散概率模型
  • MPPI (Williams et al., 2015): 模型预测路径积分控制

代码与资源