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