Download lbcache/patch.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 8.93 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/patch.py
- Command line
-
hf download hf://Cccccz/comparison/lbcache/patch.py
-
curl -L -o patch.py https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/patch.py
8.93 kB
| """Route ``WanModelFast.forward`` through a cache method. | |
| While the controller is marked active (denoising forwards) the preamble (patch | |
| embedding, time / text / camera embeddings) and the tail (``head``, ``unpatchify``) | |
| are reproduced verbatim from ``wan/modules/model_fast.py`` and only the 30-block | |
| loop is handed to the controller; an output-level method (velocity ``reuse``) is | |
| instead handed the whole stock forward as a callable. The per-chunk context pass | |
| runs the untouched original forward. | |
| """ | |
| import types | |
| import torch | |
| import torch.nn.functional as torch_F | |
| from einops import rearrange | |
| from .methods import StepCtx | |
| def parse_schedule(schedule, num_steps): | |
| # '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: | |
| 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} | |
| self.schedule = schedule_string(self.forced, num_steps) | |
| cacheable = [s for s in range(num_steps) if s not in self.forced] | |
| self.last_cacheable_step = max(cacheable) if cacheable else -1 | |
| self.first_chunk_forced = (None if first_chunk_forced_steps is None | |
| else set(first_chunk_forced_steps)) | |
| self.first_chunk_schedule = None | |
| if self.first_chunk_forced is not None: | |
| n0 = max(num_steps, max(self.first_chunk_forced, default=-1) + 1) | |
| self.first_chunk_schedule = schedule_string(self.first_chunk_forced, n0) | |
| c0 = [s for s in range(n0) if s not in self.first_chunk_forced] | |
| self.first_chunk_last_cacheable = max(c0) if c0 else -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, ctx_kwargs): | |
| first = self.first_chunk_forced is not None and self.block_idx == 0 | |
| ctx = StepCtx(block_idx=self.block_idx, step_idx=self.step_idx, | |
| forced_full=self.step_idx in self.forced_now(), | |
| last_cacheable_step=(self.first_chunk_last_cacheable if first | |
| else self.last_cacheable_step), **ctx_kwargs) | |
| out, frac = self.method.forward(ctx) | |
| self.records.append({"block": self.block_idx, "step": self.step_idx, | |
| "compute_fraction": float(frac)}) | |
| return out | |
| def summary(self): | |
| d = self.records | |
| if not d: | |
| return {} | |
| compute = sum(r["compute_fraction"] for r in d) | |
| middle = [r for r in d if r["step"] not in (self.first_chunk_forced if ( | |
| self.first_chunk_forced is not None and r["block"] == 0) else self.forced)] | |
| 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": len(d), "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": len(d) / compute if compute else float("inf")} | |
| def _cached_forward(self, x, t, context, seq_len, y=None, dit_cond_dict=None, | |
| kv_cache=None, crossattn_cache=None, current_start=0, | |
| max_attention_size=1_000_000, frame_seqlen=None, | |
| cross_attn_first_call=None): | |
| def run_full(): | |
| return self._orig_forward( | |
| x, t, context, seq_len, y=y, dit_cond_dict=dit_cond_dict, kv_cache=kv_cache, | |
| crossattn_cache=crossattn_cache, current_start=current_start, | |
| max_attention_size=max_attention_size, frame_seqlen=frame_seqlen, | |
| cross_attn_first_call=cross_attn_first_call) | |
| ctrl = getattr(self, "_cache_ctrl", None) | |
| if ctrl is None or not ctrl.active: | |
| return run_full() | |
| if getattr(ctrl.method, "level", "blocks") == "output": | |
| return ctrl.run(dict(model=self, run_full=run_full, x=None, e0=None, kwargs=None, | |
| kv_cache=kv_cache, crossattn_cache=crossattn_cache, | |
| current_start=current_start, grid_sizes=None)) | |
| from wan.modules.model import sinusoidal_embedding_1d | |
| # -- preamble, verbatim from WanModelFast.forward --------------------------------- | |
| if self.model_type == 'i2v': | |
| assert 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) | |
| if t.dim() == 1: | |
| t = t.expand(t.size(0), seq_lens) | |
| with torch.amp.autocast('cuda', dtype=torch.float32): | |
| bt = t.size(0) | |
| t = t.flatten() | |
| e = self.time_embedding( | |
| sinusoidal_embedding_1d(self.freq_dim, | |
| t).unflatten(0, (bt, seq_lens)).float()) | |
| e0 = self.time_projection(e).unflatten(2, (6, self.dim)) | |
| assert e.dtype == torch.float32 and e0.dtype == torch.float32 | |
| 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 dit_cond_dict is not None and "c2ws_plucker_emb" in dit_cond_dict: | |
| c2ws_plucker_emb = dit_cond_dict["c2ws_plucker_emb"] | |
| c2ws_plucker_emb = [ | |
| rearrange( | |
| i, | |
| '1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)', | |
| c1=self.patch_size[0], | |
| c2=self.patch_size[1], | |
| c3=self.patch_size[2], | |
| ) for i in c2ws_plucker_emb | |
| ] | |
| c2ws_plucker_emb = torch.cat(c2ws_plucker_emb, dim=1) | |
| c2ws_plucker_emb = self.patch_embedding_wancamctrl(c2ws_plucker_emb) | |
| c2ws_hidden_states = self.c2ws_hidden_states_layer2( | |
| torch_F.silu(self.c2ws_hidden_states_layer1(c2ws_plucker_emb))) | |
| dit_cond_dict = dict(dit_cond_dict) | |
| dit_cond_dict["c2ws_plucker_emb"] = ( | |
| c2ws_plucker_emb + c2ws_hidden_states) | |
| kwargs = dict( | |
| e=e0, | |
| seq_lens=seq_lens, | |
| grid_sizes=grid_sizes, | |
| freqs=self.freqs, | |
| context=context, | |
| context_lens=context_lens, | |
| dit_cond_dict=dit_cond_dict, | |
| max_attention_size=max_attention_size, | |
| frame_seqlen=frame_seqlen, | |
| cross_attn_first_call=cross_attn_first_call) | |
| x = ctrl.run(dict(model=self, run_full=run_full, x=x, e0=e0, kwargs=kwargs, | |
| kv_cache=kv_cache, crossattn_cache=crossattn_cache, | |
| current_start=current_start, grid_sizes=grid_sizes)) | |
| x = self.head(x, e) | |
| x = self.unpatchify(x, grid_sizes) | |
| return [u.float() for u in x] | |
| def install(model, controller): | |
| if not hasattr(model, "_orig_forward"): | |
| model._orig_forward = model.forward | |
| model.forward = types.MethodType(_cached_forward, model) | |
| model._cache_ctrl = controller | |
| return model | |