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