From fd9247ef4195be367da2a0ab152d934b69ade6be Mon Sep 17 00:00:00 2001 From: George Hotz Date: Wed, 24 Jun 2026 18:12:19 -0700 Subject: [PATCH] late cast --- tinygrad/codegen/late/coalese.py | 6 +++--- tinygrad/renderer/cstyle.py | 17 +++++++++++------ 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/tinygrad/codegen/late/coalese.py b/tinygrad/codegen/late/coalese.py index 16b12c11f4..c6e2d8d186 100644 --- a/tinygrad/codegen/late/coalese.py +++ b/tinygrad/codegen/late/coalese.py @@ -27,7 +27,7 @@ def transform_to_image(ctx, buf, x, valid=None): idx = buf.replace(dtype=(dtypes.imageh if buf.dtype.itemsize == 2 else dtypes.imagef)((h, w, 4))).index(cidx.src[1], cidx.src[0]) if valid is not None: # TODO: simplify valid here - idx = valid.where(idx, UOp(Ops.CONST, dtype=dtypes.idx, arg=Invalid)) + idx = valid.where(idx, UOp(Ops.CONST, dtype=idx.dtype, arg=Invalid)) return idx pm_add_image = PatternMatcher([ @@ -37,8 +37,8 @@ pm_add_image = PatternMatcher([ pm_new_gater = PatternMatcher([ # here we create the alt value for load to be 0s and remove the where Invalid - (UPat.var("gate").where(UPat.var("idx"), UPat(Ops.CONST, arg=Invalid)).load(name="l"), - lambda gate,idx,l: idx.load(l.vconst_like(0), gate)), + (UPat.var("gate").where(UPat.var("idx"), UPat(Ops.CONST, arg=Invalid)).load(), + lambda gate,idx: idx.load(idx.vconst_like(0), gate)), (UPat.var("gate").where(UPat.var("idx"), UPat(Ops.CONST, arg=Invalid)).store(UPat.var("data")), lambda gate,idx,data: idx.store(data, gate)), ]) diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index ca7be3fa88..b85ca12f33 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -314,12 +314,17 @@ class OpenCLRenderer(CStyleLanguage): lambda ctx,x: f"{(struct.unpack('I', struct.pack('f', float_to_bf16(x.arg)))[0] >> 16)}u"), # load/store image (OpenCL) (UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')), lambda ctx,buf,idx_y,idx_x: f"IMAGE<{ctx[buf]}, {ctx[idx_y]}, {ctx[idx_x]}>"), - (UPat(Ops.LOAD, dtype=dtypes.float, src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')), UPat.var("var"), UPat.var("gate"))), - lambda ctx,buf,idx_y,idx_x,var,gate: f"({ctx[gate]}?read_imagef({ctx[buf]}, smp, (int2)({ctx[idx_x]},{ctx[idx_y]})):{ctx[var]})"), - (UPat(Ops.LOAD, dtype=dtypes.float, src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')),)), - lambda ctx,buf,idx_y,idx_x: f"read_imagef({ctx[buf]}, smp, (int2)({ctx[idx_x]},{ctx[idx_y]}))"), - (UPat(Ops.STORE, src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')), UPat.var("var", dtypes.float))), - lambda ctx,buf,idx_y,idx_x,var: f"write_imagef({ctx[buf]}, (int2)({ctx[idx_x]},{ctx[idx_y]}), {ctx[var]});"), + (UPat(Ops.LOAD, name="l", dtype=(dtypes.float,dtypes.half), + src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')), UPat.var("var"), UPat.var("gate"))), + lambda ctx,buf,idx_y,idx_x,var,gate,l: + f"({ctx[gate]}?read_image{'f' if l.dtype==dtypes.float else 'h'}({ctx[buf]}, smp, (int2)({ctx[idx_x]},{ctx[idx_y]})):{ctx[var]})"), + (UPat(Ops.LOAD, name="l", dtype=(dtypes.float,dtypes.half), + src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')),)), + lambda ctx,buf,idx_y,idx_x,l: + f"read_image{'f' if l.dtype==dtypes.float else 'h'}({ctx[buf]}, smp, (int2)({ctx[idx_x]},{ctx[idx_y]}))"), + (UPat(Ops.STORE, src=(UPat.var('buf').index(UPat.var('idx_y'), UPat.var('idx_x')), UPat.var("var", (dtypes.float,dtypes.half)))), + lambda ctx,buf,idx_y,idx_x,var: + f"write_image{'f' if var.dtype==dtypes.float else 'h'}({ctx[buf]}, (int2)({ctx[idx_x]},{ctx[idx_y]}), {ctx[var]});"), ]) + base_rewrite def render_kernel(self, function_name, kernel, bufs, uops, prefix=None) -> str: