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);

Reply via email to