blob: f7f46fae78a4891bac5733ccbabbeb21dc921cd7 [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 const_fold.h
* \brief Centralized location for constant folding.
*/
#ifndef TVM_ARITH_CONST_FOLD_H_
#define TVM_ARITH_CONST_FOLD_H_
#include <tvm/ffi/optional.h>
#include <tvm/runtime/logging.h>
#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <algorithm>
#include <cmath>
#include <limits>
#include "int_operator.h"
namespace tvm {
namespace arith {
/*!
* \brief Try to run binary compute with constant folding.
*
* \param a The left operand.
* \param b The right operand.
* \tparam Op The operator type.
*
* \note a and b Must already matched data types with each other.
* \return std::nullopt if constant fold fails, otherwise return folded result.
*/
template <typename Op>
inline ffi::Optional<PrimExpr> TryConstFold(PrimExpr a, PrimExpr b);
/*!
* \brief Try to run unary compute with constant folding.
*
* \param a The left operand.
* \tparam Op The operator type.
*
* \note a and b Must already matched data types with each other.
* \return std::nullopt if constant fold fails, otherwise return folded result.
*/
template <typename Op>
inline ffi::Optional<PrimExpr> TryConstFold(PrimExpr a);
/*!
* \brief Check whether type is used to represent index.
*
* Index types are frequently used in shape computation
* and need to be aggressively constant-folded.
*
* \param type The type to represent index.
* \return the checked result.
*/
inline bool IsIndexType(DLDataType type) {
return type.code == static_cast<uint8_t>(DLDataTypeCode::kDLInt) &&
(type.bits == 32 || type.bits == 64) && type.lanes == 1;
}
inline bool IsIndexTypedExpr(const ExprNode* expr) {
TVM_FFI_DCHECK(expr != nullptr);
TVM_FFI_DCHECK(!expr->ExprNode::ty.IsMissing());
const auto* prim_ty = expr->ExprNode::ty.as<PrimTypeNode>();
TVM_FFI_DCHECK(prim_ty != nullptr);
return IsIndexType(prim_ty->dtype);
}
inline bool IsIndexTypedExpr(const PrimExpr& expr) { return IsIndexTypedExpr(expr.get()); }
/*! \brief Helper to get const folding result repr in int64. */
inline int64_t GetFoldResultInt64Repr(int64_t x, const PrimType& dtype) {
if (dtype.bits() < 64) {
x &= (1LL << dtype.bits()) - 1;
}
if (dtype.MatchesCode(DLDataTypeCode::kDLInt)) {
int64_t m = 1LL << (dtype.bits() - 1);
x = (x ^ m) - m;
}
return x;
}
/*! \brief Helper to get fp32 const folding result repr in double. */
inline double GetFoldResultDoubleRepr(float x) {
double res = static_cast<double>(x);
if (std::isinf(res) || std::isnan(res)) {
return res;
}
// certain platform (eg, on gcc7-i386) do the folding arithmetic
// on float and write back to double is optimized to double
// precision arithmetic, this is legal and we check the output
// range thus to ensure consistency when the float result is inf.
if (res < std::numeric_limits<float>::lowest()) {
LOG(WARNING) << "underlying float value overflow";
return -std::numeric_limits<double>::infinity();
} else if (res > std::numeric_limits<float>::max()) {
LOG(WARNING) << "underlying float value overflow";
return std::numeric_limits<double>::infinity();
}
return res;
}
#define TVM_ARITH_CONST_PROPAGATION(BODY) \
using tirx::FloatImmNode; \
const IntImmNode* pa = a.as<IntImmNode>(); \
const IntImmNode* pb = b.as<IntImmNode>(); \
const FloatImmNode* fa = a.as<FloatImmNode>(); \
const FloatImmNode* fb = b.as<FloatImmNode>(); \
BODY;
#define TVM_INDEX_CONST_PROPAGATION(BODY) \
const IntImmNode* pa = a.as<IntImmNode>(); \
const IntImmNode* pb = b.as<IntImmNode>(); \
if (arith::IsIndexTypedExpr(a) && arith::IsIndexTypedExpr(b)) { \
BODY; \
}
// specialization of constant folders.
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Add>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
int64_t res = pa->value + pb->value;
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa && pa->value == 0) return b;
if (pb && pb->value == 0) return a;
if (fa && fb) {
if (result_ty.bits() == 32) {
return FloatImm(result_ty, GetFoldResultDoubleRepr(static_cast<float>(fa->value) +
static_cast<float>(fb->value)));
} else if (result_ty.bits() == 64) {
return FloatImm(result_ty, fa->value + fb->value);
}
}
if (fa && fa->value == 0) return b;
if (fb && fb->value == 0) return a;
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Sub>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
TVM_FFI_ICHECK(!((pa && pa->ty.as_or_throw<PrimType>().MatchesCode(DLDataTypeCode::kDLUInt) &&
pa->value == 0U) &&
(pb && pb->ty.as_or_throw<PrimType>().MatchesCode(DLDataTypeCode::kDLUInt) &&
pb->value > 0U)))
<< "Checked failed. Minuend 's value is 0U and it's dtype is uint "
<< "while Subtrahend's dtype is uint; which will cause a negative uint";
PrimType result_ty = a.ty();
if (pa && pb) {
int64_t res = pa->value - pb->value;
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pb && pb->value == 0) return a;
if (fa && fb) {
if (result_ty.bits() == 32) {
return FloatImm(result_ty, GetFoldResultDoubleRepr(static_cast<float>(fa->value) -
static_cast<float>(fb->value)));
} else if (result_ty.bits() == 64) {
return FloatImm(result_ty, fa->value - fb->value);
}
}
if (fb && fb->value == 0) return a;
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Mul>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
int64_t res = pa->value * pb->value;
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa) {
if (pa->value == 1) return b;
if (pa->value == 0) return a;
}
if (pb) {
if (pb->value == 1) return a;
if (pb->value == 0) return b;
}
if (fa && fb) {
if (result_ty.bits() == 32) {
return FloatImm(result_ty, GetFoldResultDoubleRepr(static_cast<float>(fa->value) *
static_cast<float>(fb->value)));
} else if (result_ty.bits() == 64) {
return FloatImm(result_ty, fa->value * fb->value);
}
}
if (fa) {
if (fa->value == 1) return b;
if (fa->value == 0) return a;
}
if (fb) {
if (fb->value == 1) return a;
if (fb->value == 0) return b;
}
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Div>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
// due to division and mod can have different modes
// NOTE: this will assumes truc div.
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
int64_t res = pa->value / pb->value;
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa) {
if (pa->value == 0) return a;
}
if (pb) {
if (pb->value == 1) return a;
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
}
if (fa && fb) {
TVM_FFI_ICHECK_NE(fb->value, 0) << "Divide by zero";
if (result_ty.bits() == 32) {
return FloatImm(result_ty, GetFoldResultDoubleRepr(static_cast<float>(fa->value) /
static_cast<float>(fb->value)));
} else if (result_ty.bits() == 64) {
return FloatImm(result_ty, fa->value / fb->value);
}
}
if (fa && fa->value == 0) return a;
if (fb) {
if (fb->value == 1) return a;
TVM_FFI_ICHECK_NE(fb->value, 0) << "Divide by zero";
}
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Mod>(PrimExpr a, PrimExpr b) {
TVM_INDEX_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
int64_t res = pa->value % pb->value;
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa) {
if (pa->value == 0) return a;
}
if (pb) {
if (pb->value == 1) return IntImm(result_ty, 0);
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
}
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::FloorDiv>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
int64_t res = arith::floordiv(pa->value, pb->value);
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa) {
if (pa->value == 0) return a;
}
if (pb) {
if (pb->value == 1) return a;
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
}
if (fa && fb && fb->value != 0) {
if (result_ty.bits() == 32) {
return FloatImm(result_ty,
GetFoldResultDoubleRepr(std::floor(static_cast<float>(fa->value) /
static_cast<float>(fb->value))));
} else if (result_ty.bits() == 64) {
return FloatImm(result_ty, std::floor(fa->value / fb->value));
} else {
return std::nullopt;
}
}
if (fa && fa->value == 0) return a;
if (fb) {
if (fb->value == 1) return a;
TVM_FFI_ICHECK_NE(fb->value, 0) << "Divide by zero";
}
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::FloorMod>(PrimExpr a, PrimExpr b) {
const IntImmNode* ua = a.as<IntImmNode>();
const IntImmNode* ub = b.as<IntImmNode>();
PrimType utype = a.ty();
if (utype.MatchesCode(DLDataTypeCode::kDLUInt) && utype == b.ty() && ua && ub) {
auto as_uint = [](int64_t value, const PrimType& dtype) -> uint64_t {
uint64_t result = static_cast<uint64_t>(value);
if (dtype.bits() < 64) {
result &= (uint64_t{1} << dtype.bits()) - 1;
}
return result;
};
uint64_t lhs = as_uint(ua->value, utype);
uint64_t rhs = as_uint(ub->value, b.ty());
TVM_FFI_ICHECK_NE(rhs, 0U) << "Divide by zero";
return IntImm(utype, static_cast<int64_t>(lhs % rhs));
}
TVM_INDEX_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) {
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
int64_t res = arith::floormod(pa->value, pb->value);
return IntImm(result_ty, GetFoldResultInt64Repr(res, result_ty));
}
if (pa) {
if (pa->value == 0) return a;
}
if (pb) {
if (pb->value == 1) return IntImm(result_ty, 0);
TVM_FFI_ICHECK_NE(pb->value, 0) << "Divide by zero";
}
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Min>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) return IntImm(result_ty, std::min(pa->value, pb->value));
if (fa && fb) return FloatImm(result_ty, std::min(fa->value, fb->value));
});
if (a.same_as(b)) return a;
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Max>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
PrimType result_ty = a.ty();
if (pa && pb) return IntImm(result_ty, std::max(pa->value, pb->value));
if (fa && fb) return FloatImm(result_ty, std::max(fa->value, fb->value));
});
if (a.same_as(b)) return a;
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::GT>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value > pb->value);
if (fa && fb) return IntImm::Bool(fa->value > fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::GE>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value >= pb->value);
if (fa && fb) return IntImm::Bool(fa->value >= fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::LT>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value < pb->value);
if (fa && fb) return IntImm::Bool(fa->value < fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::LE>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value <= pb->value);
if (fa && fb) return IntImm::Bool(fa->value <= fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::EQ>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value == pb->value);
if (fa && fb) return IntImm::Bool(fa->value == fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::NE>(PrimExpr a, PrimExpr b) {
TVM_ARITH_CONST_PROPAGATION({
if (pa && pb) return IntImm::Bool(pa->value != pb->value);
if (fa && fb) return IntImm::Bool(fa->value != fb->value);
});
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::And>(PrimExpr a, PrimExpr b) {
const IntImmNode* pa = a.as<IntImmNode>();
const IntImmNode* pb = b.as<IntImmNode>();
if (pa && pa->value) return b;
if (pa && !pa->value) return a;
if (pb && pb->value) return a;
if (pb && !pb->value) return b;
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Or>(PrimExpr a, PrimExpr b) {
const IntImmNode* pa = a.as<IntImmNode>();
const IntImmNode* pb = b.as<IntImmNode>();
if (pa && pa->value) return a;
if (pa && !pa->value) return b;
if (pb && pb->value) return b;
if (pb && !pb->value) return a;
return std::nullopt;
}
template <>
inline ffi::Optional<PrimExpr> TryConstFold<tirx::Not>(PrimExpr a) {
const IntImmNode* pa = a.as<IntImmNode>();
if (pa) {
return IntImm::Bool(!(pa->value));
}
return std::nullopt;
}
/*! \brief Helper namespace for symbolic value limits */
struct SymbolicLimits {
/*! \brief positive infinity */
static PrimExpr pos_inf_;
/*! \brief negative infinity */
static PrimExpr neg_inf_;
};
/*!
* \brief Opaque expression representing positive infinity.
*
* It can only be used as parameter of by min/max
* for integer analysis and cannot be used in normal expressions.
*
* \return positive infinity.
*/
inline PrimExpr pos_inf() { return SymbolicLimits::pos_inf_; }
/*!
* \brief Check if value is positive infinity.
* \param value The value to be checked.
*
* \return The check result.
*/
inline bool is_pos_inf(const PrimExpr& value) { return value.same_as(SymbolicLimits::pos_inf_); }
/*!
* \brief Opaque expression representing negative infinity.
*
* It can only be used as parameter of by min/max
* for integer analysis and cannot be used in normal expressions.
*
* \return negative infinity.
*/
inline PrimExpr neg_inf() { return SymbolicLimits::neg_inf_; }
/*!
* \brief Check if value is negative infinity.
* \param value The value to be checked.
*
* \return The check result.
*/
inline bool is_neg_inf(const PrimExpr& value) { return value.same_as(SymbolicLimits::neg_inf_); }
} // namespace arith
} // namespace tvm
#endif // TVM_ARITH_CONST_FOLD_H_