diff --git a/src/Bounds.cpp b/src/Bounds.cpp index 6641c038f7d8..5a04e4af14d7 100644 --- a/src/Bounds.cpp +++ b/src/Bounds.cpp @@ -22,6 +22,7 @@ #include "SimplifyCorrelatedDifferences.h" #include "Solve.h" #include "StrictifyFloat.h" +#include "Substitute.h" #include "Util.h" #include "Var.h" @@ -2039,6 +2040,102 @@ class FindInnermostVar : public IRVisitor { } }; +// Bind each compound index in the body of an IfThenElse that its condition +// also mentions to a let around it, and repeat the condition in terms of the +// name inside the original one, e.g. +// +// if (min_x <= x*s + r) { f(x*s + r); g(r) } +// +// becomes +// +// let t = x*s + r in if (min_x <= x*s + r) { if (min_x <= t) { f(t); g(r) } } +// +// BoxesTouched bounds an index by an IfThenElse's condition one variable at a +// time, so it can only use a condition that relates several variables if +// they're named by one. The original condition stays outside so that it keeps +// bounding each of those variables on its own, for the other indices they +// appear in. +class NameGuardedIndices : public IRMutator { + using IRMutator::visit; + + Stmt visit(const IfThenElse *op) override { + // Name the indices here before recursing, so that the conditions of + // nested IfThenElses that bound the same index refer to the same let. + Stmt then_case = op->then_case; + vector indices; + auto add_index = [&](const Expr &e) { + if (e.type() != Int(32) || e.as() || is_const(e)) { + return; + } + for (const Expr &i : indices) { + if (equal(i, e)) { + return; + } + } + indices.push_back(e); + }; + // An index and each compound part of it, outermost first, so that the + // largest one the condition mentions gets the name. + auto add_index_and_parts = [&](const Expr &arg) { + add_index(arg); + visit_with( + arg, + [&](auto *self, const Add *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Sub *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Mul *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Select *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Div *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Mod *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Min *op) { add_index(op); self->visit_base(op); }, + [&](auto *self, const Max *op) { add_index(op); self->visit_base(op); }); + }; + visit_with( + then_case, + [&](auto *self, const Call *call) { + if (call->call_type == Call::Halide || call->call_type == Call::Image) { + for (const Expr &arg : call->args) { + add_index_and_parts(arg); + } + } + self->visit_base(call); + }, + [&](auto *self, const Provide *provide) { + for (const Expr &arg : provide->args) { + add_index_and_parts(arg); + } + self->visit_base(provide); + }); + + Expr condition = op->condition; + vector> lets; + for (const Expr &i : indices) { + string name = unique_name('t'); + Expr var = Variable::make(i.type(), name); + Expr new_condition = substitute(i, var, condition); + if (!new_condition.same_as(condition)) { + condition = new_condition; + then_case = substitute(i, var, then_case); + lets.emplace_back(name, i); + } + } + then_case = mutate(then_case); + Stmt else_case = mutate(op->else_case); + + if (lets.empty()) { + if (then_case.same_as(op->then_case) && else_case.same_as(op->else_case)) { + return op; + } + return IfThenElse::make(op->condition, then_case, else_case); + } + then_case = IfThenElse::make(condition, then_case); + Stmt stmt = IfThenElse::make(op->condition, then_case, else_case); + for (const auto &let : reverse_view(lets)) { + stmt = LetStmt::make(let.first, let.second, stmt); + } + return stmt; + } +}; + // Place innermost vars in an IfThenElse's condition as far to the left as possible. class SolveIfThenElse : public IRMutator { protected: @@ -3126,6 +3223,7 @@ map boxes_touched(const Expr &e, Stmt s, bool consider_calls, bool // as possible, so that BoxesTouched can prune the variable scope tighter // when encountering the IfThenElse. if (s.defined()) { + s = NameGuardedIndices()(s); s = SolveIfThenElse()(s); } diff --git a/src/IRMatch.h b/src/IRMatch.h index 6fa4cadc4eae..be4b6674568f 100644 --- a/src/IRMatch.h +++ b/src/IRMatch.h @@ -451,6 +451,15 @@ struct WildConst { return make_const_expr(val, type); } + // The matched value itself, no IR built. Integer constants only. + HALIDE_ALWAYS_INLINE + int64_t bound_const_int(MatcherState &state) const noexcept { + halide_scalar_value_t val; + Type type; + state.get_bound_const(i, val, type); + return val.u.i64; + } + constexpr static bool foldable = true; [[nodiscard]] HALIDE_ALWAYS_INLINE bool make_folded_const(halide_scalar_value_t &val, Type &ty, MatcherState &state) const noexcept { @@ -490,6 +499,13 @@ struct Wild { return state.get_binding(i); } + // The bound node itself. Unlike make() this doesn't even touch a reference + // count, which lets predicates inspect what matched for free. + HALIDE_ALWAYS_INLINE + const BaseExprNode *bound_node(MatcherState &state) const noexcept { + return state.get_binding(i); + } + constexpr static bool foldable = false; }; @@ -549,6 +565,12 @@ struct IntLiteral { return v == b.v; } + // The literal value itself, no IR built. + HALIDE_ALWAYS_INLINE + int64_t bound_const_int(MatcherState &state) const noexcept { + return v; + } + HALIDE_ALWAYS_INLINE Expr make(MatcherState &state, Type type_hint) const { return make_const(type_hint, v); @@ -2554,7 +2576,7 @@ struct CanProve { // Includes a raw call to an inlined make method, so don't inline. [[nodiscard]] HALIDE_NEVER_INLINE bool make_folded_const(halide_scalar_value_t &val, Type &ty, MatcherState &state) const { Expr condition = a.make(state, {}); - condition = prover->mutate(condition, nullptr); + condition = prover->simplify_can_prove_condition(condition); val.u.u64 = is_const_one(condition); ty = Bool(condition.type().lanes()); return false; @@ -2573,6 +2595,115 @@ std::ostream &operator<<(std::ostream &s, const CanProve &op) { return s; } +// Detects patterns that can hand back the node they matched without building +// anything. The predicates below are restricted to these, which is what makes +// them allocation-free: it is a compile error to ask about a derived expression +// like min_diff(x, y + 1). Put the offset on the other side of the comparison +// instead: min_diff(x, y) >= 1. +template +struct has_bound_node : std::false_type {}; + +template +struct has_bound_node().bound_node(std::declval()))>> + : std::true_type {}; + +// As has_bound_node, for terms whose constant reads out as a plain int64_t. +template +struct has_bound_const_int : std::false_type {}; + +template +struct has_bound_const_int().bound_const_int(std::declval()))>> + : std::true_type {}; + +// Bounds on the linear combination (ca * a - cb * b) of two matched +// expressions, derived from the facts the prover has learned. Used as +// (min_diff(x, y, this) >= 0) and friends. ca and cb are constants already in +// hand (matched WildConsts, typically) that sit outside a and b's own IR, so +// peeling can't find them. Allocation-free: ca/cb read as raw ints, a/b as raw +// bound nodes. When nothing is known the fold reports overflow, which the +// rewriter already treats as a failed predicate, so the rule simply doesn't +// fire. +template +struct LinearDiffBound { + struct pattern_tag {}; + A a; + CA ca; + B b; + CB cb; + Prover *prover; + + static_assert(has_bound_node::value && has_bound_node::value, + "The a/b operands of min_diff/max_diff must be wildcards, so " + "that testing the predicate doesn't have to construct any IR."); + static_assert(has_bound_const_int::value && has_bound_const_int::value, + "The coefficient operands of linear_min_diff/linear_max_diff " + "must be WildConsts or integer literals."); + + constexpr static uint32_t binds = bindings::mask | bindings::mask | bindings::mask | bindings::mask; + + // An integer-valued term of a comparison. + constexpr static IRNodeType min_node_type = IRNodeType::IntImm; + constexpr static IRNodeType max_node_type = IRNodeType::IntImm; + constexpr static bool canonical = true; + + constexpr static bool foldable = true; + + [[nodiscard]] HALIDE_ALWAYS_INLINE bool make_folded_const(halide_scalar_value_t &val, Type &ty, MatcherState &state) const noexcept { + int64_t result = 0; + bool known; + if (is_min) { + known = prover->known_min_diff(a.bound_node(state), ca.bound_const_int(state), + b.bound_node(state), cb.bound_const_int(state), &result); + } else { + known = prover->known_max_diff(a.bound_node(state), ca.bound_const_int(state), + b.bound_node(state), cb.bound_const_int(state), &result); + } + val.u.i64 = result; + ty = Int(64); + // An unknown bound reports as overflow, failing the predicate. + return !known; + } +}; + +template +HALIDE_ALWAYS_INLINE auto linear_min_diff(A &&a, CA &&ca, B &&b, CB &&cb, Prover *p) noexcept + -> LinearDiffBound { + assert_is_lvalue_if_expr(); + assert_is_lvalue_if_expr(); + return {pattern_arg(a), pattern_arg(ca), pattern_arg(b), pattern_arg(cb), p}; +} + +template +HALIDE_ALWAYS_INLINE auto linear_max_diff(A &&a, CA &&ca, B &&b, CB &&cb, Prover *p) noexcept + -> LinearDiffBound { + assert_is_lvalue_if_expr(); + assert_is_lvalue_if_expr(); + return {pattern_arg(a), pattern_arg(ca), pattern_arg(b), pattern_arg(cb), p}; +} + +// Bounds on the plain difference (a - b). +template +HALIDE_ALWAYS_INLINE auto min_diff(A &&a, B &&b, Prover *p) noexcept + -> LinearDiffBound { + assert_is_lvalue_if_expr(); + assert_is_lvalue_if_expr(); + return {pattern_arg(a), IntLiteral{1}, pattern_arg(b), IntLiteral{1}, p}; +} + +template +HALIDE_ALWAYS_INLINE auto max_diff(A &&a, B &&b, Prover *p) noexcept + -> LinearDiffBound { + assert_is_lvalue_if_expr(); + assert_is_lvalue_if_expr(); + return {pattern_arg(a), IntLiteral{1}, pattern_arg(b), IntLiteral{1}, p}; +} + +template +std::ostream &operator<<(std::ostream &s, const LinearDiffBound &op) { + s << (is_min ? "linear_min_diff(" : "linear_max_diff(") << op.a << ", " << op.ca << ", " << op.b << ", " << op.cb << ")"; + return s; +} + template struct IsFloat { struct pattern_tag {}; diff --git a/src/Simplify.cpp b/src/Simplify.cpp index 18129614aca7..67f6d3967f58 100644 --- a/src/Simplify.cpp +++ b/src/Simplify.cpp @@ -84,7 +84,300 @@ void Simplify::found_buffer_reference(const string &name, size_t dimensions) { } } +namespace { + +// Peel constant add/mul/div terms off e, maintaining +// +// denom * coeff_in * e_in == coeff * e + off + err +// +// All five accumulate into the caller's running totals, so peels compose: an +// additive term under an already-peeled factor is scaled by it first ((x + c0) +// * c1 -> coeff = c1, off = c1 * c0, e = x). Division is the inexact one -- +// e / c drops a remainder in [0, c - 1] -- so it scales everything by c and +// banks the remainder in err. Walks existing nodes; builds nothing. +void peel_affine_term(const BaseExprNode *&e, int64_t &coeff, int64_t &off, + int64_t &denom, ConstantInterval &err) { + while (true) { + if (e->node_type == IRNodeType::Add) { + const Add *add = (const Add *)e; + if (const IntImm *i = add->b.as()) { + int64_t term; + if (!mul_with_overflow(64, coeff, i->value, &term) || + !add_with_overflow(64, off, term, &off)) { + break; + } + e = add->a.get(); + continue; + } else if (const IntImm *i = add->a.as()) { + int64_t term; + if (!mul_with_overflow(64, coeff, i->value, &term) || + !add_with_overflow(64, off, term, &off)) { + break; + } + e = add->b.get(); + continue; + } + } else if (e->node_type == IRNodeType::Sub) { + const Sub *sub = (const Sub *)e; + if (const IntImm *i = sub->b.as()) { + int64_t term; + if (!mul_with_overflow(64, coeff, i->value, &term) || + !sub_with_overflow(64, off, term, &off)) { + break; + } + e = sub->a.get(); + continue; + } + } else if (e->node_type == IRNodeType::Mul) { + const Mul *mul = (const Mul *)e; + if (const IntImm *i = mul->b.as()) { + if (!mul_with_overflow(64, coeff, i->value, &coeff)) { + break; + } + e = mul->a.get(); + continue; + } else if (const IntImm *i = mul->a.as()) { + if (!mul_with_overflow(64, coeff, i->value, &coeff)) { + break; + } + e = mul->b.get(); + continue; + } + } else if (e->node_type == IRNodeType::Div) { + const Div *div = (const Div *)e; + const IntImm *i = div->b.as(); + // Positive divisors only; a negative one floors the other way. + if (i && i->value > 0) { + const int64_t c = i->value; + int64_t new_denom, new_off, tmp; + if (!mul_with_overflow(64, denom, c, &new_denom) || + !mul_with_overflow(64, off, c, &new_off) || + !mul_with_overflow(64, coeff, c - 1, &tmp)) { + break; + } + // c * coeff * (a / c) == coeff * a - coeff * r, r == a % c. + denom = new_denom; + off = new_off; + err *= c; + err -= ConstantInterval(0, c - 1) * coeff; + e = div->a.get(); + continue; + } + } + // Nothing left to peel. + break; + } +} + +// Rewrite (ca * a - cb * b) as +// +// denom * (ca * a - cb * b) == coeff_a * a' - coeff_b * b' + offset + err +// +// peeling each side independently, so that facts and queries meet at a common +// pair however each was spelled: x vs y + 3, 2 * x vs 4 * x, x vs y / c +// against c * x vs y. Absent a division denom is 1 and err is 0. +void peel_affine_terms(const BaseExprNode *&a, const BaseExprNode *&b, + int64_t &coeff_a, int64_t &coeff_b, int64_t &offset, + int64_t &denom, ConstantInterval &err) { + const BaseExprNode *const a_in = a; + const BaseExprNode *const b_in = b; + const int64_t ca_in = coeff_a, cb_in = coeff_b; + + // Peeled nothing: the pair as it came in, which every caller can still use. + auto give_up = [&]() { + a = a_in; + b = b_in; + coeff_a = ca_in; + coeff_b = cb_in; + denom = 1; + err = ConstantInterval(0, 0); + offset = 0; + }; + + int64_t off_a = 0, off_b = 0, denom_a = 1, denom_b = 1; + ConstantInterval err_a(0, 0), err_b(0, 0); + peel_affine_term(a, coeff_a, off_a, denom_a, err_a); + peel_affine_term(b, coeff_b, off_b, denom_b, err_b); + + // Put the two sides over a common denominator. + if (!mul_with_overflow(64, denom_a, denom_b, &denom)) { + // Nothing useful to say about numbers this large. + give_up(); + return; + } + if (denom == 1) { + // Common-case optimization + err = err_a - err_b; + } else { + if (!mul_with_overflow(64, coeff_a, denom_b, &coeff_a) || + !mul_with_overflow(64, coeff_b, denom_a, &coeff_b) || + !mul_with_overflow(64, off_a, denom_b, &off_a) || + !mul_with_overflow(64, off_b, denom_a, &off_b)) { + give_up(); + return; + } + err = err_a * denom_b - err_b * denom_a; + } + // An offset we can't represent has to sink the whole rewrite: dropping it + // would leave a and b peeled but the relation between them misstated. + if (!sub_with_overflow(64, off_a, off_b, &offset)) { + give_up(); + return; + } +} + +// Reduce (ca, cb) to a coprime, sign-canonical (pa, pb) and a scale s with +// (ca, cb) == s * (pa, pb). False if both coefficients are zero. Facts and +// queries both go through this, so a fact about 2 * x - 4 * y and a query +// about 3 * x - 6 * y meet at the pair (1, 2) with scales 2 and 3. +bool reduce_affine_coeffs(int64_t ca, int64_t cb, int64_t &pa, int64_t &pb, int64_t &s) { + if (ca == 0 && cb == 0) { + return false; + } + if (ca == 1 || ca == -1 || cb == 1 || cb == -1) { + // A unit coefficient, which nearly every query has, makes the pair + // coprime already: no gcd, and no divisions by it. + pa = ca; + pb = cb; + s = 1; + } else { + int64_t g = gcd(ca, cb); + pa = ca / g; + pb = cb / g; + s = g; + } + if (pa < 0 || (pa == 0 && pb < 0)) { + pa = -pa; + pb = -pb; + s = -s; + } + return true; +} + +// Solve a bound on (d * v) for a bound on v, where v is an integer. Rounding +// inwards at both ends is what makes this exact: a plain interval divide +// floors the low end where it should ceil, losing the last integer whenever d +// doesn't divide it (a bound of >= -7 on 8 * v is >= 0 on v, not >= -1). +ConstantInterval solve_scaled_bound(const ConstantInterval &bound, int64_t d) { + internal_assert(d != 0); + + // Halide division is Euclidean: div_imp floors for d > 0 and ceils for + // d < 0, leaving a remainder r = v - q * d in [0, |d|). The division was + // inexact iff r != 0, i.e. iff q * d != v. + // + // q * d itself can overflow int64 when v is near the ends of the range and + // d doesn't divide it (e.g. v = INT64_MIN, d = 3 gives q * d = INT64_MIN - + // 1), and signed overflow is undefined behavior. So the comparison must be + // done in uint64_t, where arithmetic is defined to wrap mod 2^64. The + // wrapped difference v - q * d is then r mod 2^64, and since 0 <= r < 2^64, + // that is zero exactly when r is. + auto inexact = [=](int64_t v, int64_t q) { + return (uint64_t)q * (uint64_t)d != (uint64_t)v; + }; + auto round_up = [=](int64_t v) { + int64_t q = div_imp(v, d); + return q + (d > 0 && inexact(v, q)); + }; + auto round_down = [=](int64_t v) { + int64_t q = div_imp(v, d); + return q - (d < 0 && inexact(v, q)); + }; + + ConstantInterval result; + // A negative d swaps which end is which. + const bool flip = d < 0; + const bool lo_defined = flip ? bound.max_defined : bound.min_defined; + const bool hi_defined = flip ? bound.min_defined : bound.max_defined; + // -INT64_MIN doesn't fit, so INT64_MIN / -1 has no representable + // quotient. Leave that end open: at the top that's exact (every value is + // <= 2^63), and at the bottom it only forgets that the bound was + // unsatisfiable. + auto fits = [=](int64_t v) { + return !(d == -1 && v == INT64_MIN); + }; + const int64_t lo = flip ? bound.max : bound.min; + const int64_t hi = flip ? bound.min : bound.max; + if (lo_defined && fits(lo)) { + result.min_defined = true; + result.min = round_up(lo); + } + if (hi_defined && fits(hi)) { + result.max_defined = true; + result.max = round_down(hi); + } + return result; +} + +} // namespace + +void Simplify::ScopedFact::learn_difference(const Expr &a, const Expr &b, + const ConstantInterval &diff, bool invert) { + // Differences are only meaningful where they can't wrap. + if (!simplify->no_overflow_int(a.type()) || a.type() != b.type()) { + return; + } + + const BaseExprNode *pa = a.get(), *pb = b.get(); + int64_t coeff_a = 1, coeff_b = 1, offset = 0, denom = 1; + ConstantInterval err(0, 0); + if ((pa->node_type >= IRNodeType::Add && pa->node_type <= IRNodeType::Div) || + (pb->node_type >= IRNodeType::Add && pb->node_type <= IRNodeType::Div)) { + peel_affine_terms(pa, pb, coeff_a, coeff_b, offset, denom, err); + } + + // denom * (a - b) == (coeff_a * pa - coeff_b * pb) + offset + err, so + // solve for the peeled quantity: scale the given bound up by denom and + // take back the offset and the remainder any peeled division discarded. + ConstantInterval peeled = diff * denom - offset - err; + + int64_t prim_a, prim_b, scale; + if (!reduce_affine_coeffs(coeff_a, coeff_b, prim_a, prim_b, scale)) { + // Both coefficients vanished (something peeled down to 0 * ...). + return; + } + + if (invert) { + // Only a single point is representable, and only on a lattice point: + // off-lattice, no integer primitive quantity could have hit it anyway. + if (!peeled.is_single_point() || peeled.min % scale != 0) { + return; + } + } + + ConstantInterval primitive_bound = solve_scaled_bound(peeled, scale); + + if (simplify->difference_heads.empty()) { + simplify->difference_heads.assign(Simplify::difference_buckets, -1); + } + int32_t &head = simplify->difference_heads[Simplify::difference_bucket(Simplify::difference_key(pa->hash, pb->hash))]; + simplify->known_bounds.push_back( + Simplify::KnownBound{Expr(pa), Expr(pb), primitive_bound, invert, prim_a, prim_b, pa->hash, pb->hash, head}); + head = (int32_t)simplify->known_bounds.size() - 1; +} + void Simplify::ScopedFact::learn_false(const Expr &fact) { + // Facts must already be simplified, so that they are stored in the same + // form the simplifier produces when it visits them. It never produces + // > or >=. + internal_assert(!fact.as() && !fact.as()) + << "learn_false expects a simplified fact: " << fact << "\n"; + + // Record what this says about the difference between the two sides. And, + // Not, and the tag intrinsic are handled by the recursion below instead. + if (const LT *lt = fact.as()) { + // !(a < b) -> a - b >= 0 + learn_difference(lt->a, lt->b, ConstantInterval::bounded_below(0), false); + } else if (const LE *le = fact.as()) { + // !(a <= b) -> a - b >= 1 + learn_difference(le->a, le->b, ConstantInterval::bounded_below(1), false); + } else if (const EQ *eq = fact.as()) { + // !(a == b) -> a - b is anything but zero + learn_difference(eq->a, eq->b, ConstantInterval::single_point(0), true); + } else if (const NE *ne = fact.as()) { + // !(a != b) -> a - b == 0 + learn_difference(ne->a, ne->b, ConstantInterval::single_point(0), false); + } + Simplify::VarInfo info; info.old_uses = info.new_uses = 0; if (const Variable *v = fact.as()) { @@ -172,6 +465,28 @@ void Simplify::ScopedFact::learn_lower_bound(const Variable *v, int64_t val) { } void Simplify::ScopedFact::learn_true(const Expr &fact) { + // Facts must already be simplified, so that they are stored in the same + // form the simplifier produces when it visits them. It never produces + // > or >=. + internal_assert(!fact.as() && !fact.as()) + << "learn_true expects a simplified fact: " << fact << "\n"; + + // Record what this says about the difference between the two sides. And, + // Not, and the tag intrinsic are handled by the recursion below instead. + if (const LT *lt = fact.as()) { + // a < b -> a - b <= -1 + learn_difference(lt->a, lt->b, ConstantInterval::bounded_above(-1), false); + } else if (const LE *le = fact.as()) { + // a <= b -> a - b <= 0 + learn_difference(le->a, le->b, ConstantInterval::bounded_above(0), false); + } else if (const EQ *eq = fact.as()) { + // a == b -> a - b == 0 + learn_difference(eq->a, eq->b, ConstantInterval::single_point(0), false); + } else if (const NE *ne = fact.as()) { + // a != b -> a - b is anything but zero + learn_difference(ne->a, ne->b, ConstantInterval::single_point(0), true); + } + Simplify::VarInfo info; info.old_uses = info.new_uses = 0; if (const Variable *v = fact.as()) { @@ -345,16 +660,56 @@ void Simplify::ScopedFact::learn_true(const Expr &fact) { } namespace { +// Is a boolean Expr known to be true or false? Facts are stored in the same +// form the simplifier itself produces, so a comparison has to be canonicalized +// the same way before looking it up. +std::optional lookup_fact(const Expr &e, + const std::set &truths, + const std::set &falsehoods) { + if (const Not *n = e.as()) { + auto known = lookup_fact(n->a, truths, falsehoods); + return known ? std::make_optional(!*known) : known; + } else if (const GT *gt = e.as()) { + return lookup_fact(gt->b < gt->a, truths, falsehoods); + } else if (const GE *ge = e.as()) { + return lookup_fact(!(ge->a < ge->b), truths, falsehoods); + } + + if (truths.count(e)) { + return true; + } else if (falsehoods.count(e)) { + return false; + } + + // A comparison may also be settled by the other strictness of the same + // comparison, in either direction. + if (const LT *lt = e.as()) { + // a < b is implied by !(b <= a), and ruled out by b <= a and by b < a. + if (falsehoods.count(lt->b <= lt->a)) { + return true; + } else if (truths.count(lt->b <= lt->a) || truths.count(lt->b < lt->a)) { + return false; + } + } else if (const LE *le = e.as()) { + // a <= b is implied by a < b and by !(b < a), and ruled out by b < a. + if (truths.count(le->a < le->b) || falsehoods.count(le->b < le->a)) { + return true; + } else if (truths.count(le->b < le->a)) { + return false; + } + } + + return std::nullopt; +} + template T substitute_facts_impl(const T &t, const std::set &truths, const std::set &falsehoods) { return mutate_with(t, [&](auto *self, const Expr &e) { if (e.type().is_bool()) { - if (truths.count(e)) { - return make_one(e.type()); - } else if (falsehoods.count(e)) { - return make_zero(e.type()); + if (auto known = lookup_fact(e, truths, falsehoods)) { + return *known ? make_one(e.type()) : make_zero(e.type()); } } return self->mutate_base(e); @@ -370,13 +725,263 @@ Stmt Simplify::ScopedFact::substitute_facts(const Stmt &s) { return substitute_facts_impl(s, truths, falsehoods); } +namespace { + +// Intersect acc with d, reporting whether the result would be empty rather than +// constructing it. make_intersection asserts on an empty result, and empty means +// the facts contradict each other, which means this code is unreachable. We +// don't try to exploit that here; we just decline to tighten any further. +bool intersect_if_nonempty(ConstantInterval &acc, const ConstantInterval &d) { + ConstantInterval result = acc; + if (d.min_defined && (!result.min_defined || d.min > result.min)) { + result.min = d.min; + result.min_defined = true; + } + if (d.max_defined && (!result.max_defined || d.max < result.max)) { + result.max = d.max; + result.max_defined = true; + } + if (result.min_defined && result.max_defined && result.min > result.max) { + return false; + } + acc = result; + return true; +} + +// What the shape of the two sides says about (a - b) on its own, with no facts +// involved: a min is at most either of its operands, and a max is at least +// either of them. Only the immediate operands are inspected, so this stays a +// couple of pointer comparisons rather than a search. +ConstantInterval structural_difference(const BaseExprNode *a, const BaseExprNode *b) { + ConstantInterval result; + + // Same restriction as learning a fact: a difference only means what we take + // it to mean for integers that don't wrap. It keeps floats, where a NaN + // makes even min(p, q) <= p false, out of it too. + if (!(a->type.is_int() && a->type.bits() >= 32) || a->type != b->type) { + return result; + } + + auto is_operand_of = [](const BaseExprNode *e, const BaseExprNode *node) { + if (node->node_type == IRNodeType::Min) { + const Min *m = (const Min *)node; + return equal(*m->a.get(), *e) || equal(*m->b.get(), *e); + } else if (node->node_type == IRNodeType::Max) { + const Max *m = (const Max *)node; + return equal(*m->a.get(), *e) || equal(*m->b.get(), *e); + } + return false; + }; + + // min(p, q) - b <= 0 and max(p, q) - b >= 0, when b is one of the operands. + if (a->node_type == IRNodeType::Min && is_operand_of(b, a)) { + result = ConstantInterval::bounded_above(0); + } else if (a->node_type == IRNodeType::Max && is_operand_of(b, a)) { + result = ConstantInterval::bounded_below(0); + } else if (b->node_type == IRNodeType::Min && is_operand_of(a, b)) { + // a - min(p, q) >= 0 + result = ConstantInterval::bounded_below(0); + } else if (b->node_type == IRNodeType::Max && is_operand_of(a, b)) { + result = ConstantInterval::bounded_above(0); + } + + return result; +} + +} // namespace + +ConstantInterval Simplify::known_difference(const BaseExprNode *a, const BaseExprNode *b) { + return known_linear_difference(a, 1, b, 1); +} + +ConstantInterval Simplify::known_linear_difference(const BaseExprNode *a, int64_t ca, + const BaseExprNode *b, int64_t cb) { + ConstantInterval result; + + // Canonicalize the query the way facts are canonicalized when learned. + // ca/cb seed the coefficients: they already apply to the unpeeled a/b (a + // matched WildConst, say), so peeling can't discover them itself. + int64_t coeff_a = ca, coeff_b = cb, offset = 0, denom = 1; + ConstantInterval err(0, 0); + if ((a->node_type >= IRNodeType::Add && a->node_type <= IRNodeType::Div) || + (b->node_type >= IRNodeType::Add && b->node_type <= IRNodeType::Div)) { + peel_affine_terms(a, b, coeff_a, coeff_b, offset, denom, err); + } + + if (coeff_a == coeff_b && equal(*a, *b)) { + result = ConstantInterval::single_point(0); + } else if (a->node_type == IRNodeType::IntImm && b->node_type == IRNodeType::IntImm) { + // Two constants need no facts to compare. + int64_t va = ((const IntImm *)a)->value, vb = ((const IntImm *)b)->value; + int64_t ta, tb, diff; + if (mul_with_overflow(64, coeff_a, va, &ta) && + mul_with_overflow(64, coeff_b, vb, &tb) && + sub_with_overflow(64, ta, tb, &diff)) { + result = ConstantInterval::single_point(diff); + } + } else if (coeff_a == 1 && coeff_b == 1) { + // The structural heuristic is about (a - b) alone; it doesn't + // generalize to a scaled combination. + intersect_if_nonempty(result, structural_difference(a, b)); + } + + if (!result.is_single_point() && !known_bounds.empty()) { + // Only the records in this pair's bucket can be about it. The chain + // runs newest first. Almost every query finds the bucket empty, so + // look before doing anything else. + const uint32_t fa = a->hash, fb = b->hash; + const int32_t head = difference_heads[difference_bucket(difference_key(fa, fb))]; + int64_t prim_a, prim_b, scale; + if (head >= 0 && reduce_affine_coeffs(coeff_a, coeff_b, prim_a, prim_b, scale)) { + // A hole only bites once the ends are known, so collect and apply + // them below. There are hardly ever any. + constexpr int max_holes = 4; + int64_t holes[max_holes]; + int num_holes = 0; + + for (int32_t i = head; i >= 0; i = known_bounds[i].next) { + const KnownBound &kb = known_bounds[i]; + // Hashes first: a record about another pair costs two + // integer compares, not a walk over two Exprs. + const bool same_order = (fa == kb.hash_a && fb == kb.hash_b); + const bool swapped = (fa == kb.hash_b && fb == kb.hash_a); + if (!same_order && !swapped) { + continue; + } + + ConstantInterval d; + if (same_order && equal(*a, *kb.a.get()) && equal(*b, *kb.b.get()) && + prim_a == kb.coeff_a && prim_b == kb.coeff_b) { + d = kb.diff * scale; + } else if (swapped && equal(*a, *kb.b.get()) && equal(*b, *kb.a.get())) { + // The fact runs the other way. Reduce (coeff_b, + // coeff_a) -- this query in the fact's operand order -- + // then negate to flip back. + int64_t sw_prim_a, sw_prim_b, sw_scale; + if (reduce_affine_coeffs(coeff_b, coeff_a, sw_prim_a, sw_prim_b, sw_scale) && + sw_prim_a == kb.coeff_a && sw_prim_b == kb.coeff_b) { + d = -(kb.diff * sw_scale); + } else { + continue; + } + } else { + continue; + } + + if (kb.invert) { + if (num_holes < max_holes) { + holes[num_holes++] = d.min; + } + } else if (!intersect_if_nonempty(result, d)) { + break; + } + } + + for (int i = 0; i < num_holes; i++) { + const int64_t hole = holes[i]; + // A point only narrows the bounds from an end, and only if + // something survives: a hole swallowing the interval means the + // facts contradict and the code is unreachable. Say nothing + // rather than hand back a backwards interval. + if (result.min_defined && result.max_defined && + result.min == hole && result.max == hole) { + continue; + } + int64_t new_min, new_max; + if (result.min_defined && result.min == hole && + add_with_overflow(64, hole, 1, &new_min)) { + result.min = new_min; + } + if (result.max_defined && result.max == hole && + sub_with_overflow(64, hole, 1, &new_max)) { + result.max = new_max; + } + } + } + } + + if (!result.min_defined && !result.max_defined) { + // Nothing is known, and no canonicalization changes that. + return result; + } + + // Undo the canonicalization. + result = solve_scaled_bound(result + offset + err, denom); + + return result; +} + +bool Simplify::known_min_diff(const BaseExprNode *a, int64_t ca, const BaseExprNode *b, int64_t cb, int64_t *result) { + ConstantInterval bounds = known_linear_difference(a, ca, b, cb); + if (bounds.min_defined) { + *result = bounds.min; + return true; + } + return false; +} + +bool Simplify::known_max_diff(const BaseExprNode *a, int64_t ca, const BaseExprNode *b, int64_t cb, int64_t *result) { + ConstantInterval bounds = known_linear_difference(a, ca, b, cb); + if (bounds.max_defined) { + *result = bounds.max; + return true; + } + return false; +} + +bool Simplify::is_known_true(const Expr &e) { + if (truths.empty() && falsehoods.empty()) { + return false; + } + auto known = lookup_fact(e, truths, falsehoods); + return known && *known; +} + +Expr Simplify::simplify_can_prove_condition(const Expr &e) { + if (can_prove_depth >= max_can_prove_depth) { + // Too deep to safely recurse into the full simplifier. The only thing + // the caller does with the result is check whether it is the literal + // constant true, and nothing here can fold a compound expression (an + // And of two known-true operands stays an unfolded And, not true) -- + // that folding is exactly the recursive work we're declining to do. + // So a substitute_facts tree walk can't prove anything a direct + // lookup of the condition itself couldn't already: skip the walk. + if (is_known_true(e)) { + return const_true(e.type().lanes(), nullptr); + } + return e; + } + ScopedValue guard(can_prove_depth, can_prove_depth + 1); + return mutate(substitute_facts(e), nullptr); +} + +Expr Simplify::substitute_facts(const Expr &e) { + if (truths.empty() && falsehoods.empty()) { + return e; + } + return substitute_facts_impl(e, truths, falsehoods); +} + Simplify::ScopedFact::~ScopedFact() { + if (!simplify) { + // Moved from; the object that took over owns the cleanup. + return; + } for (const auto *v : pop_list) { simplify->var_info.pop(v->name); } for (const auto *v : bounds_pop_list) { simplify->bounds_and_alignment_info.pop(v->name); } + // Unchain the records this scope pushed, newest first, which puts each + // bucket's head back to what it was. A scope that ends after an enclosing + // one (the assumptions of the public simplify() go in a vector) finds its + // records already gone. + while (simplify->known_bounds.size() > known_bounds_size) { + const KnownBound &kb = simplify->known_bounds.back(); + simplify->difference_heads[Simplify::difference_bucket(Simplify::difference_key(kb.hash_a, kb.hash_b))] = kb.next; + simplify->known_bounds.pop_back(); + } for (const auto &e : truths) { simplify->truths.erase(e); } diff --git a/src/Simplify.h b/src/Simplify.h index 64459a43c44b..2034fda91866 100644 --- a/src/Simplify.h +++ b/src/Simplify.h @@ -18,7 +18,8 @@ namespace Internal { * rearranging, etc. Simplifies across let statements, so must not be called on * stmts with dangling or repeated variable names. Can optionally be passed * known bounds of any variables, known alignment properties, and any other - * Exprs that should be assumed to be true. + * Exprs that should be assumed to be true. The assumptions must already be + * simplified. */ // @{ Stmt simplify(const Stmt &, diff --git a/src/Simplify_Add.cpp b/src/Simplify_Add.cpp index a07ad1b4464b..72716f222416 100644 --- a/src/Simplify_Add.cpp +++ b/src/Simplify_Add.cpp @@ -10,6 +10,14 @@ Expr Simplify::visit(const Add *op, ExprInfo *info) { if (info) { info->bounds = a_info.bounds + b_info.bounds; + // a + b is the difference (a - (-b)), so the facts can say more about + // it than the two sides do apart. + if (has_facts() && no_overflow_int(op->type)) { + ConstantInterval d = known_linear_difference(a.get(), 1, b.get(), -1); + if (d.min_defined || d.max_defined) { + info->bounds = ConstantInterval::make_intersection(info->bounds, d); + } + } info->alignment = a_info.alignment + b_info.alignment; info->cast_to(op->type); info->trim_bounds_using_alignment(); diff --git a/src/Simplify_Div.cpp b/src/Simplify_Div.cpp index 4098f6f027e7..89685a8ae72b 100644 --- a/src/Simplify_Div.cpp +++ b/src/Simplify_Div.cpp @@ -84,6 +84,25 @@ Expr Simplify::visit(const Div *op, ExprInfo *info) { rewrite(select(x, c0, c1) / c2, select(x, fold(c0 / c2), fold(c1 / c2))) || (!op->type.is_float() && rewrite(x / x, select(x == 0, 0, 1))) || + + // Facts learned higher up may say which side of a max or min + // survives the division. Test them before the rewrites below, which + // would destroy the form. + // + // For c0 > 0 and floor division, x >= y/c0 iff c0*x - y >= 1 - c0, + // and x <= y/c0 iff c0*x - y <= 0. The c0 on x isn't in x's own IR, + // so linear_{min,max}_diff take it explicitly. Learning peels + // divisions too, so a fact spelled either way round (c0*x < y or + // x < y/c0) is stored in the multiplied-out form asked for here. + // Unlike can_prove this only looks facts up, never building an Expr + // and so never recursing back into the simplifier. + (no_overflow(op->type) && has_facts() && + (rewrite(max(x * c0, y) / c0, x, c0 > 0 && linear_min_diff(x, c0, y, 1, this) >= fold(1 - c0)) || + rewrite(max(y, x * c0) / c0, x, c0 > 0 && linear_min_diff(x, c0, y, 1, this) >= fold(1 - c0)) || + rewrite(min(x * c0, y) / c0, x, c0 > 0 && linear_max_diff(x, c0, y, 1, this) <= 0) || + rewrite(min(y, x * c0) / c0, x, c0 > 0 && linear_max_diff(x, c0, y, 1, this) <= 0) || + false)) || + (no_overflow(op->type) && // Fold repeated division (rewrite((x / c0) / c2, x / fold(c0 * c2), c0 > 0 && c2 > 0 && !overflows(c0 * c2)) || diff --git a/src/Simplify_Internal.h b/src/Simplify_Internal.h index 94a50bebc644..35f865cf786e 100644 --- a/src/Simplify_Internal.h +++ b/src/Simplify_Internal.h @@ -441,30 +441,157 @@ class Simplify : public VariadicVisitor { std::set truths, falsehoods; + /** What we know about an affine difference between a pair of Exprs. Every + * comparison we learn from becomes a statement about (coeff_a * a - + * coeff_b * b), with a and b peeled down to base terms and (coeff_a, + * coeff_b) the coprime, sign-canonical coefficient pair (see + * peel_affine_terms). For a plain comparison both are 1: a < b puts it at + * most -1, !(a < b) at least 0, a == b exactly 0. The complement of a + * half-line is a half-line, so only a negated equality fails to be an + * interval, and that is a single point removed -- hence invert. */ + struct KnownBound { + Expr a, b; + ConstantInterval diff; + // If set, the combination is known *not* to lie in diff, which is + // then always a single point. Only a != b produces one of these. + bool invert = false; + int64_t coeff_a = 1, coeff_b = 1; + // The hashes of a and b, kept here so that walking a chain of records + // reads only the records, not the nodes they point to. + uint32_t hash_a = 0, hash_b = 0; + // The next record in the same bucket, or -1 at the end of the chain. + int32_t next = -1; + }; + std::vector known_bounds; + + // The records are chained per bucket of their pair key, so that a query + // walks the few records that could be about its pair rather than the whole + // table. A pipeline's asserts alone can put hundreds of facts in scope for + // its entire body. The heads are allocated on the first record, so a + // simplifier that never learns a difference, which is most of them, pays + // nothing for them. + static constexpr int difference_buckets = 256; + std::vector difference_heads; + + /** Everything the facts tell us about (a - b), without building any IR. + * The arguments are borrowed, so this is safe to call with the raw nodes a + * rewrite rule has bound to its wildcards. */ + ConstantInterval known_difference(const BaseExprNode *a, const BaseExprNode *b); + + /** As known_difference, but for the linear combination (ca * a - cb * b). + * For rules holding a constant multiplier (a matched WildConst, say) that + * sits outside a or b's own IR, where peeling can't find it. */ + ConstantInterval known_linear_difference(const BaseExprNode *a, int64_t ca, + const BaseExprNode *b, int64_t cb); + + // Helpers over the above, for use as rewrite rule predicates. They return + // false when nothing is known, so such a rule simply doesn't fire. + bool known_min_diff(const BaseExprNode *a, int64_t ca, const BaseExprNode *b, int64_t cb, int64_t *result); + bool known_max_diff(const BaseExprNode *a, int64_t ca, const BaseExprNode *b, int64_t cb, int64_t *result); + + // How deeply are we nested inside the conditions of can_prove predicates? + // Proving such a condition recursively invokes the simplifier on it, so a + // rule whose left-hand side also matches something built while proving its + // own predicate recurses without bound. Bound it. + // + // The work grows sharply with this limit -- on an adversarial nest of + // min(x, y) - min(z, w) it is roughly 0.02s at 1 or 2, 0.11s at 3 and 0.72s + // at 4 -- while no rule needs the depth: instrumenting every correctness + // test shows the deepest nesting any of them reaches is one. So this is + // already a level of headroom over anything observed. + int can_prove_depth = 0; + static constexpr int max_can_prove_depth = 2; + + // Is there anything a min_diff/max_diff predicate could look up? Used to + // gate rules whose predicates are only ever provable from facts learned + // higher up in the IR, so that we don't pay for them in the common case. + bool has_facts() const { + return !truths.empty() || !falsehoods.empty(); + } + + // Is there anything a min_diff or max_diff predicate could look up? Only a + // comparison of non-overflowing integers leaves a record here, so this is + // strictly narrower than has_facts: a boolean fact, or a fact about a type + // that can wrap, satisfies that one while leaving this table empty. Rules + // that ask about differences must gate on this, or they spend a lookup on + // a table that cannot answer. + bool has_difference_facts() const { + return !known_bounds.empty(); + } + + // Symmetric key for a pair. Xoring two equal hashes gives zero whatever + // they were, so key that case by the hash itself rather than letting every + // pair of equal-hashing operands share the one bit. + HALIDE_ALWAYS_INLINE + static uint32_t difference_key(uint32_t fa, uint32_t fb) { + return fa == fb ? fa * 0x9e3779b9u : (fa ^ fb); + } + + // The bucket of a pair key. It has to come from mixed bits rather than + // from the bottom of the key: an Expr's hash carries its node type in the + // low bits, so indexing by those puts every pair of the same two kinds in + // one bucket. + HALIDE_ALWAYS_INLINE + static uint32_t difference_bucket(uint32_t key) { + constexpr int bits = 8; + static_assert(difference_buckets == (1 << bits)); + return (key * 0x9e3779b9u) >> (32 - bits); + } + + // Replace exprs known to be truths or falsehoods with const_true or + // const_false. Used to inject everything currently known into the + // conditions of can_prove predicates in rewrite rules. + Expr substitute_facts(const Expr &e); + + // Simplify the condition of a can_prove predicate in a rewrite rule, using + // everything currently known. + Expr simplify_can_prove_condition(const Expr &e); + + // Is a boolean Expr already known to be true? Unlike can_prove this only + // looks the condition up in the facts, without simplifying anything. + bool is_known_true(const Expr &e); + struct ScopedFact { Simplify *simplify; std::vector pop_list; std::vector bounds_pop_list; std::set truths, falsehoods; + // Everything in the simplifier's known_bounds from this index on was + // pushed by this scope, and is popped again when it ends. + size_t known_bounds_size = 0; void learn_false(const Expr &fact); void learn_true(const Expr &fact); void learn_upper_bound(const Variable *v, int64_t val); void learn_lower_bound(const Variable *v, int64_t val); + // Record what a comparison says about the difference between its sides. + void learn_difference(const Expr &a, const Expr &b, const ConstantInterval &diff, bool invert); // Replace exprs known to be truths or falsehoods with const_true or const_false. Expr substitute_facts(const Expr &e); Stmt substitute_facts(const Stmt &s); ScopedFact(Simplify *s) - : simplify(s) { + : simplify(s), known_bounds_size(s->known_bounds.size()) { } ~ScopedFact(); // allow move but not copy ScopedFact(const ScopedFact &that) = delete; - ScopedFact(ScopedFact &&that) = default; + // Not defaulted: the moved-from object must not undo anything in its + // destructor. The containers below would be empty after a move and so + // would be harmless, but known_bounds_size would survive and truncate + // away the facts this scope had just learned. + ScopedFact(ScopedFact &&that) noexcept + : simplify(that.simplify), + pop_list(std::move(that.pop_list)), + bounds_pop_list(std::move(that.bounds_pop_list)), + truths(std::move(that.truths)), + falsehoods(std::move(that.falsehoods)), + known_bounds_size(that.known_bounds_size) { + that.simplify = nullptr; + } }; // Tell the simplifier to learn from and exploit a boolean diff --git a/src/Simplify_Max.cpp b/src/Simplify_Max.cpp index 88d3ce2cbf5e..df395b3ca178 100644 --- a/src/Simplify_Max.cpp +++ b/src/Simplify_Max.cpp @@ -58,6 +58,24 @@ Expr Simplify::visit(const Max *op, ExprInfo *info) { std::swap(a_info, b_info); } + // Or when the facts learned higher up in the IR, or the shapes of the two + // sides, order them. One lookup answers for both sides. + { + ConstantInterval d = known_difference(a.get(), b.get()); + if (d.min_defined && d.min >= 0) { + if (info) { + *info = a_info; + } + return a; + } + if (d.max_defined && d.max <= 0) { + if (info) { + *info = b_info; + } + return b; + } + } + int lanes = op->type.lanes(); auto rewrite = IRMatcher::rewriter(IRMatcher::max(a, b), op->type); diff --git a/src/Simplify_Min.cpp b/src/Simplify_Min.cpp index 5203a0c14166..c5609aa9f2d2 100644 --- a/src/Simplify_Min.cpp +++ b/src/Simplify_Min.cpp @@ -57,6 +57,24 @@ Expr Simplify::visit(const Min *op, ExprInfo *info) { std::swap(a_info, b_info); } + // Or when the facts learned higher up in the IR, or the shapes of the two + // sides, order them. One lookup answers for both sides. + { + ConstantInterval d = known_difference(a.get(), b.get()); + if (d.max_defined && d.max <= 0) { + if (info) { + *info = a_info; + } + return a; + } + if (d.min_defined && d.min >= 0) { + if (info) { + *info = b_info; + } + return b; + } + } + int lanes = op->type.lanes(); auto rewrite = IRMatcher::rewriter(IRMatcher::min(a, b), op->type); diff --git a/src/Simplify_Sub.cpp b/src/Simplify_Sub.cpp index cb3fc09a0a13..ba5add29e9cf 100644 --- a/src/Simplify_Sub.cpp +++ b/src/Simplify_Sub.cpp @@ -13,6 +13,15 @@ Expr Simplify::visit(const Sub *op, ExprInfo *info) { // cancellation rule that exploits that should always // remutate to recalculate the bounds. info->bounds = a_info.bounds - b_info.bounds; + // Except for the correlation the facts know about: x - 5 * y is + // positive wherever 5 * y < x was learned, however little the bounds + // of x and of 5 * y say on their own. + if (has_facts() && no_overflow_int(op->type)) { + ConstantInterval d = known_difference(a.get(), b.get()); + if (d.min_defined || d.max_defined) { + info->bounds = ConstantInterval::make_intersection(info->bounds, d); + } + } info->alignment = a_info.alignment - b_info.alignment; info->cast_to(op->type); info->trim_bounds_using_alignment(); diff --git a/test/correctness/CMakeLists.txt b/test/correctness/CMakeLists.txt index 9eda1bf6dd82..d6f45240b5fb 100644 --- a/test/correctness/CMakeLists.txt +++ b/test/correctness/CMakeLists.txt @@ -62,6 +62,7 @@ tests( bounds_internal.cpp bounds_of_abs.cpp bounds_of_cast.cpp + bounds_of_compound_index_from_where.cpp bounds_of_func.cpp bounds_of_monotonic_math.cpp bounds_of_multiply.cpp @@ -453,6 +454,7 @@ tests( side_effects.cpp simplified_away_embedded_image.cpp simplify.cpp + simplify_region_bound_regression.cpp skip_stages.cpp skip_stages_cse_bug.cpp skip_stages_external_array_functions.cpp diff --git a/test/correctness/bounds_of_compound_index_from_where.cpp b/test/correctness/bounds_of_compound_index_from_where.cpp new file mode 100644 index 000000000000..80c4723149ba --- /dev/null +++ b/test/correctness/bounds_of_compound_index_from_where.cpp @@ -0,0 +1,372 @@ +#include "Halide.h" +#include + +using namespace Halide; + +// A where clause on an index made of several variables bounds both the index +// as a whole and each of the variables it is made of. Bounds inference must +// use it both ways: for the input read at that index, and for any input read +// at just one of its variables. + +// A max pool in the style of hannk's: the input is read through a clamp, and +// the reduction domain is restricted by a where clause that says the same +// thing. The simplifier may remove the clamp as redundant, so bounds inference +// must get the region of the input that is read from the where clause alone. +// Otherwise the pipeline asks for input it doesn't need and fails its own +// bounds check. +int hannk_style_max_pool() { + ImageParam input(UInt(8), 2, "input"); + Param stride_x("stride_x"), stride_y("stride_y"); + Param filter_width("filter_width"), filter_height("filter_height"); + + Var x("x"), y("y"); + + Expr min_x = input.dim(0).min(); + Expr max_x = input.dim(0).max(); + Expr min_y = input.dim(1).min(); + Expr max_y = input.dim(1).max(); + + Func input_bounded("input_bounded"); + input_bounded(x, y) = input(clamp(x, min_x, max_x), clamp(y, min_y, max_y)); + + RDom r(0, filter_width, 0, filter_height); + Expr x_rx = x * stride_x + r.x; + Expr y_ry = y * stride_y + r.y; + r.where(min_x <= x_rx && x_rx <= max_x && min_y <= y_ry && y_ry <= max_y); + + Func maximum("maximum"); + maximum(x, y) = cast(0); + maximum(x, y) = max(maximum(x, y), input_bounded(x_rx, y_ry)); + + Func output("output"); + output(x, y) = maximum(x, y); + + // The input starts at 1, and the first output pixel's window starts at 0, + // off the edge of it, the way padding does. + const int w = 16, h = 12, fw = 3, fh = 3; + Buffer input_buf(w, h); + input_buf.set_min(1, 1); + input_buf.for_each_element([&](int i, int j) { + input_buf(i, j) = (uint8_t)(i * 7 + j * 13); + }); + + input.set(input_buf); + stride_x.set(1); + stride_y.set(1); + filter_width.set(fw); + filter_height.set(fh); + + Buffer out = output.realize({w, h}); + + for (int j = 0; j < h; j++) { + for (int i = 0; i < w; i++) { + uint8_t correct = 0; + for (int ry = 0; ry < fh; ry++) { + for (int rx = 0; rx < fw; rx++) { + int xi = i + rx, yi = j + ry; + if (xi >= 1 && xi <= w && yi >= 1 && yi <= h) { + correct = std::max(correct, input_buf(xi, yi)); + } + } + } + if (out(i, j) != correct) { + printf("out(%d, %d) = %d instead of %d\n", i, j, out(i, j), correct); + return 1; + } + } + } + return 0; +} + +// As above, but the window samples every other input pixel, and the where +// clause bounds the scaled reduction variable, 2 * r.x, rather than the index +// itself. The two only agree once the coefficient is taken into account. +int scaled_rdom_max_pool() { + ImageParam input(UInt(8), 1, "input"); + Param stride_x("stride_x"), filter_width("filter_width"); + + Var x("x"); + + Expr min_x = input.dim(0).min(); + Expr max_x = input.dim(0).max(); + + Func input_bounded("input_bounded"); + input_bounded(x) = input(clamp(x, min_x, max_x)); + + RDom r(0, filter_width); + Expr x_rx = x * stride_x + 2 * r.x; + r.where(min_x <= x_rx && 2 * r.x <= max_x - x * stride_x); + + Func maximum("maximum"); + maximum(x) = cast(0); + maximum(x) = max(maximum(x), input_bounded(x_rx)); + + Func output("output"); + output(x) = maximum(x); + + const int w = 16, fw = 3; + Buffer input_buf(w); + input_buf.set_min(1); + input_buf.for_each_element([&](int i) { + input_buf(i) = (uint8_t)(i * 7); + }); + + input.set(input_buf); + stride_x.set(1); + filter_width.set(fw); + + Buffer out = output.realize({w}); + + for (int i = 0; i < w; i++) { + uint8_t correct = 0; + for (int rx = 0; rx < fw; rx++) { + int xi = i + 2 * rx; + if (xi >= 1 && xi <= w) { + correct = std::max(correct, input_buf(xi)); + } + } + if (out(i) != correct) { + printf("out(%d) = %d instead of %d\n", i, out(i), correct); + return 1; + } + } + return 0; +} + +// As above, but the input is read at half the position the where clause +// bounds, so the clamped index is only a part of what the condition mentions. +int halved_index_max_pool() { + ImageParam input(UInt(8), 1, "input"); + Param stride_x("stride_x"), filter_width("filter_width"); + + Var x("x"); + + Expr min_x = input.dim(0).min(); + Expr max_x = input.dim(0).max(); + + Func input_bounded("input_bounded"); + input_bounded(x) = input(clamp(x, min_x, max_x)); + + RDom r(0, filter_width); + Expr x_rx = x * stride_x + r.x; + r.where(2 * min_x <= x_rx && x_rx <= 2 * max_x + 1); + + Func maximum("maximum"); + maximum(x) = cast(0); + maximum(x) = max(maximum(x), input_bounded(x_rx / 2)); + + Func output("output"); + output(x) = maximum(x); + + const int w = 16, fw = 3; + Buffer input_buf(w); + input_buf.set_min(1); + input_buf.for_each_element([&](int i) { + input_buf(i) = (uint8_t)(i * 7); + }); + + input.set(input_buf); + stride_x.set(1); + filter_width.set(fw); + + Buffer out = output.realize({w}); + + for (int i = 0; i < w; i++) { + uint8_t correct = 0; + for (int rx = 0; rx < fw; rx++) { + int xi = i + rx; + if (xi >= 2 && xi <= 2 * w + 1) { + correct = std::max(correct, input_buf(xi / 2)); + } + } + if (out(i) != correct) { + printf("out(%d) = %d instead of %d\n", i, out(i), correct); + return 1; + } + } + return 0; +} + +// A where clause on the sum of two reduction variables. Bounds inference names +// the sum so that the clause can bound it as a whole, but the clause must keep +// bounding each variable on its own too, for the input that is read at just +// one of them. +int where_on_sum_bounds_each_term() { + ImageParam in(Int(32), 1, "in"), in2(Int(32), 1, "in2"); + + Var x("x"); + RDom r(0, 200, 0, 200); + r.where(r.x + r.y < 50); + + Func f("f"); + f(x) = 0; + f(x) += in(r.x) + in2(r.x + r.y); + + f.infer_input_bounds({4}); + for (ImageParam *p : {&in, &in2}) { + Buffer<> b = p->get(); + if (b.dim(0).min() != 0 || b.dim(0).extent() != 50) { + printf("%s is required over [%d, %d] instead of [0, 49]\n", + p->name().c_str(), b.dim(0).min(), b.dim(0).max()); + return 1; + } + } + + Buffer in_buf(50), in2_buf(50); + in_buf.for_each_element([&](int i) { in_buf(i) = i; }); + in2_buf.for_each_element([&](int i) { in2_buf(i) = 1000 * i; }); + in.set(in_buf); + in2.set(in2_buf); + + Buffer out = f.realize({4}); + int correct = 0; + for (int ry = 0; ry < 200; ry++) { + for (int rx = 0; rx < 200; rx++) { + if (rx + ry < 50) { + correct += in_buf(rx) + in2_buf(rx + ry); + } + } + } + for (int i = 0; i < 4; i++) { + if (out(i) != correct) { + printf("out(%d) = %d instead of %d\n", i, out(i), correct); + return 1; + } + } + return 0; +} + +// The where clause bounds x*2 + r.x + r.y, and the clamp on the input read at +// it is redundant, but a second input is read at just the x*2 + r.x part. That +// part must still be bounded by the clause: through what it says about r.x +// and x on their own, the index stays within [0, 32] for a 16 wide output. +int where_on_compound_index_bounds_its_parts() { + ImageParam input(Int(32), 1, "input"), in2(Int(32), 1, "in2"); + Param min_x("min_x"), max_x("max_x"); + + Var x("x"); + RDom r(0, 200, 0, 200); + Expr xr = x * 2 + r.x; + r.where(min_x <= xr + r.y && xr + r.y <= max_x); + + Func m("m"); + m(x) = 0; + m(x) += input(clamp(xr + r.y, min_x, max_x)) + in2(xr); + + Func output("output"); + output(x) = m(x); + + const int w = 16; + min_x.set(1); + max_x.set(w); + output.infer_input_bounds({w}); + { + Buffer<> b = input.get(); + if (b.dim(0).min() != 1 || b.dim(0).max() != w) { + printf("input is required over [%d, %d] instead of [1, %d]\n", + b.dim(0).min(), b.dim(0).max(), w); + return 1; + } + b = in2.get(); + if (b.dim(0).min() < 0 || b.dim(0).max() > 32) { + printf("in2 is required over [%d, %d], which is wider than [0, 32]\n", + b.dim(0).min(), b.dim(0).max()); + return 1; + } + } + + Buffer input_buf(w), in2_buf(33); + input_buf.set_min(1); + input_buf.for_each_element([&](int i) { input_buf(i) = i; }); + in2_buf.for_each_element([&](int i) { in2_buf(i) = 1000 * i; }); + input.set(input_buf); + in2.set(in2_buf); + + Buffer out = output.realize({w}); + for (int i = 0; i < w; i++) { + int correct = 0; + for (int ry = 0; ry < 200; ry++) { + for (int rx = 0; rx < 200; rx++) { + int xi = i * 2 + rx; + if (1 <= xi + ry && xi + ry <= w) { + correct += input_buf(xi + ry) + in2_buf(xi); + } + } + } + if (out(i) != correct) { + printf("out(%d) = %d instead of %d\n", i, out(i), correct); + return 1; + } + } + return 0; +} + +// The where clause bounds a select inside the index. A select can't be solved +// for the variable it contains, so the clause only helps once the select +// itself is named. +int where_on_select_part() { + ImageParam in(Int(32), 1, "in"); + Param p("p"), fw("fw"); + + Var x("x"); + RDom r(0, fw); + Expr base = select(p > 0, x * 2, x); + Expr idx = base * 2 + r.x; + r.where(0 <= base && base <= 10); + + Func f("f"); + f(x) = 0; + f(x) += in(idx); + + p.set(1); + fw.set(3); + f.infer_input_bounds({100}); + Buffer<> b = in.get(); + if (b.dim(0).min() != 0 || b.dim(0).max() != 22) { + printf("in is required over [%d, %d] instead of [0, 22]\n", + b.dim(0).min(), b.dim(0).max()); + return 1; + } + + Buffer in_buf(23); + in_buf.for_each_element([&](int i) { in_buf(i) = i * 7; }); + in.set(in_buf); + Buffer out = f.realize({100}); + for (int i = 0; i < 100; i++) { + int correct = 0; + for (int rx = 0; rx < 3; rx++) { + int base = i * 2; + if (0 <= base && base <= 10) { + correct += in_buf(base * 2 + rx); + } + } + if (out(i) != correct) { + printf("out(%d) = %d instead of %d\n", i, out(i), correct); + return 1; + } + } + return 0; +} + +int main(int argc, char **argv) { + if (hannk_style_max_pool() != 0) { + return 1; + } + if (scaled_rdom_max_pool() != 0) { + return 1; + } + if (halved_index_max_pool() != 0) { + return 1; + } + if (where_on_sum_bounds_each_term() != 0) { + return 1; + } + if (where_on_compound_index_bounds_its_parts() != 0) { + return 1; + } + if (where_on_select_part() != 0) { + return 1; + } + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/simplify.cpp b/test/correctness/simplify.cpp index 981a00e6f0ee..61c2426431c3 100644 --- a/test/correctness/simplify.cpp +++ b/test/correctness/simplify.cpp @@ -2377,6 +2377,24 @@ void check_invariant() { } } +void check_with_assumptions(const Expr &a, const Expr &b, const std::vector &assumptions) { + // The simplifier only learns facts in the form it produces itself, so + // simplify the assumptions first, as the compiler would have. + std::vector simplified_assumptions; + for (const Expr &e : assumptions) { + simplified_assumptions.push_back(simplify(e)); + } + Expr simpler = simplify(a, Scope(), Scope(), simplified_assumptions); + if (!equal(simpler, b)) { + std::cerr + << "\nSimplification failure:\n" + << "Input: " << a << "\n" + << "Output: " << simpler << "\n" + << "Expected output: " << b << "\n"; + abort(); + } +} + void check_unreachable() { Var x("x"), y("y"); @@ -2405,6 +2423,207 @@ void check_unreachable() { Evaluate::make(0)); } +void check_facts() { + Expr x = Var("x"), y = Var("y"), z = Var("z"); + + // Several assumptions at once each leave a record. The public simplify() + // keeps their scopes in a vector, so they end in the order they were made + // rather than in reverse, and each must pop only what is still its own. + check_with_assumptions(min(x, z) + max(x, y), x + y, {x < y, x < z}); + + // A fact stated in any comparison direction should let the simplifier pick + // the winning side of a max or min. + check_with_assumptions(max(x, y), x, {x > y}); + check_with_assumptions(max(x, y), x, {y < x}); + check_with_assumptions(max(x, y), y, {x < y}); + check_with_assumptions(max(x, y), y, {y > x}); + check_with_assumptions(min(x, y), y, {x > y}); + check_with_assumptions(min(x, y), x, {x < y}); + + // A non-strict fact is enough to pick a side of a max or min, and a strict + // fact implies the non-strict one. + check_with_assumptions(max(x, y), x, {x >= y}); + check_with_assumptions(max(x, y), y, {x <= y}); + check_with_assumptions(min(x, y), x, {x <= y}); + check_with_assumptions(min(x, y), y, {x >= y}); + + // Facts about compound expressions work too. + check_with_assumptions(max(x + z, y * 3), x + z, {x + z > y * 3}); + check_with_assumptions(max(max(x, y), z), z, {max(x, y) < z}); + + // Both branches of an if learn from the condition, in opposite directions. + check(IfThenElse::make(x < y, not_no_op(max(x, y)), not_no_op(max(x, y))), + IfThenElse::make(x < y, not_no_op(y), not_no_op(x))); + + // A fact only applies where it holds. + check(Block::make(not_no_op(max(x, y)), + IfThenElse::make(x < y, not_no_op(max(x, y)))), + Block::make(not_no_op(max(x, y)), + IfThenElse::make(x < y, not_no_op(y)))); + + // A division can cancel a multiplication inside a max or min when we know + // which side wins after the division. + check_with_assumptions(max(x * 8, y) / 8, x, {x >= y / 8}); + check_with_assumptions(max(y, x * 8) / 8, x, {x >= y / 8}); + check_with_assumptions(min(x * 8, y) / 8, x, {x <= y / 8}); + check_with_assumptions(min(y, x * 8) / 8, x, {x <= y / 8}); + + // The direction in which a fact is stated doesn't matter, on either side: + // both the facts and the conditions of can_prove predicates are looked up + // in the same canonical form. + check_with_assumptions(max(x * 8, y) / 8, x, {y / 8 <= x}); + check_with_assumptions(max(x * 8, y) / 8, x, {!(x < y / 8)}); + check_with_assumptions(min(x * 8, y) / 8, x, {y / 8 >= x}); + + // A strict fact settles a non-strict predicate too. + check_with_assumptions(max(x * 8, y) / 8, x, {x > y / 8}); + check_with_assumptions(min(x * 8, y) / 8, x, {x < y / 8}); + + // Coefficients are reduced to a coprime pair and a scale, so a fact and a + // query that differ only by an overall factor meet. + check_with_assumptions(max(x * 2, y), y, {x * 4 <= y * 2}); + check_with_assumptions(min(x * 3, y * 6), x * 3, {x <= y * 2}); + + // Offsets peeled from under a factor are scaled by it on the way out, so + // the two spellings of the same affine term meet. + check_with_assumptions(max((x + 3) * 4, y), y, {x * 4 + 12 <= y}); + + // A fact over a division orders the multiplied-out terms, and vice versa: + // for c > 0, x <= y / c iff c * x <= y. + check_with_assumptions(min(x * 8, y), x * 8, {x <= y / 8}); + check_with_assumptions(max(x, y / 8), y / 8, {x * 8 <= y}); + + // Large divisors peel too. + const int big = 1 << 24; + check_with_assumptions(min(x * big, y), x * big, {x <= y / big}); + check_with_assumptions(max(x * big, y) / big, x, {x >= y / big}); + + // Denominators too large to represent stop the peeling, but a fact still + // settles a query spelled the same way. Within one term: + { + Expr x64 = Variable::make(Int(64), "x64"), y64 = Variable::make(Int(64), "y64"); + Expr c40 = make_const(Int(64), int64_t(1) << 40); + Expr c30 = make_const(Int(64), int64_t(1) << 30); + Expr e = (x64 / c40) / c30; + check_with_assumptions(min(e, y64), e, {e <= y64}); + // And across the two sides of the difference: + Expr a = x64 / c40, b = y64 / c30; + check_with_assumptions(min(a, b), a, {a <= b}); + check_with_assumptions(min(a, b), min(a, b), {a <= b + 1}); + + // A bound that solves to INT64_MIN / -1. The else branch learns + // !(-3 * x < y / 3 - k), which peels to -9 * x - y >= INT64_MIN, i.e. + // 9 * x + y <= 2^63, which says nothing. It must not turn into a bound + // on 9 * x + y (x = y = 0 satisfies the fact but not the query). + Expr k = make_const(Int(64), (int64_t)3074457345618258602); // (2^63 - 2) / 3 + Expr cond = x64 * make_const(Int(64), -3) < y64 / make_const(Int(64), 3) - k; + cond = simplify(cond); + Expr q = x64 * make_const(Int(64), 9) + y64 <= make_const(Int(64), -5); + Stmt s = IfThenElse::make(cond, not_no_op(x64), not_no_op(q)); + check(s, s); + } + + // Divisions on both sides are peeled over a common denominator. The + // remainders they discard cost a little precision, but x <= y is a wide + // enough margin to survive it. + check_with_assumptions(max(x / 4, y / 4), y / 4, {x <= y}); + + // Coefficients that don't reduce to the same coprime pair don't match: + // 2 * x <= y says nothing about 3 * x against y. + check_with_assumptions(min(x * 3, y), min(x * 3, y), {x * 2 <= y}); + + // Comparing a sum or difference against a constant is a comparison + // between its two terms, so a fact about the pair settles it. + check_with_assumptions(max(x - y * 5, 0), x - y * 5, {x > y * 5}); + check_with_assumptions(min(x - y * 5, 0), 0, {x > y * 5}); + check_with_assumptions(max(x + y * 5, 0), y * 5 + x, {x > y * -5}); + + // The constant comes back as an offset, so it needn't be zero. + check_with_assumptions(max(x - y * 5, 3), x - y * 5, {x > y * 5 + 3}); + check_with_assumptions(max(x - y * 5, -2), x - y * 5, {x >= y * 5}); + check_with_assumptions(min(x - y * 5, 7), 7, {x > y * 5 + 7}); + + // The margin has to actually cover the constant. + check_with_assumptions(max(x - y * 5, 3), max(x - y * 5, 3), {x > y * 5}); + + // A difference only means what we take it to mean where the type cannot + // wrap. Given x >= y + 5 over uint8, y = 253 makes y + 5 equal 2, so x = 10 + // satisfies it while sitting far below y: ordering the min from that would + // pick the wrong side. Only the types whose overflow is undefined, and so + // may be assumed not to happen, are eligible. + for (Type t : {UInt(8), Int(8), Int(16), UInt(32)}) { + Expr a = Variable::make(t, "wrap_a"); + Expr b = Variable::make(t, "wrap_b"); + check_with_assumptions(min(a, b), min(a, b), {a >= b + cast(t, 5)}); + check_with_assumptions(max(a, b), max(a, b), {a >= b + cast(t, 5)}); + } + for (Type t : {Int(32), Int(64)}) { + Expr a = Variable::make(t, "wrap_a"); + Expr b = Variable::make(t, "wrap_b"); + check_with_assumptions(min(a, b), b, {a >= b + cast(t, 5)}); + check_with_assumptions(max(a, b), a, {a >= b + cast(t, 5)}); + } + + // A min is at most either of its operands and a max is at least either of + // them, which needs no facts at all. That only bounds the difference on one + // side, but knowing the two are unequal removes the endpoint, and the two + // together settle a comparison that neither settles alone. + check_with_assumptions(max(min(x, y) + 1, x), x, {min(x, y) != x}); + check_with_assumptions(min(max(x, y) - 1, x), x, {max(x, y) != x}); + + // Neither ingredient is enough by itself: without the inequality the + // difference could still be zero, and without the shape there is no bound + // for the inequality to tighten. + check_with_assumptions(max(min(x, y) + 1, x), max(min(x, y) + 1, x), {z < z + 1}); + check_with_assumptions(max(y + 1, x), max(y + 1, x), {y != x}); + + // Deeply nested mins and maxes must not make the work of proving the + // predicates of the rules above blow up. + Expr nest = x; + for (int i = 0; i < 24; i++) { + nest = min(max(nest + i, y - i), z * i); + } + // The result isn't interesting; what matters is that we get one at all. + (void)simplify(nest, Scope(), Scope(), {x < y}); + + // can_prove-based rules (unlike the min_diff/max_diff ones above) recursively + // invoke the simplifier on their own predicate, and that predicate can be + // a freshly built expression rather than a piece of the original IR (e.g. + // min(x, y) - min(z, w) -> y - w, can_prove(x - y == z - w)) constructs a + // brand new subtraction). If the operands are themselves unsimplified + // instances of the same shape, this recurses; the depth limit must bound + // the work rather than let it explode. + Expr deep = min(Var("da"), Var("db")) - min(Var("dc"), Var("dd")); + for (int i = 0; i < 10; i++) { + Expr y = Var("dy" + std::to_string(i)); + Expr z = Var("dz" + std::to_string(i)); + Expr w = Var("dw" + std::to_string(i)); + deep = min(deep, y) - min(z, w); + } + (void)simplify(deep); + + // Constant offsets are peeled off both the facts and the queries, so a fact + // stated about a shifted operand still settles a predicate about the + // unshifted one, in either direction. + check_with_assumptions(max(x, y), y, {x + 1 <= y}); + check_with_assumptions(max(x, y), x, {y <= x + 0}); + check_with_assumptions(max(x + 3, y), y, {x + 4 <= y}); + check_with_assumptions(min(x, y), x, {x + 1 <= y}); + + // But an offset that leaves the order undetermined still doesn't fire. + check_with_assumptions(max(x, y), max(x, y), {x <= y + 1}); + + // Without the fact, the division stays put. + check(max(x * 8, y) / 8, max(x * 8, y) / 8); + + // Facts that don't strictly order the operands don't fire these rules. + check_with_assumptions(max(x, y), max(x, y), {x != y}); + + // A fact over a division is learned multiplied out, so it orders the + // operands as well as one spelled that way: x < y / 8 means 8 * x <= y - 8. + check_with_assumptions(max(x * 8, y) / 8, y / 8, {x < y / 8}); +} + int main(int argc, char **argv) { check_invariant(); check_casts(); @@ -2417,6 +2636,7 @@ int main(int argc, char **argv) { check_bitwise(); check_lets(); check_unreachable(); + check_facts(); // Miscellaneous cases that don't fit into one of the categories above. Expr x = Var("x"), y = Var("y"); diff --git a/test/correctness/simplify_region_bound_regression.cpp b/test/correctness/simplify_region_bound_regression.cpp new file mode 100644 index 000000000000..82159de9e7ff --- /dev/null +++ b/test/correctness/simplify_region_bound_regression.cpp @@ -0,0 +1,164 @@ +// A fact in scope must not change how an unrelated expression folds. +// +// Where this comes from +// --------------------- +// Lowering apps/local_laplacian at pyramid_levels=6 widens two allocations +// relative to main. Reduced, the cause is below. +// +// The pyramid bound is let-bound once, near the top of the IR, where no fact is +// in scope: +// +// let gPyramid4.s0.v1.max = max((E + 14)/16, (((E + 30)/32)*2) + 2) +// +// and a copy of that same expression appears further in, inside an if, where +// facts *are* in scope. Peeling constant divisors lets known_difference settle +// (E + 14)/16 against (((E + 30)/32)*2) + 2 by arithmetic alone -- but the +// min/max rules that consult it are gated on has_difference_facts(), so the +// copy inside the if folds to its second arm and the let-bound one does not. +// +// The two copies are then no longer spelled the same, and that is what matters +// here: a LetStmt-bound variable is never substituted into the body (single-use +// inlining is disabled for Stmt bodies), and a written-out copy of a let's value +// is not recognised and rewritten into the variable. So +// +// max(, min(Q, gPyramid4.s0.v1.max)) +// +// can only collapse while both copies are still the same expression. Once one +// of them folds and the other does not, nothing collapses, and the surviving +// min/max widens the allocation. +// +// On main neither copy folds -- there is no such prover at all -- so the two +// stay equal and everything downstream works. Firing in *some* scopes is what +// breaks it, not firing as such. +// +// The fold itself is correct; only its inconsistency is the problem. + +#include "Halide.h" +#include +#include + +using namespace Halide; +using namespace Halide::Internal; + +namespace { + +int failures = 0; + +// Simplify e with no fact in scope, and again with an unrelated fact in scope. +// The two must agree. +Stmt sink(const Expr &x) { + return Evaluate::make(Call::make(Int(32), "sink", {x}, Call::Extern)); +} + +void check_fact_independent(const Expr &e) { + Expr p = Variable::make(Int(32), "p"); + Expr q = Variable::make(Int(32), "q"); + + Stmt bare = simplify(sink(e)); + Stmt guarded = simplify(IfThenElse::make(p < q, sink(e))); + + // Dig the simplified expression back out of each. + const Evaluate *bare_eval = bare.as(); + const IfThenElse *ite = guarded.as(); + if (!bare_eval || !ite) { + std::cerr << "test is malformed\n"; + failures++; + return; + } + const Evaluate *guarded_eval = ite->then_case.as(); + if (!guarded_eval) { + std::cerr << "test is malformed\n"; + failures++; + return; + } + + Expr without = bare_eval->value.as()->args[0]; + Expr with = guarded_eval->value.as()->args[0]; + + if (!equal(without, with)) { + std::cerr << "\nA fact in scope changed an unrelated simplification:\n" + << "Input: " << e << "\n" + << "Without a fact: " << without << "\n" + << "With p < q: " << with << "\n"; + failures++; + } +} + +// Simplify a statement and compare against the expected result. +void check_stmt(const Stmt &a, const Stmt &b) { + std::cerr << "----\n"; + std::cerr << "Input:\n" + << a << "\n"; + Stmt simpler = simplify(a); + std::cerr << "\nOutput:\n" + << simpler << "\n"; + if (!equal(simpler, b)) { + std::cerr << "\nSimplification failure:\n" + << "Expected output:\n" + << b << "\n"; + failures++; + } else { + std::cerr << "Ok!\n"; + } +} + +} // namespace + +int main() { + Expr e = Variable::make(Int(32), "e"); // output.extent.1 + Expr m = Variable::make(Int(32), "m"); // output.min.1 + Expr E = e + m; + + // gPyramid4.s0.v1.max, verbatim. Folds to its second arm with a fact in + // scope and stays put without one. + check_fact_independent(max((E + 14) / 16, (((E + 30) / 32) * 2) + 2)); + + // The v0 counterpart, /8 and /16. + check_fact_independent(max((E + 6) / 8, (((E + 14) / 16) * 2) + 2)); + + // gPyramid4.s0.v1.min, the same shape for min. + check_fact_independent(min((m + -15) / 16, (((m + -31) / 32) * 2) + -1)); + + // The consequence. In the lowered IR the two copies of the bound are not + // in the same scope: one is written out, the other is a reference to a + // LetStmt-bound variable whose binding sits outside the `if`, where no fact + // is in scope. + // + // let V = max(A1, A2) <- no fact in scope here + // if (p < q) { <- facts in scope here + // sink(max(max(A1, A2), min(Q, V))) + // } + // + // A LetStmt-bound variable is never substituted into the body -- single-use + // inlining is disabled for Stmt bodies -- and a written-out copy of a let's + // value is not recognised and rewritten into the variable. So V stays + // opaque, and the collapse to a single term needs both copies to still be + // spelled the same. If the fold fires on the written-out copy but not on + // V's value, they differ and nothing collapses. + Expr p = Variable::make(Int(32), "p"); + Expr q = Variable::make(Int(32), "q"); + Expr Q = Variable::make(Int(32), "Q"); + Expr V = Variable::make(Int(32), "V"); + Expr bound = max((E + 14) / 16, (((E + 30) / 32) * 2) + 2); + Expr folded = (((E + 30) / 32) * 2) + 2; + + // Both copies written out, without any facts, in one scope: symmetric, and it collapses. + check_stmt(sink(max(bound, min(Q, bound))), + sink(folded)); + + // Both copies written out, in one scope: symmetric, and it collapses. + check_stmt(IfThenElse::make(p < q, sink(max(bound, min(Q, bound)))), + IfThenElse::make(p < q, sink(folded))); + + // One copy behind a LetStmt bound outside the `if`: the shape from the IR. + check_stmt(LetStmt::make("V", bound, + IfThenElse::make(p < q, sink(max(bound, min(Q, V))))), + IfThenElse::make(p < q, sink(folded))); + + if (failures) { + printf("\n%d check(s) failed\n", failures); + return 1; + } + printf("Success!\n"); + return 0; +}