DeMemWM / algorithms /dememwm /models /attention.py
BonanDing's picture
Make DeMemWM memory streams self-only
047d060
Raw
History Blame Contribute Delete
7.29 kB
"""
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