MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/sgm/models/autoencoder.py:119–174  ·  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

117 """
118
119 def __init__(
120 self,
121 *args,
122 encoder_config: Dict,
123 decoder_config: Dict,
124 loss_config: Dict,
125 regularizer_config: Dict,
126 optimizer_config: Union[Dict, None] = None,
127 lr_g_factor: float = 1.0,
128 trainable_ae_params: Optional[List[List[str]]] = None,
129 ae_optimizer_args: Optional[List[dict]] = None,
130 trainable_disc_params: Optional[List[List[str]]] = None,
131 disc_optimizer_args: Optional[List[dict]] = None,
132 disc_start_iter: int = 0,
133 diff_boost_factor: float = 3.0,
134 ckpt_engine: Union[None, str, dict] = None,
135 ckpt_path: Optional[str] = None,
136 additional_decode_keys: Optional[List[str]] = None,
137 **kwargs,
138 ):
139 super().__init__(*args, **kwargs)
140 self.automatic_optimization = False # pytorch lightning
141
142 self.encoder: torch.nn.Module = instantiate_from_config(encoder_config)
143 self.decoder: torch.nn.Module = instantiate_from_config(decoder_config)
144 self.loss: torch.nn.Module = instantiate_from_config(loss_config)
145 self.regularization: AbstractRegularizer = instantiate_from_config(regularizer_config)
146 self.optimizer_config = default(optimizer_config, {"target": "torch.optim.Adam"})
147 self.diff_boost_factor = diff_boost_factor
148 self.disc_start_iter = disc_start_iter
149 self.lr_g_factor = lr_g_factor
150 self.trainable_ae_params = trainable_ae_params
151 if self.trainable_ae_params is not None:
152 self.ae_optimizer_args = default(
153 ae_optimizer_args,
154 [{} for _ in range(len(self.trainable_ae_params))],
155 )
156 assert len(self.ae_optimizer_args) == len(self.trainable_ae_params)
157 else:
158 self.ae_optimizer_args = [{}] # makes type consitent
159
160 self.trainable_disc_params = trainable_disc_params
161 if self.trainable_disc_params is not None:
162 self.disc_optimizer_args = default(
163 disc_optimizer_args,
164 [{} for _ in range(len(self.trainable_disc_params))],
165 )
166 assert len(self.disc_optimizer_args) == len(self.trainable_disc_params)
167 else:
168 self.disc_optimizer_args = [{}] # makes type consitent
169
170 if ckpt_path is not None:
171 assert ckpt_engine is None, "Can't set ckpt_engine and ckpt_path"
172 logpy.warn("Checkpoint path is deprecated, use `checkpoint_egnine` instead")
173 self.apply_ckpt(default(ckpt_path, ckpt_engine))
174 self.additional_decode_keys = set(default(additional_decode_keys, []))
175
176 def get_input(self, batch: Dict) -> torch.Tensor:

Callers 5

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

Calls 3

instantiate_from_configFunction · 0.50
defaultFunction · 0.50
apply_ckptMethod · 0.45

Tested by

no test coverage detected