forked from tinygrad/tinygrad
remove lt/gt/le/ge from SimpleMathTrait [pr] (#8027)
just use the dunder methods
This commit is contained in:
+6
-10
@@ -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):
|
||||
|
||||
+3
-3
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user