File size: 6,239 Bytes
ec0a9aa | 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 | """
Radial Attention Processor for Ctrl-World SVD UNet.
Replaces standard SDPA with Radial Attention, using block sparse
attention kernels (FlashInfer Backend) for large spatial layers.
"""
from __future__ import annotations
import warnings
from typing import Optional
import torch
import torch.nn.functional as F
# Additive import from the local radial_attn package (added to sys.path by adapter.py)
try:
from radial_attn.attn_mask import RadialAttention, MaskMap
except ImportError:
RadialAttention = None
MaskMap = None
def _split_heads(x: torch.Tensor, num_heads: int) -> torch.Tensor:
"""(B, T, D) → (B, H, T, D//H)"""
B, T, D = x.shape
return x.reshape(B, T, num_heads, D // num_heads).permute(0, 2, 1, 3).contiguous()
class SVDRadialAttnProcessor:
"""
Drop-in replacement for diffusers AttnProcessor2_0, powered by Radial Attention.
"""
# ── shared class-level state (set by adapter.py) ──────────────
decay_factor: float = 1.0
block_size: int = 64
pad_small_layers: bool = True
first_layers_fp: int = 2
MIN_TOKENS: int = 128 # skip temporal (T=11) and mid (T=45)
# Geometry attrs
num_views: int = 3
H_per_view: int = 24
W: int = 40
def __init__(self, layer_idx: int, num_layers: int):
self.layer_idx = layer_idx
self.num_layers = num_layers
# Cache for MaskMap per padded sequence length
self._mask_map_cache = {}
def __call__(
self,
attn,
hidden_states: torch.Tensor, # (B, T, D)
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
**kwargs,
) -> torch.Tensor:
B, T_q, D = hidden_states.shape
H = attn.heads
# ── project Q, K, V ────────────────────────────────────────────────
q = attn.to_q(hidden_states)
kv_src = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
k = attn.to_k(kv_src)
v = attn.to_v(kv_src)
T_kv = kv_src.shape[1]
q = _split_heads(q, H) # (B, H, T_q, D_head)
k = _split_heads(k, H)
v = _split_heads(v, H)
# ── decide: full vs sparse ──────────────────────────────────────────
is_self_attn = (T_q == T_kv)
long_enough = (T_q >= self.MIN_TOKENS)
early_layer = (self.layer_idx < self.first_layers_fp)
use_sparse = is_self_attn and long_enough and (not early_layer) and (RadialAttention is not None)
if use_sparse:
hidden_states = self._radial_attention(q, k, v, T_q)
# Radial attention returns (B, T_q, H * D_head)
else:
hidden_states = F.scaled_dot_product_attention(
q, k, v, attn_mask=attention_mask, dropout_p=0.0
) # (B, H, T_q, D_head)
hidden_states = hidden_states.permute(0, 2, 1, 3).contiguous().reshape(B, T_q, -1)
# ── output proj ─────────────────────────────────────────────────────
hidden_states = hidden_states.to(q.dtype)
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
return hidden_states
# ── Radial Attention logic ──────────────────────────────────────────────
def _radial_attention(
self,
q: torch.Tensor, # (B, H, T, D_h)
k: torch.Tensor,
v: torch.Tensor,
T: int,
) -> torch.Tensor:
divisible = (T % self.block_size == 0)
pad_len = 0
if not divisible:
if not self.pad_small_layers:
# Option 1: Fallback to SDPA for non-divisible layers
out = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0)
B, H, _, D_h = out.shape
return out.permute(0, 2, 1, 3).contiguous().reshape(B, T, -1)
else:
# Option 2: Pad to nearest multiple of block_size
pad_len = self.block_size - (T % self.block_size)
# Pad T dimension: (pad_last_dim_left, pad_last_dim_right, pad_2nd_last_dim_left, pad_2nd_last_dim_right)
q = F.pad(q, (0, 0, 0, pad_len))
k = F.pad(k, (0, 0, 0, pad_len))
v = F.pad(v, (0, 0, 0, pad_len))
padded_T = T + pad_len
if padded_T not in self._mask_map_cache:
# Recreate mask map for this sequence length
self._mask_map_cache[padded_T] = MaskMap(video_token_num=padded_T, num_frame=self.num_views)
mask_map = self._mask_map_cache[padded_T]
# RadialAttention expects (batch, seq_len, heads, dim)
q = q.transpose(1, 2).contiguous()
k = k.transpose(1, 2).contiguous()
v = v.transpose(1, 2).contiguous()
# Extract boolean mask
video_mask = mask_map.queryLogMask(q, "radial", block_size=self.block_size, decay_factor=self.decay_factor, model_type=self.model_type)
video_mask = video_mask[:padded_T // self.block_size, :padded_T // self.block_size]
# Get flashinfer wrapper for this geometry
bsr_wrapper = mask_map.get_bsr_wrapper(video_mask, q, k, self.block_size)
# Process each batch item (Fixes a bug in the native RadialAttention implementation where batch elements are dropped)
B = q.shape[0]
out_list = []
for i in range(B):
o_i = bsr_wrapper.run(q[i], k[i], v[i]) # (padded_T, H, D_h)
out_list.append(o_i)
out_padded = torch.stack(out_list, dim=0).flatten(2, 3) # (B, padded_T, H * D_h)
# Unpad
if pad_len > 0:
out = out_padded[:, :T, :]
else:
out = out_padded
return out
|