Spaces:
Running on Zero
Running on Zero
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]
|