| 288 | |
| 289 | |
| 290 | class AutoEncoder(nn.Module): |
| 291 | def __init__(self, params: AutoEncoderParams): |
| 292 | super().__init__() |
| 293 | self.encoder = Encoder( |
| 294 | resolution=params.resolution, |
| 295 | in_channels=params.in_channels, |
| 296 | ch=params.ch, |
| 297 | ch_mult=params.ch_mult, |
| 298 | num_res_blocks=params.num_res_blocks, |
| 299 | z_channels=params.z_channels, |
| 300 | ) |
| 301 | self.decoder = Decoder( |
| 302 | resolution=params.resolution, |
| 303 | in_channels=params.in_channels, |
| 304 | ch=params.ch, |
| 305 | out_ch=params.out_ch, |
| 306 | ch_mult=params.ch_mult, |
| 307 | num_res_blocks=params.num_res_blocks, |
| 308 | z_channels=params.z_channels, |
| 309 | ) |
| 310 | self.reg = DiagonalGaussian() |
| 311 | |
| 312 | self.scale_factor = params.scale_factor |
| 313 | self.shift_factor = params.shift_factor |
| 314 | |
| 315 | def encode(self, x: Tensor) -> Tensor: |
| 316 | z = self.reg(self.encoder(x)) |
| 317 | z = self.scale_factor * (z - self.shift_factor) |
| 318 | return z |
| 319 | |
| 320 | def decode(self, z: Tensor) -> Tensor: |
| 321 | z = z / self.scale_factor + self.shift_factor |
| 322 | return self.decoder(z) |
| 323 | |
| 324 | def forward(self, x: Tensor) -> Tensor: |
| 325 | return self.decode(self.encode(x)) |
| 326 | |
| 327 | |
| 328 | def print_load_warning(missing: list[str], unexpected: list[str]) -> None: |