Spaces:
Sleeping
Sleeping
File size: 2,299 Bytes
bda104d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 | import abc
from typing import Tuple
import torch
class LazyOptimizer(abc.ABC):
r"""Lazy implementation of optimizers. Contrary to standard PyTorch optimizers, these don't require the network
parameters at initialization and can therefore be configured directly from the command line then instantiated
properly afterwards.
"""
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
self.optimizer_class = None
def __call__(self, parameters) -> torch.optim.Optimizer:
return self.optimizer_class(parameters, *self.args, **self.kwargs)
def __str__(self):
params = '\n'.join(f"\t{arg}," for arg in self.args) + '\n'.join(f"\t{k}: {v}" for k, v in self.kwargs.items())
return self.optimizer_class.__name__ + "(\n" + params + "\n)"
class Adam(LazyOptimizer):
def __init__(
self,
lr: float = 1e-3,
betas: Tuple[float, float] = (0.9, 0.999),
eps: float = 1e-8,
weight_decay: float = 0.,
amsgrad: bool = False,
**kwargs
):
super(Adam, self).__init__(
lr=lr,
betas=betas,
eps=eps,
weight_decay=weight_decay,
amsgrad=amsgrad,
**kwargs
)
self.optimizer_class = torch.optim.Adam
class LazyScheduler(abc.ABC):
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
self.scheduler_class = None
def __call__(self, optimizer):
return self.scheduler_class(optimizer, *self.args, **self.kwargs)
def __str__(self):
params = '\n'.join(f"\t{arg}," for arg in self.args) + '\n'.join(f"\t{k}: {v}" for k, v in self.kwargs.items())
return self.scheduler_class.__name__ + "(\n" + params + "\n)"
class CosineAnnealing(LazyScheduler):
def __init__(
self,
T_max: int,
eta_min: float = 0,
last_epoch: int = -1,
verbose: bool = False
):
super(CosineAnnealing, self).__init__(
T_max=T_max,
eta_min=eta_min,
last_epoch=last_epoch,
verbose=verbose
)
self.scheduler_class = torch.optim.lr_scheduler.CosineAnnealingLR
|