""" 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