MCPcopy Create free account
hub / github.com/bbaaii/DreamDiffusion / __init__

Method __init__

code/dc_ldm/models/autoencoder.py:136–182  ·  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

134
135class VQModel(pl.LightningModule):
136 def __init__(self,
137 ddconfig,
138 lossconfig,
139 n_embed,
140 embed_dim,
141 ckpt_path=None,
142 ignore_keys=[],
143 image_key="image",
144 colorize_nlabels=None,
145 monitor=None,
146 batch_resize_range=None,
147 scheduler_config=None,
148 lr_g_factor=1.0,
149 remap=None,
150 sane_index_shape=False, # tell vector quantizer to return indices as bhw
151 use_ema=False
152 ):
153 super().__init__()
154 self.embed_dim = embed_dim
155 self.n_embed = n_embed
156 self.image_key = image_key
157 self.encoder = Encoder(**ddconfig)
158 self.decoder = Decoder(**ddconfig)
159 self.loss = instantiate_from_config(lossconfig)
160 self.quantize = VectorQuantizer(n_embed, embed_dim, beta=0.25,
161 remap=remap,
162 sane_index_shape=sane_index_shape)
163 self.quant_conv = torch.nn.Conv2d(ddconfig["z_channels"], embed_dim, 1)
164 self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
165 if colorize_nlabels is not None:
166 assert type(colorize_nlabels)==int
167 self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
168 if monitor is not None:
169 self.monitor = monitor
170 self.batch_resize_range = batch_resize_range
171 if self.batch_resize_range is not None:
172 print(f"{self.__class__.__name__}: Using per-batch resizing in range {batch_resize_range}.")
173
174 self.use_ema = use_ema
175 if self.use_ema:
176 self.model_ema = LitEma(self)
177 print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
178
179 if ckpt_path is not None:
180 self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
181 self.scheduler_config = scheduler_config
182 self.lr_g_factor = lr_g_factor
183
184 @contextmanager
185 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 7

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

Tested by

no test coverage detected