| 37 | |
| 38 | |
| 39 | class Attention(nn.Module): |
| 40 | def __init__( |
| 41 | self, |
| 42 | dim: int, |
| 43 | num_heads: int = 8, |
| 44 | qkv_bias: bool = False, |
| 45 | qk_norm: bool = False, |
| 46 | attn_drop: float = 0., |
| 47 | proj_drop: float = 0., |
| 48 | norm_layer: nn.Module = nn.LayerNorm, |
| 49 | ) -> None: |
| 50 | |
| 51 | super().__init__() |
| 52 | assert dim % num_heads == 0, 'dim should be divisible by num_heads' |
| 53 | self.num_heads = num_heads |
| 54 | self.head_dim = dim // num_heads |
| 55 | self.scale = self.head_dim ** -0.5 |
| 56 | self.fused_attn = use_fused_attn() |
| 57 | |
| 58 | self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) |
| 59 | self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() |
| 60 | self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() |
| 61 | self.attn_drop = nn.Dropout(attn_drop) |
| 62 | self.proj = nn.Linear(dim, dim) |
| 63 | self.proj_drop = nn.Dropout(proj_drop) |
| 64 | |
| 65 | def forward(self, x: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor: |
| 66 | B, N, C = x.shape |
| 67 | qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) |
| 68 | q, k, v = qkv.unbind(0) |
| 69 | q, k = self.q_norm(q), self.k_norm(k) |
| 70 | |
| 71 | if self.fused_attn: |
| 72 | x = F.scaled_dot_product_attention( |
| 73 | q, k, v, |
| 74 | attn_mask = ~mask, |
| 75 | dropout_p=self.attn_drop.p if self.training else 0., |
| 76 | ) |
| 77 | else: |
| 78 | q = q * self.scale |
| 79 | attn = q @ k.transpose(-2, -1) |
| 80 | attn = attn.softmax(dim=-1) |
| 81 | attn = self.attn_drop(attn) |
| 82 | x = attn @ v |
| 83 | |
| 84 | x = x.transpose(1, 2).reshape(B, N, C) |
| 85 | x = self.proj(x) |
| 86 | x = self.proj_drop(x) |
| 87 | return x |
| 88 | |
| 89 | class TimestepEmbedder(nn.Module): |
| 90 | """ |