comparison / hycache /patch.py
Cccccz's picture
Add files using upload-large-folder tool
b34c6c3 verified
Raw
History Blame Contribute Delete
8.13 kB
"""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):
# '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
# Chunk 0 may get its own schedule (usually all-full, possibly with more
# steps than the other chunks); 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.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)]
# 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": 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"
# -- preamble, verbatim from forward_vision -------------------------------------
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")
# get_rotary_pos_embed builds the tables on CPU and upstream moves them inside
# apply_rotary_emb; the selective path indexes them with a CUDA index first.
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