comparison / lbcache /patch.py
Cccccz's picture
Add files using upload-large-folder tool
b34c6c3 verified
Raw History Blame Contribute Delete
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