MCPcopy Create free account
hub / github.com/microsoft/Cream / __init__

Method __init__

AutoFormerV2/model/SSS.py:175–224  ·  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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 5

WindowAttentionClass · 0.70
MlpClass · 0.70
window_partitionFunction · 0.70
DropPathClass · 0.50
__init__Method · 0.45

Tested by

no test coverage detected