diff --git a/tinygrad/codegen/uopgraph.py b/tinygrad/codegen/uopgraph.py index 4e4c926f48..2ce7b4ee99 100644 --- a/tinygrad/codegen/uopgraph.py +++ b/tinygrad/codegen/uopgraph.py @@ -418,10 +418,6 @@ expander = PatternMatcher([ (UPat(UOps.CONTRACT, name="con"), do_contract), # vectorize DEFINE_ACC (UPat(UOps.VECTORIZE, src=UPat(UOps.DEFINE_ACC, name="acc"), name="v"), lambda acc,v: acc.replace(dtype=v.dtype)), - # remove EXPANDs from SINK - (UPat(UOps.SINK, name="root"), - lambda root: UOp(UOps.SINK, root.dtype, a, root.arg) - if len(a:=tuple(flatten(x.src if x.op is UOps.EXPAND else (x,) for x in root.src))) != len(root.src) else None), # BARRIERs aren't actually expanded (UPat(UOps.BARRIER, src=(UPat(UOps.EXPAND, name="ex"),)), lambda ex: UOp(UOps.EXPAND, dtypes.void, (UOp(UOps.BARRIER, dtypes.void, ex.src),)*len(ex.src), ex.arg)), diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 528f08243f..790254b13c 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -285,7 +285,9 @@ class UOp(MathTrait, metaclass=UOpMetaClass): def __bool__(self): return self._eval((dtypes.bool,), bool) def __int__(self): return self._eval(dtypes.ints, int) def __float__(self): return self._eval(dtypes.floats, float) - def substitute(self, dvars:Dict[UOp, UOp]): return graph_rewrite(self, _substitute, dvars) + def substitute(self, dvars:Dict[UOp, UOp]): + with Context(TRACK_MATCH_STATS=0): + return graph_rewrite(self, _substitute, dvars) # *** uop syntactic sugar *** @@ -554,8 +556,10 @@ class UPat(MathTrait): def any(*src): return UPatAny(src=src) @staticmethod + @functools.lru_cache(None) def var(name:Optional[str]=None, dtype:Optional[Union[DType, Tuple[DType, ...]]]=None): return UPat(dtype=dtype, name=name) @staticmethod + @functools.lru_cache(None) def cvar(name:Optional[str]=None, dtype:Optional[DType]=None, vec=True): return UPat((UOps.CONST, UOps.VCONST) if vec else UOps.CONST, dtype=dtype, name=name) @staticmethod @@ -1023,7 +1027,8 @@ symbolic_simple = PatternMatcher([ (UPat.var("x", dtype=dtypes.bool).logical_not().logical_not(), lambda x: x), # ** zero folding ** (UPat.var("x") < UPat.var("x"), lambda x: UOp.const(dtypes.bool.vec(x.dtype.count), False)), # x < x -> False - (UPat.var("x", dtype=dtypes.ints) != UPat.var("x"), lambda x: UOp.const(dtypes.bool.vec(x.dtype.count), False)), # x != x -> False (only ints) + (UPat.var("x", dtype=dtypes.ints) != UPat.var("x", dtype=dtypes.ints), + lambda x: UOp.const(dtypes.bool.vec(x.dtype.count), False)), # x != x -> False (only ints) # x*0 -> 0 or 0*x -> 0 # if x is nan or inf it should render the nan value. # NOTE: this can be wrong for loaded NaN @@ -1034,9 +1039,9 @@ symbolic_simple = PatternMatcher([ # ** COMMUTATIVE flipping ** *[(UPat(UOps.ALU, arg=op, name='x'), lambda x: x.replace(src=x.src[::-1]) if x.src[1].tuplize < x.src[0].tuplize else None) for op in COMMUTATIVE], # bool MUL is AND, ADD/MAX is OR. prevents other rules to rewrite bool ADD/MUL incorrectly - (UPat.var('x', dtype=dtypes.bool) * UPat.var('y'), lambda x,y: x&y), - (UPat.var('x', dtype=dtypes.bool) + UPat.var('y'), lambda x,y: x|y), - (UPat.var('x', dtype=dtypes.bool).maximum(UPat.var('y')), lambda x,y: x|y), + (UPat.var('x', dtype=dtypes.bool) * UPat.var('y', dtype=dtypes.bool), lambda x,y: x&y), + (UPat.var('x', dtype=dtypes.bool) + UPat.var('y', dtype=dtypes.bool), lambda x,y: x|y), + (UPat.var('x', dtype=dtypes.bool).maximum(UPat.var('y', dtype=dtypes.bool)), lambda x,y: x|y), # *** cast *** (UPat(UOps.CAST, name="root", src=UPat.cvar("c")), lambda root, c: root.const_like(c.arg)), (UPat(UOps.CAST, name="root"), lambda root: root.src[0] if root.dtype == root.src[0].dtype else None),