comparison / lbcache /selective.py
Cccccz's picture
Add files using upload-large-folder tool
b34c6c3 verified
Raw History Blame Contribute Delete
6.82 kB
"""Block-level re-implementation of ``CausalWanAttentionBlock.forward`` (LingBot
``wan/modules/model_fast.py``) in three flavours:
* ``block_forward`` -- the full block, optionally capturing the
self-attn / cross-attn / FFN outputs (TaylorSeer);
* ``block_forward_selective`` -- the block on a token subset, refreshing only the
selected KV slots (MotionCache);
* ``block_cam`` -- the per-token camera scale/shift a block applies
after self-attention (constant within a chunk).
Differences from ``cachelib/selective.py`` (Self-Forcing): the AdaLN modulation
``e`` is already per token ([B, L, 6, C]), the camera injection sits between the
self-attention residual and the cross-attention, and the self-attention takes
``max_attention_size`` / ``frame_seqlen`` from the call. Exactness against the
stock forward is checked by ``lb_test_exact.py``.
"""
import torch
import torch.nn.functional as torch_F
from wan.modules.attention import attention
def block_mod(block, e0):
"""The six per-token modulation tensors [B, L, C] of one block."""
with torch.amp.autocast('cuda', dtype=torch.float32):
e = (block.modulation.unsqueeze(0) + e0).chunk(6, dim=2)
return [t.squeeze(2) for t in e]
def block_cam(block, plucker):
"""(scale, shift) [B, L, C] of the block's camera injection, or (None, None)."""
if plucker is None:
return None, None
h = block.cam_injector_layer2(torch_F.silu(block.cam_injector_layer1(plucker)))
h = h + plucker
return block.cam_scale_layer(h), block.cam_shift_layer(h)
def block_forward(block, x, e0, kw, kv_cache, crossattn_cache, current_start, features=None):
"""Stock block forward, op for op. ``features`` (dict) receives the pre-gate
self-attn / cross-attn / FFN outputs under keys ``sa`` / ``ca`` / ``ffn``."""
e = block_mod(block, e0)
y = block.self_attn(
block.norm1(x).float() * (1 + e[1]) + e[0],
kw["seq_lens"], kw["grid_sizes"], kw["freqs"], kv_cache, current_start,
kw["max_attention_size"], frame_seqlen=kw["frame_seqlen"], seq_lens_int=None)
if features is not None:
features["sa"] = y
with torch.amp.autocast('cuda', dtype=torch.float32):
x = x + y * e[2]
plucker = (kw["dit_cond_dict"] or {}).get("c2ws_plucker_emb")
scale, shift = block_cam(block, plucker)
if scale is not None:
x = (1.0 + scale) * x + shift
y = block.cross_attn(block.norm3(x), kw["context"], kw["context_lens"],
crossattn_cache=crossattn_cache,
cross_attn_first_call=kw["cross_attn_first_call"])
if features is not None:
features["ca"] = y
x = x + y
y = block.ffn(block.norm2(x).float() * (1 + e[4]) + e[3])
if features is not None:
features["ffn"] = y
with torch.amp.autocast('cuda', dtype=torch.float32):
x = x + y * e[5]
return x
def build_causal_freqs(grid_sizes, freqs, start_frame):
"""Per-token RoPE multipliers [seq_len, 1, c] for one chunk (``causal_rope_apply``)."""
c = freqs.size(1)
parts = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
f, h, w = grid_sizes[0].tolist()
seq_len = f * h * w
return torch.cat([
parts[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1),
parts[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
parts[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),
], dim=-1).reshape(seq_len, 1, -1)
def rope_apply_rows(x, freqs_rows):
b, s, n, _ = x.shape
x_c = torch.view_as_complex(x.to(torch.float64).reshape(b, s, n, -1, 2))
return torch.view_as_real(x_c * freqs_rows.unsqueeze(0)).flatten(3).type_as(x)
def self_attn_selective(attn, x_sel, sel_idx, freqs_rows, kv_cache, current_start,
block_seq_len, max_attention_size):
"""Self-attention over a token subset, refreshing only those KV slots. Only
valid on a repeat step of a chunk (its slots already exist: no eviction)."""
b, n_sel = x_sel.shape[:2]
n, d = attn.num_heads, attn.head_dim
q = attn.norm_q(attn.q(x_sel)).view(b, n_sel, n, d)
k = attn.norm_k(attn.k(x_sel)).view(b, n_sel, n, d)
v = attn.v(x_sel).view(b, n_sel, n, d)
rows = freqs_rows[sel_idx]
roped_q = rope_apply_rows(q, rows).type_as(v)
roped_k = rope_apply_rows(k, rows).type_as(v)
current_end = current_start + block_seq_len
assert current_end == kv_cache["global_end_index"].item(), "selective step must repeat a chunk"
local_end_index = kv_cache["local_end_index"].item()
local_start_index = local_end_index - block_seq_len
slots = local_start_index + sel_idx
kv_cache["k"].index_copy_(1, slots, roped_k)
kv_cache["v"].index_copy_(1, slots, v)
lo = max(0, local_end_index - max_attention_size)
out = attention(roped_q, kv_cache["k"][:, lo:local_end_index], kv_cache["v"][:, lo:local_end_index])
return attn.o(out.flatten(2))
def block_forward_selective(block, x_sel, sel_idx, e0, freqs_rows, kw, kv_cache,
crossattn_cache, current_start, block_seq_len):
e = [t[:, sel_idx] for t in block_mod(block, e0)]
y = self_attn_selective(
block.self_attn, block.norm1(x_sel).float() * (1 + e[1]) + e[0], sel_idx, freqs_rows,
kv_cache, current_start, block_seq_len, kw["max_attention_size"])
with torch.amp.autocast('cuda', dtype=torch.float32):
x_sel = x_sel + y * e[2]
plucker = (kw["dit_cond_dict"] or {}).get("c2ws_plucker_emb")
scale, shift = block_cam(block, None if plucker is None else plucker[:, sel_idx])
if scale is not None:
x_sel = (1.0 + scale) * x_sel + shift
y = block.cross_attn(block.norm3(x_sel), kw["context"], kw["context_lens"],
crossattn_cache=crossattn_cache, cross_attn_first_call=False)
x_sel = x_sel + y
y = block.ffn(block.norm2(x_sel).float() * (1 + e[4]) + e[3])
with torch.amp.autocast('cuda', dtype=torch.float32):
x_sel = x_sel + y * e[5]
return x_sel
def run_blocks_selective(model, x, sel_idx, e0, kw, kv_cache, crossattn_cache, current_start):
"""Run all blocks on ``x[:, sel_idx]``; returns the new rows [B, n_sel, C]."""
grid_sizes = kw["grid_sizes"]
f, h, w = grid_sizes[0].tolist()
frame_seqlen = h * w
block_seq_len = f * frame_seqlen
freqs_rows = build_causal_freqs(grid_sizes, kw["freqs"], current_start // frame_seqlen)
x_sel = x[:, sel_idx]
for i, block in enumerate(model.blocks):
x_sel = block_forward_selective(block, x_sel, sel_idx, e0, freqs_rows, kw,
kv_cache[i], crossattn_cache[i], current_start, block_seq_len)
return x_sel