From 693990a34639afd115e422078a4e4d7482bdae83 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Tue, 30 Jul 2024 14:04:13 -0700 Subject: [PATCH] swap src[2] and src[3] in load [run_process_replay] (#5821) * swap src[2] and src[3] in load [run_process_replay] * cleanups + bugfix * fix ptx --- test/test_uop_graph.py | 14 +++++++------- tinygrad/codegen/lowerer.py | 4 ++-- tinygrad/codegen/uopgraph.py | 22 +++++++++++----------- tinygrad/codegen/uops.py | 2 +- tinygrad/renderer/assembly.py | 10 +++++----- tinygrad/renderer/cstyle.py | 2 +- tinygrad/renderer/llvmir.py | 4 ++-- tinygrad/runtime/ops_python.py | 2 +- 8 files changed, 30 insertions(+), 30 deletions(-) diff --git a/test/test_uop_graph.py b/test/test_uop_graph.py index ecb16fc460..8d0754915e 100644 --- a/test/test_uop_graph.py +++ b/test/test_uop_graph.py @@ -231,8 +231,8 @@ class TestUOpGraph(TestUOps): glbl1 = UOp(UOps.DEFINE_GLOBAL, PtrDType(dtypes.int), (), (1, False)) glbl2 = UOp(UOps.DEFINE_GLOBAL, PtrDType(dtypes.int), (), (2, False)) idx = UOp.const(dtypes.int, 0) - ld0 = UOp(UOps.LOAD, dtypes.int, (glbl1, idx, UOp.const(dtypes.bool, False), UOp.const(dtypes.int, 2))) - ld1 = UOp(UOps.LOAD, dtypes.int, (glbl2, idx, UOp.const(dtypes.bool, True), UOp.const(dtypes.int, 3))) + ld0 = UOp(UOps.LOAD, dtypes.int, (glbl1, idx, UOp.const(dtypes.int, 2), UOp.const(dtypes.bool, False))) + ld1 = UOp(UOps.LOAD, dtypes.int, (glbl2, idx, UOp.const(dtypes.int, 3), UOp.const(dtypes.bool, True))) uops = UOpGraph([UOp(UOps.STORE, None, (glbl0, idx, ld1+ld0))]) ld0, ld1 = uops[-1].src[2].src # ld0 becomes the invalid value @@ -246,8 +246,8 @@ class TestUOpGraph(TestUOps): lidx = UOp(UOps.SPECIAL, dtypes.int, (), ("lidx0", 16)) st = UOp(UOps.STORE, None, (smem, lidx, UOp.load(glbl0, lidx, dtype=dtypes.int))) barrier = UOp(UOps.BARRIER, None, (st, )) - ld0 = UOp(UOps.LOAD, dtypes.int, (smem, lidx+1, UOp.const(dtypes.bool, False), UOp.const(dtypes.int, 2), barrier)) - ld1 = UOp(UOps.LOAD, dtypes.int, (smem, lidx+2, UOp.const(dtypes.bool, True), UOp.const(dtypes.int, 3), barrier)) + ld0 = UOp(UOps.LOAD, dtypes.int, (smem, lidx+1, UOp.const(dtypes.int, 2), UOp.const(dtypes.bool, False), barrier)) + ld1 = UOp(UOps.LOAD, dtypes.int, (smem, lidx+2, UOp.const(dtypes.int, 3), UOp.const(dtypes.bool, True), barrier)) uops = UOpGraph([UOp(UOps.STORE, None, (glbl0, lidx, ld1+ld0))]) ld0, ld1 = uops[-1].src[2].src # ld0 becomes the invalid value @@ -438,18 +438,18 @@ class TestLoadStoreFolder(unittest.TestCase): def test_simple_load_fold_gated(self): buf = UOp(UOps.DEFINE_GLOBAL, PtrDType(dtypes.float)) gate = UOp(UOps.DEFINE_VAR, dtypes.bool) - load = [UOp(UOps.LOAD, dtypes.float, (buf, UOp.const(dtypes.int, i), gate, UOp.const(dtypes.float, i))) for i in range(4)] + load = [UOp(UOps.LOAD, dtypes.float, (buf, UOp.const(dtypes.int, i), UOp.const(dtypes.float, i), gate)) for i in range(4)] sink = UOp(UOps.EXPAND, dtypes.float, tuple(load), ((0,4),)) sink = float4_rewrite(sink) assert len([x for x in sink.sparents if x.op is UOps.LOAD]) == 1 single_load = [x for x in sink.sparents if x.op is UOps.LOAD][0] - self.assertListEqual([src.arg for src in single_load.src[3].src], [0.0, 1.0, 2.0, 3.0]) + self.assertListEqual([src.arg for src in single_load.src[2].src], [0.0, 1.0, 2.0, 3.0]) def test_simple_load_dont_fold_different_gated(self): buf = UOp(UOps.DEFINE_GLOBAL, PtrDType(dtypes.float)) gate = UOp(UOps.DEFINE_VAR, dtypes.bool, arg="g1") gate2 = UOp(UOps.DEFINE_VAR, dtypes.bool, arg="g2") - load = [UOp(UOps.LOAD, dtypes.float, (buf, UOp.const(dtypes.int, i), gate if i == 0 else gate2, UOp.const(dtypes.float, i))) for i in range(4)] + load = [UOp(UOps.LOAD, dtypes.float, (buf, UOp.const(dtypes.int, i), UOp.const(dtypes.float, i), gate if i == 0 else gate2)) for i in range(4)] sink = UOp(UOps.EXPAND, dtypes.float, tuple(load), ((0,4),)) sink = float4_rewrite(sink) assert len([x for x in sink.sparents if x.op is UOps.LOAD]) == 3 diff --git a/tinygrad/codegen/lowerer.py b/tinygrad/codegen/lowerer.py index a6420c7cbd..41f0d161e9 100644 --- a/tinygrad/codegen/lowerer.py +++ b/tinygrad/codegen/lowerer.py @@ -161,10 +161,10 @@ class IndependentLowerer: if idx.dtype == dtypes.int.vec(3): # this should all simplify if there's consts for id4. if not, w/e idx, id4 = UOp(UOps.VECTORIZE, dtypes.int.vec(2), (idx.src[0], idx.src[1])), idx.src[2] - vec_load = UOp(UOps.LOAD, load_dtype.vec(4), (buf, idx) + ((valid, UOp.const(load_dtype.vec(4), 0)) if has_valid else ()) + barrier) + vec_load = UOp(UOps.LOAD, load_dtype.vec(4), (buf, idx) + ((UOp.const(load_dtype.vec(4), 0), valid) if has_valid else ()) + barrier) return functools.reduce(lambda ret, i: id4.ne(i).where(ret, UOp(UOps.GEP, load_dtype, (vec_load,), i)), range(4), UOp.const(load_dtype, float('nan'))) - return UOp(UOps.LOAD, load_dtype, (buf, idx) + ((valid, UOp.const(load_dtype, 0)) if has_valid else ()) + barrier) + return UOp(UOps.LOAD, load_dtype, (buf, idx) + ((UOp.const(load_dtype, 0), valid) if has_valid else ()) + barrier) # NOTE: only store the local reduceop in the first thread (this is wrong for non group for reduces!) if x.arg.idx >= 0: for oidx, ridx in zip(self.idxs, self.ridxs): diff --git a/tinygrad/codegen/uopgraph.py b/tinygrad/codegen/uopgraph.py index 53771be9c2..b394dc378c 100644 --- a/tinygrad/codegen/uopgraph.py +++ b/tinygrad/codegen/uopgraph.py @@ -29,7 +29,7 @@ def fold_expanded(ex, buf): # add idx and idy for image if is_image: root_src = (s.src[1].src[0:2], root_src) # add gates for gated - if len(s.src) >= 4: root_src = (s.src[2] if is_load else s.src[3], root_src) # maybe flip the gate and the const? + if len(s.src) >= 4: root_src = (s.src[3], root_src) assert arg not in offsets_rootsrc[root_src] offsets_rootsrc[root_src][arg] = i @@ -44,15 +44,15 @@ def fold_expanded(ex, buf): if not is_image and not new_src[1].divides(fold_length): continue # for images, we rewrite the index if is_image: new_src[1] = UOp(UOps.VECTORIZE, dtypes.int.vec(2), (new_src[1].src[0], new_src[1].src[1])) + # vectorize the store/loadconst + if not is_load or len(new_src) >= 4: + new_src[2] = UOp(UOps.VECTORIZE, new_src[2].dtype.vec(fold_length), tuple(new_srcs[offsets[o+i]].src[2] for i in range(fold_length))) + # generate the folded new_srcs if is_load: - # vectorize the const. if we flip const and gate it's nicer here too - if len(new_src) >= 4: - new_src[3] = UOp(UOps.VECTORIZE, load_1.dtype.vec(fold_length), tuple(new_srcs[offsets[o+i]].src[3] for i in range(fold_length))) - new_load = UOp(load_1.op, load_1.dtype.vec(fold_length), tuple(new_src), load_1.arg) + new_load = UOp(UOps.LOAD, load_1.dtype.vec(fold_length), tuple(new_src)) for i in range(fold_length): new_srcs[offsets[o+i]] = UOp(UOps.GEP, load_1.dtype, (new_load,), i) else: - new_src[2] = UOp(UOps.VECTORIZE, new_src[2].dtype.vec(fold_length), tuple(new_srcs[offsets[o+i]].src[2] for i in range(fold_length))) - for i in range(fold_length): new_srcs[offsets[o+i]] = UOp(load_1.op, None, tuple(new_src), load_1.arg) if i == 0 else None + for i in range(fold_length): new_srcs[offsets[o+i]] = UOp(UOps.STORE, None, tuple(new_src)) if i == 0 else None for i in range(fold_length): used.add((rootsrc,o+i)) # dedup expand for LOAD @@ -286,11 +286,11 @@ constant_folder = PatternMatcher([ (NOp(UOps.CAST, name="root"), lambda root: root.src[0] if str(root.dtype) == str(root.src[0].dtype) else None), (NOp(UOps.VECTORIZE, name="root"), lambda root: root.src[0] if str(root.dtype) == str(root.src[0].dtype) else None), # fold gated LOAD/STORE - (NOp.load(NOp.var("buf"), NOp.var("idx"), NOp.const(dtypes.bool, True), NOp.cvar("var")), lambda buf,idx,var: UOp.load(buf, idx, dtype=var.dtype)), - (NOp.load(NOp.var("buf"), NOp.var("idx"), NOp.const(dtypes.bool, True), NOp.cvar("var"), NOp.var("barrier")), + (NOp.load(NOp.var("buf"), NOp.var("idx"), NOp.cvar("var"), NOp.const(dtypes.bool, True)), lambda buf,idx,var: UOp.load(buf, idx, dtype=var.dtype)), + (NOp.load(NOp.var("buf"), NOp.var("idx"), NOp.cvar("var"), NOp.const(dtypes.bool, True), NOp.var("barrier")), lambda buf,idx,var,barrier: UOp.load(buf, idx, barrier, dtype=var.dtype)), - (NOp.load(NOp.var(), NOp.var(), NOp.const(dtypes.bool, False), NOp.cvar("var")), lambda var: var), - (NOp.load(NOp.var(), NOp.var(), NOp.const(dtypes.bool, False), NOp.cvar("var"), NOp.var()), lambda var: var), + (NOp.load(NOp.var(), NOp.var(), NOp.cvar("var"), NOp.const(dtypes.bool, False)), lambda var: var), + (NOp.load(NOp.var(), NOp.var(), NOp.cvar("var"), NOp.const(dtypes.bool, False), NOp.var()), lambda var: var), (NOp.store(NOp.var("buf"), NOp.var("idx"), NOp.var("val"), NOp.const(dtypes.bool, True)), UOp.store), (NOp.store(NOp.var(), NOp.var(), NOp.var(), NOp.const(dtypes.bool, False)), lambda: UOp(UOps.NOOP)), # remove NOOPs from SINK diff --git a/tinygrad/codegen/uops.py b/tinygrad/codegen/uops.py index 6028de1315..e47c9fcb58 100644 --- a/tinygrad/codegen/uops.py +++ b/tinygrad/codegen/uops.py @@ -204,7 +204,7 @@ def type_verify(uops): if uop is UOps.VECTORIZE: assert dtype.count > 1 and len(src) == dtype.count, f"dtype vectorization mismatch {dtype.count=} != {len(src)=}" assert all(dtype == x.dtype.vec(len(src)) for x in src), f"{dtype=} must be {src[0].dtype.vec(len(src))}" - if uop is UOps.LOAD and len(src) > 3 and src[2].op is UOps.ALU: assert src[2].dtype == dtypes.bool and src[3].dtype == dtype + if uop is UOps.LOAD and len(src) > 3 and src[3].op is UOps.ALU: assert src[3].dtype == dtypes.bool and src[2].dtype == dtype if uop is UOps.GEP: assert dtype == src[0].dtype.scalar(), f"GEP of {src[0].dtype=} should be {src[0].dtype.scalar()} != {dtype}" if uop is UOps.STORE: assert dtype is None, f"{uop} dtype must be None, got {dtype}" diff --git a/tinygrad/renderer/assembly.py b/tinygrad/renderer/assembly.py index aad47f0065..1c9f0b91dd 100644 --- a/tinygrad/renderer/assembly.py +++ b/tinygrad/renderer/assembly.py @@ -176,16 +176,16 @@ class PTXRenderer(Renderer): assert src[0].dtype == dtypes.int64, "load isn't int64" assert src[1].op is UOps.CONST, f"load isn't const {u}" mem_type = '.shared' if src[0].op is UOps.DEFINE_LOCAL or any(x.op is UOps.DEFINE_LOCAL for x in src[0].parents) else '.global' - has_gate = len(src) > 3 and src[2].op is UOps.ALU + has_gate = len(src) > 3 and src[3].op is UOps.ALU if dtype.count > 1: r[u] = [ssa('val', dtype=self.types[dtype.scalar()]) for _ in range(dtype.count)] if has_gate: for v in r[u]: kk(f"mov.{self.mem_types[dtype.scalar()]} {v}, {render_val(0, dtype.scalar())};") - kk((f"@{r[src[2]]}"if has_gate else "") + kk((f"@{r[src[3]]}"if has_gate else "") + f" ld{mem_type}.v{dtype.count}.{self.mem_types[dtype.scalar()]} {{{', '.join(r[u])}}}, [{r[src[0]]}+{src[1].arg}];") else: - kk(*self.render_load(r[src[0]], ssa('val', u), dtype, gate=r[src[2]] if has_gate else None, - alt=r[src[3]] if has_gate else None, ss=mem_type, offset=src[1].arg)) + kk(*self.render_load(r[src[0]], ssa('val', u), dtype, gate=r[src[3]] if has_gate else None, + alt=r[src[2]] if has_gate else None, ss=mem_type, offset=src[1].arg)) elif uop is UOps.PHI: if dtype.count > 1: for x0, x1 in zip(r[src[0]], r[src[1]]): kk(f"mov.b{self.types[dtype.scalar()][1:]} {x0}, {x1};") @@ -239,7 +239,7 @@ ptx_matcher = PatternMatcher([ (UPat(UOps.ALU, name="x", dtype=dtypes.bool, arg=BinaryOps.MAX), lambda x: UOp(UOps.ALU, dtypes.uint8, tuple(s.cast(dtypes.uint8) for s in x.src), x.arg).cast(dtypes.bool)), (UPat(UOps.LOAD, name="root", dtype=dtypes.bool, src=(UPat(name="x"),UPat(name="y"),UPat(name="z"),UPat(name="k"))), - lambda root,x,y,z,k: UOp(root.op, dtypes.uint8, (x,y,z,k.cast(dtypes.uint8))).cast(dtypes.bool)), + lambda root,x,y,z,k: UOp(root.op, dtypes.uint8, (x,y,z.cast(dtypes.uint8),k)).cast(dtypes.bool)), (UPat(UOps.LOAD, name="root", dtype=dtypes.bool, src=(UPat(),UPat())), lambda root: UOp(root.op, dtypes.uint8, root.src, root.arg).cast(dtypes.bool)), (UPat(UOps.STORE, name="root", src=(UPat(),UPat(),UPat(name="z",dtype=dtypes.bool), UPat())), diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 76ab3c5a17..6a5e33e6b4 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -150,7 +150,7 @@ class CStyleLanguage(Renderer): elif uop is UOps.LOAD: val = self.render_load(dtype, r[src[0]], src[0].dtype, strip_parens(r[src[1]]), src[0].op is UOps.DEFINE_LOCAL) # NOTE: this relies on the load not happening if it's in the unselected branch - if len(src) > 3 and src[2].op is UOps.ALU: val = self.code_for_op[TernaryOps.WHERE](r[src[2]], val, r[src[3]], dtype) + if len(src) > 3 and src[3].op is UOps.ALU: val = self.code_for_op[TernaryOps.WHERE](r[src[3]], val, r[src[2]], dtype) kk(f"{self.render_dtype(dtype)} {ssa('val',u)} = {val};") elif uop is UOps.PHI: kk(f"{r[src[0]]} = {r[src[1]]};") diff --git a/tinygrad/renderer/llvmir.py b/tinygrad/renderer/llvmir.py index 5cf313eee6..44ea9db30d 100644 --- a/tinygrad/renderer/llvmir.py +++ b/tinygrad/renderer/llvmir.py @@ -133,9 +133,9 @@ class LLVMRenderer(Renderer): reduce_phis.append(u) elif uop is UOps.LOAD: if len(src) > 2: - aug_idx = bb[-1].select(lvars[src[2]], lvars[src[1]], ir.Constant(ir.IntType(32), 0)) + aug_idx = bb[-1].select(lvars[src[3]], lvars[src[1]], ir.Constant(ir.IntType(32), 0)) val = bb[-1].load(bb[-1].gep(lvars[src[0]], [aug_idx], inbounds=True)) - val = bb[-1].select(lvars[src[2]], val, lvars[src[3]]) + val = bb[-1].select(lvars[src[3]], val, lvars[src[2]]) else: val = bb[-1].load(bb[-1].gep(lvars[src[0]], [lvars[src[1]]], inbounds=True)) lvars[u] = val diff --git a/tinygrad/runtime/ops_python.py b/tinygrad/runtime/ops_python.py index 4c1a3e3683..66a1c96bba 100644 --- a/tinygrad/runtime/ops_python.py +++ b/tinygrad/runtime/ops_python.py @@ -18,7 +18,7 @@ def _load(m, i): return m[i] def load(inp, j=0): - if len(inp) == 4: return [_load(m, x+j) if gate else default for m,x,gate,default in zip(*inp)] + if len(inp) == 4: return [_load(m, x+j) if gate else default for m,x,default,gate in zip(*inp)] return [_load(m, x+j) for m,x in zip(inp[0], inp[1])] def _store(m, i, v):