blob: 45a561610e907fa86e53ca1152b84df755908d41 [file]
/*
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
/*!
* \file int_set.cc
* \brief The integer set functions
*/
#include <tvm/arith/int_set.h>
#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/runtime/logging.h>
#include <tvm/tirx/expr.h>
#include <tvm/tirx/expr_functor.h>
#include <tvm/tirx/op.h>
#include <algorithm>
#include <unordered_map>
#include <unordered_set>
#include <utility>
#include "constraint_extract.h"
#include "int_operator.h"
#include "interval_set.h"
#include "pattern_match.h"
namespace tvm {
namespace arith {
using tirx::is_one;
using tirx::is_zero;
using tirx::MakeConst;
TVM_FFI_STATIC_INIT_BLOCK() { IntervalSetNode::RegisterReflection(); }
PrimExpr SymbolicLimits::pos_inf_ = tirx::PrimVar("pos_inf", PrimType::Int(64));
PrimExpr SymbolicLimits::neg_inf_ = tirx::PrimVar("neg_inf", PrimType::Int(64));
IntervalSet::IntervalSet(PrimExpr min_value, PrimExpr max_value) {
auto node = ffi::make_object<IntervalSetNode>();
node->min_value = std::move(min_value);
node->max_value = std::move(max_value);
data_ = std::move(node);
}
IntervalSet MakeIntervalSet(PrimExpr min_value, PrimExpr max_value) {
return IntervalSet(min_value, max_value);
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("arith.IntervalSet", MakeIntervalSet);
}
IntervalSet Intersect(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b) {
PrimExpr max_value = min(a->max_value, b->max_value);
PrimExpr min_value = max(a->min_value, b->min_value);
PrimType max_ty = max_value.ty();
PrimType min_ty = min_value.ty();
if (max_ty.MatchesCode(DLDataTypeCode::kDLInt, DLDataTypeCode::kDLUInt) &&
min_ty.MatchesCode(DLDataTypeCode::kDLInt, DLDataTypeCode::kDLUInt) &&
analyzer->CanProve(max_value < min_value)) {
return IntervalSet::Empty();
} else {
return IntervalSet(min_value, max_value);
}
}
IntervalSet Union(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b) {
if (a->IsEmpty()) return b;
if (b->IsEmpty()) return a;
PrimExpr max_value = max(a->max_value, b->max_value);
PrimExpr min_value = min(a->min_value, b->min_value);
return IntervalSet(min_value, max_value);
}
// type traits
template <typename OP>
struct is_logical_op {
static const bool value = false;
};
#define TVM_DECLARE_LOGICAL_OP(OP) \
template <> \
struct is_logical_op<tirx::OP> { \
static const bool value = true; \
};
TVM_DECLARE_LOGICAL_OP(And);
TVM_DECLARE_LOGICAL_OP(Or);
TVM_DECLARE_LOGICAL_OP(EQ);
TVM_DECLARE_LOGICAL_OP(NE);
TVM_DECLARE_LOGICAL_OP(GE);
TVM_DECLARE_LOGICAL_OP(GT);
TVM_DECLARE_LOGICAL_OP(LE);
TVM_DECLARE_LOGICAL_OP(LT);
TVM_DECLARE_LOGICAL_OP(Not);
/*!
* \brief Combine two interval set under arithmetic operations.
* \param analyzer The analyzer for simplification and proving
* \param a The first interval set
* \param b The second interval set
* \param op The operation node, used to extract dtype and other properties
* \note this can possibly relax the set.
*/
template <typename Op, typename OpNode>
inline IntervalSet Combine(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b, const OpNode* op) {
PrimType dtype = op->ty.template as_or_throw<PrimType>();
if (a->IsSinglePoint() && b->IsSinglePoint()) {
PrimExpr expr;
if (auto res = TryConstFold<Op>(a->min_value, b->min_value)) {
expr = res.value();
} else {
expr = Op(a->min_value, b->min_value);
}
return IntervalSet::SinglePoint(expr);
}
if (is_logical_op<Op>::value) {
return IntervalSet(IntImm(dtype, 0), IntImm(dtype, 1));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (a->IsEverything()) return a;
if (b->IsEverything()) return b;
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::Add>(AnalyzerObj* analyer, IntervalSet a, IntervalSet b,
const tirx::AddNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(a->min_value + b->min_value);
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
PrimExpr min_value =
a->HasLowerBound() && b->HasLowerBound() ? a->min_value + b->min_value : neg_inf();
PrimExpr max_value =
a->HasUpperBound() && b->HasUpperBound() ? a->max_value + b->max_value : pos_inf();
return IntervalSet(min_value, max_value);
}
template <>
inline IntervalSet Combine<tirx::Sub>(AnalyzerObj* analyer, IntervalSet a, IntervalSet b,
const tirx::SubNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(a->min_value - b->min_value);
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
PrimExpr min_value =
a->HasLowerBound() && b->HasUpperBound() ? a->min_value - b->max_value : neg_inf();
PrimExpr max_value =
a->HasUpperBound() && b->HasLowerBound() ? a->max_value - b->min_value : pos_inf();
return IntervalSet(min_value, max_value);
}
template <>
inline IntervalSet Combine<tirx::Mul>(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b,
const tirx::MulNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(a->min_value * b->min_value);
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (a->IsSinglePoint()) {
std::swap(a, b);
}
if (b->IsSinglePoint()) {
if (is_zero(b->min_value)) return b;
if (is_one(b->min_value)) return a;
if (analyzer->CanProveGreaterEqual(b->min_value, 0)) {
PrimExpr min_value = a->HasLowerBound() ? a->min_value * b->min_value : neg_inf();
PrimExpr max_value = a->HasUpperBound() ? a->max_value * b->min_value : pos_inf();
return IntervalSet(min_value, max_value);
} else if (analyzer->CanProveGreaterEqual(-b->min_value, 1)) {
PrimExpr min_value = a->HasUpperBound() ? a->max_value * b->min_value : neg_inf();
PrimExpr max_value = a->HasLowerBound() ? a->min_value * b->min_value : pos_inf();
return IntervalSet(min_value, max_value);
} else if (a->HasUpperBound() && a->HasLowerBound()) {
using tirx::Select;
PrimExpr sign = b->min_value >= IntImm(b->min_value.ty().WithLanes(1), 0);
PrimExpr e1 = a->min_value * b->min_value;
PrimExpr e2 = a->max_value * b->min_value;
return IntervalSet(Select(sign, e1, e2), Select(sign, e2, e1));
}
}
DLOG(WARNING) << "Return Everything in CombineInterval Mul";
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::Div>(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b,
const tirx::DivNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(a->min_value / b->min_value);
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (b->IsSinglePoint()) {
if (is_zero(b->min_value)) {
TVM_FFI_THROW(InternalError) << "Divide by zero in CombineInterval Div";
}
if (is_one(b->min_value)) return a;
// no relaxation is needed in here due to set is inclusive
if (analyzer->CanProveGreaterEqual(b->min_value, 0)) {
PrimExpr min_value = a->HasLowerBound() ? a->min_value / b->min_value : neg_inf();
PrimExpr max_value = a->HasUpperBound() ? a->max_value / b->min_value : pos_inf();
return IntervalSet(min_value, max_value);
} else if (analyzer->CanProveGreaterEqual(-b->min_value, 1)) {
PrimExpr min_value = a->HasUpperBound() ? a->max_value / b->min_value : neg_inf();
PrimExpr max_value = a->HasLowerBound() ? a->min_value / b->min_value : pos_inf();
return IntervalSet(min_value, max_value);
} else if (a->HasUpperBound() && a->HasLowerBound()) {
using tirx::Select;
PrimExpr sign = b->min_value >= IntImm(b->min_value.ty().WithLanes(1), 0);
PrimExpr e1 = a->min_value / b->min_value;
PrimExpr e2 = a->max_value / b->min_value;
return IntervalSet(Select(sign, e1, e2), Select(sign, e2, e1));
}
}
DLOG(WARNING) << "Return Everything in CombineInterval Div";
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::Mod>(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b,
const tirx::ModNode* op) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(truncmod(a->min_value, b->min_value));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (b->IsSinglePoint()) {
const PrimExpr& divisor = b->min_value;
if (is_zero(divisor)) {
TVM_FFI_THROW(InternalError) << "Modular by zero in CombineInterval Mod";
}
// We need to add more bound constraints throughout the code.
// The logic below assumes a is non-negative, which usually
// is the case of our application.
// TODO(tqchen): add bound constraints for a.
if (analyzer->CanProveGreaterEqual(divisor, 0)) {
return IntervalSet(IntImm(divisor.ty(), 0), divisor - 1);
} else {
PrimExpr bound = abs(divisor) - 1;
return IntervalSet(-bound, bound);
}
}
DLOG(WARNING) << "Return Everything in CombineInterval Mod";
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::FloorDiv>(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b,
const tirx::FloorDivNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(floordiv(a->min_value, b->min_value));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (b->IsSinglePoint()) {
if (is_zero(b->min_value)) {
TVM_FFI_THROW(InternalError) << "Divide by zero in CombineInterval Div";
}
if (is_one(b->min_value)) return a;
// no relaxation is needed in here due to set is inclusive
if (analyzer->CanProveGreaterEqual(b->min_value, 0)) {
PrimExpr min_value = a->HasLowerBound() ? floordiv(a->min_value, b->min_value) : neg_inf();
PrimExpr max_value = a->HasUpperBound() ? floordiv(a->max_value, b->min_value) : pos_inf();
return IntervalSet(min_value, max_value);
} else if (analyzer->CanProveGreaterEqual(-b->min_value, 1)) {
PrimExpr min_value = a->HasUpperBound() ? floordiv(a->max_value, b->min_value) : neg_inf();
PrimExpr max_value = a->HasLowerBound() ? floordiv(a->min_value, b->min_value) : pos_inf();
return IntervalSet(min_value, max_value);
} else if (a->HasUpperBound() && a->HasLowerBound()) {
using tirx::Select;
PrimExpr sign = b->min_value >= IntImm(b->min_value.ty().WithLanes(1), 0);
PrimExpr e1 = floordiv(a->min_value, b->min_value);
PrimExpr e2 = floordiv(a->max_value, b->min_value);
return IntervalSet(Select(sign, e1, e2), Select(sign, e2, e1));
}
}
DLOG(WARNING) << "Return Everything in CombineInterval Div";
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::FloorMod>(AnalyzerObj* analyzer, IntervalSet a, IntervalSet b,
const tirx::FloorModNode* op) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(floormod(a->min_value, b->min_value));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
if (b->IsSinglePoint()) {
const PrimExpr& divisor = b->min_value;
if (is_zero(divisor)) {
TVM_FFI_THROW(InternalError) << "Modular by zero in CombineInterval Mod";
}
if (analyzer->CanProveGreaterEqual(divisor, 0)) {
if (divisor.as<tirx::IntImmNode>()) {
// a mod b = a - (a / b) * b if a_max / b == a_min / b
auto qmax = a->HasUpperBound() ? floordiv(a->max_value, divisor) : pos_inf();
auto qmin = a->HasLowerBound() ? floordiv(a->min_value, divisor) : neg_inf();
// We can compare +/- inf against each other, but cannot use
// operator== between the symbolic limits and an integer.
bool qmin_is_symbolic_limit = is_pos_inf(qmin) || is_neg_inf(qmin);
bool qmax_is_symbolic_limit = is_pos_inf(qmax) || is_neg_inf(qmax);
bool compatible_dtypes = qmin_is_symbolic_limit == qmax_is_symbolic_limit;
if (compatible_dtypes && analyzer->CanProve(qmax == qmin)) {
auto tmax = a->max_value - divisor * qmin;
auto tmin = a->min_value - divisor * qmin;
return IntervalSet(tmin, tmax);
}
}
// Enhanced: Use ModularSet analysis for better bounds
if (auto* div_imm = divisor.as<tirx::IntImmNode>()) {
int64_t div_val = div_imm->value;
// Analyze the modular properties of the dividend
ModularSet dividend_mod = analyzer->modular_set(op->a);
if (dividend_mod.defined() && dividend_mod->coeff > 0) {
// Calculate GCD of dividend coefficient and divisor
int64_t gcd = ZeroAwareGCD(dividend_mod->coeff, div_val);
if (gcd > 1 && div_val % gcd == 0) {
// The dividend is a multiple of gcd, and divisor is also a multiple of gcd
// So the result is also a multiple of gcd, with max value = (div_val/gcd - 1) * gcd
int64_t max_quotient = (div_val / gcd) - 1;
int64_t max_mod_result = max_quotient * gcd + (dividend_mod->base % gcd);
if (max_mod_result >= 0 && max_mod_result < div_val) {
PrimType result_ty = op->ty.as_or_throw<PrimType>();
return IntervalSet(IntImm(result_ty, 0), IntImm(result_ty, max_mod_result));
}
}
}
}
return IntervalSet(IntImm(divisor.ty(), 0), divisor - 1);
} else {
PrimExpr bound = abs(divisor) - 1;
return IntervalSet(-bound, bound);
}
}
DLOG(WARNING) << "Return Everything in CombineInterval Mod";
return IntervalSet::Everything();
}
template <>
inline IntervalSet Combine<tirx::Max>(AnalyzerObj* analzyer, IntervalSet a, IntervalSet b,
const tirx::MaxNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(max(a->min_value, b->min_value));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
return IntervalSet(max(a->min_value, b->min_value), max(a->max_value, b->max_value));
}
template <>
inline IntervalSet Combine<tirx::Min>(AnalyzerObj* analzyer, IntervalSet a, IntervalSet b,
const tirx::MinNode* /* op */) {
if (a->IsSinglePoint() && b->IsSinglePoint()) {
return IntervalSet::SinglePoint(min(a->min_value, b->min_value));
}
if (a->IsEmpty()) return a;
if (b->IsEmpty()) return b;
return IntervalSet(min(a->min_value, b->min_value), min(a->max_value, b->max_value));
}
// internal helper function to get an interval set
IntervalSet ToIntervalSet(IntSet set) {
if (auto node = set.as<IntervalSet>()) {
return node.value();
}
DLOG(INFO) << "cannot resolve int set " << set;
return IntervalSet::Everything();
}
using namespace tirx;
// Simplified version of int set evaluator that operates on IntervalSet
// We might use better set analysis in the future to replace the intervalset.
class IntervalSetEvaluator : public ExprFunctor<IntervalSet(const Expr&)> {
public:
IntervalSetEvaluator(AnalyzerObj* analyzer, const ffi::Map<Var, IntSet>& dom_map,
const std::vector<std::pair<Var, IntSet>>* dom_constraints = nullptr,
bool eval_vec = false)
: analyzer_(analyzer),
dom_map_(dom_map),
dom_constraints_(dom_constraints),
eval_vec_(eval_vec) {}
IntervalSet Eval(const PrimExpr& val) { return this->VisitExpr(val); }
// evaluate and relax the set
IntervalSet Eval(IntervalSet val) {
// avoid recursive indefinite recursive expansion.
if (static_cast<size_t>(recur_depth_) >= dom_map_.size()) return val;
++recur_depth_;
IntervalSet min_set = this->Eval(val->min_value);
IntervalSet max_set = this->Eval(val->max_value);
--recur_depth_;
return IntervalSet(min_set->min_value, max_set->max_value);
}
IntervalSet VisitExpr_(const IntImmNode* op) final {
return IntervalSet::SinglePoint(ffi::GetRef<PrimExpr>(op));
}
IntervalSet VisitExpr_(const VarNode* op) final {
Var var = ffi::GetRef<Var>(op);
auto it = dom_map_.find(var);
// Scoped constraints refine explicit relaxation domains. Variables that
// are absent from dom_map_ remain free parameters of the relaxed result.
if (it == dom_map_.end()) {
return IntervalSet::SinglePoint(var.as_or_throw<PrimExpr>());
}
ffi::Array<IntSet> values;
if (dom_constraints_) {
for (const auto& constraint : *dom_constraints_) {
if (var.same_as(constraint.first)) {
values.push_back(constraint.second);
}
}
}
values.push_back((*it).second);
IntSet intersection = [&]() {
if (values.size() == 1) {
return values.front();
} else {
return Intersect(values);
}
}();
IntervalSet res = ToIntervalSet(intersection);
if (res->min_value.same_as(var) && res->max_value.same_as(var)) {
return res;
}
// Recursively relax the mapped interval, since the domain bounds may
// themselves reference other variables that need to be relaxed.
//
// Memoize the fully-relaxed interval per variable, and guard against
// cyclic variable dependencies with an in-progress set. Without this,
// diamond-shaped variable dependencies (var a -> {b, c}, b -> {d, e}, ...)
// are re-expanded along every path: each level evaluates both the min and
// max sub-expressions, so the cost is exponential (2^depth) in the length
// of the variable dependency chain rather than linear.
auto memo_it = relax_memo_.find(op);
if (memo_it != relax_memo_.end()) {
return memo_it->second;
}
if (relax_in_progress_.count(op)) {
// Cyclic dependency among variable bounds: stop relaxing here to keep
// the recursion finite, keeping this variable symbolic.
return res;
}
relax_in_progress_.insert(op);
IntervalSet relaxed = Eval(res);
relax_in_progress_.erase(op);
relax_memo_[op] = relaxed;
return relaxed;
}
IntervalSet VisitExpr_(const AddNode* op) final { return VisitBinaryExpr_<Add>(op); }
IntervalSet VisitExpr_(const SubNode* op) final { return VisitBinaryExpr_<Sub>(op); }
IntervalSet VisitExpr_(const MulNode* op) final { return VisitBinaryExpr_<Mul>(op); }
IntervalSet VisitExpr_(const DivNode* op) final { return VisitBinaryExpr_<Div>(op); }
IntervalSet VisitExpr_(const ModNode* op) final { return VisitBinaryExpr_<Mod>(op); }
IntervalSet VisitExpr_(const FloorDivNode* op) final { return VisitBinaryExpr_<FloorDiv>(op); }
IntervalSet VisitExpr_(const FloorModNode* op) final { return VisitBinaryExpr_<FloorMod>(op); }
IntervalSet VisitExpr_(const MinNode* op) final { return VisitBinaryExpr_<Min>(op); }
IntervalSet VisitExpr_(const MaxNode* op) final { return VisitBinaryExpr_<Max>(op); }
IntervalSet VisitExpr_(const EQNode* op) final { return VisitBinaryExpr_<EQ>(op); }
IntervalSet VisitExpr_(const NENode* op) final { return VisitBinaryExpr_<NE>(op); }
IntervalSet VisitExpr_(const LTNode* op) final { return VisitBinaryExpr_<LT>(op); }
IntervalSet VisitExpr_(const LENode* op) final { return VisitBinaryExpr_<LE>(op); }
IntervalSet VisitExpr_(const GTNode* op) final { return VisitBinaryExpr_<GT>(op); }
IntervalSet VisitExpr_(const GENode* op) final { return VisitBinaryExpr_<GE>(op); }
IntervalSet VisitExpr_(const AndNode* op) final { return VisitBinaryExpr_<And>(op); }
IntervalSet VisitExpr_(const OrNode* op) final { return VisitBinaryExpr_<Or>(op); }
IntervalSet VisitExpr_(const RampNode* op) final {
TVM_FFI_ICHECK(eval_vec_);
IntervalSet base = Eval(op->base);
PVar<IntImm> stride;
if (stride.Match(op->stride)) {
PrimType t = op->base.ty();
int64_t vstride = stride.Eval()->value;
if (op->lanes->IsInstance<IntImmNode>()) {
int lanes = static_cast<int>(op->lanes.as_or_throw<IntImm>()->value);
if (vstride > 0) {
PrimExpr stride_expr = MakeConst(t, vstride * (lanes - 1));
auto add_op = tirx::Add(op->base, stride_expr);
auto add_node = add_op.as<tirx::AddNode>();
return Combine<Add>(analyzer_, base, IntervalSet(IntImm(t, 0), stride_expr), add_node);
} else {
PrimExpr stride_expr = MakeConst(t, vstride * (lanes - 1));
auto add_op = tirx::Add(op->base, stride_expr);
auto add_node = add_op.as<tirx::AddNode>();
return Combine<Add>(analyzer_, base, IntervalSet(stride_expr, IntImm(t, 0)), add_node);
}
} else { /* Scalable vector */
if (vstride > 0) {
auto add_op = tirx::Add(op->base, IntImm(t, 0));
auto add_node = add_op.as<tirx::AddNode>();
return Combine<Add>(analyzer_, base, IntervalSet(IntImm(t, 0), pos_inf()), add_node);
} else {
auto add_op = tirx::Add(op->base, IntImm(t, 0));
auto add_node = add_op.as<tirx::AddNode>();
return Combine<Add>(analyzer_, base, IntervalSet(neg_inf(), IntImm(t, 0)), add_node);
}
}
}
DLOG(WARNING) << "cannot evaluate set on expression " << ffi::GetRef<PrimExpr>(op);
return IntervalSet::Everything();
}
IntervalSet VisitExpr_(const BroadcastNode* op) final {
TVM_FFI_ICHECK(eval_vec_);
return VisitExpr(op->value);
}
IntervalSet VisitExpr_(const SelectNode* op) final {
IntervalSet true_set = this->Eval(op->true_value);
IntervalSet false_set = this->Eval(op->false_value);
return Union(analyzer_, false_set, true_set);
}
IntervalSet VisitExpr_(const CastNode* op) final {
IntervalSet value_set = this->Eval(op->value);
// short cut for the int set.
if (value_set->min_value.same_as(value_set->max_value)) {
if (value_set->IsEmpty()) return value_set;
return IntervalSet::SinglePoint(cast(op->ty.as_or_throw<PrimType>(), value_set->min_value));
}
PrimExpr min_value = value_set->HasLowerBound()
? cast(op->ty.as_or_throw<PrimType>(), value_set->min_value)
: neg_inf();
PrimExpr max_value = value_set->HasUpperBound()
? cast(op->ty.as_or_throw<PrimType>(), value_set->max_value)
: pos_inf();
return IntervalSet(min_value, max_value);
}
IntervalSet VisitExpr_(const BufferLoadNode* op) final {
PrimType op_ty = op->ty.as_or_throw<PrimType>();
if (!op_ty.MatchesCode(DLDataTypeCode::kDLInt, DLDataTypeCode::kDLUInt)) {
DLOG(WARNING) << "cannot evaluate set BufferLoad which loads from a " << op_ty->dtype
<< " buffer";
return IntervalSet::Everything();
}
// If the indices do not contain any variables to be relaxed, return the BufferLoad itself.
// Otherwise return `IntervalSet::everything()` since we have no knowledge on the buffer data.
for (const PrimExpr& index : op->indices) {
if (UsesVar(index, [dom_map = &this->dom_map_](const VarNode* var) {
return dom_map->find(ffi::GetRef<Var>(var)) != dom_map->end();
})) {
return IntervalSet::Everything();
}
}
return IntervalSet::SinglePoint(ffi::GetRef<PrimExpr>(op));
}
IntervalSet VisitExpr_(const CallNode* op) final {
if (op->op.same_as(tirx::builtin::vscale())) {
PrimExpr call = ffi::GetRef<Call>(op).as_or_throw<PrimExpr>();
return IntervalSet(call, call);
}
return IntervalSet::Everything();
}
IntervalSet VisitExprDefault_(const ffi::Object* op) final {
DLOG(WARNING) << "cannot evaluate set type " << op->GetTypeKey();
return IntervalSet::Everything();
}
private:
// whether set is exactly single point that equals value.
bool MatchPoint(const IntervalSet& set, const PrimExpr& value) const {
return set->min_value.same_as(value) && set->max_value.same_as(value);
}
template <typename TOp, typename T>
inline IntervalSet VisitBinaryExpr_(const T* op) {
static_assert(std::is_same<typename TOp::ContainerType, T>::value, "constraint");
IntervalSet a = this->Eval(op->a);
IntervalSet b = this->Eval(op->b);
if (MatchPoint(a, op->a) && MatchPoint(b, op->b)) {
return IntervalSet::SinglePoint(ffi::GetRef<PrimExpr>(op));
}
return Combine<TOp>(analyzer_, a, b, op);
}
// recursive depth
int recur_depth_{0};
// Memo of fully-relaxed interval sets per variable, to avoid exponential
// re-expansion of diamond-shaped variable dependencies.
std::unordered_map<const VarNode*, IntervalSet> relax_memo_;
// Variables currently being relaxed, used to break cyclic dependencies.
std::unordered_set<const VarNode*> relax_in_progress_;
// analyzer
AnalyzerObj* analyzer_;
const ffi::Map<Var, IntSet>& dom_map_;
const std::vector<std::pair<Var, IntSet>>* dom_constraints_;
bool eval_vec_{false};
};
class IntSetAnalyzer::Impl {
public:
explicit Impl(AnalyzerObj* analyzer) : analyzer_(analyzer) {}
IntSet Eval(const PrimExpr& expr, const ffi::Map<Var, IntSet>& dom_map) const {
return IntervalSetEvaluator(analyzer_, dom_map).Eval(expr);
}
IntSet Eval(const PrimExpr& expr) const {
return IntervalSetEvaluator(analyzer_, dom_map_, &dom_constraints_, true).Eval(expr);
}
void Bind(const Var& var, const Range& range, bool allow_override) {
Update(var, IntSet::FromRange(range), allow_override);
}
void Update(const Var& var, const IntSet& info, bool override_info);
void Bind(const Var& var, const PrimExpr& expr, bool override_info);
std::function<void()> EnterConstraint(const PrimExpr& constraint);
void CopyFrom(const Impl& other) {
dom_map_ = other.dom_map_;
dom_constraints_ = other.dom_constraints_;
}
private:
// Utility function to split a boolean condition into the domain
// bounds implied by that condition.
static std::vector<std::pair<Var, IntSet>> DetectBoundInfo(const PrimExpr& cond);
// The parent arith::Analyzer
AnalyzerObj* analyzer_;
// Map of variables to global variable bounds (e.g. loop iterator
// ranges)
ffi::Map<Var, IntSet> dom_map_;
// List of implicit scope-dependent bounds (e.g. inside the body of
// an if-statement). Maintained as a list of constraints, rather
// than as a `ffi::Map<Var,IntSet>`, to avoid computing an Intersection
// until required.
std::vector<std::pair<Var, IntSet>> dom_constraints_;
};
IntSetAnalyzer::IntSetAnalyzer(AnalyzerObj* parent) : impl_(new Impl(parent)) {}
IntSetAnalyzer::~IntSetAnalyzer() { delete impl_; }
void IntSetAnalyzer::CopyFrom(const IntSetAnalyzer& other) { impl_->CopyFrom(*other.impl_); }
IntSet IntSetAnalyzer::operator()(const PrimExpr& expr, const ffi::Map<Var, IntSet>& dom_map) {
return impl_->Eval(expr, dom_map);
}
IntSet IntSetAnalyzer::operator()(const PrimExpr& expr) { return impl_->Eval(expr); }
void IntSetAnalyzer::Update(const Var& var, const IntSet& info, bool allow_override) {
impl_->Update(var, info, allow_override);
}
void IntSetAnalyzer::Bind(const Var& var, const Range& range, bool allow_override) {
impl_->Bind(var, range, allow_override);
}
void IntSetAnalyzer::Impl::Update(const Var& var, const IntSet& info, bool can_override) {
if (!can_override) {
auto it = dom_map_.find(var);
if (it != dom_map_.end()) {
const IntSet& old_info = (*it).second;
TVM_FFI_ICHECK(ExprDeepEqual()(old_info.min(), info.min()))
<< "Trying to update var \'" << var << "\'"
<< " with a different minimum value: "
<< "original=" << old_info.min() << ", new=" << info.min();
TVM_FFI_ICHECK(ExprDeepEqual()(old_info.max(), info.max()))
<< "Trying to update var \'" << var << "\'"
<< " with a different maximum value: "
<< "original=" << old_info.max() << ", new=" << info.max();
}
}
dom_map_.Set(var, info);
}
void IntSetAnalyzer::Impl::Bind(const Var& var, const PrimExpr& expr, bool can_override) {
Update(var, Eval(expr), can_override);
}
std::vector<std::pair<Var, IntSet>> IntSetAnalyzer::Impl::DetectBoundInfo(
const PrimExpr& constraint) {
PVar<Var> x;
PVar<PrimExpr> limit;
std::vector<std::pair<Var, IntSet>> bounds;
for (const PrimExpr& subconstraint : ExtractConstraints(constraint)) {
if ((x <= limit).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(SymbolicLimits::neg_inf_, limit.Eval())});
} else if ((x < limit).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(SymbolicLimits::neg_inf_, limit.Eval() - 1)});
} else if ((x >= limit).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(limit.Eval(), SymbolicLimits::pos_inf_)});
} else if ((x > limit).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(limit.Eval() + 1, SymbolicLimits::pos_inf_)});
} else if ((x == limit).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::SinglePoint(limit.Eval())});
}
if ((limit >= x).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(SymbolicLimits::neg_inf_, limit.Eval())});
} else if ((limit > x).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(SymbolicLimits::neg_inf_, limit.Eval() - 1)});
} else if ((limit <= x).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(limit.Eval(), SymbolicLimits::pos_inf_)});
} else if ((limit < x).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::Interval(limit.Eval() + 1, SymbolicLimits::pos_inf_)});
} else if ((limit == x).Match(subconstraint)) {
bounds.push_back({x.Eval(), IntSet::SinglePoint(limit.Eval())});
}
}
return bounds;
}
std::function<void()> IntSetAnalyzer::EnterConstraint(const PrimExpr& constraint) {
return impl_->EnterConstraint(constraint);
}
std::function<void()> IntSetAnalyzer::Impl::EnterConstraint(const PrimExpr& constraint) {
auto bounds = DetectBoundInfo(constraint);
if (bounds.size() == 0) return nullptr;
size_t old_size = dom_constraints_.size();
dom_constraints_.insert(dom_constraints_.end(), bounds.begin(), bounds.end());
size_t new_size = dom_constraints_.size();
auto frecover = [old_size, new_size, this]() {
TVM_FFI_ICHECK_EQ(dom_constraints_.size(), new_size);
dom_constraints_.resize(old_size);
};
return frecover;
}
// Quickly adapt to IntSet interface
// TODO(tqchen): revisit IntSet interface as well.
Range IntSet::CoverRange(Range max_range) const {
IntSet temp;
Analyzer analyzer;
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
TVM_FFI_ICHECK(s_int != nullptr);
if (s_int->HasUpperBound() && s_int->HasLowerBound()) {
return Range::FromMinExtent(analyzer->Simplify(s_int->min_value),
analyzer->Simplify(s_int->max_value + 1 - s_int->min_value));
}
return max_range;
}
PrimExpr IntSet::min() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
TVM_FFI_ICHECK(s_int);
return s_int->min_value;
}
PrimExpr IntSet::max() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
TVM_FFI_ICHECK(s_int);
return s_int->max_value;
}
bool IntSet::IsNothing() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
return (s_int && s_int->IsEmpty());
}
bool IntSet::IsEverything() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
return (s_int && s_int->IsEverything());
}
bool IntSet::IsSinglePoint() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
return (s_int && s_int->IsSinglePoint());
}
bool IntSet::CanProveSinglePoint(const Analyzer& ana) const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
if (!s_int) return false;
if (s_int->IsSinglePoint()) return true;
return ana->CanProveEqual(s_int->min_value, s_int->max_value);
}
bool IntSet::CanProvePositive() const {
Analyzer analyzer;
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
return (s_int && is_positive_const(analyzer->Simplify(s_int->min_value)));
}
bool IntSet::CanProveNegative() const {
Analyzer analyzer;
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
return (s_int && is_negative_const(analyzer->Simplify(s_int->max_value)));
}
bool IntSet::CanProveNonPositive() const {
Analyzer analyzer;
if (const auto* s_int = (*this).as<IntervalSetNode>()) {
auto max = analyzer->Simplify(s_int->max_value);
return is_zero(max) || is_negative_const(max);
}
return false;
}
bool IntSet::CanProveNonNegative() const {
Analyzer analyzer;
if (const IntervalSetNode* s_int = (*this).as<IntervalSetNode>()) {
auto min = analyzer->Simplify(s_int->min_value);
return is_zero(min) || is_positive_const(min);
}
return false;
}
bool IntSet::HasLowerBound() const {
if (const IntervalSetNode* s_int = (*this).as<IntervalSetNode>()) {
return s_int->HasLowerBound();
}
return false;
}
bool IntSet::HasUpperBound() const {
if (const IntervalSetNode* s_int = (*this).as<IntervalSetNode>()) {
return s_int->HasUpperBound();
}
return false;
}
SignType IntSet::GetSignType() const {
if (CanProvePositive()) {
return kPositive;
} else if (CanProveNegative()) {
return kNegative;
} else if (IsSinglePoint() && is_zero(PointValue())) {
return kZero;
} else {
return kUnknown;
}
}
PrimExpr IntSet::PointValue() const {
const IntervalSetNode* s_int = (*this).as<IntervalSetNode>();
TVM_FFI_ICHECK(s_int && s_int->IsSinglePoint());
return s_int->min_value;
}
IntSet IntSet::Nothing() { return IntervalSet::Empty(); }
IntSet IntSet::Everything() { return IntervalSet::Everything(); }
IntSet IntSet::SinglePoint(PrimExpr x) { return IntervalSet::SinglePoint(x); }
IntSet IntSet::Interval(PrimExpr min, PrimExpr max) {
if (min.same_as(max)) {
return IntSet::SinglePoint(min);
}
return IntervalSet(min, max);
}
// Range related code
inline bool ProveEqual(AnalyzerObj* analyzer, PrimExpr lhs, PrimExpr rhs) {
return is_zero(analyzer->Simplify(lhs - rhs));
}
IntSet IntSet::FromMinExtent(PrimExpr min, PrimExpr extent) {
if (is_one(extent)) {
return IntSet::SinglePoint(min);
}
return IntervalSet(min, extent + min - 1);
}
IntSet IntSet::FromRange(Range r) {
// must make sure it can be matched back by MatchRange.
if (is_one(r->extent)) {
return IntSet::SinglePoint(r->min);
}
return IntervalSet(r->min, r->extent + r->min - 1);
}
bool IntSet::MatchRange(const Range& b) const {
const IntSet& a = *this;
const IntervalSetNode* a_int = a.as<IntervalSetNode>();
if (!a_int) return false;
if (!a_int->HasUpperBound() || !a_int->HasLowerBound()) return false;
Analyzer ana;
return ProveEqual(ana.get(), a_int->min_value, b->min) &&
ProveEqual(ana.get(), a_int->max_value, b->extent + b->min - 1);
}
IntSet Union(const ffi::Array<IntSet>& sets) {
if (sets.size() == 0) return IntSet::Nothing();
if (sets.size() == 1) return sets[0];
Analyzer ana;
IntervalSet x = ToIntervalSet(sets[0]);
for (size_t i = 1; i < sets.size(); ++i) {
x = Union(ana.get(), x, ToIntervalSet(sets[i]));
}
return IntervalSet(ana->Simplify(x->min_value), ana->Simplify(x->max_value));
}
ffi::Array<IntSet> UnionRegion(const ffi::Array<ffi::Array<IntSet>>& nd_int_sets) {
if (nd_int_sets.empty()) {
return {};
}
int n = nd_int_sets.size();
int ndim = nd_int_sets[0].size();
ffi::Array<IntSet> result;
result.reserve(ndim);
for (int i = 0; i < ndim; ++i) {
ffi::Array<IntSet> candidates;
candidates.reserve(n);
for (int j = 0; j < n; ++j) {
candidates.push_back(nd_int_sets[j][i]);
}
result.push_back(Union(candidates));
}
return result;
}
IntSet UnionLowerBound(const ffi::Array<IntSet>& sets) {
if (sets.size() == 0) return IntSet::Nothing();
if (sets.size() == 1) return sets[0];
Analyzer analyzer;
bool is_first_interval = true;
PrimExpr min_inclusive{nullptr};
PrimExpr max_inclusive(nullptr);
for (const IntSet& int_set : sets) {
if (int_set.IsNothing()) continue;
if (const auto* interval_set = int_set.as<IntervalSetNode>()) {
PrimExpr new_min_inclusive = interval_set->min_value;
PrimExpr new_max_inclusive = interval_set->max_value;
if (is_first_interval) {
is_first_interval = false;
min_inclusive = std::move(new_min_inclusive);
max_inclusive = std::move(new_max_inclusive);
continue;
}
bool bound_1 = is_neg_inf(new_min_inclusive) || is_pos_inf(max_inclusive) ||
analyzer->CanProve(new_min_inclusive <= max_inclusive + 1);
bool bound_2 = is_neg_inf(min_inclusive) || is_pos_inf(new_max_inclusive) ||
analyzer->CanProve(min_inclusive <= new_max_inclusive + 1);
if (bound_1 && bound_2) {
min_inclusive = min(min_inclusive, new_min_inclusive);
max_inclusive = max(max_inclusive, new_max_inclusive);
}
}
}
if (is_first_interval) {
return IntSet::Nothing();
}
return IntSet::Interval(min_inclusive, max_inclusive);
}
ffi::Array<IntSet> UnionRegionLowerBound(const ffi::Array<ffi::Array<IntSet>>& nd_int_sets) {
if (nd_int_sets.empty()) {
return {};
}
int n = nd_int_sets.size();
int ndim = nd_int_sets[0].size();
ffi::Array<IntSet> result;
result.reserve(ndim);
for (int i = 0; i < ndim; ++i) {
ffi::Array<IntSet> candidates;
candidates.reserve(n);
for (int j = 0; j < n; ++j) {
candidates.push_back(nd_int_sets[j][i]);
}
result.push_back(UnionLowerBound(candidates));
}
return result;
}
IntSet Intersect(const ffi::Array<IntSet>& sets) {
if (sets.size() == 0) return IntSet::Nothing();
if (sets.size() == 1) return sets[0];
Analyzer ana;
IntervalSet x = ToIntervalSet(sets[0]);
for (size_t i = 1; i < sets.size(); ++i) {
x = Intersect(ana.get(), x, ToIntervalSet(sets[i]));
}
return IntervalSet(ana->Simplify(x->min_value), ana->Simplify(x->max_value));
}
ffi::Map<Var, IntSet> ConvertDomMap(const ffi::Map<IterVar, IntSet>& dom_map) {
ffi::Map<Var, IntSet> dmap;
for (auto kv : dom_map) {
dmap.Set(kv.first->var, kv.second);
}
return dmap;
}
ffi::Map<Var, IntSet> ConvertDomMap(const std::unordered_map<const VarNode*, IntSet>& dom_map) {
ffi::Map<Var, IntSet> dmap;
for (auto kv : dom_map) {
dmap.Set(ffi::GetRef<Var>(kv.first), kv.second);
}
return dmap;
}
IntSet EvalSet(PrimExpr e, const ffi::Map<Var, IntSet>& dom_map) {
Analyzer ana;
return IntervalSetEvaluator(ana.get(), dom_map, {}, false).Eval(e);
}
IntSet IntSet::Vector(PrimExpr x) {
// short cut: simply get single point
if (!x.ty().IsScalableVector() && !x.ty().IsFixedLengthVector()) {
return IntSet::SinglePoint(x);
} else {
// vector case.
Analyzer ana;
ffi::Map<Var, IntSet> dmap;
return IntervalSetEvaluator(ana.get(), dmap, {}, true).Eval(x);
}
}
IntSet EvalSet(PrimExpr e, const ffi::Map<IterVar, IntSet>& dom_map) {
return EvalSet(e, ConvertDomMap(dom_map));
}
IntSet EvalSet(PrimExpr e, const std::unordered_map<const VarNode*, IntSet>& dom_map) {
return EvalSet(e, ConvertDomMap(dom_map));
}
IntSet EvalSet(Range r, const ffi::Map<Var, IntSet>& dom_map) {
Analyzer ana;
PrimType min_ty = r->min.ty();
if (min_ty.MatchesCode(DLDataTypeCode::kDLInt, DLDataTypeCode::kDLUInt) &&
ana->CanProveEqual(r->extent, 1)) {
return EvalSet(r->min, dom_map);
}
IntervalSetEvaluator m(ana.get(), dom_map);
// Simplifying first can give tighter bounds if r->min and r->extent share variables
PrimExpr sum = r->min + r->extent - 1;
auto res = m.Eval(IntervalSet(r->min, ana->Simplify(sum)));
return res;
}
IntSet EvalSet(Range r, const std::unordered_map<const VarNode*, IntSet>& dom_map) {
return EvalSet(r, ConvertDomMap(dom_map));
}
ffi::Array<IntSet> EvalSet(const ffi::Array<Range>& region, const ffi::Map<Var, IntSet>& dom_map) {
Analyzer ana;
IntervalSetEvaluator m(ana.get(), dom_map);
ffi::Array<IntSet> result;
result.reserve(region.size());
for (const Range& r : region) {
PrimExpr sum = r->min + (r->extent - 1);
result.push_back(m.Eval(IntervalSet(r->min, ana->Simplify(sum))));
}
return result;
}
IntSet EvalSet(IntSet s, const std::unordered_map<const VarNode*, IntSet>& dom_map) {
Analyzer ana;
auto dmap = ConvertDomMap(dom_map);
IntervalSetEvaluator m(ana.get(), dmap);
const IntervalSetNode* s_int = s.as<IntervalSetNode>();
PrimExpr vmax = s_int->HasUpperBound() ? m.Eval(s_int->max_value).max() : s_int->max_value;
PrimExpr vmin = s_int->HasLowerBound() ? m.Eval(s_int->min_value).min() : s_int->min_value;
return IntervalSet(vmin, vmax);
}
class SubExprIntervalSetEvaluator : public IntervalSetEvaluator {
public:
explicit SubExprIntervalSetEvaluator(AnalyzerObj* analyzer, const ffi::Map<Var, IntSet>& dom_map)
: IntervalSetEvaluator(analyzer, dom_map) {}
IntervalSet VisitExpr(const Expr& n) final {
IntervalSet ret = IntervalSetEvaluator::VisitExpr(n);
expr_map[n.as_or_throw<PrimExpr>()] = ret;
return ret;
}
ExprIntSetMap expr_map;
};
ExprIntSetMap EvalSetForEachSubExpr(PrimExpr e,
const std::unordered_map<const VarNode*, IntSet>& dom_map) {
Analyzer ana;
auto dmap = ConvertDomMap(dom_map);
SubExprIntervalSetEvaluator m(ana.get(), dmap);
m.Eval(e);
return m.expr_map;
}
IntSet EvalSet(Range r, const ffi::Map<IterVar, IntSet>& dom_map) {
return EvalSet(r, ConvertDomMap(dom_map));
}
ffi::Map<Var, arith::IntSet> AsIntSet(const ffi::Map<Var, Range>& var_dom) {
ffi::Map<Var, arith::IntSet> result;
for (auto kv : var_dom) {
const Var& var = kv.first;
const Range& range = kv.second;
result.Set(var, arith::IntSet::FromRange(range));
}
return result;
}
/*! \brief Helper function to convert IterSumExpr to the actual touched range. */
static ffi::Optional<IntSet> EvalIterSum(const IterSumExpr& iter_min, const PrimExpr& extent,
AnalyzerObj* analyzer) {
if (analyzer->CanProve(extent == 0)) {
return IntSet::Nothing();
}
if (iter_min->args.empty()) {
return IntSet::FromMinExtent(iter_min->base, extent);
}
TVM_FFI_ICHECK_EQ(iter_min->args.size(), 1) << "The `EvalIterSum` expects fused iter sum expr";
const IterSplitExpr& split = iter_min->args[0];
if (analyzer->CanProve(split->extent == 0)) {
return IntSet::Nothing();
}
if (!analyzer->CanProve(extent >= split->scale)) {
return std::nullopt;
}
const PrimExpr& base = iter_min->base;
// IterSplitExpr: (source // lower_factor) % extent * scale
// where `(source // lower_factor) % extent` is within [0, extent - 1]
if (analyzer->CanProve(split->scale < 0)) {
// If scale is negative, the var dom is [(extent - 1) * scale, 0]
// The total base is `base + (extent - 1) * scale`,
// while total extent is `dom_extent + (extent - 1) * (-scale)`
const PrimExpr& var_extent = (split->extent - 1) * split->scale;
return IntSet::FromMinExtent(base + var_extent, extent - var_extent);
} else {
// If scale is positive, the var dom is [0, (extent - 1) * scale]
// The total dom is [base, dom_extent + (extent - 1) * scale]
return IntSet::FromMinExtent(base, extent + (split->extent - 1) * split->scale);
}
}
ffi::Optional<ffi::Array<IntSet>> EstimateRegionStrictBound(const ffi::Array<Range>& region,
const ffi::Map<Var, Range>& var_dom,
const PrimExpr& predicate,
const Analyzer& analyzer) {
ffi::Map<PrimVar, Range> input_iters;
for (const auto& [var, range] : var_dom) {
input_iters.Set(var.as_or_throw<PrimVar>(), range);
}
AnalyzerObj* analyzer_ptr = analyzer.get();
int ndim = region.size();
ffi::Array<IterSumExpr> iter_sum_exprs{nullptr};
{
ffi::Array<PrimExpr> affine_indices;
affine_indices.reserve(ndim);
for (const Range& range : region) {
if (!is_const_number(range->extent)) {
// dynamic extent is not supported yet.
return std::nullopt;
}
affine_indices.push_back(range->min);
}
auto res = DetectIterMap(
/*indices=*/affine_indices, /*input_iters=*/input_iters,
/*predicate=*/predicate, /*check_level=*/IterMapLevel::Surjective, analyzer);
iter_sum_exprs = res->indices;
}
if (iter_sum_exprs.empty()) {
return std::nullopt;
}
TVM_FFI_ICHECK_EQ(iter_sum_exprs.size(), ndim);
ffi::Array<IntSet> result;
result.reserve(ndim);
for (int i = 0; i < ndim; ++i) {
const IterSumExpr& sum_expr = iter_sum_exprs[i];
const Range& range = region[i];
ffi::Optional<IntSet> int_set = EvalIterSum(sum_expr, range->extent, analyzer_ptr);
if (int_set.has_value()) {
result.push_back(int_set.value());
} else {
return std::nullopt;
}
}
return result;
}
ffi::Optional<ffi::Array<IntSet>> EstimateRegionLowerBound(const ffi::Array<Range>& region,
const ffi::Map<Var, Range>& var_dom,
const PrimExpr& predicate,
const Analyzer& analyzer) {
return EstimateRegionStrictBound(region, var_dom, predicate, analyzer);
}
ffi::Array<IntSet> EstimateRegionUpperBound(const ffi::Array<Range>& region,
const ffi::Map<Var, Range>& var_dom,
const PrimExpr& predicate, const Analyzer& analyzer) {
AnalyzerObj* analyzer_ptr = analyzer.get();
if (ffi::Optional<ffi::Array<arith::IntSet>> result = EstimateRegionStrictBound(
/*region=*/region,
/*var_dom=*/var_dom,
/*predicate=*/predicate, /*analyzer=*/analyzer)) {
return result.value();
}
ffi::Array<IntSet> result;
result.reserve(region.size());
ffi::Map<PrimVar, Range> input_iters;
for (const auto& [var, range] : var_dom) {
input_iters.Set(var.as_or_throw<PrimVar>(), range);
}
// try estimate each dimension independently
for (const Range& range : region) {
auto res = DetectIterMap(
/*indices=*/{range->min}, /*input_iters=*/input_iters,
/*predicate=*/predicate, /*check_level=*/IterMapLevel::Surjective, analyzer);
if (!res->indices.empty()) {
TVM_FFI_ICHECK_EQ(res->indices.size(), 1U);
IterSumExpr sum_expr = res->indices[0];
// dynamic extent is not supported yet.
PrimExpr extent = range->extent;
if (!is_const_number(extent)) {
IntSet relaxed = EvalSet(extent, AsIntSet(var_dom));
TVM_FFI_ICHECK(relaxed.HasUpperBound());
extent = relaxed.max();
}
if (ffi::Optional<IntSet> int_set = EvalIterSum(sum_expr, range->extent, analyzer_ptr)) {
result.push_back(int_set.value());
continue;
}
}
// fallback to coarse grained evalset
result.push_back(EvalSet(range, AsIntSet(var_dom)));
}
return result;
}
// Pattern A (RM): auto-default repr from reflection.
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef()
.def("arith.intset_single_point", IntSet::SinglePoint)
.def("arith.intset_vector", IntSet::Vector)
.def("arith.intset_interval", IntSet::Interval)
.def_method("arith.IntervalSetGetMin", &IntSet::min)
.def_method("arith.IntervalSetGetMax", &IntSet::max)
.def_method("arith.IntSetIsNothing", &IntSet::IsNothing)
.def_method("arith.IntSetIsEverything", &IntSet::IsEverything)
.def("arith.EstimateRegionLowerBound",
[](ffi::Array<Range> region, ffi::Map<Var, Range> var_dom, PrimExpr predicate,
ffi::Optional<Analyzer> opt_analyzer) -> ffi::Optional<ffi::Array<IntSet>> {
Analyzer analyzer = opt_analyzer.has_value() ? opt_analyzer.value() : Analyzer();
return EstimateRegionLowerBound(region, var_dom, predicate, analyzer);
})
.def("arith.EstimateRegionStrictBound",
[](ffi::Array<Range> region, ffi::Map<Var, Range> var_dom, PrimExpr predicate,
ffi::Optional<Analyzer> opt_analyzer) -> ffi::Optional<ffi::Array<IntSet>> {
Analyzer analyzer = opt_analyzer.has_value() ? opt_analyzer.value() : Analyzer();
return EstimateRegionStrictBound(region, var_dom, predicate, analyzer);
})
.def("arith.EstimateRegionUpperBound",
[](ffi::Array<Range> region, ffi::Map<Var, Range> var_dom, PrimExpr predicate,
ffi::Optional<Analyzer> opt_analyzer) -> ffi::Optional<ffi::Array<IntSet>> {
Analyzer analyzer = opt_analyzer.has_value() ? opt_analyzer.value() : Analyzer();
return EstimateRegionUpperBound(region, var_dom, predicate, analyzer);
})
.def("arith.PosInf", []() { return SymbolicLimits::pos_inf_; })
.def("arith.NegInf", []() { return SymbolicLimits::neg_inf_; })
.def("arith.UnionLowerBound", UnionLowerBound);
}
} // namespace arith
} // namespace tvm