Builds the Slot Attention-based Auto-encoder. Args: resolution: Tuple of integers specifying width and height of input image num_slots: Number of slots in Slot Attention. iters: Number of iterations in Slot Attention.
(self, resolution, num_slots, in_out_channels=3, iters=5)
| 156 | class SlotAttentionAutoEncoder(nn.Module): |
| 157 | """Slot Attention-based auto-encoder for object discovery.""" |
| 158 | def __init__(self, resolution, num_slots, in_out_channels=3, iters=5): |
| 159 | """Builds the Slot Attention-based Auto-encoder. |
| 160 | |
| 161 | Args: |
| 162 | resolution: Tuple of integers specifying width and height of input image |
| 163 | num_slots: Number of slots in Slot Attention. |
| 164 | iters: Number of iterations in Slot Attention. |
| 165 | """ |
| 166 | super(SlotAttentionAutoEncoder, self).__init__() |
| 167 | |
| 168 | self.iters = iters |
| 169 | self.num_slots = num_slots |
| 170 | self.resolution = resolution |
| 171 | self.in_out_channels = in_out_channels |
| 172 | |
| 173 | self.encoder_arch = [64, 'MP', 128, 'MP', 256] |
| 174 | self.encoder_dims = self.encoder_arch[-1] |
| 175 | self.encoder_cnn, ratio = self.make_encoder(self.in_out_channels, self.encoder_arch) |
| 176 | self.encoder_end_size = (int(resolution[0] / ratio), int(resolution[1] / ratio)) |
| 177 | self.encoder_pos = SoftPositionEmbed(self.encoder_dims, self.encoder_end_size) |
| 178 | self.decoder_initial_size = (int(resolution[0] / 8), int(resolution[1] / 8)) |
| 179 | self.decoder_pos = SoftPositionEmbed(self.encoder_dims, self.decoder_initial_size) |
| 180 | |
| 181 | self.layer_norm = nn.LayerNorm(self.encoder_dims) |
| 182 | |
| 183 | self.mlp = nn.Sequential( |
| 184 | nn.Linear(self.encoder_dims, self.encoder_dims), |
| 185 | nn.ReLU(inplace=True), |
| 186 | nn.Linear(self.encoder_dims, self.encoder_dims) |
| 187 | ) |
| 188 | |
| 189 | self.slot_attention = SlotAttention( |
| 190 | iters=self.iters, |
| 191 | num_slots=self.num_slots, |
| 192 | encoder_dims=self.encoder_dims, |
| 193 | hidden_dim=self.encoder_dims) |
| 194 | |
| 195 | self.decoder_cnn = nn.Sequential( |
| 196 | nn.ConvTranspose2d(self.encoder_dims, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 197 | nn.InstanceNorm2d(64, affine=True), |
| 198 | nn.ReLU(inplace=True), |
| 199 | # |
| 200 | nn.ConvTranspose2d(64, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 201 | nn.InstanceNorm2d(64, affine=True), |
| 202 | nn.ReLU(inplace=True), |
| 203 | # |
| 204 | nn.ConvTranspose2d(64, 64, kernel_size=5, padding=2, output_padding=1, stride=2), |
| 205 | nn.InstanceNorm2d(64, affine=True), |
| 206 | nn.ReLU(inplace=True), |
| 207 | # |
| 208 | nn.Conv2d(64, 64, kernel_size=5, padding=2, stride=1), |
| 209 | nn.InstanceNorm2d(64, affine=True), |
| 210 | nn.ReLU(inplace=True), |
| 211 | nn.Conv2d(64, in_out_channels + 1, kernel_size=5, padding=2, stride=1) |
| 212 | ) |
| 213 | |
| 214 | def make_encoder(self, in_channels, encoder_arch): |
| 215 | layers = [] |
nothing calls this directly
no test coverage detected