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