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

Method get_param_groups

sat/sgm/models/autoencoder.py:340–358  ·  view source on GitHub ↗
(
        self, parameter_names: List[List[str]], optimizer_args: List[dict]
    )

Source from the content-addressed store, hash-verified

338 return full_log_dict
339
340 def get_param_groups(
341 self, parameter_names: List[List[str]], optimizer_args: List[dict]
342 ) -> Tuple[List[Dict[str, Any]], int]:
343 groups = []
344 num_params = 0
345 for names, args in zip(parameter_names, optimizer_args):
346 params = []
347 for pattern_ in names:
348 pattern_params = []
349 pattern = re.compile(pattern_)
350 for p_name, param in self.named_parameters():
351 if re.match(pattern, p_name):
352 pattern_params.append(param)
353 num_params += param.numel()
354 if len(pattern_params) == 0:
355 logpy.warn(f"Did not find parameters for pattern {pattern_}")
356 params.extend(pattern_params)
357 groups.append({"params": params, **args})
358 return groups, num_params
359
360 def configure_optimizers(self) -> List[torch.optim.Optimizer]:
361 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