MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / WindowAttention

Class WindowAttention

Image_Classification/src/models/swin.py:72–136  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

70
71
72class 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)

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected