From 49778d9a487b04bd60b3cd451cf57381c08fce6d Mon Sep 17 00:00:00 2001 From: chenyu Date: Tue, 18 Aug 2026 17:43:27 -0400 Subject: [PATCH] start renderer casted const migration [pr] (#17582) before rendering, rewrite strong typed const to casted weak const and have renderer adopt the new UOp. starting with PYTHON --- test/backend/test_linearizer.py | 3 ++- tinygrad/codegen/__init__.py | 11 +++++++++-- tinygrad/renderer/__init__.py | 2 ++ tinygrad/runtime/ops_python.py | 1 + tinygrad/uop/render.py | 2 +- tinygrad/uop/spec.py | 6 ++++++ tinygrad/uop/validate.py | 2 ++ 7 files changed, 23 insertions(+), 4 deletions(-) diff --git a/test/backend/test_linearizer.py b/test/backend/test_linearizer.py index 0f36811baa..ab63e732b5 100644 --- a/test/backend/test_linearizer.py +++ b/test/backend/test_linearizer.py @@ -252,7 +252,7 @@ class TestLinearizer(unittest.TestCase): for u in uops: if u.op is Ops.STORE and u.src[0].addrspace is AddrSpace.REG: if uops.index(u) < begin_range: - assert u.src[1].op is Ops.CONST + assert u.src[1].op not in GroupOp.ALU else: assert u.src[1].op in GroupOp.ALU assert begin_range < uops.index(u) < end_range @@ -261,6 +261,7 @@ class TestLinearizer(unittest.TestCase): assert end_range < uops.index(u) @unittest.skipUnless(Device[Device.DEFAULT].renderer.has_local, "test requires locals") + @unittest.skipIf(Device[Device.DEFAULT].renderer.casted_consts, "reads a literal, which is casted here. TODO: flip this") def test_default_global_reversed(self): # shrink so that the dims do not collapse t = Tensor.ones(5, 6, 7).contiguous().realize().shrink(((0, 4), (0, 5), (0, 6))) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 982ab35fa9..fb62365aa4 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -5,7 +5,7 @@ from tinygrad.helpers import ALLOW_TF32, DEFAULT_FLOAT, DEFAULT_INT, TracingKey, from tinygrad.uop.ops import PatternMatcher, graph_rewrite, UOp, Ops, UPat, rewrite_group, KernelInfo, ProgramInfo, GroupOp, AxisType from tinygrad.uop.weak import pm_lower_index_dtype, pm_commit_weak, pm_cast_weak from tinygrad.uop.render import pyrender -from tinygrad.uop.spec import type_verify, spec_tensor, spec_program +from tinygrad.uop.spec import type_verify, spec_tensor, spec_program, spec_program_casted_consts from tinygrad.renderer import Renderer, Estimates from tinygrad.renderer.isa import ISARenderer, IselContext, PreRegAllocContext from tinygrad.dtype import dtypes, AddrSpace @@ -281,6 +281,10 @@ pm_implicit_barriers = PatternMatcher([ (UPat(Ops.END, name="end"), add_war_barrier), ]) +pm_casted_consts = PatternMatcher([ + (UPat(Ops.CONST, dtypes.all, name="c"), lambda c: UOp(Ops.CAST, c.dtype, src=(UOp.const(c.val),), arg=c.dtype)), +]) + def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp: if VIZ: graph_rewrite(ast, PatternMatcher([]), name="View Base AST") if DEBUG >= 5: print(pyrender(ast)) @@ -383,8 +387,11 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp: num_params = len([x for x in sink.toposort() if x.op is Ops.PARAM and x.arg.slot != -1]) sink = graph_rewrite(sink, pm_number_params, ctx=[num_params], name="number params with -1", walk=True) + # TODO: delete once migration are done + if ren.casted_consts: sink = graph_rewrite(sink, pm_casted_consts, name="casted consts", walk=True) + if VIZ: graph_rewrite(sink, PatternMatcher([]), name="View Output AST") - if SPEC: type_verify(sink, spec_program) + if SPEC: type_verify(sink, spec_program_casted_consts if ren.casted_consts else spec_program) # return the rewritten sink return sink diff --git a/tinygrad/renderer/__init__.py b/tinygrad/renderer/__init__.py index 728e6148d3..ec42643ade 100644 --- a/tinygrad/renderer/__init__.py +++ b/tinygrad/renderer/__init__.py @@ -72,6 +72,8 @@ class Renderer: tensor_cores: list[TensorCore] = [] extra_matcher: PatternMatcher|None = None code_for_op: dict[Ops, Callable] = {} + # migration: this renderer consumes every literal as a casted const CAST(dt, CONST(value)) + casted_consts: bool = False compiler: Compiler = Compiler() diff --git a/tinygrad/runtime/ops_python.py b/tinygrad/runtime/ops_python.py index 0171dabd58..b9b383e021 100644 --- a/tinygrad/runtime/ops_python.py +++ b/tinygrad/runtime/ops_python.py @@ -214,6 +214,7 @@ class PythonCompiler(Compiler): class PythonRenderer(Renderer): code_for_op = python_alu compiler = PythonCompiler() + casted_consts: bool = True def __init__(self, target:Target): assert (emu:=getenv("EMULATE", "")) == "", ("EMULATE is deprecated, use DEV=PYTHON::" + diff --git a/tinygrad/uop/render.py b/tinygrad/uop/render.py index 6718d4636e..2cb516c960 100644 --- a/tinygrad/uop/render.py +++ b/tinygrad/uop/render.py @@ -81,7 +81,7 @@ sugar = {Ops.SINK, Ops.END, Ops.STORE, Ops.LOAD, Ops.SQRT, Ops.INDEX, Ops.REDUCE Ops.RECIPROCAL, Ops.EXP2, Ops.LOG2, Ops.SIN, Ops.CONTIGUOUS, Ops.BARRIER, Ops.DETACH} pm_pyrender_extra = PatternMatcher([ (UPat(Ops.CONST, src=(), name="x"), lambda x: f"UOp.const({x.val}, {x.dtype})"), - (UPat((Ops.CAST, Ops.BITCAST), name="x"), lambda ctx,x: f"{ctx[x.src[0]]}.{x.op.name.lower()}({x.dtype})"), + (UPat((Ops.CAST, Ops.BITCAST), name="x"), lambda ctx,x: f"{ctx[x.src[0]]}.{x.op.name.lower()}({x.dtype})" if x.dtype != x.src[0].dtype else None), (UPat(Ops.SPECIAL, src=(UPat(Ops.CONST),), name="x"), lambda x: f"UOp.special({x.src[0].val}, {repr(x.arg)}, dtype={x.dtype})"), (UPat(Ops.BUFFER, src=(UPat(),), name="x"), lambda x: f"UOp.new_buffer({repr(x.arg.device)}, {x.max_numel()}, {x.dtype}, {x.arg.slot})" diff --git a/tinygrad/uop/spec.py b/tinygrad/uop/spec.py index 35fa4bb94f..fae9daff7a 100644 --- a/tinygrad/uop/spec.py +++ b/tinygrad/uop/spec.py @@ -223,6 +223,12 @@ spec_program = PatternMatcher([ (UPat(Ops.SPECIAL, src=(UPat.var("x", dtypes.int32),), name="s"), lambda s,x: matches_dtype(x, s.dtype) and isinstance(s.arg, str)), ])+spec_shared +# migration: on a casted_consts renderer every literal is CAST(dt, CONST(value)) with a weak inner CONST +spec_program_casted_consts = PatternMatcher([ + (UPat(Ops.CONST, dtype=dtypes.weaks, name="x"), lambda x: x.dtype is dtypes.from_py(x.val)), + (UPat(Ops.SHRINK, src=(UPat((Ops.PARAM, Ops.BUFFER, Ops.AFTER)), UPat(), UPat(Ops.CAST, src=(UPat(Ops.CONST),)))), lambda: True), +])+spec_program + spec_hcq = PatternMatcher([ (UPat(Ops.GETADDR, dtypes.uint64, src=(UPat((Ops.BUFFER, Ops.PARAM)).or_after(),), name="x"), lambda x: is_device(x.arg)), (UPat(Ops.PROGRAM, dtypes.void, src=(UPat((Ops.BUFFER, Ops.PARAM)).or_after(),)), lambda: True), diff --git a/tinygrad/uop/validate.py b/tinygrad/uop/validate.py index 2dfb9025ea..05c8cd3b01 100644 --- a/tinygrad/uop/validate.py +++ b/tinygrad/uop/validate.py @@ -52,6 +52,8 @@ z3_renderer = PatternMatcher([ create_bounded(f"cast{len(ctx[1])}", x.dtype.min, x.dtype.max, ctx[0])), # A comparison between floats introduces a new bool variable (UPat(GroupOp.Comparison, src=UPat(dtype=dtypes.floats)), lambda ctx: (z3.Bool(f"float_cmp{len(ctx[1])}", ctx=ctx[0]), None)), + # a same-dtype cast states a width, which z3 does not model: identity. must precede the rules below (bool->bool) + (UPat(Ops.CAST, name="x"), lambda x,ctx: (ctx[1][x.src[0]], None) if x.dtype == x.src[0].dtype else None), # casts from bool/int to int/bool (UPat(Ops.CAST, dtypes.ints+(dtypes.weakint,),src=(UPat.var("x", dtypes.bool),)), lambda x,ctx: (z3.If(ctx[1][x], 1, 0), None)), (UPat(Ops.CAST, dtypes.ints+(dtypes.weakint,), src=(UPat.var("x", dtypes.ints+(dtypes.weakint,)),)), lambda x,ctx: (ctx[1][x], None)),