From 176377ff6eb98032e65efdb7da04d11d68de0aa9 Mon Sep 17 00:00:00 2001 From: chenyu Date: Fri, 21 Aug 2026 09:33:43 -0400 Subject: [PATCH] weak 1 for FDIV in get_late_rewrite_patterns [PR] (#17655) --- tinygrad/codegen/__init__.py | 8 ++++---- tinygrad/codegen/decomp/op.py | 4 ++-- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 16afcdefdb..0a177904c0 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -377,6 +377,10 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp: pm_final_rewrite = pm_commit_weak+pm_cast_weak+pm_decomp+extra_matcher+pm_split_ends sink = graph_rewrite(sink, pm_final_rewrite+pm_remove_invalid, ctx=ren, name="final rewrite") + # spell every literal as a casted const CAST(dt, CONST(value)) + # TODO: remove once consts are always weak + sink = graph_rewrite(sink, pm_casted_consts, name="casted consts", walk=True) + # add implicit barriers (stores/loads through LOCAL memory ordered by AFTER or across loop iterations need workgroup barriers) sink = graph_rewrite(sink, pm_implicit_barriers, name="add implicit barriers") @@ -387,10 +391,6 @@ 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) - # spell every literal as a casted const CAST(dt, CONST(value)) - # TODO: remove once consts are always weak - 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) diff --git a/tinygrad/codegen/decomp/op.py b/tinygrad/codegen/decomp/op.py index 6a48cdca53..a23142809f 100644 --- a/tinygrad/codegen/decomp/op.py +++ b/tinygrad/codegen/decomp/op.py @@ -128,6 +128,6 @@ def get_late_rewrite_patterns(ops:tuple[Ops, ...], disable_fast_idiv:bool) -> Pa if Ops.SHL in ops: pat += [(UPat.var('x').alu(Ops.SHL, UPat.cvar('n'))+UPat.var('c'), lambda x,n,c: x.alu(Ops.MULACC, x.const_like(1< a/b if Ops.FDIV in ops: - pat += [(UPat.var("x").reciprocal(), lambda x: x.const_like(1).alu(Ops.FDIV, x))] - pat += [(UPat.var("a", dtypes.floats) * UPat(Ops.FDIV, dtypes.floats, src=(UPat.const(1), UPat.var("b"))), lambda a,b: a.alu(Ops.FDIV, b))] + pat += [(UPat.var("x").reciprocal(), lambda x: UOp.const(1.0).alu(Ops.FDIV, x))] + pat += [(UPat.var("a") * UPat(Ops.FDIV, dtypes.floats, src=(UPat.const(1), UPat.var("b"))), lambda a,b: a.alu(Ops.FDIV, b))] return PatternMatcher(pat)