This is an automated email from the ASF dual-hosted git repository.
lunderberg pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new a543f6bebf [Arith] Implemented PMatchesOneOf and matches_one_of
(#13933)
a543f6bebf is described below
commit a543f6bebfc546b56c37b68c4ef6484535be65b9
Author: Eric Lunderberg <[email protected]>
AuthorDate: Fri Feb 17 15:23:42 2023 -0600
[Arith] Implemented PMatchesOneOf and matches_one_of (#13933)
* [Arith] Implemented PMatchesOneOf and matches_one_of
Frequently, several related rewrite rules will have identical output
patterns and preconditions, differing only by the form of the input
pattern. For example, both `(x + y) - y` and `(y + x) - y` can be
simplified to `x`. Previously, these would be written as separate
rewrite rules, which required duplicating the output and
preconditions.
This commit introduces `PMatchesOneOf`, which attempts a series of
pattern matches, stopping after the first successful match. The
`PMatchesOneOf` instance can be used as the input pattern for rewrite
rules. A helper function `matches_one_of` is also added, which
returns an instance of `PMatchesOneOf`.
```c++
// Before
TVM_TRY_REWRITE( (x + y) - y, x);
TVM_TRY_REWRITE( (y + x) - y, x);
// After
TVM_TRY_REWRITE( matches_one_of( (x + y) - y,
(y + x) - y),
x);
```
---
src/arith/pattern_match.h | 106 +++++++-
src/arith/rewrite_simplify.cc | 546 +++++++++++++++++++++---------------------
2 files changed, 377 insertions(+), 275 deletions(-)
diff --git a/src/arith/pattern_match.h b/src/arith/pattern_match.h
index 69f064e119..55b51d7a31 100644
--- a/src/arith/pattern_match.h
+++ b/src/arith/pattern_match.h
@@ -104,18 +104,40 @@ class Pattern {
* and nest by reference for PVars.
*/
using Nested = Derived;
+
/*!
* \brief Check if value matches the current pattern.
*
* This call also populates the PVars with matched value.
* The values in PVars are valid until the next call to Match.
*
+ * \param value The value to be matched against
+ *
* \return whether value matches the pattern.
*/
template <typename NodeType>
- bool Match(const NodeType& value) const {
+ inline bool Match(const NodeType& value) const {
+ return Match(value, []() { return true; });
+ }
+
+ /*!
+ * \brief Check if value matches the current pattern.
+ *
+ * This call also populates the PVars with matched value.
+ * The values in PVars are valid until the next call to Match.
+ *
+ * \param value The value to be matched against
+ *
+ * \param cond A callable that performs additional validation,
+ * returning true if the match passes. This will typically be a
+ * lambda function written in terms of the filled PVars.
+ *
+ * \return whether value matches the pattern.
+ */
+ template <typename NodeType, typename Condition>
+ bool Match(const NodeType& value, Condition cond) const {
derived().InitMatch_();
- return derived().Match_(value);
+ return derived().Match_(value) && cond();
}
/*! \return Derived instance of current class. */
const Derived& derived() const { return *static_cast<const Derived*>(this); }
@@ -813,6 +835,86 @@ inline PCallExpr<PIfThenElseOp, TCond, TA, TB>
if_then_else(const Pattern<TCond>
false_value.derived());
}
+template <typename... TPattern>
+class PMatchesOneOf {
+ public:
+ explicit PMatchesOneOf(const TPattern&... patterns) : patterns_{patterns...}
{}
+
+ /*! \brief Check if value matches one of the patterns.
+ *
+ * This call also populates the PVars with matched value based on
+ * the first successful match. The values in PVars are valid until
+ * the next call to Match.
+ *
+ * \param value The value to be matched against.
+ *
+ * \return Whether value matches the pattern.
+ */
+ template <typename NodeType>
+ inline bool Match(const NodeType& value) const {
+ return Match(value, []() { return true; });
+ }
+
+ /*! \brief Check if value matches one of the patterns.
+ *
+ * This call also populates the PVars with matched value based on
+ * the first successful match. The values in PVars are valid until
+ * the next call to Match.
+ *
+ * \param value The value to be matched against.
+ *
+ * \param cond A callable that performs additional validation,
+ * returning true if the match passes. This will typically be a
+ * lambda function written in terms of the filled PVars. This will
+ * be called once for each successful pattern match. If `cond()`
+ * returns false, the next match will be attempted.
+ *
+ * \return Whether value matches the pattern.
+ */
+ template <typename NodeType, typename Condition>
+ inline bool Match(const NodeType& value, Condition cond) const {
+ return MatchImpl(value, cond,
std::make_index_sequence<sizeof...(TPattern)>());
+ }
+
+ private:
+ template <typename NodeType, typename Condition>
+ inline bool MatchImpl(const NodeType& value, Condition cond,
std::index_sequence<>) const {
+ return false;
+ }
+
+ template <typename NodeType, typename Condition, size_t FirstIndex,
size_t... RemainingIndices>
+ inline bool MatchImpl(const NodeType& value, Condition cond,
+ std::index_sequence<FirstIndex, RemainingIndices...>)
const {
+ return std::get<FirstIndex>(patterns_).Match(value, cond) ||
+ MatchImpl(value, cond, std::index_sequence<RemainingIndices...>());
+ }
+
+ // Hold the patterns by const&. This follows the same usage as both
+ // the `PVar`, which occurs as `const PVar<T>&` when it appears
+ // inside other patterns. Because the `PVar<T>::value_` field is
+ // mutable, it can still be updated through these const references.
+ // So long as the call to `Match()` occurs within the same
+ // expression as created the patterns, this avoids accidental copies
+ // without creating dangling references. This may be improved in
+ // the future by use of `constexpr` constructors/operators, allowing
+ // more typical value semantics.
+ std::tuple<const TPattern&...> patterns_;
+};
+
+/* \brief Return a proxy object that returns true after the first match
+ *
+ * In the RewriteSimplifier, there are often several expressions that
+ * simplify to the same resulting expression. This utility allows
+ * them to be specified as a single rule, reducing duplication of the
+ * result/condition of a rewrite.
+ */
+template <typename... TPattern>
+inline std::enable_if_t<(std::is_base_of_v<Pattern<TPattern>, TPattern> && ...
&& true),
+ PMatchesOneOf<TPattern...>>
+matches_one_of(const TPattern&... patterns) {
+ return PMatchesOneOf<TPattern...>(patterns...);
+}
+
} // namespace arith
} // namespace tvm
#endif // TVM_ARITH_PATTERN_MATCH_H_
diff --git a/src/arith/rewrite_simplify.cc b/src/arith/rewrite_simplify.cc
index 6418f28cf8..ce2c3e1a96 100644
--- a/src/arith/rewrite_simplify.cc
+++ b/src/arith/rewrite_simplify.cc
@@ -42,6 +42,19 @@ namespace arith {
using namespace tir;
+// Note: When using matches_one_of or PMatchesOneOf alongside these
+// macros, be careful which patterns are used in the ResExpr. While
+// the different source expressions may be in terms of different PVar,
+// the ResExpr should only contain patterns that are defined in
+// *every* SrcExpr given.
+//
+// Allowed (replacement does not use either c1 or y):
+// TVM_TRY_REWRITE(matches_one_of(x + c1 - c1, x + y - y), x)
+//
+// Forbidden (c3 undefined if the first pattern matches):
+// TVM_TRY_REWRITE(matches_one_of(floormod(x*c1,c2), floormod(x*c1 + c3,
c2)),
+// floormod(x*floormod(c1,c2) + floormod(c3,c2), c2))
+
// macro for doing simple rewrite
#define TVM_TRY_REWRITE(SrcExpr, ResExpr) \
if ((SrcExpr).Match(ret)) { \
@@ -55,14 +68,14 @@ using namespace tir;
}
// macro rewrite only if CondExor is true after match.
-#define TVM_TRY_REWRITE_IF(SrcExpr, ResExpr, CondExpr) \
- if ((SrcExpr).Match(ret) && (CondExpr)) { \
- return (ResExpr).Eval(); \
+#define TVM_TRY_REWRITE_IF(SrcExpr, ResExpr, CondExpr) \
+ if ((SrcExpr).Match(ret, [&]() { return (CondExpr); })) { \
+ return (ResExpr).Eval(); \
}
// macro rewrite + recursive_rewrite only if CondExor is true after match.
#define TVM_TRY_RECURSIVE_REWRITE_IF(SrcExpr, ResExpr, CondExpr) \
- if ((SrcExpr).Match(ret) && (CondExpr)) { \
+ if ((SrcExpr).Match(ret, [&]() { return (CondExpr); })) { \
return RecursiveRewrite((ResExpr).Eval()); \
}
@@ -228,10 +241,13 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
AddNode* op) {
TVM_TRY_REWRITE_IF(max(y + z * c1, x) + z * c2, max(x + z * c2, y),
c1.Eval()->value == -c2.Eval()->value);
- TVM_TRY_REWRITE(max(x, y) + min(x, y), x + y);
- TVM_TRY_REWRITE(min(x, y) + max(x, y), x + y);
- TVM_TRY_REWRITE(max(x, y) + min(y, x), x + y);
- TVM_TRY_REWRITE(min(x, y) + max(y, x), x + y);
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(x, y) + min(x, y),
+ min(x, y) + max(x, y),
+ max(x, y) + min(y, x),
+ min(x, y) + max(y, x),
+ }),
+ x + y);
TVM_TRY_REWRITE_IF(min(x, y + c1) + c2, min(x + c2, y), c1.Eval()->value
== -c2.Eval()->value);
TVM_TRY_REWRITE_IF(min(x + c1, y) + c2, min(x, y + c2), c1.Eval()->value
== -c2.Eval()->value);
@@ -244,33 +260,28 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
AddNode* op) {
// mul co-efficient folding
TVM_TRY_REWRITE(x + x, x * 2);
- TVM_TRY_REWRITE(x * y + x, x * (y + 1));
- TVM_TRY_REWRITE(y * x + x, x * (y + 1));
- TVM_TRY_REWRITE(x + y * x, x * (1 + y));
- TVM_TRY_REWRITE(x + x * y, x * (1 + y));
- TVM_TRY_REWRITE(x * y + x * z, x * (y + z));
- TVM_TRY_REWRITE(y * x + x * z, x * (y + z));
- TVM_TRY_REWRITE(x * y + z * x, x * (y + z));
- TVM_TRY_REWRITE(y * x + z * x, x * (y + z));
+
+ TVM_TRY_REWRITE(matches_one_of(x * y + x, y * x + x, x + y * x, x + x *
y), x * (y + 1));
+
+ TVM_TRY_REWRITE(matches_one_of(x * y + x * z, y * x + x * z, x * y + z *
x, y * x + z * x),
+ x * (y + z));
// DivMod rules
// truc div
TVM_TRY_REWRITE(truncdiv(x, c1) * c1 + truncmod(x, c1), x);
// floor div
- TVM_TRY_REWRITE(floordiv(x, y) * y + floormod(x, y), x);
- TVM_TRY_REWRITE(y * floordiv(x, y) + floormod(x, y), x);
- TVM_TRY_REWRITE(floormod(x, y) + floordiv(x, y) * y, x);
- TVM_TRY_REWRITE(floormod(x, y) + y * floordiv(x, y), x);
+ TVM_TRY_REWRITE(
+ matches_one_of(floordiv(x, y) * y + floormod(x, y), y * floordiv(x, y)
+ floormod(x, y),
+ floormod(x, y) + floordiv(x, y) * y, floormod(x, y) + y
* floordiv(x, y)),
+ x);
TVM_TRY_REWRITE_IF(floordiv(floormod(x, c2) + c1, c2) + floordiv(x, c2),
floordiv(x + c1, c2),
c2.Eval()->value > 0);
// canonicalization rule
// will try rewrite again after canonicalization.
- TVM_TRY_RECURSIVE_REWRITE(x + (c1 - y), (x - y) + c1);
- TVM_TRY_RECURSIVE_REWRITE((c1 - y) + x, (x - y) + c1);
- TVM_TRY_RECURSIVE_REWRITE(x + c1 + y, (x + y) + c1);
- TVM_TRY_RECURSIVE_REWRITE(x + (c1 + y), (x + y) + c1);
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(x + (c1 - y), (c1 - y) + x), (x -
y) + c1);
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(x + c1 + y, x + (c1 + y)), (x +
y) + c1);
TVM_TRY_RECURSIVE_REWRITE(x + max(y, z), max(y, z) + x);
TVM_TRY_RECURSIVE_REWRITE(x + min(y, z), min(y, z) + x);
@@ -344,69 +355,47 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
SubNode* op) {
if (IsIndexType(op->dtype)) {
// Index rules
// cancelation rules
- TVM_TRY_REWRITE((x + y) - y, x);
- TVM_TRY_REWRITE((x + y) - x, y);
- TVM_TRY_REWRITE(x - (y + x), 0 - y);
- TVM_TRY_REWRITE(x - (x + y), 0 - y);
+ TVM_TRY_REWRITE(matches_one_of((x + y) - y, (y + x) - y), x);
+ TVM_TRY_REWRITE(matches_one_of(x - (y + x), x - (x + y)), 0 - y);
- TVM_TRY_REWRITE(min(x, y) - x, min(0, y - x));
- TVM_TRY_REWRITE(min(x, y) - y, min(x - y, 0));
- TVM_TRY_REWRITE(max(x, y) - x, max(0, y - x));
- TVM_TRY_REWRITE(max(x, y) - y, max(x - y, 0));
-
- TVM_TRY_REWRITE(x - max(x, y), min(0, x - y));
- TVM_TRY_REWRITE(y - max(x, y), min(y - x, 0));
- TVM_TRY_REWRITE(x - min(x, y), max(0, x - y));
- TVM_TRY_REWRITE(y - min(x, y), max(y - x, 0));
+ TVM_TRY_REWRITE(matches_one_of(min(x, y) - y, x - max(y, x)), min(x - y,
0));
+ TVM_TRY_REWRITE(matches_one_of(x - max(x, y), min(y, x) - y), min(0, x -
y));
+ TVM_TRY_REWRITE(matches_one_of(max(x, y) - y, x - min(y, x)), max(x - y,
0));
+ TVM_TRY_REWRITE(matches_one_of(x - min(x, y), max(y, x) - y), max(0, x -
y));
// mul co-efficient folding
TVM_TRY_REWRITE(x - x, ZeroWithTypeLike(x));
- TVM_TRY_REWRITE(x * y - x, x * (y - 1));
- TVM_TRY_REWRITE(y * x - x, x * (y - 1));
- TVM_TRY_REWRITE(x - y * x, x * (1 - y));
- TVM_TRY_REWRITE(x - x * y, x * (1 - y));
- TVM_TRY_REWRITE(x * y - x * z, x * (y - z));
- TVM_TRY_REWRITE(y * x - x * z, x * (y - z));
- TVM_TRY_REWRITE(x * y - z * x, x * (y - z));
- TVM_TRY_REWRITE(y * x - z * x, x * (y - z));
+ TVM_TRY_REWRITE(matches_one_of(x * y - x, y * x - x), x * (y - 1));
+ TVM_TRY_REWRITE(matches_one_of(x - y * x, x - x * y), x * (1 - y));
+ TVM_TRY_REWRITE(matches_one_of(x * y - x * z, y * x - x * z, x * y - z *
x, y * x - z * x),
+ x * (y - z));
// constant cancelation
TVM_TRY_REWRITE((x + c1) - c2, x + (c1 - c2));
TVM_TRY_REWRITE((c1 - x) - (c2 - y), (y - x) + (c1 - c2));
// cancelization rule involving 4 operands
- TVM_TRY_REWRITE((x + y) - (x + z), y - z);
- TVM_TRY_REWRITE((x + y) - (z + x), y - z);
- TVM_TRY_REWRITE((y + x) - (z + x), y - z);
- TVM_TRY_REWRITE((y + x) - (x + z), y - z);
-
- TVM_TRY_REWRITE(min(x + y, z) - x, min(y, z - x));
- TVM_TRY_REWRITE(min(y + x, z) - x, min(y, z - x));
- TVM_TRY_REWRITE(min(z, x + y) - x, min(z - x, y));
- TVM_TRY_REWRITE(min(z, y + x) - x, min(z - x, y));
-
- TVM_TRY_REWRITE(max(x + y, z) - x, max(y, z - x));
- TVM_TRY_REWRITE(max(y + x, z) - x, max(y, z - x));
- TVM_TRY_REWRITE(max(z, x + y) - x, max(z - x, y));
- TVM_TRY_REWRITE(max(z, y + x) - x, max(z - x, y));
-
- TVM_TRY_REWRITE(x - min(x + y, z), max(0 - y, x - z));
- TVM_TRY_REWRITE(x - min(y + x, z), max(0 - y, x - z));
- TVM_TRY_REWRITE(x - min(z, x + y), max(x - z, 0 - y));
- TVM_TRY_REWRITE(x - min(z, y + x), max(x - z, 0 - y));
+ TVM_TRY_REWRITE(
+ matches_one_of((x + y) - (x + z), (x + y) - (z + x), (y + x) - (z +
x), (y + x) - (x + z)),
+ y - z);
+
+ TVM_TRY_REWRITE(matches_one_of(min(x + y, z) - x, min(y + x, z) - x),
min(y, z - x));
+ TVM_TRY_REWRITE(matches_one_of(min(z, x + y) - x, min(z, y + x) - x),
min(z - x, y));
+
+ TVM_TRY_REWRITE(matches_one_of(max(x + y, z) - x, max(y + x, z) - x),
max(y, z - x));
+ TVM_TRY_REWRITE(matches_one_of(max(z, x + y) - x, max(z, y + x) - x),
max(z - x, y));
+
+ TVM_TRY_REWRITE(matches_one_of(x - min(x + y, z), x - min(y + x, z)),
max(0 - y, x - z));
+ TVM_TRY_REWRITE(matches_one_of(x - min(z, x + y), x - min(z, y + x)),
max(x - z, 0 - y));
TVM_TRY_REWRITE(min(x, y) - min(y, x), ZeroWithTypeLike(x));
TVM_TRY_REWRITE(max(x, y) - max(y, x), ZeroWithTypeLike(x));
- TVM_TRY_REWRITE_IF(min(b1, b2) - min(s1, s2), b1 - s1,
- CanProveEqual(((b1 - s1) - (b2 - s2)).Eval(), 0));
+ TVM_TRY_REWRITE_IF(matches_one_of(min(b1, b2) - min(s1, s2), min(b1, b2) -
min(s2, s1)),
+ b1 - s1, CanProveEqual(((b1 - s1) - (b2 - s2)).Eval(),
0));
- TVM_TRY_REWRITE_IF(min(b1, b2) - min(s1, s2), b1 - s2,
- CanProveEqual(((b1 - s2) - (b2 - s1)).Eval(), 0));
- TVM_TRY_REWRITE_IF(max(b1, b2) - max(s1, s2), b1 - s1,
- CanProveEqual(((b1 - s1) - (b2 - s2)).Eval(), 0));
- TVM_TRY_REWRITE_IF(max(b1, b2) - max(s1, s2), b1 - s2,
- CanProveEqual(((b1 - s2) - (b2 - s1)).Eval(), 0));
+ TVM_TRY_REWRITE_IF(matches_one_of(max(b1, b2) - max(s1, s2), max(b1, b2) -
max(s2, s1)),
+ b1 - s1, CanProveEqual(((b1 - s1) - (b2 - s2)).Eval(),
0));
// DivMod rules
// trucdiv
@@ -529,8 +518,9 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const MulNode*
op) {
// Vector rules
if (op->dtype.lanes() != 1) {
TVM_TRY_REWRITE(broadcast(x, lanes) * broadcast(y, lanes), broadcast(x *
y, lanes));
- TVM_TRY_REWRITE(ramp(b1, s1, lanes) * broadcast(x, lanes), ramp(b1 * x, s1
* x, lanes));
- TVM_TRY_REWRITE(broadcast(x, lanes) * ramp(b1, s1, lanes), ramp(b1 * x, s1
* x, lanes));
+ TVM_TRY_REWRITE(matches_one_of(ramp(b1, s1, lanes) * broadcast(x, lanes),
+ broadcast(x, lanes) * ramp(b1, s1, lanes)),
+ ramp(b1 * x, s1 * x, lanes));
TVM_TRY_REWRITE_IF(broadcast(c3, lanes) * x, broadcast(c3, lanes),
c3.Eval()->value == 0.0f);
}
@@ -538,8 +528,7 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const MulNode*
op) {
// constant simplification rule
TVM_TRY_REWRITE((x + c1) * c2, x * c2 + c1 * c2);
TVM_TRY_REWRITE((x * c1) * c2, x * (c1 * c2));
- TVM_TRY_REWRITE(min(x, y) * max(x, y), x * y);
- TVM_TRY_REWRITE(max(x, y) * min(x, y), x * y);
+ TVM_TRY_REWRITE(matches_one_of(min(x, y) * max(x, y), max(x, y) * min(x,
y)), x * y);
// Two representations of const*ceildiv(x, c1)
TVM_TRY_REWRITE_IF(floordiv(x - floormod(x, c2), c1) * c1, x - floormod(x,
c2),
@@ -629,8 +618,7 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const DivNode*
op) {
}
TVM_TRY_REWRITE(truncdiv(x, x), OneWithTypeLike(x));
- TVM_TRY_REWRITE(truncdiv(x * c1, x), c1);
- TVM_TRY_REWRITE(truncdiv(c1 * x, x), c1);
+ TVM_TRY_REWRITE(matches_one_of(truncdiv(x * c1, x), truncdiv(c1 * x, x)),
c1);
// Rules involving 2-operands.
TVM_TRY_REWRITE_IF(truncdiv(x * c1 + y, c2), x * truncdiv(c1, c2) +
truncdiv(y, c2),
@@ -689,39 +677,24 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
DivNode* op) {
c1.Eval()->value % c2.Eval()->value == 0 &&
CanProveGreaterEqual(x.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(x + y, x), truncdiv(y, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(y + x, x), truncdiv(y, x) + 1,
+ TVM_TRY_REWRITE_IF(matches_one_of(truncdiv(x + y, x), truncdiv(y + x, x)),
truncdiv(y, x) + 1,
CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0));
TVM_TRY_REWRITE_IF(
- truncdiv((x + y) + z, x), truncdiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0) && CanProveGreaterEqual((y +
z).Eval(), 0));
- TVM_TRY_REWRITE_IF(
- truncdiv((y + x) + z, x), truncdiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0) && CanProveGreaterEqual((y +
z).Eval(), 0));
- TVM_TRY_REWRITE_IF(
- truncdiv(y + (z + x), x), truncdiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0) && CanProveGreaterEqual((y +
z).Eval(), 0));
- TVM_TRY_REWRITE_IF(
- truncdiv(y + (x + z), x), truncdiv(y + z, x) + 1,
+ matches_one_of(truncdiv((x + y) + z, x), truncdiv((y + x) + z, x),
truncdiv(y + (z + x), x),
+ truncdiv(y + (x + z), x)),
+ truncdiv(y + z, x) + 1,
CanProveGreaterEqual(x.Eval(), 0) && CanProveGreaterEqual((y +
z).Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(x * y, y), x,
- CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(y * x, y), x,
+ TVM_TRY_REWRITE_IF(matches_one_of(truncdiv(x * y, y), truncdiv(y * x, y)),
x,
CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(x * z + y, z), x + truncdiv(y, z),
- CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0) &&
- CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(z * x + y, z), x + truncdiv(y, z),
+ TVM_TRY_REWRITE_IF(matches_one_of(truncdiv(x * z + y, z), truncdiv(z * x +
y, z)),
+ x + truncdiv(y, z),
CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0) &&
CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(y + x * z, z), truncdiv(y, z) + x,
- CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0) &&
- CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(truncdiv(y + z * x, z), truncdiv(y, z) + x,
+ TVM_TRY_REWRITE_IF(matches_one_of(truncdiv(y + x * z, z), truncdiv(y + z *
x, z)),
+ truncdiv(y, z) + x,
CanProveGreaterEqual(x.Eval(), 0) &&
CanProveGreaterEqual(y.Eval(), 0) &&
CanProveGreaterEqual(z.Eval(), 0));
}
@@ -887,8 +860,7 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
FloorDivNode* op) {
}
TVM_TRY_REWRITE(floordiv(x, x), OneWithTypeLike(x));
- TVM_TRY_REWRITE(floordiv(x * c1, x), c1);
- TVM_TRY_REWRITE(floordiv(c1 * x, x), c1);
+ TVM_TRY_REWRITE(matches_one_of(floordiv(x * c1, x), floordiv(c1 * x, x)),
c1);
// Rules involving 2-operands.
TVM_TRY_REWRITE_IF(floordiv(min(x * c1, y), c2), min(x * floordiv(c1, c2),
floordiv(y, c2)),
@@ -911,10 +883,8 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
FloorDivNode* op) {
c2.Eval()->value % c1.Eval()->value == 0 &&
CanProveEqual(floordiv(y.Eval() + z.Eval(),
c1.Eval()), 0));
- TVM_TRY_REWRITE_IF(floordiv(x * c1 - y + z, c2), x * floordiv(c1, c2) +
floordiv(z - y, c2),
- c2.Eval()->value > 0 && c1.Eval()->value %
c2.Eval()->value == 0);
-
- TVM_TRY_REWRITE_IF(floordiv(x * c1 + y - z, c2), x * floordiv(c1, c2) +
floordiv(y - z, c2),
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv(x * c1 - y + z, c2), floordiv(x
* c1 + z - y, c2)),
+ x * floordiv(c1, c2) + floordiv(z - y, c2),
c2.Eval()->value > 0 && c1.Eval()->value %
c2.Eval()->value == 0);
TVM_TRY_REWRITE_IF(floordiv(y + x * c1 + z, c2), x * floordiv(c1, c2) +
floordiv(y + z, c2),
@@ -925,30 +895,20 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
FloorDivNode* op) {
TVM_TRY_REWRITE_IF(floordiv(x * c1, x * c2), floordiv(c1, c2),
c2.Eval()->value > 0);
- TVM_TRY_REWRITE_IF(floordiv(x + y, x), floordiv(y, x) + 1,
CanProveGreaterEqual(x.Eval(), 0));
-
- TVM_TRY_REWRITE_IF(floordiv(y + x, x), floordiv(y, x) + 1,
CanProveGreaterEqual(x.Eval(), 0));
-
- TVM_TRY_REWRITE_IF(floordiv((x + y) + z, x), floordiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv((y + x) + z, x), floordiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(y + (z + x), x), floordiv(y + z, x) + 1,
- CanProveGreaterEqual(x.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(y + (x + z), x), floordiv(y + z, x) + 1,
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv(x + y, x), floordiv(y + x, x)),
floordiv(y, x) + 1,
CanProveGreaterEqual(x.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(x * y, y), x, CanProveGreaterEqual(y.Eval(),
0));
- TVM_TRY_REWRITE_IF(floordiv(y * x, y), x, CanProveGreaterEqual(y.Eval(),
0));
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv((x + y) + z, x), floordiv((y +
x) + z, x),
+ floordiv(y + (z + x), x), floordiv(y +
(x + z), x)),
+ floordiv(y + z, x) + 1, CanProveGreaterEqual(x.Eval(),
0));
+
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv(x * y, y), floordiv(y * x, y)),
x,
+ CanProveGreaterEqual(y.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(x * z + y, z), x + floordiv(y, z),
- CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(z * x + y, z), x + floordiv(y, z),
- CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(y + x * z, z), floordiv(y, z) + x,
- CanProveGreaterEqual(z.Eval(), 0));
- TVM_TRY_REWRITE_IF(floordiv(y + z * x, z), floordiv(y, z) + x,
- CanProveGreaterEqual(z.Eval(), 0));
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv(x * z + y, z), floordiv(z * x +
y, z)),
+ x + floordiv(y, z), CanProveGreaterEqual(z.Eval(), 0));
+ TVM_TRY_REWRITE_IF(matches_one_of(floordiv(y + x * z, z), floordiv(y + z *
x, z)),
+ floordiv(y, z) + x, CanProveGreaterEqual(z.Eval(), 0));
TVM_TRY_REWRITE_IF(floordiv(x - floormod(x, c1), c1), floordiv(x, c1),
c1.Eval()->value != 0);
}
@@ -1021,8 +981,7 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
FloorModNode* op) {
TVM_TRY_REWRITE_IF(floormod(x * c1, x * c2), x * floormod(c1, c2),
c2.Eval()->value != 0);
- TVM_TRY_REWRITE(floormod(x * y, y), ZeroWithTypeLike(x));
- TVM_TRY_REWRITE(floormod(y * x, y), ZeroWithTypeLike(y));
+ TVM_TRY_REWRITE(matches_one_of(floormod(x * y, y), floormod(y * x, y)),
ZeroWithTypeLike(y));
// try modular analysis
if (floormod(x, c1).Match(ret)) {
@@ -1090,70 +1049,71 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
MinNode* op) {
}
// DivMod rules
- // Divide up rounding: truc div
// NOTE: trucdiv(x, y) >= floordiv(x, y)
- TVM_TRY_REWRITE_IF(min(truncdiv(x + c1, c2) * c2, x), x,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(min(truncdiv(x + c1, c2) * c2, max(x, c2)), max(x, c2),
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value &&
- CanProveGreaterEqual(x.Eval(), 1));
-
- TVM_TRY_REWRITE_IF(min(x, truncdiv(x + c1, c2) * c2), x,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(min(max(x, c2), truncdiv(x + c1, c2) * c2), max(x, c2),
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value &&
- CanProveGreaterEqual(x.Eval(), 1));
-
- // Divide up rounding: floor div
- TVM_TRY_REWRITE_IF(min(floordiv(x + c1, c2) * c2, x), x,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(min(floordiv(x + c1, c2) * c2, max(x, c2)), max(x, c2),
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value &&
- CanProveGreaterEqual(x.Eval(), 1));
-
- TVM_TRY_REWRITE_IF(min(x, floordiv(x + c1, c2) * c2), x,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(min(max(x, c2), floordiv(x + c1, c2) * c2), max(x, c2),
+ TVM_TRY_REWRITE_IF(
+ matches_one_of(min(truncdiv(x + c1, c2) * c2, x), min(x, truncdiv(x +
c1, c2) * c2),
+ min(floordiv(x + c1, c2) * c2, x), min(x, floordiv(x +
c1, c2) * c2)),
+ x, c2.Eval()->value > 0 && c1.Eval()->value + 1 == c2.Eval()->value);
+
+ TVM_TRY_REWRITE_IF(matches_one_of(min(truncdiv(x + c1, c2) * c2, max(x,
c2)),
+ min(max(x, c2), truncdiv(x + c1, c2) *
c2),
+ min(floordiv(x + c1, c2) * c2, max(x,
c2)),
+ min(max(x, c2), floordiv(x + c1, c2) *
c2)),
+ max(x, c2),
c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value &&
CanProveGreaterEqual(x.Eval(), 1));
- TVM_TRY_REWRITE_IF(min(x, floordiv(x, c2) * c2), floordiv(x, c2) * c2,
c2.Eval()->value > 0);
- TVM_TRY_REWRITE_IF(min(floordiv(x, c2) * c2, x), floordiv(x, c2) * c2,
c2.Eval()->value > 0);
-
- TVM_TRY_REWRITE(min(max(x, y), min(x, y)), min(x, y));
- TVM_TRY_REWRITE(min(max(x, y), min(y, x)), min(x, y));
- TVM_TRY_REWRITE(min(min(x, y), max(x, y)), min(x, y));
- TVM_TRY_REWRITE(min(min(x, y), max(y, x)), min(x, y));
-
- TVM_TRY_REWRITE(min(max(x, y), x), x);
- TVM_TRY_REWRITE(min(max(x, y), y), y);
- TVM_TRY_REWRITE(min(min(x, y), x), min(x, y));
- TVM_TRY_REWRITE(min(min(x, y), y), min(x, y));
-
- TVM_TRY_REWRITE(min(x, max(x, y)), x);
- TVM_TRY_REWRITE(min(y, max(x, y)), y);
- TVM_TRY_REWRITE(min(x, min(x, y)), min(x, y));
- TVM_TRY_REWRITE(min(y, min(x, y)), min(x, y));
+ TVM_TRY_REWRITE_IF(matches_one_of(min(x, floordiv(x, c2) * c2),
min(floordiv(x, c2) * c2, x)),
+ floordiv(x, c2) * c2, c2.Eval()->value > 0);
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ min(max(x, y), min(x, y)),
+ min(max(x, y), min(y, x)),
+ min(min(x, y), max(x, y)),
+ min(min(x, y), max(y, x)),
+ min(min(x, y), x),
+ min(min(x, y), y),
+ min(x, min(x, y)),
+ min(y, min(x, y)),
+ }),
+ min(x, y));
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ min(max(x, y), x),
+ min(max(y, x), x),
+ min(x, max(x, y)),
+ min(x, max(y, x)),
+ }),
+ x);
TVM_TRY_REWRITE(min(min(min(x, y), z), y), min(min(x, y), z));
TVM_TRY_REWRITE(min(min(min(min(x, y), z), s1), y), min(min(min(x, y), z),
s1));
TVM_TRY_REWRITE(min(min(min(min(min(x, y), z), s1), s2), y),
min(min(min(min(x, y), z), s1), s2));
- TVM_TRY_REWRITE(min(max(x, y), max(x, z)), max(min(y, z), x));
- TVM_TRY_REWRITE(min(max(x, y), max(z, x)), max(min(y, z), x));
- TVM_TRY_REWRITE(min(max(y, x), max(x, z)), max(min(y, z), x));
- TVM_TRY_REWRITE(min(max(y, x), max(z, x)), max(min(y, z), x));
-
- TVM_TRY_REWRITE(min(min(x, y), min(x, z)), min(min(y, z), x));
- TVM_TRY_REWRITE(min(min(x, y), min(z, x)), min(min(y, z), x));
- TVM_TRY_REWRITE(min(min(y, x), min(x, z)), min(min(y, z), x));
- TVM_TRY_REWRITE(min(min(y, x), min(z, x)), min(min(y, z), x));
-
- TVM_TRY_REWRITE(min(y + x, z + x), min(y, z) + x);
- TVM_TRY_REWRITE(min(y + x, x + z), min(y, z) + x);
- TVM_TRY_REWRITE(min(x + y, x + z), min(y, z) + x);
- TVM_TRY_REWRITE(min(x + y, z + x), min(y, z) + x);
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ min(max(x, y), max(x, z)),
+ min(max(x, y), max(z, x)),
+ min(max(y, x), max(x, z)),
+ min(max(y, x), max(z, x)),
+ }),
+ max(min(y, z), x));
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ min(min(x, y), min(x, z)),
+ min(min(x, y), min(z, x)),
+ min(min(y, x), min(x, z)),
+ min(min(y, x), min(z, x)),
+ }),
+ min(min(y, z), x));
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ min(y + x, z + x),
+ min(y + x, x + z),
+ min(x + y, x + z),
+ min(x + y, z + x),
+ }),
+ min(y, z) + x);
// sub distribution
TVM_TRY_REWRITE(min(y - x, z - x), min(y, z) - x);
@@ -1265,34 +1225,46 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
MaxNode* op) {
// DivMod rules
// Divide up rounding: truc div
// NOTE: trucdiv(x, y) >= floordiv(x, y)
- TVM_TRY_REWRITE_IF(max(truncdiv(x + c1, c2) * c2, x), truncdiv(x + c1, c2)
* c2,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(max(x, truncdiv(x + c1, c2) * c2), truncdiv(x + c1, c2)
* c2,
+ TVM_TRY_REWRITE_IF((PMatchesOneOf{
+ max(truncdiv(x + c1, c2) * c2, x),
+ max(x, truncdiv(x + c1, c2) * c2),
+ }),
+ truncdiv(x + c1, c2) * c2,
c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
// Divide up rounding: floor div
- TVM_TRY_REWRITE_IF(max(floordiv(x + c1, c2) * c2, x), floordiv(x + c1, c2)
* c2,
+ TVM_TRY_REWRITE_IF((PMatchesOneOf{
+ max(floordiv(x + c1, c2) * c2, x),
+ max(x, floordiv(x + c1, c2) * c2),
+ }),
+ floordiv(x + c1, c2) * c2,
c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
- TVM_TRY_REWRITE_IF(max(x, floordiv(x + c1, c2) * c2), floordiv(x + c1, c2)
* c2,
- c2.Eval()->value > 0 && c1.Eval()->value + 1 ==
c2.Eval()->value);
-
- TVM_TRY_REWRITE_IF(max(floordiv(x, c2) * c2, x), x, c2.Eval()->value > 0);
- TVM_TRY_REWRITE_IF(max(x, floordiv(x, c2) * c2), x, c2.Eval()->value > 0);
- TVM_TRY_REWRITE(max(min(x, y), max(x, y)), max(x, y));
- TVM_TRY_REWRITE(max(min(x, y), max(y, x)), max(x, y));
- TVM_TRY_REWRITE(max(max(x, y), min(x, y)), max(x, y));
- TVM_TRY_REWRITE(max(max(x, y), min(y, x)), max(x, y));
-
- TVM_TRY_REWRITE(max(min(x, y), x), x);
- TVM_TRY_REWRITE(max(min(x, y), y), y);
- TVM_TRY_REWRITE(max(max(x, y), x), max(x, y));
- TVM_TRY_REWRITE(max(max(x, y), y), max(x, y));
-
- TVM_TRY_REWRITE(max(x, min(x, y)), x);
- TVM_TRY_REWRITE(max(y, min(x, y)), y);
- TVM_TRY_REWRITE(max(x, max(x, y)), max(x, y));
- TVM_TRY_REWRITE(max(y, max(x, y)), max(x, y));
+ TVM_TRY_REWRITE_IF((PMatchesOneOf{
+ max(floordiv(x, c2) * c2, x),
+ max(x, floordiv(x, c2) * c2),
+ }),
+ x, c2.Eval()->value > 0);
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(min(x, y), x),
+ max(min(y, x), x),
+ max(x, min(x, y)),
+ max(x, min(y, x)),
+ }),
+ x);
+
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(min(x, y), max(x, y)),
+ max(min(x, y), max(y, x)),
+ max(max(x, y), min(x, y)),
+ max(max(x, y), min(y, x)),
+ max(max(x, y), x),
+ max(max(x, y), y),
+ max(x, max(x, y)),
+ max(y, max(x, y)),
+ }),
+ max(x, y));
TVM_TRY_REWRITE(max(max(max(x, y), z), y), max(max(x, y), z));
TVM_TRY_REWRITE(max(max(max(max(x, y), z), s1), y), max(max(max(x, y), z),
s1));
@@ -1300,22 +1272,31 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
MaxNode* op) {
max(max(max(max(x, y), z), s1), s2));
// max/max cancelation
- TVM_TRY_REWRITE(max(max(x, y), max(x, z)), max(max(y, z), x));
- TVM_TRY_REWRITE(max(max(x, y), max(z, x)), max(max(y, z), x));
- TVM_TRY_REWRITE(max(max(y, x), max(x, z)), max(max(y, z), x));
- TVM_TRY_REWRITE(max(max(y, x), max(z, x)), max(max(y, z), x));
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(max(x, y), max(x, z)),
+ max(max(x, y), max(z, x)),
+ max(max(y, x), max(x, z)),
+ max(max(y, x), max(z, x)),
+ }),
+ max(max(y, z), x));
// max/min distribution
- TVM_TRY_REWRITE(max(min(x, y), min(x, z)), min(max(y, z), x));
- TVM_TRY_REWRITE(max(min(x, y), min(z, x)), min(max(y, z), x));
- TVM_TRY_REWRITE(max(min(y, x), min(x, z)), min(max(y, z), x));
- TVM_TRY_REWRITE(max(min(y, x), min(z, x)), min(max(y, z), x));
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(min(x, y), min(x, z)),
+ max(min(x, y), min(z, x)),
+ max(min(y, x), min(x, z)),
+ max(min(y, x), min(z, x)),
+ }),
+ min(max(y, z), x));
// add distribution
- TVM_TRY_REWRITE(max(y + x, z + x), max(y, z) + x);
- TVM_TRY_REWRITE(max(y + x, x + z), max(y, z) + x);
- TVM_TRY_REWRITE(max(x + y, x + z), max(y, z) + x);
- TVM_TRY_REWRITE(max(x + y, z + x), max(y, z) + x);
+ TVM_TRY_REWRITE((PMatchesOneOf{
+ max(y + x, z + x),
+ max(y + x, x + z),
+ max(x + y, x + z),
+ max(x + y, z + x),
+ }),
+ max(y, z) + x);
// sub distribution
TVM_TRY_REWRITE(max(y - x, z - x), max(y, z) - x);
@@ -1404,7 +1385,7 @@ PrimExpr RewriteSimplifier::Impl::ApplyRewriteRules(EQ
ret) {
// Pattern var to match any expression
PVar<PrimExpr> x, y;
// Pattern var match IntImm
- PVar<IntImm> c1;
+ PVar<IntImm> c1, c2;
PVar<int> lanes;
// vector rule
@@ -1422,9 +1403,9 @@ PrimExpr RewriteSimplifier::Impl::ApplyRewriteRules(EQ
ret) {
}
TVM_TRY_REWRITE(c1 == x, x == c1);
- TVM_TRY_REWRITE(x - c1 == 0, x == c1);
- TVM_TRY_REWRITE(c1 - x == 0, x == c1);
- TVM_TRY_REWRITE(x + c1 == 0, x == 0 - c1);
+ TVM_TRY_REWRITE(x - c1 == c2, x == c2 + c1);
+ TVM_TRY_REWRITE(c1 - x == c2, x == c1 - c2);
+ TVM_TRY_REWRITE(x + c1 == c2, x == c2 - c1);
TVM_TRY_RECURSIVE_REWRITE(x * y == 0, x == 0 || y == 0);
}
return std::move(ret);
@@ -1562,8 +1543,6 @@ PrimExpr RewriteSimplifier::Impl::ApplyRewriteRules(LT
ret) {
TVM_TRY_REWRITE(x < x + z, 0 < z);
TVM_TRY_REWRITE(x < z + x, 0 < z);
TVM_TRY_REWRITE(x < x - z, z < 0);
- TVM_TRY_REWRITE(c1 < x + c2, c1 - c2 < x);
- TVM_TRY_REWRITE(c1 < c2 - x, x < c2 - c1);
TVM_TRY_REWRITE_IF(x * c1 < y * c1, x < y, c1.Eval()->value > 0);
TVM_TRY_REWRITE_IF(x * c1 < y * c1, y < x, c1.Eval()->value < 0);
@@ -1642,20 +1621,34 @@ PrimExpr RewriteSimplifier::Impl::ApplyRewriteRules(LT
ret) {
TVM_TRY_RECURSIVE_REWRITE(z < min(x, y), z < x && z < y);
TVM_TRY_RECURSIVE_REWRITE(z < max(x, y), z < x || z < y);
+ // clang-format on
+
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(c1 < x + c2, c1 - x < c2), c1 -
c2 < x);
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(c1 < c2 - x, x + c1 < c2), x < c2
- c1);
+ TVM_TRY_RECURSIVE_REWRITE(c1 < x - c2, c1 + c2 < x);
+ TVM_TRY_RECURSIVE_REWRITE(x - c2 < c1, x < c1 + c2);
+
TVM_TRY_RECURSIVE_REWRITE(x < c1 - y, x + y < c1);
- TVM_TRY_RECURSIVE_REWRITE(x < c1 + y, x - y < c1);
TVM_TRY_RECURSIVE_REWRITE(c1 - y < x, c1 < x + y);
- TVM_TRY_RECURSIVE_REWRITE(c1 + y < x, c1 < x - y);
- TVM_TRY_RECURSIVE_REWRITE(x + c1 < c2, x < c2 - c1);
- TVM_TRY_RECURSIVE_REWRITE(x - c1 < c2, x < c2 + c1);
- TVM_TRY_REWRITE(x - c1 < 0, x < c1);
+ TVM_TRY_RECURSIVE_REWRITE(x < c1 + y, x - y < c1);
+ TVM_TRY_RECURSIVE_REWRITE(c1 + y < x, c1 < x - y);
- TVM_TRY_RECURSIVE_REWRITE(x - 1 < y, x <= y);
- TVM_TRY_RECURSIVE_REWRITE(x < y + 1, x <= y);
- TVM_TRY_RECURSIVE_REWRITE(x + (-1) < y, x <= y);
- TVM_TRY_RECURSIVE_REWRITE(x < y - (-1), x <= y);
- // clang-format on
+ if ((x + c1 < y + c2).Match(ret)) {
+ int64_t diff = c2.Eval()->value - c1.Eval()->value;
+ PrimExpr out = [&]() {
+ if (diff == 0) {
+ return (x < y).Eval();
+ } else if (diff == 1) {
+ return (x <= y).Eval();
+ } else if (diff < 0) {
+ return (x + (-diff) < y).Eval();
+ } else {
+ return (x < y + diff).Eval();
+ }
+ }();
+ return RecursiveRewrite(out);
+ }
}
return std::move(ret);
}
@@ -1762,31 +1755,42 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
AndNode* op) {
TVM_TRY_REWRITE_IF(x < c1 && c2 < x, cfalse, c2.Eval()->value + 1 >=
c1.Eval()->value);
TVM_TRY_REWRITE_IF(c2 < x && x < c1, cfalse, c2.Eval()->value + 1 >=
c1.Eval()->value);
- TVM_TRY_REWRITE_IF(x < c1 && c2 <= x, cfalse, c2.Eval()->value >=
c1.Eval()->value);
- TVM_TRY_REWRITE_IF(c2 <= x && x < c1, cfalse, c2.Eval()->value >=
c1.Eval()->value);
- TVM_TRY_REWRITE_IF(x <= c1 && c2 < x, cfalse, c2.Eval()->value >=
c1.Eval()->value);
- TVM_TRY_REWRITE_IF(c2 < x && x <= c1, cfalse, c2.Eval()->value >=
c1.Eval()->value);
-
- TVM_TRY_REWRITE_IF(x <= c1 && c2 <= x, cfalse, c2.Eval()->value >
c1.Eval()->value);
- TVM_TRY_REWRITE_IF(c2 <= x && x <= c1, cfalse, c2.Eval()->value >
c1.Eval()->value);
-
- TVM_TRY_REWRITE(x == c1 && x != c2, x == c1 && c1 != c2);
- TVM_TRY_REWRITE(x != c2 && x == c1, x == c1 && c1 != c2);
-
- TVM_TRY_RECURSIVE_REWRITE(floordiv(x, c2) == c1 && floormod(x, c2) == c3, x
== c1 * c2 + c3);
- TVM_TRY_RECURSIVE_REWRITE(floormod(x, c2) == c3 && floordiv(x, c2) == c1, x
== c1 * c2 + c3);
-
- TVM_TRY_RECURSIVE_REWRITE_IF(0 <= x - y * c1 &&
- x - y * c1<c1, y == floordiv(x, c1),
c1.Eval()->value> 0);
- TVM_TRY_RECURSIVE_REWRITE_IF(x - y * c1 < c1 && 0 <= x - y * c1, y ==
floordiv(x, c1),
- c1.Eval()->value > 0);
-
- TVM_TRY_RECURSIVE_REWRITE(c1 < x - y * c1 && x - y * c1 <= 0, y ==
floordiv(x, c1));
- TVM_TRY_RECURSIVE_REWRITE(x - y * c1 < c1 && 0 <= x - y * c1, y ==
floordiv(x, c1));
- TVM_TRY_RECURSIVE_REWRITE_IF(0 <= x + y * c2 && x + y * c2 < c1, y ==
floordiv(x, c1),
- c2.Eval()->value == -c1.Eval()->value);
- TVM_TRY_RECURSIVE_REWRITE_IF(x + y * c2 < c1 && 0 <= x + y * c2, y ==
floordiv(x, c1),
- c2.Eval()->value == -c1.Eval()->value);
+ TVM_TRY_REWRITE_IF((PMatchesOneOf{
+ x < c1 && c2 <= x,
+ c2 <= x && x < c1,
+ x <= c1 && c2 < x,
+ c2 < x && x <= c1,
+ }),
+ cfalse, c2.Eval()->value >= c1.Eval()->value);
+
+ TVM_TRY_REWRITE_IF((PMatchesOneOf{
+ x <= c1 && c2 <= x,
+ c2 <= x && x <= c1,
+ }),
+ cfalse, c2.Eval()->value > c1.Eval()->value);
+
+ TVM_TRY_REWRITE(matches_one_of(x == c1 && x != c2, x != c2 && x == c1), x ==
c1 && c1 != c2);
+
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(floordiv(x, c2) == c1 &&
floormod(x, c2) == c3,
+ floormod(x, c2) == c3 &&
floordiv(x, c2) == c1),
+ x == c1 * c2 + c3);
+
+ TVM_TRY_RECURSIVE_REWRITE_IF((PMatchesOneOf{
+ 0 <= x - y * c1 && x - y * c1 < c1,
+ x - y * c1 < c1 && 0 <= x - y * c1,
+ }),
+ y == floordiv(x, c1), c1.Eval()->value > 0);
+
+ TVM_TRY_RECURSIVE_REWRITE((PMatchesOneOf{
+ c1 < x - y * c1 && x - y * c1 <= 0,
+ x - y * c1 < c1 && 0 <= x - y * c1,
+ }),
+ y == floordiv(x, c1));
+ TVM_TRY_RECURSIVE_REWRITE_IF((PMatchesOneOf{
+ 0 <= x + y * c2 && x + y * c2 < c1,
+ x + y * c2 < c1 && 0 <= x + y * c2,
+ }),
+ y == floordiv(x, c1), c2.Eval()->value ==
-c1.Eval()->value);
TVM_TRY_RECURSIVE_REWRITE_IF(x < c1 && floormod(x, c2) < c3,
x < c1 - c2 + c3 && floormod(x, c2) < c3,
@@ -1804,22 +1808,18 @@ PrimExpr RewriteSimplifier::Impl::VisitExpr_(const
AndNode* op) {
(((c1.Eval()->value + 1) % c2.Eval()->value) + c2.Eval()->value) %
c2.Eval()->value >
c3.Eval()->value);
- TVM_TRY_RECURSIVE_REWRITE(floordiv(x, c2) == c1 && floormod(x, c2) < c3,
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(floordiv(x, c2) == c1 &&
floormod(x, c2) < c3,
+ floormod(x, c2) < c3 && floordiv(x,
c2) == c1),
c1 * c2 <= x && x < c1 * c2 + c3);
- TVM_TRY_RECURSIVE_REWRITE(floormod(x, c2) < c3 && floordiv(x, c2) == c1,
- c1 * c2 <= x && x < c1 * c2 + c3);
- TVM_TRY_RECURSIVE_REWRITE(floordiv(x, c2) == c1 && floormod(x, c2) <= c3,
- c1 * c2 <= x && x <= c1 * c2 + c3);
- TVM_TRY_RECURSIVE_REWRITE(floormod(x, c2) <= c3 && floordiv(x, c2) == c1,
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(floordiv(x, c2) == c1 &&
floormod(x, c2) <= c3,
+ floormod(x, c2) <= c3 &&
floordiv(x, c2) == c1),
c1 * c2 <= x && x <= c1 * c2 + c3);
- TVM_TRY_RECURSIVE_REWRITE(floordiv(x, c2) == c1 && c3 <= floormod(x, c2),
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(floordiv(x, c2) == c1 && c3 <=
floormod(x, c2),
+ c3 <= floormod(x, c2) &&
floordiv(x, c2) == c1),
c1 * c2 + c3 <= x && x < (c1 + 1) * c2);
- TVM_TRY_RECURSIVE_REWRITE(c3 <= floormod(x, c2) && floordiv(x, c2) == c1,
- c1 * c2 + c3 <= x && x < (c1 + 1) * c2);
- TVM_TRY_RECURSIVE_REWRITE(floordiv(x, c2) == c1 && c3 < floormod(x, c2),
- c1 * c2 + c3 < x && x < (c1 + 1) * c2);
- TVM_TRY_RECURSIVE_REWRITE(c3 < floormod(x, c2) && floordiv(x, c2) == c1,
+ TVM_TRY_RECURSIVE_REWRITE(matches_one_of(floordiv(x, c2) == c1 && c3 <
floormod(x, c2),
+ c3 < floormod(x, c2) && floordiv(x,
c2) == c1),
c1 * c2 + c3 < x && x < (c1 + 1) * c2);
TVM_TRY_RECURSIVE_REWRITE(x && (y && z), (x && y) && z);