(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')
| 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): |
nothing calls this directly
no test coverage detected