late cast

This commit is contained in:
2026-06-24 18:12:19 -07:00
parent bc1c45fb75
commit fd9247ef41
2 changed files with 14 additions and 9 deletions
+3 -3
View File
@@ -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)),
])
+11 -6
View File
@@ -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: