三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

【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 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_dtype="bfloat16") cache = MagCacheManager(pipe, threshold=0.1) out = pipe(prompt="a cat jumping", num_inference_steps=50, magcache=cache).videos[0] print(cache.actual_steps) # 打印 63,但用户要的是 50

第二类,在4-step 蒸馏模型上几乎没效果甚至变差:

pipe = WanPipeline.from_pretrained("wan-ai/Wan2.2-distilled-4step", torch_dtype="bfloat16") cache = MagCacheManager(pipe, threshold=0.1) # 沿用文生图经验阈值 out = pipe(prompt="a cat jumping", num_inference_steps=4, magcache=cache).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 上同时被打破:

  1. 双 transformer 计步错位:MagCache 的step_counter是全局的,两个 transformer 共用,导致「transformer A 的第 i 步」和「transformer B 的第 i 步」被当成同一个 step 决策,实际执行步数比预期多(每个 transformer 都各自推进了一次计数器加一),于是actual_steps膨胀。
  2. 4-step 蒸馏失效:蒸馏模型把 50 步压缩成 4 步,每步承担的信息量极大,模长变化天然剧烈,MagCache 的「变化小才跳过」几乎永远不成立,或成立后引入的近似误差无法被后续 step 修正,结果又糊又省不了时间。

三、根因

根因两点:

  1. 步计数没有按 transformer 隔离:MagCacheManager 内部只有一个self.step,而 Wan 2.2 pipeline 的transformer_(a|b)各自在 forward 时调用cache.maybe_skip(),每次都self.step += 1,于是两个 transformer 把同一个计数器各加一遍,actual_steps翻倍计数。
  2. 阈值与步数解耦不当:MagCache 用固定threshold判断跳步,但蒸馏低步数模型每步模长变化大,固定阈值要么从不触发(没加速),要么触发后误差不可恢复。它缺少「步数越少、越不敢跳」的感知,也没对双 transformer 分别维护各自的收敛状态。

本质:MagCache 的「单计数器 + 固定阈值」假设与「双 transformer + 极低步数蒸馏」的现实不匹配

四、最小可运行复现

用真实 pytorch 通信原语(这里用普通累加模拟)复现「双 transformer 计步翻倍」:

class MagCacheManager: def __init__(self, threshold=0.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(threshold=0.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_magnitude=0.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, threshold=0.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(threshold=0.1) cache.bind(total_steps=4) for inference_step in range(4): for tf in ("a", "b"): cache.maybe_skip(tf, attn_magnitude=0.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 在接线时显式声明有几个 transformer:

from dataclasses import dataclass, field from typing import Dict, List @dataclass(frozen=True) 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, threshold=0.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 False

pipeline 接线时传入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, threshold=0.1) cache.bind(total_steps=4) for _ in range(4): for tf in POLICY.transformer_names: cache.maybe_skip(tf, attn_magnitude=0.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, threshold=0.1) cache.bind(total_steps=4) # 蒸馏 4-step skips = 0 for _ in range(4): for tf in POLICY.transformer_names: if cache.maybe_skip(tf, attn_magnitude=0.05): skips += 1 assert skips == 0, "4-step 蒸馏模型不应跳步" def test_skip_ratio_capped(): cache = MagCacheManagerV3(POLICY, threshold=0.001) # 低阈值,制造大量可跳 cache.bind(total_steps=50) for _ in range(50): for tf in POLICY.transformer_names: cache.maybe_skip(tf, attn_magnitude=0.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(prompt="x", num_inference_steps=4).videos[0] cached = pipe(prompt="x", num_inference_steps=4, magcache=MagCacheManagerV3(POLICY)).videos[0] assert _sharpness(cached) >= _sharpness(base) * 0.95

CI 里把test_distilled_model_rarely_skips作为 MagCache × Wan 的必过项,防止再有人把文生图阈值直接套到蒸馏视频模型上。

八、排查清单

MagCache 接双 transformer / 蒸馏模型异常按顺序查:

  1. actual_steps是否等于「transformer 数 × 推理步数」?比这还多就是计数器被重复加。
  2. 是否有按 transformer 隔离的计数器?全局单计数器在双 transformer 下必然翻倍统计。
  3. 阈值是否随步数自适应?固定阈值在 4-step 蒸馏模型上要么不触发、要么触发即糊。
  4. 蒸馏模型(总步数 <=distilled_step_cap)是否基本不跳?低步数下跳步误差不可恢复。
  5. 跳步比例是否有上限?无上限可能在某 transformer 上跳太多导致结构崩坏。
  6. 两个 transformer 的模长分布是否差异大?差异大就要分别维护counters/skips,不能用同一份状态。

九、小结

MagCache 在 Wan 2.2 双 transformer + 4-step 蒸馏上的「Bug」本质是**「单全局计数器 + 固定阈值」假设与「双 transformer 独立计步 + 极低步数」现实不匹配**。第一层用按 transformer 隔离的计数器 + 随步数自适应的阈值让计数不再错乱、蒸馏模型不再乱跳;第二层把 transformer 命名、蒸馏保护、跳步上限收敛到MagCacheWanPolicy单一真源,由 pipeline 显式声明结构;第三层用 pytest 守住「计步隔离、蒸馏不跳、比例封顶、质量不降」。通用教训:任何「跳步/缓存」加速都必须感知它所服务的模型结构(几个 transformer、几步去噪),否则假设一错,加速变减速、清晰变模糊

← 返回列表