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", ]