From 78de0ea9ee8cbd65dce28bd4abcc131c98451aa2 Mon Sep 17 00:00:00 2001 From: David Hou Date: Thu, 29 Feb 2024 13:45:24 -0800 Subject: [PATCH] don't increment num_batches_tracked if not tracking running stats --- tinygrad/nn/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tinygrad/nn/__init__.py b/tinygrad/nn/__init__.py index 4e81e9eb0a..8eb9c015f2 100644 --- a/tinygrad/nn/__init__.py +++ b/tinygrad/nn/__init__.py @@ -39,7 +39,7 @@ class BatchNorm2d: self.running_mean.assign((1-self.momentum) * self.running_mean + self.momentum * batch_mean.detach()) self.running_var.assign((1-self.momentum) * self.running_var + self.momentum * prod(y.shape[1:])/(prod(y.shape[1:])-y.shape[2]) * batch_var.detach()) - self.num_batches_tracked += 1 + self.num_batches_tracked += 1 else: batch_mean = self.running_mean # NOTE: this can be precomputed for static inference. we expand it here so it fuses