(
self,
in_channels=3,
out_channels=3,
down_block_types=("DownEncoderBlock2D",),
block_out_channels=(64,),
layers_per_block=2,
norm_num_groups=32,
act_fn="silu",
double_z=True,
)
| 38 | |
| 39 | class Encoder(nn.Module): |
| 40 | def __init__( |
| 41 | self, |
| 42 | in_channels=3, |
| 43 | out_channels=3, |
| 44 | down_block_types=("DownEncoderBlock2D",), |
| 45 | block_out_channels=(64,), |
| 46 | layers_per_block=2, |
| 47 | norm_num_groups=32, |
| 48 | act_fn="silu", |
| 49 | double_z=True, |
| 50 | ): |
| 51 | super().__init__() |
| 52 | self.layers_per_block = layers_per_block |
| 53 | |
| 54 | self.conv_in = torch.nn.Conv2d( |
| 55 | in_channels, |
| 56 | block_out_channels[0], |
| 57 | kernel_size=3, |
| 58 | stride=1, |
| 59 | padding=1, |
| 60 | ) |
| 61 | |
| 62 | self.mid_block = None |
| 63 | self.down_blocks = nn.ModuleList([]) |
| 64 | |
| 65 | # down |
| 66 | output_channel = block_out_channels[0] |
| 67 | for i, down_block_type in enumerate(down_block_types): |
| 68 | input_channel = output_channel |
| 69 | output_channel = block_out_channels[i] |
| 70 | is_final_block = i == len(block_out_channels) - 1 |
| 71 | |
| 72 | down_block = get_down_block( |
| 73 | down_block_type, |
| 74 | num_layers=self.layers_per_block, |
| 75 | in_channels=input_channel, |
| 76 | out_channels=output_channel, |
| 77 | add_downsample=not is_final_block, |
| 78 | resnet_eps=1e-6, |
| 79 | downsample_padding=0, |
| 80 | resnet_act_fn=act_fn, |
| 81 | resnet_groups=norm_num_groups, |
| 82 | attn_num_head_channels=None, |
| 83 | temb_channels=None, |
| 84 | ) |
| 85 | self.down_blocks.append(down_block) |
| 86 | |
| 87 | # mid |
| 88 | self.mid_block = UNetMidBlock2D( |
| 89 | in_channels=block_out_channels[-1], |
| 90 | resnet_eps=1e-6, |
| 91 | resnet_act_fn=act_fn, |
| 92 | output_scale_factor=1, |
| 93 | resnet_time_scale_shift="default", |
| 94 | attn_num_head_channels=None, |
| 95 | resnet_groups=norm_num_groups, |
| 96 | temb_channels=None, |
| 97 | ) |
no test coverage detected