From 7cc973a4e033e5c497eaa96de1b5b286103519e6 Mon Sep 17 00:00:00 2001 From: chenyu Date: Sun, 30 Aug 2026 09:28:21 -0400 Subject: [PATCH] clean up unneeded dtype check in rules [PR] (#17845) --- tinygrad/codegen/decomp/op.py | 2 +- tinygrad/codegen/simplify.py | 2 +- tinygrad/renderer/isa/x86.py | 10 +++++----- tinygrad/schedule/prepare.py | 2 +- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/tinygrad/codegen/decomp/op.py b/tinygrad/codegen/decomp/op.py index 4e1d5f81c4..1604321225 100644 --- a/tinygrad/codegen/decomp/op.py +++ b/tinygrad/codegen/decomp/op.py @@ -84,7 +84,7 @@ def get_simplifying_rewrite_patterns(ops:tuple[Ops, ...]) -> PatternMatcher: if Ops.AND in ops: pat.append((UPat.var("x", dtypes.ints)%UPat.cvar("c"), lambda x,c: x & (c.val-1) if c.val in powers_of_two else None)) pat.append((UPat.var("a")%UPat.var("b"), floormod_to_mod)) # no real hardware supports THREEFRY, but NullRenderer does - if Ops.THREEFRY not in ops: pat.append((UPat(Ops.THREEFRY, dtype=dtypes.uint64, src=(UPat.var("x"), UPat.var("key"))), threefry2x32)) + if Ops.THREEFRY not in ops: pat.append((UPat(Ops.THREEFRY, src=(UPat.var("x"), UPat.var("key"))), threefry2x32)) # MAX can be rewritten as CMPLT + WHERE (max function is annoying on many cstyle backends) if Ops.MAX not in ops and Ops.CMPLT in ops: pat.append((UPat(Ops.MAX, name="m"), lambda m: (m.src[0] < m.src[1]).where(m.src[1], m.src[0]))) return PatternMatcher(pat) diff --git a/tinygrad/codegen/simplify.py b/tinygrad/codegen/simplify.py index 582e721662..21c928a769 100644 --- a/tinygrad/codegen/simplify.py +++ b/tinygrad/codegen/simplify.py @@ -84,7 +84,7 @@ def reduce_unparented(red:UOp) -> UOp|None: assert all(x.op is Ops.RANGE for x in red.src[1:]), "some reduce srcs aren't ranges" reduce_parented, reduce_unparented = partition(red.src[1:], lambda x: x in red.src[0].ranges) if len(reduce_unparented) == 0: return None - ret = red.replace(src=(red.src[0],)+tuple(reduce_parented)) if len(reduce_parented) or red.dtype != red.src[0].dtype else red.src[0] + ret = red.replace(src=(red.src[0],)+tuple(reduce_parented)) if len(reduce_parented) else red.src[0] if red.arg[0] is Ops.ADD: for r in reduce_unparented: ret = ret * r.src[0] if red.arg[0] is Ops.MUL: diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index 3bef96b2fa..042486d3d7 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -377,7 +377,7 @@ isel_matcher = PatternMatcher([ (UPat(GroupOp.Comparison, src=(UPat(dtype=dtypes.float64), UPat()), name="m").where(UPat.var("a", dtypes.float64), UPat.var("b")), lambda m,a,b: a.ins(X86Ops.VBLENDVPD, src=(b, a, mask(m)))), # in this case we have a mask producing comparison whose user expects a bool, so we convert to bool - (UPat(GroupOp.Comparison, dtypes.bool, (UPat.var("y", (dtypes.float32, dtypes.float64)), UPat()), name="x"), lambda y,x: + (UPat(GroupOp.Comparison, src=(UPat.var("y", (dtypes.float32, dtypes.float64)), UPat()), name="x"), lambda y,x: UOp(Ops.AND, src=(mask(x).bitcast(dt:=to_int(y.dtype)), UOp.cconst(1, dt))).bitcast(dtypes.bool)), # conditional moves that use flags # TODO: remove this once we allow all flag producing ops in cmove @@ -394,10 +394,10 @@ isel_matcher = PatternMatcher([ (UPat(Ops.IF, src=(UPat(Ops.CMPEQ, name="y"),), name="x"), lambda y,x: x.ins(X86Ops.JE, src=(cmp(y),))), (UPat(Ops.IF, src=(UPat(Ops.CMPNE, name="y"),), name="x"), lambda y,x: x.ins(X86Ops.JNE, src=(cmp(y),))), # comparisons whose user doesn't use the flag, move flag result to register - (UPat(Ops.CMPLT, dtypes.bool, (UPat(dtype=dtypes.uints), UPat()), name="x"), lambda x: x.ins(X86Ops.SETB, src=(cmp(x),))), - (UPat(Ops.CMPLT, dtypes.bool, name="x"), lambda x: x.ins(X86Ops.SETL, src=(cmp(x),))), - (UPat(Ops.CMPEQ, dtypes.bool, name="x"), lambda x: x.ins(X86Ops.SETE, src=(cmp(x),))), - (UPat(Ops.CMPNE, dtypes.bool, name="x"), lambda x: x.ins(X86Ops.SETNE, src=(cmp(x),))), + (UPat(Ops.CMPLT, src=(UPat(dtype=dtypes.uints), UPat()), name="x"), lambda x: x.ins(X86Ops.SETB, src=(cmp(x),))), + (UPat(Ops.CMPLT, name="x"), lambda x: x.ins(X86Ops.SETL, src=(cmp(x),))), + (UPat(Ops.CMPEQ, name="x"), lambda x: x.ins(X86Ops.SETE, src=(cmp(x),))), + (UPat(Ops.CMPNE, name="x"), lambda x: x.ins(X86Ops.SETNE, src=(cmp(x),))), # float unary (UPat.var("y", dtypes.float32).sqrt().named("x"), lambda y,x: x.ins(X86Ops.VSQRTSS, src=(y, y)) if x.max_numel() == 1 else x.ins(X86Ops.VSQRTPS)), (UPat.var("y", dtypes.float64).sqrt().named("x"), lambda y,x: x.ins(X86Ops.VSQRTSD, src=(y, y)) if x.max_numel() == 1 else x.ins(X86Ops.VSQRTPD)), diff --git a/tinygrad/schedule/prepare.py b/tinygrad/schedule/prepare.py index 8d9f4bc398..037afb6df7 100644 --- a/tinygrad/schedule/prepare.py +++ b/tinygrad/schedule/prepare.py @@ -38,7 +38,7 @@ def _mop_index(r:UOp, idx:UOp): if r.op is Ops.RESHAPE: src_prefix = len(r.src[0].shape) - len(r.shape[len(idxs):]) if src_prefix >= 0 and r.src[0].shape[src_prefix:] == r.shape[len(idxs):]: - if src_prefix == 0: return r.src[0] if r.src[0].dtype == idx.dtype else None + if src_prefix == 0: return r.src[0] ret = r.src[0].index(*apply_movement_op(r.op, r.src[0].shape[:src_prefix], r.shape[:len(idxs)], idxs), arg=idx.arg) return ret if ret.shape == idx.shape else None