(self,
ddconfig,
lossconfig,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
test=False,
logdir=None,
input_dim=4,
test_args=None,
)
| 14 | |
| 15 | class AutoencoderKL(pl.LightningModule): |
| 16 | def __init__(self, |
| 17 | ddconfig, |
| 18 | lossconfig, |
| 19 | embed_dim, |
| 20 | ckpt_path=None, |
| 21 | ignore_keys=[], |
| 22 | image_key="image", |
| 23 | colorize_nlabels=None, |
| 24 | monitor=None, |
| 25 | test=False, |
| 26 | logdir=None, |
| 27 | input_dim=4, |
| 28 | test_args=None, |
| 29 | ): |
| 30 | super().__init__() |
| 31 | self.image_key = image_key |
| 32 | self.encoder = Encoder(**ddconfig) |
| 33 | self.decoder = Decoder(**ddconfig) |
| 34 | self.loss = instantiate_from_config(lossconfig) |
| 35 | assert ddconfig["double_z"] |
| 36 | self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) |
| 37 | self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) |
| 38 | self.embed_dim = embed_dim |
| 39 | self.input_dim = input_dim |
| 40 | self.test = test |
| 41 | self.test_args = test_args |
| 42 | self.logdir = logdir |
| 43 | if colorize_nlabels is not None: |
| 44 | assert type(colorize_nlabels)==int |
| 45 | self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1)) |
| 46 | if monitor is not None: |
| 47 | self.monitor = monitor |
| 48 | if ckpt_path is not None: |
| 49 | self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) |
| 50 | if self.test: |
| 51 | self.init_test() |
| 52 | |
| 53 | def init_test(self,): |
| 54 | self.test = True |
no test coverage detected