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):
|
2024-05-06 20:40:06 -03:00
|
|
|
for module, name, parameter in self.parameters:
|
2024-05-06 19:47:27 -03:00
|
|
|
parameter.zero_grad()
|