From 9152bb5b4a03c1dd5abe2f99ddef2d894d95a3bf Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sat, 11 Feb 2023 10:22:37 -0800 Subject: [PATCH] momentum support in SGD --- examples/hlb_cifar10.py | 3 ++- tinygrad/nn/optim.py | 16 +++++++++++----- 2 files changed, 13 insertions(+), 6 deletions(-) diff --git a/examples/hlb_cifar10.py b/examples/hlb_cifar10.py index 3a783a75b3..45cc02a201 100644 --- a/examples/hlb_cifar10.py +++ b/examples/hlb_cifar10.py @@ -87,7 +87,8 @@ def train_cifar(): if getenv("ADAM"): optimizer = optim.Adam(get_parameters(model), lr=3e-4) else: - optimizer = optim.SGD(get_parameters(model), lr=0.001) + #optimizer = optim.SGD(get_parameters(model), lr=0.001) + optimizer = optim.SGD(get_parameters(model), lr=0.001, momentum=0.85, nesterov=True) # 97 steps in 2 seconds = 20ms / step # step is 1163.42 GOPS = 56 TFLOPS!!!, 41% of max 136 diff --git a/tinygrad/nn/optim.py b/tinygrad/nn/optim.py index 564312d67f..ce7fb19bb5 100644 --- a/tinygrad/nn/optim.py +++ b/tinygrad/nn/optim.py @@ -28,15 +28,21 @@ class Optimizer: p.realize() class SGD(Optimizer): - def __init__(self, params : List[Tensor], lr=0.001): + def __init__(self, params : List[Tensor], lr=0.001, momentum=0, nesterov=False): super().__init__(params) - self.lr = lr + self.lr, self.momentum, self.nesterov = lr, momentum, nesterov + self.b = [Tensor.zeros(*t.shape, device=params[0].device, requires_grad=False) for t in self.params] if self.momentum else [] + # https://pytorch.org/docs/stable/generated/torch.optim.SGD.html def step(self) -> None: - for t in self.params: + for i, t in enumerate(self.params): assert t.grad is not None - t.assign(t.detach() - t.grad * self.lr) - self.realize() + g = t.grad + if self.momentum: + self.b[i].assign(self.momentum * self.b[i] + g) + g = (g + self.momentum * self.b[i]) if self.nesterov else self.b[i] + t.assign(t.detach() - g * self.lr) + self.realize(self.b) class RMSprop(Optimizer): def __init__(self, params : List[Tensor], lr=0.001, decay=0.9, eps=1e-8):