diff --git a/extra/gemm/metal_uop_matmul.py b/extra/gemm/metal_uop_matmul.py index 431a645b32..266f7f6dfa 100644 --- a/extra/gemm/metal_uop_matmul.py +++ b/extra/gemm/metal_uop_matmul.py @@ -20,8 +20,8 @@ def hand_spec_tc_cores(): gk = UOp.range(N // 8, 0, AxisType.REDUCE) - a_tc = UOp.vectorize(*[mat_idx(a, gx, gk, warp, i) for i in range(2)]) - b_tc = UOp.vectorize(*[mat_idx(b, gk, gy, warp, i) for i in range(2)]) + a_tc = UOp.stack(*[mat_idx(a, gx, gk, warp, i) for i in range(2)]) + b_tc = UOp.stack(*[mat_idx(b, gk, gy, warp, i) for i in range(2)]) acc = UOp.placeholder((2,), dtypes.float, slot=0, addrspace=AddrSpace.REG) acc = acc[0].set(0.0) @@ -30,7 +30,7 @@ def hand_spec_tc_cores(): # TODO: make this simple wmma_arg = ('WMMA_8_8_8_float_float', (8, 8, 8), dtypes.float, dtypes.float, 'METAL', 32, (((3, 2),), ((3, 2),), ((3, 2),)), ()) - acc_load = UOp.vectorize(acc.after(gk)[0], acc.after(gk)[1]) + acc_load = UOp.stack(acc.after(gk)[0], acc.after(gk)[1]) out = UOp(Ops.WMMA, dtypes.float.vec(2), (a_tc, b_tc, acc_load), arg=wmma_arg) end_loop = UOp.group(*[acc[i].store(out.index(i)) for i in range(2)]).end(gk) diff --git a/extra/gemm/mi350x_uop_matmul.py b/extra/gemm/mi350x_uop_matmul.py index 6d10c59852..2366aa7529 100644 --- a/extra/gemm/mi350x_uop_matmul.py +++ b/extra/gemm/mi350x_uop_matmul.py @@ -192,7 +192,7 @@ acc = UOp.placeholder((4,), dtypes.float, 0, AddrSpace.REG) acc = acc[init_l:=UOp.range(4, 1)].set(0.0, end=init_l) # do the wmma -acc_load = UOp.vectorize(*[acc.after(K_loop)[i] for i in range(4)]) +acc_load = UOp.stack(*[acc.after(K_loop)[i] for i in range(4)]) wmma_arg = ('WMMA_16_16_32_half_float', (16, 16, 32), dtypes.half, dtypes.float, 'AMD', 64, ((), (), ((3, 2), (2, 2))), ()) out = UOp(Ops.WMMA, dtypes.float.vec(4), (A_in, B_in, acc_load), arg=wmma_arg) diff --git a/extra/thunder/tiny/tk/group.py b/extra/thunder/tiny/tk/group.py index 389f564097..4d598d3cf5 100644 --- a/extra/thunder/tiny/tk/group.py +++ b/extra/thunder/tiny/tk/group.py @@ -84,13 +84,13 @@ class Group: for width in self.ker.range(c.shape[-2], track=False): for inner in self.ker.range(a.shape[-2], axis_type=AxisType.REDUCE, track=False): if a_base_shape.cols == 16: - a_in = UOp.vectorize(*[a[height, inner, i] for i in range(4)]) - b_in = UOp.vectorize(*[b[inner, width, i] for i in range(4)]) + a_in = UOp.stack(*[a[height, inner, i] for i in range(4)]) + b_in = UOp.stack(*[b[inner, width, i] for i in range(4)]) elif a_base_shape.cols == 32: - a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)]) - b_in = UOp.vectorize(*[b[inner, width, i] for i in range(8)]) + a_in = UOp.stack(*[a[height, inner, i] for i in range(8)]) + b_in = UOp.stack(*[b[inner, width, i] for i in range(8)]) else: raise NotImplementedError(f"mma_AB not implemented for {a_base_shape.cols=}") - d_in = UOp.vectorize(*[c[height, width, i] for i in range(4)]) + d_in = UOp.stack(*[c[height, width, i] for i in range(4)]) out = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in, d_in), arg=wmma_arg) c_i = [c[height, width, i].store(out.index(i)) for i in range(4)] @@ -114,13 +114,13 @@ class Group: for width in self.ker.range(c.shape[-2], track=False): for inner in self.ker.range(a.shape[-2], axis_type=AxisType.REDUCE, track=False): if a_base_shape.cols == 16: - a_in = UOp.vectorize(*[a[height, inner, i] for i in range(4)]) - b_in = UOp.vectorize(*[b[width, inner, i] for i in range(4)]) + a_in = UOp.stack(*[a[height, inner, i] for i in range(4)]) + b_in = UOp.stack(*[b[width, inner, i] for i in range(4)]) elif a_base_shape.cols == 32: - a_in = UOp.vectorize(*[a[height, inner, i] for i in range(8)]) - b_in = UOp.vectorize(*[b[width, inner, i] for i in range(8)]) + a_in = UOp.stack(*[a[height, inner, i] for i in range(8)]) + b_in = UOp.stack(*[b[width, inner, i] for i in range(8)]) else: raise NotImplementedError(f"mma_ABt not implemented for {a_base_shape.cols=}") - d_in = UOp.vectorize(*[c[height, width, i] for i in range(4)]) + d_in = UOp.stack(*[c[height, width, i] for i in range(4)]) out = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in, d_in), arg=wmma_arg) c_i = [c[height, width, i].store(out.index(i)) for i in range(4)] @@ -144,13 +144,13 @@ class Group: for width in self.ker.range(c.shape[-2], track=False): for inner in self.ker.range(a.shape[-3], axis_type=AxisType.REDUCE, track=False): if a_base_shape.cols == 16: - a_in = UOp.vectorize(*[a[inner, height, i] for i in range(4)]) - b_in = UOp.vectorize(*[b[inner, width, i] for i in range(4)]) + a_in = UOp.stack(*[a[inner, height, i] for i in range(4)]) + b_in = UOp.stack(*[b[inner, width, i] for i in range(4)]) elif a_base_shape.cols == 32: - a_in = UOp.vectorize(*[a[inner, height, i] for i in range(8)]) - b_in = UOp.vectorize(*[b[inner, width, i] for i in range(8)]) + a_in = UOp.stack(*[a[inner, height, i] for i in range(8)]) + b_in = UOp.stack(*[b[inner, width, i] for i in range(8)]) else: raise NotImplementedError(f"mma_AtB not implemented for {a_base_shape.cols=}") - d_in = UOp.vectorize(*[c[height, width, i] for i in range(4)]) + d_in = UOp.stack(*[c[height, width, i] for i in range(4)]) out = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in, d_in), arg=wmma_arg) c_i = [c[height, width, i].store(out.index(i)) for i in range(4)] @@ -174,13 +174,13 @@ class Group: for width in self.ker.range(c.shape[-2], track=False): for inner in self.ker.range(a.shape[-3], axis_type=AxisType.REDUCE, track=False): if a_base_shape.cols == 16: - a_in = UOp.vectorize(*[a[inner, height, i] for i in range(4)]) - b_in = UOp.vectorize(*[b[width, inner, i] for i in range(4)]) + a_in = UOp.stack(*[a[inner, height, i] for i in range(4)]) + b_in = UOp.stack(*[b[width, inner, i] for i in range(4)]) elif a_base_shape.cols == 32: - a_in = UOp.vectorize(*[a[inner, height, i] for i in range(8)]) - b_in = UOp.vectorize(*[b[width, inner, i] for i in range(8)]) + a_in = UOp.stack(*[a[inner, height, i] for i in range(8)]) + b_in = UOp.stack(*[b[width, inner, i] for i in range(8)]) else: raise NotImplementedError(f"mma_AtBt not implemented for {a_base_shape.cols=}") - d_in = UOp.vectorize(*[c[height, width, i] for i in range(4)]) + d_in = UOp.stack(*[c[height, width, i] for i in range(4)]) out = UOp(Ops.WMMA, dtypes.float32.vec(4), (a_in, b_in, d_in), arg=wmma_arg) c_i = [c[height, width, i].store(out.index(i)) for i in range(4)] diff --git a/test/backend/test_isel.py b/test/backend/test_isel.py index db11bed97c..ebf71c0d8f 100644 --- a/test/backend/test_isel.py +++ b/test/backend/test_isel.py @@ -39,8 +39,8 @@ class TestIselX86(unittest.TestCase): c = UOp.variable("c", 0, 0, dtypes.float32.vec(4)) d = UOp.variable("e", 0, 0, dtypes.float32) - valid = [UOp.vectorize(lane(a, 0), lane(b, 1), lane(a, 2), lane(b, 3)), - UOp.vectorize(lane(a, 3), lane(b, 2), lane(c, 1), d)] + valid = [UOp.stack(lane(a, 0), lane(b, 1), lane(a, 2), lane(b, 3)), + UOp.stack(lane(a, 3), lane(b, 2), lane(c, 1), d)] for shuf in valid: self.assertIs(self.isel_rewrite(shuf).arg, X86Ops.VINSERTPS) # complex address is [base + index*scale + displacement] diff --git a/test/null/test_graph_rewrite.py b/test/null/test_graph_rewrite.py index 4404d394b2..243e7d0ca3 100644 --- a/test/null/test_graph_rewrite.py +++ b/test/null/test_graph_rewrite.py @@ -188,7 +188,7 @@ class TestGEPAndVectorizeRewrite(unittest.TestCase): def test_gep_tuple_extraction(self): # GEP on a vector dtype to extract multiple elements as a vector base_vector = UOp.const(dtypes.float32, (1.0, 2.0, 3.0, 4.0)) - self.assertEqual(list(apply_rewrite_values(UOp.vectorize(*[base_vector.index(i) for i in (2, 3)]))), [3.0, 4.0]) + self.assertEqual(list(apply_rewrite_values(UOp.stack(*[base_vector.index(i) for i in (2, 3)]))), [3.0, 4.0]) def test_gep_on_const_stack(self): # GEP on a const STACK to extract a single element @@ -198,7 +198,7 @@ class TestGEPAndVectorizeRewrite(unittest.TestCase): def test_gep_tuple_on_const_stack(self): # GEP on a const STACK using a tuple to extract multiple elements const_stack = UOp.const(dtypes.float32, (7.0, 8.0, 9.0, 10.0)) - self.assertEqual(list(apply_rewrite_values(UOp.vectorize(*[const_stack.index(i) for i in (1, 3)]))), [8.0, 10.0]) + self.assertEqual(list(apply_rewrite_values(UOp.stack(*[const_stack.index(i) for i in (1, 3)]))), [8.0, 10.0]) def test_vectorize_multiple_elements(self): # Vectorizing multiple elements using GEP diff --git a/test/null/test_viz.py b/test/null/test_viz.py index 2ab0b5b66a..c2734b2150 100644 --- a/test/null/test_viz.py +++ b/test/null/test_viz.py @@ -268,12 +268,12 @@ class TestViz(unittest.TestCase): def test_stack_movement_not_folded_unless_all_const(self): a = UOp.variable("a", 0, 10, dtype=dtypes.int) c = UOp.const(dtypes.int, 1) - stack = a.vectorize(c) + stack = a.stack(c) reshaped = stack.reshape((1, 2)) graph = uop_to_json(VizData(), reshaped) self.assertFalse(graph[id(stack)]["exclude"]) - const_stack = c.vectorize(UOp.const(dtypes.int, 2)) + const_stack = c.stack(UOp.const(dtypes.int, 2)) const_reshaped = const_stack.reshape((1, 2)) const_graph = uop_to_json(VizData(), const_reshaped) self.assertTrue(const_graph[id(const_stack)]["exclude"]) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 46e73017fd..aec6491ded 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -111,7 +111,7 @@ def broadcast_and_devec_wmma(b:UOp): for idx in itertools.product(*[range(i) for i in b.shape[:-1]]): idx_c = [UOp.const(dtypes.index, i) for i in idx] src.append(b.replace(src=tuple([x.index(*idx_c) for x in src_reshaped]))) - return UOp.vectorize(*src).reshape(b.shape) + return UOp.stack(*src).reshape(b.shape) pm_wmma_add = PatternMatcher([ (UPat(Ops.WMMA, name="wmma") + UPat.var("add"), @@ -136,7 +136,7 @@ def do_devectorize(b:UOp): for idx in itertools.product(*[range(x) for x in b.shape]): idx_c = [UOp.const(dtypes.index, i) for i in idx] src.append(b.replace(src=tuple([x.index(*idx_c) for x in b.src]))) - return UOp.vectorize(*src).reshape(b.shape) if b.op is not Ops.STORE else UOp.group(*src) + return UOp.stack(*src).reshape(b.shape) if b.op is not Ops.STORE else UOp.group(*src) def do_stack_wmma(u:UOp): if all(x.op in (Ops.STACK, Ops.WMMA) for x in u.src): return None @@ -144,7 +144,7 @@ def do_stack_wmma(u:UOp): src = [] for b in u.src: if b.op != Ops.STACK: - src.append(UOp._stack(*[b.index(UOp.const(dtypes.index, i)) for i in range(b.max_numel())])) + src.append(UOp.stack(*[b.index(UOp.const(dtypes.index, i)) for i in range(b.max_numel())])) else: src.append(b) return u.replace(src=tuple(src)) @@ -163,7 +163,7 @@ devectorizer2 = mop_cleanup+pm_mops+PatternMatcher([ (UPat(Ops.WMMA, name="u"), do_stack_wmma), # stacked INDEX is many INDEX (UPat(Ops.INDEX, src=(UPat((Ops.PARAM, Ops.BUFFER), name="b"), UPat(Ops.STACK, name="s"))), - lambda b,s: UOp.vectorize(*[b.index(u) for u in s.src])), + lambda b,s: UOp.stack(*[b.index(u) for u in s.src])), # INDEX into RESHAPE moves the RESHAPE (UPat(Ops.INDEX, src=(UPat((Ops.PARAM, Ops.BUFFER), name="b"), UPat(Ops.RESHAPE, name="s"))), lambda b,s: b.index(s.src[0]).reshape(s.shape)), @@ -173,7 +173,7 @@ devectorizer2 = mop_cleanup+pm_mops+PatternMatcher([ (UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0].index(UOp.const(dtypes.index, 0)) if x.marg == () and x.src[0].shape == (1,) else None), # EXPAND on scalar -> STACK (UPat(Ops.EXPAND, src=(UPat.var("x"), UPat()), name="out"), - lambda x,out: UOp.vectorize(*([x]*out.max_numel())) if x.shape == () and out.shape == (out.max_numel(),) else None), + lambda x,out: UOp.stack(*([x]*out.max_numel())) if x.shape == () and out.shape == (out.max_numel(),) else None), # INDEX on INDEX is INDEX (UPat(Ops.INDEX, src=(UPat(Ops.INDEX, name="idx1", allow_any_len=True),), allow_any_len=True, name="idx2"), lambda idx1, idx2: idx1.src[0].index(*idx1.src[1:], *idx2.src[1:])), diff --git a/tinygrad/codegen/late/coalese.py b/tinygrad/codegen/late/coalese.py index 60d5c96435..eeac316af4 100644 --- a/tinygrad/codegen/late/coalese.py +++ b/tinygrad/codegen/late/coalese.py @@ -41,7 +41,7 @@ def simplify_valid_load(buf:UOp, start_idx:UOp, valid:UOp) -> UOp|None: def simplify_valid_image_load(buf:UOp, idx_y:UOp, idx_x:UOp, valid:UOp) -> UOp|None: if not is_image_shape(buf._shape): return None if idx_x.dtype != idx_y.dtype: idx_x, idx_y = idx_x.cast(dtypes.int), idx_y.cast(dtypes.int) - start_idx = idx_x._stack(idx_y) + start_idx = idx_x.stack(idx_y) idx = uop_given_valid(valid, start_idx) drop_stmt = _drop_valid_stmts(valid, idx, buf._shape[0], buf._shape[1]) @@ -74,7 +74,7 @@ def transform_to_image(ctx, buf:UOp, x:UOp) -> UOp|None: # search for dims that drop the most valid statements best_drop, cands = -1, [] for ch, cw in [shapes[buf.arg.slot]] if buf.arg.slot in shapes else image_valid_dims(buf.dtype, buf.max_numel(), ren.target.arch): - cidx = uop_given_valid(valid, ((x//4)%cw)._stack(x//(4*cw))) + cidx = uop_given_valid(valid, ((x//4)%cw).stack(x//(4*cw))) dropped = len(_drop_valid_stmts(valid, cidx, ch, cw)) if dropped > best_drop: best_drop, cands = dropped, [(ch, cw, cidx)] elif dropped == best_drop: cands.append((ch, cw, cidx)) @@ -152,7 +152,7 @@ def memory_coalesing(sink:UOp, ctx:Renderer) -> UOp: for i,g in enumerate(grp): assert len(offsets[g]) == 1, f"attempting multiple stores: {len(offsets[g])}" datas.append(offsets[g][0].src[1]) - store = idx.store(UOp._stack(*datas) if len(datas) > 1 else datas[0]) + store = idx.store(UOp.stack(*datas) if len(datas) > 1 else datas[0]) for i,g in enumerate(grp): replacements[offsets[g][0]] = store else: ld = idx.load() diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 3415cba3a2..c9a52f83de 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -518,8 +518,6 @@ class UOp(RandMixin, metaclass=UOpMetaClass): def group(*srcs:UOp|None): # pylint: disable=no-self-argument if len(srcs) == 1 and isinstance(srcs[0], UOp): return srcs[0] return UOp(Ops.GROUP, src=tuple([x for x in srcs if x is not None])) - def _stack(self, *srcs): return self.stack(*srcs) - def vectorize(self, *srcs): return self._stack(*srcs) def index(self, *srcs:UOp|int|None, **kwargs): new_srcs: list[UOp] = [UOp.const(dtypes.index, x) if isinstance(x, int) else x for x in srcs if x is not None] if len(new_srcs) == 1 and new_srcs[0].op is Ops.CONST and self.op is Ops.STACK: return self.src[new_srcs[0].arg] @@ -577,7 +575,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): def ins(self, arg, **kwargs): return UOp(Ops.INS, kwargs.pop("dtype", self.dtype), kwargs.pop("src", self.src), arg, kwargs.pop("tag", self.tag)) def contract(self, *rngs:UOp): assert all(x.arg[-1] == AxisType.UPCAST for x in rngs), "all contract ranges must be upcast" - return UOp.vectorize(*[self.substitute(dict(zip(rngs, [r.const_like(i) for r,i in zip(rngs, idx)]))) + return UOp.stack(*[self.substitute(dict(zip(rngs, [r.const_like(i) for r,i in zip(rngs, idx)]))) for idx in itertools.product(*[range(int(r.vmax)+1) for r in rngs])]) @staticmethod def wmma(a:UOp, b:UOp, acc:UOp, arg:tuple[tuple[int, int, int], str, int]): @@ -599,7 +597,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): # NOTE: it always has to be STACK now, even if they are all the same if isinstance(b, tuple): stk = [UOp(Ops.CONST, dtype, arg=dtype.const(c), src=()) for c in b] - ret = UOp.vectorize(*stk) + ret = UOp.stack(*stk) else: ret = UOp(Ops.CONST, dtype, arg=dtype.const(b), src=()) return ret._mop(Ops.EXPAND, arg=shape) if shape is not None and shape != () and ret.shape != shape else ret @@ -624,11 +622,11 @@ class UOp(RandMixin, metaclass=UOpMetaClass): return cond.where(self, self.const_like(Invalid)) def get_idx(self) -> UOp: assert dtypes.is_int(self.dtype), "Can only call get_idx on index dtype" - if self.op is Ops.STACK: return UOp.vectorize(*(x.get_idx() for x in self.src)) + if self.op is Ops.STACK: return UOp.stack(*(x.get_idx() for x in self.src)) return self.src[1] if self.op is Ops.WHERE and self.src[2].arg is Invalid else self def get_valid(self) -> UOp: assert dtypes.is_int(self.dtype), "Can only call get_valid on index dtype" - if self.op is Ops.STACK: return UOp.vectorize(*(x.get_valid() for x in self.src)) + if self.op is Ops.STACK: return UOp.stack(*(x.get_valid() for x in self.src)) return self.src[0] if self.op is Ops.WHERE and self.src[2].arg is Invalid else UOp.const(dtypes.bool, self.arg is not Invalid) def reduce(self, *src:UOp, **kwargs): arg = kwargs.pop('arg', None)