MCPcopy Create free account
hub / github.com/IceClear/StableSR / __init__

Method __init__

ldm/models/autoencoder.py:23–69  ·  view source on GitHub ↗
(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
                 )

Source from the content-addressed store, hash-verified

21
22class 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):

Callers 4

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 6

init_from_ckptMethod · 0.95
EncoderClass · 0.90
DecoderClass · 0.90
instantiate_from_configFunction · 0.90
LitEmaClass · 0.85
register_bufferMethod · 0.45

Tested by

no test coverage detected