diff --git a/tinygrad/codegen/simplify.py b/tinygrad/codegen/simplify.py index 495841e14a..5d315be081 100644 --- a/tinygrad/codegen/simplify.py +++ b/tinygrad/codegen/simplify.py @@ -98,17 +98,14 @@ pm_reduce_collapse = pm_reduce_unparented + PatternMatcher([ # lift x*y out of reduce ((UPat.var("x")*UPat.var("y")) < UPat.var("c"), lambda x,y,c: (x < ((c+y-1) // y)) if no_range(y) and no_range(c) and dtypes.is_int(y.dtype) and y.vmin > 0 else None), - # fold the range - # bound from below - ((UPat(Ops.RANGE, name="r") < UPat.var("cut")).where(0, UPat.var("val")).reduce(UPat.var("r"), arg=Ops.ADD), - lambda r,cut,val: (r.src[0]-cut).maximum(0).minimum(r.src[0]).cast(val.dtype) * val if no_range(val) else None), - # bound from two sides - (((UPat.var("r") clamp(min(upper,N) - max(lower,0), 0, N) * val + (UPat.any( + (UPat(Ops.RANGE, name="r") < UPat.var("upper")).where(UPat.var("val"), 0), + (UPat(Ops.RANGE, name="r") < UPat.var("lower")).where(0, UPat.var("val")), + ((UPat.var("r")