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