(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str, ...] = ("DownEncoderBlock2D",),
up_block_types: Tuple[str, ...] = ("UpDecoderBlock2D",),
block_out_channels: Tuple[int, ...] = (64,),
layers_per_block: int = 1,
act_fn: str = "silu",
latent_channels: int = 3,
sample_size: int = 32,
num_vq_embeddings: int = 256,
norm_num_groups: int = 32,
vq_embed_dim: Optional[int] = None,
scaling_factor: float = 0.18215,
norm_type: str = "group", # group, spatial
mid_block_add_attention=True,
lookup_from_codebook=False,
force_upcast=False,
)
| 73 | |
| 74 | @register_to_config |
| 75 | def __init__( |
| 76 | self, |
| 77 | in_channels: int = 3, |
| 78 | out_channels: int = 3, |
| 79 | down_block_types: Tuple[str, ...] = ("DownEncoderBlock2D",), |
| 80 | up_block_types: Tuple[str, ...] = ("UpDecoderBlock2D",), |
| 81 | block_out_channels: Tuple[int, ...] = (64,), |
| 82 | layers_per_block: int = 1, |
| 83 | act_fn: str = "silu", |
| 84 | latent_channels: int = 3, |
| 85 | sample_size: int = 32, |
| 86 | num_vq_embeddings: int = 256, |
| 87 | norm_num_groups: int = 32, |
| 88 | vq_embed_dim: Optional[int] = None, |
| 89 | scaling_factor: float = 0.18215, |
| 90 | norm_type: str = "group", # group, spatial |
| 91 | mid_block_add_attention=True, |
| 92 | lookup_from_codebook=False, |
| 93 | force_upcast=False, |
| 94 | ): |
| 95 | super().__init__() |
| 96 | |
| 97 | # pass init params to Encoder |
| 98 | self.encoder = Encoder( |
| 99 | in_channels=in_channels, |
| 100 | out_channels=latent_channels, |
| 101 | down_block_types=down_block_types, |
| 102 | block_out_channels=block_out_channels, |
| 103 | layers_per_block=layers_per_block, |
| 104 | act_fn=act_fn, |
| 105 | norm_num_groups=norm_num_groups, |
| 106 | double_z=False, |
| 107 | mid_block_add_attention=mid_block_add_attention, |
| 108 | ) |
| 109 | |
| 110 | vq_embed_dim = vq_embed_dim if vq_embed_dim is not None else latent_channels |
| 111 | |
| 112 | self.quant_conv = nn.Conv2d(latent_channels, vq_embed_dim, 1) |
| 113 | self.quantize = VectorQuantizer(num_vq_embeddings, vq_embed_dim, beta=0.25, remap=None, sane_index_shape=False) |
| 114 | self.post_quant_conv = nn.Conv2d(vq_embed_dim, latent_channels, 1) |
| 115 | |
| 116 | # pass init params to Decoder |
| 117 | self.decoder = Decoder( |
| 118 | in_channels=latent_channels, |
| 119 | out_channels=out_channels, |
| 120 | up_block_types=up_block_types, |
| 121 | block_out_channels=block_out_channels, |
| 122 | layers_per_block=layers_per_block, |
| 123 | act_fn=act_fn, |
| 124 | norm_num_groups=norm_num_groups, |
| 125 | norm_type=norm_type, |
| 126 | mid_block_add_attention=mid_block_add_attention, |
| 127 | ) |
| 128 | |
| 129 | @apply_forward_hook |
| 130 | def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> VQEncoderOutput: |
nothing calls this directly
no test coverage detected