comparison / hycache /blocks.py
Cccccz's picture
Add files using upload-large-folder tool
b34c6c3 verified
Raw History Blame Contribute Delete
4.79 kB
"""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