| 18 | |
| 19 | |
| 20 | class WrappedLightningCLI(LightningCLI): |
| 21 | def before_instantiate_classes(self) -> None: |
| 22 | self.config = format_with_env(self.config) |
| 23 | |
| 24 | # Changing the lr_scheduler interval to step instead of epoch |
| 25 | @staticmethod |
| 26 | def configure_optimizers( |
| 27 | lightning_module: LightningModule, |
| 28 | optimizer: Optimizer, |
| 29 | lr_scheduler: Optional[LRSchedulerTypeUnion] = None, |
| 30 | ) -> Any: |
| 31 | optimizer_list, lr_scheduler_list = LightningCLI.configure_optimizers( |
| 32 | lightning_module, optimizer=optimizer, lr_scheduler=lr_scheduler |
| 33 | ) |
| 34 | |
| 35 | for idx in range(len(lr_scheduler_list)): |
| 36 | if not isinstance(lr_scheduler_list[idx], dict): |
| 37 | lr_scheduler_list[idx] = { |
| 38 | "scheduler": lr_scheduler_list[idx], |
| 39 | "interval": "step", |
| 40 | } |
| 41 | return optimizer_list, lr_scheduler_list |
| 42 | |
| 43 | |
| 44 | def main_cli(args: ArgsType = None, run: bool = True): |