跳转至

Rethinking Token Reduction for Diffusion Models via Output-Similarity-Awareness

会议: ECCV 2026
论文: ECCV 2026 Poster
领域: 模型压缩
关键词: 扩散模型, Token 缩减, 输出相似度感知, 匹配重用, 频次惩罚

一句话总结

针对扩散模型中传统基于输入相似度的 Token 缩减导致恢复误差偏高的问题,提出输出相似度感知缩减方法 DiTo,利用相邻时间步输出的高时序一致性,在匹配步利用前序输出相似度建立索引并在多个缩减步复用,结合 PMR 自适应调度与频次感知惩罚,在 Flux 和 SD3 上取得 1.6–3.9 dB 的 PSNR 提升及高达 18.6% 的延迟降低。

研究背景与动机

扩散 Transformer(Diffusion Transformers, DiTs)凭借其基于 Transformer 堆叠架构的高可扩展性,在高质量图像生成领域取得了突破性进展。然而,Transformer 中自注意力机制的计算复杂度与序列 Token 长度呈二次方关系(\(O(N^2)\)),在高分辨率图像生成任务中带来了巨大的计算开销与推理时延。尽管 FlashAttention 等算子级优化在底层缓解了内存访问瓶颈,但并未改变核心的二次方算法复杂度。为此,直接修剪或合并冗余 Token 的 Token 缩减(Token Reduction, TR)技术成为加速扩散模型极具吸引力的方向。

然而,现有的扩散模型 Token 缩减方法(如 ToMeSD、ToFu、SiTo、ToMA 等)直接继承了传统判别式 Vision Transformer(ViT)的设计范式,完全依赖当前层的输入 Token 相似度来进行匹配。在判别式分类任务中,模型仅需移除冗余信息;但在扩散生成模型中,Token 对应于画布上的精确空间位置,在缩减计算后必须通过恢复阶段(Recovery Stage)将减少的 Token 通过复制匹配目标 Token 的特征进行空间重建。因此,生成模型中 Token 缩减的首要目标是最小化恢复输出与密集原始输出之间的恢复误差(Recovery Error)。由于恢复过程本质上是用匹配 Token 的输出去近似被裁剪 Token 的输出,因此决定重建质量的关键在于输出 Token 之间的相似度,而非输入 Token。实验证明,现有完全基于输入的匹配方式与最小化恢复误差的根本目标存在错位,造成了严重的结构失真与纹理模糊。

要根据输出相似度进行匹配,最棘手的问题在于当前步的稠密输出在计算完成之前是不可知的。本文的切入角度是观察到扩散模型去噪轨迹中隐藏的时间连续性——虽然激活特征在相邻步高度相关已被熟知,但输出 Token 间的配对相似度在相邻步间同样具有强一致性。核心 idea:提出面向扩散模型的输出相似度感知 Token 缩减框架 DiTo,将时间步划分为匹配步与缩减步,利用前一时间步的输出相似度作为当前步金标准输出匹配的精准代理,并将匹配索引在后续多个缩减步中高效复用,同时引入 PMR 指导的离线调度与频次感知惩罚以消除局部伪影。

方法详解

整体框架

DiTo 将原本耦合在每个扩散时间步内的匹配、缩减与恢复操作进行解耦,通过在时间维度上交替调度「匹配步(Matching Step)」与「缩减步(Reduction Step)」来打破效率瓶颈。在匹配步中,模型执行完整的全量计算,并基于该步产生的输出特征计算 Token 间的空间相似度,构建高精度的源-目标映射关系;在随后的多个缩减步中,模型无需重复进行高开销的 Token 匹配计算,直接复用已保存的紧凑索引元数据进行 Token 缩减、高效注意力计算与空间恢复重建。

为了使该流程兼顾高保真度与低计算开销,DiTo 结合了两个关键机制:一是离线基于 Top-\(k\) 对匹配率(Pair Match Ratio, PMR)评估时序退化规律,自适应规划每个时间步允许的最大复用间隔;二是在匹配阶段引入 Token 被选频次的历史衰减惩罚,防止特定空间区域因反复缩减累积过大局部误差而产生块状伪影。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["输入噪声潜变量 x_t"] --> B["输出相似度感知 Token 匹配<br/>利用前序输出计算代理相似度"]
    B --> C["频次感知惩罚校正<br/>扣减累计选择频次得分以防偏置"]
    C --> D["PMR 引导的动态间隔调度<br/>判定本步为匹配步还是缩减步"]
    D -->|匹配步: 刷新对应关系| E["执行稠密计算并暂存匹配索引"]
    D -->|缩减步: 复用先验索引| F["Token 缩减与稀疏注意力计算"]
    F --> G["Token 恢复重建<br/>依据索引复制输出特征回填原位置"]
    E --> H["当前时间步去噪潜变量 x_{t-1}"]
    G --> H

关键设计

1. 输出相似度感知 Token 缩减:以前序输出作为当前输出匹配的金标准代理

传统方法在当前时间步 \(t\) 直接利用该步的注意力输入特征计算相似度矩阵 \(\mathbf{A} \in \mathbb{R}^{|D| \times |S|}\),但在恢复阶段将未保留的源 Token \(s \in S\) 通过其匹配的目标 Token \(d^\star(s)\) 的输出进行复制回填时,恢复误差 \(\mathcal{L}_{\text{out}} = \|Y - \tilde{Y}\|_F^2\) 直接取决于输出特征空间中 \(Y_s\) 与 \(Y_{d^\star(s)}\) 的接近程度。既然当前稠密输出 \(Y^{(t)}\) 无法预先获取,DiTo 利用去噪轨迹中相邻时间步特征高度一致的时序先验,将前一时间步的输出特征相似度作为代用指标。作者通过统计表明,使用前序时间步输出(\(\Delta \ge 1\))计算出的匹配对应关系与当前真实输出金标准匹配结果的重合度,大幅且稳定地高于直接用当前时间步输入(\(\Delta = 0\))计算得到的结果,从而在根本上对齐了最小化恢复误差的数学目标。

2. PMR 引导的自适应时间步调度:兼顾跨步复用收益与时序漂移界限

虽然复用单次匹配结果能够避免频繁的 Token 配对开销,但复用时间间隔 \(\Delta\) 越长,时序漂移引起的匹配精度衰减越明显。为了科学量化匹配有效性,DiTo 提出了 Top-\(k\) 对匹配率(Pair Match Ratio, PMR)指标,其定义为源 Token 在时间步 \(t-\Delta\) 下预测的前 \(k\) 个候选目标集合中包含时间步 \(t\) 真实单目标金标准 Token 的比例: $\(\mathrm{PMR}_{\text{Top-}k}(t, b, \Delta) = \frac{1}{|S|} \sum_{s \in S} \mathbb{I}\left[ \left| T^{(1)}_{t, b}(s) \cap T^{(k)}_{t-\Delta, b}(s) \right| > 0 \right]\)$ 对所有 Transformer 块求均值后得到平均 \(\overline{\mathrm{PMR}}(t, \Delta)\)。在离线标定阶段,给定质量阈值 \(\tau\),找出每个时间步 \(t\) 满足 \(\overline{\mathrm{PMR}}(t, \Delta) \ge \tau\) 的最大允许复用步长 \(\Delta_t^{\max}\)。在推理时,采用单步前瞻策略(One-step Look-ahead):实时检查下一步预测间隔 \(\Delta_{t+1} = (t+1) - m\) 是否超过 \(\Delta_{t+1}^{\max}\)(其中 \(m\) 为最近匹配步编号),若超限则将当前步设为新的匹配步以更新缓存映射,反之设为快速缩减步。连续多个匹配步时仅保留末尾步,彻底消除了冗余匹配计算。

3. 频次感知 Token 匹配:消除局部误差累积与块状伪影

当同一匹配关系跨越连续多个时间步复用时,部分具有微弱语义差异但相似度偏高的 Token 会被连续选为缩减对象,导致其空间对应区域的近似误差按复用步数线性叠加,最终在生成图像的特定空间坐标处引发肉眼可见的马赛克与块状伪影(Blocking Artifacts)。为解决这一偏置,DiTo 维护一个全局历史选择频次向量 \(\mathbf{C} \in \mathbb{R}^N\),记录每个 Token 被作为源 Token 裁剪的累积次数。在匹配步骤中,通过归一化候选相似度动态范围 \(\Delta s = \max(\hat{\mathbf{s}}) - \min(\hat{\mathbf{s}})\) 和缩减率 \(r\),构造尺度无关的惩罚项并对候选相似度打分进行自适应扣减: $\(\mathbf{p} = \lambda \cdot r \cdot \Delta s \cdot \mathbf{C}[\mathcal{S}], \qquad \hat{\mathbf{s}}^{\mathrm{pen}} = \hat{\mathbf{s}} - \mathbf{p}\)$ 其中 \(\lambda\) 为惩罚强度超参数。该设计迫使模型在多个周期中动态轮换被缩减的局部 Token,有效打散了误差空间分布,在保持计算加速的同时彻底平抑了局部几何畸变。

损失函数 / 训练策略

DiTo 属于完全免训练(Training-Free)与即插即用的推理加速方案,不需要对原扩散模型进行任何权重微调或梯度反向传播。匹配调度规则基于固定扩散步数下离线校准得到的 \(\Delta_t^{\max}\) 清单执行。在运行时,模型仅在匹配步计算并暂存整数索引数组(空间复杂度为 \(O(N)\),显存占用远低于需要缓存高维全张量特征的缓存类方案如 ToCa),在缩减步无需重复构建相似度图,开销近乎为零。

实验关键数据

主实验

评估采用 1024×1024 分辨率文本生成图像基准,在 FLUX.1-dev(35 步)与 Stable Diffusion 3(SD3, 50 步)上测试 ImageNet-1k 类别提示词(共 3,000 张图像),以原生密集采样输出为黄金基准计算重建保真度,在单张 NVIDIA RTX 6000 Ada GPU 上测定真实时延。

模型 缩减率 方法 FID↓ PSNR↑ SSIM↑ LPIPS↓ CLIP↑ 延迟 (s)↓ 相对加速比 (Δ)↓
Flux Baseline 密集全量 31.62 – – – 23.71 26.64 0%
Flux 0.25 DiTo (本文) 31.95 25.21 0.8552 0.2385 23.85 24.25 -9.0%
Flux 0.25 ToMeSD 33.76 21.59 0.7817 0.3395 23.93 24.53 -7.9%
Flux 0.25 ToFu 33.75 21.70 0.7851 0.3354 23.94 24.47 -8.1%
Flux 0.25 SiTo 32.30 22.39 0.7965 0.3462 23.93 25.25 -5.2%
Flux 0.25 ToMA 32.03 22.97 0.8112 0.2716 23.78 24.28 -8.9%
Flux 0.50 DiTo (本文) 33.47 22.04 0.7757 0.3329 24.01 21.69 -18.6%
Flux 0.50 ToMeSD 90.00 19.73 0.6172 0.5384 23.04 22.00 -17.4%
Flux 0.50 ToFu 83.79 19.81 0.6328 0.5302 23.05 22.02 -17.3%
Flux 0.50 SiTo 61.05 20.17 0.6503 0.5245 23.45 22.54 -15.4%
Flux 0.50 ToMA 43.71 20.49 0.6936 0.4386 23.90 21.62 -18.8%
SD3 Baseline 密集全量 27.67 – – – 24.67 9.32 0%
SD3 0.25 DiTo (本文) 27.83 26.83 0.9160 0.1492 24.77 8.63 -7.4%
SD3 0.25 ToMeSD 28.42 23.32 0.8697 0.2045 24.74 8.82 -5.3%
SD3 0.25 ToFu 27.99 23.83 0.8835 0.1922 24.72 8.87 -4.9%
SD3 0.25 SiTo 27.92 23.01 0.8682 0.2211 24.79 9.30 -0.2%
SD3 0.25 ToMA 28.47 22.93 0.8624 0.2029 24.63 9.08 -2.5%
SD3 0.50 DiTo (本文) 28.89 23.16 0.8510 0.2577 24.87 7.79 -16.4%
SD3 0.50 ToMeSD 31.45 21.22 0.8090 0.3139 24.80 8.00 -14.2%
SD3 0.50 ToFu 31.89 21.21 0.8112 0.3187 24.79 8.05 -13.6%
SD3 0.50 SiTo 33.13 21.23 0.8046 0.3170 24.82 8.42 -9.7%
SD3 0.50 ToMA 33.37 20.91 0.8049 0.3053 24.62 8.17 -12.3%

消融实验

消融实验对比了恢复误差与关键模块(频次感知惩罚、PMR 调度阈值)的影响:

评估配置 / 变体 核心观察指标与数值 机制分析与结论
基于输出匹配 vs 基于输入匹配 恢复误差 \(\mathcal{L}_{\text{out}}\):500 个样本在散点图中全部位于 \(y=x\) 下方 证明输出空间特征相似度是决定空间恢复误差的决定性因变量,输入相似度存在显著对齐偏差
DiTo w/o 频次惩罚 (\(\lambda=0\)) 局部 Token 连续选中计数峰值高达 100 次,产生明显肉眼可见的块状伪影 长期复用导致误差高度局部聚集,破坏了高频细节与空间平滑度
DiTo w/ 频次惩罚 (\(\lambda>0\)) 局部选中计数峰值从 100 骤降至 40 以下,视觉块状伪影被完全抑制 引入尺度无关频次衰减后,被剪枝 Token 在全图平摊分布,消除了局部畸变
高缩减率鲁棒性 (Flux 50% 缩减) FID 为 33.47(对比基线 ToMeSD 90.00, ToFu 83.79, SiTo 61.05, ToMA 43.71) 传统输入匹配在激进裁剪下完全瓦解,而输出相似度保持了高阶语义结构不崩塌

关键发现

  • 恢复误差的决定性因子在于输出而非输入:直接分析匹配误差散点图可知,100% 的评估样本采用输出相似度代理后恢复误差更低,这解释了为什么输入驱动的方法在大缩减率下会出现严重伪影。
  • 高剪切率下的极端鲁棒性:当剪切率达到 50% 时,传统方法的 FID 出现断崖式恶化(如 ToMeSD 从 33.76 暴跌至 90.00),而 DiTo 依然保持在 33.47,PSNR 仍有 22.04 dB,展现出极其稳固的 Pareto 前沿。
  • 极低的元数据存储开销:与整层缓存激活特征的特征缓存方案(如 ToCa)相比,DiTo 仅跨步传递整数索引映射,推理内存开销可忽略不计,兼具极高的工程部署友好度。

亮点与洞察

  • 切中生成任务本质的输出中心视角:指出了判别式 ViT 缩减与生成式扩散模型缩减的核心区别在于「恢复阶段」的引入,从而将优化核心从输入相关性转向最小化输出恢复误差。
  • 巧妙的时序一致性代理与频次衰减控制:利用前序时间步输出作为当前输出未决时的天然无损代理,并配以尺度无关的频次扣减惩罚,以近乎零的计算成本实现了高质量的自适应时空采样。
  • 通用性可扩展至多种骨干模型:该方法不局限于特定的注意力拓扑,不仅适用于 MM-DiT(Flux、SD3),也能无缝迁移至传统 U-Net(SD1.5)和标准 DiT,具有极强的通用价值。

局限与展望

  • 离线调度的条件适配性:PMR 调度虽然对 Prompt、随机种子和 CFG 展现出较好的泛化性,但在更换全新的采样步数 \(T\) 或截然不同的调度器(如 Flow Matching 与 DDIM 混合)时,仍需进行一次轻量的离线 PMR 曲线标定。
  • 极高动态场景的潜在边界:若应用于极少数跨步特征突变显著的极端瞬态阶段,前序时间步输出代理的准确率会有所下降,未来可探索动态自适应在线步长调控机制。

相关工作与启发

  • vs ToMeSD / ToFu: ToMeSD 与 ToFu 依赖输入 Token 的相似度计算软匹配或混合裁剪,未考虑恢复阶段输出特征的匹配误差,在大裁剪率下结构严重畸变;DiTo 转向输出相似度导向,在 50% 裁剪率下 FID 提升 56+ 点。
  • vs ToMA: ToMA 采用局部窗口内的并行贪婪匹配以加速计算,依然局限于当前步输入;DiTo 进一步通过跨时间步索引复用降低计算复杂度,在达到同等或更优加速比的同时 PSNR 高出 1.5–2.2 dB。
  • vs ToCa: ToCa 通过高维特征缓存实现跳步加速,但伴随巨大的显存占用;DiTo 仅复用轻量级的整数索引映射,显存增加几乎为零,更利于端侧边缘设备部署。

评分

  • 新颖性: ⭐⭐⭐⭐☆ 敏锐抓住了扩散模型中恢复误差与输出相似度的本质联系,打破了延续自 ViT 的输入匹配思维惯性。
  • 实验充分度: ⭐⭐⭐⭐⭐ 涵盖 Flux 与 SD3 两大主流高分辨率模型,对比全面,评价指标丰富,消融深入。
  • 写作质量: ⭐⭐⭐⭐⭐ 逻辑结构严谨,问题陈述一针见血,Mermaid 图文呼应清晰。
  • 价值: ⭐⭐⭐⭐⭐ 免训练、低显存、高性能,为大模型文生图工业界端侧落地提供了关键的算力压缩支撑。