"""One HY-WorldPlay double-stream block's vision path, re-implemented so the cache methods can (a) run the residual stream on a token subset while the keys/values of the other tokens come from a bank filled at the last full step, (b) capture those keys/values, and (c) capture or replay the pre-gate attention / MLP features TaylorSeer forecasts. With ``sel=None`` and no recording it must match ``MMDoubleStreamBlock.forward_vision`` bit for bit; ``hy_test_exact.py`` checks that. """ import os import torch from einops import rearrange from hyvideo.models.transformers.modules.attention import sequence_parallel_attention_vision from hyvideo.models.transformers.modules.modulate_layers import apply_gate, modulate from hyvideo.models.transformers.modules.posemb_layers import apply_rotary_emb from hyvideo.prope.camera_rope import _prepare_apply_fns_all_dim def _prope_fns(head_dim, viewmats, Ks): # Same call prope_qkv makes; patches / image size are None in upstream too. return _prepare_apply_fns_all_dim(head_dim=head_dim, viewmats=viewmats, Ks=Ks, patches_x=None, patches_y=None, image_width=None, image_height=None) def _bhld(x): return x.permute(0, 2, 1, 3) def block_forward(block, idx, img, vec, freqs_cis, viewmats, Ks, kv_cache, sel=None, bank=None, features=None): """Return the block output for ``img``. img [B, L, C] residual stream (L = all tokens, or len(sel) when sel given) vec [B*L, C] per-token time+action modulation input (rows match img) freqs_cis (cos, sin) RoPE tables for the rows of img viewmats [B, L, 4, 4], Ks [B, L, 3, 3] per-token cameras for the rows of img sel LongTensor of the token indices img holds, when it is a subset. The other tokens' k/v come from ``bank[idx]`` (full-sequence tensors). bank dict idx -> [k, v] ([B, S, H, D], post-norm, pre-RoPE). Written at full steps, updated in place at selected rows during selective steps. features dict idx -> {"attn": ..., "mlp": ...}: pre-gate features are stored here when given (TaylorSeer's recording step). """ heads = block.heads_num q, k, v, g1, s2, sc2, g2 = block.modulate_img(vec, img) head_dim = q.shape[-1] if sel is None: if bank is not None: bank[idx] = [k, v] k_all, v_all = k, v cam_q, K_q = viewmats, Ks cam_kv, K_kv = viewmats, Ks cos, sin = freqs_cis cos_q, sin_q = cos, sin else: bank_k, bank_v = bank[idx] bank_k[:, sel] = k bank_v[:, sel] = v k_all, v_all = bank_k, bank_v cam_q, K_q = viewmats, Ks # already the subset's cameras cam_kv, K_kv = bank["viewmats"], bank["Ks"] # every token's camera cos, sin = bank["freqs_cis"] cos_q, sin_q = freqs_cis fq, _, fo = _prope_fns(head_dim, cam_q, K_q) _, fkv, _ = _prope_fns(head_dim, cam_kv, K_kv) q_p = _bhld(fq(_bhld(q))) k_p = _bhld(fkv(_bhld(k_all))) v_p = _bhld(fkv(_bhld(v_all))) q_r, _ = apply_rotary_emb(q, q, (cos_q, sin_q), head_first=False) _, k_r = apply_rotary_emb(k_all, k_all, (cos, sin), head_first=False) attn, attn_p, _ = sequence_parallel_attention_vision( (q_r, q_p), (k_r, k_p), (v_all, v_p), block_idx=idx, kv_cache=kv_cache, cache_vision=False) attn_p = rearrange(attn_p, "B L (H D) -> B H L D", H=heads) attn_p = rearrange(fo(attn_p), "B H L D -> B L (H D)") attn_feat = block.img_attn_proj(attn) + block.img_attn_prope_proj(attn_p) img = img + apply_gate(attn_feat, gate=g1) mlp_feat = block.img_mlp(modulate(block.img_norm2(img), shift=s2, scale=sc2)) img = img + apply_gate(mlp_feat, gate=g2) if features is not None: features[idx] = {"attn": attn_feat, "mlp": mlp_feat} return img def block_gates(block, vec): """The two gates a forecast step needs (same chunking as modulate_img).""" _, _, g1, _, _, g2 = block.img_mod(vec).chunk(6, dim=-1) return g1, g2 # -- fused forecast adds (TaylorSeer's cached step) ------------------------------ def _taylor_add_o0(img, attn0, mlp0, g1, g2): img = img + apply_gate(attn0, gate=g1) return img + apply_gate(mlp0, gate=g2) def _taylor_add_o1(img, attn0, attn1, mlp0, mlp1, g1, g2, d): img = img + apply_gate(attn0 + d * attn1, gate=g1) return img + apply_gate(mlp0 + d * mlp1, gate=g2) if os.environ.get("HYCACHE_NO_COMPILE", "0") != "1": taylor_add_o0 = torch.compile(_taylor_add_o0, dynamic=False) taylor_add_o1 = torch.compile(_taylor_add_o1, dynamic=False) else: taylor_add_o0, taylor_add_o1 = _taylor_add_o0, _taylor_add_o1