(self,
ddconfig,
lossconfig,
n_embed,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
batch_resize_range=None,
scheduler_config=None,
lr_g_factor=1.0,
remap=None,
sane_index_shape=False, # tell vector quantizer to return indices as bhw
use_ema=False
)
| 21 | |
| 22 | class VQModel(pl.LightningModule): |
| 23 | def __init__(self, |
| 24 | ddconfig, |
| 25 | lossconfig, |
| 26 | n_embed, |
| 27 | embed_dim, |
| 28 | ckpt_path=None, |
| 29 | ignore_keys=[], |
| 30 | image_key="image", |
| 31 | colorize_nlabels=None, |
| 32 | monitor=None, |
| 33 | batch_resize_range=None, |
| 34 | scheduler_config=None, |
| 35 | lr_g_factor=1.0, |
| 36 | remap=None, |
| 37 | sane_index_shape=False, # tell vector quantizer to return indices as bhw |
| 38 | use_ema=False |
| 39 | ): |
| 40 | super().__init__() |
| 41 | self.embed_dim = embed_dim |
| 42 | self.n_embed = n_embed |
| 43 | self.image_key = image_key |
| 44 | self.encoder = Encoder(**ddconfig) |
| 45 | self.decoder = Decoder(**ddconfig) |
| 46 | self.loss = instantiate_from_config(lossconfig) |
| 47 | self.quantize = VectorQuantizer(n_embed, embed_dim, beta=0.25, |
| 48 | remap=remap, |
| 49 | sane_index_shape=sane_index_shape) |
| 50 | self.quant_conv = torch.nn.Conv2d(ddconfig["z_channels"], embed_dim, 1) |
| 51 | self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) |
| 52 | if colorize_nlabels is not None: |
| 53 | assert type(colorize_nlabels)==int |
| 54 | self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1)) |
| 55 | if monitor is not None: |
| 56 | self.monitor = monitor |
| 57 | self.batch_resize_range = batch_resize_range |
| 58 | if self.batch_resize_range is not None: |
| 59 | print(f"{self.__class__.__name__}: Using per-batch resizing in range {batch_resize_range}.") |
| 60 | |
| 61 | self.use_ema = use_ema |
| 62 | if self.use_ema: |
| 63 | self.model_ema = LitEma(self) |
| 64 | print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.") |
| 65 | |
| 66 | if ckpt_path is not None: |
| 67 | self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) |
| 68 | self.scheduler_config = scheduler_config |
| 69 | self.lr_g_factor = lr_g_factor |
| 70 | |
| 71 | @contextmanager |
| 72 | def ema_scope(self, context=None): |
no test coverage detected