MCPcopy Create free account
hub / github.com/Little-Podi/AdaWorld / Decoder

Class Decoder

worldmodel/vwm/modules/diffusionmodules/model.py:358–484  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

356
357
358class Decoder(nn.Module):
359 def __init__(
360 self,
361 *,
362 ch,
363 out_channels,
364 ch_mult=(1, 2, 4, 8),
365 num_res_blocks,
366 attn_resolutions,
367 dropout=0.0,
368 resamp_with_conv=True,
369 in_channels,
370 resolution,
371 z_channels,
372 give_pre_end=False,
373 tanh_out=False,
374 use_linear_attn=False,
375 attn_type="vanilla",
376 **ignorekwargs
377 ):
378 super(Decoder, self).__init__()
379 if use_linear_attn:
380 attn_type = "linear"
381 self.ch = ch
382 self.temb_ch = 0
383 self.num_resolutions = len(ch_mult)
384 self.num_res_blocks = num_res_blocks
385 self.resolution = resolution
386 self.in_channels = in_channels
387 self.give_pre_end = give_pre_end
388 self.tanh_out = tanh_out
389
390 # Compute in_ch_mult, block_in and curr_res at lowest res
391 in_ch_mult = (1,) + tuple(ch_mult)
392 block_in = ch * ch_mult[self.num_resolutions - 1]
393 curr_res = resolution // 2 ** (self.num_resolutions - 1)
394 z_shape = (1, z_channels, curr_res, curr_res)
395 print(f"Working with z of shape {z_shape} = {np.prod(z_shape)} dimensions")
396
397 make_attn_cls = self._make_attn()
398 make_resblock_cls = self._make_resblock()
399 make_conv_cls = self._make_conv()
400 # z to block_in
401 self.conv_in = nn.Conv2d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
402
403 self.mid = nn.Module()
404 self.mid.block_1 = make_resblock_cls(
405 in_channels=block_in,
406 out_channels=block_in,
407 temb_channels=self.temb_ch,
408 dropout=dropout
409 )
410 self.mid.attn_1 = make_attn_cls(block_in, attn_type=attn_type)
411 self.mid.block_2 = make_resblock_cls(
412 in_channels=block_in,
413 out_channels=block_in,
414 temb_channels=self.temb_ch,
415 dropout=dropout

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected