| 195 | |
| 196 | |
| 197 | class SwinTransformer(nn.Module): |
| 198 | def __init__(self, *, hidden_dim, layers, heads, channels=3, num_classes=1000, head_dim=32, window_size=7, |
| 199 | downscaling_factors=(4, 2, 2, 2), relative_pos_embedding=True): |
| 200 | super().__init__() |
| 201 | |
| 202 | self.stage1 = StageModule(in_channels=channels, hidden_dimension=hidden_dim, layers=layers[0], |
| 203 | downscaling_factor=downscaling_factors[0], num_heads=heads[0], head_dim=head_dim, |
| 204 | window_size=window_size, relative_pos_embedding=relative_pos_embedding) |
| 205 | self.stage2 = StageModule(in_channels=hidden_dim, hidden_dimension=hidden_dim * 2, layers=layers[1], |
| 206 | downscaling_factor=downscaling_factors[1], num_heads=heads[1], head_dim=head_dim, |
| 207 | window_size=window_size, relative_pos_embedding=relative_pos_embedding) |
| 208 | self.stage3 = StageModule(in_channels=hidden_dim * 2, hidden_dimension=hidden_dim * 4, layers=layers[2], |
| 209 | downscaling_factor=downscaling_factors[2], num_heads=heads[2], head_dim=head_dim, |
| 210 | window_size=window_size, relative_pos_embedding=relative_pos_embedding) |
| 211 | self.stage4 = StageModule(in_channels=hidden_dim * 4, hidden_dimension=hidden_dim * 8, layers=layers[3], |
| 212 | downscaling_factor=downscaling_factors[3], num_heads=heads[3], head_dim=head_dim, |
| 213 | window_size=window_size, relative_pos_embedding=relative_pos_embedding) |
| 214 | |
| 215 | self.mlp_head = nn.Sequential( |
| 216 | nn.LayerNorm(hidden_dim * 8), |
| 217 | nn.Linear(hidden_dim * 8, num_classes) |
| 218 | ) |
| 219 | |
| 220 | def forward(self, img): |
| 221 | x = self.stage1(img) |
| 222 | x = self.stage2(x) |
| 223 | x = self.stage3(x) |
| 224 | x = self.stage4(x) |
| 225 | x = x.mean(dim=[2, 3]) |
| 226 | return self.mlp_head(x) |
| 227 | |
| 228 | |
| 229 | def swin_t(num_classes:int, hidden_dim=96, layers=(2, 2, 6, 2), heads=(3, 6, 12, 24), **kwargs): |