comparison / cachelib /methods.py
Cccccz's picture
Add files using upload-large-folder tool
f87692b verified
Raw History Blame Contribute Delete
25.1 kB
"""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)