MCPcopy Create free account
hub / github.com/bbaaii/DreamDiffusion / UNetModel

Class UNetModel

code/dc_ldm/modules/diffusionmodules/openaimodel.py:415–761  ·  view source on GitHub ↗

The full UNet model with attention and timestep embedding. :param in_channels: channels in the input Tensor. :param model_channels: base channel count for the model. :param out_channels: channels in the output Tensor. :param num_res_blocks: number of residual blocks per downsamp

Source from the content-addressed store, hash-verified

413
414
415class UNetModel(nn.Module):
416 """
417 The full UNet model with attention and timestep embedding.
418 :param in_channels: channels in the input Tensor.
419 :param model_channels: base channel count for the model.
420 :param out_channels: channels in the output Tensor.
421 :param num_res_blocks: number of residual blocks per downsample.
422 :param attention_resolutions: a collection of downsample rates at which
423 attention will take place. May be a set, list, or tuple.
424 For example, if this contains 4, then at 4x downsampling, attention
425 will be used.
426 :param dropout: the dropout probability.
427 :param channel_mult: channel multiplier for each level of the UNet.
428 :param conv_resample: if True, use learned convolutions for upsampling and
429 downsampling.
430 :param dims: determines if the signal is 1D, 2D, or 3D.
431 :param num_classes: if specified (as an int), then this model will be
432 class-conditional with `num_classes` classes.
433 :param use_checkpoint: use gradient checkpointing to reduce memory usage.
434 :param num_heads: the number of attention heads in each attention layer.
435 :param num_heads_channels: if specified, ignore num_heads and instead use
436 a fixed channel width per attention head.
437 :param num_heads_upsample: works with num_heads to set a different number
438 of heads for upsampling. Deprecated.
439 :param use_scale_shift_norm: use a FiLM-like conditioning mechanism.
440 :param resblock_updown: use residual blocks for up/downsampling.
441 :param use_new_attention_order: use a different attention pattern for potentially
442 increased efficiency.
443 """
444
445 def __init__(
446 self,
447 image_size,
448 in_channels,
449 model_channels,
450 out_channels,
451 num_res_blocks,
452 attention_resolutions,
453 dropout=0,
454 channel_mult=(1, 2, 4, 8),
455 conv_resample=True,
456 dims=2,
457 num_classes=None,
458 use_checkpoint=False,
459 use_fp16=False,
460 num_heads=-1,
461 num_head_channels=-1,
462 num_heads_upsample=-1,
463 use_scale_shift_norm=False,
464 resblock_updown=False,
465 use_new_attention_order=False,
466 use_spatial_transformer=False, # custom transformer support
467 transformer_depth=1, # custom transformer support
468 context_dim=None, # custom transformer support
469 n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model
470 legacy=True,
471 cond_scale=1.0,
472 global_pool=False,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected