From 004b2ecff5a3dc4867656a942fe076a54f079db0 Mon Sep 17 00:00:00 2001 From: chenyu Date: Wed, 4 Dec 2024 00:24:33 -0500 Subject: [PATCH] remove lt/gt/le/ge from SimpleMathTrait [pr] (#8027) just use the dunder methods --- tinygrad/ops.py | 16 ++++++---------- tinygrad/tensor.py | 6 +++--- 2 files changed, 9 insertions(+), 13 deletions(-) diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 0d49080f22..8f00161d18 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -57,18 +57,14 @@ class SimpleMathTrait: def __ror__(self, x): return self.bitwise_or(x, True) def __rxor__(self, x): return self.xor(x, True) - def lt(self, x): return self.alu(Ops.CMPLT, self.ufix(x)) - def gt(self, x): return self.ufix(x).alu(Ops.CMPLT, self) - def ne(self, x): return self.alu(Ops.CMPNE, self.ufix(x)) - def ge(self, x): return self.lt(x).logical_not() - def le(self, x): return self.gt(x).logical_not() - def eq(self, x): return self.ne(x).logical_not() + def __lt__(self, x): return self.alu(Ops.CMPLT, self.ufix(x)) + def __gt__(self, x): return self.ufix(x).alu(Ops.CMPLT, self) + def __ge__(self, x): return (self < x).logical_not() + def __le__(self, x): return (self > x).logical_not() - def __lt__(self, x): return self.lt(x) - def __gt__(self, x): return self.gt(x) + def ne(self, x): return self.alu(Ops.CMPNE, self.ufix(x)) + def eq(self, x): return self.ne(x).logical_not() def __ne__(self, x): return self.ne(x) - def __ge__(self, x): return self.ge(x) - def __le__(self, x): return self.le(x) # NOTE: __eq__ isn't overridden, and means the same thing as is by default class MathTrait(SimpleMathTrait): diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 253a9839a0..f7986cf0a0 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -3206,7 +3206,7 @@ class Tensor(SimpleMathTrait): """ return -((-self).maximum(-x)) - def where(self:Tensor, x:Union[Tensor, ConstType], y:Union[Tensor, ConstType]): + def where(self:Tensor, x:Union[Tensor, ConstType, sint], y:Union[Tensor, ConstType, sint]): """ Return a tensor of elements selected from either `x` or `y`, depending on `self`. `output_i = x_i if self_i else y_i`. @@ -3258,8 +3258,8 @@ class Tensor(SimpleMathTrait): def __ilshift__(self, x) -> Tensor: return self.assign(self.lshift(x)) def __irshift__(self, x) -> Tensor: return self.assign(self.rshift(x)) - def lt(self, x) -> Tensor: return F.Less.apply(*self._broadcasted(x, False)) - def gt(self, x) -> Tensor: return F.Less.apply(*self._broadcasted(x, True)) + def __lt__(self, x) -> Tensor: return F.Less.apply(*self._broadcasted(x, False)) + def __gt__(self, x) -> Tensor: return F.Less.apply(*self._broadcasted(x, True)) def ne(self, x) -> Tensor: return F.Neq.apply(*self._broadcasted(x)) def __eq__(self, x) -> Tensor: return self.eq(x) # type: ignore[override]