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

Method get_param_groups

sat/vae_modules/autoencoder.py:355–373  ·  view source on GitHub ↗
(
        self, parameter_names: List[List[str]], optimizer_args: List[dict]
    )

Source from the content-addressed store, hash-verified

353 return full_log_dict
354
355 def get_param_groups(
356 self, parameter_names: List[List[str]], optimizer_args: List[dict]
357 ) -> Tuple[List[Dict[str, Any]], int]:
358 groups = []
359 num_params = 0
360 for names, args in zip(parameter_names, optimizer_args):
361 params = []
362 for pattern_ in names:
363 pattern_params = []
364 pattern = re.compile(pattern_)
365 for p_name, param in self.named_parameters():
366 if re.match(pattern, p_name):
367 pattern_params.append(param)
368 num_params += param.numel()
369 if len(pattern_params) == 0:
370 logpy.warn(f"Did not find parameters for pattern {pattern_}")
371 params.extend(pattern_params)
372 groups.append({"params": params, **args})
373 return groups, num_params
374
375 def configure_optimizers(self) -> List[torch.optim.Optimizer]:
376 if self.trainable_ae_params is None:

Callers 1

configure_optimizersMethod · 0.95

Calls 2

appendMethod · 0.80
extendMethod · 0.80

Tested by

no test coverage detected