(self, encoder_dims, resolution, num_slots, iters)
| 409 | |
| 410 | class SlotAttentionModule2(nn.Module): |
| 411 | def __init__(self, encoder_dims, resolution, num_slots, iters): |
| 412 | super(SlotAttentionModule2, self).__init__() |
| 413 | self.resolution = resolution |
| 414 | self.num_slots = num_slots |
| 415 | self.encoder_pos = SoftPositionEmbed(encoder_dims, ((int(resolution), int(resolution)))) |
| 416 | self.layer_norm = nn.LayerNorm(encoder_dims) |
| 417 | self.mlp = nn.Sequential( |
| 418 | nn.Linear(encoder_dims, encoder_dims), |
| 419 | nn.ReLU(inplace=True), |
| 420 | nn.Linear(encoder_dims, encoder_dims) |
| 421 | ) |
| 422 | self.slot_attention = SlotAttention(iters=iters, |
| 423 | num_slots=num_slots, |
| 424 | encoder_dims=encoder_dims, |
| 425 | hidden_dim=encoder_dims) |
| 426 | self.decoder_pos = SoftPositionEmbed(encoder_dims, (int(resolution), int(resolution))) |
| 427 | |
| 428 | self.decoder_cnn = nn.Sequential( |
| 429 | nn.ConvTranspose2d(encoder_dims, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 430 | nn.InstanceNorm2d(64, affine=True), |
| 431 | nn.ReLU(inplace=True), |
| 432 | # |
| 433 | nn.ConvTranspose2d(64, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 434 | nn.InstanceNorm2d(64, affine=True), |
| 435 | nn.ReLU(inplace=True), |
| 436 | # |
| 437 | nn.ConvTranspose2d(64, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 438 | nn.InstanceNorm2d(64, affine=True), |
| 439 | nn.ReLU(inplace=True), |
| 440 | # |
| 441 | nn.Conv2d(64, 64, kernel_size=5, padding=2, stride=1), |
| 442 | nn.InstanceNorm2d(64, affine=True), |
| 443 | nn.ReLU(inplace=True), |
| 444 | nn.Conv2d(64, 3 + 1, kernel_size=5, padding=2, stride=1) |
| 445 | ) |
| 446 | self.convModules = nn.ModuleList([nn.Conv2d(encoder_dims, encoder_dims//num_slots, kernel_size=1, padding=0, stride=1) for i in range (num_slots)]) |
| 447 | |
| 448 | def forward(self, x): |
| 449 | batch = x.shape[0] |
nothing calls this directly
no test coverage detected