"""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