diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 0e5681c6d6..3f147d74ba 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -17,7 +17,7 @@ from tinygrad.codegen.late.expander import expander, pm_pre_expander, pm_group_f from tinygrad.codegen.late.devectorizer import load_store_folding, load_store_indexing, devectorize, pm_reduce, \ ReduceContext, correct_load_store, pm_render, pm_add_loads, pm_make_images from tinygrad.codegen.opt.postrange import apply_opts -from tinygrad.codegen.late.gater import pm_move_gates_from_index +from tinygrad.codegen.late.gater import pm_image_index, pm_move_gates_from_index from tinygrad.codegen.simplify import pm_simplify_ranges, pm_flatten_range, pm_split_ranges, pm_load_collapse from tinygrad.schedule.rangeify import pm_add_buffers_local, rangeify_codegen, pm_mops, pm_syntactic_sugar, pm_store_ranges from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize @@ -77,6 +77,9 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp: else: pm_devectorize = sym+load_store_folding+correct_load_store+load_store_indexing if DEVECTORIZE >= 0: sink = graph_rewrite(sink, pm_devectorize, ctx=ren, name="devectorize") + # convert image linear offsets to image coordinates before symbolic/index dtype cleanup + sink = graph_rewrite(sink, pm_image_index, name="image indexing") + # 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") diff --git a/tinygrad/codegen/late/devectorizer.py b/tinygrad/codegen/late/devectorizer.py index 4d079ef992..8c494f4f5e 100644 --- a/tinygrad/codegen/late/devectorizer.py +++ b/tinygrad/codegen/late/devectorizer.py @@ -207,28 +207,9 @@ def split_load_store(ctx:Renderer|None, ls:UOp, idx:UOp): if len(ret) <= 1: return None return UOp(Ops.VCAT, ls.dtype, tuple(ret)) if ls.op is Ops.LOAD else UOp.group(*ret) -def get_image_idx(idx:UOp, width:int): - x, valid = idx.src[1].get_idx(), idx.src[1].get_valid() - idx_x, idx_y = (x // 4) % width, x // (4*width) - return idx.replace(src=(idx.src[0], idx_x.valid(valid), idx_y.valid(valid))) - -def image_fixup(ls:UOp): - # normal image load or store, with the CAST from expand_index - if isinstance(dt:=ls.src[0].src[0].dtype, ImageDType) and ls.src[0].op is Ops.CAST: - assert ls.src[0].dtype.count == 4, "image must be casted to 4" - return ls.replace(src=(get_image_idx(ls.src[0].src[0], dt.shape[1]),)+ls.src[1:]) - - # this is an unprocessed image without a cast, we should just make it a buffer - if isinstance(dt, ImageDType) and len(ls.src[0].src) != 3: - off = ls.src[0].src[1] - idx = ls.src[0].src[0].replace(dtype=(new_dt:=dtypes.half if dt.itemsize == 2 else dtypes.float).ptr(dt.size)).index(off) - return ls.replace(src=(idx,), dtype=new_dt).cast(dtypes.float) if ls.op is Ops.LOAD else ls.replace(src=(idx, ls.src[1].cast(new_dt))) - correct_load_store = PatternMatcher([ # split LOAD/STORE (UPat((Ops.LOAD, Ops.STORE), src=(UPat(Ops.INDEX, name="idx").cast(),), name="ls", allow_any_len=True), split_load_store), - # image indexing, including unfoldable images - (UPat((Ops.LOAD, Ops.STORE), name="ls"), image_fixup), ]) # *** uop expander *** diff --git a/tinygrad/codegen/late/gater.py b/tinygrad/codegen/late/gater.py index 1bab753cb8..2321b973c4 100644 --- a/tinygrad/codegen/late/gater.py +++ b/tinygrad/codegen/late/gater.py @@ -14,6 +14,37 @@ def image_coords_to_int(idx:UOp, buf:UOp, x:UOp, y:UOp): if not isinstance(buf.dtype, ImageDType) or (x.dtype != dtypes.long and y.dtype != dtypes.long): return None return idx.replace(src=(buf, x.cast(dtypes.int) if x.dtype == dtypes.long else x, y.cast(dtypes.int) if y.dtype == dtypes.long else y)) +def index_and_valid(idx:UOp) -> tuple[UOp, UOp]: + if idx.dtype.scalar() is dtypes.weakint: return idx.get_idx(), idx.get_valid() + if idx.op is Ops.WHERE and idx.src[2].arg is Invalid: return idx.src[1], idx.src[0] + return idx, UOp.const(dtypes.bool, idx.arg is not Invalid) + +def valid_idx(idx:UOp, valid:UOp) -> UOp: + return idx if valid.op is Ops.CONST and valid.arg is True else valid.where(idx, idx.const_like(Invalid)) + +def get_image_idx(idx:UOp, width:int) -> UOp: + x, valid = index_and_valid(idx.src[1]) + idx_x, idx_y = (x.gep(0), x.gep(1)) if x.dtype.count == 2 else ((x // 4) % width, x // (4*width)) + return idx.replace(src=(idx.src[0], valid_idx(idx_x, valid), valid_idx(idx_y, valid))) + +def image_fixup(ls:UOp): + # normal image load/store from split_load_store: casted linear offset -> image x/y coordinates + if ls.src[0].op is Ops.CAST and (cast_idx:=ls.src[0].src[0]).op is Ops.INDEX and isinstance(dt:=cast_idx.src[0].dtype, ImageDType): + assert ls.src[0].dtype.count == 4, "image must be casted to 4" + return ls.replace(src=(cast_idx if len(cast_idx.src) == 3 else get_image_idx(cast_idx, dt.shape[1]),)+ls.src[1:]) + + if ls.src[0].op is not Ops.INDEX or not isinstance(dt:=ls.src[0].src[0].dtype, ImageDType) or len(ls.src[0].src) == 3: return None + off, _ = index_and_valid(ls.src[0].src[1]) + if off.dtype.count == 2: return ls.replace(src=(get_image_idx(ls.src[0], dt.shape[1]),)+ls.src[1:]) + + # this is an unprocessed image without a cast, we should just make it a buffer + idx = ls.src[0].src[0].replace(dtype=(new_dt:=dtypes.half if dt.itemsize == 2 else dtypes.float).ptr(dt.size)).index(ls.src[0].src[1]) + return ls.replace(src=(idx,), dtype=new_dt).cast(dtypes.float) if ls.op is Ops.LOAD else ls.replace(src=(idx, ls.src[1].cast(new_dt))) + +pm_image_index = PatternMatcher([ + (UPat((Ops.LOAD, Ops.STORE), name="ls"), image_fixup), +]) + pm_move_gates_from_index = PatternMatcher([ # here we create the alt value for load to be 0s and remove the where Invalid (UPat.var("buf").index(UPat.var("gate").where(UPat.var("idx"), UPat(arg=Invalid))).or_casted(name="cast").load(name="l"),