From 8950942e75fe861a615b53696db6facbfd0f7ca7 Mon Sep 17 00:00:00 2001 From: chenyu Date: Fri, 21 Aug 2026 22:27:36 -0400 Subject: [PATCH] remove explicit dtype for NOOP and decomp [PR] (#17678) --- tinygrad/codegen/decomp/dtype.py | 4 ++-- tinygrad/engine/jit.py | 2 +- tinygrad/schedule/rangeify.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tinygrad/codegen/decomp/dtype.py b/tinygrad/codegen/decomp/dtype.py index 00f7e8a059..7dcdf4dd04 100644 --- a/tinygrad/codegen/decomp/dtype.py +++ b/tinygrad/codegen/decomp/dtype.py @@ -139,7 +139,7 @@ pm_long_decomp = PatternMatcher([ (UPat(GroupOp.Defines, src=(UPat.var("sz"),), name="x"), lambda x,sz: x.replace(dtype=l2i_dt[x.dtype], arg=replace(x.arg, dtype=l2i_dt[x.dtype]), src=(sz*2,)) if x.dtype in l2i_dt else None), (UPat(Ops.INDEX, tuple(l2i_dt.keys()), name='x'), lambda x: - reindex(x, x.tag[0]).replace(dtype=x.tag[1], tag=None) if x.tag is not None else None), + reindex(x, x.tag[0]).replace(tag=None) if x.tag is not None else None), (UPat(Ops.STORE, src=(UPat.var('idx', tuple(l2i_dt.keys())), UPat.var('val')), name='st'), lambda st,idx,val: st.replace(src=(idx.rtag((0, dt:=l2i_dt[idx.dtype])), val.rtag((0, dt)))).group( st.replace(src=(idx.rtag((1, dt)), val.rtag((1, dt))))) if val.tag is None else None), @@ -160,7 +160,7 @@ pm_long_decomp = PatternMatcher([ split_l2i(ctx, x.op, l2i_dt[x.dtype], *flatten((a.rtag((0, l2i_dt[x.dtype])), a.rtag((1, l2i_dt[x.dtype]))) for a in x.src))[x.tag[0]] if x.tag is not None else None), (UPat(Ops.LOAD, tuple(l2i_dt.keys()), src=(UPat.var('idx'),), name='x'), lambda x,idx: - x.replace(dtype=l2i_dt[x.dtype], src=(reindex(idx, x.tag[0]).replace(dtype=l2i_dt[x.dtype], tag=None),), tag=None) if x.tag is not None else None), + x.replace(dtype=l2i_dt[x.dtype], src=(reindex(idx, x.tag[0]).replace(tag=None),), tag=None) if x.tag is not None else None), (UPat(Ops.CONST, tag={(w, dt) for w in (0, 1) for dt in l2i_dt.values()}, name='x'), lambda x: UOp.const(truncate[x.tag[1]]((x.val >> 32) if x.tag[0] == 1 else (x.val & 0xFFFFFFFF)), x.tag[1])) ]) diff --git a/tinygrad/engine/jit.py b/tinygrad/engine/jit.py index a543623072..221889f7d6 100644 --- a/tinygrad/engine/jit.py +++ b/tinygrad/engine/jit.py @@ -211,7 +211,7 @@ def _prepare_jit_inputs(args, kwargs): # collect buffer UOps (including MultiBuffer) input_buf_uops: list[UOp] = [u.base for u in input_uops if u.base.realized is not None] if len(set(input_buf_uops)) != len(input_buf_uops): raise JitError("duplicate inputs to JIT") - inputs = [(*(u.substitute({u.base:UOp(Ops.NOOP, u.base.dtype)}, extra_pm=mop_cleanup).unbind_all()), u.dtype, u.device) for u in input_uops] + inputs = [(*(u.substitute({u.base:UOp(Ops.NOOP)}, extra_pm=mop_cleanup).unbind_all()), u.dtype, u.device) for u in input_uops] _var_vals = merge_dicts([x[1] for x in inputs] + [dict(v.unbind() for v in (args + tuple(kwargs.values())) if isinstance(v, UOp))]) var_vals = {k.expr:v for k,v in _var_vals.items()} expected_input_info = [(x[0], tuple(sorted(x[1].keys(), key=lambda v: v.expr)), x[2], x[3]) for x in inputs] diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 979cef04c0..4751e44212 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -79,7 +79,7 @@ def split_reduceop(reduce:UOp, x:UOp): # get expanded by rangeifying the UOp x indexed = x.index(*[UOp.range(s, i) if resolve(s>1) else 0 for i,s in enumerate(x.shape)]) - range_nums = [y.arg[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP, x.base.dtype)}, extra_pm=pm_mops).ranges] + range_nums = [y.arg[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP)}, extra_pm=pm_mops).ranges] is_expanded = [i not in range_nums for i in range(len(x.shape))] if not (split_candidates:=[(i,d) for i in range(reduce.arg[1])