MCPcopy Create free account
hub / github.com/Francis-Rings/StableAnimator / DownBlock3D

Class DownBlock3D

animation/modules/unet_3d_blocks.py:553–639  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

551
552
553class DownBlock3D(nn.Module):
554 def __init__(
555 self,
556 in_channels: int,
557 out_channels: int,
558 temb_channels: int,
559 dropout: float = 0.0,
560 num_layers: int = 1,
561 resnet_eps: float = 1e-6,
562 resnet_time_scale_shift: str = "default",
563 resnet_act_fn: str = "swish",
564 resnet_groups: int = 32,
565 resnet_pre_norm: bool = True,
566 output_scale_factor: float = 1.0,
567 add_downsample: bool = True,
568 downsample_padding: int = 1,
569 ):
570 super().__init__()
571 resnets = []
572 temp_convs = []
573
574 for i in range(num_layers):
575 in_channels = in_channels if i == 0 else out_channels
576 resnets.append(
577 ResnetBlock2D(
578 in_channels=in_channels,
579 out_channels=out_channels,
580 temb_channels=temb_channels,
581 eps=resnet_eps,
582 groups=resnet_groups,
583 dropout=dropout,
584 time_embedding_norm=resnet_time_scale_shift,
585 non_linearity=resnet_act_fn,
586 output_scale_factor=output_scale_factor,
587 pre_norm=resnet_pre_norm,
588 )
589 )
590 temp_convs.append(
591 TemporalConvLayer(
592 out_channels,
593 out_channels,
594 dropout=0.1,
595 norm_num_groups=resnet_groups,
596 )
597 )
598
599 self.resnets = nn.ModuleList(resnets)
600 self.temp_convs = nn.ModuleList(temp_convs)
601
602 if add_downsample:
603 self.downsamplers = nn.ModuleList(
604 [
605 Downsample2D(
606 out_channels,
607 use_conv=True,
608 out_channels=out_channels,
609 padding=downsample_padding,
610 name="op",

Callers 1

get_down_blockFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected