diff --git a/test/test_uop_graph.py b/test/test_uop_graph.py index 86dcba1798..ea0ba740df 100644 --- a/test/test_uop_graph.py +++ b/test/test_uop_graph.py @@ -245,34 +245,34 @@ class TestUOpGraph(unittest.TestCase): # possible val = UOp(UOps.LOAD, dtypes.float.vec(4), (d1, idx)) - xyzw = tuple(UOp(UOps.GEP, dtypes.float, (val,), i) for i in range(4)) + xyzw = tuple(UOp(UOps.GEP, dtypes.float, (val,), (i,)) for i in range(4)) assert_equiv_uops(_test_vec(xyzw), val) # unaligned val = UOp(UOps.LOAD, dtypes.float.vec(4), (d1, idx)) - wzyx = tuple(UOp(UOps.GEP, dtypes.float, (val,), i) for i in reversed(range(4))) + wzyx = tuple(UOp(UOps.GEP, dtypes.float, (val,), (i,)) for i in reversed(range(4))) self.assertIs(_test_vec(wzyx).op, UOps.VECTORIZE) # different_size val = UOp(UOps.LOAD, dtypes.float.vec(2), (d1, idx)) - xy = tuple(UOp(UOps.GEP, dtypes.float, (val, ), i) for i in range(2)) + xy = tuple(UOp(UOps.GEP, dtypes.float, (val, ), (i,)) for i in range(2)) self.assertIs(_test_vec(xy+xy).op, UOps.VECTORIZE) val = UOp(UOps.LOAD, dtypes.float.vec(4), (d1, idx)) - xy = tuple(UOp(UOps.GEP, dtypes.float, (val, ), i) for i in range(2)) + xy = tuple(UOp(UOps.GEP, dtypes.float, (val, ), (i,)) for i in range(2)) self.assertIs(_test_vec(xy, count=2).op, UOps.VECTORIZE) # different vals val1 = UOp(UOps.LOAD, dtypes.float.vec(2), (d1, idx)) val2 = UOp(UOps.LOAD, dtypes.float.vec(2), (d2, idx)) - xy1 = tuple(UOp(UOps.GEP, dtypes.float, (val1, ), i) for i in range(2)) - xy2 = tuple(UOp(UOps.GEP, dtypes.float, (val2, ), i) for i in range(2)) + xy1 = tuple(UOp(UOps.GEP, dtypes.float, (val1, ), (i,)) for i in range(2)) + xy2 = tuple(UOp(UOps.GEP, dtypes.float, (val2, ), (i,)) for i in range(2)) self.assertIs(_test_vec(xy1+xy2).op, UOps.VECTORIZE) def test_gep_vec_const_fold(self): for vec_size in [2, 4, 8]: consts = [UOp.const(dtypes.float, float(i)) for i in range(vec_size)] vec = UOp(UOps.VECTORIZE, dtypes.float.vec(vec_size), tuple(consts)) - uops = to_uops_list([UOp(UOps.GEP, dtypes.float, (vec,), i) for i in range(vec_size)]) + uops = to_uops_list([UOp(UOps.GEP, dtypes.float, (vec,), (i,)) for i in range(vec_size)]) for uop, const in zip(uops, consts): assert_equiv_uops(uop, const) diff --git a/tinygrad/codegen/lowerer.py b/tinygrad/codegen/lowerer.py index f8f516761c..6c55ddca0e 100644 --- a/tinygrad/codegen/lowerer.py +++ b/tinygrad/codegen/lowerer.py @@ -117,7 +117,7 @@ class IndependentLowerer: UOp(UOps.CONTRACT, dtype=in_uops[0].dtype.vec(wmma_sz[0]), src=(in_uops[0],), arg=upcast_axes[0]), UOp(UOps.CONTRACT, dtype=in_uops[1].dtype.vec(wmma_sz[1]), src=(in_uops[1],), arg=upcast_axes[1]), UOp.const(x.dtype.vec(wmma_sz[2]), 0.0)), arg=x.arg) - return UOp(UOps.EXPAND, x.dtype, tuple(UOp(UOps.GEP, x.dtype, (ret,), i) for i in range(wmma_sz[2])), arg=upcast_axes[2]) + return UOp(UOps.EXPAND, x.dtype, tuple(UOp(UOps.GEP, x.dtype, (ret,), (i,)) for i in range(wmma_sz[2])), arg=upcast_axes[2]) if x.op is UOps.REDUCE_AXIS: # NOTE: always using ridxs is fine here reduce_range, reduce_expand = partition([self.ridxs[i] for i in x.arg[1]], lambda y: y.op is UOps.RANGE) diff --git a/tinygrad/codegen/uopgraph.py b/tinygrad/codegen/uopgraph.py index ab9360319e..a3c9d56a56 100644 --- a/tinygrad/codegen/uopgraph.py +++ b/tinygrad/codegen/uopgraph.py @@ -49,7 +49,7 @@ def fold_expanded(ex, buf): # generate the folded new_srcs if is_load: 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) + for i in range(fold_length): new_srcs[offsets[o+i]] = UOp(UOps.GEP, load_1.dtype, (new_load,), (i,)) else: for i in range(fold_length): new_srcs[offsets[o+i]] = UOp(UOps.STORE, dtypes.void, tuple(new_src)) if i == 0 else None for i in range(fold_length): used.add((rootsrc,o+i)) @@ -79,7 +79,7 @@ def fix_unfoldable_image_load(load:UOp, buf:UOp): if len(new_src) >= 4: new_src[2] = UOp(UOps.VECTORIZE, new_src[2].dtype.vec(4), tuple(new_src[2] for _ in range(4))) vec_load = UOp(UOps.LOAD, load.dtype.vec(4), tuple(new_src)) - return functools.reduce(lambda ret, i: id4.ne(i).where(ret, UOp(UOps.GEP, load.dtype, (vec_load,), i)), range(4), load.const_like(float('nan'))) + return functools.reduce(lambda ret, i: id4.ne(i).where(ret, UOp(UOps.GEP, load.dtype, (vec_load,), (i,))), range(4), load.const_like(float('nan'))) float4_folding = PatternMatcher([ (UPat(UOps.EXPAND, src=UPat(UOps.LOAD, src=(UPat.var("buf"), UPat()), allow_any_len=True), name="ex"), fold_expanded), @@ -198,7 +198,7 @@ def reduce_before_expand(reduce, expand, x): expands = flatten([x.arg for x in reduce.src[1:] if x.op is UOps.EXPAND]) if any(x in expands for x in expand.arg): return None red = UOp(UOps.REDUCE, x.dtype, (x,)+reduce.src[1:], reduce.arg) - return UOp(expand.op, expand.dtype, tuple(UOp(UOps.GEP, reduce.dtype, (red,), i) for i in range(x.dtype.count)), expand.arg) + return UOp(expand.op, expand.dtype, tuple(UOp(UOps.GEP, reduce.dtype, (red,), (i,)) for i in range(x.dtype.count)), expand.arg) def loop_collapse(loop_start, loop_end, compval, idx, mval, multconst, rng, reduce, idx2=None, idx3=None, extra=None): if getenv("DISABLE_LOOP_COLLAPSE") or rng not in reduce.src: return None # must be the right REDUCE @@ -224,19 +224,22 @@ constant_folder = PatternMatcher([ # bool ADD is OR, MUL is AND. prevents other rules to rewrite bool ADD/MUL incorrectly (UPat(UOps.ALU, dtypes.bool, arg=BinaryOps.ADD, name="x"), lambda x: UOp(x.op, x.dtype, x.src, BinaryOps.OR)), (UPat(UOps.ALU, dtypes.bool, arg=BinaryOps.MUL, name="x"), lambda x: UOp(x.op, x.dtype, x.src, BinaryOps.AND)), - # VECTORIZE/GEP - (UPat(UOps.GEP, src=(UPat(UOps.VECTORIZE, name="cast"),), name="gep"), lambda gep, cast: cast.src[gep.arg]), - *[(UPat(UOps.VECTORIZE, dtypes.float.vec(i), tuple(UPat(UOps.GEP, dtypes.float, - src=(UPat.var("x", dtype=dtypes.float.vec(i)),), arg=j) for j in range(i))), lambda x: x) for i in ([2, 4, 8, 16] + ([256] if AMX else []))], - *[(UPat(UOps.VECTORIZE, dtypes.half.vec(i), tuple(UPat(UOps.GEP, dtypes.half, - src=(UPat.var("x", dtype=dtypes.half.vec(i)),), arg=j) for j in range(i))), lambda x: x) for i in [2, 4, 8, 16]], + # VECTORIZE/GEP: the expander rule allows tuple GEP creation, this is just for removal + (UPat(UOps.VECTORIZE, src=UPat(UOps.GEP, src=(UPat(name="x"),)), name="vec"), + lambda vec,x: x if x.dtype == vec.dtype and tuple(y.arg[0] for y in vec.src) == tuple(range(len(vec.src))) else None), + # GEP/VECTORIZE, GEP/GEP, GEP/CONST, GEP/VCONST + (UPat(UOps.GEP, src=(UPat(UOps.GEP, name='g2'),), name='g1'), + lambda g1, g2: g2.src[0].gep(tuple(g2.arg[g1.arg[i]] for i in range(g1.dtype.count)))), + (UPat(UOps.GEP, src=(UPat(UOps.VECTORIZE, name="vec"),), name="gep"), + lambda gep, vec: UOp(UOps.VECTORIZE, gep.dtype, tuple(vec.src[i] for i in gep.arg)) if len(gep.arg) > 1 else vec.src[gep.arg[0]]), + (UPat(UOps.GEP, src=(UPat.cvar("c"),), name="gep"), lambda gep, c: gep.const_like(c.arg)), + (UPat(UOps.GEP, src=(UPat(UOps.VCONST, name="c"),), name="gep"), lambda gep, c: gep.const_like(tuple(c.arg[x] for x in gep.arg))), # tensor core with a 0 input is acc - *[(UPat(UOps.WMMA, src=(UPat(UOps.VECTORIZE, src=tuple(UPat.const(None, 0.0) for _ in range(i))), UPat.var(), UPat.var("acc"))), - lambda acc: acc) for i in [2, 4, 8]], - *[(UPat(UOps.WMMA, src=(UPat.var(), UPat(UOps.VECTORIZE, src=tuple(UPat.const(None, 0.0) for _ in range(i))), UPat.var("acc"))), - lambda acc: acc) for i in [2, 4, 8]], + *[(UPat(UOps.WMMA, src=(UPat.const(None, 0.0), UPat.var(), UPat.var("acc"))), lambda acc: acc) for i in [2, 4, 8]], + *[(UPat(UOps.WMMA, src=(UPat.var(), UPat.const(None, 0.0), UPat.var("acc"))), lambda acc: acc) for i in [2, 4, 8]], # tensor core cleanups - *[(UPat(UOps.REDUCE, src=(UPat(UOps.EXPAND, src=tuple(UPat(UOps.GEP, dtypes.float, src=(UPat.var("x"),), arg=i) for i in range(j)), name="expand"),) + *[(UPat(UOps.REDUCE, src=(UPat(UOps.EXPAND, + src=tuple(UPat(UOps.GEP, dtypes.float, src=(UPat.var("x"),), arg=(i,)) for i in range(j)), name="expand"),) ,name="reduce", allow_any_len=True), reduce_before_expand) for j in ([2,4,8] + ([16,256] if AMX else []))], (UPat.var("add") + UPat(UOps.WMMA, name="wmma"), lambda add, wmma: UOp(wmma.op, wmma.dtype, (wmma.src[0], wmma.src[1], wmma.src[2]+add), wmma.arg)), @@ -273,8 +276,6 @@ constant_folder = PatternMatcher([ # max folding (UPat.max(UPat.var("x"), UPat.var("y")), lambda x,y: x if x.vmin >= y.vmax else y if x.vmax <= y.vmin else None), # GEP/CAST const rules - (UPat(UOps.GEP, src=(UPat.cvar("c"),), name="root"), lambda root, c: root.const_like(c.arg)), - (UPat(UOps.GEP, src=(UPat(UOps.VCONST, name="c"),), name="root"), lambda root, c: root.const_like(c.arg[root.arg])), (UPat(UOps.CAST, name="root", src=UPat.cvar("c")), lambda root, c: root.const_like(c.arg)), # a conditional with the same results either way is a noop, also fold const conditionals (UPat.var().where(UPat.var("val"), UPat.var("val")), lambda val: val), @@ -436,7 +437,7 @@ def do_contract(con:UOp): def no_vectorized_alu(alu): if alu.dtype.count == 1: return None alus = tuple(UOp(alu.op, alu.dtype.scalar(), - tuple(UOp(UOps.GEP, s.dtype.scalar(), (s,), i) for s in alu.src), alu.arg) for i in range(alu.dtype.count)) + tuple(UOp(UOps.GEP, s.dtype.scalar(), (s,), (i,)) for s in alu.src), alu.arg) for i in range(alu.dtype.count)) return UOp(UOps.VECTORIZE, alu.dtype, alus) def create_gate(root:UOp) -> Optional[UOp]: @@ -450,6 +451,7 @@ def create_gate(root:UOp) -> Optional[UOp]: expander = PatternMatcher([ (UPat(UOps.VECTORIZE, src=UPat(UOps.CONST), name="vec"), lambda vec: UOp.const(vec.dtype, tuple(x.arg for x in vec.src))), + (UPat(UOps.VECTORIZE, src=UPat(UOps.GEP, src=(UPat(name="x"),)), name="vec"), lambda vec,x: x.gep(tuple(y.arg[0] for y in vec.src))), # create gate MUST BE BEFORE expander (UPat(UOps.STORE, name="root"), create_gate), # do expansion @@ -483,6 +485,7 @@ reducer = PatternMatcher([ (UPat(UOps.CONST, name='c'), lambda c: UOp(UOps.VECTORIZE, c.dtype, (UOp.const(c.dtype.scalar(), c.arg),)*c.dtype.count) if c.dtype.count > 1 else None), (UPat(UOps.VCONST, name='c'), lambda c: UOp(UOps.VECTORIZE, c.dtype, tuple(UOp.const(c.dtype.scalar(), x) for x in c.arg))), + (UPat(UOps.GEP, name='gep'), lambda gep: UOp(UOps.VECTORIZE, gep.dtype, tuple(gep.src[0].gep(x) for x in gep.arg)) if len(gep.arg) > 1 else None), # no ALU on vectorized dtypes (UPat((UOps.ALU, UOps.CAST, UOps.BITCAST), name="alu"), no_vectorized_alu), # delete_redundant_gates (after expand, is this still needed?) diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 5b4d699d90..4a0321144d 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -381,7 +381,10 @@ class UOp(MathTrait): def const_like(self, b:ConstType|Variable): return type(self).const(self.dtype, b) def cast(self, dtype:DType): return type(self)(UOps.CAST, dtype, (self,)) def bitcast(self, dtype:DType): return type(self)(UOps.BITCAST, dtype, (self,)) - def gep(self, i:int): return type(self)(UOps.GEP, self.dtype.scalar(), (self,), i) + def gep(self, i:Union[Tuple[int, ...], int]): + if isinstance(i, int): i = (i,) + if i == tuple(range(len(i))) and self.dtype.count == len(i): return self + return UOp(UOps.GEP, self.dtype.scalar().vec(len(i)) if len(i) > 1 else self.dtype.scalar(), (self,), i) @classmethod def load(cls, *src:UOp, dtype:DType): return cls(UOps.LOAD, dtype, src) @classmethod @@ -648,7 +651,7 @@ class UPat(MathTrait): # copied from UOp def cast(self, dtype=None): return type(self)(UOps.CAST, dtype, (self,)) def bitcast(self, dtype=None): return type(self)(UOps.BITCAST, dtype, (self,)) - def gep(self, i:int): return type(self)(UOps.GEP, None, (self,), i) + def gep(self, i:int): return type(self)(UOps.GEP, None, (self,), (i,)) @classmethod def load(cls, *src:UPat, dtype:Optional[DType]=None): return cls(UOps.LOAD, dtype, src) @classmethod diff --git a/tinygrad/renderer/assembly.py b/tinygrad/renderer/assembly.py index ce5e2aabd3..d6c9bfbac6 100644 --- a/tinygrad/renderer/assembly.py +++ b/tinygrad/renderer/assembly.py @@ -213,7 +213,9 @@ class PTXRenderer(Renderer): r[u] = f"%{args[0]}" kk(*self.render_load(args[0], ssa('dat', u, self.types[dtype]), dtype, ss=".param")) elif uop is UOps.CONST: r[u] = const(args, dtype, mov=True) - elif uop is UOps.GEP: r[u] = r[src[0]][u.arg] + elif uop is UOps.GEP: + assert len(u.arg) == 1 + r[u] = r[src[0]][u.arg[0]] elif uop is UOps.LOAD: assert src[0].dtype == dtypes.int64, "load isn't int64" assert src[1].op is UOps.CONST, f"load isn't const {u}" diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index ad014dee33..db2a96d572 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -177,9 +177,11 @@ class CStyleLanguage(Renderer): elif uop is UOps.DEFINE_ACC: kk(f"{self.render_dtype(dtype)} {ssa('acc',u)} = {r[src[0]]};") elif uop is UOps.CONST: r[u] = self.render_const(args, dtype) if args >= 0 else f"({self.render_const(args, dtype)})" elif uop is UOps.GEP: + assert len(args) == 1 from_ssa = src[0].op in {UOps.LOAD, UOps.WMMA, UOps.DEFINE_ACC} r[u] = (r[src[0]] if from_ssa else f"{(r[src[0]])}") + \ - (f"[{args}]" if src[0].dtype.count > (8 if self.device in {"CUDA", "NV"} else 4) or self.device == 'CLANG' else f".{'xyzwabcd'[args]}") + (f"[{args[0]}]" if src[0].dtype.count > (8 if self.device in {"CUDA", "NV"} else 4) \ + or self.device == 'CLANG' else f".{'xyzwabcd'[args[0]]}") else: raise RuntimeError(f"failed to render {u}") # NOTE: this relies on bufs dict preserving order diff --git a/tinygrad/runtime/ops_python.py b/tinygrad/runtime/ops_python.py index 222f393eb5..a66165b9e0 100644 --- a/tinygrad/runtime/ops_python.py +++ b/tinygrad/runtime/ops_python.py @@ -127,7 +127,8 @@ class PythonProgram: for j in range(len(inp[0])): inp[0][j] = inp[1][j] ul[i] = inp[0] elif uop is UOps.GEP: - ul[i] = inp[0][arg] + assert len(arg) == 1 + ul[i] = inp[0][arg[0]] elif uop is UOps.WMMA: # here are the models for the WMMA instruction on the different hardware def wmma_helper(WARP_THREADS, K, NUM_A, NUM_B, NUM_C, a_elem, b_elem, c_map):