(self, embed_dim: int, **kwargs)
| 417 | |
| 418 | class AutoencodingEngineLegacy(AutoencodingEngine): |
| 419 | def __init__(self, embed_dim: int, **kwargs): |
| 420 | self.max_batch_size = kwargs.pop("max_batch_size", None) |
| 421 | ddconfig = kwargs.pop("ddconfig") |
| 422 | ckpt_path = kwargs.pop("ckpt_path", None) |
| 423 | ckpt_engine = kwargs.pop("ckpt_engine", None) |
| 424 | super().__init__( |
| 425 | encoder_config={ |
| 426 | "target": "sgm.modules.diffusionmodules.model.Encoder", |
| 427 | "params": ddconfig, |
| 428 | }, |
| 429 | decoder_config={ |
| 430 | "target": "sgm.modules.diffusionmodules.model.Decoder", |
| 431 | "params": ddconfig, |
| 432 | }, |
| 433 | **kwargs, |
| 434 | ) |
| 435 | self.quant_conv = torch.nn.Conv2d( |
| 436 | (1 + ddconfig["double_z"]) * ddconfig["z_channels"], |
| 437 | (1 + ddconfig["double_z"]) * embed_dim, |
| 438 | 1, |
| 439 | ) |
| 440 | self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) |
| 441 | self.embed_dim = embed_dim |
| 442 | |
| 443 | self.apply_ckpt(default(ckpt_path, ckpt_engine)) |
| 444 | |
| 445 | def get_autoencoder_params(self) -> list: |
| 446 | params = super().get_autoencoder_params() |
nothing calls this directly
no test coverage detected