| 9 | |
| 10 | |
| 11 | class Attention(nn.Module): |
| 12 | |
| 13 | def __init__( |
| 14 | self, |
| 15 | dim, |
| 16 | num_heads=8, |
| 17 | qkv_bias=False, |
| 18 | qk_scale=None, |
| 19 | attn_drop=0.0, |
| 20 | proj_drop=0.0, |
| 21 | ): |
| 22 | super().__init__() |
| 23 | self.num_heads = num_heads |
| 24 | head_dim = dim // num_heads |
| 25 | self.scale = qk_scale or head_dim**-0.5 |
| 26 | |
| 27 | self.q = nn.Linear(dim, dim, bias=qkv_bias) |
| 28 | self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias) |
| 29 | self.attn_drop = nn.Dropout(attn_drop) |
| 30 | self.proj = nn.Linear(dim, dim) |
| 31 | self.proj_drop = nn.Dropout(proj_drop) |
| 32 | |
| 33 | def forward(self, q, kv, key_mask=None): |
| 34 | N, C = kv.shape[1:] |
| 35 | QN = q.shape[1] |
| 36 | q = self.q(q).reshape([-1, QN, self.num_heads, |
| 37 | C // self.num_heads]).transpose(1, 2) |
| 38 | q = q * self.scale |
| 39 | k, v = self.kv(kv).reshape( |
| 40 | [-1, N, 2, self.num_heads, |
| 41 | C // self.num_heads]).permute(2, 0, 3, 1, 4) |
| 42 | |
| 43 | attn = q.matmul(k.transpose(2, 3)) |
| 44 | |
| 45 | if key_mask is not None: |
| 46 | attn = attn + key_mask.unsqueeze(1) |
| 47 | |
| 48 | attn = F.softmax(attn, -1) |
| 49 | # if not self.training: |
| 50 | # self.attn_map = attn |
| 51 | attn = self.attn_drop(attn) |
| 52 | |
| 53 | x = (attn.matmul(v)).transpose(1, 2).reshape((-1, QN, C)) |
| 54 | x = self.proj(x) |
| 55 | x = self.proj_drop(x) |
| 56 | return x |
| 57 | |
| 58 | |
| 59 | class EdgeDecoderLayer(nn.Module): |