| 193 | """ |
| 194 | |
| 195 | def __init__(self, dim, window_size, num_heads, qkv_bias=True, rel_pos_spatial=False): |
| 196 | super().__init__() |
| 197 | self.dim = dim |
| 198 | self.window_size = window_size # Wh, Ww |
| 199 | self.num_heads = num_heads |
| 200 | head_dim = dim // num_heads |
| 201 | self.scale = head_dim ** -0.5 |
| 202 | self.rel_pos_spatial=rel_pos_spatial |
| 203 | |
| 204 | if COMPAT: |
| 205 | q_size = window_size[0] |
| 206 | kv_size = window_size[1] |
| 207 | rel_sp_dim = 2 * q_size - 1 |
| 208 | self.rel_pos_h = nn.Parameter(torch.zeros(rel_sp_dim, head_dim)) |
| 209 | self.rel_pos_w = nn.Parameter(torch.zeros(rel_sp_dim, head_dim)) |
| 210 | |
| 211 | self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) |
| 212 | self.proj = nn.Linear(dim, dim) |
| 213 | |
| 214 | def forward(self, x, H, W): |
| 215 | """ Forward function. |