(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)
| 173 | """ |
| 174 | |
| 175 | def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0, |
| 176 | mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0., |
| 177 | act_layer=nn.GELU, norm_layer=nn.LayerNorm): |
| 178 | super().__init__() |
| 179 | self.dim = dim |
| 180 | self.input_resolution = input_resolution |
| 181 | self.num_heads = num_heads |
| 182 | self.window_size = window_size |
| 183 | self.shift_size = shift_size |
| 184 | self.mlp_ratio = mlp_ratio |
| 185 | if min(self.input_resolution) <= self.window_size: |
| 186 | # if window size is larger than input resolution, we don't partition windows |
| 187 | self.shift_size = 0 |
| 188 | self.window_size = min(self.input_resolution) |
| 189 | assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size" |
| 190 | |
| 191 | self.norm1 = norm_layer(dim) |
| 192 | self.attn = WindowAttention( |
| 193 | dim, window_size=to_2tuple(self.window_size), num_heads=num_heads, |
| 194 | qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop) |
| 195 | |
| 196 | self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() |
| 197 | self.norm2 = norm_layer(dim) |
| 198 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 199 | self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
| 200 | |
| 201 | if self.shift_size > 0: |
| 202 | # calculate attention mask for SW-MSA |
| 203 | H, W = self.input_resolution |
| 204 | img_mask = torch.zeros((1, H, W, 1)) # 1 H W 1 |
| 205 | h_slices = (slice(0, -self.window_size), |
| 206 | slice(-self.window_size, -self.shift_size), |
| 207 | slice(-self.shift_size, None)) |
| 208 | w_slices = (slice(0, -self.window_size), |
| 209 | slice(-self.window_size, -self.shift_size), |
| 210 | slice(-self.shift_size, None)) |
| 211 | cnt = 0 |
| 212 | for h in h_slices: |
| 213 | for w in w_slices: |
| 214 | img_mask[:, h, w, :] = cnt |
| 215 | cnt += 1 |
| 216 | |
| 217 | mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1 |
| 218 | mask_windows = mask_windows.view(-1, self.window_size * self.window_size) |
| 219 | attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2) |
| 220 | attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)) |
| 221 | else: |
| 222 | attn_mask = None |
| 223 | |
| 224 | self.register_buffer("attn_mask", attn_mask) |
| 225 | |
| 226 | def forward(self, x): |
| 227 | H, W = self.input_resolution |
nothing calls this directly
no test coverage detected