From cee9fc7540b26cdd715bb67075e8bfbd9256c4f0 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Thu, 27 Mar 2025 13:16:30 +0800 Subject: [PATCH] new_ignore --- extra/replay_pkl.py | 4 ++-- tinygrad/codegen/devectorizer.py | 8 +++++--- tinygrad/codegen/expander.py | 34 +++++++++++++++++++++++++++++++- tinygrad/codegen/linearize.py | 26 +++++++++++++----------- tinygrad/codegen/lowerer.py | 21 ++++++++++---------- tinygrad/viz/serve.py | 2 +- 6 files changed, 67 insertions(+), 28 deletions(-) diff --git a/extra/replay_pkl.py b/extra/replay_pkl.py index 9177df4aa5..29472f4d5f 100644 --- a/extra/replay_pkl.py +++ b/extra/replay_pkl.py @@ -34,13 +34,13 @@ if __name__ == "__main__": k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) elif k.full_shape[-3:] == (32,3,3): - #if k.full_shape[-4]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, len(k.full_shape)-4, 4)) + if k.full_shape[-4]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, len(k.full_shape)-4, 4)) # 3x3 dwconv k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-3, 32)) if k.full_shape[-4]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 4)) - #elif k.full_shape[-4] == 7: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 7)) + #elif k.full_shap[-4] == 7: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 7)) #elif k.full_shape[-4] == 14: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 2)) elif len(k.full_shape) == 3 and k.full_shape[1] == 32: #if k.full_shape[0]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, 0, 4)) diff --git a/tinygrad/codegen/devectorizer.py b/tinygrad/codegen/devectorizer.py index 9b7b05ba0a..0099fc9b62 100644 --- a/tinygrad/codegen/devectorizer.py +++ b/tinygrad/codegen/devectorizer.py @@ -13,6 +13,8 @@ from tinygrad.renderer import Renderer def expand_index(buf:UOp, vec:UOp, mask:UOp|None=None): if getenv("UNSAFE_DISABLE_MASK", 0): mask = None + if buf.arg == 0: mask = None + # generate the individual indexes midx = graph_rewrite(UOp.sink(*[buf.index(vec.gep(i), mask.gep(i) if mask is not None else None) for i in range(vec.dtype.count)]), symbolic_flat+commutative+load_store_indexing, name=f"index_buf_{buf.arg}") @@ -163,7 +165,7 @@ def split_load_store(ctx:Renderer|None, ls:UOp, idx:UOp): must_divide = True if ctx is not None and ctx.device == "DSP": lengths = [128,64,32,16,8,4] - if ls.dtype.count in [192, 288, 160, 96, 544]: return None # leave these as loads + #if ls.dtype.count in [192, 288, 160, 96, 544]: return None # leave these as loads must_divide = False elif buf.dtype.base != dtypes.float and buf.dtype.base != dtypes.half and not isinstance(buf.dtype, ImageDType): pass @@ -276,8 +278,8 @@ pm_render = PatternMatcher([ # give any loads that are masked an alt value (UPat(Ops.LOAD, src=(UPat(Ops.INDEX, src=(UPat(), UPat(), UPat())).or_casted(),), name="x"), lambda x: x.replace(src=x.src+(x.const_like(0),))), # gate any stores that aren't gated with ifs - (UPat(Ops.STORE, dtype=dtypes.void, src=(UPat(src=(UPat(), UPat(), UPat(dtype=dtypes.bool)), name="idx").or_casted(), UPat()), name="store"), - lambda store,idx: UOp(Ops.STORE, src=store.src+(UOp(Ops.IF, src=(idx.src[2],)),))), + #(UPat(Ops.STORE, dtype=dtypes.void, src=(UPat(src=(UPat(), UPat(), UPat(dtype=dtypes.bool)), name="idx").or_casted(), UPat()), name="store"), + # lambda store,idx: UOp(Ops.STORE, src=store.src+(UOp(Ops.IF, src=(idx.src[2],)),))), ]) # *** uop graph *** diff --git a/tinygrad/codegen/expander.py b/tinygrad/codegen/expander.py index 8585a7e92e..23cc4a0ebb 100644 --- a/tinygrad/codegen/expander.py +++ b/tinygrad/codegen/expander.py @@ -116,9 +116,41 @@ migrate_indexing = PatternMatcher([ (UPat(Ops.STORE, name="root"), create_gate), ]) +pm_store_ignore = PatternMatcher([ + (UPat().index(UPat(), UPat(name="mask")).store(UPat()).named("store"), + lambda store,mask: store.replace(src=(store.src[0], UOp(Ops.IGNORE, src=(store.src[1], mask)))) if store.src[1].op is not Ops.IGNORE else None), +]) + +pm_move_ignore = PatternMatcher([ + # IGNORE on SELF is nothing + (UPat(Ops.IGNORE, src=(UPat(name="x"), UPat(name="x"))), lambda x: x.const_like(True)), + # IGNORE on a CONST is nothing + (UPat(Ops.IGNORE, src=(UPat((Ops.CONST, Ops.VCONST), name="c"), UPat())), lambda c: c), + # move the IGNOREs + (UPat(Ops.IGNORE, src=(UPat((*GroupOp.ALU, Ops.CAST, Ops.VECTORIZE), name="alu"), UPat.var("mask")), name="ig"), + lambda ig,alu,mask: alu.replace(src=tuple(UOp(Ops.IGNORE, x.dtype, (x, mask)) for x in alu.src))), +]) + +pm_delete_ignore = PatternMatcher([ + # IGNORE on SELF is nothing + (UPat(Ops.IGNORE, src=(UPat(name="x"), UPat())), lambda x: x), +]) + def expand_rewrite(sink:UOp) -> UOp: # initial symbolic + migrate indexing (remove this) sink = graph_rewrite(sink, sym+migrate_indexing) + # store IGNORE + sink = graph_rewrite(sink, pm_store_ignore, name="store_ignore") + + # move IGNORE + sink = graph_rewrite(sink, pm_move_ignore, name="move_ignore") + + # remove surviving ignores + sink = graph_rewrite(sink, sym+pm_delete_ignore, name="ignore_done") + # expand - return graph_rewrite(sink, sym+expander) + sink = graph_rewrite(sink, sym+expander) + + return sink + diff --git a/tinygrad/codegen/linearize.py b/tinygrad/codegen/linearize.py index 413b391b2e..e56375efec 100644 --- a/tinygrad/codegen/linearize.py +++ b/tinygrad/codegen/linearize.py @@ -34,23 +34,27 @@ def append_to_block(ctx:tuple[dict[UOp, tuple[UOp, ...]], dict[UOp, list[UOp]]], old_blocks: dict[tuple[UOp, ...], UOp] = {} new_blocks: dict[tuple[UOp, ...], list[UOp]] = {} + seen_u = set() for u in x.src: if u.op is Ops.BLOCK: - # merge sibling blocks. NOTE: blocks must only have one output source - assert u.arg.ctx not in old_blocks, "sibling should never have been created" - old_blocks[u.arg.ctx] = u + if u not in seen_u: + # merge sibling blocks. NOTE: blocks must only have one output source + assert u.arg.ctx not in old_blocks, "sibling should never have been created" + old_blocks[u.arg.ctx] = u elif u.op not in DONT_PLACE_IN_BLOCK and set(children[u]).issubset(in_this_block): - # if it can go in blocks and all its children are in the block, we add it to the block - if (block_ctx:=block_ctxs[u]) == x.arg.ctx: - # if it's the same context, we place the UOp in this block and append the parents to its srcs - new_srcs.extend(u.src) - to_append.append(u) - else: - # if it's a different context, we create a new block with this UOp - new_blocks.setdefault(block_ctx, []).append(u) + if u not in seen_u: + # if it can go in blocks and all its children are in the block, we add it to the block + if (block_ctx:=block_ctxs[u]) == x.arg.ctx: + # if it's the same context, we place the UOp in this block and append the parents to its srcs + new_srcs.extend(u.src) + to_append.append(u) + else: + # if it's a different context, we create a new block with this UOp + new_blocks.setdefault(block_ctx, []).append(u) else: # otherwise, we keep it in the srcs new_srcs.append(u) + seen_u.add(u) if len(to_append) == 0 and len(new_blocks) == 0: return None for rng,lst in new_blocks.items(): diff --git a/tinygrad/codegen/lowerer.py b/tinygrad/codegen/lowerer.py index c0ec02d4cd..c7df4882e4 100644 --- a/tinygrad/codegen/lowerer.py +++ b/tinygrad/codegen/lowerer.py @@ -187,7 +187,8 @@ pm_lowerer = PatternMatcher([ def view_to_mask(x:UOp): from tinygrad.shape.shapetracker import ShapeTracker, View - st = cast(ShapeTracker, x.st) + st = cast(ShapeTracker, x.st).simplify() + print("view_to_mask", st.views) if len(st.views) > 1: return None if st.views[-1].mask is None: return None return ShapeTracker((View(st.shape, (0,)*len(st.shape), 0, st.views[-1].mask, False),)) @@ -241,15 +242,15 @@ pm_quant = symbolic+PatternMatcher([ ((UPat.var('x') * UPat.var('v1').where(UPat(Ops.CONST, arg=1), UPat(Ops.CONST, arg=0)) * UPat.var('v2').where(UPat(Ops.CONST, arg=1), UPat(Ops.CONST, arg=0))).named("mul"), lambda x, mul, v1, v2: x * (v1&v2).where(UOp.const(mul.dtype, 1), UOp.const(mul.dtype, 0))), - # don't care - (UPat(Ops.STORE, name="x"), lambda x: - x.replace(src=(x.src[0], UOp(Ops.IGNORE, src=(x.src[1],), arg=mm), UOp(Ops.IGNORE, x.src[2].dtype, src=(x.src[2],), arg=mm),)) \ - if x.src[1].op is not Ops.IGNORE and (mm:=view_to_mask(x.src[1])) is not None else None), - (UPat(Ops.IGNORE, src=(UPat((*GroupOp.ALU, Ops.CAST), name="alu"),), name="ig"), - lambda ig,alu: alu.replace(src=tuple(UOp(Ops.IGNORE, x.dtype, (x,), ig.arg) for x in alu.src))), - (UPat(Ops.IGNORE, src=(UPat.cvar("c"),), name="ig"), lambda ig, c: c), - (UPat(Ops.IGNORE, src=(UPat(Ops.VALID, name="v"),), name="ig"), lambda ig, v: UOp.const(dtypes.bool, True) if v.src[0].arg == ig.arg else None), - (UPat(Ops.IGNORE, src=(UPat(Ops.REDUCE_AXIS, name="r"),), name="ig"), ignore_on_reduce), + # don't care (moved from here) + #(UPat(Ops.STORE, name="x"), lambda x: + # x.replace(src=(x.src[0], UOp(Ops.IGNORE, src=(x.src[1],), arg=mm), UOp(Ops.IGNORE, x.src[2].dtype, src=(x.src[2],), arg=mm),)) \ + # if x.src[1].op is not Ops.IGNORE and (mm:=view_to_mask(x.src[1])) is not None else None), + #(UPat(Ops.IGNORE, src=(UPat((*GroupOp.ALU, Ops.CAST), name="alu"),), name="ig"), + # lambda ig,alu: alu.replace(src=tuple(UOp(Ops.IGNORE, x.dtype, (x,), ig.arg) for x in alu.src))), + #(UPat(Ops.IGNORE, src=(UPat.cvar("c"),), name="ig"), lambda ig, c: c), + #(UPat(Ops.IGNORE, src=(UPat(Ops.VALID, name="v"),), name="ig"), lambda ig, v: UOp.const(dtypes.bool, True) if v.src[0].arg == ig.arg else None), + #(UPat(Ops.IGNORE, src=(UPat(Ops.REDUCE_AXIS, name="r"),), name="ig"), ignore_on_reduce), # put add in REDUCE #(UPat(Ops.REDUCE_AXIS, name="r")+UPat.var("x"), lambda r,x: r.replace(src=(r.src[0], (r.src[1]+x) if len(r.src) == 2 else x))), # distribute on casted MUL diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index 5e09abccb6..025fc924f3 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -12,7 +12,7 @@ from tinygrad.dtype import dtypes uops_colors = {Ops.LOAD: "#ffc0c0", Ops.STORE: "#87CEEB", Ops.CONST: "#e0e0e0", Ops.VCONST: "#e0e0e0", Ops.REDUCE: "#FF5B5B", Ops.DEFINE_GLOBAL: "#ffe0b0", Ops.DEFINE_LOCAL: "#ffe0d0", Ops.DEFINE_ACC: "#f0ffe0", Ops.REDUCE_AXIS: "#FF6B6B", Ops.RANGE: "#c8a0e0", Ops.ASSIGN: "#e0ffc0", Ops.BARRIER: "#ff8080", Ops.IF: "#c8b0c0", Ops.SPECIAL: "#c0c0ff", - Ops.INDEX: "#e8ffa0", Ops.WMMA: "#efefc0", Ops.VIEW: "#C8F9D4", Ops.MULTI: "#f6ccff", Ops.KERNEL: "#3e7f55", + Ops.INDEX: "#e8ffa0", Ops.WMMA: "#efefc0", Ops.VIEW: "#C8F9D4", Ops.MULTI: "#f6ccff", Ops.KERNEL: "#3e7f55", Ops.IGNORE: "#00C000", **{x:"#D8F9E4" for x in GroupOp.Movement}, **{x:"#ffffc0" for x in GroupOp.ALU}, Ops.THREEFRY:"#ffff80", Ops.BUFFER_VIEW: "#E5EAFF", Ops.BLOCK: "#C4A484", Ops.BLOCKEND: "#C4A4A4", Ops.BUFFER: "#B0BDFF", Ops.COPY: "#a040a0", Ops.NAME:"#808080"}