论文

StitchVM: 扩散对齐的拼接价值模型

StitchVM: 扩散对齐的拼接价值模型

StitchVM 是一种轻量级模型拼接框架,旨在高效地将预训练的像素空间奖励模型迁移至噪声潜空间,用于扩散模型的对齐。 当前,扩散或基于流的生成模型需要对齐任务特定奖励(如提示保真度或美学偏好),但挑战在于奖励定义于干净输出图像,而对齐过程需要估计噪声中间潜变量的价值函数。现有方法采用 Tweedie 或蒙特卡洛近似,在偏差与计算成本之间权衡:Tweedie 高效但有偏,蒙特卡洛更准确但需昂贵 rollout。一种自然替代是学习价值函数,但如何为噪声潜变量有效训练通用价值模型仍是开放问题。 StitchVM 从现有截断的像素空间奖励模型出发,将冻结的扩散骨干作为其头部附加,形成混合模型。它保留了像素空间模型的强大奖励能力,同时继承了扩散骨干处理噪声潜变量的原生能力。拼接过程极其轻量:例如拼接 CLIP ViT-L 和 SD 3.5 Medium 只需 10 GPU 小时。通过将强大的像素空间奖励模型提升至潜空间,StitchVM 开创了新的扩散对齐范式:不再对每个样本进行粗糙且昂贵的价值函数近似,而是一次性构建适用于真实噪声潜变量的正确函数,并在多个样本和迭代中分摊成本。 实验表明,StitchVM 在多种下游导向和后训练方法中带来改进:DPS 加速 3.2 倍且 GPU 峰值内存减半,DiffusionNFT 加速 2.3 倍。

论文精读

TL;DR StitchVM 通过轻量模型拼接,将像素空间奖励模型迁移至噪声潜在空间,解决扩散对齐中的价值函数估计难题,以极低训练成本实现显著加速(DPS 快 3.2 倍)。

问题

问题背景

扩散模型与流匹配模型在图像生成中广泛应用,但实际部署时需对齐特定奖励信号(如提示词一致性、美学偏好),以控制输出质量。这类奖励通常定义在干净图像上,而对齐过程(无论训练时或推理时)需要在噪声中间潜变量上估计价值函数(value function),二者之间存在天然的不匹配。

现有方法局限

目前主流方案分为两类,各有显著缺陷:

  • Tweedie 式近似:通过单步去噪估计将干净图像奖励回传到噪声状态。计算高效但估计有偏,偏差随噪声增大而恶化,最终损害对齐精度。
  • 蒙特卡洛近似:从噪声状态多次采样至干净图像后计算期望奖励。估计更准确但计算代价极高,每次更新需执行完整去噪过程,单样本开销巨大。

此外,直接为噪声潜变量训练一个强泛化的价值模型属于开放难题,因为训练信号稀疏、状态空间高维且分布随时间步推移剧烈变化,现有工作缺乏有效且通用的训练范式

为什么这个问题难且重要

挑战在于:奖励函数与扩散过程中的中间表征存在语义鸿沟。噪声潜变量缺乏明确的图像级语义,而预训练奖励模型(如 CLIP)仅在像素空间中有效,难以直接迁移。若训练专用价值模型,则面临数据效率、泛化能力与训练稳定性的多重困境,且每个新奖励函数都需要从头训练,成本不可接受。

业界对低开销、高保真的扩散对齐方法有迫切需求:推理阶段的重采样或微调若依赖蒙特卡洛估计,会拖慢生成速度且占用大量显存,阻碍实时应用。因此,如何一次性构建噪声潜变量上的高精度价值函数,并在后续采样中重复摊销成本,成为提升对齐效率的关键。

行业类比

如同在强化学习训练中需要一个低成本、低偏差的 critic 网络来指导策略更新,扩散对齐也需要一个能直接评估“半成品”潜变量质量的模块;StitchVM 通过模型拼接将成熟像素空间奖励模型迁移到噪声潜变量空间,类似于用预训练视觉 backbone 作为特征提取器来快速构建新任务的评估器,从而避免为每个中间状态重新设计或采样估计。

核心洞察

  • 价值函数估计从“每样本近似”转向“一次构建、多次摊销”:StitchVM 直接为扩散过程中的真实噪声潜在变量学习价值函数,而非在推理时重复执行有偏(Tweedie)或高成本(蒙特卡洛)的近似。这一范式转变消除了现有方法在估计偏差与计算开销之间的根本权衡,使得对齐过程中的价值评估既准确又高效,为实现大规模、低延迟的扩散对齐提供了新的可能性。
  • 通过模型拼接实现零样本能力迁移:StitchVM 将冻结的扩散骨干作为轻量“头”拼接在预训练像素空间奖励模型之上,彻底绕开了直接为噪声潜在变量训练价值模型的开放难题。它保留奖励模型的强判别能力的同时,自然获得处理噪声潜在变量的能力,训练开销仅 10 GPU 小时,显著低于从头训练或适应的方法,为跨空间知识迁移提供了可复用的范式。

方法

方法整体思路

StitchVM 旨在高效地将像素空间预训练奖励模型迁移到扩散/流模型的噪声潜在空间,从而在扩散对齐中无需昂贵近似即可获得可靠的价值估计。核心思想是通过模型拼接(model stitching),将处理噪声潜在的能力与已有的奖励能力结合,构建一个直接映射 $(\text{noisy latent}, t) \mapsto \text{reward}$ 的拼接价值模型

输入 → 关键模块 → 输出

  1. 输入:扩散或流模型的中间噪声潜在变量 $\mathbf{z}_t$,以及时间步 $t$(可选,通常由扩散主干内部管理)。
  2. 关键模块
    • 冻结的扩散骨干(如 Stable Diffusion 3.5 Medium 的 UNet):作为处理噪声潜能的“头”,将任意噪声水平的潜在表示转换为适合奖励模型理解的特征。
    • 可学习的拼接层(stitching layer):将扩散骨干的输出特征与截断的奖励模型的中间层对齐。通常是一个轻量的线性投影或小型 Transformer,用于维度匹配和语义桥接。
    • 预训练奖励模型的主体(如 CLIP ViT-L 的 Transformer 层,去掉原始图像输入层):提供强大的美学或语义奖励信号,其参数可以部分微调(例如最后几层)。
  3. 输出:一个标量值,即该噪声潜在状态下的期望奖励估计,用于后续的梯度引导或强化学习更新。

训练过程

  • 数据:从扩散采样轨迹中收集 $\mathbf{z}_t$,并用对应的干净图像的真实奖励(如 CLIP score、美学评分 Oracle)作为监督目标。
  • 优化目标:最小化拼接层输出的奖励估计与真实奖励之间的均方误差,同时可选地加入正则化(如 KL 散度)以保持奖励模型原始能力。
  • 高效性:仅更新拼接层和奖励模型的顶层,扩散骨干完全冻结。在 CLIP ViT-L + SD 3.5 Medium 配置下,全程仅需约 10 GPU 小时

与同类方法的差异

  • 相对于 Tweedie 估计:Tweedie 利用评分函数进行一步近似,计算快但偏差大;StitchVM 通过学习直接映射消除偏差,且一次训练后可摊销使用,无逐样本计算开销。
  • 相对于 Monte Carlo 采样:MC 方法需多次模拟扩散轨迹来估计价值,内存与时间成本高(如 DPS 的多次反向传播);StitchVM 提供解析的梯度,使 DPS 提速 3.2 倍,峰值 GPU 内存减半。
  • 相对于普通潜在空间价值模型:从头训练一个处理噪声潜在的价值模型困难且泛化性差;StitchVM 通过拼接迁移已有强奖励模型,极大降低了训练难度和数据需求。

实验

实验设计

实验覆盖推理时对齐与训练时对齐两类主流范式。推理时方法选取 DPSFK steering,训练时方法选取 DiffusionNFT 及直接奖励微调(含 AlignProp / DRaFT / Flow‑GRPO‑Fast),统一验证 StitchVM 替代传统近似的作用。缝合对象为冻结的扩散主干(SD 3.5 Medium)与预训练像素空间奖励模型(CLIP ViT‑L),仅训练极轻量的缝合层,将奖励值迁移到噪声潜在空间。所有实验在标准提示集上评估图像质量与奖励对齐度,重点测量实际运行时间与 GPU 内存占用。

关键发现

  • DPS 中,StitchVM 直接输出噪声潜在值的精确价值估计,消除逐样本 Monte‑Carlo 展开,实现 3.2 倍加速同时峰值 GPU 内存减半
  • DiffusionNFT 中,用 StitchVM 替换原价值函数后训练流程 加速 2.3 倍,且收敛曲线更稳定。
  • 缝合本身的算力开销极低,仅需 10 GPU‑小时,却能形成一个可摊销到任意后续采样与微调迭代的通用值模型。
  • 跨主干泛化显示,缝合值模型可适配不同扩散体系,不依赖特定去噪器结构。

与基线对比解读

传统 DPS 中常用的 Tweedie 估计有偏且噪声大,Monte‑Carlo rollout 较准确但每次采样都需数十步模拟,计算代价高昂。StitchVM 通过一次性的模型缝合,将像素空间奖励稳健地迁移至潜在空间,既不引入额外展开开销,又能提供远比 Tweedie 精准的梯度方向,因而在保持奖励对齐质量的同时实现大幅加速与内存优化。相对于最近提出的直接训练潜在空间价值模型,StitchVM 避免了对大规模噪声‑奖励配对数据的依赖,直接复用现有高质量奖励模型的知识,工程落地更友好。该范式为扩散对齐提供了一条“一次构建、重复利用”的工程路径,尤其适合需要频繁调用价值函数的迭代式微调场景。

行业影响

落地场景

StitchVM 为需要任务特异性奖励对齐的扩散/流生成模型提供了高效的值函数估计器,可直接嵌入文本生成图像、可控图像编辑、商业视觉素材生成等产品线。例如:

  • 电商商品图生成:需要同时满足高美学评分与提示一致性,StitchVM 可加速采样阶段的奖励引导,使实时批量生成可行。
  • 社交媒体内容工具:滤镜、广告素材生成常需风格化或脸部保真度约束,StitchVM 允许在推理时快速施加多种奖励,无需为每种奖励重复训练近似器。

商业价值

核心增益来自计算资源降本与体验提速

  • 推理时方法如 DPS 加速 3.2 倍、GPU 内存减半;训练时方法如 DiffusionNFT 加速 2.3 倍,显著降低部署成本。
  • 值模型一经训练可摊销到多轮采样和多个样本,避免每样本昂贵的蒙特卡洛展开或偏差大的 Tweedie 估计,提升生成质量与一致性,直接改善用户体验。
  • 仅需 10 GPU-hour 的轻量拼接微调,可复用企业已有的预训练奖励模型(如 CLIP ViT-L)与扩散骨干,保护投资。

与现有产品/工作流的接口

StitchVM 可作为即插即用模块嵌入主流扩散管线:

  1. 推理时对齐:替换原有值函数近似(如 Tweedie、蒙特卡洛),直接用于 DPSFK steering 等算法。只需加载拼接后的值模型,推理过程中在噪声潜在空间调用,无需改动生成器结构。
  2. 训练时对齐:在 DiffusionNFTFlow-GRPO 等强化微调框架中充当值网络,提供稳定、低偏差的梯度信号,支持离策略训练降低采样负担。
  • 集成方式:对现有扩散模型(如 SD 3.5)的中间层输出增加拼接层,与预训练奖励模型的视觉编码器桥接,微调后导出为独立 checkpoint,兼容 Hugging Face Diffusers 等生态。

局限

  • **StitchVM 严重依赖可用的预训练像素空间奖励模型**:该方法通过拼接现有奖励模型与扩散骨干构造价值函数,若目标任务没有成熟的像素空间奖励模型(如定制化美感指标或新设计约束),则需从头训练,失去轻量化优势。此外,预训练奖励模型本身的量级和领域(如 CLIP ViT-L 偏向语义对齐)可能限制 StitchVM 在低层次纹理保真等奖励上的表现。
  • **拼接界面的选择依赖人工经验且缺乏理论指导**:论文附录提到对拼接层位置的搜索(见 E.3),但当前方案仍为启发式,并未给出一套系统自动化的拼接层选择策略。对于不同扩散骨干(UNet / DiT)和奖励模型架构,最优拼接位置可能变化,导致工程部署时需要额外调参,限制了即插即用的通用性。
  • **在高噪声隐状态下价值估计的精度未充分验证**:虽然 StitchVM 在 DPS、DiffusionNFT 等下游任务上带来加速,但实验主要在中低噪声水平采样步数上评估。对于扩散前期高噪声隐状态的 reward tilting,拼接模型可能因奖励头未曾在极端噪声图像上训练而出现分布外退化,这一潜在偏差在直接使用 Tweedie 或 MC 近似时相对透明,但 StitchVM 的偏差模式需要更系统的标定研究。
论文Hyojun Go2026-05-19原文

相关内容