From c56c9c7519fb56fd2b3d8d8ddee35aaa0d07292f Mon Sep 17 00:00:00 2001 From: chenyu Date: Sun, 21 Jul 2024 17:30:37 -0400 Subject: [PATCH] move ufix inside UOp [run_process_replay] (#5621) dtype is always the dtype of the caller --- tinygrad/codegen/uops.py | 29 +++++++++++++++-------------- 1 file changed, 15 insertions(+), 14 deletions(-) diff --git a/tinygrad/codegen/uops.py b/tinygrad/codegen/uops.py index 7331012162..dcbc2f7eb8 100644 --- a/tinygrad/codegen/uops.py +++ b/tinygrad/codegen/uops.py @@ -28,7 +28,6 @@ class UOps(Enum): END_FOR_UOP = {UOps.IF:(UOps.STORE, UOps.ENDIF), UOps.RANGE:(UOps.PHI, UOps.ENDRANGE)} -def ufix(dtype: Optional[DType], x): return UOp.const(dtype, x) if not isinstance(x, UOp) else x @dataclass(frozen=True, eq=False) class UOp: op: UOps @@ -45,24 +44,26 @@ class UOp: self.arg.value, self.dtype, self.src) def __lt__(self, x:UOp): return self.cmp_tuple < x.cmp_tuple def __repr__(self): return pretty_print(self, lambda x: f"UOp({x.op}, {x.dtype}, arg={x.arg}, src=(%s))") + # *** uop syntactic sugar + def ufix(self, x): return UOp.const(self.dtype, x) if not isinstance(x, UOp) else x def cast(self, dtype=None): return UOp(UOps.CAST, dtype, (self,)) def bitcast(self, dtype=None): return UOp(UOps.BITCAST, dtype, (self,)) def name(self, name:Optional[str]): return UOp(UOps.VAR, src=(self,), arg=name) def __neg__(self): return UOp.alu(UnaryOps.NEG, self) - def __add__(self, x): return UOp.alu(BinaryOps.ADD, self, ufix(self.dtype, x)) - def __radd__(self, x): return UOp.alu(BinaryOps.ADD, ufix(self.dtype, x), self) - def __sub__(self, x): return UOp.alu(BinaryOps.ADD, self, -ufix(self.dtype, x)) - def __mul__(self, x): return UOp.alu(BinaryOps.MUL, self, ufix(self.dtype, x)) - def __rmul__(self, x): return UOp.alu(BinaryOps.MUL, ufix(self.dtype, x), self) - def __floordiv__(self, x): return UOp.alu(BinaryOps.IDIV, self, ufix(self.dtype, x)) - def __truediv__(self, x): return UOp.alu(BinaryOps.MUL, self, UOp.alu(UnaryOps.RECIP, ufix(self.dtype, x))) - def __mod__(self, x): return UOp.alu(BinaryOps.MOD, self, ufix(self.dtype, x)) - def __xor__(self, x): return UOp.alu(BinaryOps.XOR, self, ufix(self.dtype, x)) - def __and__(self, x): return UOp.alu(BinaryOps.AND, self, ufix(self.dtype, x)) - def __or__(self, x): return UOp.alu(BinaryOps.OR, self, ufix(self.dtype, x)) - def ne(self, x): return UOp.alu(BinaryOps.CMPNE, self, ufix(self.dtype, x)) + def __add__(self, x): return UOp.alu(BinaryOps.ADD, self, self.ufix(x)) + def __radd__(self, x): return UOp.alu(BinaryOps.ADD, self, self.ufix(x)) + def __sub__(self, x): return UOp.alu(BinaryOps.ADD, self, self.ufix(-x)) + def __mul__(self, x): return UOp.alu(BinaryOps.MUL, self, self.ufix(x)) + def __rmul__(self, x): return UOp.alu(BinaryOps.MUL, self.ufix(x), self) + def __floordiv__(self, x): return UOp.alu(BinaryOps.IDIV, self, self.ufix(x)) + def __truediv__(self, x): return UOp.alu(BinaryOps.MUL, self, UOp.alu(UnaryOps.RECIP, self.ufix(x))) + def __mod__(self, x): return UOp.alu(BinaryOps.MOD, self, self.ufix(x)) + def __xor__(self, x): return UOp.alu(BinaryOps.XOR, self, self.ufix(x)) + def __and__(self, x): return UOp.alu(BinaryOps.AND, self, self.ufix(x)) + def __or__(self, x): return UOp.alu(BinaryOps.OR, self, self.ufix(x)) + def ne(self, x): return UOp.alu(BinaryOps.CMPNE, self, self.ufix(x)) def eq(self, x): return -self.ne(x) - def lt(self, x): return UOp.alu(BinaryOps.CMPLT, self, ufix(self.dtype, x)) + def lt(self, x): return UOp.alu(BinaryOps.CMPLT, self, self.ufix(x)) def ge(self, x): return -self.lt(x) def max(self, x): return UOp.alu(BinaryOps.MAX, self, x) def min(self, x): return -UOp.alu(BinaryOps.MAX, -self, -x)