MCPcopy Create free account
hub / github.com/ChunmingHe/WS-SAM / __init__

Method __init__

lib/Slot.py:158–212  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

156class 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 = []

Callers

nothing calls this directly

Calls 4

make_encoderMethod · 0.95
SoftPositionEmbedClass · 0.85
SlotAttentionClass · 0.70
__init__Method · 0.45

Tested by

no test coverage detected