| 189 | |
| 190 | |
| 191 | class SingleStreamAttention(nn.Module): |
| 192 | def __init__( |
| 193 | self, |
| 194 | dim: int, |
| 195 | encoder_hidden_states_dim: int, |
| 196 | num_heads: int, |
| 197 | qkv_bias: bool, |
| 198 | qk_norm: bool, |
| 199 | norm_layer: nn.Module, |
| 200 | attn_drop: float = 0.0, |
| 201 | proj_drop: float = 0.0, |
| 202 | eps: float = 1e-6, |
| 203 | ) -> None: |
| 204 | super().__init__() |
| 205 | assert dim % num_heads == 0, "dim should be divisible by num_heads" |
| 206 | self.dim = dim |
| 207 | self.encoder_hidden_states_dim = encoder_hidden_states_dim |
| 208 | self.num_heads = num_heads |
| 209 | self.head_dim = dim // num_heads |
| 210 | self.scale = self.head_dim**-0.5 |
| 211 | self.qk_norm = qk_norm |
| 212 | |
| 213 | self.q_linear = nn.Linear(dim, dim, bias=qkv_bias) |
| 214 | |
| 215 | self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity() |
| 216 | self.k_norm = norm_layer(self.head_dim,eps=eps) if qk_norm else nn.Identity() |
| 217 | |
| 218 | self.attn_drop = nn.Dropout(attn_drop) |
| 219 | self.proj = nn.Linear(dim, dim) |
| 220 | self.proj_drop = nn.Dropout(proj_drop) |
| 221 | |
| 222 | self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias) |
| 223 | |
| 224 | self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() |
| 225 | self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() |
| 226 | |
| 227 | def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: |
| 228 | |
| 229 | N_t, N_h, N_w = shape |
| 230 | if not enable_sp: |
| 231 | x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) |
| 232 | |
| 233 | # get q for hidden_state |
| 234 | B, N, C = x.shape |
| 235 | q = self.q_linear(x) |
| 236 | q_shape = (B, N, self.num_heads, self.head_dim) |
| 237 | q = q.view(q_shape).permute((0, 2, 1, 3)) |
| 238 | |
| 239 | if self.qk_norm: |
| 240 | q = self.q_norm(q) |
| 241 | |
| 242 | # get kv from encoder_hidden_states |
| 243 | _, N_a, _ = encoder_hidden_states.shape |
| 244 | encoder_kv = self.kv_linear(encoder_hidden_states) |
| 245 | encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim) |
| 246 | encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4)) |
| 247 | encoder_k, encoder_v = encoder_kv.unbind(0) |
| 248 |
nothing calls this directly
no outgoing calls
no test coverage detected