MCPcopy Create free account
hub / github.com/SwinTransformer/Transformer-SSL / __init__

Method __init__

models/swin_transformer.py:183–238  ·  view source on GitHub ↗
(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,
                 mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
                 act_layer=nn.GELU, norm_layer=nn.LayerNorm, norm_before_mlp='ln')

Source from the content-addressed store, hash-verified

181 """
182
183 def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,
184 mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
185 act_layer=nn.GELU, norm_layer=nn.LayerNorm, norm_before_mlp='ln'):
186 super().__init__()
187 self.dim = dim
188 self.input_resolution = input_resolution
189 self.num_heads = num_heads
190 self.window_size = window_size
191 self.shift_size = shift_size
192 self.mlp_ratio = mlp_ratio
193 self.norm_before_mlp = norm_before_mlp
194 if min(self.input_resolution) <= self.window_size:
195 # if window size is larger than input resolution, we don't partition windows
196 self.shift_size = 0
197 self.window_size = min(self.input_resolution)
198 assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
199
200 self.norm1 = norm_layer(dim)
201 self.attn = WindowAttention(
202 dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,
203 qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
204
205 self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
206 if self.norm_before_mlp == 'ln':
207 self.norm2 = nn.LayerNorm(dim)
208 elif self.norm_before_mlp == 'bn':
209 self.norm2 = lambda x: nn.BatchNorm1d(dim)(x.transpose(1, 2)).transpose(1, 2)
210 else:
211 raise NotImplementedError
212 mlp_hidden_dim = int(dim * mlp_ratio)
213 self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
214
215 if self.shift_size > 0:
216 # calculate attention mask for SW-MSA
217 H, W = self.input_resolution
218 img_mask = torch.zeros((1, H, W, 1)) # 1 H W 1
219 h_slices = (slice(0, -self.window_size),
220 slice(-self.window_size, -self.shift_size),
221 slice(-self.shift_size, None))
222 w_slices = (slice(0, -self.window_size),
223 slice(-self.window_size, -self.shift_size),
224 slice(-self.shift_size, None))
225 cnt = 0
226 for h in h_slices:
227 for w in w_slices:
228 img_mask[:, h, w, :] = cnt
229 cnt += 1
230
231 mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
232 mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
233 attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
234 attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
235 else:
236 attn_mask = None
237
238 self.register_buffer("attn_mask", attn_mask)
239
240 def forward(self, x):

Callers

nothing calls this directly

Calls 4

WindowAttentionClass · 0.85
MlpClass · 0.85
window_partitionFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected