MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / get_optimizer_scheduler

Function get_optimizer_scheduler

optimizers/__init__.py:12–110  ·  view source on GitHub ↗
(
        optim_method: str,
        lr_scheduler: str,
        init_lr: float,
        net: Any,
        train_loader_len: int,
        max_epochs: int,
        optimizer_kwargs=dict(),
        scheduler_kwargs=dict())

Source from the content-addressed store, hash-verified

10
11
12def 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":

Callers 1

resetMethod · 0.90

Calls 12

AdaFisherClass · 0.90
AdaFisherWClass · 0.90
SGDClass · 0.90
AdamClass · 0.90
AdamWClass · 0.90
AdahessianClass · 0.90
StepLRClass · 0.90
MultiStepLRClass · 0.90
CosineAnnealingLRClass · 0.90
OneCycleLRClass · 0.90
LinearLRClass · 0.90

Tested by

no test coverage detected