(
optim_method: str,
lr_scheduler: str,
init_lr: float,
net: Any,
train_loader_len: int,
max_epochs: int,
optimizer_kwargs=dict(),
scheduler_kwargs=dict())
| 10 | |
| 11 | |
| 12 | def get_optimizer_scheduler( |
| 13 | optim_method: str, |
| 14 | lr_scheduler: str, |
| 15 | init_lr: float, |
| 16 | net: Any, |
| 17 | train_loader_len: int, |
| 18 | max_epochs: int, |
| 19 | optimizer_kwargs=dict(), |
| 20 | scheduler_kwargs=dict()) -> Tuple[torch.optim.Optimizer, Any]: |
| 21 | optimizer = None |
| 22 | scheduler = None |
| 23 | optim_processed_kwargs = { |
| 24 | k: v for k, v in optimizer_kwargs.items() if v is not None} |
| 25 | scheduler_processed_kwargs = { |
| 26 | k: v for k, v in scheduler_kwargs.items() if v is not None} |
| 27 | if optim_method == "AdaFisher": |
| 28 | optimizer = AdaFisher(model=net, lr=init_lr, |
| 29 | **optim_processed_kwargs) |
| 30 | elif optim_method == "AdaFisherW": |
| 31 | optimizer = AdaFisherW(model=net, lr=init_lr, |
| 32 | **optim_processed_kwargs) |
| 33 | elif optim_method == 'SGD': |
| 34 | if 'momentum' not in optim_processed_kwargs.keys() or \ |
| 35 | 'weight_decay' not in optim_processed_kwargs.keys(): |
| 36 | raise ValueError( |
| 37 | "'momentum' and 'weight_decay' need to be specified for" |
| 38 | " SGD optimizer in config.yaml::**kwargs") |
| 39 | optimizer = SGD( |
| 40 | net.parameters(), lr=init_lr, |
| 41 | **optim_processed_kwargs) |
| 42 | elif optim_method == 'Adam': |
| 43 | optimizer = Adam(net.parameters(), lr=init_lr, |
| 44 | **optim_processed_kwargs) |
| 45 | elif optim_method == 'AdamW': |
| 46 | optimizer = AdamW(net.parameters(), lr=init_lr, |
| 47 | **optim_processed_kwargs) |
| 48 | elif optim_method == 'AdaHessian': |
| 49 | optimizer = Adahessian(net.parameters(), lr=init_lr, |
| 50 | **optim_processed_kwargs) |
| 51 | elif optim_method in ['Shampoo', 'kfac']: |
| 52 | optimizer = SGD( |
| 53 | net.parameters(), |
| 54 | lr=init_lr, |
| 55 | weight_decay=optim_processed_kwargs["weight_decay"], |
| 56 | momentum=optim_processed_kwargs["momentum"] |
| 57 | ) |
| 58 | else: |
| 59 | raise ValueError(f"Warning: Unknown optimizer {optim_method}") |
| 60 | if lr_scheduler == 'StepLR': |
| 61 | if 'step_size' not in scheduler_processed_kwargs.keys() or \ |
| 62 | 'gamma' not in scheduler_processed_kwargs.keys(): |
| 63 | raise ValueError( |
| 64 | "'step_size' and 'gamma' need to be specified for" |
| 65 | "StepLR scheduler in config.yaml::**kwargs") |
| 66 | scheduler = StepLR( |
| 67 | optimizer, |
| 68 | **scheduler_processed_kwargs) |
| 69 | elif lr_scheduler == "MultiStepLR": |
no test coverage detected