Download lbcache/methods.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 17.6 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/methods.py
- Command line
-
hf download hf://Cccccz/comparison/lbcache/methods.py
-
curl -L -o methods.py https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/methods.py
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 | |
| 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}) | |