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