Compare commits

...
Author SHA1 Message Date
George HotzandGitHub 7cb8fdd7a6 Merge branch 'master' into late_lower_idx 2026-06-25 14:49:21 -07:00
geohot f64908b6c5 late lowering for weakint 2026-06-25 14:47:27 -07:00
2 changed files with 6 additions and 9 deletions
+6 -6
View File
@@ -102,12 +102,8 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
sink = graph_rewrite(sink, pm_make_images, name="create image buffers", bottom_up=True, ctx=ren.target.arch)
# devectorize
sink = graph_rewrite(sink, sym+devectorize_alu+devectorize_buf_and_index+load_store_folding+correct_load_store+load_store_indexing,
ctx=ren, name="devectorize")
# lower the index dtype to a concrete int
sink = graph_rewrite(sink, pm_lower_index_dtype+load_store_indexing+gep_pushing, name="lower all index dtypes")
sink = graph_rewrite(sink, symbolic, name="post index symbolic")
sink = graph_rewrite(sink, sym+devectorize_alu+devectorize_buf_and_index+load_store_folding+correct_load_store, ctx=ren, name="devectorize")
sink = graph_rewrite(sink, gep_pushing, name="gep pushing")
# optional pre matcher
if ren.pre_matcher is not None: sink = graph_rewrite(sink, ren.pre_matcher, name="pre_matcher")
@@ -119,6 +115,10 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
# do memory coalesing (late)
sink = memory_coalesing(sink, ren)
# lower the index dtype to a concrete int
sink = graph_rewrite(sink, pm_lower_index_dtype+load_store_indexing, name="lower all index dtypes")
sink = graph_rewrite(sink, symbolic, name="post index symbolic")
# floordiv+mod / dtype decomp (early)
supported_ops = tuple(ren.code_for_op.keys())
pm_decomp = symbolic_simple+get_simplifying_rewrite_patterns(supported_ops)
-3
View File
@@ -320,9 +320,6 @@ symbolic = symbolic_simple+commutative+PatternMatcher([
else y.src for y in x.src[1:]]))))),
# after with 1 src is just src[0]
(UPat(Ops.AFTER, src=(UPat.var("s"),)), lambda s: s),
# VECTORIZE/CONST
(UPat(Ops.STACK, src=UPat(Ops.CONST), name="vec"),
lambda vec: UOp.const(vec.dtype, tuple(x.arg for x in vec.src)) if len(vec.src) > 0 else None),
])+div_and_mod_symbolic+gep_pushing
# ******** we take a small aside to "simplify_valid" to rewrite valids ********