File size: 7,286 Bytes
ea5e8ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3df329c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
047d060
 
3df329c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ea5e8ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3df329c
 
 
 
 
 
 
 
 
 
 
 
ea5e8ea
 
 
 
 
 
 
 
 
 
 
1470930
ea5e8ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
"""
Based on https://github.com/buoyancy99/diffusion-forcing/blob/main/algorithms/diffusion_forcing/models/attention.py
"""

from typing import Optional
import torch
from torch import nn
from torch.nn import functional as F
from einops import rearrange
from .rotary_embedding_torch import RotaryEmbedding, apply_rotary_emb

class TemporalAxialAttention(nn.Module):
    def __init__(
        self,
        dim: int,
        heads: int,
        dim_head: int,
        reference_length: int,
        rotary_emb: RotaryEmbedding,
        is_causal: bool = True,
        is_temporal_independent: bool = False,
        use_domain_adapter = False
    ):
        super().__init__()
        self.inner_dim = dim_head * heads
        self.heads = heads
        self.head_dim = dim_head
        self.inner_dim = dim_head * heads
        self.to_qkv = nn.Linear(dim, self.inner_dim * 3, bias=False)

        self.use_domain_adapter = use_domain_adapter
        if self.use_domain_adapter:
            lora_rank = 8
            self.lora_A = nn.Linear(dim, lora_rank, bias=False)
            self.lora_B = nn.Linear(lora_rank, self.inner_dim * 3, bias=False)

        self.to_out = nn.Linear(self.inner_dim, dim)

        self.rotary_emb = rotary_emb
        self.is_causal = is_causal
        self.is_temporal_independent = is_temporal_independent

        self.reference_length = reference_length

    def _frame_memory_attn_bias(
        self,
        B,
        T,
        H,
        W,
        dtype,
        device,
        frame_memory_segments,
        frame_memory_masks,
    ):
        allow = torch.zeros((B, T, T), dtype=torch.bool, device=device)

        target_frames = frame_memory_segments["target"] if frame_memory_segments else T
        cursor = 0
        target_slice = slice(cursor, cursor + target_frames)
        target_idx = torch.arange(target_frames, device=device)
        allow[:, target_slice, target_slice] = target_idx[:, None] >= target_idx[None, :]
        cursor += target_frames

        segment_slices = {"target": target_slice}
        for segment in ("anchor", "dynamic", "revisit"):
            length = frame_memory_segments.get(segment, 0)
            segment_slice = slice(cursor, cursor + length)
            segment_slices[segment] = segment_slice
            if length > 0:
                idx = torch.arange(cursor, cursor + length, device=device)
                allow[:, idx, idx] = True
            cursor += length

        if frame_memory_masks is not None:
            valid = torch.ones((B, T), dtype=torch.bool, device=device)
            for segment, segment_slice in segment_slices.items():
                mask = frame_memory_masks.get(segment)
                if mask is not None:
                    valid[:, segment_slice] = mask.to(device=device, dtype=torch.bool)
            allow = allow & valid[:, :, None] & valid[:, None, :]

        # Invalid padded rows keep only self-attention finite so SDPA never sees
        # an all -inf query row.
        diag_idx = torch.arange(T, device=device)
        allow[:, diag_idx, diag_idx] = True
        attn_bias = torch.zeros((B, T, T), dtype=dtype, device=device)
        attn_bias = attn_bias.masked_fill(~allow, float("-inf"))
        return attn_bias.repeat_interleave(H * W, dim=0)[:, None]

    def forward(self, x: torch.Tensor, frame_memory_segments=None, frame_memory_masks=None):
        B, T, H, W, D = x.shape

        q, k, v = self.to_qkv(x).chunk(3, dim=-1)

        if self.use_domain_adapter:
            q_lora, k_lora, v_lora = self.lora_B(self.lora_A(x)).chunk(3, dim=-1)
            q = q+q_lora
            k = k+k_lora
            v = v+v_lora

        q = rearrange(q, "B T H W (h d) -> (B H W) h T d", h=self.heads)
        k = rearrange(k, "B T H W (h d) -> (B H W) h T d", h=self.heads)
        v = rearrange(v, "B T H W (h d) -> (B H W) h T d", h=self.heads)

        q = self.rotary_emb.rotate_queries_or_keys(q, self.rotary_emb.freqs)
        k = self.rotary_emb.rotate_queries_or_keys(k, self.rotary_emb.freqs)

        q, k, v = map(lambda t: t.contiguous(), (q, k, v))

        if frame_memory_segments is not None:
            attn_bias = self._frame_memory_attn_bias(
                B,
                T,
                H,
                W,
                q.dtype,
                q.device,
                frame_memory_segments,
                frame_memory_masks,
            )
        elif self.is_temporal_independent:
            attn_bias = torch.ones((T, T), dtype=q.dtype, device=q.device)
            attn_bias = attn_bias.masked_fill(attn_bias == 1, float('-inf'))
            attn_bias[range(T), range(T)] = 0
        elif self.is_causal:
            attn_bias = torch.triu(torch.ones((T, T), dtype=q.dtype, device=q.device), diagonal=1)
            attn_bias = attn_bias.masked_fill(attn_bias == 1, float('-inf'))
            attn_bias[(T-self.reference_length):] = float('-inf')
            attn_bias[range(T), range(T)] = 0
        else:
            attn_bias = None

        x = F.scaled_dot_product_attention(query=q, key=k, value=v, attn_mask=attn_bias)

        x = rearrange(x, "(B H W) h T d -> B T H W (h d)", B=B, H=H, W=W)
        x = x.to(q.dtype)

        # linear proj
        x = self.to_out(x)
        return x

class SpatialAxialAttention(nn.Module):
    def __init__(
        self,
        dim: int,
        heads: int,
        dim_head: int,
        rotary_emb: RotaryEmbedding,
        use_domain_adapter = False
    ):
        super().__init__()
        self.inner_dim = dim_head * heads
        self.heads = heads
        self.head_dim = dim_head
        self.inner_dim = dim_head * heads
        self.to_qkv = nn.Linear(dim, self.inner_dim * 3, bias=False)
        self.use_domain_adapter = use_domain_adapter
        if self.use_domain_adapter:
            lora_rank = 8
            self.lora_A = nn.Linear(dim, lora_rank, bias=False)
            self.lora_B = nn.Linear(lora_rank, self.inner_dim * 3, bias=False)

        self.to_out = nn.Linear(self.inner_dim, dim)

        self.rotary_emb = rotary_emb

    def forward(self, x: torch.Tensor):
        B, T, H, W, D = x.shape

        q, k, v = self.to_qkv(x).chunk(3, dim=-1)

        if self.use_domain_adapter:
            q_lora, k_lora, v_lora = self.lora_B(self.lora_A(x)).chunk(3, dim=-1)
            q = q+q_lora
            k = k+k_lora
            v = v+v_lora

        q = rearrange(q, "B T H W (h d) -> (B T) h H W d", h=self.heads)
        k = rearrange(k, "B T H W (h d) -> (B T) h H W d", h=self.heads)
        v = rearrange(v, "B T H W (h d) -> (B T) h H W d", h=self.heads)

        freqs = self.rotary_emb.get_axial_freqs(H, W)
        q = apply_rotary_emb(freqs, q)
        k = apply_rotary_emb(freqs, k)

        # prepare for attn
        q = rearrange(q, "(B T) h H W d -> (B T) h (H W) d", B=B, T=T, h=self.heads)
        k = rearrange(k, "(B T) h H W d -> (B T) h (H W) d", B=B, T=T, h=self.heads)
        v = rearrange(v, "(B T) h H W d -> (B T) h (H W) d", B=B, T=T, h=self.heads)

        x = F.scaled_dot_product_attention(query=q, key=k, value=v, is_causal=False)

        x = rearrange(x, "(B T) h (H W) d -> B T H W (h d)", B=B, H=H, W=W)
        x = x.to(q.dtype)

        # linear proj
        x = self.to_out(x)
        return x