diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 3d3722de92..da79a9aae2 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -492,8 +492,7 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]: tsink, rctx = run_rangeify(tsink, getenv("DEBUG_RANGEIFY", 0)) # NOTE: sym (vs symbolic_simple) breaks things here because ranges with len 1 aren't handled right - tsink = graph_rewrite(tsink, symbolic_simple+pm_reduce_unparented, name="symbolic") # this supports const folding - tsink = graph_rewrite(tsink, pm_cleanups, bottom_up=True, name="remove costly buffers") + tsink = graph_rewrite(tsink, symbolic_simple+pm_reduce_unparented+pm_cleanups, bottom_up=True, name="remove costly buffers") # TODO: can you substitute and remove costly buffers at the same time? tsink = graph_rewrite(tsink, pm_substitute_recurse, bottom_up=True, name="run substitutes") tsink = graph_rewrite(tsink, pm_limit_bufs, ctx=rctx, name="limit buffers") diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index cc25ab6d8a..529f2a5161 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -86,13 +86,16 @@ symbolic_simple = propagate_invalid + PatternMatcher([ (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), # *** div rules *** + (UPat.cvar('x', arg=0) / 0, lambda x: x.const_like(float('nan'))), # 0/0 -> nan + ((UPat.var("x") * 0) / 0, lambda x: x.const_like(float('nan'))), # (x*0)/0 -> nan # can be wrong if x or x2 is 0 - (UPat.var("x") / UPat.var("x"), lambda x: x.const_like(1)), # x/x -> 1 + (UPat.var("x") / UPat.var("x"), lambda x: x.const_like(1)), # x/x -> 1 ((UPat.var("x") * UPat.var("x2")) / UPat.var("x2"), lambda x,x2: x), # (x*x2)/x2 -> x # 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 - (UPat.var("x") * 0, lambda x: x.const_like(float("nan") if isinstance(x.arg, float) and (math.isnan(x.arg) or math.isinf(x.arg)) else 0)), + (UPat.var("x") * 0, lambda x: x.const_like(float("nan") if x.op is Ops.CONST + and isinstance(x.arg, float) and (math.isnan(x.arg) or math.isinf(x.arg)) else 0)), # *** cast/bitcast *** (UPat(Ops.CAST, name="root", src=(UPat.cvar("c"),)), lambda root, c: root.const_like(c.arg)), (UPat((Ops.CAST, Ops.BITCAST), name="root"), lambda root: root.src[0] if root.dtype == root.src[0].dtype else None),