MCPcopy Create free account
hub / github.com/DPS2022/diffusion-posterior-sampling / __init__

Method __init__

guided_diffusion/unet.py:498–687  ·  view source on GitHub ↗
(
        self,
        image_size,
        in_channels,
        model_channels,
        out_channels,
        num_res_blocks,
        attention_resolutions,
        dropout=0,
        channel_mult=(1, 2, 4, 8),
        conv_resample=True,
        dims=2,
        num_classes=None,
        use_checkpoint=False,
        use_fp16=False,
        num_heads=1,
        num_head_channels=-1,
        num_heads_upsample=-1,
        use_scale_shift_norm=False,
        resblock_updown=False,
        use_new_attention_order=False,
    )

Source from the content-addressed store, hash-verified

496 """
497
498 def __init__(
499 self,
500 image_size,
501 in_channels,
502 model_channels,
503 out_channels,
504 num_res_blocks,
505 attention_resolutions,
506 dropout=0,
507 channel_mult=(1, 2, 4, 8),
508 conv_resample=True,
509 dims=2,
510 num_classes=None,
511 use_checkpoint=False,
512 use_fp16=False,
513 num_heads=1,
514 num_head_channels=-1,
515 num_heads_upsample=-1,
516 use_scale_shift_norm=False,
517 resblock_updown=False,
518 use_new_attention_order=False,
519 ):
520 super().__init__()
521
522 if num_heads_upsample == -1:
523 num_heads_upsample = num_heads
524
525 self.image_size = image_size
526 self.in_channels = in_channels
527 self.model_channels = model_channels
528 self.out_channels = out_channels
529 self.num_res_blocks = num_res_blocks
530 self.attention_resolutions = attention_resolutions
531 self.dropout = dropout
532 self.channel_mult = channel_mult
533 self.conv_resample = conv_resample
534 self.num_classes = num_classes
535 self.use_checkpoint = use_checkpoint
536 self.dtype = th.float16 if use_fp16 else th.float32
537 self.num_heads = num_heads
538 self.num_head_channels = num_head_channels
539 self.num_heads_upsample = num_heads_upsample
540
541 time_embed_dim = model_channels * 4
542 self.time_embed = nn.Sequential(
543 linear(model_channels, time_embed_dim),
544 nn.SiLU(),
545 linear(time_embed_dim, time_embed_dim),
546 )
547
548 if self.num_classes is not None:
549 self.label_emb = nn.Embedding(num_classes, time_embed_dim)
550
551 ch = input_ch = int(channel_mult[0] * model_channels)
552 self.input_blocks = nn.ModuleList(
553 [TimestepEmbedSequential(conv_nd(dims, in_channels, ch, 3, padding=1))]
554 )
555 self._feature_size = ch

Callers

nothing calls this directly

Calls 10

conv_ndFunction · 0.85
ResBlockClass · 0.85
AttentionBlockClass · 0.85
DownsampleClass · 0.85
UpsampleClass · 0.85
normalizationFunction · 0.85
zero_moduleFunction · 0.85
linearFunction · 0.70
__init__Method · 0.45

Tested by

no test coverage detected