This is an automated email from the ASF dual-hosted git repository.
tqchen pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new f94f8a8ca6 [REFACTOR][IR] Share checked PrimVar view across dialects
(#20348)
f94f8a8ca6 is described below
commit f94f8a8ca684356eb2bba0793ccac5af360e763b
Author: Tianqi Chen <[email protected]>
AuthorDate: Tue Sep 15 17:16:48 2026 -0400
[REFACTOR][IR] Share checked PrimVar view across dialects (#20348)
Move the zero-state `PrimVar` checked view and its FFI traits beside the
shared IR `Var` definitions. Arithmetic and direct C++ callers now use
`tvm::PrimVar`, and arithmetic headers drop aliases and includes made
redundant by the move.
The class and checked conversion behavior remain unchanged, including
nullability, exact variable/type checks, vector-type acceptance and
reference identity. IterVar remains in TIRX. Existing tests receive only
required namespace changes.
---
include/tvm/arith/analyzer.h | 2 -
include/tvm/arith/int_set.h | 4 +-
include/tvm/arith/int_solver.h | 4 -
include/tvm/arith/iter_affine_map.h | 14 ++-
include/tvm/arith/pattern.h | 5 +-
include/tvm/ir/expr.h | 71 ++++++++++++++
include/tvm/ir/object_functor.h | 2 +-
include/tvm/relax/transform.h | 5 +-
include/tvm/te/operation.h | 9 +-
include/tvm/tirx/var.h | 73 ---------------
include/tvm/topi/broadcast.h | 10 +-
include/tvm/topi/detail/broadcast.h | 16 ++--
include/tvm/topi/detail/strided_slice.h | 2 +-
include/tvm/topi/nn.h | 40 ++++----
include/tvm/topi/transform.h | 14 +--
src/arith/analyzer.cc | 2 +-
src/arith/int_set.cc | 4 +-
src/arith/pattern_match.h | 6 +-
src/arith/presburger_set.h | 14 +--
src/arith/solve_linear_equation.cc | 2 +-
src/arith/transitive_comparison_analyzer.cc | 8 +-
src/backend/cuda/codegen/codegen_cuda.cc | 8 +-
src/relax/analysis/layout_transformation.cc | 10 +-
src/relax/analysis/tir_op_pattern_kind.cc | 12 +--
src/relax/analysis/type_analysis.cc | 17 ++--
src/relax/backend/contrib/tensorrt/codegen.cc | 2 +-
src/relax/backend/vm/vm_shape_lower.cc | 6 +-
src/relax/ir/block_builder.cc | 16 ++--
src/relax/ir/dataflow_matcher.cc | 2 +-
src/relax/ir/expr.cc | 2 +-
src/relax/script/printer/dependent_type.cc | 4 +-
src/relax/script/printer/tir.cc | 2 +-
src/relax/transform/alter_op_impl.cc | 2 +-
src/relax/transform/bind_symbolic_vars.cc | 15 ++-
src/relax/transform/bundle_model_params.cc | 4 +-
src/relax/transform/canonicalize_bindings.cc | 6 +-
src/relax/transform/convert_layout.cc | 6 +-
src/relax/transform/fuse_tir.cc | 12 +--
src/relax/transform/lazy_transform_params.cc | 2 +-
src/relax/transform/lift_transform_params.cc | 4 +-
src/relax/transform/rewrite_cuda_graph.cc | 39 ++++----
src/relax/utils.cc | 4 +-
.../multi_level_tiling_tensor_core.cc | 16 ++--
src/s_tir/transform/default_gpu_schedule.cc | 8 +-
src/te/operation/create_primfunc.cc | 2 +-
src/tirx/ir/buffer.cc | 2 +-
src/tirx/op/op.cc | 1 -
src/tirx/script/builder/ir.cc | 104 ++++++++++-----------
src/tirx/script/printer/block.cc | 2 +-
src/tirx/script/printer/expr.cc | 2 +-
tests/cpp/ir_functor_test.cc | 4 +-
tests/cpp/pattern_match_test.cc | 20 ++--
tests/cpp/tir_analysis_side_effect.cc | 2 +-
53 files changed, 311 insertions(+), 334 deletions(-)
diff --git a/include/tvm/arith/analyzer.h b/include/tvm/arith/analyzer.h
index 08f85948ff..47be76e898 100644
--- a/include/tvm/arith/analyzer.h
+++ b/include/tvm/arith/analyzer.h
@@ -57,8 +57,6 @@ class AnalyzerObj;
class Analyzer;
class ConstraintContext;
-using tirx::Var;
-
enum DivMode {
/*! \brief Truncated division. */
kTruncDiv,
diff --git a/include/tvm/arith/int_set.h b/include/tvm/arith/int_set.h
index 4901bf5696..16a55bdf01 100644
--- a/include/tvm/arith/int_set.h
+++ b/include/tvm/arith/int_set.h
@@ -34,8 +34,6 @@ namespace tvm {
namespace arith {
using tirx::IterVar;
-using tirx::Var;
-using tirx::VarNode;
class AnalyzerObj;
class Analyzer;
@@ -201,7 +199,7 @@ IntSet EvalSet(PrimExpr e, const ffi::Map<Var, IntSet>&
dom_map);
* \param dom_map The domain of each variable.
* \return An integer set that can cover all the possible values of e.
*/
-IntSet EvalSet(PrimExpr e, const std::unordered_map<const tirx::VarNode*,
IntSet>& dom_map);
+IntSet EvalSet(PrimExpr e, const std::unordered_map<const VarNode*, IntSet>&
dom_map);
/*!
* \brief Find an symbolic integer set that contains is union over
* all the possible conditional values in dom_map.
diff --git a/include/tvm/arith/int_solver.h b/include/tvm/arith/int_solver.h
index c2b2574109..92a790830e 100644
--- a/include/tvm/arith/int_solver.h
+++ b/include/tvm/arith/int_solver.h
@@ -26,7 +26,6 @@
#include <tvm/ir/expr.h>
#include <tvm/ir/prim/expr.h>
-#include <tvm/tirx/op.h>
#include <unordered_map>
#include <utility>
@@ -38,9 +37,6 @@ namespace tvm {
namespace arith {
using tirx::IterVar;
-using tirx::PrimVar;
-using tirx::Var;
-using tirx::VarNode;
// According to experiments two best simplifications orders were can->rw and
rw->can->rw,
// but rw->can->rw is better for a couple of cases.
diff --git a/include/tvm/arith/iter_affine_map.h
b/include/tvm/arith/iter_affine_map.h
index 6e5e14b236..f47492f497 100644
--- a/include/tvm/arith/iter_affine_map.h
+++ b/include/tvm/arith/iter_affine_map.h
@@ -52,7 +52,6 @@
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/cow.h>
#include <tvm/ir/expr.h>
-#include <tvm/tirx/var.h>
namespace tvm {
namespace arith {
@@ -316,9 +315,8 @@ class IterMapResult : public ffi::ObjectRef {
* The return object's .indices is empty on failure.
*/
IterMapResult DetectIterMap(const ffi::Array<PrimExpr>& indices,
- const ffi::Map<tirx::PrimVar, Range>& input_iters,
- const PrimExpr& predicate, IterMapLevel
check_level,
- const arith::Analyzer& analyzer,
+ const ffi::Map<PrimVar, Range>& input_iters, const
PrimExpr& predicate,
+ IterMapLevel check_level, const arith::Analyzer&
analyzer,
bool simplify_trivial_iterators = true);
/*!
@@ -333,7 +331,7 @@ IterMapResult DetectIterMap(const ffi::Array<PrimExpr>&
indices,
* \return The indices after rewrite
*/
ffi::Array<PrimExpr> IterMapSimplify(const ffi::Array<PrimExpr>& indices,
- const ffi::Map<tirx::PrimVar, Range>&
input_iters,
+ const ffi::Map<PrimVar, Range>&
input_iters,
const PrimExpr& input_pred, IterMapLevel
check_level,
const arith::Analyzer& analyzer,
bool simplify_trivial_iterators = true);
@@ -389,8 +387,8 @@ ffi::Map<Var, PrimExpr> InverseAffineIterMap(const
ffi::Array<IterSumExpr>& iter
Empty array if no match can be found.
*/
ffi::Array<ffi::Array<IterMark>> SubspaceDivide(const ffi::Array<PrimExpr>&
bindings,
- const ffi::Map<tirx::PrimVar,
Range>& input_iters,
- const
ffi::Array<tirx::PrimVar>& sub_iters,
+ const ffi::Map<PrimVar,
Range>& input_iters,
+ const ffi::Array<PrimVar>&
sub_iters,
const PrimExpr& predicate,
IterMapLevel check_level,
const arith::Analyzer&
analyzer,
bool
simplify_trivial_iterators = true);
@@ -418,7 +416,7 @@ PrimExpr NormalizeIterMapToExpr(const PrimExpr& expr);
* \param analyzer The input analyzer.
* \note This function is useful to detect iterator stride patterns.
*/
-IterSumExpr NormalizeToIterSum(PrimExpr index, const ffi::Map<tirx::PrimVar,
Range>& input_iters,
+IterSumExpr NormalizeToIterSum(PrimExpr index, const ffi::Map<PrimVar, Range>&
input_iters,
const arith::Analyzer& analyzer);
} // namespace arith
diff --git a/include/tvm/arith/pattern.h b/include/tvm/arith/pattern.h
index 51f927f60b..dd48abb763 100644
--- a/include/tvm/arith/pattern.h
+++ b/include/tvm/arith/pattern.h
@@ -25,7 +25,6 @@
#define TVM_ARITH_PATTERN_H_
#include <tvm/ir/expr.h>
-#include <tvm/ir/prim/expr.h>
namespace tvm {
namespace arith {
@@ -37,7 +36,7 @@ namespace arith {
* \param vars List of variables to be used in detection.
* \return [coeff[i]] if it is possible, empty array if it is not.
*/
-ffi::Array<PrimExpr> DetectLinearEquation(const PrimExpr& e, const
ffi::Array<tirx::PrimVar>& vars);
+ffi::Array<PrimExpr> DetectLinearEquation(const PrimExpr& e, const
ffi::Array<PrimVar>& vars);
/*!
* \brief Detect if expression corresponds to clip bound of the vars
@@ -47,7 +46,7 @@ ffi::Array<PrimExpr> DetectLinearEquation(const PrimExpr& e,
const ffi::Array<ti
* \return concat([min_value[i], max_value[i]]), None is returned if there is
no min or max value
* return empty if the e does not match the pattern.
*/
-ffi::Array<PrimExpr> DetectClipBound(const PrimExpr& e, const
ffi::Array<tirx::PrimVar>& vars);
+ffi::Array<PrimExpr> DetectClipBound(const PrimExpr& e, const
ffi::Array<PrimVar>& vars);
} // namespace arith
} // namespace tvm
diff --git a/include/tvm/ir/expr.h b/include/tvm/ir/expr.h
index acc28562c3..62db747e8b 100644
--- a/include/tvm/ir/expr.h
+++ b/include/tvm/ir/expr.h
@@ -38,6 +38,7 @@
#include <limits>
#include <optional>
#include <string>
+#include <utility>
namespace tvm {
@@ -388,6 +389,37 @@ class Var : public Expr {
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(Var, Expr, VarNode);
};
+/*!
+ * \brief Checked scalar view over a VarNode.
+ *
+ * PrimVar is a zero-state reference view over the same VarNode as Var. It
additionally
+ * guarantees that the inherited ExprNode::ty is PrimType.
+ */
+class PrimVar : public PrimExpr {
+ public:
+ /*! \brief Construct a scalar variable directly from a primitive type. */
+ explicit PrimVar(ffi::String name, PrimType dtype = PrimType::Int(32), Span
span = Span())
+ : PrimExpr(Var(std::move(name), std::move(dtype),
std::move(span)).as_or_throw<PrimExpr>()) {}
+
+ /*! \brief Construct a scalar variable directly from a checked type
annotation. */
+ explicit PrimVar(ffi::String name, Type type_annotation, Span span = Span())
+ : PrimExpr(Var(std::move(name), std::move(type_annotation),
std::move(span))
+ .as_or_throw<PrimExpr>()) {}
+
+ /*! \brief Safe widening to a general Var view over the same node. */
+ operator Var() const { return this->as_or_throw<Var>(); }
+
+ PrimVar CopyWithSuffix(const ffi::String& suffix) const {
+ return
this->as_or_throw<Var>().CopyWithSuffix(suffix).as_or_throw<PrimVar>();
+ }
+ PrimVar CopyWithDType(PrimType dtype) const {
+ return
this->as_or_throw<Var>().CopyWithDType(dtype).as_or_throw<PrimVar>();
+ }
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(PrimVar, PrimExpr, VarNode);
+ static constexpr bool _type_container_is_exact = false;
+};
+
class GlobalVar;
/*!
* \brief Global variable that lives in the top-level module.
@@ -635,6 +667,45 @@ class Range : public ffi::ObjectRef {
};
namespace ffi {
+
+template <>
+inline constexpr bool use_default_type_traits_v<PrimVar> = false;
+
+template <>
+struct TypeTraits<PrimVar> : public ObjectRefTypeTraitsBase<PrimVar> {
+ using Base = ObjectRefTypeTraitsBase<PrimVar>;
+ using Base::CopyFromAnyViewAfterCheck;
+ using Base::CopyToAnyView;
+ using Base::GetMismatchTypeInfo;
+ using Base::MoveFromAnyAfterCheck;
+ using Base::MoveToAny;
+ using Base::TypeSchema;
+ using Base::TypeStr;
+
+ TVM_FFI_INLINE static bool CheckAnyStrict(const TVMFFIAny* src) {
+ if (src->type_index == TypeIndex::kTVMFFINone) {
+ return PrimVar::_type_is_nullable;
+ }
+ if (src->type_index != VarNode::RuntimeTypeIndex()) {
+ return false;
+ }
+ const auto* var = static_cast<const VarNode*>(
+ details::ObjectUnsafe::ObjectPtrFromUnowned<Object>(src->v_obj).get());
+ return details::AnyUnsafe::CheckAnyStrict<PrimType>(var->ExprNode::ty);
+ }
+
+ TVM_FFI_INLINE static std::optional<PrimVar> TryCastFromAnyView(const
TVMFFIAny* src) {
+ if (CheckAnyStrict(src)) {
+ if (src->type_index == TypeIndex::kTVMFFINone) {
+ return details::ObjectUnsafe::ObjectRefFromObjectPtr<PrimVar>(nullptr);
+ }
+ return details::ObjectUnsafe::ObjectRefFromObjectPtr<PrimVar>(
+ details::ObjectUnsafe::ObjectPtrFromUnowned<VarNode>(src->v_obj));
+ }
+ return std::nullopt;
+ }
+};
+
template <>
inline constexpr bool object_ref_contains_v<PrimExpr, IntImmNode> = true;
template <>
diff --git a/include/tvm/ir/object_functor.h b/include/tvm/ir/object_functor.h
index 33a0a33ae2..5c2493307d 100644
--- a/include/tvm/ir/object_functor.h
+++ b/include/tvm/ir/object_functor.h
@@ -49,7 +49,7 @@ namespace tvm {
* return prefix + "IntImm";
* });
*
- * tirx::PrimVar x("x");
+ * PrimVar x("x");
* PrimExpr y = x + 1;
* // dispatch to IntImm, outputs "MyIntImm"
* LOG(INFO) << tostr(IntImm::Int32(1), "My");
diff --git a/include/tvm/relax/transform.h b/include/tvm/relax/transform.h
index df66358455..d1b88ac93c 100644
--- a/include/tvm/relax/transform.h
+++ b/include/tvm/relax/transform.h
@@ -215,9 +215,8 @@ TVM_DLL Pass BindParams(ffi::String func_name,
ffi::Map<Any, ffi::ObjectRef> par
*
* \return The Pass.
*/
-TVM_DLL Pass
-BindSymbolicVars(ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>, PrimExpr>
binding_map,
- ffi::Optional<ffi::String> func_name = std::nullopt);
+TVM_DLL Pass BindSymbolicVars(ffi::Map<ffi::Variant<PrimVar, ffi::String>,
PrimExpr> binding_map,
+ ffi::Optional<ffi::String> func_name =
std::nullopt);
/*!
* \brief Fold constant expressions within dataflow blocks.
diff --git a/include/tvm/te/operation.h b/include/tvm/te/operation.h
index 356763b63b..42800875ef 100644
--- a/include/tvm/te/operation.h
+++ b/include/tvm/te/operation.h
@@ -49,9 +49,9 @@ namespace te {
class CommReducerNode : public ffi::Object {
public:
/*! \brief The left argument of reducer */
- ffi::Array<tirx::PrimVar> lhs;
+ ffi::Array<PrimVar> lhs;
/*! \brief The right argument of reducer */
- ffi::Array<tirx::PrimVar> rhs;
+ ffi::Array<PrimVar> rhs;
/*! \brief The result of reducer */
ffi::Array<PrimExpr> result;
/*!
@@ -88,9 +88,8 @@ class CommReducerNode : public ffi::Object {
*/
class CommReducer : public ffi::ObjectRef {
public:
- TVM_DLL CommReducer(ffi::Array<tirx::PrimVar> lhs, ffi::Array<tirx::PrimVar>
rhs,
- ffi::Array<PrimExpr> result, ffi::Array<PrimExpr>
identity_element,
- Span span = Span());
+ TVM_DLL CommReducer(ffi::Array<PrimVar> lhs, ffi::Array<PrimVar> rhs,
ffi::Array<PrimExpr> result,
+ ffi::Array<PrimExpr> identity_element, Span span =
Span());
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(CommReducer, ffi::ObjectRef,
CommReducerNode);
};
diff --git a/include/tvm/tirx/var.h b/include/tvm/tirx/var.h
index d89c283eb2..63330ecdc3 100644
--- a/include/tvm/tirx/var.h
+++ b/include/tvm/tirx/var.h
@@ -37,37 +37,6 @@ namespace tirx {
using VarNode = tvm::VarNode;
using Var = tvm::Var;
-/*!
- * \brief Checked scalar view over a VarNode.
- *
- * PrimVar is a zero-state reference view over the same VarNode as Var. It
additionally
- * guarantees that the inherited ExprNode::ty is PrimType.
- */
-class PrimVar : public PrimExpr {
- public:
- /*! \brief Construct a scalar variable directly from a primitive type. */
- explicit PrimVar(ffi::String name, PrimType dtype = PrimType::Int(32), Span
span = Span())
- : PrimExpr(Var(std::move(name), std::move(dtype),
std::move(span)).as_or_throw<PrimExpr>()) {}
-
- /*! \brief Construct a scalar variable directly from a checked type
annotation. */
- explicit PrimVar(ffi::String name, Type type_annotation, Span span = Span())
- : PrimExpr(Var(std::move(name), std::move(type_annotation),
std::move(span))
- .as_or_throw<PrimExpr>()) {}
-
- /*! \brief Safe widening to a general Var view over the same node. */
- operator Var() const { return this->as_or_throw<Var>(); }
-
- PrimVar CopyWithSuffix(const ffi::String& suffix) const {
- return
this->as_or_throw<Var>().CopyWithSuffix(suffix).as_or_throw<PrimVar>();
- }
- PrimVar CopyWithDType(PrimType dtype) const {
- return
this->as_or_throw<Var>().CopyWithDType(dtype).as_or_throw<PrimVar>();
- }
-
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(PrimVar, PrimExpr, VarNode);
- static constexpr bool _type_container_is_exact = false;
-};
-
using Region = ffi::Array<Range>;
/*!
@@ -234,48 +203,6 @@ inline const char* IterVarType2String(IterVarType t) {
} // namespace tvm
-namespace tvm::ffi {
-
-template <>
-inline constexpr bool use_default_type_traits_v<tirx::PrimVar> = false;
-
-template <>
-struct TypeTraits<tirx::PrimVar> : public
ObjectRefTypeTraitsBase<tirx::PrimVar> {
- using Base = ObjectRefTypeTraitsBase<tirx::PrimVar>;
- using Base::CopyFromAnyViewAfterCheck;
- using Base::CopyToAnyView;
- using Base::GetMismatchTypeInfo;
- using Base::MoveFromAnyAfterCheck;
- using Base::MoveToAny;
- using Base::TypeSchema;
- using Base::TypeStr;
-
- TVM_FFI_INLINE static bool CheckAnyStrict(const TVMFFIAny* src) {
- if (src->type_index == TypeIndex::kTVMFFINone) {
- return tirx::PrimVar::_type_is_nullable;
- }
- if (src->type_index != tirx::VarNode::RuntimeTypeIndex()) {
- return false;
- }
- const auto* var = static_cast<const tirx::VarNode*>(
- details::ObjectUnsafe::ObjectPtrFromUnowned<Object>(src->v_obj).get());
- return details::AnyUnsafe::CheckAnyStrict<PrimType>(var->ExprNode::ty);
- }
-
- TVM_FFI_INLINE static std::optional<tirx::PrimVar> TryCastFromAnyView(const
TVMFFIAny* src) {
- if (CheckAnyStrict(src)) {
- if (src->type_index == TypeIndex::kTVMFFINone) {
- return
details::ObjectUnsafe::ObjectRefFromObjectPtr<tirx::PrimVar>(nullptr);
- }
- return details::ObjectUnsafe::ObjectRefFromObjectPtr<tirx::PrimVar>(
-
details::ObjectUnsafe::ObjectPtrFromUnowned<tirx::VarNode>(src->v_obj));
- }
- return std::nullopt;
- }
-};
-
-} // namespace tvm::ffi
-
namespace std {
template <>
struct hash<::tvm::tirx::IterVar> : public ::tvm::ffi::ObjectPtrHash {};
diff --git a/include/tvm/topi/broadcast.h b/include/tvm/topi/broadcast.h
index ada383a97b..a114bf3287 100644
--- a/include/tvm/topi/broadcast.h
+++ b/include/tvm/topi/broadcast.h
@@ -63,7 +63,7 @@ inline tvm::te::Tensor broadcast_to(const tvm::te::Tensor& t,
oshape.push_back(bh.common_shape[i]);
}
}
- auto l = [&](tvm::ffi::Array<tvm::tirx::PrimVar> ovars) {
+ auto l = [&](tvm::ffi::Array<tvm::PrimVar> ovars) {
return t(detail::InputIndexFromBroadcast(ovars, t, bh.vars2, bh.all_vars));
};
return tvm::te::compute(oshape, l, name, tag);
@@ -80,15 +80,15 @@ inline tvm::te::Tensor broadcast_to(const tvm::te::Tensor&
t,
std::string name = "T_" #Name, std::string tag =
kElementWise) { \
auto l = [](tvm::PrimExpr a, tvm::PrimExpr b) { ComputeRule; };
\
return tvm::te::compute(
\
- A->shape, [&](const ::tvm::ffi::Array<::tvm::tirx::PrimVar>& i) {
return l(A(i), B); }, \
- name, tag);
\
+ A->shape, [&](const ::tvm::ffi::Array<::tvm::PrimVar>& i) { return
l(A(i), B); }, name, \
+ tag);
\
}
\
inline tvm::te::Tensor Name(const tvm::PrimExpr& A, const tvm::te::Tensor&
B, \
std::string name = "T_" #Name, std::string tag =
kElementWise) { \
auto l = [&](tvm::PrimExpr a, tvm::PrimExpr b) { ComputeRule; };
\
return tvm::te::compute(
\
- B->shape, [&](const ::tvm::ffi::Array<::tvm::tirx::PrimVar>& i) {
return l(A, B(i)); }, \
- name, tag);
\
+ B->shape, [&](const ::tvm::ffi::Array<::tvm::PrimVar>& i) { return
l(A, B(i)); }, name, \
+ tag);
\
}
#define TOPI_DEFINE_OP_OVERLOAD(Name, OpName)
\
diff --git a/include/tvm/topi/detail/broadcast.h
b/include/tvm/topi/detail/broadcast.h
index 9e388fb963..2d6b873bcd 100644
--- a/include/tvm/topi/detail/broadcast.h
+++ b/include/tvm/topi/detail/broadcast.h
@@ -37,9 +37,9 @@ namespace detail {
struct BroadcastHelper {
std::deque<tvm::PrimExpr> common_shape;
- std::deque<tvm::tirx::PrimVar> all_vars;
- std::deque<tvm::tirx::PrimVar> vars1;
- std::deque<tvm::tirx::PrimVar> vars2;
+ std::deque<tvm::PrimVar> all_vars;
+ std::deque<tvm::PrimVar> vars1;
+ std::deque<tvm::PrimVar> vars2;
};
static inline PrimType CommonType(const PrimType& type1, const PrimType&
type2) {
@@ -66,7 +66,7 @@ inline BroadcastHelper BroadcastShape(const
tvm::ffi::Array<tvm::PrimExpr>& shap
const IntImmNode* static_size2 = shape2[s2_size - i].as<IntImmNode>();
PrimType common_type = CommonType(shape1[s1_size - i].ty(), shape2[s2_size
- i].ty());
- bh.all_vars.push_front(tvm::tirx::PrimVar("dim", common_type));
+ bh.all_vars.push_front(tvm::PrimVar("dim", common_type));
if (topi::detail::EqualCheck(shape1[s1_size - i], shape2[s2_size - i])) {
bh.common_shape.push_front(cast_if_needed(common_type, shape1[s1_size -
i]));
bh.vars1.push_front(bh.all_vars[0]);
@@ -104,7 +104,7 @@ inline BroadcastHelper BroadcastShape(const
tvm::ffi::Array<tvm::PrimExpr>& shap
auto& shape = (s1_size > s2_size) ? shape1 : shape2;
auto& vars = (s1_size > s2_size) ? bh.vars1 : bh.vars2;
for (; i <= max_size; ++i) {
- bh.all_vars.push_front(tvm::tirx::PrimVar("v", shape[max_size - 1].ty()));
+ bh.all_vars.push_front(tvm::PrimVar("v", shape[max_size - 1].ty()));
bh.common_shape.push_front(shape[max_size - i]);
vars.push_front(bh.all_vars[0]);
}
@@ -112,8 +112,8 @@ inline BroadcastHelper BroadcastShape(const
tvm::ffi::Array<tvm::PrimExpr>& shap
}
inline tvm::ffi::Array<tvm::PrimExpr> InputIndexFromBroadcast(
- const tvm::ffi::Array<tvm::tirx::PrimVar>& ovars, const tvm::te::Tensor& T,
- const std::deque<tvm::tirx::PrimVar>& my_vars, const
std::deque<tvm::tirx::PrimVar>& all_vars) {
+ const tvm::ffi::Array<tvm::PrimVar>& ovars, const tvm::te::Tensor& T,
+ const std::deque<tvm::PrimVar>& my_vars, const std::deque<tvm::PrimVar>&
all_vars) {
tvm::ffi::Array<tvm::PrimExpr> ivars;
TVM_FFI_ICHECK_EQ(ovars.size(), all_vars.size());
// N^2, could use a map but NBD.
@@ -142,7 +142,7 @@ inline tvm::te::Tensor WithBroadcast(FBinaryExpr op, const
tvm::te::Tensor& A,
const tvm::te::Tensor& B, const
std::string& name = "tensor",
const std::string& tag = "") {
auto bh = BroadcastShape(A->shape, B->shape);
- auto l = [&](tvm::ffi::Array<tvm::tirx::PrimVar> ovars) {
+ auto l = [&](tvm::ffi::Array<tvm::PrimVar> ovars) {
return op(A(InputIndexFromBroadcast(ovars, A, bh.vars1, bh.all_vars)),
B(InputIndexFromBroadcast(ovars, B, bh.vars2, bh.all_vars)));
};
diff --git a/include/tvm/topi/detail/strided_slice.h
b/include/tvm/topi/detail/strided_slice.h
index 29f8257c82..d5a9b2f854 100644
--- a/include/tvm/topi/detail/strided_slice.h
+++ b/include/tvm/topi/detail/strided_slice.h
@@ -142,7 +142,7 @@ inline ffi::Array<PrimExpr> StridedSliceOutputShape(
<< ": Input [Begin=" << begin[i] << ", End=" << end[i] << "] is
invalid for axis=" << i;
out_shape.Set(ax, cast(out_shape[i].ty(), PrimExpr(slice_size)));
} else {
- out_shape.Set(ax, tvm::tirx::PrimVar("dim", out_shape[i].ty()));
+ out_shape.Set(ax, tvm::PrimVar("dim", out_shape[i].ty()));
}
}
diff --git a/include/tvm/topi/nn.h b/include/tvm/topi/nn.h
index a6b20151cd..c7a778cc75 100644
--- a/include/tvm/topi/nn.h
+++ b/include/tvm/topi/nn.h
@@ -56,7 +56,7 @@ inline tvm::te::Tensor relu(const tvm::te::Tensor& t, T
threshold = static_cast<
std::string name = "T_relu", std::string tag =
kElementWise) {
return tvm::te::compute(
t->shape,
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& i) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& i) {
auto threshold_const = tvm::tirx::MakeConst(tvm::PrimType(t->dtype),
threshold);
return tvm::max(t(i), threshold_const);
},
@@ -78,7 +78,7 @@ inline tvm::te::Tensor leaky_relu(const tvm::te::Tensor& t,
double alpha = 0.1,
std::string tag = kElementWise) {
return tvm::te::compute(
t->shape,
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& i) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& i) {
auto value = t(i);
auto calpha = tvm::tirx::MakeConst(value.ty(), alpha);
return tvm::prim::Select(value > 0, value, value * calpha);
@@ -107,7 +107,7 @@ inline tvm::te::Tensor prelu(const tvm::te::Tensor& x,
const tvm::te::Tensor& sl
return tvm::te::compute(
x->shape,
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& indices) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& indices) {
auto xval = x(indices);
return tvm::prim::Select(xval > 0, xval, xval * slope(indices[axis]));
},
@@ -197,7 +197,7 @@ inline tvm::te::Tensor pad(
pad_value = tvm::tirx::MakeConst(tvm::PrimType(t->dtype), 0);
}
- auto l = [&](tvm::ffi::Array<tvm::tirx::PrimVar> ovars) {
+ auto l = [&](tvm::ffi::Array<tvm::PrimVar> ovars) {
tvm::ffi::Array<tvm::PrimExpr> indices;
tvm::ffi::Array<tvm::PrimExpr> sel;
tvm::ffi::Array<tvm::PrimExpr> pad_idx;
@@ -286,8 +286,7 @@ inline tvm::te::Tensor conv2d_nchw(const tvm::te::Tensor&
I, const tvm::te::Tens
auto kw = tvm::te::reduce_axis(tvm::Range{0, W->shape[3]}, "kw");
auto T =
(pad_h == 0 && pad_w == 0) ? I : pad(I, {tvm::PrimExpr(0),
tvm::PrimExpr(0), pad_h, pad_w});
- auto l = [&](tvm::tirx::PrimVar b, tvm::tirx::PrimVar o, tvm::tirx::PrimVar
h,
- tvm::tirx::PrimVar w) {
+ auto l = [&](tvm::PrimVar b, tvm::PrimVar o, tvm::PrimVar h, tvm::PrimVar w)
{
return tvm::sum(T(b, i, stride_h * h + kh, stride_w * w + kw) * W(o, i,
kh, kw), {i, kh, kw});
};
return tvm::te::compute(output_shape, l, name, tag);
@@ -330,8 +329,7 @@ inline tvm::te::Tensor conv2d_hwcn(const tvm::te::Tensor&
I, const tvm::te::Tens
auto kh = tvm::te::reduce_axis(tvm::Range{0, W->shape[0]}, "kh");
auto kw = tvm::te::reduce_axis(tvm::Range{0, W->shape[1]}, "kw");
auto T = (pad_h == 0 && pad_w == 0) ? I : pad(I, {pad_h, pad_w});
- auto l = [&](tvm::tirx::PrimVar b, tvm::tirx::PrimVar o, tvm::tirx::PrimVar
h,
- tvm::tirx::PrimVar w) {
+ auto l = [&](tvm::PrimVar b, tvm::PrimVar o, tvm::PrimVar h, tvm::PrimVar w)
{
return tvm::sum(T(stride_h * h + kh, stride_w * w + kw, i, b) * W(kh, kw,
i, o), {i, kh, kw});
};
return tvm::te::compute(output_shape, l, name, tag);
@@ -378,8 +376,7 @@ inline tvm::te::Tensor depthwise_conv2d_nchw(const
tvm::te::Tensor& I, const tvm
auto kw = tvm::te::reduce_axis(tvm::Range{0, W->shape[3]}, "kw");
auto T =
(pad_h == 0 && pad_w == 0) ? I : pad(I, {tvm::PrimExpr(0),
tvm::PrimExpr(0), pad_h, pad_w});
- auto l = [&](tvm::tirx::PrimVar b, tvm::tirx::PrimVar o, tvm::tirx::PrimVar
h,
- tvm::tirx::PrimVar w) {
+ auto l = [&](tvm::PrimVar b, tvm::PrimVar o, tvm::PrimVar h, tvm::PrimVar w)
{
return tvm::sum(T(b, indexdiv(i, pCM), stride_h * h + kh, stride_w * w +
kw) *
W(indexdiv(i, pCM), indexmod(o, pCM), kh, kw),
{i, kh, kw});
@@ -408,8 +405,7 @@ inline tvm::te::Tensor depthwise_conv2d_nhwc(const
tvm::te::Tensor& I, const tvm
auto kw = tvm::te::reduce_axis(tvm::Range{0, W->shape[1]}, "kw");
auto T =
(pad_h == 0 && pad_w == 0) ? I : pad(I, {tvm::PrimExpr(0), pad_h, pad_w,
tvm::PrimExpr(0)});
- auto l = [&](tvm::tirx::PrimVar b, tvm::tirx::PrimVar h, tvm::tirx::PrimVar
w,
- tvm::tirx::PrimVar o) {
+ auto l = [&](tvm::PrimVar b, tvm::PrimVar h, tvm::PrimVar w, tvm::PrimVar o)
{
return tvm::sum(T(b, stride_h * h + kh, stride_w * w + kw, indexdiv(i,
pCM)) *
W(kh, kw, indexdiv(i, pCM), indexmod(o, pCM)),
{kh, kw, i});
@@ -460,12 +456,12 @@ inline tvm::te::Tensor group_conv2d_ngchw(const
tvm::te::Tensor& I, const tvm::t
auto T = (pad_h == 0 && pad_w == 0)
? I
: pad(I, {tvm::PrimExpr(0), tvm::PrimExpr(0), tvm::PrimExpr(0),
pad_h, pad_w});
- auto l = [&](tvm::ffi::Array<tvm::tirx::PrimVar> args) {
- tvm::tirx::PrimVar b = args[0];
- tvm::tirx::PrimVar g = args[1];
- tvm::tirx::PrimVar o = args[2];
- tvm::tirx::PrimVar h = args[3];
- tvm::tirx::PrimVar w = args[4];
+ auto l = [&](tvm::ffi::Array<tvm::PrimVar> args) {
+ tvm::PrimVar b = args[0];
+ tvm::PrimVar g = args[1];
+ tvm::PrimVar o = args[2];
+ tvm::PrimVar h = args[3];
+ tvm::PrimVar w = args[4];
return tvm::sum(I(b, g, i, stride_h * h + kh, stride_w * w + kw) * W(g, i,
o, kh, kw),
{i, kh, kw});
};
@@ -679,7 +675,7 @@ inline Tensor nll_loss(const Tensor& predictions, const
Tensor& targets, const T
// prediction->shape = (C,), targets->shape = (), weights->shape = (C,)
auto T = tvm::te::compute(
{},
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& target_indices) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& target_indices) {
auto c = targets();
return tvm::prim::Select(c != ignore_index, -predictions(c) *
weights(c),
tvm::tirx::MakeConst(tvm::PrimType(predictions->dtype), 0));
@@ -688,7 +684,7 @@ inline Tensor nll_loss(const Tensor& predictions, const
Tensor& targets, const T
if (reduction == "mean") {
auto W = tvm::te::compute(
{},
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& target_indices) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& target_indices) {
auto c = targets();
return tvm::prim::Select(c != ignore_index, weights(c),
tvm::tirx::MakeConst(tvm::PrimType(predictions->dtype), 0));
@@ -701,7 +697,7 @@ inline Tensor nll_loss(const Tensor& predictions, const
Tensor& targets, const T
}
auto T = tvm::te::compute(
targets->shape,
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& target_indices) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& target_indices) {
auto c = targets(target_indices);
tvm::ffi::Array<tvm::PrimExpr> pred_indices;
pred_indices.push_back(target_indices[0]); // batch index
@@ -717,7 +713,7 @@ inline Tensor nll_loss(const Tensor& predictions, const
Tensor& targets, const T
if (reduction == "mean") {
auto W = tvm::te::compute(
targets->shape,
- [&](const tvm::ffi::Array<tvm::tirx::PrimVar>& target_indices) {
+ [&](const tvm::ffi::Array<tvm::PrimVar>& target_indices) {
auto c = targets(target_indices);
return tvm::prim::Select(c != ignore_index, weights(c),
tvm::tirx::MakeConst(tvm::PrimType(predictions->dtype), 0));
diff --git a/include/tvm/topi/transform.h b/include/tvm/topi/transform.h
index 7092eaa42c..6ac02f7844 100644
--- a/include/tvm/topi/transform.h
+++ b/include/tvm/topi/transform.h
@@ -738,7 +738,7 @@ inline te::Tensor dynamic_strided_slice_with_axes(
return te::compute(
out_shape,
- [&](const ffi::Array<tvm::tirx::PrimVar>& indices) {
+ [&](const ffi::Array<tvm::PrimVar>& indices) {
ffi::Array<PrimExpr> real_indices =
indices.Map([](const auto& var) -> PrimExpr { return var; });
@@ -793,7 +793,7 @@ inline Tensor dynamic_strided_slice(const Tensor& x, const
ffi::Array<PrimExpr>&
out_shape.push_back(
analyzer->Simplify(GetLength(begin[i], end[i], strides[i],
x->shape[i], assume_inbound)));
} else {
- out_shape.push_back(tvm::tirx::PrimVar("dim"));
+ out_shape.push_back(tvm::PrimVar("dim"));
}
}
@@ -803,7 +803,7 @@ inline Tensor dynamic_strided_slice(const Tensor& x, const
ffi::Array<PrimExpr>&
return te::compute(
out_shape,
- [&](const ffi::Array<tvm::tirx::PrimVar>& indices) {
+ [&](const ffi::Array<tvm::PrimVar>& indices) {
ffi::Array<PrimExpr> real_indices;
for (size_t i = 0; i < num_slice_axes; ++i) {
PrimExpr begin_index = tvm::min(begin[i], x->shape[i] - 1);
@@ -938,7 +938,7 @@ inline Tensor strided_slice_with_axes(
return te::compute(
out_shape,
- [&](const ffi::Array<tirx::PrimVar>& indices) {
+ [&](const ffi::Array<PrimVar>& indices) {
ffi::Array<PrimExpr> real_indices;
for (size_t i = 0; i < out_shape.size(); ++i)
real_indices.push_back(indices[i]);
for (size_t i = 0; i < normalized_axes.size(); ++i) {
@@ -1354,7 +1354,7 @@ inline Tensor where(const Tensor& condition, const
Tensor& x, const Tensor& y,
auto x_bh = detail::BroadcastShape(x->shape, oshape);
auto y_bh = detail::BroadcastShape(y->shape, oshape);
- auto select = [&](tvm::ffi::Array<tvm::tirx::PrimVar> ovars) {
+ auto select = [&](tvm::ffi::Array<tvm::PrimVar> ovars) {
auto c = condition(InputIndexFromBroadcast(ovars, condition, c_bh.vars1,
c_bh.all_vars));
auto true_val = x(InputIndexFromBroadcast(ovars, x, x_bh.vars1,
x_bh.all_vars));
auto false_val = y(InputIndexFromBroadcast(ovars, y, y_bh.vars1,
y_bh.all_vars));
@@ -1643,7 +1643,7 @@ inline tvm::te::Tensor matmul(const tvm::te::Tensor& A,
const tvm::te::Tensor& B
std::string name = "T_matmul", std::string tag =
kMatMul) {
tvm::ffi::Array<tvm::PrimExpr> output_shape{A->shape[trans_a ? 1 : 0],
B->shape[trans_b ? 0 : 1]};
auto k = tvm::te::reduce_axis(tvm::Range{0, A->shape[trans_a ? 0 : 1]}, "k");
- auto l = [&](tvm::tirx::PrimVar i, tvm::tirx::PrimVar j) {
+ auto l = [&](tvm::PrimVar i, tvm::PrimVar j) {
return tvm::sum((trans_a ? A[k][i] : A[i][k]) * (trans_b ? B[j][k] :
B[k][j]), {k});
};
return tvm::te::compute(output_shape, l, name, tag);
@@ -2320,7 +2320,7 @@ inline te::Tensor dynamic_strided_slice(const te::Tensor&
x, const te::Tensor& b
return te::compute(
output_shape,
- [&](const ffi::Array<tvm::tirx::PrimVar>& indices) {
+ [&](const ffi::Array<tvm::PrimVar>& indices) {
ffi::Array<PrimExpr> real_indices;
for (size_t i = 0; i < num_dynamic_axes; ++i) {
auto ind = IntImm::Int64(i);
diff --git a/src/arith/analyzer.cc b/src/arith/analyzer.cc
index 31f4087b00..3819017b46 100644
--- a/src/arith/analyzer.cc
+++ b/src/arith/analyzer.cc
@@ -107,7 +107,7 @@ void AnalyzerObj::MarkGlobalNonNegValue(const PrimExpr&
value) {
//
// We may consider enhance the sub analyzer to directly take
// MarkPositiveVar so their bounds do not overlap
- if (auto prim_var = symbol.as<tirx::PrimVar>()) {
+ if (auto prim_var = symbol.as<PrimVar>()) {
Var var = *prim_var;
// skip non-index type, keep it to be compatible
// with any_dim that do not represent any value
diff --git a/src/arith/int_set.cc b/src/arith/int_set.cc
index 58f289540d..dd39c34775 100644
--- a/src/arith/int_set.cc
+++ b/src/arith/int_set.cc
@@ -51,8 +51,8 @@ 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));
+PrimExpr SymbolicLimits::pos_inf_ = PrimVar("pos_inf", PrimType::Int(64));
+PrimExpr SymbolicLimits::neg_inf_ = PrimVar("neg_inf", PrimType::Int(64));
IntervalSet::IntervalSet(PrimExpr min_value, PrimExpr max_value) {
auto node = ffi::make_object<IntervalSetNode>();
diff --git a/src/arith/pattern_match.h b/src/arith/pattern_match.h
index 7152a62475..680afcc947 100644
--- a/src/arith/pattern_match.h
+++ b/src/arith/pattern_match.h
@@ -44,7 +44,7 @@
* return (max(x, y) + z).Eval();
* }
*
- * tvm::tirx::Var tx, ty;
+ * tvm::Var tx, ty;
* arith::PVar<IntImm> c;
* arith::PVar<Var> v;
* // We can match integer and Var, both of which are
@@ -179,9 +179,9 @@ class PEqualChecker<FloatImm> {
};
template <>
-class PEqualChecker<tirx::Var> {
+class PEqualChecker<Var> {
public:
- bool operator()(const tirx::Var& lhs, const tirx::Var& rhs) const { return
lhs.same_as(rhs); }
+ bool operator()(const Var& lhs, const Var& rhs) const { return
lhs.same_as(rhs); }
};
/*!
diff --git a/src/arith/presburger_set.h b/src/arith/presburger_set.h
index b56d9c09f2..61a39e710b 100644
--- a/src/arith/presburger_set.h
+++ b/src/arith/presburger_set.h
@@ -61,10 +61,10 @@ using namespace presburger;
class PresburgerSetNode : public IntSetNode {
public:
PresburgerSetNode() : space(PresburgerSpace::getRelationSpace()) {}
- explicit PresburgerSetNode(const PresburgerSpace& space, const
ffi::Array<tirx::PrimVar>& vars)
+ explicit PresburgerSetNode(const PresburgerSpace& space, const
ffi::Array<PrimVar>& vars)
: disjuncts({}), space(space), vars(vars) {}
explicit PresburgerSetNode(const std::vector<IntegerRelation>& disjuncts,
- const PresburgerSpace& space, const
ffi::Array<tirx::PrimVar>& vars)
+ const PresburgerSpace& space, const
ffi::Array<PrimVar>& vars)
: disjuncts(disjuncts), space(space), vars(vars) {}
/*! \brief Represent the union of multiple IntegerRelation */
@@ -92,7 +92,7 @@ class PresburgerSetNode : public IntSetNode {
* \param constraint The added constraint to the PresburgerSet.
* \param vars The specified domain vars in constraint expression.
*/
- void UpdateConstraint(const PrimExpr& constraint, const
ffi::Array<tirx::PrimVar>& vars);
+ void UpdateConstraint(const PrimExpr& constraint, const ffi::Array<PrimVar>&
vars);
/*!
* \brief Generate expression that represents the constraint
@@ -104,13 +104,13 @@ class PresburgerSetNode : public IntSetNode {
* \brief Set domain vars
* \param new_vars Vars that will be taken as the domain vars
*/
- void SetVars(const ffi::Array<tirx::PrimVar>& new_vars) { vars = new_vars; }
+ void SetVars(const ffi::Array<PrimVar>& new_vars) { vars = new_vars; }
/*!
* \brief Get the current domain vars
* \return The current doamin vars
*/
- ffi::Array<tirx::PrimVar> GetVars() const { return vars; }
+ ffi::Array<PrimVar> GetVars() const { return vars; }
/*! \return whether integer set is empty */
bool IsEmpty() const {
@@ -120,7 +120,7 @@ class PresburgerSetNode : public IntSetNode {
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("arith.PresburgerSet", PresburgerSetNode,
IntSetNode);
private:
- ffi::Array<tirx::PrimVar> vars;
+ ffi::Array<PrimVar> vars;
};
/*!
@@ -136,7 +136,7 @@ class PresburgerSet : public IntSet {
* \return The created PresburgerSet.
*/
TVM_DLL PresburgerSet(const std::vector<IntegerRelation>& disjuncts,
- const ffi::Array<tirx::PrimVar>& vars);
+ const ffi::Array<PrimVar>& vars);
/*!
* \brief Make a new instance of PresburgerSet, collect all vars as space
vars.
diff --git a/src/arith/solve_linear_equation.cc
b/src/arith/solve_linear_equation.cc
index 6f1d2625d9..40e5e0883e 100644
--- a/src/arith/solve_linear_equation.cc
+++ b/src/arith/solve_linear_equation.cc
@@ -393,7 +393,7 @@ IntConstraintsTransform SolveLinearEquations(const
IntConstraints& system_to_sol
// The j-th variable can take any integer value, create a tvm variable
for it
PrimExpr to_old = analyzer_problem->Simplify(V_inv_x[j]);
std::string name_hint = "n" + std::to_string(new_vars.size());
- if (auto old_var = to_old.as<tirx::PrimVar>()) {
+ if (auto old_var = to_old.as<PrimVar>()) {
name_hint += "_" + (*old_var)->name;
}
PrimVar v(name_hint, V_inv_x[j].ty());
diff --git a/src/arith/transitive_comparison_analyzer.cc
b/src/arith/transitive_comparison_analyzer.cc
index d8687ecdb5..73eb597728 100644
--- a/src/arith/transitive_comparison_analyzer.cc
+++ b/src/arith/transitive_comparison_analyzer.cc
@@ -63,7 +63,7 @@ class TransitiveComparisonAnalyzer::Impl {
* \param expr The bound expression
* \param allow_override Whether to allow override of existing information.
*/
- void Bind(const tirx::Var& var, const PrimExpr& expr, bool allow_override =
false);
+ void Bind(const Var& var, const PrimExpr& expr, bool allow_override = false);
/*! \brief Bind a variable as being within a specified range
*
@@ -71,7 +71,7 @@ class TransitiveComparisonAnalyzer::Impl {
* \param range The known range
* \param allow_override Whether to allow override of existing information.
*/
- void Bind(const tirx::Var& var, const Range& expr, bool allow_override =
false);
+ void Bind(const Var& var, const Range& expr, bool allow_override = false);
/*!
* \brief Update the internal state to enter constraint.
@@ -566,7 +566,7 @@ void TransitiveComparisonAnalyzer::Impl::AddKnown(const
PrimExpr& expr,
}
}
-void TransitiveComparisonAnalyzer::Impl::Bind(const tirx::Var& var, const
Range& range,
+void TransitiveComparisonAnalyzer::Impl::Bind(const Var& var, const Range&
range,
bool allow_override) {
auto it = prev_bindings_.find(var);
if (it != prev_bindings_.end()) {
@@ -595,7 +595,7 @@ void TransitiveComparisonAnalyzer::Impl::Bind(const
tirx::Var& var, const Range&
}
}
-void TransitiveComparisonAnalyzer::Impl::Bind(const tirx::Var& var, const
PrimExpr& expr,
+void TransitiveComparisonAnalyzer::Impl::Bind(const Var& var, const PrimExpr&
expr,
bool allow_override) {
Bind(var, Range::FromMinExtent(expr, 1), allow_override);
}
diff --git a/src/backend/cuda/codegen/codegen_cuda.cc
b/src/backend/cuda/codegen/codegen_cuda.cc
index 6c59a0515f..428f860719 100644
--- a/src/backend/cuda/codegen/codegen_cuda.cc
+++ b/src/backend/cuda/codegen/codegen_cuda.cc
@@ -1395,7 +1395,7 @@ void CodeGenCUDA::VisitExpr_(const CallNode* op,
std::ostream& os) {
// and finally reinterpret the result as fp4x2.
value =
Call(PrimType::UInt(16), tirx::builtin::reinterpret(),
{value}).as_or_throw<PrimExpr>();
- tirx::PrimVar temp_var("temp_var", PrimType::UInt(16));
+ PrimVar temp_var("temp_var", PrimType::UInt(16));
value = prim::Let(temp_var, value,
prim::Cast(PrimType::UInt(8),
(temp_var & IntImm(PrimType::UInt(16),
0xF)) |
@@ -1404,7 +1404,7 @@ void CodeGenCUDA::VisitExpr_(const CallNode* op,
std::ostream& os) {
value = prim::Cast(
PrimType::UInt(16),
Call(PrimType::UInt(8), tirx::builtin::reinterpret(),
{value}).as_or_throw<PrimExpr>());
- tirx::PrimVar temp_var("temp_var", PrimType::UInt(16));
+ PrimVar temp_var("temp_var", PrimType::UInt(16));
value = prim::Let(temp_var, value,
(temp_var & IntImm(PrimType::UInt(16), 0xF)) |
((temp_var & IntImm(PrimType::UInt(16), 0xF0))
<< 4));
@@ -1416,7 +1416,7 @@ void CodeGenCUDA::VisitExpr_(const CallNode* op,
std::ostream& os) {
// and finally reinterpret the result as fp4x4.
value =
Call(PrimType::UInt(32), tirx::builtin::reinterpret(),
{value}).as_or_throw<PrimExpr>();
- tirx::PrimVar temp_var("temp_var", PrimType::UInt(32));
+ PrimVar temp_var("temp_var", PrimType::UInt(32));
value = prim::Let(temp_var, value,
prim::Cast(PrimType::UInt(16),
(temp_var & IntImm(PrimType::UInt(32),
0xF)) |
@@ -1427,7 +1427,7 @@ void CodeGenCUDA::VisitExpr_(const CallNode* op,
std::ostream& os) {
value = prim::Cast(PrimType::UInt(32),
Call(PrimType::UInt(16),
tirx::builtin::reinterpret(), {value})
.as_or_throw<PrimExpr>());
- tirx::PrimVar temp_var("temp_var", PrimType::UInt(32));
+ PrimVar temp_var("temp_var", PrimType::UInt(32));
value = prim::Let(temp_var, value,
(temp_var & IntImm(PrimType::UInt(32), 0xF)) |
((temp_var & IntImm(PrimType::UInt(32), 0xF0))
<< 4) |
diff --git a/src/relax/analysis/layout_transformation.cc
b/src/relax/analysis/layout_transformation.cc
index 06757527f5..acde9f9f95 100644
--- a/src/relax/analysis/layout_transformation.cc
+++ b/src/relax/analysis/layout_transformation.cc
@@ -41,7 +41,7 @@ using namespace tirx;
/*! \brief Checks if a transformation is bijective affine over the given
ranges */
static bool IsBijectiveAffine(const IndexMap& m, const ffi::Array<Range>&
ranges) {
- ffi::Map<tirx::PrimVar, Range> input_iters;
+ ffi::Map<PrimVar, Range> input_iters;
TVM_FFI_ICHECK_EQ(m->initial_indices.size(), ranges.size());
for (size_t i = 0; i < ranges.size(); i++) {
input_iters.Set(m->initial_indices[i], ranges[i]);
@@ -85,7 +85,7 @@ class IndexAnalyzer : public ExprVisitor {
}
void VisitIterMark(const arith::IterMark& op) {
- if (auto var = op->source.as<tirx::PrimVar>())
+ if (auto var = op->source.as<PrimVar>())
iterators_.push_back(var.value());
else
VisitExpr(op->source);
@@ -298,7 +298,7 @@ static ffi::Optional<IndexMap>
InferLayoutTransformation(const SpatialLayout& sr
continue;
}
- tirx::PrimVar new_dim("d");
+ PrimVar new_dim("d");
PrimExpr new_dim_expr = new_dim;
initial_indices_it = initial_indices.insert(initial_indices_it, new_dim);
final_indices_it = final_indices.insert(final_indices_it, new_dim_expr);
@@ -309,7 +309,7 @@ static ffi::Optional<IndexMap>
InferLayoutTransformation(const SpatialLayout& sr
ffi::Array<tirx::Var> initial_array(initial_indices.begin(),
initial_indices.end());
ffi::Array<PrimExpr> final_array(final_indices.begin(), final_indices.end());
- return IndexMap(initial_array.Map([](tirx::Var var) { return
var.as_or_throw<tirx::PrimVar>(); }),
+ return IndexMap(initial_array.Map([](tirx::Var var) { return
var.as_or_throw<PrimVar>(); }),
final_array);
}
@@ -534,7 +534,7 @@ class BlockAnalyzer : public StmtExprVisitor {
private:
bool can_transform_block_;
IndexMap write_transformation_;
- ffi::Map<tirx::PrimVar, Range> spatial_dom_;
+ ffi::Map<PrimVar, Range> spatial_dom_;
arith::Analyzer arith_analyzer_;
SBlock block_;
diff --git a/src/relax/analysis/tir_op_pattern_kind.cc
b/src/relax/analysis/tir_op_pattern_kind.cc
index 3e7531d373..7b70bf9fd2 100644
--- a/src/relax/analysis/tir_op_pattern_kind.cc
+++ b/src/relax/analysis/tir_op_pattern_kind.cc
@@ -231,7 +231,7 @@ class PatternKindAnalyzer : public StmtExprVisitor {
static bool IsInjectivePattern(const BufferStore& store, const TensorLoad&
load) {
std::unordered_set<const tirx::VarNode*> vars;
for (const PrimExpr& store_index : store->indices) {
- if (auto var = store_index.as<tirx::PrimVar>()) {
+ if (auto var = store_index.as<PrimVar>()) {
vars.insert(var.value().get());
} else {
return false;
@@ -259,14 +259,14 @@ class PatternKindAnalyzer : public StmtExprVisitor {
static bool IsAllowReusePattern(const BufferStore& store, const TensorLoad&
load) {
std::unordered_set<const tirx::VarNode*> vars;
for (const PrimExpr& index : store->indices) {
- if (auto var = index.as<tirx::PrimVar>()) {
+ if (auto var = index.as<PrimVar>()) {
vars.insert(var.value().get());
} else {
return false;
}
}
auto walk_fn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult>
{
- if (auto prim_var = var.as<tirx::PrimVar>()) {
+ if (auto prim_var = var.as<PrimVar>()) {
vars.erase(prim_var.value().get());
}
return ffi::WalkResult::Advance();
@@ -417,7 +417,7 @@ bool HasReshapePattern(const PrimFunc& func) {
return;
}
- ffi::Map<tirx::PrimVar, Range> var_range;
+ ffi::Map<PrimVar, Range> var_range;
for (const IterVar& v : block->iter_vars) {
ana_->Bind(v->var, Range::FromMinExtent(v->dom->min, v->dom->extent));
var_range.Set(v->var, Range::FromMinExtent(v->dom->min,
v->dom->extent));
@@ -497,7 +497,7 @@ bool HasReshapePattern(const PrimFunc& func) {
if (nontrivial_indices.defined() && !has_zero_extent) {
PrimType dtype =
!block->iter_vars.empty() ? block->iter_vars[0]->var.ty() :
PrimType::Int(64);
- tirx::PrimVar fused_var("fused", dtype);
+ PrimVar fused_var("fused", dtype);
ffi::Map<tirx::Var, PrimExpr> inverse_indices_map;
PrimExpr stride = IntImm(dtype, /*value=*/1);
for (int i = static_cast<int>(block->iter_vars.size()) - 1; i >= 0;
--i) {
@@ -519,7 +519,7 @@ bool HasReshapePattern(const PrimFunc& func) {
ffi::Array<PrimExpr> simplify_res = arith::IterMapSimplify(
/*indices=*/{flattened_idx},
/*input_iters=*/
- ffi::Map<tirx::PrimVar, Range>{{fused_var, Range(IntImm(dtype,
/*value=*/0), stride)}},
+ ffi::Map<PrimVar, Range>{{fused_var, Range(IntImm(dtype,
/*value=*/0), stride)}},
/*input_pred=*/IntImm::Bool(true),
/*check_level=*/arith::IterMapLevel::Surjective,
/*analyzer=*/this->ana_,
diff --git a/src/relax/analysis/type_analysis.cc
b/src/relax/analysis/type_analysis.cc
index 773b1423a8..a599eb61c7 100644
--- a/src/relax/analysis/type_analysis.cc
+++ b/src/relax/analysis/type_analysis.cc
@@ -286,7 +286,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
[](const Type& info, ffi::Map<tirx::Var, PrimExpr>
shape_var_map,
ffi::Map<Var, Expr> var_map) {
for (const auto& [var, value] : shape_var_map) {
- TVM_FFI_CHECK(var.as<tirx::PrimVar>(), TypeError)
+ TVM_FFI_CHECK(var.as<PrimVar>(), TypeError)
<< "Expected an exact primitive Var, but
received " << var;
var_map.Set(var, value);
}
@@ -882,7 +882,7 @@ class CallRetTypeDeriver : public TypeBaseChecker {
return TypeBaseChecker::PrimExprMatchCheck(param, arg);
}
- if (auto var = param.as<tirx::PrimVar>()) {
+ if (auto var = param.as<PrimVar>()) {
auto it = var_map_.find(var.value());
// not populated
if (it == var_map_.end()) {
@@ -908,8 +908,7 @@ class CallRetTypeDeriver : public TypeBaseChecker {
return TypeBaseChecker::ShapeMatchCheck(lhs, rhs);
}
- if (auto* ptr = lhs.as<VarNode>();
- ptr && !lhs.as<DataflowVarNode>() && !lhs.as<tirx::PrimVar>()) {
+ if (auto* ptr = lhs.as<VarNode>(); ptr && !lhs.as<DataflowVarNode>() &&
!lhs.as<PrimVar>()) {
auto var = ffi::GetRef<Var>(ptr);
auto it = var_map_.find(var);
// not populated
@@ -1184,12 +1183,12 @@ class TIRVarsDetector : public TypeVisitor {
private:
void VisitTypePrimExprField(PrimExpr expr) {
if (collection_type == VarType::Definition) {
- if (auto opt = expr.as<tirx::PrimVar>()) {
+ if (auto opt = expr.as<PrimVar>()) {
RecordTIRVar(opt.value());
}
} else if (collection_type == VarType::Usage) {
for (const tirx::Var& tir_var : tirx::UndefinedVars(expr)) {
- if (auto prim_var = tir_var.as<tirx::PrimVar>()) {
+ if (auto prim_var = tir_var.as<PrimVar>()) {
RecordTIRVar(prim_var.value());
}
}
@@ -1402,7 +1401,7 @@ class SymbolicVarCollector : public relax::ExprVisitor,
public relax::TypeVisito
void VisitTypeExprField(const PrimExpr& expr) final {
if (mode_ & VisitMode::kProvideDefinition) {
- if (auto var = expr.as<tirx::PrimVar>()) {
+ if (auto var = expr.as<PrimVar>()) {
defined_symbolic_var_.insert(var.value());
}
}
@@ -1415,7 +1414,7 @@ class SymbolicVarCollector : public relax::ExprVisitor,
public relax::TypeVisito
if (!op->ty.as<PrimTypeNode>()) {
return;
}
- tirx::PrimVar var = ffi::GetRef<Var>(op).as_or_throw<tirx::PrimVar>();
+ PrimVar var = ffi::GetRef<Var>(op).as_or_throw<PrimVar>();
// default mode, check defined.
if (defined_symbolic_var_.count(var) == 0) {
free_symbolic_var_.insert(var);
@@ -1426,7 +1425,7 @@ class SymbolicVarCollector : public relax::ExprVisitor,
public relax::TypeVisito
void VisitVarDef_(const VarNode* op) final {
if (op->ty.as<PrimTypeNode>()) {
-
defined_symbolic_var_.insert(ffi::GetRef<Var>(op).as_or_throw<tirx::PrimVar>());
+
defined_symbolic_var_.insert(ffi::GetRef<Var>(op).as_or_throw<PrimVar>());
}
relax::ExprVisitor::VisitVarDef_(op);
}
diff --git a/src/relax/backend/contrib/tensorrt/codegen.cc
b/src/relax/backend/contrib/tensorrt/codegen.cc
index 4d40ec9185..74fae94f7c 100644
--- a/src/relax/backend/contrib/tensorrt/codegen.cc
+++ b/src/relax/backend/contrib/tensorrt/codegen.cc
@@ -197,7 +197,7 @@ class CollectFromCompositeFunctionBody : public ExprVisitor
{
if (initial.size() != final_indices.size()) return true;
ffi::Array<int64_t> permutation;
for (const PrimExpr& expr : final_indices) {
- auto var = expr.as<tirx::PrimVar>();
+ auto var = expr.as<PrimVar>();
if (!var.has_value()) return true;
int64_t pos = -1;
for (size_t j = 0; j < initial.size(); ++j) {
diff --git a/src/relax/backend/vm/vm_shape_lower.cc
b/src/relax/backend/vm/vm_shape_lower.cc
index 09053a1df9..b30c83d293 100644
--- a/src/relax/backend/vm/vm_shape_lower.cc
+++ b/src/relax/backend/vm/vm_shape_lower.cc
@@ -136,7 +136,7 @@ class PrimExprSlotCollector : public ExprVisitor, public
TypeVisitor {
void VisitExpr_(const VarNode* op) final {
Var var = ffi::GetRef<Var>(op);
if (collect_scalar_ && !var.as<DataflowVarNode>()) {
- if (auto prim_var = var.as<tirx::PrimVar>();
+ if (auto prim_var = var.as<PrimVar>();
prim_var && prim_var.value().ty()->dtype == DLDataType{kDLInt, 64,
1}) {
HandlePrimExpr(prim_var.value());
}
@@ -301,7 +301,7 @@ class VMShapeLowerMutator
Expr VisitExpr_(const VarNode* op) final {
Var var = ffi::GetRef<Var>(op);
if (!var.as<DataflowVarNode>()) {
- if (auto prim_var = var.as<tirx::PrimVar>(); prim_var &&
slot_map_.count(*prim_var)) {
+ if (auto prim_var = var.as<PrimVar>(); prim_var &&
slot_map_.count(*prim_var)) {
return RewritePrimValue(*prim_var);
}
}
@@ -419,7 +419,7 @@ class VMShapeLowerMutator
PrimExprSlot* GetPrimValueSlot(const Var& var) const {
if (var.as<DataflowVarNode>()) return nullptr;
- auto prim_var = var.as<tirx::PrimVar>();
+ auto prim_var = var.as<PrimVar>();
if (!prim_var) return nullptr;
auto it = slot_map_.find(PrimExpr(*prim_var));
return it == slot_map_.end() ? nullptr : it->second;
diff --git a/src/relax/ir/block_builder.cc b/src/relax/ir/block_builder.cc
index 115e582f0d..7b9205dfc1 100644
--- a/src/relax/ir/block_builder.cc
+++ b/src/relax/ir/block_builder.cc
@@ -197,9 +197,9 @@ class BlockBuilderImpl : public BlockBuilderNode {
// defined in parameter type annotations. The implementation
// is correct (since we will simply erase all relax Vars in
EraseToWellDefined),
// but can be further improved.
- ffi::Map<tirx::PrimVar, PrimExpr> var_map =
TypeVarCollector::Collect(GetType(var));
+ ffi::Map<PrimVar, PrimExpr> var_map =
TypeVarCollector::Collect(GetType(var));
for (const auto& kv : var_map) {
- const tirx::PrimVar& shape_var = kv.first;
+ const PrimVar& shape_var = kv.first;
const PrimExpr& shape_expr = kv.second;
auto it = shape_var_map.find(shape_var);
if (it == shape_var_map.end()) {
@@ -333,7 +333,7 @@ class BlockBuilderImpl : public BlockBuilderNode {
//
// TODO(relax-team) tracks the var defined also through match-cast.
/*! \brief set of defined symbolic vars, value as themself. */
- ffi::Map<tirx::PrimVar, PrimExpr> shape_var_map;
+ ffi::Map<PrimVar, PrimExpr> shape_var_map;
};
/*! \brief A stack to store block frames. */
@@ -472,7 +472,7 @@ class BlockBuilderImpl : public BlockBuilderNode {
// shape vars as defined when calling BeginScope(params)
class TypeVarCollector : public TypeVisitor {
public:
- static ffi::Map<tirx::PrimVar, PrimExpr> Collect(const Type& ty) {
+ static ffi::Map<PrimVar, PrimExpr> Collect(const Type& ty) {
TypeVarCollector collector;
collector(ty);
return collector.shape_var_map_;
@@ -483,7 +483,7 @@ class BlockBuilderImpl : public BlockBuilderNode {
if (const auto* shape_expr = op->shape.as<ShapeExprNode>()) {
for (const PrimExpr& s : shape_expr->values) {
// Only collect single var defined shape. Ignore something like
`R.Tensor((m + 1, n + 1))
- if (auto var = s.as<tirx::PrimVar>()) {
+ if (auto var = s.as<PrimVar>()) {
shape_var_map_.Set(var.value(), s);
}
}
@@ -493,14 +493,14 @@ class BlockBuilderImpl : public BlockBuilderNode {
void VisitType_(const ShapeTypeNode* op) final {
for (const PrimExpr& s : op->values.value_or(ffi::Array<PrimExpr>())) {
// Only collect single var defined shape. Ignore something like
`R.Shape((m + 1, n + 1))
- if (auto var = s.as<tirx::PrimVar>()) {
+ if (auto var = s.as<PrimVar>()) {
shape_var_map_.Set(var.value(), s);
}
}
}
private:
- ffi::Map<tirx::PrimVar, PrimExpr> shape_var_map_;
+ ffi::Map<PrimVar, PrimExpr> shape_var_map_;
};
};
@@ -864,7 +864,7 @@ class Normalizer : public BlockBuilderImpl, private
ExprFunctor<Expr(const Expr&
}
auto* curr_scope = CurrentScopeFrame();
auto f_var_map = [curr_scope](const Var& var) -> ffi::Optional<Expr> {
- auto prim_var = var.as<tirx::PrimVar>();
+ auto prim_var = var.as<PrimVar>();
if (!prim_var) return std::nullopt;
auto it = curr_scope->shape_var_map.find(prim_var.value());
if (it != curr_scope->shape_var_map.end()) return (*it).second;
diff --git a/src/relax/ir/dataflow_matcher.cc b/src/relax/ir/dataflow_matcher.cc
index 8de1416d50..ad0b721b87 100644
--- a/src/relax/ir/dataflow_matcher.cc
+++ b/src/relax/ir/dataflow_matcher.cc
@@ -462,7 +462,7 @@ PrimExpr DFPatternMatcher::SimplifyCondition(PrimExpr
condition) {
auto sort_key = [](PrimExpr expr) -> ffi::String {
if (const auto* equal = expr.as<prim::EQNode>()) {
- if (auto var = equal->a.as<tirx::PrimVar>()) {
+ if (auto var = equal->a.as<PrimVar>()) {
return var.value()->name;
}
}
diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc
index 3ab6670fd5..fb40dc735f 100644
--- a/src/relax/ir/expr.cc
+++ b/src/relax/ir/expr.cc
@@ -731,7 +731,7 @@ Function::Function(ffi::Array<Var> params, Expr body,
ffi::Optional<Type> ret_ty
auto tir_vars = DefinableTIRVarsInType(TupleType(params.Map(GetType)));
std::unordered_set<tirx::Var> lookup(tir_vars.begin(), tir_vars.end());
return [lookup = std::move(lookup)](const Var& var) ->
ffi::Optional<Expr> {
- if (auto prim_var = var.as<tirx::PrimVar>(); prim_var &&
lookup.count(prim_var.value())) {
+ if (auto prim_var = var.as<PrimVar>(); prim_var &&
lookup.count(prim_var.value())) {
return prim_var.value().as_or_throw<PrimExpr>();
}
return std::nullopt;
diff --git a/src/relax/script/printer/dependent_type.cc
b/src/relax/script/printer/dependent_type.cc
index 615b342ac5..41d4beecab 100644
--- a/src/relax/script/printer/dependent_type.cc
+++ b/src/relax/script/printer/dependent_type.cc
@@ -47,7 +47,7 @@ ExprDoc PrintShapeVar(const PrimExpr& e, const AccessPath&
e_p, const IRDocsifie
bool func_var_mode = false;
if (f != nullptr) {
auto walk_fn = [f, &func_var_mode](const tirx::Var& var) ->
ffi::Expected<ffi::WalkResult> {
- if (auto prim_var = var.as<tirx::PrimVar>()) {
+ if (auto prim_var = var.as<PrimVar>()) {
if (f->func_vars->count(prim_var.value().get())) {
func_var_mode = true;
}
@@ -59,7 +59,7 @@ ExprDoc PrintShapeVar(const PrimExpr& e, const AccessPath&
e_p, const IRDocsifie
// Step 3. Stringify the PrimExpr if func var exists
bool is_bare_type_var = false;
if (f != nullptr && f->type_vars != nullptr) {
- if (auto var = e.as<tirx::PrimVar>()) {
+ if (auto var = e.as<PrimVar>()) {
is_bare_type_var = f->type_vars->count(var.value().get());
}
}
diff --git a/src/relax/script/printer/tir.cc b/src/relax/script/printer/tir.cc
index d79d41eeb1..8cdd946620 100644
--- a/src/relax/script/printer/tir.cc
+++ b/src/relax/script/printer/tir.cc
@@ -45,7 +45,7 @@ Doc PrintCanonicalVar(Var n, AccessPath n_p, IRDocsifier d) {
if (!n->ty.as<PrimTypeNode>()) {
return PrintRelaxVar(n, n_p, d);
}
- tirx::PrimVar prim_var = n.as_or_throw<tirx::PrimVar>();
+ PrimVar prim_var = n.as_or_throw<PrimVar>();
if (!d->IsVarDefined(n)) {
PrimType n_ty = n->ty.as_or_throw<PrimType>();
TVM_FFI_CHECK(!n_ty.IsScalableVector() && !n_ty.IsFixedLengthVector(),
TypeError)
diff --git a/src/relax/transform/alter_op_impl.cc
b/src/relax/transform/alter_op_impl.cc
index 19b65a5532..01e8e4714c 100644
--- a/src/relax/transform/alter_op_impl.cc
+++ b/src/relax/transform/alter_op_impl.cc
@@ -208,7 +208,7 @@ class AlterOpImplMutator : public ExprMutator {
// Output tensor of remove_pad op
te::Tensor output_tensor = te::compute(
dyn_old_shape,
- [&placeholder_tensor](const ffi::Array<tirx::PrimVar>& indices) {
+ [&placeholder_tensor](const ffi::Array<PrimVar>& indices) {
return placeholder_tensor(indices);
},
"output", topi::kElementWise);
diff --git a/src/relax/transform/bind_symbolic_vars.cc
b/src/relax/transform/bind_symbolic_vars.cc
index 05a898d7c7..62fdcb7e91 100644
--- a/src/relax/transform/bind_symbolic_vars.cc
+++ b/src/relax/transform/bind_symbolic_vars.cc
@@ -31,7 +31,7 @@ namespace tvm {
namespace relax {
Function FunctionBindSymbolicVars(
- Function func, ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>,
PrimExpr> obj_remap) {
+ Function func, ffi::Map<ffi::Variant<PrimVar, ffi::String>, PrimExpr>
obj_remap) {
// Early bail-out if no updates need to be made.
if (obj_remap.empty()) {
return func;
@@ -65,7 +65,7 @@ Function FunctionBindSymbolicVars(
TVM_FFI_ICHECK(!var_remap.count(var))
<< "Remap of variable " << var << " was defined multiple times";
var_remap.Set(var, replacement);
- } else if (auto opt = key.as<tirx::PrimVar>()) {
+ } else if (auto opt = key.as<PrimVar>()) {
auto var = opt.value();
TVM_FFI_ICHECK(!var_remap.count(var))
@@ -94,7 +94,7 @@ Function FunctionBindSymbolicVars(
namespace {
IRModule ModuleBindSymbolicVars(
- IRModule mod, ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>, PrimExpr>
binding_map) {
+ IRModule mod, ffi::Map<ffi::Variant<PrimVar, ffi::String>, PrimExpr>
binding_map) {
std::unordered_set<ffi::Any, ffi::AnyHash, ffi::AnyEqual> used;
IRModule updates;
for (const auto& [gvar, base_func] : mod->functions) {
@@ -102,8 +102,7 @@ IRModule ModuleBindSymbolicVars(
auto func = opt.value();
// Collect bindings that are used by this function.
- auto func_binding_map =
- [&]() -> ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>,
PrimExpr> {
+ auto func_binding_map = [&]() -> ffi::Map<ffi::Variant<PrimVar,
ffi::String>, PrimExpr> {
std::unordered_set<std::string> var_names;
std::unordered_set<const tirx::VarNode*> vars;
for (const auto& var : DefinedSymbolicVars(func)) {
@@ -111,12 +110,12 @@ IRModule ModuleBindSymbolicVars(
vars.insert(var.get());
}
- ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>, PrimExpr> out;
+ ffi::Map<ffi::Variant<PrimVar, ffi::String>, PrimExpr> out;
for (const auto& [key, replacement] : binding_map) {
bool used_by_function = false;
if (auto opt = key.as<ffi::String>()) {
used_by_function = var_names.count(opt.value());
- } else if (auto var = key.as<tirx::PrimVar>()) {
+ } else if (auto var = key.as<PrimVar>()) {
used_by_function = vars.count(var.value().get());
} else {
TVM_FFI_THROW(InternalError)
@@ -162,7 +161,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
namespace transform {
-Pass BindSymbolicVars(ffi::Map<ffi::Variant<tirx::PrimVar, ffi::String>,
PrimExpr> binding_map,
+Pass BindSymbolicVars(ffi::Map<ffi::Variant<PrimVar, ffi::String>, PrimExpr>
binding_map,
ffi::Optional<ffi::String> func_name) {
auto pass_func = [=](IRModule mod, PassContext context) -> IRModule {
if (func_name) {
diff --git a/src/relax/transform/bundle_model_params.cc
b/src/relax/transform/bundle_model_params.cc
index 417400fb1a..a46d9e94a2 100644
--- a/src/relax/transform/bundle_model_params.cc
+++ b/src/relax/transform/bundle_model_params.cc
@@ -67,7 +67,7 @@ class ModelParamBundler : public ExprMutator {
std::unordered_set<const VarNode*> bundled_prim_params;
ffi::Array<Var> bundled_prim_params_in_order;
for (size_t i = num_input; i < func->params.size(); i++) {
- if (func->params[i].as<tirx::PrimVar>()) {
+ if (func->params[i].as<PrimVar>()) {
bundled_prim_params.insert(func->params[i].get());
bundled_prim_params_in_order.push_back(func->params[i]);
}
@@ -138,7 +138,7 @@ class ModelParamBundler : public ExprMutator {
return ExprMutator::VisitExpr_(op);
}
if (auto it = var_to_expr_.find(var); it != var_to_expr_.end()) {
- bool is_prim_param = var.as<tirx::PrimVar>().has_value();
+ bool is_prim_param = var.as<PrimVar>().has_value();
if (is_prim_param) {
if (auto cached = var_remap_.find(var); cached != var_remap_.end()) {
return cached->second;
diff --git a/src/relax/transform/canonicalize_bindings.cc
b/src/relax/transform/canonicalize_bindings.cc
index 464e829e95..e72627bfe1 100644
--- a/src/relax/transform/canonicalize_bindings.cc
+++ b/src/relax/transform/canonicalize_bindings.cc
@@ -77,7 +77,7 @@ class SymbolicVarCanonicalizer : public ExprMutator {
bool has_runtime_use = false;
for (const auto& [var, value] : tir_var_map) {
if (var.same_as(binding->var)) continue;
- auto tir_var = var.as<tirx::PrimVar>();
+ auto tir_var = var.as<PrimVar>();
if (!tir_var) continue;
has_runtime_use = has_runtime_use ||
runtime_prim_var_uses_.count(*tir_var);
PrimExpr prim_expr = value.as_or_throw<PrimExpr>();
@@ -166,7 +166,7 @@ class SymbolicVarCanonicalizer : public ExprMutator {
void VisitExpr_(const VarNode* op) final {
Var var = ffi::GetRef<Var>(op);
- if (auto prim_var = var.as<tirx::PrimVar>()) {
+ if (auto prim_var = var.as<PrimVar>()) {
uses_.insert(*prim_var);
}
}
@@ -177,7 +177,7 @@ class SymbolicVarCanonicalizer : public ExprMutator {
PrimExpr CanonicalizeShapeValue(const PrimExpr& expr) {
auto f_substitute = [this](const Var& var) ->
ffi::Expected<ffi::UnchangedOr<ffi::Any>> {
- auto prim_var = var.as<tirx::PrimVar>();
+ auto prim_var = var.as<PrimVar>();
if (!prim_var) return ffi::Unchanged();
auto it = known_values_.find(*prim_var);
if (it == known_values_.end()) return ffi::Unchanged();
diff --git a/src/relax/transform/convert_layout.cc
b/src/relax/transform/convert_layout.cc
index f33f70e8ac..9759f5c9e5 100644
--- a/src/relax/transform/convert_layout.cc
+++ b/src/relax/transform/convert_layout.cc
@@ -107,9 +107,9 @@ class LayoutConvertMutator : public ExprMutator {
initial_indices_expr.push_back(var.as_or_throw<PrimExpr>());
}
ffi::Array<PrimExpr> desired_shape =
todesired.ForwardIndex(initial_indices_expr);
- return IndexMap(initial_indices.Map(
- [](tvm::tirx::Var var) { return
var.as_or_throw<tvm::tirx::PrimVar>(); }),
- desired_shape, std::move(inverse_index_map));
+ return IndexMap(
+ initial_indices.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::PrimVar>(); }),
+ desired_shape, std::move(inverse_index_map));
}
Expr RewriteExpr(const Expr& expr, const NLayout& to) {
diff --git a/src/relax/transform/fuse_tir.cc b/src/relax/transform/fuse_tir.cc
index e955d6fd3d..166a782de7 100644
--- a/src/relax/transform/fuse_tir.cc
+++ b/src/relax/transform/fuse_tir.cc
@@ -563,7 +563,7 @@ class FusedTIRConstructor : public ExprVisitor {
void VisitExpr_(const FunctionNode* func) final {
auto relax_to_tir_var_map =
RelaxToTIRVarMapCollector::Collect(mod_, ffi::GetRef<Function>(func));
- std::vector<ffi::Variant<tirx::PrimVar, tirx::BufferVar>> prim_func_params;
+ std::vector<ffi::Variant<PrimVar, tirx::BufferVar>> prim_func_params;
for (const Var& relax_param : func->params) {
size_t size_before = prim_func_params.size();
CollectPrimFuncParams(relax_param, &prim_func_params,
relax_to_tir_var_map.Get(relax_param));
@@ -596,7 +596,7 @@ class FusedTIRConstructor : public ExprVisitor {
tirx::Var param = tirx::Var("p_" + buffer.name(),
PointerType::VoidPointerTy());
func_info_.params.push_back(param);
func_info_.buffer_map.Set(param, buffer);
- } else if (auto var = param.as<tirx::PrimVar>()) {
+ } else if (auto var = param.as<PrimVar>()) {
func_info_.params.push_back(var.value());
}
}
@@ -927,7 +927,7 @@ class FusedTIRConstructor : public ExprVisitor {
* \param out The vector into which to collect the params/buffers
*/
static void CollectPrimFuncParams(const Var& relax_param,
- std::vector<ffi::Variant<tirx::PrimVar,
tirx::BufferVar>>* out,
+ std::vector<ffi::Variant<PrimVar,
tirx::BufferVar>>* out,
const ffi::Optional<tirx::BufferVar>&
tir_buffer_param) {
auto ty = GetType(relax_param);
@@ -953,12 +953,12 @@ class FusedTIRConstructor : public ExprVisitor {
} else if (ty.as<PrimTypeNode>()) {
// Case 2. The relax param is a scalar, so its canonical Var is a TIR
parameter.
- out->push_back(relax_param.as_or_throw<tirx::PrimVar>());
+ out->push_back(relax_param.as_or_throw<PrimVar>());
} else if (const auto* shape_expr = ty.as<ShapeTypeNode>()) {
// Case 3. The relax param is a tuple of scalars, each represented as a
tirx var
for (const auto& var : shape_expr->values.value()) {
- auto prim_var = var.as<tirx::PrimVar>();
+ auto prim_var = var.as<PrimVar>();
TVM_FFI_ICHECK(prim_var.has_value());
out->push_back(prim_var.value());
}
@@ -1228,7 +1228,7 @@ class TIRFuseMutator : public ExprMutator {
TVM_FFI_ICHECK(shape->values.has_value())
<< "FuseTIR requires all shape input has ty value.";
for (const PrimExpr& prim_value : shape->values.value()) {
- TVM_FFI_ICHECK(prim_value.as<tirx::PrimVar>())
+ TVM_FFI_ICHECK(prim_value.as<PrimVar>())
<< "All shape inputs are expected to be single tirx var.";
arg_list.push_back(prim_value);
}
diff --git a/src/relax/transform/lazy_transform_params.cc
b/src/relax/transform/lazy_transform_params.cc
index e2ce85896b..7d2a0251c9 100644
--- a/src/relax/transform/lazy_transform_params.cc
+++ b/src/relax/transform/lazy_transform_params.cc
@@ -74,7 +74,7 @@ class LazyInputMutator : public ExprMutator {
std::unordered_set<tirx::Var>
externally_visible_vars(array_externally_visible_vars.begin(),
array_externally_visible_vars.end());
Type new_ret_ty = EraseToWellDefined(func->ret_ty, [&](const Var& var) ->
ffi::Optional<Expr> {
- if (auto prim_var = var.as<tirx::PrimVar>();
+ if (auto prim_var = var.as<PrimVar>();
prim_var && externally_visible_vars.count(prim_var.value())) {
return prim_var.value().as_or_throw<PrimExpr>();
}
diff --git a/src/relax/transform/lift_transform_params.cc
b/src/relax/transform/lift_transform_params.cc
index 84e77c6394..58c5f8fa7b 100644
--- a/src/relax/transform/lift_transform_params.cc
+++ b/src/relax/transform/lift_transform_params.cc
@@ -240,7 +240,7 @@ struct LocalCollectInfo : public BaseCollectInfo {
ffi::Array<tirx::Var> global_tir_vars =
global_info->GetPropagatedSymbolicVariables();
global_tir_vars = global_tir_vars.Map([&](const tirx::Var& var) ->
tirx::Var {
if (auto it = global_to_local.find(var); it != global_to_local.end()) {
- return (*it).second.as_or_throw<tirx::PrimVar>();
+ return (*it).second.as_or_throw<PrimVar>();
} else {
// This is the case when the some of the outputs of the shared
transform is not used in
// this function.
@@ -251,7 +251,7 @@ struct LocalCollectInfo : public BaseCollectInfo {
}();
if (propagated_tir_vars.size()) {
ShapeType shape_ty(propagated_tir_vars.Map(
- [](tirx::Var var) { return
var.as_or_throw<tirx::PrimVar>().as_or_throw<PrimExpr>(); }));
+ [](tirx::Var var) { return
var.as_or_throw<PrimVar>().as_or_throw<PrimExpr>(); }));
Var shape_expr("vars_from_compile_time_params", shape_ty);
params.push_back(shape_expr);
}
diff --git a/src/relax/transform/rewrite_cuda_graph.cc
b/src/relax/transform/rewrite_cuda_graph.cc
index 99b0f2d2d6..5f7b7563c9 100644
--- a/src/relax/transform/rewrite_cuda_graph.cc
+++ b/src/relax/transform/rewrite_cuda_graph.cc
@@ -113,7 +113,7 @@ class FuncBuilder : public ExprMutator {
* \brief Mark a TIR variable as the ShapeExpr input of the new function.
* \param var The variable to mark as input
*/
- void MarkShapeExprInput(const tirx::PrimVar& var) {
shape_expr_inputs_.push_back(var); }
+ void MarkShapeExprInput(const PrimVar& var) {
shape_expr_inputs_.push_back(var); }
/*!
* \brief Mark a variable as the output of the new function. The variable
must be the LHS of an
* existing binding in the new function.
@@ -130,8 +130,8 @@ class FuncBuilder : public ExprMutator {
ffi::Optional<Var> shape_expr = std::nullopt;
if (shape_expr_inputs_.size()) {
ffi::Array<PrimExpr> tir_vars;
- for (const tirx::PrimVar& var : shape_expr_inputs_) {
- tirx::PrimVar new_var = var.CopyWithSuffix("");
+ for (const PrimVar& var : shape_expr_inputs_) {
+ PrimVar new_var = var.CopyWithSuffix("");
var_remap_[var] = new_var;
tir_vars.push_back(new_var);
}
@@ -168,7 +168,7 @@ class FuncBuilder : public ExprMutator {
support::OrderedSet<const VarNode*> inputs_;
support::OrderedSet<const VarNode*> outputs_;
- support::OrderedSet<tirx::PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
shape_expr_inputs_;
+ support::OrderedSet<PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
shape_expr_inputs_;
std::vector<const VarBindingNode*> bindings_;
};
@@ -252,7 +252,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
DefinableTIRVarsInType(func->params[i]->ty.as_or_throw<Type>());
if (i < num_inputs) {
for (const auto& symbolic_var : symbolic_vars) {
- auto prim_var = symbolic_var.as_or_throw<tirx::PrimVar>();
+ auto prim_var = symbolic_var.as_or_throw<PrimVar>();
if (capture_symbolic_var_name_hints.count(symbolic_var->name)) {
capture_symbolic_vars_.insert(prim_var);
}
@@ -260,7 +260,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
} else {
static_vars_.insert(func->params[i].get());
for (const auto& symbolic_var : symbolic_vars) {
-
capture_symbolic_vars_.insert(symbolic_var.as_or_throw<tirx::PrimVar>());
+
capture_symbolic_vars_.insert(symbolic_var.as_or_throw<PrimVar>());
}
}
}
@@ -278,7 +278,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
plan->lifted_bindings = std::move(region->bindings_);
if (region->shape_expr_inputs_.size()) {
ffi::Array<PrimExpr> tir_vars;
- for (const tirx::PrimVar& var : region->shape_expr_inputs_) {
+ for (const PrimVar& var : region->shape_expr_inputs_) {
tir_vars.push_back(var);
}
plan->propogated_tir_vars = ShapeExpr(tir_vars);
@@ -371,7 +371,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
// Check whether the call can be lifted to the capture function. It
requires all the arguments
// to be static and the call to be a kernel launch or a pure operation
(e.g. memory view).
std::vector<const VarNode*> args;
- std::vector<tirx::PrimVar> tir_vars;
+ std::vector<PrimVar> tir_vars;
bool is_all_static = [&]() {
if (!IsStatic(call->args, &args, &tir_vars)) {
return false;
@@ -418,7 +418,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
}
void MarkAsFuncInput(const std::vector<const VarNode*>& vars,
- const std::vector<tirx::PrimVar>& tir_vars = {}) {
+ const std::vector<PrimVar>& tir_vars = {}) {
if (current_block_scope_.capture_builder == nullptr) {
return;
}
@@ -428,7 +428,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
current_block_scope_.capture_builder->MarkInput(var);
}
}
- for (const tirx::PrimVar& tir_var : tir_vars) {
+ for (const PrimVar& tir_var : tir_vars) {
current_block_scope_.capture_builder->MarkShapeExprInput(tir_var);
}
}
@@ -458,7 +458,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
void VisitBinding_(const VarBindingNode* binding, const TupleNode* tuple)
final {
std::vector<const VarNode*> args;
- std::vector<tirx::PrimVar> tir_vars;
+ std::vector<PrimVar> tir_vars;
if (IsStatic(tuple->fields, &args, &tir_vars)) {
AddStaticBinding(binding, false);
MarkAsFuncInput(args, tir_vars);
@@ -482,11 +482,11 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
bool IsStatic(const PrimExpr& expr,
[[maybe_unused]] std::vector<const VarNode*>* vars_collector =
nullptr,
- std::vector<tirx::PrimVar>* tir_vars_collector = nullptr) {
+ std::vector<PrimVar>* tir_vars_collector = nullptr) {
bool is_static = true;
std::unordered_set<const ffi::Object*> visited;
auto walk_fn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult>
{
- auto prim_var = var.as<tirx::PrimVar>();
+ auto prim_var = var.as<PrimVar>();
if (!prim_var || !visited.insert(prim_var.value().get()).second) {
return ffi::WalkResult::Advance();
}
@@ -504,7 +504,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
}
bool IsStatic(const Expr& expr, std::vector<const VarNode*>* vars_collector
= nullptr,
- std::vector<tirx::PrimVar>* tir_vars_collector = nullptr) {
+ std::vector<PrimVar>* tir_vars_collector = nullptr) {
if (expr->IsInstance<ConstantNode>() ||
expr->IsInstance<DataTypeImmNode>() ||
expr->IsInstance<StringImmNode>() ||
expr->IsInstance<GlobalVarNode>()) {
return true;
@@ -530,7 +530,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
}
bool IsStaticRuntimeVar(const VarNode* var, std::vector<const VarNode*>*
vars_collector,
- std::vector<tirx::PrimVar>* tir_vars_collector) {
+ std::vector<PrimVar>* tir_vars_collector) {
if (vars_collector != nullptr) {
vars_collector->push_back(var);
}
@@ -540,7 +540,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
template <typename T>
bool IsStatic(const ffi::Array<T>& exprs, std::vector<const VarNode*>*
vars_collector = nullptr,
- std::vector<tirx::PrimVar>* tir_vars_collector = nullptr) {
+ std::vector<PrimVar>* tir_vars_collector = nullptr) {
bool result = true;
for (const auto& expr : exprs) {
// If vars_collector is provided, we will collect all the vars in the
exprs and we should
@@ -554,7 +554,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
}
bool IsStatic(const Type& ty, std::vector<const VarNode*>* vars_collector =
nullptr,
- std::vector<tirx::PrimVar>* tir_vars_collector = nullptr) {
+ std::vector<PrimVar>* tir_vars_collector = nullptr) {
if (const auto* tensor_ty = ty.as<TensorTypeNode>()) {
if (auto shape = tensor_ty->GetShape()) {
return IsStatic(shape.value(), vars_collector, tir_vars_collector);
@@ -630,7 +630,7 @@ class CUDAGraphRewritePlanner : public ExprVisitor {
std::unordered_set<const VarNode*> static_vars_;
// Symbolic variables that are allowed to be captured. This can come from
symbolic shapes of
// weights or hints in the function annotations.
- std::unordered_set<tirx::PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
capture_symbolic_vars_;
+ std::unordered_set<PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
capture_symbolic_vars_;
// Binding to the FuncBuilder if the binding is lifted. This is used to
update the inputs/outputs
// of the lifted function when its binding is used outside.
std::unordered_map<const VarNode*, FuncBuilder*> binding_to_region_;
@@ -819,8 +819,7 @@ class CUDAGraphRewriter : public ExprMutator {
ffi::Map<Var, Expr> var_remap;
TVM_FFI_ICHECK_EQ(symbolic_params.size(),
propogated_tir_vars->values.size());
for (int i = 0; i < static_cast<int>(symbolic_params.size()); ++i) {
- var_remap.Set(symbolic_params[i].as_or_throw<tirx::PrimVar>(),
- propogated_tir_vars->values[i]);
+ var_remap.Set(symbolic_params[i].as_or_throw<PrimVar>(),
propogated_tir_vars->values[i]);
}
call_ty = Bind(call_ty, var_remap);
}
diff --git a/src/relax/utils.cc b/src/relax/utils.cc
index 5021494196..f20d048130 100644
--- a/src/relax/utils.cc
+++ b/src/relax/utils.cc
@@ -122,7 +122,7 @@ tvm::ffi::Map<Var, Expr> InferSymbolicVarMap(
tvm::ffi::Map<Var, Expr> var_remap = relax_var_remap;
for (const auto& [var, value] : relax_var_remap) {
- if (!var.as<tirx::PrimVar>()) continue;
+ if (!var.as<PrimVar>()) continue;
TVM_FFI_CHECK(value.as<PrimExpr>().has_value(), ValueError)
<< "Explicit binding for symbolic variable " << var
<< " must be a primitive expression, but received " << value;
@@ -130,7 +130,7 @@ tvm::ffi::Map<Var, Expr> InferSymbolicVarMap(
auto bind_from_prim_expr = [&relax_var_remap, &var_remap, &analyzer](const
PrimExpr& var_shape,
const
PrimExpr& expr_shape) {
- if (auto var = var_shape.as<tirx::PrimVar>()) {
+ if (auto var = var_shape.as<PrimVar>()) {
if (auto it = relax_var_remap.find(var.value()); it !=
relax_var_remap.end()) {
auto explicit_value = (*it).second.as<PrimExpr>();
TVM_FFI_CHECK(explicit_value.has_value(), ValueError)
diff --git
a/src/s_tir/meta_schedule/schedule_rule/multi_level_tiling_tensor_core.cc
b/src/s_tir/meta_schedule/schedule_rule/multi_level_tiling_tensor_core.cc
index d117db7854..e8613fb0c0 100644
--- a/src/s_tir/meta_schedule/schedule_rule/multi_level_tiling_tensor_core.cc
+++ b/src/s_tir/meta_schedule/schedule_rule/multi_level_tiling_tensor_core.cc
@@ -799,11 +799,11 @@ ffi::Optional<LoopRV>
MultiLevelTilingTensorCoreNode::TransformWithTensorIntrin(
const tirx::IndexMap& index_map = mapping_info->mappings[0];
// Find the correspondence between block iters and the iters in the index
map.
- std::unordered_map<tirx::PrimVar, tirx::PrimVar, ffi::ObjectPtrHash,
ffi::ObjectPtrEqual>
+ std::unordered_map<PrimVar, PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
lhs_to_index_map_src;
- std::unordered_map<tirx::PrimVar, PrimExpr, ffi::ObjectPtrHash,
ffi::ObjectPtrEqual>
+ std::unordered_map<PrimVar, PrimExpr, ffi::ObjectPtrHash,
ffi::ObjectPtrEqual>
rhs_to_index_map_tgt;
- std::unordered_set<tirx::PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
unmapped_index_map_src;
+ std::unordered_set<PrimVar, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
unmapped_index_map_src;
TVM_FFI_ICHECK_EQ(mapping_info->lhs_iters.size(),
index_map->initial_indices.size());
for (int i = 0; i < static_cast<int>(mapping_info->lhs_iters.size()); ++i) {
lhs_to_index_map_src[mapping_info->lhs_iters[i]->var] =
index_map->initial_indices[i];
@@ -817,7 +817,7 @@ ffi::Optional<LoopRV>
MultiLevelTilingTensorCoreNode::TransformWithTensorIntrin(
static_cast<int>(mapping_info->rhs_iters.size());
TVM_FFI_ICHECK_GE(offset, 0);
for (int i = 0; i < offset; ++i) {
- auto var = index_map->final_indices[i].as<tirx::PrimVar>();
+ auto var = index_map->final_indices[i].as<PrimVar>();
TVM_FFI_ICHECK(var.has_value());
unmapped_index_map_src.insert(var.value());
}
@@ -827,21 +827,21 @@ ffi::Optional<LoopRV>
MultiLevelTilingTensorCoreNode::TransformWithTensorIntrin(
auto f_get_sub_index_map = [&](const tirx::BufferVar& lhs_buffer,
const ffi::Array<Range>& lhs_region) {
- std::vector<tirx::PrimVar> sub_index_map_src;
+ std::vector<PrimVar> sub_index_map_src;
std::vector<PrimExpr> sub_index_map_tgt;
const tirx::BufferVar& rhs_buffer =
mapping_info->lhs_buffer_map[lhs_buffer];
for (const Range& range : lhs_region) {
TVM_FFI_ICHECK(tirx::is_one(range->extent));
- auto var = range->min.as<tirx::PrimVar>();
+ auto var = range->min.as<PrimVar>();
TVM_FFI_ICHECK(var.has_value());
- const tirx::PrimVar& lhs_representer = lhs_to_index_map_src[var.value()];
+ const PrimVar& lhs_representer = lhs_to_index_map_src[var.value()];
sub_index_map_src.push_back(lhs_representer);
if (unmapped_index_map_src.count(lhs_representer)) {
sub_index_map_tgt.push_back(lhs_representer);
}
}
for (size_t i = 0; i <
mapping_info->rhs_buffer_indices[rhs_buffer].size(); ++i) {
- auto var =
mapping_info->rhs_buffer_indices[rhs_buffer][i].as<tirx::PrimVar>();
+ auto var = mapping_info->rhs_buffer_indices[rhs_buffer][i].as<PrimVar>();
TVM_FFI_ICHECK(var.has_value());
sub_index_map_tgt.push_back(rhs_to_index_map_tgt[var.value()]);
}
diff --git a/src/s_tir/transform/default_gpu_schedule.cc
b/src/s_tir/transform/default_gpu_schedule.cc
index eb33710e09..59e8834dd0 100644
--- a/src/s_tir/transform/default_gpu_schedule.cc
+++ b/src/s_tir/transform/default_gpu_schedule.cc
@@ -136,15 +136,15 @@ tirx::PrimFunc WrapBareSBlockBody(const tirx::PrimFunc&
func) {
tvm::IntImm one(tvm::PrimType::Int(32), 1);
tirx::Var loop_var("u", tvm::PrimType::Int(32));
tirx::Var iter_var_var("vu", tvm::PrimType::Int(32));
- tirx::IterVar new_iter(tvm::Range::FromMinExtent(zero, one),
- iter_var_var.as_or_throw<tirx::PrimVar>(),
tirx::IterVarType::kDataPar);
+ tirx::IterVar new_iter(tvm::Range::FromMinExtent(zero, one),
iter_var_var.as_or_throw<PrimVar>(),
+ tirx::IterVarType::kDataPar);
tirx::SBlock inner_block = realize->block;
inner_block.CopyOnWrite()->iter_vars = ffi::Array<tirx::IterVar>{new_iter};
tirx::SBlockRealize inner_realize(
/*iter_values=*/ffi::Array<tvm::PrimExpr>{loop_var.as_or_throw<tvm::PrimExpr>()},
/*predicate=*/realize->predicate, inner_block);
- tirx::Stmt for_stmt = tirx::For(loop_var.as_or_throw<tirx::PrimVar>(), zero,
one,
- tirx::ForKind::kSerial, inner_realize);
+ tirx::Stmt for_stmt =
+ tirx::For(loop_var.as_or_throw<PrimVar>(), zero, one,
tirx::ForKind::kSerial, inner_realize);
tirx::SBlock root_block(/*iter_vars=*/ffi::Array<tirx::IterVar>{},
/*reads=*/ffi::Array<tirx::BufferRegion>{},
/*writes=*/ffi::Array<tirx::BufferRegion>{},
diff --git a/src/te/operation/create_primfunc.cc
b/src/te/operation/create_primfunc.cc
index 938ca0648a..09df9b42e4 100644
--- a/src/te/operation/create_primfunc.cc
+++ b/src/te/operation/create_primfunc.cc
@@ -904,7 +904,7 @@ PrimFunc GenerateAndCompletePrimFunc(const
ffi::Array<ffi::ObjectRef>& arg_tir_v
auto it = info->tensor2buffers.find(tensor);
TVM_FFI_ICHECK(it != info->tensor2buffers.end());
parameters.push_back(it->second.var());
- } else if (auto var = arg.as<tirx::PrimVar>()) {
+ } else if (auto var = arg.as<PrimVar>()) {
parameters.push_back(var.value());
}
}
diff --git a/src/tirx/ir/buffer.cc b/src/tirx/ir/buffer.cc
index ffbea28fab..623f461e72 100644
--- a/src/tirx/ir/buffer.cc
+++ b/src/tirx/ir/buffer.cc
@@ -731,7 +731,7 @@ tirx::BufferVar
BufferWithOffsetAlignment(ffi::Array<PrimExpr> shape, PrimType d
std::string memory_scope) {
PrimExpr elem_offset;
if (offset_factor != 0) {
- elem_offset = tirx::PrimVar(name + "_elem_offset", shape[0].ty());
+ elem_offset = PrimVar(name + "_elem_offset", shape[0].ty());
} else {
elem_offset = PrimExpr();
}
diff --git a/src/tirx/op/op.cc b/src/tirx/op/op.cc
index ad0469c341..50a163f756 100644
--- a/src/tirx/op/op.cc
+++ b/src/tirx/op/op.cc
@@ -47,7 +47,6 @@ using tirx::CallEffectKind;
using tirx::is_const_int;
using tirx::IterVar;
using tirx::MakeConst;
-using tirx::PrimVar;
using tirx::TCallEffectKind;
using tirx::TGlobalSymbol;
using tirx::TIRxOpCategory;
diff --git a/src/tirx/script/builder/ir.cc b/src/tirx/script/builder/ir.cc
index f505fc57b9..79a29a157d 100644
--- a/src/tirx/script/builder/ir.cc
+++ b/src/tirx/script/builder/ir.cc
@@ -58,7 +58,7 @@ BufferVar BufferDecl(ffi::Array<PrimExpr> shape, PrimType
dtype, ffi::String buf
}
if (!elem_offset.has_value() && offset_factor) {
PrimType shape_dtype = shape.empty() ? PrimType::Int(32) : shape[0].ty();
- elem_offset = tvm::tirx::PrimVar("elem_offset", shape_dtype);
+ elem_offset = tvm::PrimVar("elem_offset", shape_dtype);
}
return BufferVar(buffer_name, tvm::tirx::BufferType(storage_scope, dtype,
shape,
strides.value_or(ffi::Array<PrimExpr>()),
@@ -207,14 +207,14 @@ ffi::Array<tvm::tirx::Var>
ScopeId(ffi::Optional<ffi::Array<PrimExpr>> extents,
}
ffi::Array<tvm::tirx::Var> scope_ids;
for (size_t i = 0; i < n_vars; ++i) {
- scope_ids.push_back(tvm::tirx::PrimVar("", dtype));
+ scope_ids.push_back(tvm::PrimVar("", dtype));
}
// Emit a standalone ScopeIdDefStmt to the current TIRFrame's stmts list.
// The def is visible to all subsequent stmts within the same enclosing
// scope (PrimFunc body, AttrStmt body, ExecScope body, etc.).
tvm::tirx::ScopeIdDef def(
- scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::tirx::PrimVar>(); }),
- extents, tvm::tirx::StringPairToScopeBinding(parent, cur));
+ scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::PrimVar>(); }), extents,
+ tvm::tirx::StringPairToScopeBinding(parent, cur));
AddToParent(tvm::tirx::ScopeIdDefStmt(def));
return scope_ids;
}
@@ -235,11 +235,11 @@ ffi::Array<tvm::tirx::Var>
CtaId(ffi::Optional<ffi::Array<PrimExpr>> extents, ff
<< "ValueError: preferred=... requires explicit extents (deferred form
is incompatible)";
ffi::Array<tvm::tirx::Var> scope_ids;
for (size_t i = 0; i < extents.value().size(); ++i) {
- scope_ids.push_back(tvm::tirx::PrimVar("", dtype));
+ scope_ids.push_back(tvm::PrimVar("", dtype));
}
tvm::tirx::ScopeIdDef def(
- scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::tirx::PrimVar>(); }),
- extents, tvm::tirx::StringPairToScopeBinding(parent, "cta"),
preferred);
+ scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::PrimVar>(); }), extents,
+ tvm::tirx::StringPairToScopeBinding(parent, "cta"), preferred);
AddToParent(tvm::tirx::ScopeIdDefStmt(def));
return scope_ids;
}
@@ -248,9 +248,9 @@ ffi::Array<tvm::tirx::Var>
CtaId(ffi::Optional<ffi::Array<PrimExpr>> extents, ff
ffi::Array<tvm::tirx::Var> CtaIdInPair(PrimType dtype) {
CheckExplicitIndexDtype(dtype);
- ffi::Array<tvm::tirx::Var> scope_ids{tvm::tirx::PrimVar("", dtype)};
+ ffi::Array<tvm::tirx::Var> scope_ids{tvm::PrimVar("", dtype)};
tvm::tirx::ScopeIdDef def(
- scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::tirx::PrimVar>(); }),
+ scope_ids.Map([](tvm::tirx::Var var) { return
var.as_or_throw<tvm::PrimVar>(); }),
ffi::Array<PrimExpr>{IntImm::Int32(2)},
tvm::tirx::ScopeBinding::kClusterCtaPair);
AddToParent(tvm::tirx::ScopeIdDefStmt(def));
return scope_ids;
@@ -421,17 +421,17 @@ IterVar PushBlockVar(IterVar iter_var, PrimExpr binding) {
return iter_var;
}
-#define TVM_TIRX_IR_BUILDER_AXIS(Method, Kind, Name)
\
- Var Method(Range dom, PrimExpr binding, PrimType dtype) {
\
- TVM_FFI_ICHECK(dom.defined()) << Name << " axis must have a domain";
\
- PrimType min_ty = dom->min.ty();
\
- PrimType extent_ty = dom->extent.ty();
\
- int bits = std::max({min_ty.bits(), extent_ty.bits(), dtype.bits()});
\
- PrimType var_ty = dtype.WithBits(bits);
\
- return PushBlockVar(IterVar(/*dom=*/dom, /*var=*/tvm::tirx::PrimVar("",
var_ty), \
- /*iter_type=*/Kind, /*thread_tag=*/""),
\
- binding)
\
- ->var;
\
+#define TVM_TIRX_IR_BUILDER_AXIS(Method, Kind, Name)
\
+ Var Method(Range dom, PrimExpr binding, PrimType dtype) {
\
+ TVM_FFI_ICHECK(dom.defined()) << Name << " axis must have a domain";
\
+ PrimType min_ty = dom->min.ty();
\
+ PrimType extent_ty = dom->extent.ty();
\
+ int bits = std::max({min_ty.bits(), extent_ty.bits(), dtype.bits()});
\
+ PrimType var_ty = dtype.WithBits(bits);
\
+ return PushBlockVar(IterVar(/*dom=*/dom, /*var=*/tvm::PrimVar("", var_ty),
/*iter_type=*/Kind, \
+ /*thread_tag=*/""),
\
+ binding)
\
+ ->var;
\
}
TVM_TIRX_IR_BUILDER_AXIS(Spatial, tvm::tirx::IterVarType::kDataPar, "Spatial");
TVM_TIRX_IR_BUILDER_AXIS(Reduce, tvm::tirx::IterVarType::kCommReduce,
"Reduction");
@@ -470,14 +470,14 @@ ffi::Array<Var> Remap(ffi::String kinds,
ffi::Array<PrimExpr> bindings, PrimType
PrimType dtype = v.value().ty();
if (c == 'S') {
results.push_back(PushBlockVar(IterVar(/*dom=*/dom,
- /*var=*/tvm::tirx::PrimVar("",
dtype),
+ /*var=*/tvm::PrimVar("", dtype),
/*iter_type=*/IterVarType::kDataPar,
/*thread_tag=*/""),
e)
->var);
} else if (c == 'R') {
results.push_back(PushBlockVar(IterVar(/*dom=*/dom,
- /*var=*/tvm::tirx::PrimVar("",
dtype),
+ /*var=*/tvm::PrimVar("", dtype),
/*iter_type=*/IterVarType::kCommReduce,
/*thread_tag=*/""),
e)
@@ -527,31 +527,31 @@ PrimExpr ConvertLoopBound(const PrimExpr& e, const
PrimType& var_ty) {
return tvm::prim::Cast(var_ty, e);
}
-#define TVM_TIRX_IR_BUILDER_FOR_FRAME(Method, Kind)
\
- ForFrame Method(PrimExpr start, PrimExpr stop,
\
- ffi::Optional<ffi::Map<ffi::String, Any>> annotations,
\
- ffi::Optional<PrimExpr> step, ffi::Optional<PrimType> dtype)
{ \
- PrimType var_ty = InferLoopVarDtype(start, stop, dtype);
\
- PrimExpr min = ConvertLoopBound(start, var_ty);
\
- PrimExpr extent = arith::Analyzer()->Simplify(ConvertLoopBound(stop,
var_ty) - min); \
- if (step.has_value()) {
\
- step = ConvertLoopBound(step.value(), var_ty);
\
- }
\
- ffi::ObjectPtr<ForFrameNode> n = ffi::make_object<ForFrameNode>();
\
- n->vars = {Var("v", var_ty)};
\
- n->doms = {Range::FromMinExtent(min, extent)};
\
- n->steps = {step};
\
- n->f_make_for_loop = [annotations](ffi::Array<Var> vars, ffi::Array<Range>
doms, \
- ffi::Array<ffi::Optional<PrimExpr>>
steps, \
- tvm::tirx::Stmt body) {
\
- TVM_FFI_ICHECK_EQ(vars.size(), 1);
\
- TVM_FFI_ICHECK_EQ(doms.size(), 1);
\
- TVM_FFI_ICHECK_EQ(steps.size(), 1);
\
- return tvm::tirx::For(vars[0].as_or_throw<tvm::tirx::PrimVar>(),
doms[0]->min, \
- doms[0]->extent, Kind, body, std::nullopt,
\
- annotations.value_or(ffi::Map<ffi::String,
Any>()), steps[0]); \
- };
\
- return ForFrame(n);
\
+#define TVM_TIRX_IR_BUILDER_FOR_FRAME(Method, Kind)
\
+ ForFrame Method(PrimExpr start, PrimExpr stop,
\
+ ffi::Optional<ffi::Map<ffi::String, Any>> annotations,
\
+ ffi::Optional<PrimExpr> step, ffi::Optional<PrimType> dtype)
{ \
+ PrimType var_ty = InferLoopVarDtype(start, stop, dtype);
\
+ PrimExpr min = ConvertLoopBound(start, var_ty);
\
+ PrimExpr extent = arith::Analyzer()->Simplify(ConvertLoopBound(stop,
var_ty) - min); \
+ if (step.has_value()) {
\
+ step = ConvertLoopBound(step.value(), var_ty);
\
+ }
\
+ ffi::ObjectPtr<ForFrameNode> n = ffi::make_object<ForFrameNode>();
\
+ n->vars = {Var("v", var_ty)};
\
+ n->doms = {Range::FromMinExtent(min, extent)};
\
+ n->steps = {step};
\
+ n->f_make_for_loop = [annotations](ffi::Array<Var> vars, ffi::Array<Range>
doms, \
+ ffi::Array<ffi::Optional<PrimExpr>>
steps, \
+ tvm::tirx::Stmt body) {
\
+ TVM_FFI_ICHECK_EQ(vars.size(), 1);
\
+ TVM_FFI_ICHECK_EQ(doms.size(), 1);
\
+ TVM_FFI_ICHECK_EQ(steps.size(), 1);
\
+ return tvm::tirx::For(vars[0].as_or_throw<tvm::PrimVar>(), doms[0]->min,
doms[0]->extent, \
+ Kind, body, std::nullopt,
\
+ annotations.value_or(ffi::Map<ffi::String,
Any>()), steps[0]); \
+ };
\
+ return ForFrame(n);
\
}
TVM_TIRX_IR_BUILDER_FOR_FRAME(Serial, tvm::tirx::ForKind::kSerial);
@@ -580,9 +580,9 @@ ForFrame ThreadBinding(PrimExpr start, PrimExpr stop,
ffi::String thread,
TVM_FFI_ICHECK_EQ(vars.size(), 1);
TVM_FFI_ICHECK_EQ(doms.size(), 1);
TVM_FFI_ICHECK(steps.size() == 1 && (!steps[0].has_value() ||
is_one(*steps[0])));
- IterVar iter_var(Range(nullptr), tvm::tirx::PrimVar("iter", dtype),
IterVarType::kThreadIndex,
+ IterVar iter_var(Range(nullptr), tvm::PrimVar("iter", dtype),
IterVarType::kThreadIndex,
thread);
- return For(vars[0].as_or_throw<tvm::tirx::PrimVar>(), doms[0]->min,
doms[0]->extent,
+ return For(vars[0].as_or_throw<tvm::PrimVar>(), doms[0]->min,
doms[0]->extent,
ForKind::kThreadBinding, body, iter_var,
annotations.value_or(ffi::Map<ffi::String, ffi::Any>()),
std::nullopt);
};
@@ -623,7 +623,7 @@ ForFrame Grid(ffi::Array<ffi::Variant<PrimExpr,
ffi::Tuple<PrimExpr, PrimExpr>>>
for (int i = n - 1; i >= 0; --i) {
Range dom = doms[i];
Var var = vars[i];
- body = For(var.as_or_throw<tvm::tirx::PrimVar>(), dom->min, dom->extent,
ForKind::kSerial,
+ body = For(var.as_or_throw<tvm::PrimVar>(), dom->min, dom->extent,
ForKind::kSerial,
std::move(body),
/*thread_binding=*/std::nullopt, /*annotations=*/{},
/*step=*/steps[i]);
}
@@ -774,8 +774,8 @@ ComposeOpFrame ComposeOp(ffi::Map<ffi::String, BufferVar>
workspace,
}
Var EnvThread(ffi::String thread_tag, PrimType dtype) {
- IterVar iter_var(Range{nullptr}, tvm::tirx::PrimVar("", dtype),
- tvm::tirx::IterVarType::kThreadIndex, thread_tag);
+ IterVar iter_var(Range{nullptr}, tvm::PrimVar("", dtype),
tvm::tirx::IterVarType::kThreadIndex,
+ thread_tag);
Var var = iter_var->var;
if (ffi::Optional<PrimFuncFrame> opt_frame =
IRBuilder::Current()->FindFrame<PrimFuncFrame>()) {
opt_frame.value()->env_threads.Set(var, iter_var);
diff --git a/src/tirx/script/printer/block.cc b/src/tirx/script/printer/block.cc
index 18e77fe8fe..a8627a51e5 100644
--- a/src/tirx/script/printer/block.cc
+++ b/src/tirx/script/printer/block.cc
@@ -53,7 +53,7 @@ Doc PrintBlock(IRDocsifier d, tirx::SBlock block, AccessPath
block_p, //
PrimExpr value = realize->iter_values[i];
if (iter_var->iter_type == tirx::IterVarType::kDataPar ||
iter_var->iter_type == tirx::IterVarType::kCommReduce) {
- if (auto var = value.as<tirx::PrimVar>()) {
+ if (auto var = value.as<PrimVar>()) {
if (loop_vars.count(var.value().get())) {
tirx::For for_loop = loop_vars.at(var.value().get());
if (expr_equal(for_loop->min, iter_var->dom->min) &&
diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc
index 09982511a5..c9174c2040 100644
--- a/src/tirx/script/printer/expr.cc
+++ b/src/tirx/script/printer/expr.cc
@@ -259,7 +259,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
});
}
-LambdaDoc PrintIndexMap(const ffi::ObjectRef& map, const
ffi::Array<tirx::PrimVar>& vs,
+LambdaDoc PrintIndexMap(const ffi::ObjectRef& map, const ffi::Array<PrimVar>&
vs,
const AccessPath& vs_p, const ffi::Array<PrimExpr>& es,
const AccessPath& es_p, const IRDocsifier& d) {
With<TIRFrame> f(d, map);
diff --git a/tests/cpp/ir_functor_test.cc b/tests/cpp/ir_functor_test.cc
index 8fd773fe26..c3f6b33b16 100644
--- a/tests/cpp/ir_functor_test.cc
+++ b/tests/cpp/ir_functor_test.cc
@@ -51,7 +51,7 @@ TEST(IRF, Basic) {
TEST(IRF, ObjectFunctorDispatch) {
using namespace tvm;
- tirx::PrimVar x("x");
+ PrimVar x("x");
ObjectFunctor<int(const ffi::ObjectRef&)> f;
EXPECT_FALSE(f.CanDispatch(x));
@@ -79,7 +79,7 @@ TEST(IRF, ObjectFunctorDispatch) {
TEST(IRF, ObjectFunctorFinalize) {
using namespace tvm;
- tirx::PrimVar x("x");
+ PrimVar x("x");
PrimExpr z = x + 1;
ObjectFunctor<int(const ffi::ObjectRef&, int)> f;
f.SetDispatch<ExprNode>([](const ffi::ObjectRef&, int b) {
diff --git a/tests/cpp/pattern_match_test.cc b/tests/cpp/pattern_match_test.cc
index bd03e8c236..89973c7868 100644
--- a/tests/cpp/pattern_match_test.cc
+++ b/tests/cpp/pattern_match_test.cc
@@ -26,7 +26,7 @@ TEST(Pattern, Basic) {
using namespace tvm;
using namespace tvm::tirx;
using namespace tvm::arith;
- tvm::tirx::PrimVar x("x"), y("y"), z("z");
+ tvm::PrimVar x("x"), y("y"), z("z");
PrimExpr scalable_lanes = prim::Mul(Call(PrimType::Int(32),
prim::builtin::vscale(), {}), 4);
arith::PVar<PrimExpr> px, py, pz;
arith::PVar<DLDataType> pt;
@@ -131,7 +131,7 @@ TEST(Pattern, Basic) {
TEST(Pattern, IntImm) {
using namespace tvm;
- tirx::PrimVar tx("tx"), ty("ty");
+ PrimVar tx("tx"), ty("ty");
arith::PVar<IntImm> c;
arith::PVar<tirx::Var> v;
{
@@ -151,20 +151,20 @@ TEST(Pattern, MatchWithType) {
using namespace tvm;
// match expr with specified dtype
arith::PVarWithDataType<PrimExpr, arith::PConst<DLDataType>>
pat(DLDataType{kDLFloat, 32, 1});
- tirx::PrimVar x("x", PrimType::Float(32));
- tirx::PrimVar y("y", PrimType::Float(32));
- tirx::PrimVar x_int("x", PrimType::Int(32));
- tirx::PrimVar y_int("y", PrimType::Int(32));
+ PrimVar x("x", PrimType::Float(32));
+ PrimVar y("y", PrimType::Float(32));
+ PrimVar x_int("x", PrimType::Int(32));
+ PrimVar y_int("y", PrimType::Int(32));
TVM_FFI_ICHECK(pat.Match(x + y * 2.0f));
TVM_FFI_ICHECK(!pat.Match(x_int + y_int * 2));
// match vectorized expr with specified element dtype
arith::PVecDataType vec_ty(DLDataType{kDLFloat, 32, 1});
arith::PVarWithDataType<PrimExpr, arith::PVecDataType> vpat(vec_ty);
- tirx::PrimVar vx("x", PrimType::Float(32, 8));
- tirx::PrimVar vy("y", PrimType::Float(32, 8));
- tirx::PrimVar vx_int("x", PrimType::Int(32, 8));
- tirx::PrimVar vy_int("y", PrimType::Int(32, 8));
+ PrimVar vx("x", PrimType::Float(32, 8));
+ PrimVar vy("y", PrimType::Float(32, 8));
+ PrimVar vx_int("x", PrimType::Int(32, 8));
+ PrimVar vy_int("y", PrimType::Int(32, 8));
TVM_FFI_ICHECK(vpat.Match(vx + vy * prim::Broadcast(2.0f, 8)));
TVM_FFI_ICHECK(!vpat.Match(vx_int + vy_int * prim::Broadcast(2, 8)));
}
diff --git a/tests/cpp/tir_analysis_side_effect.cc
b/tests/cpp/tir_analysis_side_effect.cc
index da55576861..6432c0716d 100644
--- a/tests/cpp/tir_analysis_side_effect.cc
+++ b/tests/cpp/tir_analysis_side_effect.cc
@@ -27,7 +27,7 @@
TEST(SimplePasses, SideEffect) {
using namespace tvm;
auto buf = tirx::decl_buffer({16}, PrimType::Float(32));
- auto i = tirx::PrimVar("i", PrimType::Int(32));
+ auto i = PrimVar("i", PrimType::Int(32));
TVM_FFI_ICHECK(tirx::SideEffect(tirx::BufferLoad(buf, {i})) ==
tirx::CallEffectKind::kReadState);
TVM_FFI_ICHECK(tirx::SideEffect(exp(prim::Cast(PrimType::Float(32), i + 1)))
==
tirx::CallEffectKind::kPure);