File size: 8,933 Bytes
b34c6c3 | 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 201 202 203 204 205 | """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
|