归一化轨迹模型

Normalizing Trajectory Models

arXiv: 2605.08078v1

论文信息

标题: Normalizing Trajectory Models

作者: Jiatao Gu, Tianrong Chen, Ying Shen, et al.

发布日期: 2026-05-08

arXiv ID: 2605.08078v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:扩散模型在减少采样步数(如从 50 步压缩到 4 步)时,原本假设每一步为 “简单高斯去噪” 的近似会严重失效,导致生成质量急剧下降。论文要解决这一根本性瓶颈。
  • 核心方法:提出归一化轨迹模型(NTM),将生成轨迹中的每一步条件分布 p(xs∣xt)p(\boldsymbol{x}_s \mid \boldsymbol{x}_t) 显式建模为一个可精确计算似然的条件归一化流(Normalizing Flow),替代传统的单高斯假设。
  • 关键结果:在文生图基准 GenEval 上,仅需 4 步去噪,从零训练的 NTM 即获得 0.82 的综合分数,不仅远超此前归一化流模型 STARFlow 的 0.56(需 256 步),也匹配或超越了许多需要更多步数的强扩散基线(见表 1)。
  • 主要局限:单步生成(T=1T=1)效果会严重劣化,因为此时整个数据分布的非高斯结构完全由轻量化的传输器承担,超出了其有限容量。论文需更多步数来分摊非高斯建模压力(见第 5 节)。
  • 适合读者:对生成模型(扩散模型、流匹配、归一化流)、少步推理加速、表示学习与概率建模结合感兴趣的学界与业界研究者。

论文背景和研究动机

基于扩散的模型及流匹配模型已成为高保真图像生成的主流范式。这些方法将生成过程拆解为大量微小的时间步,每一步被建模为一个高斯转移,其均值由神经网络预测。当步数足够多、步长足够小时,真实的逆向条件分布的确接近高斯,因此这一近似是准确的。然而,为了提高推理效率而将步数压缩到极少数(如 4 步)时,每一步必须跨越一个很大的时间区间,此时真实的条件分布 p(xs∣xt)p(\boldsymbol{x}_s \mid \boldsymbol{x}_t) 实际上是多个高斯分布的混合,可能呈现多模态和重尾特征,单高斯假设就成为了制约少步生成质量的根本瓶颈(见论文第 2.2 节)。

现有的少步方法多通过在预训练模型上进行蒸馏、一致性训练或对抗性目标来避开这一问题,但它们都牺牲了似然(likelihood)框架,不再提供对生成过程可追踪的概率密度。DDGAN 等工作尝试用对抗网络学习隐式非高斯分布,但引入了模式坍塌和训练不稳定等问题。

基于此背景,论文提出了一个根本性思路:既然单高斯预测不够,那就将每一步的逆向条件分布升级为一个强大的条件归一化流,且该流可进行精确的似然训练。由此,归一化轨迹模型(Normalizing Trajectory Models, NTM) 应运而生。

核心方法和技术细节

NTM 的核心思想,是在生成轨迹的每一对相邻时间步 (t,s)(t, s) 之间(s<ts < t),将条件分布 p(xs∣xt)p(\boldsymbol{x}_s \mid \boldsymbol{x}_t) 建模为一个条件归一化流,从而获得模型在该步上精确的对数似然。它由两个关键部件组成(图 3):

  1. 可逆传输器(Invertible Transporter)fTf_{\mathcal{T}}:一个由若干浅层自回归流模块(TarFlow 风格)堆叠而成的空间映射。它将原始数据空间中的样本 xt\boldsymbol{x}_t 和 xs\boldsymbol{x}_s 分别映射到一个隐空间表示 ut\boldsymbol{u}_t 和 us\boldsymbol{u}_s。该映射是维度不变的,并且其雅可比行列式 log⁡∣det⁡JfT∣\log|\det J_{f_{\mathcal{T}}}| 可被精确计算。
  2. 高斯预测器(Gaussian Predictor)fPf_{\mathcal{P}}:一个深度 Transformer 网络,它从更嘈杂的隐表示 ut\boldsymbol{u}_t 出发,预测目标隐表示 us\boldsymbol{u}_s 的均值 μP\boldsymbol{\mu}_{\mathcal{P}} 和对角方差 σP2\boldsymbol{\sigma}_{\mathcal{P}}^2,从而在隐空间形成一个简单的高斯条件分布 N(μP,diag⁡(σP2))\mathcal{N}(\boldsymbol{\mu}_{\mathcal{P}}, \operatorname{diag}(\boldsymbol{\sigma}_{\mathcal{P}}^2))。

模型的训练目标不是常见的均方误差(MSE),而是精确的负对数似然(NLL)。在隐空间用预测器计算高斯对数概率,再加上传输器带来的雅可比对数行列式,就得到了 p(xs∣xt)p(\boldsymbol{x}_s \mid \boldsymbol{x}_t) 的精确似然。通过变量替换公式,整个轨迹的负对数似然损失最终可被简化为一个优雅的形式:

LNTM=∑k=1T[12∥zk∥2+∑n(log⁡σP(k,n)+∑ℓlog⁡σT(k,ℓ,n))],\mathcal{L}_{\text{NTM}} = \sum_{k=1}^{T}\Big[\frac{1}{2}\|\boldsymbol{z}_k\|^2 + \sum_{n}\Big(\log\boldsymbol{\sigma}_{\mathcal{P}}^{(k,n)} + \sum_{\ell}\log\boldsymbol{\sigma}_{\mathcal{T}}^{(k,\ell,n)}\Big)\Big],

其中 zk=(utk−1−μP)/σP\boldsymbol{z}_k = (\boldsymbol{u}_{t_{k-1}} - \boldsymbol{\mu}_{\mathcal{P}})/\boldsymbol{\sigma}_{\mathcal{P}}。这意味着,NTM 将一个本属于表示学习范畴的 “预测器-编码器” 架构,通过可逆性约束,巧妙地转化为了一个可精确优化似然的归一化流框架(见公式 3.4)。

此外,该框架具备两大扩展能力:

  • 从预训练模型初始化:通过将传输器初始化为恒等映射,并将预测器的输出初始化为预训练流匹配模型的去噪后验均值,同时设置一个零初始化的尺度修正项,可使 NTM 在训练之初完全等价于原预训练模型。再辅以均值对齐辅助损失,便可稳定地将强大的预训练模型转化为 NTM(见第 3.3 节)。
  • 轨迹级得分去噪与蒸馏:NTM 为整个生成轨迹提供了精确的联合分布。对任意生成的噪声轨迹,NTM 损失函数关于轨迹样本的梯度,就是该轨迹的 “联合得分”。利用这个得分以及根据前向过程推导出的轨迹协方差矩阵,可以对整条轨迹进行精细的梯度修正。这一计算代价较高的自修正过程,可以进一步通过一个轻量化的 “去噪器” 网络进行蒸馏学习,从而在 4 步推理的最后,以一个简单的前馈网络替代复杂迭代的梯度修正,实现高质量生成(见第 3.4 节)。

创新点和贡献

与已有工作相比,NTM 做出了以下几点核心贡献:

  1. 首个具备精确似然的少步生成框架:NTM 首次将每个逆向条件分布建模为精确的条件归一化流,不同于蒸馏、一致性模型和 GAN 方法,它在实现少步生成的同时,完整保留了概率框架下的精确似然(见论文第 1 节与第 3 节)。
  2. 架构层面的 “深度-宽度” 权衡再思考:论文将 NTM 定位为 “纯归一化流”(如 STARFlow)与 “纯流匹配” 之间的折衷方案。它将非高斯建模的深度分配在 “轨迹宽度”(多个时间步)上,而非 “每一步的内部深度” 上,从而能以每步极浅的传输器实现强大的整体表达能力,为生成模型架构设计提供了新思路(见第 5 节)。
  3. 稳定且可扩展的训练策略:论文提供了一套从零训练(图 6)和从预训练模型微调(图 7)的完备方案。特别是微调方案中的 “均值对齐辅助损失”,被证明是防止灾难性遗忘和保证训练稳定性的关键,解决了直接微调大型归一化流时容易出现的发散问题(见第 4.4 节)。
  4. 轨迹级别的自修正机制:利用了生成轨迹在马尔可夫前向过程中的天然相关性,通过协方差矩阵进行跨时间步的联合梯度修正,这比传统方法中对每个样本独立进行得分去噪更为有效,并能进一步蒸馏加速,形成完整的训练-加速链条(见第 3.4 节)。

实验结果分析

NTM 在文生图任务上展现出其 “少步高效” 的显著优势。

  • 从零训练:在 GenEval 基准上,仅用 4 步的 NTM(256×256 分辨率)即取得 0.82 的综合分数。这不仅显著超越了同为归一化流的 STARFlow(0.56,256 步),更直接媲美当下许多主流的扩散模型,如 SD3-Medium(0.62)、FLUX.1-dev(0.66)等(表 1)。这充分证明了精确似然训练在少步生成场景下的潜力。
  • 微调预训练模型:通过微调一个 40 亿参数的流匹配模型(FLUX.2-klein),NTM 在 512×512 分辨率下同样取得强大效果,验证了其方法的可扩展性(表 1)。

论文中的消融研究同样关键,揭示了其设计的合理性:

  • 辅助损失的必要性: 若不使用均值对齐损失(λ=0\lambda=0),从预训练模型微调的 NTM 将在训练早期发散(图 7a)。这证明了在优化复杂似然目标时,需要一个 “锚点” 来维系从预训练模型继承的知识。
  • 少步优于单步: 实验明确指出,当步数缩减至 T=1T=1 时,NTM 的生成质量会严重退化(图 8)。这表明此时浅层传输器的容量成为了瓶颈,无法独立承担整个数据分布的非高斯建模任务,侧面验证了其 “用轨迹分摊复杂度” 的设计哲学。

实践建议

NTM 的设计理念和训练策略,为实际工程落地提供了明确的指引。

  1. 在资源受限下权衡成本与质量:NTM 的深度预测器 + 浅层传输器结构,本质上是用一个可并行计算的深度网络(预测器)来处理跨时间步的推理,而仅用少量串行计算(传输器)处理每步局部的非高斯细节。在实际部署时,若想追求极致速度,可像论文一样训练一个去噪器替代传输器解码,将推理简化为 “4 步预测器 + 1 次前馈去噪”,实现约 9 倍的速度提升(1.88 img/s vs 0.20 img/s,表 2),代价是样本质量上微小的损失(LPIPS 0.121)。
  2. 微调预训练模型时的稳定技巧:如果你的任务是微调一个现有的扩散或流匹配模型以实现少步生成,建立 “恒等初始化” + “均值对齐辅助损失” 的策略极为关键。具体来说:
    • 将原模型的输出重新参数化为 NTM 中预测器的均值部分,并保证在初始化时数学上完全等价。
    • 设置一个输出为 0 的线性层来学习对原方差的修正。
    • 使用一个较大的系数(如 λ=2.5\lambda=2.5)在训练初期约束预测器均值,使其不要偏离原模型太远,并适时退火,以确保模型在继承全部旧能力的基础上,平稳地学习非高斯结构。
  3. 利用轨迹得分进行后处理优化:即使不依赖蒸馏好的去噪器,NTM 的轨迹级得分修正也是一个强大的后处理工具。在生成一条样本轨迹后,可以利用 Pytorch 的 autograd 机制,对轨迹上的所有样本点进行 backward() 操作,获取联合梯度,并根据论文附录 A.6 中推导的解析协方差矩阵 S\boldsymbol{S} 进行一步修正。该步骤无需任何额外训练,且充分利用了模型内部的信息,能有效提升最终输出样本的保真度。