From 9b49a07d2a001b7ce65cd0b69c2ef3b962f36096 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Thu, 23 May 2024 12:11:59 -0300 Subject: [PATCH] fix max min sum access elements autograd backward running on device --- norch/autograd/functions.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/norch/autograd/functions.py b/norch/autograd/functions.py index 917bd94..477837d 100644 --- a/norch/autograd/functions.py +++ b/norch/autograd/functions.py @@ -136,7 +136,7 @@ class SumBackward: input_shape = self.input[0].shape.copy() if self.axis == -1: # If axis is None, sum reduces the tensor to a scalar. - grad_output = float(gradient.tensor.contents.data[0]) * self.input[0].ones_like() + grad_output = float(gradient[[0] * len(gradient.shape)]) * self.input[0].ones_like() else: if self.keepdim: @@ -214,9 +214,9 @@ class MaxBackward: max_value = self.input[0].max() mask = self.input[0].equal(max_value) - grad_output = float(gradient.tensor.contents.data[0]) * self.input[0].ones_like() + grad_output = float(gradient[[0] * len(gradient.shape)]) * self.input[0].ones_like() - grad_output = (grad_output * mask) / mask.sum().tensor.contents.data[0] + grad_output = (grad_output * mask) / mask.sum()[0] else: @@ -250,9 +250,9 @@ class MinBackward: min_value = self.input[0].min() mask = self.input[0].equal(min_value) - grad_output = float(gradient.tensor.contents.data[0]) * self.input[0].ones_like() + grad_output = float(gradient[[0] * len(gradient.shape)]) * self.input[0].ones_like() - grad_output = (grad_output * mask) / mask.sum().tensor.contents.data[0] + grad_output = (grad_output * mask) / mask.sum()[0] else: if self.keepdim: