| 303 | """ |
| 304 | |
| 305 | def __init__(self, dim, input_resolution, num_heads, window_size=7, |
| 306 | mlp_ratio=4., drop=0., drop_path=0., |
| 307 | local_conv_size=3, |
| 308 | activation=nn.GELU, |
| 309 | ): |
| 310 | super().__init__() |
| 311 | self.dim = dim |
| 312 | self.input_resolution = input_resolution |
| 313 | self.num_heads = num_heads |
| 314 | assert window_size > 0, 'window_size must be greater than 0' |
| 315 | self.window_size = window_size |
| 316 | self.mlp_ratio = mlp_ratio |
| 317 | |
| 318 | self.drop_path = DropPath( |
| 319 | drop_path) if drop_path > 0. else nn.Identity() |
| 320 | |
| 321 | assert dim % num_heads == 0, 'dim must be divisible by num_heads' |
| 322 | head_dim = dim // num_heads |
| 323 | |
| 324 | window_resolution = (window_size, window_size) |
| 325 | self.attn = Attention(dim, head_dim, num_heads, |
| 326 | attn_ratio=1, resolution=window_resolution) |
| 327 | |
| 328 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 329 | mlp_activation = activation |
| 330 | self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, |
| 331 | act_layer=mlp_activation, drop=drop) |
| 332 | |
| 333 | pad = local_conv_size // 2 |
| 334 | self.local_conv = Conv2d_BN( |
| 335 | dim, dim, ks=local_conv_size, stride=1, pad=pad, groups=dim) |
| 336 | |
| 337 | def forward(self, x): |
| 338 | H, W = self.input_resolution |