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