MCPcopy Create free account
hub / github.com/YvanYin/DrivingWorld / Decoder

Class Decoder

modules/tokenizers/vq_model.py:207–256  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

205 return h
206
207class Decoder(nn.Module):
208 def __init__(self, z_channels=256, ch=128, ch_mult=(1,1,2,2,4), num_res_blocks=2, norm_type="group",
209 dropout=0.0, resamp_with_conv=True, out_channels=3):
210 super().__init__()
211 self.num_resolutions = len(ch_mult)
212 self.num_res_blocks = num_res_blocks
213 block_in = ch*ch_mult[self.num_resolutions-1]
214 self.conv_in = nn.Conv2d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
215 self.mid = nn.ModuleList()
216 self.mid.append(ResnetBlock(block_in, block_in, dropout=dropout, norm_type=norm_type))
217 self.mid.append(AttnBlock(block_in, norm_type=norm_type))
218 self.mid.append(ResnetBlock(block_in, block_in, dropout=dropout, norm_type=norm_type))
219 self.conv_blocks = nn.ModuleList()
220 for i_level in reversed(range(self.num_resolutions)):
221 conv_block = nn.Module()
222 res_block = nn.ModuleList()
223 attn_block = nn.ModuleList()
224 block_out = ch*ch_mult[i_level]
225 for _ in range(self.num_res_blocks + 1):
226 res_block.append(ResnetBlock(block_in, block_out, dropout=dropout, norm_type=norm_type))
227 block_in = block_out
228 if i_level == self.num_resolutions - 1:
229 attn_block.append(AttnBlock(block_in, norm_type))
230 conv_block.res = res_block
231 conv_block.attn = attn_block
232 if i_level != 0:
233 conv_block.upsample = Upsample(block_in, resamp_with_conv)
234 self.conv_blocks.append(conv_block)
235 self.norm_out = Normalize(block_in, norm_type)
236 self.conv_out = nn.Conv2d(block_in, out_channels, kernel_size=3, stride=1, padding=1)
237
238 @property
239 def last_layer(self):
240 return self.conv_out.weight
241
242 def forward(self, z):
243 h = self.conv_in(z)
244 for mid_block in self.mid:
245 h = mid_block(h)
246 for i_level, block in enumerate(self.conv_blocks):
247 for i_block in range(self.num_res_blocks + 1):
248 h = block.res[i_block](h)
249 if len(block.attn) > 0:
250 h = block.attn[i_block](h)
251 if i_level != self.num_resolutions - 1:
252 h = block.upsample(h)
253 h = self.norm_out(h)
254 h = nonlinearity(h)
255 h = self.conv_out(h)
256 return h
257
258
259class VectorQuantizer(nn.Module):

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected