(
self, parameter_names: List[List[str]], optimizer_args: List[dict]
)
| 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: |
no test coverage detected