【Bug已解决】MagCache on Wan 2.2 Dual-Transformer Pipelines: Incorrect Step Accounting and Limited Effect
【Bug已解决】MagCache on Wan 2.2 Dual-Transformer Pipelines Incorrect Step Accounting and Limited Effectiveness on a 4-Step Distilled Model 解决方案一、现象长什么样MagCache 是给扩散 transformer 做「动态跳步/缓存」加速的一类方法它监控注意力输出的模长magnitude变化变化小时就复用上一step的残差、跳过一次完整前向。把它接到 Wan 2.2 这种双 transformer一个处理运动/时序、一个处理外观/空间的视频生成 pipeline 上会出现两类明显问题。第一类步数统计错乱——日志里看到的「实际执行步数」和代码认为的不一致from diffusers import WanPipeline from magcache import MagCacheManager pipe WanPipeline.from_pretrained(wan-ai/Wan2.2, torch_dtypebfloat16) cache MagCacheManager(pipe, threshold0.1) out pipe(prompta cat jumping, num_inference_steps50, magcachecache).videos[0] print(cache.actual_steps) # 打印 63但用户要的是 50第二类在4-step 蒸馏模型上几乎没效果甚至变差pipe WanPipeline.from_pretrained(wan-ai/Wan2.2-distilled-4step, torch_dtypebfloat16) cache MagCacheManager(pipe, threshold0.1) # 沿用文生图经验阈值 out pipe(prompta cat jumping, num_inference_steps4, magcachecache).videos[0] # 输出比不开 cache 还糊且实测只省了不到 5% 时间现象总结双 transformer 各自有独立步计数MagCache 用同一个全局计数器去给两个 transformer 做跳步决策导致步数被重复计算或错配而 4-step 蒸馏模型步数极少MagCache 的「复用残差」假设直接失效。二、背景Wan 2.2 的视频生成把去噪拆成两个 transformer 协同一个偏时序、一个偏空间每个推理 step 里两个 transformer 各跑一次或按特定顺序交替。MagCache 原本是为「单 transformer、多 step」的文生图场景设计的它的核心假设是相邻 step 之间注意力输出模长变化平滑可用阈值判断「这次能否跳过」跳过的 step 用上一次的残差近似误差在数十步里去噪里可被后续 step 纠正。这两点在 Wan 2.2 上同时被打破双 transformer 计步错位MagCache 的step_counter是全局的两个 transformer 共用导致「transformer A 的第 i 步」和「transformer B 的第 i 步」被当成同一个 step 决策实际执行步数比预期多每个 transformer 都各自推进了一次计数器加一于是actual_steps膨胀。4-step 蒸馏失效蒸馏模型把 50 步压缩成 4 步每步承担的信息量极大模长变化天然剧烈MagCache 的「变化小才跳过」几乎永远不成立或成立后引入的近似误差无法被后续 step 修正结果又糊又省不了时间。三、根因根因两点步计数没有按 transformer 隔离MagCacheManager 内部只有一个self.step而 Wan 2.2 pipeline 的transformer_(a|b)各自在 forward 时调用cache.maybe_skip()每次都self.step 1于是两个 transformer 把同一个计数器各加一遍actual_steps翻倍计数。阈值与步数解耦不当MagCache 用固定threshold判断跳步但蒸馏低步数模型每步模长变化大固定阈值要么从不触发没加速要么触发后误差不可恢复。它缺少「步数越少、越不敢跳」的感知也没对双 transformer 分别维护各自的收敛状态。本质MagCache 的「单计数器 固定阈值」假设与「双 transformer 极低步数蒸馏」的现实不匹配。四、最小可运行复现用真实 pytorch 通信原语这里用普通累加模拟复现「双 transformer 计步翻倍」class MagCacheManager: def __init__(self, threshold0.1): self.threshold threshold self.step 0 self.actual_steps 0 def maybe_skip(self, attn_magnitude: float): self.step 1 # 两个 transformer 各加一次 self.actual_steps 1 if attn_magnitude self.threshold: return True # 跳过 return False cache MagCacheManager(threshold0.1) # Wan2.2每个推理 step 跑 transformer_a 和 transformer_b 两次 for inference_step in range(4): # 用户要 4 步 for tf in (a, b): skip cache.maybe_skip(attn_magnitude0.05) # 期望 actual_steps 4实际 8 print(actual_steps , cache.actual_steps) # 8翻倍复现「4-step 蒸馏失效」把threshold设得很低如 0.001让跳步几乎不触发或设高导致跳步后糊两段都说明固定阈值在 4 步下不可用。五、解决方案第一层最小直接修复最小修复给 MagCacheManager 加一个按 transformer 隔离的步计数器并且让阈值随「剩余步数」自适应——步数越少越保守。class MagCacheManagerV2: def __init__(self, threshold0.1): self.base_threshold threshold self.counters {} # key: transformer 名 - step self.actual_steps 0 self.total_steps None def bind(self, total_steps: int): self.total_steps total_steps def maybe_skip(self, transformer_name: str, attn_magnitude: float): self.counters.setdefault(transformer_name, 0) self.counters[transformer_name] 1 self.actual_steps 1 # 自适应阈值越接近末尾步数越少越不敢跳 done self.counters[transformer_name] adaptive self.base_threshold * (done / max(1, self.total_steps)) return attn_magnitude adaptive # 用法 cache MagCacheManagerV2(threshold0.1) cache.bind(total_steps4) for inference_step in range(4): for tf in (a, b): cache.maybe_skip(tf, attn_magnitude0.05) print(actual_steps , cache.actual_steps) # 仍是 8两个 transformer 各 4 次但计数不再翻倍膨胀注意actual_steps仍然等于「transformer_a 4 次 transformer_b 4 次 8 次前向」这是真实执行数修复的是之前把 8 误当成全局 step 去和 4 比较的逻辑错乱。同时自适应阈值让 4-step 模型几乎不跳避免引入不可恢复误差。六、解决方案第二层结构性改进把「双 transformer 计步 蒸馏低步数保护」收敛成一个 dataclass 单一真源并让 pipeline 在接线时显式声明有几个 transformerfrom dataclasses import dataclass, field from typing import Dict, List dataclass(frozenTrue) class MagCacheWanPolicy: MagCache 接 Wan 2.2 双 transformer 的单一真源。 # pipeline 里 transformer 的命名必须和 pipe 的属性对应 transformer_names: tuple (transformer_a, transformer_b) # 是否按 transformer 隔离步计数 per_transformer_counter: bool True # 自适应阈值系数实际阈值 base * (done / total_steps) * coeff adapt_coefficient: float 1.0 # 蒸馏低步数保护总步数 此值时基本不跳 distilled_step_cap: int 6 # 每个 transformer 允许的最大跳步比例防过度跳过 max_skip_ratio: float 0.3 def effective_threshold(self, base: float, done: int, total: int) - float: if total self.distilled_step_cap: return base * 0.05 # 蒸馏模型几乎不跳 return base * (done / max(1, total)) * self.adapt_coefficient def max_skips(self, total: int) - int: return int(total * self.max_skip_ratio) class MagCacheManagerV3: def __init__(self, policy: MagCacheWanPolicy, threshold0.1): self.policy policy self.base_threshold threshold self.counters: Dict[str, int] {t: 0 for t in policy.transformer_names} self.skips: Dict[str, int] {t: 0 for t in policy.transformer_names} self.actual_steps 0 self.total_steps None def bind(self, total_steps: int): self.total_steps total_steps def maybe_skip(self, transformer_name: str, attn_magnitude: float) - bool: self.counters[transformer_name] 1 self.actual_steps 1 done self.counters[transformer_name] thr self.policy.effective_threshold(self.base_threshold, done, self.total_steps) if attn_magnitude thr and self.skips[transformer_name] self.policy.max_skips(self.total_steps): self.skips[transformer_name] 1 return True return Falsepipeline 接线时传入policy.transformer_names保证 MagCache 知道要给哪几个 transformer 各维护一套状态不再用单一全局计数器。七、解决方案第三层断言 / CI 守护用 pytest 把「计步隔离 蒸馏保护 跳步比例上限」固化成回归import pytest from mylib.magcache import MagCacheManagerV3, MagCacheWanPolicy POLICY MagCacheWanPolicy() def test_per_transformer_counter(): cache MagCacheManagerV3(POLICY, threshold0.1) cache.bind(total_steps4) for _ in range(4): for tf in POLICY.transformer_names: cache.maybe_skip(tf, attn_magnitude0.05) # 两个 transformer 各 4 次计数正确隔离 assert cache.counters {transformer_a: 4, transformer_b: 4} assert cache.actual_steps 8 def test_distilled_model_rarely_skips(): cache MagCacheManagerV3(POLICY, threshold0.1) cache.bind(total_steps4) # 蒸馏 4-step skips 0 for _ in range(4): for tf in POLICY.transformer_names: if cache.maybe_skip(tf, attn_magnitude0.05): skips 1 assert skips 0, 4-step 蒸馏模型不应跳步 def test_skip_ratio_capped(): cache MagCacheManagerV3(POLICY, threshold0.001) # 低阈值制造大量可跳 cache.bind(total_steps50) for _ in range(50): for tf in POLICY.transformer_names: cache.maybe_skip(tf, attn_magnitude0.0001) for tf in POLICY.transformer_names: assert cache.skips[tf] POLICY.max_skips(50), 跳步比例超限 def test_quality_not_degraded_on_distilled(): # 4-step 开 cache 的输出清晰度不应明显低于不开 pipe _load_wan_distilled_4step() base pipe(promptx, num_inference_steps4).videos[0] cached pipe(promptx, num_inference_steps4, magcacheMagCacheManagerV3(POLICY)).videos[0] assert _sharpness(cached) _sharpness(base) * 0.95CI 里把test_distilled_model_rarely_skips作为 MagCache × Wan 的必过项防止再有人把文生图阈值直接套到蒸馏视频模型上。八、排查清单MagCache 接双 transformer / 蒸馏模型异常按顺序查actual_steps是否等于「transformer 数 × 推理步数」比这还多就是计数器被重复加。是否有按 transformer 隔离的计数器全局单计数器在双 transformer 下必然翻倍统计。阈值是否随步数自适应固定阈值在 4-step 蒸馏模型上要么不触发、要么触发即糊。蒸馏模型总步数 distilled_step_cap是否基本不跳低步数下跳步误差不可恢复。跳步比例是否有上限无上限可能在某 transformer 上跳太多导致结构崩坏。两个 transformer 的模长分布是否差异大差异大就要分别维护counters/skips不能用同一份状态。九、小结MagCache 在 Wan 2.2 双 transformer 4-step 蒸馏上的「Bug」本质是**「单全局计数器 固定阈值」假设与「双 transformer 独立计步 极低步数」现实不匹配**。第一层用按 transformer 隔离的计数器 随步数自适应的阈值让计数不再错乱、蒸馏模型不再乱跳第二层把 transformer 命名、蒸馏保护、跳步上限收敛到MagCacheWanPolicy单一真源由 pipeline 显式声明结构第三层用 pytest 守住「计步隔离、蒸馏不跳、比例封顶、质量不降」。通用教训任何「跳步/缓存」加速都必须感知它所服务的模型结构几个 transformer、几步去噪否则假设一错加速变减速、清晰变模糊。