"""Route ``CausalWanModel._forward_inference`` through a cache method. The DiT preamble (patch embed, time embed, text embed) and the head/unpatchify tail are reproduced verbatim from ``wan/modules/causal_model.py``; only the 30-block loop in between is handed to the active cache method. """ import torch from wan.modules.causal_model import CausalWanModel, sinusoidal_embedding_1d from .methods import StepCtx, run_blocks_full DEFAULT_SCHEDULE = "FxxF" def parse_schedule(schedule, num_steps): """``"FxxF"`` -> ``(0, 3)``: the steps marked ``F`` always run the full DiT, the ones marked ``x`` (or ``?``) may be served from cache. Step 0 must be ``F`` -- nothing is cached yet when a chunk starts.""" # 'R' spells out "reuse" in the naive-cache baselines (FRFF / FRRF / FRRR); # it is the same thing as 'x': a step the method may serve from cache. s = schedule.strip().upper().replace("?", "X").replace("R", "X") if len(s) != num_steps or set(s) - {"F", "X"}: raise ValueError(f"schedule {schedule!r} must be {num_steps} chars of F/x") if s[0] != "F": raise ValueError(f"schedule {schedule!r}: step 0 must be F") return tuple(i for i, c in enumerate(s) if c == "F") def schedule_string(forced, num_steps): return "".join("F" if i in forced else "x" for i in range(num_steps)) class CacheController: """Tracks where in the (chunk, denoising step) grid the model currently is. ``active`` is set only around the four denoising forwards of a chunk. The KV-cache refresh pass and the initial-latent context passes run the full DiT and are neither cached nor timed. ``forced_steps`` is the schedule: those steps always run the full DiT and every other step is the method's to decide. ``(0, -1)`` is ``FxxF``. """ def __init__(self, method, num_steps=4, forced_steps=(0, -1), first_chunk_forced_steps=None): self.method = method self.num_steps = num_steps self.forced = {s % num_steps for s in forced_steps} # Chunk 0 has no previous chunk to reuse from and sets the appearance the # whole video inherits, so it can be given its own (usually full) schedule. # The first chunk's schedule may be longer than the others' (naive-N-step # rows with chunk 0 at the full four steps), so it is not reduced mod num_steps. self.first_chunk_forced = (None if first_chunk_forced_steps is None else set(first_chunk_forced_steps)) self.schedule = schedule_string(self.forced, num_steps) self.first_chunk_schedule = (None if self.first_chunk_forced is None else schedule_string(self.first_chunk_forced, max(num_steps, max(self.first_chunk_forced, default=-1) + 1))) self.active = False self.block_idx = -1 self.step_idx = -1 self.records = [] def reset_video(self): self.method.reset_video() self.records = [] def begin_chunk(self, block_idx): self.block_idx = block_idx self.method.begin_chunk(block_idx) def denoise_step(self, step_idx): self.step_idx = step_idx self.active = True def end_step(self): self.active = False def forced_now(self): if self.first_chunk_forced is not None and self.block_idx == 0: return self.first_chunk_forced return self.forced def run(self, model, x, e0, kwargs, kv_cache, crossattn_cache, current_start, cache_start, grid_sizes): if not self.active: return run_blocks_full(model, x, kwargs, kv_cache, crossattn_cache, current_start, cache_start) ctx = StepCtx( model=model, x=x, e0=e0, kwargs=kwargs, kv_cache=kv_cache, crossattn_cache=crossattn_cache, current_start=current_start, cache_start=cache_start, grid_sizes=grid_sizes, block_idx=self.block_idx, step_idx=self.step_idx, forced_full=self.step_idx in self.forced_now(), ) x, frac = self.method.forward(ctx) self.records.append({ "block": self.block_idx, "step": self.step_idx, "compute_fraction": float(frac), "forced": self.step_idx in self.forced_now(), }) return x def summary(self): denoise = self.records if not denoise: return {} total = len(denoise) compute = sum(r["compute_fraction"] for r in denoise) middle = [r for r in denoise if not r.get("forced", r["step"] in self.forced)] # A token-wise method is only doing token-wise work while its cacheable # steps are *partial*. Steps that select nothing (or everything) are # behaviourally a whole-step skip (or a full step), so record how the # budget is spread, not just how large it is. active = [r["compute_fraction"] for r in middle if 1e-9 < r["compute_fraction"] < 1 - 1e-9] n_mid = len(middle) or 1 return { "denoise_forwards": total, "compute_equivalent_forwards": compute, "middle_steps": len(middle), "middle_compute_equivalent": sum(r["compute_fraction"] for r in middle), "active_step_ratio": len(active) / n_mid, "empty_step_ratio": sum(r["compute_fraction"] <= 1e-9 for r in middle) / n_mid, "full_step_ratio": sum(r["compute_fraction"] >= 1 - 1e-9 for r in middle) / n_mid, "mean_selected_fraction_active": (sum(active) / len(active)) if active else 0.0, "flops_speedup_estimate": total / compute if compute else float("inf"), } def _cached_forward_inference(self, x, t, context, seq_len, clip_fea=None, y=None, kv_cache=None, crossattn_cache=None, current_start=0, cache_start=0): ctrl = getattr(self, "_cache_ctrl", None) if ctrl is None: return self._orig_forward_inference( x, t, context, seq_len, clip_fea, y, kv_cache, crossattn_cache, current_start, cache_start) if self.model_type == "i2v": assert clip_fea is not None and y is not None device = self.patch_embedding.weight.device if self.freqs.device != device: self.freqs = self.freqs.to(device) if y is not None: x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] x = [self.patch_embedding(u.unsqueeze(0)) for u in x] grid_sizes = torch.stack( [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) x = [u.flatten(2).transpose(1, 2) for u in x] seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) assert seq_lens.max() <= seq_len x = torch.cat(x) e = self.time_embedding( sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x)) e0 = self.time_projection(e).unflatten( 1, (6, self.dim)).unflatten(dim=0, sizes=t.shape) context_lens = None context = self.text_embedding( torch.stack([ torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context ])) if clip_fea is not None: context = torch.concat([self.img_emb(clip_fea), context], dim=1) kwargs = dict( e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=self.freqs, context=context, context_lens=context_lens, block_mask=self.block_mask, ) x = ctrl.run(self, x, e0, kwargs, kv_cache, crossattn_cache, current_start, cache_start, grid_sizes) x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2)) x = self.unpatchify(x, grid_sizes) return torch.stack(x) def install(model, controller): """Attach ``controller`` to a ``CausalWanModel`` and patch the class once.""" if not hasattr(CausalWanModel, "_orig_forward_inference"): CausalWanModel._orig_forward_inference = CausalWanModel._forward_inference CausalWanModel._forward_inference = _cached_forward_inference model._cache_ctrl = controller return model