| 194 | return x |
| 195 | |
| 196 | class Attention(nn.Module): |
| 197 | def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.): |
| 198 | super().__init__() |
| 199 | self.num_heads = num_heads |
| 200 | head_dim = dim // num_heads |
| 201 | self.scale = qk_scale or head_dim ** -0.5 |
| 202 | |
| 203 | self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) |
| 204 | self.attn_drop = nn.Dropout(attn_drop) |
| 205 | self.proj = nn.Linear(dim, dim) |
| 206 | self.proj_drop = nn.Dropout(proj_drop) |
| 207 | |
| 208 | def forward(self, x): |
| 209 | B, N, C = x.shape |
| 210 | qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) |
| 211 | q, k, v = qkv[0], qkv[1], qkv[2] |
| 212 | |
| 213 | attn = (q @ k.transpose(-2, -1)) * self.scale |
| 214 | attn = attn.softmax(dim=-1) |
| 215 | attn = self.attn_drop(attn) |
| 216 | |
| 217 | x = (attn @ v).transpose(1, 2).reshape(B, N, C) |
| 218 | x = self.proj(x) |
| 219 | x = self.proj_drop(x) |
| 220 | return x, attn |
| 221 | |
| 222 | |
| 223 | class Block(nn.Module): |