Source code for numpynet.optimizers
from .base import Optimizer
from .sgd import SGD
from .adam import Adam, AdamW
from .rmsprop import RMSprop
from .adagrad import Adagrad, Adadelta
_REGISTRY = {
"sgd": SGD,
"adam": Adam,
"adamw": AdamW,
"rmsprop": RMSprop,
"adagrad": Adagrad,
"adadelta": Adadelta,
}
[docs]
def get(identifier):
if isinstance(identifier, Optimizer):
return identifier
if isinstance(identifier, str):
key = identifier.lower()
if key in _REGISTRY:
return _REGISTRY[key]()
raise ValueError(f"Unknown optimizer: '{identifier}'. Available: {list(_REGISTRY)}")
raise TypeError(f"Could not interpret optimizer: {identifier}")
__all__ = [
"Optimizer",
"SGD",
"Adam",
"AdamW",
"RMSprop",
"Adagrad",
"Adadelta",
"get",
]