From 22ab17369d75985da7baee86f73e66deee03b189 Mon Sep 17 00:00:00 2001 From: wrongtest Date: Sun, 13 Feb 2022 00:13:31 +0800 Subject: [PATCH] add a simplify rule for floordiv(x*8+7, 16) => floordiv(x, 2) --- src/arith/rewrite_simplify.cc | 14 ++++++ tests/python/unittest/test_arith_intset.py | 10 +++-- .../unittest/test_arith_rewrite_simplify.py | 4 ++ .../unittest/test_tir_schedule_compute_at.py | 45 ++++++++++++++++--- 4 files changed, 64 insertions(+), 9 deletions(-) diff --git a/src/arith/rewrite_simplify.cc b/src/arith/rewrite_simplify.cc index 4a99e10211b7..b7acff1e59d5 100644 --- a/src/arith/rewrite_simplify.cc +++ b/src/arith/rewrite_simplify.cc @@ -762,6 +762,20 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const FloorDivNode* op) { if (c2val % c1val == 0) return floordiv(x, floordiv(c2, c1)).Eval(); } } + if (floordiv(x * c1 + c2, c3).Match(ret)) { + int64_t c1val = c1.Eval()->value; + int64_t c2val = c2.Eval()->value; + int64_t c3val = c3.Eval()->value; + if (c1val > 0 && c3val > 0 && c3val % c1val == 0 && floormod(c2val, c3val) < c1val) { + // assume c3 == a * c1, x == a * y + b, c2 = d * c3 + e then + // (x * c1 + c2) // c3 + // ==> ((a * y + b) * c1 + d * a * c1 + e) // (a * c1) + // ==> y + d + (b * c1 + e) // c3 + // ==> y + d since 0 <= b * c1 <= (a-1) * c1, 0 <= e < c1 + // ==> x // (c3 // c1) + (c2 // c3) + return (floordiv(x, floordiv(c3, c1)) + floordiv(c2, c3)).Eval(); + } + } TVM_TRY_REWRITE(floordiv(x, x), OneWithTypeLike(x)); TVM_TRY_REWRITE(floordiv(x * c1, x), c1); diff --git a/tests/python/unittest/test_arith_intset.py b/tests/python/unittest/test_arith_intset.py index b40f3c9f56ea..e6738543a6aa 100644 --- a/tests/python/unittest/test_arith_intset.py +++ b/tests/python/unittest/test_arith_intset.py @@ -31,7 +31,7 @@ def err_msg(): return "\ndata={}\ndmap={}\nres={}\nexpected={}".format(data, dmap, res, expected) def equal(x, y): - res = self.analyzer.canonical_simplify(x - y) + res = self.analyzer.simplify(x - y) return tvm.tir.analysis.expr_deep_equal(res, 0) assert equal(res.min_value, expected[0]), err_msg() @@ -99,10 +99,14 @@ def test_mod(): ck.verify(flm(x, 10), {x: tvm.arith.IntervalSet(3, 11)}, (0, 9)) ck.verify(flm(x, 10), {x: tvm.arith.IntervalSet(1, 21)}, (0, 9)) - floordiv = tvm.te.floordiv + fld = tvm.te.floordiv z = te.var("z") ck.analyzer.bind(x, tvm.ir.Range.from_min_extent(0, 3)) - ck.verify(flm(y, 8), {y: tvm.arith.IntervalSet(z * 8 + x * 4, z * 8 + x * 4 + 3)}, (0, 7)) + ck.verify( + flm(y, 8), + {y: tvm.arith.IntervalSet(z * 8 + x * 4, z * 8 + x * 4 + 3)}, + (x * 4 - 8 * fld(x * 4, 8), x * 4 - 8 * fld(x * 4, 8) + 3), + ) ck1 = IntSetChecker() ck1.analyzer.bind(x, tvm.ir.Range.from_min_extent(0, 2)) ck1.verify( diff --git a/tests/python/unittest/test_arith_rewrite_simplify.py b/tests/python/unittest/test_arith_rewrite_simplify.py index 6ca2a2a5fcb0..662038e129f7 100644 --- a/tests/python/unittest/test_arith_rewrite_simplify.py +++ b/tests/python/unittest/test_arith_rewrite_simplify.py @@ -462,6 +462,10 @@ def test_floordiv_index_simplify(): ck.verify(fld(x * 2, 4), fld(x, 2)) ck.verify(fld(x * 4, 2), x * 2) + ck.verify(fld(x * 8 + 7, 16), fld(x, 2)) + ck.verify(fld(x * 8 + 39, 16), fld(x, 2) + 2) + ck.verify(fld(x * 8 - 1, 16), fld(x * 8 + -1, 16)) + ck.verify(fld(x * 8 - 9, 16), fld(x, 2) + -1) ck.verify(fld(x * 4 + y, 2), x * 2 + fld(y, 2)) ck.verify(fld(tvm.te.min(x * 6, y), 2), tvm.te.min(x * 3, fld(y, 2))) diff --git a/tests/python/unittest/test_tir_schedule_compute_at.py b/tests/python/unittest/test_tir_schedule_compute_at.py index dc0812780503..4d081b507403 100644 --- a/tests/python/unittest/test_tir_schedule_compute_at.py +++ b/tests/python/unittest/test_tir_schedule_compute_at.py @@ -912,7 +912,7 @@ def concat_two_elemwise(x: T.Buffer[(16,), "float32"], for i in T.serial(24): with T.block("T_concat"): ax = T.axis.spatial(24, i) - T_concat[ax] = T.if_then_else(16 <= ax, T_add_1[ax - 16], T_add_2[ax], dtype="float32") + T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax], dtype="float32") @T.prim_func def concat_two_elemwise_after_compute_at(x: T.Buffer[(16,), "float32"], @@ -922,16 +922,16 @@ def concat_two_elemwise_after_compute_at(x: T.Buffer[(16,), "float32"], T_add_2 = T.alloc_buffer([8], dtype="float32") for i in T.serial(24): with T.block("T_add_1"): - ax = T.axis.spatial(16, i - 16) - T.where(16 <= i) + ax = T.axis.spatial(16, i) + T.where(i < 16) T_add_1[ax] = x[ax] + T.float32(1) with T.block("T_add_2"): - ax = T.axis.spatial(8, i) - T.where(i < 8) + ax = T.axis.spatial(8, i - 16) + T.where(16 <= i) T_add_2[ax] = y[ax] + T.float32(2) with T.block("T_concat"): ax = T.axis.spatial(24, i) - T_concat[ax] = T.if_then_else(16 <= ax, T_add_1[ax - 16], T_add_2[ax], dtype="float32") + T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax], dtype="float32") @T.prim_func def floordiv_and_floormod_indices(a: T.handle, b: T.handle) -> None: @@ -962,6 +962,31 @@ def floordiv_and_floormod_indices_after_reverse_compute_at(a: T.handle, b: T.han v_i = T.axis.spatial(256, i * 16 + ax0) Y[v_i] = temp[v_i // 16, v_i % 16] + +@T.prim_func +def tiled_repeat_op(x: T.Buffer[(4,), "float32"], T_repeat: T.Buffer[(64,), "float32"]) -> None: + T_add = T.alloc_buffer([4], dtype="float32") + for i0 in T.serial(4): + with T.block("T_add"): + ax0 = T.axis.spatial(4, i0) + T_add[ax0] = x[ax0] + 1.0 + for i0_0, i0_1 in T.grid(8, 8): + with T.block("T_repeat"): + ax0 = T.axis.spatial(64, i0_0 * 8 + i0_1) + T_repeat[ax0] = T_add[ax0 // 16] + +@T.prim_func +def tiled_repeat_op_after_compute_at(x: T.Buffer[(4,), "float32"], T_repeat: T.Buffer[(64,), "float32"]) -> None: + T_add = T.alloc_buffer([4], dtype="float32") + for i0_0 in T.serial(8): + with T.block("T_add"): + ax0 = T.axis.spatial(4, i0_0 // 2) + T_add[ax0] = x[ax0] + T.float32(1) + for i0_1 in T.serial(8): + with T.block("T_repeat"): + ax0 = T.axis.spatial(64, i0_0 * 8 + i0_1) + T_repeat[ax0] = T_add[ax0 // 16] + # pylint: enable=no-member,invalid-name,unused-variable,line-too-long,redefined-outer-name,unexpected-keyword-arg,too-many-nested-blocks # fmt: on @@ -1077,6 +1102,14 @@ def test_compute_at_concat(): verify_trace_roundtrip(sch=sch, mod=concat_two_elemwise) +def test_compute_at_tiled_repeat_op(): + sch = tir.Schedule(tiled_repeat_op, debug_mask="all") + outer_ax, _ = sch.get_loops(sch.get_block("T_repeat")) + sch.compute_at(sch.get_block("T_add"), outer_ax) + tvm.ir.assert_structural_equal(tiled_repeat_op_after_compute_at, sch.mod["main"]) + verify_trace_roundtrip(sch=sch, mod=tiled_repeat_op) + + def test_reverse_compute_at_tiled(): sch = tir.Schedule(tiled, debug_mask="all") block = sch.get_block("C")