"""The four cache methods, ported to the block-causal 4-step Self-Forcing / Causal-Forcing DiT. Schedule contract: the controller marks some steps of every chunk as forced-full (``ctx.forced_full``) and the method decides the rest. The default ``FxxF`` (step 0 and step 3 always full, steps 1-2 cacheable) is equivalent to the original implementations' ``ret_steps=1`` / ``cutoff_steps=num_steps-1`` warm-up and cool-down guards. ``FFxx`` gives TaylorSeer two computed points before its first forecast, so the first-order term is actually exercised; ``Fxxx`` lifts the cool-down guard and raises the arithmetic ceiling from 2x to 4x. Indicator note -------------- ``TeaCache4Wan2.1`` derives its indicator from ``e0`` alone (the timestep modulation). That is sound for bidirectional Wan, where a forward covers the whole video at one timestep, but it degenerates here: every chunk runs the same four timesteps, so an ``e0``-only indicator is *identical* for every chunk and the cache decision cannot adapt to content. We therefore default to TeaCache's original definition -- the relative L1 distance of the **timestep-modulated noisy input** (the tensor block 0 feeds to its self-attention), which is what TeaCache uses on HunyuanVideo/FLUX. ``--indicator e0`` restores the literal Wan port. """ import math import os from dataclasses import dataclass, field from typing import Optional import numpy as np import torch from .selective import run_blocks_selective # --------------------------------------------------------------------------- # # shared helpers # --------------------------------------------------------------------------- # def run_blocks_full(model, x, kwargs, kv_cache, crossattn_cache, current_start, cache_start): for i, block in enumerate(model.blocks): kw = dict(kwargs) kw.update(kv_cache=kv_cache[i], crossattn_cache=crossattn_cache[i], current_start=current_start, cache_start=cache_start) x = block(x, **kw) return x def modulated_input(model, x, e0): """The tensor block 0 hands to its self-attention: norm1(x) * (1+scale) + shift.""" b0 = model.blocks[0] e = (b0.modulation.unsqueeze(1) + e0).chunk(6, dim=2) num_frames = e0.shape[1] frame_seqlen = x.shape[1] // num_frames return (b0.norm1(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2) def rel_l1(cur, prev): # Reduce in fp32: a bf16 mean quantises the distance to ~3 significant digits, # which is coarser than the spread the threshold has to discriminate. 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 x: torch.Tensor # [B, L, C] after patch-embed e0: torch.Tensor # [B, F, 6, C] kwargs: dict kv_cache: list crossattn_cache: list current_start: int cache_start: int grid_sizes: torch.Tensor block_idx: int # which chunk of the video step_idx: int # denoising step inside the chunk forced_full: bool class CacheMethod: """Base class. ``forward`` returns (x_after_blocks, compute_fraction).""" name = "base" 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.model, ctx.x, ctx.e0) 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} # --------------------------------------------------------------------------- # # 1. TeaCache -- one decision for the whole forward, residual copy # --------------------------------------------------------------------------- # class TeaCache(CacheMethod): """Timestep-Embedding Aware caching (Liu et al.). A single global accumulator over the rescaled relative L1 distance of the modulated input; when it stays below ``thresh`` the whole 30-block residual from the last computed step is reused. """ 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.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) 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. FlowCache -- independent policy per chunk-group, residual copy per group # --------------------------------------------------------------------------- # class FlowCache(CacheMethod): """Chunkwise adaptive caching (FlowCache, ICLR 2026). FlowCache's premise is that the units a single forward covers denoise at different rates and therefore need *independent* caching policies. In SkyReels-V2 those units are the chunks of the diffusion-forcing window; here a forward covers one chunk of ``num_frame_per_block`` latent frames, so the corresponding units are frame groups inside that chunk (``group_size`` frames each, 1 by default). Each group keeps its own accumulator, previous indicator and cached residual, exactly as ``transformer_flowcache.py`` keeps ``acc[g] / prev[g] / res[g]``. """ name = "flowcache" def __init__(self, thresh=0.0, group_size=1, **kw): super().__init__(**kw) self.thresh = float(thresh) self.group_size = int(group_size) self.reset_video() def reset_video(self): self.begin_chunk(-1) def begin_chunk(self, block_idx): # A chunk's groups live for that chunk's whole denoising trajectory and # then leave the window, so their policy state starts fresh per chunk. self.acc = {} self.prev_ind = {} self.residual = None def _groups(self, ctx): num_frames = ctx.e0.shape[1] frame_seqlen = ctx.x.shape[1] // num_frames gs = min(self.group_size, num_frames) bounds = [] for start in range(0, num_frames, gs): end = min(start + gs, num_frames) bounds.append((start * frame_seqlen, end * frame_seqlen)) return bounds def forward(self, ctx): ind = self.indicator(ctx) groups = self._groups(ctx) if ctx.forced_full or self.residual is None: ori = ctx.x x = run_blocks_full(ctx.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) self.residual = x - ori for g, (lo, hi) in enumerate(groups): self.acc[g] = 0.0 self.prev_ind[g] = ind[:, lo:hi] return x, 1.0 recompute = [] for g, (lo, hi) in enumerate(groups): cur = ind[:, lo:hi] self.acc[g] = self.acc.get(g, 0.0) + self.rescale(rel_l1(cur, self.prev_ind[g])) if self.acc[g] >= self.thresh: self.acc[g] = 0.0 recompute.append(g) self.prev_ind[g] = cur x = ctx.x + self.residual if not recompute: return x, 0.0 sel = torch.cat([torch.arange(groups[g][0], groups[g][1], device=ctx.x.device) for g in recompute]) x_new = run_blocks_selective(ctx.model, ctx.x, sel, ctx.e0, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start, ctx.grid_sizes) x = x.index_copy(1, sel, x_new) self.residual = self.residual.index_copy(1, sel, x_new - ctx.x[:, sel]) return x, sel.numel() / ctx.x.shape[1] def config(self): return dict(super().config(), thresh=self.thresh, group_size=self.group_size) # --------------------------------------------------------------------------- # # 3. 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, num_frames, frame_seqlen): x = x + (sa.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e2).flatten(1, 2) x = x + ca x = x + (ffn.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e5).flatten(1, 2) return x def _taylor_block_o1(x, sa0, sa1, ca0, ca1, f0, f1, e2, e5, d, num_frames, frame_seqlen): return _taylor_block_add_eager(x, sa0 + d * sa1, ca0 + d * ca1, f0 + d * f1, e2, e5, num_frames, frame_seqlen) def _taylor_block_o0(x, sa0, ca0, f0, e2, e5, num_frames, frame_seqlen): return _taylor_block_add_eager(x, sa0, ca0, f0, e2, e5, num_frames, frame_seqlen) # Upstream applies @torch.compile to the equivalent `wan_attention_cache_forward`: # without fusion the forecast is ~90 separate bandwidth-bound kernels per step and # eats most of the saving it is supposed to deliver. if os.environ.get("CACHELIB_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): """From Reusing to Forecasting (TaylorSeer, ICCV 2025). Caches each block's self-attention / cross-attention / FFN output together with finite-difference derivatives, and on a skipped step *extrapolates* those features with a Taylor series rather than copying them. The current step's AdaLN modulation is still applied to the forecast, and the hidden state still flows through the residual stream, so the skipped step is not a pure copy. The schedule is TaylorSeer's uniform ``fresh_threshold`` interval. A fractional interval is supported so intermediate speedups are reachable: the per-chunk pattern alternates between ceil and floor intervals with a Bresenham-style accumulator, giving the requested *average* interval over the video instead of a content-adaptive threshold (TaylorSeer has none). """ 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.activated = [] # Bresenham pick between the two integer intervals bracketing self.interval. 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): should_calc = self._should_calc(ctx) if should_calc: 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 # -- full step: run the blocks and update the derivative cache ------------- def _forward_record(self, ctx): if len(self.activated) >= 1: dt = ctx.step_idx - self.activated[-1] else: dt = 1 x = ctx.x num_frames = ctx.e0.shape[1] frame_seqlen = x.shape[1] // num_frames for i, block in enumerate(ctx.model.blocks): e = (block.modulation.unsqueeze(1) + ctx.e0).chunk(6, dim=2) slot = self.cache.setdefault(i, {}) y = block.self_attn( (block.norm1(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2), ctx.kwargs["seq_lens"], ctx.grid_sizes, ctx.kwargs["freqs"], ctx.kwargs["block_mask"], ctx.kv_cache[i], ctx.current_start, ctx.cache_start) self._update_derivatives(slot, "sa", y, dt) x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2) y = block.cross_attn(block.norm3(x), ctx.kwargs["context"], None, crossattn_cache=ctx.crossattn_cache[i]) self._update_derivatives(slot, "ca", y, dt) x = x + y y = block.ffn( (block.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[4]) + e[3]).flatten(1, 2)) self._update_derivatives(slot, "ffn", y, dt) x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]).flatten(1, 2) 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 # -- cached step: forecast each module output ------------------------------ def _forward_taylor(self, ctx, distance): x = ctx.x num_frames = ctx.e0.shape[1] frame_seqlen = x.shape[1] // num_frames d = float(distance) for i, block in enumerate(ctx.model.blocks): e = (block.modulation.unsqueeze(1) + ctx.e0).chunk(6, dim=2) slot = self.cache[i] sa, ca, ffn = slot["sa"], slot["ca"], slot["ffn"] 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], d, num_frames, frame_seqlen) elif self.max_order <= 1: x = _taylor_block_o0(x, sa[0], ca[0], ffn[0], e[2], e[5], num_frames, frame_seqlen) else: x = _taylor_block_add_eager( x, taylor_formula(sa, distance), taylor_formula(ca, distance), taylor_formula(ffn, distance), e[2], e[5], num_frames, frame_seqlen) return x def config(self): return dict(super().config(), interval=self.interval, max_order=self.max_order) # --------------------------------------------------------------------------- # # 4. MotionCache -- motion-weighted, token-level reuse # --------------------------------------------------------------------------- # class MotionCache(CacheMethod): """Motion-aware token-level caching (MotionCache, ICML 2026). Inter-frame differences of the last computed output give each spatial token a motion weight; the per-token rescaled distance is multiplied by that weight and accumulated, so high-motion tokens cross the threshold sooner and are recomputed more often while static tokens are reused for longer. Tokens that do cross the threshold go through the selective forward. """ 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): """Inter-frame relative differences of the computed output -> [L] weights.""" 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.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) 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, ctx.cache_start, ctx.grid_sizes) x = x.index_copy(1, sel, x_new) self.residual = self.residual.index_copy(1, sel, x_new - ctx.x[:, sel]) 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) # --------------------------------------------------------------------------- # class NoCache(CacheMethod): """Baseline: every step runs the full DiT.""" name = "none" def forward(self, ctx): x = run_blocks_full(ctx.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) return x, 1.0 class DirectReuse(CacheMethod): """The naive cache baseline: every cacheable step reuses the residual of the last computed step, with no indicator and no threshold. `FRFF` / `FRRF` / `FRRR` are this method under the corresponding schedule, i.e. 3 / 2 / 1 full DiT evaluations per chunk. It is the fixed-schedule ablation the adaptive methods have to beat.""" name = "reuse" def reset_video(self): self.prev_residual = None def forward(self, ctx): if ctx.forced_full or self.prev_residual is None: ori = ctx.x x = run_blocks_full(ctx.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) self.prev_residual = x - ori return x, 1.0 return ctx.x + self.prev_residual, 0.0 class CalibrationProbe(CacheMethod): """Always full, but records the pairs TeaCache fits its rescale polynomial on. For every pair of consecutive denoising steps inside a chunk it logs ``(rel_l1(indicator_t, indicator_{t-1}), rel_l1(residual_t, residual_{t-1}))``. The residual is exactly the tensor a skipping method would reuse, so the fitted polynomial turns the cheap input-side distance into an estimate of the error that reuse would incur -- which is what the accumulator is meant to bound. """ name = "calibrate" def __init__(self, **kw): super().__init__(**kw) self.samples = [] self.prev_ind = None self.prev_res = None def begin_chunk(self, block_idx): self.prev_ind = None self.prev_res = None def forward(self, ctx): ind = self.indicator(ctx) ori = ctx.x x = run_blocks_full(ctx.model, ctx.x, ctx.kwargs, ctx.kv_cache, ctx.crossattn_cache, ctx.current_start, ctx.cache_start) 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 REGISTRY = { "none": NoCache, # Same forward as "none"; a distinct label for the deterministic flow-Euler # sampler rows so tooling never mistakes them for the FFFF reference. "euler": NoCache, "reuse": DirectReuse, "calibrate": CalibrationProbe, "teacache": TeaCache, "flowcache": FlowCache, "taylorseer": TaylorSeer, "motioncache": MotionCache, } def build_method(name, **kw): if name not in REGISTRY: raise KeyError(f"unknown cache method {name!r}; have {sorted(REGISTRY)}") cls = REGISTRY[name] sig_free = {k: v for k, v in kw.items() if v is not None} return cls(**sig_free)