File size: 8,267 Bytes
f87692b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 | """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
|