Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions src/arith/rewrite_simplify.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
10 changes: 7 additions & 3 deletions tests/python/unittest/test_arith_intset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand Down
4 changes: 4 additions & 0 deletions tests/python/unittest/test_arith_rewrite_simplify.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)))
Expand Down
45 changes: 39 additions & 6 deletions tests/python/unittest/test_tir_schedule_compute_at.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
Expand All @@ -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:
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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")
Expand Down