MCPcopy Create free account
hub / github.com/TencentARC/AnimeGamer / __init__

Method __init__

VDM_Decoder/vae_modules/autoencoder.py:135–190  ·  view source on GitHub ↗
(
        self,
        *args,
        encoder_config: Dict,
        decoder_config: Dict,
        loss_config: Dict,
        regularizer_config: Dict,
        optimizer_config: Union[Dict, None] = None,
        lr_g_factor: float = 1.0,
        trainable_ae_params: Optional[List[List[str]]] = None,
        ae_optimizer_args: Optional[List[dict]] = None,
        trainable_disc_params: Optional[List[List[str]]] = None,
        disc_optimizer_args: Optional[List[dict]] = None,
        disc_start_iter: int = 0,
        diff_boost_factor: float = 3.0,
        ckpt_engine: Union[None, str, dict] = None,
        ckpt_path: Optional[str] = None,
        additional_decode_keys: Optional[List[str]] = None,
        **kwargs,
    )

Source from the content-addressed store, hash-verified

133 """
134
135 def __init__(
136 self,
137 *args,
138 encoder_config: Dict,
139 decoder_config: Dict,
140 loss_config: Dict,
141 regularizer_config: Dict,
142 optimizer_config: Union[Dict, None] = None,
143 lr_g_factor: float = 1.0,
144 trainable_ae_params: Optional[List[List[str]]] = None,
145 ae_optimizer_args: Optional[List[dict]] = None,
146 trainable_disc_params: Optional[List[List[str]]] = None,
147 disc_optimizer_args: Optional[List[dict]] = None,
148 disc_start_iter: int = 0,
149 diff_boost_factor: float = 3.0,
150 ckpt_engine: Union[None, str, dict] = None,
151 ckpt_path: Optional[str] = None,
152 additional_decode_keys: Optional[List[str]] = None,
153 **kwargs,
154 ):
155 super().__init__(*args, **kwargs)
156 self.automatic_optimization = False # pytorch lightning
157
158 self.encoder = instantiate_from_config(encoder_config)
159 self.decoder = instantiate_from_config(decoder_config)
160 self.loss = instantiate_from_config(loss_config)
161 self.regularization = instantiate_from_config(regularizer_config)
162 self.optimizer_config = default(optimizer_config, {"target": "torch.optim.Adam"})
163 self.diff_boost_factor = diff_boost_factor
164 self.disc_start_iter = disc_start_iter
165 self.lr_g_factor = lr_g_factor
166 self.trainable_ae_params = trainable_ae_params
167 if self.trainable_ae_params is not None:
168 self.ae_optimizer_args = default(
169 ae_optimizer_args,
170 [{} for _ in range(len(self.trainable_ae_params))],
171 )
172 assert len(self.ae_optimizer_args) == len(self.trainable_ae_params)
173 else:
174 self.ae_optimizer_args = [{}] # makes type consitent
175
176 self.trainable_disc_params = trainable_disc_params
177 if self.trainable_disc_params is not None:
178 self.disc_optimizer_args = default(
179 disc_optimizer_args,
180 [{} for _ in range(len(self.trainable_disc_params))],
181 )
182 assert len(self.disc_optimizer_args) == len(self.trainable_disc_params)
183 else:
184 self.disc_optimizer_args = [{}] # makes type consitent
185
186 if ckpt_path is not None:
187 assert ckpt_engine is None, "Can't set ckpt_engine and ckpt_path"
188 logpy.warn("Checkpoint path is deprecated, use `checkpoint_egnine` instead")
189 self.apply_ckpt(default(ckpt_path, ckpt_engine))
190 self.additional_decode_keys = set(default(additional_decode_keys, []))
191
192 def get_input(self, batch: Dict) -> torch.Tensor:

Callers

nothing calls this directly

Calls 4

instantiate_from_configFunction · 0.70
defaultFunction · 0.70
__init__Method · 0.45
apply_ckptMethod · 0.45

Tested by

no test coverage detected