PyNorch/norch/optim/optimizer.py

22 lines
611 B
Python
Raw Permalink Normal View History

2024-05-06 19:47:27 -03:00
from abc import ABC
from norch.tensor import Tensor
class Optimizer(ABC):
"""
Abstract class for optimizers
"""
def __init__(self, parameters):
if isinstance(parameters, Tensor):
raise TypeError("parameters should be an iterable but got {}".format(type(parameters)))
elif isinstance(parameters, dict):
parameters = parameters.values()
self.parameters = list(parameters)
def step(self):
raise NotImplementedError
def zero_grad(self):
for module, name, parameter in self.parameters:
2024-05-06 19:47:27 -03:00
parameter.zero_grad()