remove lt/gt/le/ge from SimpleMathTrait [pr] (#8027)

just use the dunder methods
This commit is contained in:
chenyu
2024-12-04 00:24:33 -05:00
committed by GitHub
parent 39e0fc05f5
commit 004b2ecff5
2 changed files with 9 additions and 13 deletions
+6 -10
View File
@@ -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
View File
@@ -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]