| """Route ``HunyuanVideo_1_5_DiffusionTransformer.forward_vision`` through a cache method. |
| |
| The preamble (patch embed, time/action modulation vector, RoPE tables, per-token |
| cameras) and the tail (``final_layer``, ``unpatchify``) are reproduced from the |
| upstream method; only the 54-block loop is handed to the controller, and only on |
| denoising forwards -- context passes (``cache_vision=True``) run untouched. |
| """ |
|
|
| import types |
|
|
| import torch |
| from einops import rearrange, repeat |
|
|
| from hyvideo.commons.parallel_states import get_parallel_state |
|
|
| from .methods import StepCtx |
|
|
| DEFAULT_SCHEDULE = "FxxF" |
|
|
|
|
| def parse_schedule(schedule, num_steps): |
| |
| |
| 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) |
| img, frac = self.method.forward(ctx) |
| self.records.append({"block": self.block_idx, "step": self.step_idx, |
| "compute_fraction": float(frac)}) |
| return img |
|
|
| 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_vision(self, hidden_states, timestep, timestep_r=None, freqs_cos=None, |
| freqs_sin=None, return_dict=False, mask_type="t2v", |
| extra_kwargs=None, action=None, viewmats=None, Ks=None, |
| kv_cache=None, cache_vision=False, rope_temporal_size=4, |
| start_rope_start_idx=0): |
| ctrl = getattr(self, "_cache_ctrl", None) |
| if ctrl is None or not ctrl.active or cache_vision: |
| return self._orig_forward_vision( |
| hidden_states=hidden_states, timestep=timestep, timestep_r=timestep_r, |
| freqs_cos=freqs_cos, freqs_sin=freqs_sin, return_dict=return_dict, |
| mask_type=mask_type, extra_kwargs=extra_kwargs, action=action, |
| viewmats=viewmats, Ks=Ks, kv_cache=kv_cache, cache_vision=cache_vision, |
| rope_temporal_size=rope_temporal_size, start_rope_start_idx=start_rope_start_idx) |
| assert not get_parallel_state().sp_enabled, "cache methods are single-GPU" |
|
|
| |
| img = x = hidden_states |
| t = timestep |
| bs, _, ot, oh, ow = x.shape |
| tt, th, tw = ot // self.patch_size[0], oh // self.patch_size[1], ow // self.patch_size[2] |
| self.attn_param["thw"] = [tt, th, tw] |
| rope_temporal_size = rope_temporal_size // self.patch_size[0] |
| if freqs_cos is None and freqs_sin is None: |
| freqs_cos, freqs_sin = self.get_rotary_pos_embed((rope_temporal_size, th, tw)) |
| per_latent_size = th * tw |
| start_index = start_rope_start_idx * per_latent_size |
| end_index = (start_rope_start_idx + tt) * per_latent_size |
| freqs_cos = freqs_cos[start_index:end_index, ...] |
| freqs_sin = freqs_sin[start_index:end_index, ...] |
| img = self.img_in(img) |
| action = action.reshape(-1) |
| t = t.reshape(-1) |
| vec = self.time_in(t) |
| vec = vec + self.action_in(action) |
| vec = repeat(vec, "(B T) C->B (T H W) C", B=img.shape[0], H=th, W=tw) |
| viewmats = repeat(viewmats, "B T M N->B (T H W) M N", H=th, W=tw) |
| Ks = repeat(Ks, "B T M N->B (T H W) M N", H=th, W=tw) |
| vec = rearrange(vec, "B S C->(B S) C") |
| |
| |
| freqs_cis = (freqs_cos.to(img.device), freqs_sin.to(img.device)) |
|
|
| img = ctrl.run(dict(model=self, img=img, vec=vec, freqs_cis=freqs_cis, |
| viewmats=viewmats, Ks=Ks, kv_cache=kv_cache, grid=(tt, th, tw))) |
|
|
| img = self.final_layer(img, vec) |
| img = self.unpatchify(img, tt, th, tw) |
| assert return_dict is False |
| return (img, None) |
|
|
|
|
| def install(transformer, controller): |
| if not hasattr(transformer, "_orig_forward_vision"): |
| transformer._orig_forward_vision = transformer.forward_vision |
| transformer.forward_vision = types.MethodType(_cached_forward_vision, transformer) |
| transformer._cache_ctrl = controller |
| return transformer |
|
|