File size: 5,857 Bytes
36cdb93
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Block-causal streaming forward for Wan2.1, corrected and de-overheaded.

Corrections relative to streaming_blocks.py, each backed by a measurement
recorded in PROGRESS.md:

  1. TIMESTEP SCALE. final.pt was fine-tuned with t in [0,1]
     (train_streaming.py:152 draws torch.rand()), not the scheduler's [0,1000].
     Callers must pass t in [0,1]; `timestep_to_train_scale` does the conversion.
     Measured: normalised flow error 0.499 -> 0.168.
  2. TEMPORAL RoPE. The old path gave every frame temporal index 0. Here each
     latent frame is rotated at its true absolute index, so cached keys carry
     real temporal position. Measured: 0.180 -> 0.107 (uniform, t=0.5).
  3. NO CAUSAL MASK. The K/V cache holds only past+current frames, so full
     attention over it is already block-causal. The old code passed causal=True
     with Lq != Lk, which imposes a spurious raster ordering *within* a frame.
     Measured: slightly better error AND 1.56x faster (no mask materialisation).

Overhead reductions, targeting the profile
"""
import torch

from wan.modules.attention import flash_attention

from .rope import RopeTable, apply_rope
from .kvcache import StreamingKVCache


def timestep_to_train_scale(t_val, num_train_timesteps=1000):
    """Scheduler timestep (0..1000) -> the t in [0,1] final.pt was trained on."""
    return float(t_val) / float(num_train_timesteps)


class ModulationCache:
    """Per-block AdaLN modulation chunks for one timestep.

    `(block.modulation + e0).chunk(6)` is identical for every frame at a given
    denoising step, but the old code recomputed it per block *per frame*.
    """

    def __init__(self, blocks, e0):
        self.chunks = []
        with torch.amp.autocast('cuda', dtype=torch.float32):
            for blk in blocks:
                self.chunks.append((blk.modulation + e0).chunk(6, dim=1))

    def __getitem__(self, i):
        return self.chunks[i]


def block_forward(blk, x, mod, rope_tbl, t_index, cache, layer, ctx, ctx_lens):
    """One WanAttentionBlock in block-causal streaming mode.

    x: [B, S, dim] tokens of the CURRENT latent frame only.
    Returns the updated x. K/V for this frame is written (uncommitted) to `cache`.
    """
    ec = mod
    sa_in = blk.norm1(x).float() * (1 + ec[1]) + ec[0]

    b, s = sa_in.shape[0], sa_in.shape[1]
    n = blk.num_heads
    d = blk.dim // n
    sa = blk.self_attn

    q = sa.norm_q(sa.q(sa_in)).view(b, s, n, d)
    k = sa.norm_k(sa.k(sa_in)).view(b, s, n, d)
    v = sa.v(sa_in).view(b, s, n, d)

    tbl = rope_tbl.frame(t_index)
    q = apply_rope(q, tbl)
    k = apply_rope(k, tbl)

    # Write current chunk, then attend over past+current as one view.
    cache.write(layer, k.to(cache.k.dtype), v.to(cache.v.dtype))
    ctx_k, ctx_v = cache.context(layer, s)

    y = flash_attention(q=q.to(ctx_k.dtype), k=ctx_k, v=ctx_v,
                        window_size=(-1, -1), causal=False)
    y = sa.o(y.flatten(2))

    with torch.amp.autocast('cuda', dtype=torch.float32):
        x = x + y * ec[2]
    x = x + blk.cross_attn(blk.norm3(x), ctx, ctx_lens)
    y = blk.ffn(blk.norm2(x).float() * (1 + ec[4]) + ec[3])
    with torch.amp.autocast('cuda', dtype=torch.float32):
        x = x + y * ec[5]
    return x


@torch.no_grad()
def frame_forward(model, latent_frame, e, mod, rope_tbl, t_index, cache,
                  ctx, ctx_lens):
    """Denoise ONE latent frame against the cached clean past.

    latent_frame: [B, C, 1, H, W]. Returns predicted velocity [C, 1, H, W].
    Does not commit the frame's K/V -- the caller commits once the chunk is final.
    """
    x = model.patch_embedding(latent_frame)
    grid = torch.stack([torch.tensor(x.shape[2:], dtype=torch.long, device=x.device)
                        for _ in range(x.shape[0])])
    x = x.flatten(2).transpose(1, 2)

    for i, blk in enumerate(model.blocks):
        x = block_forward(blk, x, mod[i], rope_tbl, t_index, cache, i, ctx, ctx_lens)

    return model.unpatchify(model.head(x, e), grid)[0]


@torch.no_grad()
def sequence_forward(model, latents, e, mod, rope_tbl, cache, ctx, ctx_lens,
                     start_index=0, commit=True):
    """Run a run of latent frames causally, committing each as it completes.

    Used to (a) prime the cache from clean context frames and (b) evaluate the
    model over a whole sequence for diagnostics. latents: [B, C, F, H, W].
    """
    outs = []
    for f in range(latents.shape[2]):
        out = frame_forward(model, latents[:, :, f:f + 1], e, mod, rope_tbl,
                            start_index + f, cache, ctx, ctx_lens)
        outs.append(out)
        if commit:
            cache.commit(rope_tbl.seq)
    return torch.cat(outs, dim=1)


def make_rope_table(model, h_patches, w_patches, max_frames, device):
    return RopeTable(model.freqs.to(device), h_patches, w_patches, max_frames, device)


def make_cache(model, tokens_per_frame, max_frames, device, dtype=torch.bfloat16):
    n = model.num_heads
    d = model.dim // n
    return StreamingKVCache(
        num_layers=len(model.blocks),
        max_tokens=tokens_per_frame * max_frames,
        num_heads=n, head_dim=d, device=device, dtype=dtype,
        scratch_tokens=tokens_per_frame)


def latent_geometry(width, height, vae_stride=(4, 8, 8), patch_size=(1, 2, 2)):
    """Pixel (W,H) -> latent (H_lat, W_lat) and patch grid (H_p, W_p).

    Mirrors WanModel/run_demo: target_shape[2] = size[1]//vae_stride[1] (height),
    target_shape[3] = size[0]//vae_stride[2] (width).
    """
    h_lat = height // vae_stride[1]
    w_lat = width // vae_stride[2]
    if h_lat % patch_size[1] or w_lat % patch_size[2]:
        raise ValueError(f'{width}x{height} -> latent {h_lat}x{w_lat} not '
                         f'divisible by patch size {patch_size[1:]}')
    return h_lat, w_lat, h_lat // patch_size[1], w_lat // patch_size[2]