File size: 4,789 Bytes
b34c6c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""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