From 3cc05081f44bf6ca2e48d06b622ffc0513a9fc19 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Thu, 6 Feb 2025 22:53:49 +0800 Subject: [PATCH] llvm no devectorize, the right way (#8901) * closer * env flag + transcendental issue --- tinygrad/codegen/rewriter.py | 19 +++++++++++++++---- 1 file changed, 15 insertions(+), 4 deletions(-) diff --git a/tinygrad/codegen/rewriter.py b/tinygrad/codegen/rewriter.py index c49bcf2e4d..471249e64d 100644 --- a/tinygrad/codegen/rewriter.py +++ b/tinygrad/codegen/rewriter.py @@ -421,7 +421,8 @@ expander = PatternMatcher([ Ops.VECTORIZE, Ops.IF), name="root", custom_early_reject=set([Ops.UNROLL])), do_expand), (UPat(Ops.CONTRACT, name="con"), do_contract), # vectorize DEFINE_ACC - (UPat(Ops.VECTORIZE, src=UPat(Ops.DEFINE_ACC, name="acc"), name="v"), lambda acc,v: acc.replace(dtype=v.dtype)), + (UPat(Ops.VECTORIZE, src=UPat(Ops.DEFINE_ACC, name="acc"), name="v"), + lambda acc,v: acc.replace(dtype=v.dtype, src=(acc.src[0].broadcast(v.dtype.count),)+acc.src[1:])), # BARRIERs aren't actually expanded (UPat(Ops.BARRIER, src=(UPat(Ops.UNROLL, name="ex"),)), lambda ex: UOp(Ops.UNROLL, dtypes.void, (UOp(Ops.BARRIER, dtypes.void, ex.src),)*len(ex.src), ex.arg)), @@ -453,6 +454,12 @@ devectorize = PatternMatcher([ (UPat((Ops.LOAD, Ops.STORE), name="ls"), no_vectorized_load_store), ]) +devectorize_load_store = PatternMatcher([ + # TODO: add vectorized support to transcendental + (UPat((Ops.INDEX, Ops.EXP2, Ops.LOG2, Ops.SIN), name="alu"), no_vectorized_alu), + (UPat((Ops.LOAD, Ops.STORE), name="ls"), no_vectorized_load_store), +]) + def delete_redundant_gates(buf:UOp, idx:UOp, val:UOp, store_gate:UOp, cast:UOp|None=None) -> UOp|None: if store_gate not in [gate.src[0] for gate in val.toposort if gate.op is Ops.IF]: return None # remove the gate from the index @@ -508,9 +515,13 @@ def full_graph_rewrite(sink:UOp, opts:Optional[Renderer]=None) -> UOp: # expand sink = graph_rewrite(sink, sym+expander) - # devectorize + load_store_indexing + mulacc_unrolled, mulacc_unrolled must be last because it can break loop_collapse - sink = graph_rewrite(sink, sym+(devectorize+float4_folding if opts is not None and opts.supports_float4 else devectorize)+load_store_indexing+ - mulacc_unrolled) + if getenv("NO_DEVECTORIZE"): + # new devectorize for load/store + sink = graph_rewrite(sink, sym+devectorize_load_store) + else: + # devectorize + load_store_indexing + mulacc_unrolled, mulacc_unrolled must be last because it can break loop_collapse + sink = graph_rewrite(sink, sym+(devectorize+float4_folding if opts is not None and opts.supports_float4 else devectorize)+load_store_indexing+ + mulacc_unrolled) # final rules for the renderer (without sym) sink = graph_rewrite(sink, symbolic_simple+get_late_rewrite_patterns(supported_ops, TRANSCENDENTAL>=2)+pm_render+extra_matcher)