diff --git a/tinygrad/codegen/uopgraph.py b/tinygrad/codegen/uopgraph.py index aabc947569..80e1f62007 100644 --- a/tinygrad/codegen/uopgraph.py +++ b/tinygrad/codegen/uopgraph.py @@ -290,15 +290,12 @@ sym = symbolic_flat+PatternMatcher([ (UPat.store(UPat.var("buf"), UPat.var("idx"), UPat.var("gate").where(UPat.var("alt"), UPat.load(UPat.var("buf"), UPat.var("idx")))), lambda buf, idx, gate, alt: UOp.store(buf, idx, alt, gate)), # fold gated LOAD/STORE - (UPat.load(UPat.var("buf"), UPat.var("idx"), UPat.var("var"), UPat.const(dtypes.bool, True)), - lambda buf,idx,var: UOp.load(buf, idx, dtype=var.dtype)), - (UPat.load(UPat.var("buf"), UPat.var("idx"), UPat.var("var"), UPat.const(dtypes.bool, True), UPat.var("barrier")), - lambda buf,idx,var,barrier: UOp.load(buf, idx, barrier, dtype=var.dtype)), - (UPat.load(UPat.var(), UPat.var(), UPat.var("var"), UPat.const(dtypes.bool, False)), lambda var: var), - (UPat.load(UPat.var(), UPat.var(), UPat.var("var"), UPat.const(dtypes.bool, False), UPat.var()), lambda var: var), - (UPat.store(UPat.var("buf"), UPat.var("idx"), UPat.var("val"), UPat.const(dtypes.bool, True)), - lambda buf,idx,val: UOp.store(buf, idx, val)), # pylint: disable=unnecessary-lambda - (UPat.store(UPat.var(), UPat.var(), UPat.var(), UPat.const(dtypes.bool, False)), lambda: UOp(UOps.NOOP)), + (UPat.load(UPat(), UPat(), UPat(), UPat.const(dtypes.bool, True), name="ld"), lambda ld: ld.replace(src=ld.src[:2])), + (UPat.load(UPat(), UPat(), UPat(), UPat.const(dtypes.bool, True), UPat.var("bar"), name="ld"), lambda ld,bar: ld.replace(src=ld.src[:2]+(bar,))), + (UPat.load(UPat(), UPat(), UPat.var("var"), UPat.const(dtypes.bool, False)), lambda var: var), + (UPat.load(UPat(), UPat(), UPat.var("var"), UPat.const(dtypes.bool, False), UPat()), lambda var: var), + (UPat.store(UPat(), UPat(), UPat(), UPat.const(dtypes.bool, True), name="store"), lambda store: store.replace(src=store.src[:3])), + (UPat.store(UPat(), UPat(), UPat(), UPat.const(dtypes.bool, False)), lambda: UOp(UOps.NOOP)), # remove NOOPs from SINK (UPat(UOps.SINK, name="root"), lambda root: UOp(UOps.SINK, root.dtype, a, root.arg) if len(a:=tuple(x for x in root.src if x.op is not UOps.NOOP)) != len(root.src) else None), diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 6417a29063..bf13455b8f 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -542,9 +542,9 @@ class UPat(MathTrait): def bitcast(self, dtype=None): return UPat(UOps.BITCAST, dtype, (self,)) def gep(self, i:int): return UPat(UOps.GEP, None, (self,), (i,)) @staticmethod - def load(*src:UPat, dtype:Optional[DType]=None): return UPat(UOps.LOAD, dtype, src) + def load(*src:UPat, **kwargs): return UPat(UOps.LOAD, src=src, **kwargs) @staticmethod - def store(*src:UPat): return UPat(UOps.STORE, dtypes.void, src) + def store(*src:UPat, **kwargs): return UPat(UOps.STORE, dtypes.void, src, **kwargs) def const_like(self, b:ConstType|Variable|Tuple[ConstType]): return UPat.const(self.dtype, cast(ConstType, b)) def alu(self, arg, *src:UPat):