"""Cache methods on the LingBot-World-V2 DiT (``WanModelFast``): the same policies as ``cachelib/methods.py`` (Self-Forcing / Causal-Forcing) -- see there for the rationale of every deviation from the upstream implementations. Block-level plumbing lives in ``selective.py``. Two levels of method: * block level (TeaCache / TaylorSeer / MotionCache / none / calibrate): a step is the 30-block loop; the preamble (patch embedding, time / text / camera embeddings) and the tail (head, unpatchify) always run; * output level (``reuse``): on an ``R`` step the DiT is not called at all and the previous computed step's **velocity** (flow prediction) is returned, so ``x0 = x_t - sigma_t * v_prev`` and the renoise proceeds as usual. This is the user's definition of the naive-cache baseline on this base model (the other three bases reuse the block-stack residual instead). """ import math import os from dataclasses import dataclass from typing import Callable import numpy as np import torch from .selective import block_forward, block_mod, block_cam, run_blocks_selective # --------------------------------------------------------------------------- # # shared helpers # --------------------------------------------------------------------------- # def run_blocks_full(ctx, features=None): """The stock block loop; with ``features`` (list of dicts, one per block) the op-for-op re-implementation that also captures per-block outputs.""" model, x = ctx.model, ctx.x if features is None: kwargs = dict(ctx.kwargs) for i, block in enumerate(model.blocks): kwargs.update(kv_cache=ctx.kv_cache[i], crossattn_cache=ctx.crossattn_cache[i], current_start=ctx.current_start) x = block(x, **kwargs) return x for i, block in enumerate(model.blocks): x = block_forward(block, x, ctx.e0, ctx.kwargs, ctx.kv_cache[i], ctx.crossattn_cache[i], ctx.current_start, features=features[i]) return x def modulated_input(ctx): """The tensor block 0 hands to its self-attention: norm1(x) * (1+scale) + shift.""" b0 = ctx.model.blocks[0] e = block_mod(b0, ctx.e0) return b0.norm1(ctx.x).float() * (1 + e[1]) + e[0] def rel_l1(cur, prev): diff = (cur - prev).abs().float().mean() base = prev.abs().float().mean() + 1e-8 return (diff / base).item() def rel_l1_per_token(cur, prev): """[B, L, C] -> [L] relative L1 distance, one value per token.""" diff = (cur - prev).abs().float().mean(dim=(0, -1)) base = prev.abs().float().mean(dim=(0, -1)) + 1e-8 return diff / base @dataclass class StepCtx: model: object run_full: Callable # () -> the stock forward's output (output-level methods) x: torch.Tensor # [B, L, C] after patch embedding (block-level methods) e0: torch.Tensor # [B, L, 6, C] kwargs: dict # block kwargs (e, seq_lens, grid_sizes, freqs, context, ...) kv_cache: list crossattn_cache: list current_start: int grid_sizes: torch.Tensor block_idx: int step_idx: int forced_full: bool last_cacheable_step: int class CacheMethod: """``forward`` returns (x_after_blocks, compute_fraction) for block-level methods, (output_list, compute_fraction) for output-level ones.""" name = "base" level = "blocks" def __init__(self, coefficients=None, indicator="modulated_input", **kw): self.coefficients = list(coefficients) if coefficients else None self.indicator_kind = indicator self.extra = kw def rescale(self, value): if self.coefficients is None: return value return float(np.poly1d(self.coefficients)(value)) def indicator(self, ctx): if self.indicator_kind == "e0": return ctx.e0 return modulated_input(ctx) def reset_video(self): pass def begin_chunk(self, block_idx): pass def forward(self, ctx): raise NotImplementedError def config(self): return {"name": self.name, "indicator": self.indicator_kind} class NoCache(CacheMethod): name = "none" def forward(self, ctx): return run_blocks_full(ctx), 1.0 # --------------------------------------------------------------------------- # # 1. TeaCache -- one decision for the whole forward, residual copy # --------------------------------------------------------------------------- # class TeaCache(CacheMethod): name = "teacache" def __init__(self, thresh=0.0, **kw): super().__init__(**kw) self.thresh = float(thresh) self.reset_video() def reset_video(self): self.acc = 0.0 self.prev_ind = None self.prev_residual = None def forward(self, ctx): ind = self.indicator(ctx) if ctx.forced_full or self.prev_ind is None or self.prev_residual is None: should_calc = True self.acc = 0.0 else: self.acc += self.rescale(rel_l1(ind, self.prev_ind)) should_calc = self.acc >= self.thresh if should_calc: self.acc = 0.0 self.prev_ind = ind if should_calc: ori = ctx.x x = run_blocks_full(ctx) self.prev_residual = x - ori return x, 1.0 return ctx.x + self.prev_residual, 0.0 def config(self): return dict(super().config(), thresh=self.thresh) # --------------------------------------------------------------------------- # # 2. TaylorSeer -- forecast skipped features instead of copying them # --------------------------------------------------------------------------- # def taylor_formula(derivatives, distance): out = 0 for i in range(len(derivatives)): out = out + (1.0 / math.factorial(i)) * derivatives[i] * (distance ** i) return out def _taylor_block_add_eager(x, sa, ca, ffn, e2, e5, scale, shift): x = x + sa * e2 if scale is not None: x = (1.0 + scale) * x + shift x = x + ca x = x + ffn * e5 return x def _taylor_block_o1(x, sa0, sa1, ca0, ca1, f0, f1, e2, e5, scale, shift, d): return _taylor_block_add_eager(x, sa0 + d * sa1, ca0 + d * ca1, f0 + d * f1, e2, e5, scale, shift) def _taylor_block_o0(x, sa0, ca0, f0, e2, e5, scale, shift): return _taylor_block_add_eager(x, sa0, ca0, f0, e2, e5, scale, shift) if os.environ.get("LBCACHE_NO_COMPILE", "0") != "1": _taylor_block_o1 = torch.compile(_taylor_block_o1, dynamic=False) _taylor_block_o0 = torch.compile(_taylor_block_o0, dynamic=False) class TaylorSeer(CacheMethod): name = "taylorseer" def __init__(self, interval=1.0, max_order=1, **kw): super().__init__(**kw) self.interval = float(interval) self.max_order = int(max_order) self.reset_video() def reset_video(self): self._pattern_acc = 0.0 self.begin_chunk(-1) def begin_chunk(self, block_idx): self.cache = {} self.cam = {} self.activated = [] lo = int(math.floor(self.interval)) hi = int(math.ceil(self.interval)) if lo == hi: self.chunk_interval = max(1, lo) else: frac = self.interval - lo self._pattern_acc += frac if self._pattern_acc >= 1.0 - 1e-9: self._pattern_acc -= 1.0 self.chunk_interval = max(1, hi) else: self.chunk_interval = max(1, lo) self._since_full = 0 def _should_calc(self, ctx): if ctx.forced_full or not self.activated: return True return (self._since_full + 1) >= self.chunk_interval def forward(self, ctx): if self._should_calc(ctx): x = self._forward_record(ctx) self.activated.append(ctx.step_idx) self._since_full = 0 return x, 1.0 self._since_full += 1 distance = ctx.step_idx - self.activated[-1] return self._forward_taylor(ctx, distance), 0.0 def _forward_record(self, ctx): dt = ctx.step_idx - self.activated[-1] if self.activated else 1 feats = [{} for _ in ctx.model.blocks] x = run_blocks_full(ctx, features=feats) for i, f in enumerate(feats): slot = self.cache.setdefault(i, {}) for key in ("sa", "ca", "ffn"): self._update_derivatives(slot, key, f[key], dt) return x def _update_derivatives(self, slot, key, feature, dt): prev = slot.get(key) new = {0: feature} if prev is not None and dt > 0: for order in range(self.max_order): if order in prev: new[order + 1] = (new[order] - prev[order]) / dt else: break slot[key] = new def _forward_taylor(self, ctx, distance): x = ctx.x d = float(distance) plucker = (ctx.kwargs["dit_cond_dict"] or {}).get("c2ws_plucker_emb") for i, block in enumerate(ctx.model.blocks): e = block_mod(block, ctx.e0) if i not in self.cam: self.cam[i] = block_cam(block, plucker) scale, shift = self.cam[i] slot = self.cache[i] sa, ca, ffn = slot["sa"], slot["ca"], slot["ffn"] with torch.amp.autocast('cuda', dtype=torch.float32): if 1 in sa and 1 in ca and 1 in ffn: x = _taylor_block_o1(x, sa[0], sa[1], ca[0], ca[1], ffn[0], ffn[1], e[2], e[5], scale, shift, d) elif self.max_order <= 1: x = _taylor_block_o0(x, sa[0], ca[0], ffn[0], e[2], e[5], scale, shift) else: x = _taylor_block_add_eager(x, taylor_formula(sa, distance), taylor_formula(ca, distance), taylor_formula(ffn, distance), e[2], e[5], scale, shift) return x def config(self): return dict(super().config(), interval=self.interval, max_order=self.max_order) # --------------------------------------------------------------------------- # # 3. MotionCache -- motion-weighted, token-level reuse # --------------------------------------------------------------------------- # class MotionCache(CacheMethod): name = "motioncache" def __init__(self, thresh=0.0, weight_norm="mean", weight_floor=0.3, min_update_ratio=0.0, temporal_consistency=False, temporal_thresh=0.5, **kw): super().__init__(**kw) self.thresh = float(thresh) self.weight_norm = weight_norm self.weight_floor = float(weight_floor) self.min_update_ratio = float(min_update_ratio) self.temporal_consistency = bool(temporal_consistency) self.temporal_thresh = float(temporal_thresh) self.reset_video() def reset_video(self): self.prev_chunk_last_frame = None self.begin_chunk(-1) def begin_chunk(self, block_idx): self.acc = None self.prev_ind = None self.residual = None self.weights = None self.chunk_idx = block_idx def _motion_weights(self, out, grid_sizes): f, h, w = grid_sizes[0].tolist() spatial = h * w o = out.view(out.shape[0], f, spatial, out.shape[-1]) diffs = [] for fi in range(f): cur = o[:, fi] if fi == 0: if self.prev_chunk_last_frame is not None: prev = self.prev_chunk_last_frame else: diffs.append(None) continue else: prev = o[:, fi - 1] d = (cur - prev).abs().mean(dim=(0, 2)) base = prev.abs().mean(dim=(0, 2)) + 1e-8 diffs.append(d / base) if diffs[0] is None: diffs[0] = diffs[1].clone() if f > 1 else torch.ones(spatial, device=out.device, dtype=torch.float32) frame_diff = torch.stack(diffs).float() self.prev_chunk_last_frame = o[:, -1].detach().clone() if self.weight_norm == "max": weights = frame_diff / (frame_diff.max(dim=1, keepdim=True)[0] + 1e-8) elif self.weight_norm == "max_rescale": lo = frame_diff.min(dim=1, keepdim=True)[0] hi = frame_diff.max(dim=1, keepdim=True)[0] weights = self.weight_floor + (1 - self.weight_floor) * (frame_diff - lo) / (hi - lo + 1e-8) else: weights = frame_diff / (frame_diff.mean(dim=1, keepdim=True) + 1e-8) return weights.reshape(-1) def forward(self, ctx): ind = self.indicator(ctx) L = ctx.x.shape[1] if ctx.forced_full or self.residual is None: ori = ctx.x x = run_blocks_full(ctx) self.residual = x - ori self.prev_ind = ind self.acc = torch.zeros(L, device=x.device, dtype=torch.float32) self.weights = self._motion_weights(x, ctx.grid_sizes) return x, 1.0 dist = rel_l1_per_token(ind, self.prev_ind) if self.coefficients is not None: c = self.coefficients dist = sum(c[i] * dist ** (len(c) - 1 - i) for i in range(len(c))) self.acc = self.acc + dist * self.weights self.prev_ind = ind need = self.acc >= self.thresh if self.temporal_consistency: f, h, w = ctx.grid_sizes[0].tolist() spatial = h * w ratio = need.view(f, spatial).float().mean(dim=0) need = (ratio > self.temporal_thresh).unsqueeze(0).expand(f, -1).reshape(-1) sel = torch.nonzero(need, as_tuple=False).flatten() if 0 < sel.numel() < self.min_update_ratio * L: sel = sel[:0] else: self.acc = torch.where(need, torch.zeros_like(self.acc), self.acc) x = ctx.x + self.residual if sel.numel() == 0: return x, 0.0 x_new = run_blocks_selective(ctx.model, ctx.x, sel, ctx.e0, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start) x = x.index_copy(1, sel, x_new.to(x.dtype)) self.residual = self.residual.index_copy(1, sel, (x_new - ctx.x[:, sel]).to(self.residual.dtype)) return x, sel.numel() / L def config(self): return dict(super().config(), thresh=self.thresh, weight_norm=self.weight_norm, min_update_ratio=self.min_update_ratio, temporal_consistency=self.temporal_consistency) # --------------------------------------------------------------------------- # # 4. naive baselines and probes # --------------------------------------------------------------------------- # class DirectReuse(CacheMethod): """Naive cache baseline (velocity reuse): every cacheable step returns the last computed step's flow prediction without calling the DiT (`FRFF`/`FRRF`/`FRRR`).""" name = "reuse" level = "output" def reset_video(self): self.prev_velocity = None def forward(self, ctx): if ctx.forced_full or self.prev_velocity is None: out = ctx.run_full() self.prev_velocity = [u.detach() for u in out] return out, 1.0 return [u.clone() for u in self.prev_velocity], 0.0 def config(self): return dict(super().config(), reuse="velocity") class CalibrationProbe(CacheMethod): """Always full; records (rel_l1 of indicator, rel_l1 of residual) pairs for consecutive steps inside a chunk -- TeaCache's rescale-polynomial fit.""" name = "calibrate" def __init__(self, **kw): super().__init__(**kw) self.samples = [] self.prev_ind = self.prev_res = None def begin_chunk(self, block_idx): self.prev_ind = self.prev_res = None def forward(self, ctx): ind = self.indicator(ctx) ori = ctx.x x = run_blocks_full(ctx) res = x - ori if self.prev_ind is not None: self.samples.append({"x": rel_l1(ind, self.prev_ind), "y": rel_l1(res, self.prev_res), "block": ctx.block_idx, "step": ctx.step_idx}) self.prev_ind, self.prev_res = ind, res return x, 1.0 class ExactCheck(CacheMethod): """Test-only: cacheable steps go through the re-implemented full loop (``mode="reimpl"``) or the selective path with *every* token selected (``mode="selective"``). Both must reproduce the stock forward bit for bit.""" name = "exactcheck" def __init__(self, mode="reimpl", **kw): super().__init__(**kw) self.mode = mode def forward(self, ctx): if self.mode == "reimpl": return run_blocks_full(ctx, features=[{} for _ in ctx.model.blocks]), 1.0 if ctx.forced_full: return run_blocks_full(ctx), 1.0 sel = torch.arange(ctx.x.shape[1], device=ctx.x.device) return run_blocks_selective(ctx.model, ctx.x, sel, ctx.e0, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start), 1.0 REGISTRY = {"none": NoCache, "reuse": DirectReuse, "calibrate": CalibrationProbe, "teacache": TeaCache, "taylorseer": TaylorSeer, "motioncache": MotionCache, "exactcheck": ExactCheck} def build_method(name, **kw): if name not in REGISTRY: raise KeyError(f"unknown cache method {name!r}; have {sorted(REGISTRY)}") return REGISTRY[name](**{k: v for k, v in kw.items() if v is not None})