From 138fb4a783d82f4e877ad2fe3692aaf8d1de2e46 Mon Sep 17 00:00:00 2001 From: chenyu Date: Sun, 16 Aug 2026 21:12:17 -0400 Subject: [PATCH] delete dead DType.scalar [PR] (#17561) --- extra/thunder/tiny/tk/tiles.py | 2 +- test/backend/test_isel.py | 2 +- test/backend/test_quantize_onnx.py | 2 +- test/null/test_dtype_spec.py | 4 ---- test/null/test_tensor.py | 2 +- test/opt/test_gen_float4.py | 8 ++++---- test/unit/test_allreduce.py | 2 +- tinygrad/dtype.py | 1 - tinygrad/renderer/__init__.py | 8 ++++---- tinygrad/renderer/cstyle.py | 12 ++++++------ tinygrad/renderer/isa/x86.py | 8 ++++---- tinygrad/renderer/ptx.py | 26 +++++++++++++------------- tinygrad/runtime/support/hcq2.py | 2 +- 13 files changed, 37 insertions(+), 42 deletions(-) diff --git a/extra/thunder/tiny/tk/tiles.py b/extra/thunder/tiny/tk/tiles.py index 98cff727c7..f27ee77fc7 100644 --- a/extra/thunder/tiny/tk/tiles.py +++ b/extra/thunder/tiny/tk/tiles.py @@ -209,7 +209,7 @@ class ST: return cls(uop, rows, cols, layout, base_shape, ker) def swizzle(self, row, col): - swizzled_offset = self.base_shape.swizzle(row, col, self._uop.dtype.scalar()) + swizzled_offset = self.base_shape.swizzle(row, col, self._uop.dtype) row = swizzled_offset // self.base_shape.cols col = swizzled_offset % self.base_shape.cols diff --git a/test/backend/test_isel.py b/test/backend/test_isel.py index 6965a5db17..ed38be7c3b 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.const(i, dtypes.int), dtype=y.dtype.scalar()) +def lane(y:UOp, i:int) -> UOp: return y.index(UOp.const(i, dtypes.int), dtype=y.dtype) @unittest.skipUnless(isinstance(Device[Device.DEFAULT].renderer, X86Renderer), "only x86") class TestIselX86(unittest.TestCase): diff --git a/test/backend/test_quantize_onnx.py b/test/backend/test_quantize_onnx.py index 587b03f884..8528caa9c4 100644 --- a/test/backend/test_quantize_onnx.py +++ b/test/backend/test_quantize_onnx.py @@ -82,7 +82,7 @@ class TestQuantizeOnnxCPU(unittest.TestCase): linear = run_onnx({"input":inp})["output"].schedule_linear() prg = to_program(linear.src[-2].src[0], renderer=Device[Device.DEFAULT].renderer) daccs = [u for u in tuple(prg.src[1].src) if u.op is Ops.BUFFER and u.addrspace is AddrSpace.REG] - assert all(u.dtype.scalar() is dtypes.int for u in daccs) + assert all(u.dtype is dtypes.int for u in daccs) @unittest.skipIf(Device.DEFAULT != "DSP", "only tests for DSP") class TestQuantizeOnnx(unittest.TestCase): diff --git a/test/null/test_dtype_spec.py b/test/null/test_dtype_spec.py index 2aaebefe05..e963ffb6c0 100644 --- a/test/null/test_dtype_spec.py +++ b/test/null/test_dtype_spec.py @@ -51,10 +51,6 @@ class TestHelpers(unittest.TestCase): assert dtypes.is_float(dtypes.fp8e4m3) assert dtypes.is_float(dtypes.fp8e5m2) - @given(strat.sampled_from([d for d in DTYPES_DICT.values() if dtypes.is_float(d) or dtypes.is_int(d)])) - def test_scalar(self, dtype): - assert dtype.scalar() == dtype - def test_from_py(self): assert dtypes.from_py(True) == dtypes.bool assert dtypes.from_py(Invalid) == dtypes.bool diff --git a/test/null/test_tensor.py b/test/null/test_tensor.py index 4156b37177..32ccc255e3 100644 --- a/test/null/test_tensor.py +++ b/test/null/test_tensor.py @@ -69,7 +69,7 @@ class TestIdxUpcast(unittest.TestCase): if not isinstance(Device[Device.DEFAULT].renderer, (PTXRenderer, NIRRenderer)): assert idx.op is Ops.INDEX idx_val = idx.src[1] - self.assertFalse(idx_val.overflows(idx_val.dtype.scalar())) + self.assertFalse(idx_val.overflows(idx_val.dtype)) # use expand to generate kernel that uses large idx def do_op_then_assert(self, dtype: DType, dim1, dim2, dim3): diff --git a/test/opt/test_gen_float4.py b/test/opt/test_gen_float4.py index b946cb6bd1..b09aabe50f 100644 --- a/test/opt/test_gen_float4.py +++ b/test/opt/test_gen_float4.py @@ -10,12 +10,12 @@ from test.helpers import replace_opts class TestFloat4(unittest.TestCase): @staticmethod def count_float4(uops: list[UOp], n=4): - return (len([uop for uop in uops if uop.op is Ops.LOAD and uop.dtype.scalar() == dtypes.float and uop.shape == (4,)]), - len([uop for uop in uops if uop.op is Ops.STORE and uop.src[1].dtype.scalar() == dtypes.float and uop.shape == (4,)])) + return (len([uop for uop in uops if uop.op is Ops.LOAD and uop.dtype == dtypes.float and uop.shape == (4,)]), + len([uop for uop in uops if uop.op is Ops.STORE and uop.src[1].dtype == dtypes.float and uop.shape == (4,)])) @staticmethod def count_half4(uops: list[UOp]): - return (len([uop for uop in uops if uop.op is Ops.LOAD and uop.dtype.scalar() == dtypes.half and uop.shape == (4,)]), - len([uop for uop in uops if uop.op is Ops.STORE and uop.src[1].dtype.scalar() == dtypes.half and uop.shape == (4,)])) + return (len([uop for uop in uops if uop.op is Ops.LOAD and uop.dtype == dtypes.half and uop.shape == (4,)]), + len([uop for uop in uops if uop.op is Ops.STORE and uop.src[1].dtype == dtypes.half and uop.shape == (4,)])) def test_float4_basic(self): a = Tensor.empty(2, 8).realize() diff --git a/test/unit/test_allreduce.py b/test/unit/test_allreduce.py index 0997e55979..84dce58242 100644 --- a/test/unit/test_allreduce.py +++ b/test/unit/test_allreduce.py @@ -64,7 +64,7 @@ class TestAllreduceCast(unittest.TestCase): with Context(ALLREDUCE_CAST=allreduce_cast, RING=0, SCACHE=0): t = Tensor.empty(4, 4, dtype=dtype).shard(ds, axis=0) linear = t.sum(0).linear_with_vars()[0] - return {si.src[1].buffer.dtype.scalar() for si in linear.src if si.src[0].op is Ops.COPY} + return {si.src[1].buffer.dtype for si in linear.src if si.src[0].op is Ops.COPY} def test_allreduce_cast_bf16(self): # with ALLREDUCE_CAST, allreduce copies stay in bfloat16 instead of promoting to float32 diff --git a/tinygrad/dtype.py b/tinygrad/dtype.py index 9c0b336a58..e160cb6cdd 100644 --- a/tinygrad/dtype.py +++ b/tinygrad/dtype.py @@ -66,7 +66,6 @@ class DType(metaclass=DTypeMetaClass): def __reduce__(self): return type(self), tuple(getattr(self, f.name) for f in fields(self)) def __repr__(self): return f"dtypes.{INVERSE_DTYPES_DICT[self.name]}" def __lt__(self, o:DType): return (self.priority, self.bitsize, self.name, self.fmt) < (o.priority, o.bitsize, o.name, o.fmt) - def scalar(self) -> DType: return self @functools.cached_property def min(self): if dtypes.is_int(self): return 0 if dtypes.is_unsigned(self) else -2**(self.bitsize-1) diff --git a/tinygrad/renderer/__init__.py b/tinygrad/renderer/__init__.py index 438af612be..728e6148d3 100644 --- a/tinygrad/renderer/__init__.py +++ b/tinygrad/renderer/__init__.py @@ -35,8 +35,8 @@ class Estimates: while len(buf.src) and buf.op is not Ops.PARAM: buf = buf.src[0] if buf.op is Ops.PARAM: # u.src[0] is INDEX, cap at buffer size for re-reads (e.g. matmul) - accessed = mem.get((buf, u.op), 0) + u.src[0].max_numel() * u.src[0].dtype.scalar().itemsize * mults - mem[(buf, u.op)] = smin(accessed, buf.max_numel() * buf.dtype.scalar().itemsize) + accessed = mem.get((buf, u.op), 0) + u.src[0].max_numel() * u.src[0].dtype.itemsize * mults + mem[(buf, u.op)] = smin(accessed, buf.max_numel() * buf.dtype.itemsize) if u.op is Ops.RANGE: mult_stack.append(mults) if u.dtype is not dtypes.void: # unbounded loop, unknown trip count @@ -47,9 +47,9 @@ class Estimates: elif u.op is Ops.SPECIAL: mults *= cast(sint, u.src[0].ssimplify()) # NOTE: we don't push to the mult_stack here, you can't end these elif u.op is Ops.PARAM and u.arg.addrspace == AddrSpace.ALU and u.expr == 'core_id': mults *= int(u.vmax) + 1 elif u.op is Ops.LOAD and u.src[0].addrspace != AddrSpace.REG: - lds += u.max_numel() * u.dtype.scalar().itemsize * mults + lds += u.max_numel() * u.dtype.itemsize * mults elif u.op is Ops.STORE and u.src[0].addrspace != AddrSpace.REG: - lds += u.max_numel() * u.src[1].dtype.scalar().itemsize * mults + lds += u.max_numel() * u.src[1].dtype.itemsize * mults elif u.op in GroupOp.ALU and u not in excluded: flops += (mults * (2 if u.op is Ops.MULACC else 1)) * u.max_numel() elif u.op is Ops.WMMA and u not in excluded: diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 3b555ac177..7e220512ac 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -107,11 +107,11 @@ def uops_to_dtypes(uops:list[UOp]) -> list[tuple[DType, int]]: def _wmma_name(u:UOp) -> str: # sanitize spaces in DType.name (int8 = "signed char") - return f"WMMA_{'_'.join(map(str, u.arg[0]))}_{u.arg[1].name}_{u.dtype.scalar().name}".replace(" ", "_") + return f"WMMA_{'_'.join(map(str, u.arg[0]))}_{u.arg[1].name}_{u.dtype.name}".replace(" ", "_") # (name, dims, dtype_in, dtype_out, device, threads, upcast_sizes) def wmma_args(uops:list[UOp]): - return dedup((_wmma_name(uop), uop.arg[0], uop.arg[1], uop.dtype.scalar(), *(uop.arg[2:4]), + return dedup((_wmma_name(uop), uop.arg[0], uop.arg[1], uop.dtype, *(uop.arg[2:4]), tuple(uop.src[i].shape[-1] for i in range(3))) for uop in uops if uop.op is Ops.WMMA) @@ -182,8 +182,8 @@ class CStyleLanguage(Renderer): if addrspace in (AddrSpace.LOCAL, AddrSpace.GLOBAL) or override_ptr: suffix = "*" if sz > 1: - return prefix + self.type_map.get(scalar:=dtype.scalar(), scalar.name).replace(" ", "_") + str(sz) + suffix - return prefix + self.type_map.get(scalar:=dtype.scalar(), scalar.name) + suffix + return prefix + self.type_map.get(dtype, dtype.name).replace(" ", "_") + str(sz) + suffix + return prefix + self.type_map.get(dtype, dtype.name) + suffix def render_type(self, u:UOp): return self._render_dtype(u.dtype, u.max_numel(), u.addrspace, shape=u._shape) def render_access(self, u:UOp): @@ -472,7 +472,7 @@ class CUDARenderer(CStyleLanguage): class NVCCRenderer(CUDARenderer): def __init__(self, target:Target): super().__init__(target, use_nvcc=True) -def fp8_index(dtype: DType): return (dtypes.fp8e4m3, dtypes.fp8e5m2).index(dtype.scalar()) +def fp8_index(dtype: DType): return (dtypes.fp8e4m3, dtypes.fp8e5m2).index(dtype) def _ocml(op): return lambda x,dtype: f"__ocml_{op}_f{ {dtypes.half:16, dtypes.double:64}.get(dtype, 32)}({x})" class HIPRenderer(CStyleLanguage): @@ -546,7 +546,7 @@ class HIPRenderer(CStyleLanguage): ockl = [(f"__ockl_get_{name}", "unsigned int", "size_t", "const") for name in ["local_id", "group_id", "local_size"]] ocml_ops = {Ops.EXP2: ("exp2", "pure"), Ops.LOG2: ("log2", "pure"), Ops.SQRT: ("sqrt", "const"), Ops.SIN: ("sin", ""), Ops.TRUNC: ("trunc", "")} ocml = [(f"__ocml_{ocml_ops[op][0]}_f{dt.bitsize}", dt.name, dt.name, ocml_ops[op][1]) - for op, dt in dedup((u.op, u.dtype.scalar()) for u in uops) if op in ocml_ops and dt in (dtypes.half, dtypes.float, dtypes.double)] + for op, dt in dedup((u.op, u.dtype) for u in uops) if op in ocml_ops and dt in (dtypes.half, dtypes.float, dtypes.double)] if any(dt == dtypes.bfloat16 for dt, _ in used_dtypes): prefix.append(f"typedef {'__bf16' if self.is_cdna4(self.target.arch) else 'unsigned short'} hip_bfloat16;") if any(dt == dtypes.half for dt, _ in used_dtypes): prefix.append("#define half _Float16") diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index 1346a02678..85b4ece0d2 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -165,7 +165,7 @@ def scratch_buffer(elem_dt:DType, count:int, slot:int) -> UOp: return UOp.placeholder((count,), elem_dt, slot, AddrSpace.LOCAL) def gated_load(ctx, addr:UOp, alt:UOp, gate:UOp, x:UOp): - local = scratch_buffer(addr.src[0].dtype.scalar(), x.max_numel(), next(ctx)) + local = scratch_buffer(addr.src[0].dtype, x.max_numel(), next(ctx)) local_idx = local.index(UOp.const(0, dtypes.int32), dtype=dtypes.uint64) # the selected address is a 64bit value, the AFTER orders the load after the scratch store and carries the element dtype for the encoder sel = gate.where(addr.replace(dtype=dtypes.uint64), local_idx) @@ -173,7 +173,7 @@ def gated_load(ctx, addr:UOp, alt:UOp, gate:UOp, x:UOp): return ptr.load(dtype=x.dtype) def gated_store(addr:UOp, gate:UOp, val:UOp): - local = scratch_buffer(addr.src[0].dtype.scalar(), val.max_numel(), -1) + local = scratch_buffer(addr.src[0].dtype, val.max_numel(), -1) sel = gate.where(addr.replace(dtype=dtypes.uint64), local.index(UOp.const(0, dtypes.int32), dtype=dtypes.uint64)) return UOp(Ops.AFTER, addr.dtype, (sel,)).store(val) @@ -237,7 +237,7 @@ def cmp(x:UOp) -> UOp: return x.ins(X86Ops.CMP, dtype=dtypes.void) if (i:=to_imm(x.src[1])) is None else x.ins(X86Ops.CMPi, dtype=dtypes.void, src=(x.src[0], i)) def vcmp(x:UOp) -> UOp: v = imm(dtypes.uint8, {Ops.CMPLT: 1, Ops.CMPNE: 4, Ops.CMPEQ: 0}[x.op]) - if x.dtype.scalar() is dtypes.float32: return x.ins(X86Ops.VCMPSS if x.max_numel() == 1 else X86Ops.VCMPPS, src=x.src + (v,)) + if x.dtype is dtypes.float32: return x.ins(X86Ops.VCMPSS if x.max_numel() == 1 else X86Ops.VCMPPS, src=x.src + (v,)) return x.ins(X86Ops.VCMPSD if x.max_numel() == 1 else X86Ops.VCMPPD, src=x.src + (v,)) # vinsertps xmm2, xmm0, xmm1, imm @@ -252,7 +252,7 @@ def vinsertps(x:UOp) -> UOp: # vpinsq xmm2, xmm0, rax, imm # inserts element in rax into any position in xmm0, result is written to xmm2 according to imm def vpins(x:UOp) -> UOp: - op = {1: X86Ops.VPINSRB, 2: X86Ops.VPINSRW, 4: X86Ops.VPINSRD, 8: X86Ops.VPINSRQ}[x.dtype.scalar().itemsize] + op = {1: X86Ops.VPINSRB, 2: X86Ops.VPINSRW, 4: X86Ops.VPINSRD, 8: X86Ops.VPINSRQ}[x.dtype.itemsize] return functools.reduce(lambda ret,i: x.ins(op, src=(ret, x.src[i], imm(dtypes.uint8, i))), range(len(x.src)), def_reg(x.dtype)) # we don't call ctx.vreg on the srcs to avoid duplicates, a rewrite will assign the tuple of valid registers to a vreg diff --git a/tinygrad/renderer/ptx.py b/tinygrad/renderer/ptx.py index 1b6b77859a..30c73c2f93 100644 --- a/tinygrad/renderer/ptx.py +++ b/tinygrad/renderer/ptx.py @@ -64,7 +64,7 @@ def render_wmma(ctx: "PTXRenderer", wmma: UOp): for src, regs in zip(wmma.src, ctx.wmma_r): for i, reg in enumerate(regs): # pack input and acc registers - if (elems_per_reg := 4 // src.dtype.scalar().itemsize) == 1: yield f"mov.b32 {reg}, {ctx.r[src][i]};" + if (elems_per_reg := 4 // src.dtype.itemsize) == 1: yield f"mov.b32 {reg}, {ctx.r[src][i]};" else: yield f"mov.b32 {reg}, {{{', '.join(ctx.r[src][i * elems_per_reg : (i+1) * elems_per_reg])}}};" dt_map_in, dt_map_out = {dtypes.float: "tf32", dtypes.half: "f16"}, {dtypes.float: "f32", dtypes.half: "f16"} @@ -101,17 +101,17 @@ string_rewrite = PatternMatcher([ if loc.addrspace == AddrSpace.REG else None), (UPat(Ops.STORE, src=(UPat((Ops.INDEX, Ops.SHRINK), name="loc"), UPat.var("var"))), lambda ctx, loc, var: f"st.{mem_type(loc)}" + \ - f"{f'.v{cnt}' if ((cnt:=var.max_numel())>1) else ''}.{ctx.mem_types[var.dtype.scalar()]} " + \ + f"{f'.v{cnt}' if ((cnt:=var.max_numel())>1) else ''}.{ctx.mem_types[var.dtype]} " + \ f"[{ctx.r[loc]}+0], {('{' + ', '.join(ctx.r[var]) + '}') if var.max_numel() > 1 else ctx.r[var]};"), (UPat(Ops.LOAD, name="x", src=(UPat((Ops.INDEX, Ops.SHRINK), name="loc"), UPat.var("alt"), UPat.var("gate"))), lambda ctx, x, loc, alt, gate: 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(loc)}.v{x.max_numel()}.{ctx.mem_types[x.dtype.scalar()]} {{{', '.join(ctx.r[x])}}}, [{ctx.r[loc]}+0];"] + [f"mov.{ctx.mem_types[x.dtype]} {v}, {render_val(0, x.dtype)};" for v in ctx.r[x]], + [f"@{ctx.r[gate]} ld.{mem_type(loc)}.v{x.max_numel()}.{ctx.mem_types[x.dtype]} {{{', '.join(ctx.r[x])}}}, [{ctx.r[loc]}+0];"] ]) if alt.max_numel() > 1 else [ - f"@{ctx.r[gate]} ld.{mem_type(loc)}.{ctx.mem_types[x.dtype.scalar()]} {ctx.r[x]}, [{ctx.r[loc]}+0];", - f"@!{ctx.r[gate]} mov.b{ctx.types[x.dtype.scalar()][1:]} {ctx.r[x]}, {ctx.r[alt]};"]), + f"@{ctx.r[gate]} ld.{mem_type(loc)}.{ctx.mem_types[x.dtype]} {ctx.r[x]}, [{ctx.r[loc]}+0];", + f"@!{ctx.r[gate]} mov.b{ctx.types[x.dtype][1:]} {ctx.r[x]}, {ctx.r[alt]};"]), (UPat(Ops.LOAD, name="x", src=(UPat((Ops.INDEX, Ops.SHRINK), name="loc"),)), - lambda ctx, x, loc: f"ld.{mem_type(loc)}.v{x.max_numel()}.{ctx.mem_types[x.dtype.scalar()]} {{{', '.join(ctx.r[x])}}}, [{ctx.r[loc]}+0];" \ + lambda ctx, x, loc: f"ld.{mem_type(loc)}.v{x.max_numel()}.{ctx.mem_types[x.dtype]} {{{', '.join(ctx.r[x])}}}, [{ctx.r[loc]}+0];" \ if x.max_numel() > 1 else f"ld.{mem_type(loc)}.{ctx.mem_types[x.dtype]} {ctx.r[x]}, [{ctx.r[loc]}+0];"), # simple (UPat(Ops.BUFFER, name="x"), lambda ctx, x: [] if x.addrspace == AddrSpace.REG else [ @@ -197,7 +197,7 @@ class PTXRenderer(Renderer): r[u] = [cast(str,r[x]) for x in u.src] continue if u.op is Ops.BUFFER and u.addrspace == AddrSpace.REG: - r[u] = [ssa("reg", u, self.types[u.dtype.scalar()]) for _ in range(u.max_numel())] + r[u] = [ssa("reg", u, self.types[u.dtype]) for _ in range(u.max_numel())] continue if u.op in {Ops.INDEX, Ops.SHRINK, Ops.LOAD} and u.src[0].addrspace in (AddrSpace.REG, AddrSpace.ALU): # on REG, INDEX/SHRINK pick the register (must be CONST) and LOAD is a noop @@ -207,14 +207,14 @@ class PTXRenderer(Renderer): continue if u.op is Ops.SPECIAL: r[u] = "%" + u.arg elif u.op is Ops.LOAD: - r[u] = [ssa('val', dtype=self.types[u.dtype.scalar()]) for _ in range(u.max_numel())] if u.max_numel() > 1 else ssa('val', u) + r[u] = [ssa('val', dtype=self.types[u.dtype]) for _ in range(u.max_numel())] if u.max_numel() > 1 else ssa('val', u) elif u.op is Ops.PARAM: bufs.append((f"data{u.arg.slot}", u)) elif u.op is Ops.WMMA: # registers for packing/unpacking input and acc - self.wmma_r = [[ssa("wmma_in", dtype="b32") for _ in range(0, len(r[u.src[0]]), 4 // u.src[0].dtype.scalar().itemsize)], - [ssa("wmma_in", dtype="b32") for _ in range(0, len(r[u.src[1]]), 4 // u.src[0].dtype.scalar().itemsize)], - [ssa("wmma_acc", dtype="b32") for _ in range(0, len(r[u.src[2]]), 4 // u.dtype.scalar().itemsize)]] - r[u] = [ssa("wmma", dtype=self.types[u.dtype.scalar()]) for _ in range(u.max_numel())] + self.wmma_r = [[ssa("wmma_in", dtype="b32") for _ in range(0, len(r[u.src[0]]), 4 // u.src[0].dtype.itemsize)], + [ssa("wmma_in", dtype="b32") for _ in range(0, len(r[u.src[1]]), 4 // u.src[0].dtype.itemsize)], + [ssa("wmma_acc", dtype="b32") for _ in range(0, len(r[u.src[2]]), 4 // u.dtype.itemsize)]] + r[u] = [ssa("wmma", dtype=self.types[u.dtype]) for _ in range(u.max_numel())] prefix, dtype = {Ops.CAST: ("cast", None), Ops.BITCAST: ("cast", None), Ops.END: ("pred", "pred"), Ops.RANGE: ("ridx", None), Ops.CONST: ("const", None), Ops.BUFFER: ("local", "u64"), Ops.INDEX: ("bidx", "u64"), Ops.SHRINK: ("bidx", "u64"), Ops.PARAM: ("dat", "u64" if u.addrspace is AddrSpace.GLOBAL else None), **{op: ("alu", None) for op in GroupOp.ALU}}.get(u.op, (None, None)) diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index 5b468a8741..0d42960a39 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -466,7 +466,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.scalar(), src=tuple(x if y is s else y for y in op.src)) for x in s.src)) + 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)) def fold_binary(buf:UOp, blob:UOp) -> UOp: for b in (m.bufs if isinstance(m:=buf.buffer, MultiBuffer) else (m,)):