diff --git a/examples/mlperf/optim.py b/examples/mlperf/optim.py index f93ccef119..5e01319301 100644 --- a/examples/mlperf/optim.py +++ b/examples/mlperf/optim.py @@ -15,7 +15,7 @@ def stochastic_round_bf16(x:Tensor) -> Tensor: bits = x.bitcast(dtypes.uint32) if isinstance(x.device, tuple): shape = x.uop.shard_shape if x.uop.axis is not None else x.shape - noise = Tensor(UOp(Ops.MSTACK, dtypes.default_float, tuple(Tensor.rand(*shape, device=d).uop for d in x.device))) + noise = Tensor(UOp(Ops.MSTACK, src=tuple(Tensor.rand(*shape, device=d).uop for d in x.device))) else: noise = x.rand_like() noise = (noise * 0xFFFF).cast(dtypes.uint32) diff --git a/extra/gemm/moe_routing.py b/extra/gemm/moe_routing.py index 30ba001c3f..eb4982c6e6 100644 --- a/extra/gemm/moe_routing.py +++ b/extra/gemm/moe_routing.py @@ -53,7 +53,7 @@ def _ggather_bwd(gradient:UOp, kernel:UOp) -> tuple: g, m, j, jo, ji = _kv_ranges(Gk, M, Dk, _blk_for(Dk)) row = idx.index(g, m).cast(dtypes.weakint) val = gout.index(g, m, j).load().cast(dtypes.float32) - atomic = UOp(Ops.CUSTOM, dtypes.void, (gtab.index(g, row, j), val), arg=atomic_str) + atomic = UOp(Ops.CUSTOM, src=(gtab.index(g, row, j), val), arg=atomic_str) return atomic.end(g, m, jo, ji).sink(arg=KernelInfo(name=f"ggather_bwd_{M}_{Dk}", opts_to_apply=())) grad_table = Tensor.custom_kernel(gt, go, Tensor(idx_u, device=dev), fxn=_bwd_kernel)[0] return (None, grad_table.cast(table_u.dtype).uop, None) diff --git a/extra/hcq2/ops_amd2.py b/extra/hcq2/ops_amd2.py index c4f7b6bc3b..bfd238210a 100644 --- a/extra/hcq2/ops_amd2.py +++ b/extra/hcq2/ops_amd2.py @@ -87,7 +87,7 @@ def release_mem(ctx, address=0x0, value=0, data_sel=0, int_sel=2, ctxid=0, cache def memory_barrier(ctx): pf = '' if ctx.nbio.version[0] == 2 else '0' if ctx.nbio.version[:2] != (7, 11) else '1' - return UOp(Ops.LINEAR, dtypes.void, ( + return UOp(Ops.LINEAR, src=( wait_reg_mem(ctx, reg=getattr(ctx.nbio, f'regBIF_BX_PF{pf}_GPU_HDP_FLUSH_REQ').addr[0], reg_done=getattr(ctx.nbio, f'regBIF_BX_PF{pf}_GPU_HDP_FLUSH_DONE').addr[0], value=0xffffffff), acquire_mem(ctx))) @@ -135,7 +135,7 @@ def pm4_program(ctx, call, prg): wreg(ctx, ctx.gc.regCOMPUTE_START_X, 0, 0, 0, *(info.local_size or (1, 1, 1)), 0, 0), pkt3(ctx, PM4Ops.DISPATCH_DIRECT, *info.global_size, dispatch_init), pkt3(ctx, PM4Ops.EVENT_WRITE, ctx.pm4.EVENT_TYPE(ctx.soc.CS_PARTIAL_FLUSH) | ctx.pm4.EVENT_INDEX(EVENT_INDEX_PARTIAL_FLUSH))] - return UOp(Ops.LINEAR, dtypes.void, tuple(ins)) + return UOp(Ops.LINEAR, src=tuple(ins)) pm_pm4_opsel = PatternMatcher([ (UPat(Ops.CALL, src=(UPat(Ops.PROGRAM, name="prg"),), name="call", allow_any_len=True), pm4_program), @@ -207,7 +207,7 @@ def sdma_timestamp(ctx, ins, dst): pm_sdma_opsel = PatternMatcher([ (UPat(Ops.CALL, src=(UPat(Ops.COPY),), name="call", allow_any_len=True), sdma_copy), - (UPat(Ops.INS, arg="barrier"), lambda: UOp(Ops.NOOP, dtypes.void, ())), + (UPat(Ops.INS, arg="barrier"), lambda: UOp(Ops.NOOP)), (UPat(Ops.INS, arg="wait", src=(UPat(name="dst"), UPat(name="val")), name="ins"), sdma_wait), (UPat(Ops.INS, arg="timestamp", src=(UPat(name="dst"),), name="ins"), sdma_timestamp), (UPat(Ops.INS, arg="store", src=(UPat((Ops.BUFFER, Ops.PARAM), name="dst"), UPat(name="val")), name="ins"), sdma_store), diff --git a/extra/llama_kernels/quantize_fp8_delayed/__init__.py b/extra/llama_kernels/quantize_fp8_delayed/__init__.py index 6e1c25a1fc..db9bb5fd68 100644 --- a/extra/llama_kernels/quantize_fp8_delayed/__init__.py +++ b/extra/llama_kernels/quantize_fp8_delayed/__init__.py @@ -50,7 +50,7 @@ def _custom_quantize_fp8_with_amax(fp8_out:UOp, amax_out:UOp, x:UOp, amax_state: else: raise NotImplementedError(f"no atomic max for device {device}") amax_idx = amax_out.reshape((1,)).index(UOp.const(0)) max_val = lds[0].load() - atomic = UOp(Ops.CUSTOM, dtypes.void, (amax_idx, max_val.bitcast(dtypes.int32), max_val, amax_idx.load()), arg=atomic_arg) + atomic = UOp(Ops.CUSTOM, src=(amax_idx, max_val.bitcast(dtypes.int32), max_val, amax_idx.load()), arg=atomic_arg) return atomic.end(tid, wg).sink(arg=KernelInfo(f"quantize_fp8_with_amax_{n_elems}", opts_to_apply=())) @functools.cache diff --git a/extra/llama_kernels/quantize_mxfp4/__init__.py b/extra/llama_kernels/quantize_mxfp4/__init__.py index 6cdd983c6f..e68fbb146b 100644 --- a/extra/llama_kernels/quantize_mxfp4/__init__.py +++ b/extra/llama_kernels/quantize_mxfp4/__init__.py @@ -12,7 +12,7 @@ def _custom_quantize_mxfp4(row_fp4:UOp, row_scale:UOp, col_fp4:UOp, col_scale:UO mem = M*N*2 + M*N + M*N//16 # read bf16, write row+col fp4 + e8m0 outputs = (row_fp4, row_scale, col_fp4, col_scale) sink = UOp.sink(*(o.base for o in outputs), x.base, - *(UOp(Ops.CUSTOM, dtypes.void, (o.base.index(0),), arg="") for o in outputs), + *(UOp(Ops.CUSTOM, src=(o.base.index(0),), arg="") for o in outputs), UOp.special(256, "lidx0"), UOp.special(M//128, "gidx0"), UOp.special(N//64, "gidx1"), arg=KernelInfo(name, estimates=Estimates(ops=12*M*N, mem=mem))) src = (pathlib.Path(__file__).parent/"quantize_mxfp4.cpp").read_text() diff --git a/test/backend/test_isel.py b/test/backend/test_isel.py index a3b986c619..f219195a52 100644 --- a/test/backend/test_isel.py +++ b/test/backend/test_isel.py @@ -7,7 +7,7 @@ from tinygrad.renderer.isa.x86 import X86Renderer, X86Ops from tinygrad.renderer.isa import IselContext # INDEX on a register value with a constant index extracts a single element (the old GEP) -def lane(y:UOp, i:int) -> UOp: return y.index(UOp.cconst(i, dtypes.int), dtype=y.dtype) +def lane(y:UOp, i:int) -> UOp: return y.index(UOp.cconst(i, dtypes.int)) @unittest.skipUnless(isinstance(Device[Device.DEFAULT].renderer, X86Renderer), "only x86") class TestIselX86(unittest.TestCase): diff --git a/test/null/test_uop_vmin_vmax.py b/test/null/test_uop_vmin_vmax.py index 2dc3f27c35..6398c28bb2 100644 --- a/test/null/test_uop_vmin_vmax.py +++ b/test/null/test_uop_vmin_vmax.py @@ -82,7 +82,7 @@ class TestVminVmaxProperties(unittest.TestCase): def test_vmin_vmax_multiplication_0_inf(self): # vmin and vmax for multiplication with a variable x = UOp.const(0.0) - y = UOp.load(UOp.param(0, dtypes.float, (1,)), UOp.const(0), dtype=dtypes.float) + y = UOp.load(UOp.param(0, dtypes.float, (1,)), UOp.const(0)) uop = x * y # TODO: these should be 0, but definitely should not be nan self.assertEqual(uop.vmin, -math.inf) diff --git a/test/null/test_uops.py b/test/null/test_uops.py index aaded43fdf..f4230bf9e8 100644 --- a/test/null/test_uops.py +++ b/test/null/test_uops.py @@ -56,7 +56,7 @@ class TestDTypeFromUOp(unittest.TestCase): if u.is_invalid)), (dtypes.float32, dtypes.float32, dtypes.bool)) invalid, value = UOp.invalid(), UOp.const(1, dtypes.float32) for u in (UOp.param(0, dtypes.bool, ()).where(value, invalid), value+invalid, UOp.stack(value, invalid)): self.assertIs(u.src[-1], invalid) - for u in (UOp(Ops.STACK, dtypes.float32, src=(value, invalid)), UOp(Ops.ADD, dtypes.float32, src=(value, invalid)), + for u in (UOp(Ops.STACK, src=(value, invalid)), UOp(Ops.ADD, src=(value, invalid)), UOp.const(True).where(value, invalid), UOp(Ops.CMPLT, src=(invalid, value)), UOp(Ops.CMPLT, src=(value, invalid)), UOp.param(0, dtypes.float32, (4,)).index(invalid)): type_verify(u, spec_shared) gate, value = UOp.param(0, dtypes.bool, ()), UOp.param(1, dtypes.float, ()) @@ -64,7 +64,7 @@ class TestDTypeFromUOp(unittest.TestCase): type_verify(out.sink(), spec_program) def test_remove_invalid_stack_lanes(self): - stack = UOp(Ops.STACK, dtypes.half, (UOp.const(1, dtypes.half), UOp.invalid())) + stack = UOp(Ops.STACK, src=(UOp.const(1, dtypes.half), UOp.invalid())) out = graph_rewrite(stack, pm_remove_invalid) self.assertEqual(out.src, (UOp.const(1, dtypes.half), UOp.const(0, dtypes.half))) type_verify(out.sink(), spec_program) @@ -284,7 +284,7 @@ class TestFastIdiv(unittest.TestCase): g = UOp.param(0, dt, (3,)) c = UOp.const(2) l = g.index(c) - a = UOp(Ops.CDIV, dt, (l, c)) + a = UOp(Ops.CDIV, src=(l, c)) uops = to_uops_list([a], ren=Device[Device.DEFAULT].renderer) Device[Device.DEFAULT].renderer.render(uops) ops = [x.op for x in uops] @@ -296,7 +296,7 @@ class TestFastIdiv(unittest.TestCase): for dt in (dtypes.int32, dtypes.uint32): g = UOp.param(0, dt, (9,)) c = UOp.const(8) - a = UOp(Ops.FLOORMOD, dt, (g.index(c), c)) + a = UOp(Ops.FLOORMOD, src=(g.index(c), c)) uops = to_uops_list([a], ren=Device[Device.DEFAULT].renderer) ops = [x.op for x in uops] self.assertIn(Ops.AND, ops, f"For dtype={dt} FLOORMOD by pow2 did not simplify to AND") @@ -308,7 +308,7 @@ class TestFastIdiv(unittest.TestCase): for dt in (dtypes.int32, dtypes.uint32, dtypes.int64, dtypes.uint64): g = UOp.param(0, dt, (3,)) c = UOp.const(2) - a = UOp(Ops.FLOORDIV, dt, (g.index(c), c)) + a = UOp(Ops.FLOORDIV, src=(g.index(c), c)) uops = to_uops_list([a], ren=Device[Device.DEFAULT].renderer) ops = [x.op for x in uops] self.assertIn(Ops.SHR, ops, f"For dtype={dt} FLOORDIV by power of two did not simplify to shift") @@ -469,10 +469,10 @@ class TestUOpRender(unittest.TestCase): self.assertEqual(UOp.range(1, 0, src=(shrink,), dtype=dtypes.int).render(simplify=False), "r0") def test_render_vectorize_empty(self): - u = UOp(Ops.STACK, dtype=dtypes.void, src=()) + u = UOp(Ops.STACK, src=()) self.assertEqual(u.render(simplify=False), "{}") def test_render_vectorize_empty_simplified(self): - u = UOp(Ops.STACK, dtype=dtypes.void, src=()) + u = UOp(Ops.STACK, src=()) self.assertEqual(u.render(), "{}") def test_render_vectorize_same(self): u = UOp(Ops.STACK, src=(UOp.const(0),)*3) diff --git a/test/null/test_validate_oob.py b/test/null/test_validate_oob.py index abc7a9bc06..5a256f8db6 100644 --- a/test/null/test_validate_oob.py +++ b/test/null/test_validate_oob.py @@ -14,37 +14,37 @@ class TestValidateOOB(unittest.TestCase): def test_const_index(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) - to_uops_list([buf.index(UOp.const(0)).load(dtype=dtypes.int)]) # valid - to_uops_list([buf.index(UOp.const(15)).load(dtype=dtypes.int)]) # valid (last element) + to_uops_list([buf.index(UOp.const(0)).load()]) # valid + to_uops_list([buf.index(UOp.const(15)).load()]) # valid (last element) with self.assertRaises(RuntimeError): - to_uops_list([buf.index(UOp.const(16)).load(dtype=dtypes.int)]) # off by one + to_uops_list([buf.index(UOp.const(16)).load()]) # off by one with self.assertRaises(RuntimeError): - to_uops_list([buf.index(UOp.const(42)).load(dtype=dtypes.int)]) # way out + to_uops_list([buf.index(UOp.const(42)).load()]) # way out def test_variable_index(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) - to_uops_list([buf.index(Variable("i", 0, 15)).load(dtype=dtypes.int)]) # valid + to_uops_list([buf.index(Variable("i", 0, 15)).load()]) # valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(Variable("i", 0, 20)).load(dtype=dtypes.int)]) # oob + to_uops_list([buf.index(Variable("i", 0, 20)).load()]) # oob with self.assertRaises(RuntimeError): - to_uops_list([buf.index(Variable("i", -5, 10)).load(dtype=dtypes.int)]) # negative + to_uops_list([buf.index(Variable("i", -5, 10)).load()]) # negative def test_range_with_mask(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) r = UOp.range(42, 0, AxisType.GLOBAL) - to_uops_list([buf.index(r.valid(r < 16)).load(dtype=dtypes.int)]) # valid + to_uops_list([buf.index(r.valid(r < 16)).load()]) # valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r.valid(r < 17)).load(dtype=dtypes.int)]) # oob + to_uops_list([buf.index(r.valid(r < 17)).load()]) # oob def test_variable_with_mask(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) v = Variable("v", -5, 80) - to_uops_list([buf.index(v.valid((v >= 0) & (v < 16))).load(dtype=dtypes.int)]) # valid + to_uops_list([buf.index(v.valid((v >= 0) & (v < 16))).load()]) # valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(v.valid(v < 20)).load(dtype=dtypes.int)]) # negative not masked + to_uops_list([buf.index(v.valid(v < 20)).load()]) # negative not masked def test_gated_store(self): with Context(CHECK_OOB=1, SPEC=2): @@ -58,62 +58,62 @@ class TestValidateOOB(unittest.TestCase): def test_floordiv(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) - to_uops_list([buf.index(UOp.range(32, 0, AxisType.GLOBAL) // 2).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(UOp.range(32, 0, AxisType.GLOBAL) // 2).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(UOp.range(34, 0, AxisType.GLOBAL) // 2).load(dtype=dtypes.int)]) # 0..16 oob + to_uops_list([buf.index(UOp.range(34, 0, AxisType.GLOBAL) // 2).load()]) # 0..16 oob def test_mod(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) r = UOp.range(100, 0, AxisType.GLOBAL) - to_uops_list([buf.index(r % 16).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(r % 16).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r % 20).load(dtype=dtypes.int)]) # 0..19 oob + to_uops_list([buf.index(r % 20).load()]) # 0..19 oob def test_shr(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) - to_uops_list([buf.index(UOp.range(64, 0, AxisType.GLOBAL) >> 2).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(UOp.range(64, 0, AxisType.GLOBAL) >> 2).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(UOp.range(128, 0, AxisType.GLOBAL) >> 2).load(dtype=dtypes.int)]) # 0..31 oob + to_uops_list([buf.index(UOp.range(128, 0, AxisType.GLOBAL) >> 2).load()]) # 0..31 oob def test_shl(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (64,)) r = UOp.range(8, 0, AxisType.GLOBAL) - to_uops_list([buf.index(r << 2).load(dtype=dtypes.int)]) # 0..28 valid + to_uops_list([buf.index(r << 2).load()]) # 0..28 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r << 4).load(dtype=dtypes.int)]) # 0..112 oob + to_uops_list([buf.index(r << 4).load()]) # 0..112 oob def test_and(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) r = UOp.range(100, 0, AxisType.GLOBAL) - to_uops_list([buf.index(r & 15).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(r & 15).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r & 31).load(dtype=dtypes.int)]) # 0..31 oob + to_uops_list([buf.index(r & 31).load()]) # 0..31 oob # align masks round down to a multiple of 2^k - to_uops_list([buf.index((r & -4).valid(r < 16)).load(dtype=dtypes.int)]) # 0..12 valid + to_uops_list([buf.index((r & -4).valid(r < 16)).load()]) # 0..12 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r & -2).load(dtype=dtypes.int)]) # 0..100 oob + to_uops_list([buf.index(r & -2).load()]) # 0..100 oob # other masks can't be modeled as mod with self.assertRaisesRegex(RuntimeError, "z3 int AND only supports"): - to_uops_list([buf.index(r & 21).load(dtype=dtypes.int)]) + to_uops_list([buf.index(r & 21).load()]) def test_max(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) - to_uops_list([buf.index(Variable("v", -10, 15).maximum(0)).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(Variable("v", -10, 15).maximum(0)).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(Variable("v2", -10, 20).maximum(0)).load(dtype=dtypes.int)]) # 0..20 oob + to_uops_list([buf.index(Variable("v2", -10, 20).maximum(0)).load()]) # 0..20 oob def test_xor_in_mask(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (16,)) r = UOp.range(32, 0, AxisType.GLOBAL) - to_uops_list([buf.index(r.valid((r < 8) ^ ((r >= 8) & (r < 16)))).load(dtype=dtypes.int)]) # 0..15 valid + to_uops_list([buf.index(r.valid((r < 8) ^ ((r >= 8) & (r < 16)))).load()]) # 0..15 valid with self.assertRaises(RuntimeError): - to_uops_list([buf.index(r.valid((r < 10) ^ (r >= 20))).load(dtype=dtypes.int)]) # 0..9,20..31 oob + to_uops_list([buf.index(r.valid((r < 10) ^ (r >= 20))).load()]) # 0..9,20..31 oob # cast patterns def test_float_cast_in_index(self): @@ -121,13 +121,13 @@ class TestValidateOOB(unittest.TestCase): buf = UOp.param(0, dtypes.int, (16,)) r = UOp.range(20, 0) i = (r.cast(dtypes.float) * 0.68).trunc().cast(dtypes.int) - to_uops_list([buf.index(i.valid((i >= 0) & (i < 16))).load(dtype=dtypes.int)]) + to_uops_list([buf.index(i.valid((i >= 0) & (i < 16))).load()]) def test_bool_cast_in_mask(self): with Context(CHECK_OOB=1, SPEC=2): buf = UOp.param(0, dtypes.int, (1,)) r = UOp.range(20, 0) - to_uops_list([buf.index(r.valid(r.cast(dtypes.bool).logical_not())).load(dtype=dtypes.int)]) # only r=0 valid + to_uops_list([buf.index(r.valid(r.cast(dtypes.bool).logical_not())).load()]) # only r=0 valid # load result as index/mask def test_load_as_index(self): @@ -135,18 +135,18 @@ class TestValidateOOB(unittest.TestCase): buf0 = UOp.param(0, dtypes.int, (16,)) buf1 = UOp.param(1, dtypes.int, (64,)) r = UOp.range(42, 0, AxisType.GLOBAL) - ld0 = buf0.index(r.valid(r < 8)).load(dtype=dtypes.int).cast(dtypes.weakint) - to_uops_list([buf1.index((ld0 * 2).valid((ld0 >= 0) & (ld0 < 32))).load(dtype=dtypes.int)]) # valid + ld0 = buf0.index(r.valid(r < 8)).load().cast(dtypes.weakint) + to_uops_list([buf1.index((ld0 * 2).valid((ld0 >= 0) & (ld0 < 32))).load()]) # valid with self.assertRaises(RuntimeError): - to_uops_list([buf1.index((ld0 * 2).valid((ld0 >= 0) & (ld0 < 64))).load(dtype=dtypes.int)]) # oob + to_uops_list([buf1.index((ld0 * 2).valid((ld0 >= 0) & (ld0 < 64))).load()]) # oob def test_load_from_shrink_as_index(self): with Context(CHECK_OOB=1, SPEC=2): buf0 = UOp.param(0, dtypes.int, (16,)) buf1 = UOp.param(1, dtypes.int, (64,)) shrink = UOp(Ops.SHRINK, src=(buf0, UOp.const(0, dtypes.int), UOp.const(4))) - ld0 = shrink.load(dtype=dtypes.int).index(0) - to_uops_list([buf1.index(ld0.valid((ld0 >= 0) & (ld0 < 64))).load(dtype=dtypes.int)]) + ld0 = shrink.load().index(0) + to_uops_list([buf1.index(ld0.valid((ld0 >= 0) & (ld0 < 64))).load()]) def test_load_bool_as_mask(self): with Context(CHECK_OOB=1, SPEC=2): diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index fd3ee107e7..47f7ba5941 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -502,7 +502,7 @@ pm_bufferize = PatternMatcher([(UPat(Ops.PARAM, name="buf"), bufferize_buf)]) # 7. resolve patches def push_stack(op, s): return UOp(Ops.STACK, - src=tuple(op.replace(dtype=op.dtype, src=tuple(x if y is s else y for y in op.src)) for x in s.src)) + src=tuple(op.replace(src=tuple(x if y is s else y for y in op.src)) for x in s.src)) def fold_binary(buf:UOp, blob:UOp) -> UOp: for b in (m.bufs if isinstance(m:=buf.buffer, MultiBuffer) else (m,)):