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