From 437205ae0384352e6a7cf7004e5005be709b1dbf Mon Sep 17 00:00:00 2001 From: George Hotz Date: Mon, 4 May 2026 20:02:41 -0700 Subject: [PATCH] flip order, this is simpler --- test/null/test_uop_graph.py | 6 +++--- tinygrad/codegen/late/gater.py | 2 +- tinygrad/renderer/cstyle.py | 4 ++-- tinygrad/renderer/llvmir.py | 2 +- tinygrad/renderer/ptx.py | 4 ++-- tinygrad/runtime/ops_python.py | 2 +- tinygrad/uop/spec.py | 2 +- 7 files changed, 11 insertions(+), 11 deletions(-) diff --git a/test/null/test_uop_graph.py b/test/null/test_uop_graph.py index bb7414f35a..ea84737d98 100644 --- a/test/null/test_uop_graph.py +++ b/test/null/test_uop_graph.py @@ -428,7 +428,7 @@ class TestUOpGraph(unittest.TestCase): uops = to_uops_list([w, red]) for u in uops: assert u.op is not Ops.WHERE - if u.op is Ops.LOAD and u.src[0].src[0].op is Ops.PARAM: assert u.src[2].arg==5 + if u.op is Ops.LOAD and u.src[0].src[0].op is Ops.PARAM: assert u.src[1].arg==5 def test_where_on_gated_load_folds_swapped_branches(self): ridx0 = UOp.range(100, 0) @@ -438,7 +438,7 @@ class TestUOpGraph(unittest.TestCase): uops = to_uops_list([w]) for u in uops: assert u.op is not Ops.WHERE - if u.op is Ops.LOAD: assert u.src[2].arg==5 + if u.op is Ops.LOAD: assert u.src[1].arg==5 def test_where_on_gated_load_with_cast(self): ridx0 = UOp.range(100, 0) @@ -451,7 +451,7 @@ class TestUOpGraph(unittest.TestCase): uops = to_uops_list([w, red]) for u in uops: assert u.op is not Ops.WHERE - if u.op is Ops.LOAD and u.src[0].src[0].op is Ops.PARAM: assert u.src[2].arg == 5 + if u.op is Ops.LOAD and u.src[0].src[0].op is Ops.PARAM: assert u.src[1].arg == 5 def test_where_on_casted_gated_load_extra_cond(self): ridx0 = UOp.range(100, 0) diff --git a/tinygrad/codegen/late/gater.py b/tinygrad/codegen/late/gater.py index 24a02a8f1e..e3808bc5e5 100644 --- a/tinygrad/codegen/late/gater.py +++ b/tinygrad/codegen/late/gater.py @@ -3,7 +3,7 @@ from tinygrad.dtype import Invalid pm_move_gates_from_index = PatternMatcher([ (UPat.var("buf").index(UPat.var("gate").where(UPat.var("idx"), UPat(arg=Invalid))).or_casted(name="cast").load(name="l"), - lambda buf,gate,idx,cast,l: buf.index(idx).cast(cast.dtype).load(gate, l.const_like(0), dtype=l.dtype)), + lambda buf,gate,idx,cast,l: buf.index(idx).cast(cast.dtype).load(l.const_like(0), gate, dtype=l.dtype)), (UPat.var("buf").index(UPat.var("gate").where(UPat.var("idx"), UPat(arg=Invalid))).or_casted(name="cast").store(UPat.var("data")), lambda buf,gate,idx,cast,data: buf.index(idx).cast(cast.dtype).store(data, gate)), ]) diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 1654521d14..8f526b94d7 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -47,7 +47,7 @@ base_rewrite = PatternMatcher([ (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var('idx')), allow_any_len=True), lambda ctx,buf,idx: f"({ctx[buf]}+{strip_parens(ctx[idx]) if idx.arg == Ops.ADD else ctx[idx]})"), (UPat(Ops.LOAD, src=(UPat.var('bidx'),)), lambda ctx,bidx: f"(*{ctx[bidx]})"), - (UPat(Ops.LOAD, src=(UPat.var("bidx"), UPat.var("gate"), UPat.var("var"))), lambda ctx,bidx,var,gate: f"({ctx[gate]}?*{ctx[bidx]}:{ctx[var]})"), + (UPat(Ops.LOAD, src=(UPat.var("bidx"), UPat.var("var"), UPat.var("gate"))), lambda ctx,bidx,var,gate: f"({ctx[gate]}?*{ctx[bidx]}:{ctx[var]})"), (UPat(Ops.STORE, src=(UPat.var('bidx'), UPat.var("var")), allow_any_len=True), lambda ctx,bidx,var: f"*{ctx[bidx]} = {ctx[var]};"), # alu/gep # TODO: look for left-associative @@ -301,7 +301,7 @@ class OpenCLRenderer(CStyleLanguage): (UPat(Ops.CONST, dtypes.bfloat16, name="x"), lambda ctx,x: f"{(struct.unpack('I', struct.pack('f', float_to_bf16(x.arg)))[0] >> 16)}u"), # load/store image (OpenCL) - (UPat(Ops.LOAD, dtype=dtypes.float.vec(4), src=(UPat.var('buf').index(UPat.var('idx', dtypes.int.vec(2))), UPat.var("gate"), UPat.var("var"))), + (UPat(Ops.LOAD, dtype=dtypes.float.vec(4), src=(UPat.var('buf').index(UPat.var('idx', dtypes.int.vec(2))), UPat.var("var"), UPat.var("gate"))), lambda ctx,buf,idx,var,gate: f"({ctx[gate]}?read_imagef({ctx[buf]}, smp, {ctx[idx]}):{ctx[var]})"), (UPat(Ops.LOAD, dtype=dtypes.float.vec(4), src=(UPat.var('buf').index(UPat.var('idx', dtypes.int.vec(2))),)), lambda ctx,buf,idx: f"read_imagef({ctx[buf]}, smp, {ctx[idx]})"), diff --git a/tinygrad/renderer/llvmir.py b/tinygrad/renderer/llvmir.py index 5008e4e660..a65f9ca951 100644 --- a/tinygrad/renderer/llvmir.py +++ b/tinygrad/renderer/llvmir.py @@ -76,7 +76,7 @@ base_rewrite = PatternMatcher([ # memory load/store (UPat(Ops.INDEX, name="x"), lambda ctx,x: f" {ctx[x]} = getelementptr inbounds {ldt(x.dtype.base)}, {ldt(x.src[0].dtype)} {ctx[x.src[0]]}, {ldt(x.src[1].dtype)} {ctx[x.src[1]]}"), - (UPat(Ops.LOAD, src=(UPat(Ops.INDEX).or_casted("idx"), UPat.var("mask"), UPat.var("alt")), allow_any_len=True, name="x"), + (UPat(Ops.LOAD, src=(UPat(Ops.INDEX).or_casted("idx"), UPat.var("alt"), UPat.var("mask")), allow_any_len=True, name="x"), lambda ctx,x,idx,alt,mask: f" br label {ctx[x]}_entry\n{ctx[x][1:]}_entry:\n" f" br i1 {ctx[mask]}, label {ctx[x]}_load, label {ctx[x]}_exit\n{ctx[x][1:]}_load:\n" diff --git a/tinygrad/renderer/ptx.py b/tinygrad/renderer/ptx.py index 3798b8c43c..e07262418a 100644 --- a/tinygrad/renderer/ptx.py +++ b/tinygrad/renderer/ptx.py @@ -47,7 +47,7 @@ ptx_matcher = PatternMatcher([ lambda x: (UOp(x.op, dtypes.float32, tuple(vv.cast(dtypes.float32) for vv in x.src), x.arg).cast(dtypes.half))), # load/store bool -> uint8 (UPat(Ops.LOAD, dtypes.bool, src=(UPat(dtype=dtypes.int64),), name="x", allow_any_len=True), - lambda x: UOp(x.op, dtypes.uint8, x.src[0:2] + ((x.src[2].cast(dtypes.uint8),) if len(x.src) >= 3 else ()) + x.src[3:]).cast(dtypes.bool)), + lambda x: UOp(x.op, dtypes.uint8, x.src[0:1] + ((x.src[1].cast(dtypes.uint8),) if len(x.src) >= 2 else ()) + x.src[2:]).cast(dtypes.bool)), (UPat(Ops.STORE, src=(UPat(dtype=dtypes.int64), UPat(dtype=dtypes.bool)), name="x", allow_any_len=True), lambda x: UOp(x.op, dtypes.void, (x.src[0], x.src[1].cast(dtypes.uint8)))), # indexing on PTX is in uint64, we do the math while it's still in the graph @@ -106,7 +106,7 @@ string_rewrite = PatternMatcher([ lambda ctx, loc, var, buf: f"st.{mem_type(buf)}" + \ f"{f'.v{cnt}' if ((cnt:=var.dtype.count)>1) else ''}.{ctx.mem_types[var.dtype.scalar()]} " + \ f"[{ctx.r[loc]}+0], {('{' + ', '.join(ctx.r[var]) + '}') if var.dtype.count > 1 else ctx.r[var]};"), - (UPat(Ops.LOAD, name="x", src=(UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("loc"))), UPat.var("gate"), UPat.var("alt")), allow_any_len=True), + (UPat(Ops.LOAD, name="x", src=(UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("loc"))), UPat.var("alt"), UPat.var("gate")), allow_any_len=True), lambda ctx, x, loc, alt, gate, buf: flatten([ [f"mov.{ctx.mem_types[x.dtype.scalar()]} {v}, {render_val(0, x.dtype.scalar())};" for v in ctx.r[x]], [f"@{ctx.r[gate]} ld.{mem_type(buf)}.v{x.dtype.count}.{ctx.mem_types[x.dtype.scalar()]} {{{', '.join(ctx.r[x])}}}, [{ctx.r[loc]}+0];"] diff --git a/tinygrad/runtime/ops_python.py b/tinygrad/runtime/ops_python.py index ae9ec7d432..a891fe31a0 100644 --- a/tinygrad/runtime/ops_python.py +++ b/tinygrad/runtime/ops_python.py @@ -18,7 +18,7 @@ def _load(m, i, dtype: DType): return from_storage_scalar(m[i], dtype) def load(inp, j, dtype: DType): - if len(inp) >= 3: return [_load(m, x+j if x is not None else None, dtype) if gate else default for (m,x),gate,default in zip(*inp[:3])] + if len(inp) >= 3: return [_load(m, x+j if x is not None else None, dtype) if gate else default for (m,x),default,gate in zip(*inp[:3])] return [_load(m, x+j if x is not None else None, dtype) for m,x in inp[0]] def _store(m, i, v, dtype: DType): diff --git a/tinygrad/uop/spec.py b/tinygrad/uop/spec.py index 4a63e5b06f..60ec10b97c 100644 --- a/tinygrad/uop/spec.py +++ b/tinygrad/uop/spec.py @@ -176,7 +176,7 @@ shared_codegen_spec = PatternMatcher([ # LOAD(idx) / STORE(idx, val) with gates on the LOAD/STORE (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("idx"))).or_casted().load(), validate_index), - (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("idx"))).or_casted().load(UPat.var("gate", dtype=dtypes.bool), UPat()), validate_index), + (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("idx"))).or_casted().load(UPat(), UPat.var("gate", dtype=dtypes.bool)), validate_index), (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("idx"))).or_casted().store(UPat()), validate_index), (UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("idx"))).or_casted().store(UPat(), UPat.var("gate", dtype=dtypes.bool)), validate_index),