| 70 | |
| 71 | |
| 72 | class WindowAttention(nn.Module): |
| 73 | def __init__(self, dim, heads, head_dim, shifted, window_size, relative_pos_embedding): |
| 74 | super().__init__() |
| 75 | inner_dim = head_dim * heads |
| 76 | |
| 77 | self.heads = heads |
| 78 | self.scale = head_dim ** -0.5 |
| 79 | self.window_size = window_size |
| 80 | self.relative_pos_embedding = relative_pos_embedding |
| 81 | self.shifted = shifted |
| 82 | |
| 83 | if self.shifted: |
| 84 | displacement = window_size // 2 |
| 85 | self.cyclic_shift = CyclicShift(-displacement) |
| 86 | self.cyclic_back_shift = CyclicShift(displacement) |
| 87 | self.upper_lower_mask = nn.Parameter(create_mask(window_size=window_size, displacement=displacement, |
| 88 | upper_lower=True, left_right=False), requires_grad=False) |
| 89 | self.left_right_mask = nn.Parameter(create_mask(window_size=window_size, displacement=displacement, |
| 90 | upper_lower=False, left_right=True), requires_grad=False) |
| 91 | |
| 92 | self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False) |
| 93 | |
| 94 | if self.relative_pos_embedding: |
| 95 | self.relative_indices = get_relative_distances(window_size) + window_size - 1 |
| 96 | self.pos_embedding = nn.Parameter(torch.randn(2 * window_size - 1, 2 * window_size - 1)) |
| 97 | else: |
| 98 | self.pos_embedding = nn.Parameter(torch.randn(window_size ** 2, window_size ** 2)) |
| 99 | |
| 100 | self.to_out = nn.Linear(inner_dim, dim) |
| 101 | |
| 102 | def forward(self, x): |
| 103 | if self.shifted: |
| 104 | x = self.cyclic_shift(x) |
| 105 | |
| 106 | b, n_h, n_w, _, h = *x.shape, self.heads |
| 107 | |
| 108 | qkv = self.to_qkv(x).chunk(3, dim=-1) |
| 109 | nw_h = n_h // self.window_size |
| 110 | nw_w = n_w // self.window_size |
| 111 | |
| 112 | q, k, v = map( |
| 113 | lambda t: rearrange(t, 'b (nw_h w_h) (nw_w w_w) (h d) -> b h (nw_h nw_w) (w_h w_w) d', |
| 114 | h=h, w_h=self.window_size, w_w=self.window_size), qkv) |
| 115 | |
| 116 | dots = einsum('b h w i d, b h w j d -> b h w i j', q, k) * self.scale |
| 117 | |
| 118 | if self.relative_pos_embedding: |
| 119 | dots += self.pos_embedding[self.relative_indices[:, :, 0], self.relative_indices[:, :, 1]] |
| 120 | else: |
| 121 | dots += self.pos_embedding |
| 122 | |
| 123 | if self.shifted: |
| 124 | dots[:, :, -nw_w:] += self.upper_lower_mask |
| 125 | dots[:, :, nw_w - 1::nw_w] += self.left_right_mask |
| 126 | |
| 127 | attn = dots.softmax(dim=-1) |
| 128 | |
| 129 | out = einsum('b h w i j, b h w j d -> b h w i d', attn, v) |