CoolFace
Apppublic

Doubiiu/DynamiCrafter_interp_loop

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
165likes
autoencoder.py219 linesDownload Raw Back to models
1import os2from contextlib import contextmanager3import torch4import numpy as np5from einops import rearrange6import torch.nn.functional as F7import pytorch_lightning as pl8from lvdm.modules.networks.ae_modules import Encoder, Decoder9from lvdm.distributions import DiagonalGaussianDistribution10from utils.utils import instantiate_from_config11 12 13class AutoencoderKL(pl.LightningModule):14    def __init__(self,15                 ddconfig,16                 lossconfig,17                 embed_dim,18                 ckpt_path=None,19                 ignore_keys=[],20                 image_key="image",21                 colorize_nlabels=None,22                 monitor=None,23                 test=False,24                 logdir=None,25                 input_dim=4,26                 test_args=None,27                 ):28        super().__init__()29        self.image_key = image_key30        self.encoder = Encoder(**ddconfig)31        self.decoder = Decoder(**ddconfig)32        self.loss = instantiate_from_config(lossconfig)33        assert ddconfig["double_z"]34        self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)35        self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)36        self.embed_dim = embed_dim37        self.input_dim = input_dim38        self.test = test39        self.test_args = test_args40        self.logdir = logdir41        if colorize_nlabels is not None:42            assert type(colorize_nlabels)==int43            self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))44        if monitor is not None:45            self.monitor = monitor46        if ckpt_path is not None:47            self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)48        if self.test:49            self.init_test()50    51    def init_test(self,):52        self.test = True53        save_dir = os.path.join(self.logdir, "test")54        if 'ckpt' in self.test_args:55            ckpt_name = os.path.basename(self.test_args.ckpt).split('.ckpt')[0] + f'_epoch{self._cur_epoch}'56            self.root = os.path.join(save_dir, ckpt_name)57        else:58            self.root = save_dir59        if 'test_subdir' in self.test_args:60            self.root = os.path.join(save_dir, self.test_args.test_subdir)61 62        self.root_zs = os.path.join(self.root, "zs")63        self.root_dec = os.path.join(self.root, "reconstructions")64        self.root_inputs = os.path.join(self.root, "inputs")65        os.makedirs(self.root, exist_ok=True)66 67        if self.test_args.save_z:68            os.makedirs(self.root_zs, exist_ok=True)69        if self.test_args.save_reconstruction:70            os.makedirs(self.root_dec, exist_ok=True)71        if self.test_args.save_input:72            os.makedirs(self.root_inputs, exist_ok=True)73        assert(self.test_args is not None)74        self.test_maximum = getattr(self.test_args, 'test_maximum', None) 75        self.count = 076        self.eval_metrics = {}77        self.decodes = []78        self.save_decode_samples = 204879 80    def init_from_ckpt(self, path, ignore_keys=list()):81        sd = torch.load(path, map_location="cpu")82        try:83            self._cur_epoch = sd['epoch']84            sd = sd["state_dict"]85        except:86            self._cur_epoch = 'null'87        keys = list(sd.keys())88        for k in keys:89            for ik in ignore_keys:90                if k.startswith(ik):91                    print("Deleting key {} from state_dict.".format(k))92                    del sd[k]93        self.load_state_dict(sd, strict=False)94        # self.load_state_dict(sd, strict=True)95        print(f"Restored from {path}")96 97    def encode(self, x, **kwargs):98        99        h = self.encoder(x)100        moments = self.quant_conv(h)101        posterior = DiagonalGaussianDistribution(moments)102        return posterior103 104    def decode(self, z, **kwargs):105        z = self.post_quant_conv(z)106        dec = self.decoder(z)107        return dec108 109    def forward(self, input, sample_posterior=True):110        posterior = self.encode(input)111        if sample_posterior:112            z = posterior.sample()113        else:114            z = posterior.mode()115        dec = self.decode(z)116        return dec, posterior117 118    def get_input(self, batch, k):119        x = batch[k]120        if x.dim() == 5 and self.input_dim == 4:121            b,c,t,h,w = x.shape122            self.b = b123            self.t = t 124            x = rearrange(x, 'b c t h w -> (b t) c h w')125 126        return x127 128    def training_step(self, batch, batch_idx, optimizer_idx):129        inputs = self.get_input(batch, self.image_key)130        reconstructions, posterior = self(inputs)131 132        if optimizer_idx == 0:133            # train encoder+decoder+logvar134            aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,135                                            last_layer=self.get_last_layer(), split="train")136            self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)137            self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)138            return aeloss139 140        if optimizer_idx == 1:141            # train the discriminator142            discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,143                                                last_layer=self.get_last_layer(), split="train")144 145            self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)146            self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)147            return discloss148 149    def validation_step(self, batch, batch_idx):150        inputs = self.get_input(batch, self.image_key)151        reconstructions, posterior = self(inputs)152        aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,153                                        last_layer=self.get_last_layer(), split="val")154 155        discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,156                                            last_layer=self.get_last_layer(), split="val")157 158        self.log("val/rec_loss", log_dict_ae["val/rec_loss"])159        self.log_dict(log_dict_ae)160        self.log_dict(log_dict_disc)161        return self.log_dict162    163    def configure_optimizers(self):164        lr = self.learning_rate165        opt_ae = torch.optim.Adam(list(self.encoder.parameters())+166                                  list(self.decoder.parameters())+167                                  list(self.quant_conv.parameters())+168                                  list(self.post_quant_conv.parameters()),169                                  lr=lr, betas=(0.5, 0.9))170        opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),171                                    lr=lr, betas=(0.5, 0.9))172        return [opt_ae, opt_disc], []173 174    def get_last_layer(self):175        return self.decoder.conv_out.weight176 177    @torch.no_grad()178    def log_images(self, batch, only_inputs=False, **kwargs):179        log = dict()180        x = self.get_input(batch, self.image_key)181        x = x.to(self.device)182        if not only_inputs:183            xrec, posterior = self(x)184            if x.shape[1] > 3:185                # colorize with random projection186                assert xrec.shape[1] > 3187                x = self.to_rgb(x)188                xrec = self.to_rgb(xrec)189            log["samples"] = self.decode(torch.randn_like(posterior.sample()))190            log["reconstructions"] = xrec191        log["inputs"] = x192        return log193 194    def to_rgb(self, x):195        assert self.image_key == "segmentation"196        if not hasattr(self, "colorize"):197            self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))198        x = F.conv2d(x, weight=self.colorize)199        x = 2.*(x-x.min())/(x.max()-x.min()) - 1.200        return x201 202class IdentityFirstStage(torch.nn.Module):203    def __init__(self, *args, vq_interface=False, **kwargs):204        self.vq_interface = vq_interface  # TODO: Should be true by default but check to not break older stuff205        super().__init__()206 207    def encode(self, x, *args, **kwargs):208        return x209 210    def decode(self, x, *args, **kwargs):211        return x212 213    def quantize(self, x, *args, **kwargs):214        if self.vq_interface:215            return x, None, [None, None, None]216        return x217 218    def forward(self, x, *args, **kwargs):219        return x