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