| 6 | |
| 7 | class MultiHeadAttention(nn.Module): |
| 8 | def __init__(self, embed_dim, num_heads, dropout=0.1, bias=True): |
| 9 | super().__init__() |
| 10 | assert embed_dim % num_heads == 0 |
| 11 | |
| 12 | self.embed_dim = embed_dim |
| 13 | self.num_heads = num_heads |
| 14 | self.head_dim = embed_dim // num_heads |
| 15 | self.scale = self.head_dim ** -0.5 |
| 16 | |
| 17 | self.qkv = nn.Linear(embed_dim, embed_dim * 3, bias=bias) |
| 18 | self.proj = nn.Linear(embed_dim, embed_dim, bias=bias) |
| 19 | self.dropout = nn.Dropout(dropout) |
| 20 | |
| 21 | def forward(self, x, mask=None): |
| 22 | B, N, C = x.shape |