| 69 | |
| 70 | |
| 71 | class Attention(nn.Module): |
| 72 | def __init__( |
| 73 | self, dim, num_heads=8, qkv_bias=False, window_size=None, rel_pos_spatial=False): |
| 74 | super().__init__() |
| 75 | self.num_heads = num_heads |
| 76 | head_dim = dim // num_heads |
| 77 | self.scale = head_dim ** -0.5 |
| 78 | self.rel_pos_spatial = rel_pos_spatial |
| 79 | self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) |
| 80 | self.window_size = window_size |
| 81 | if COMPAT: |
| 82 | if COMPAT == 2: |
| 83 | self.rel_pos_h = nn.Parameter(torch.zeros(2 * window_size[0] - 1, head_dim)) |
| 84 | self.rel_pos_w = nn.Parameter(torch.zeros(2 * window_size[1] - 1, head_dim)) |
| 85 | else: |
| 86 | q_size = window_size[0] |
| 87 | kv_size = q_size |
| 88 | rel_sp_dim = 2 * q_size - 1 |
| 89 | self.rel_pos_h = nn.Parameter(torch.zeros(rel_sp_dim, head_dim)) |
| 90 | self.rel_pos_w = nn.Parameter(torch.zeros(rel_sp_dim, head_dim)) |
| 91 | self.proj = nn.Linear(dim, dim) |
| 92 | |
| 93 | def forward(self, x, H, W): |
| 94 | B, N, C = x.shape |
| 95 | qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) |
| 96 | q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple) |
| 97 | |
| 98 | attn = ((q * self.scale) @ k.transpose(-2, -1)) |
| 99 | if self.rel_pos_spatial: |
| 100 | raise |
| 101 | attn = calc_rel_pos_spatial(attn, q, self.window_size, self.window_size, self.rel_pos_h, self.rel_pos_w) |
| 102 | |
| 103 | attn = attn.softmax(dim=-1) |
| 104 | |
| 105 | x = (attn @ v).transpose(1, 2).reshape(B, N, C) |
| 106 | x = self.proj(x) |
| 107 | return x |
| 108 | |
| 109 | |
| 110 | def window_partition(x, window_size): |