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