Download lbcache/selective.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 6.82 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/selective.py
- Command line
-
hf download hf://Cccccz/comparison/lbcache/selective.py
-
curl -L -o selective.py https://huggingface.co/Cccccz/comparison/resolve/main/lbcache/selective.py
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 | |