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 cc0f9f07c1 [REFACTOR][IR] Unify structural mutation modes and native
traversal entrypoints (#20338)
cc0f9f07c1 is described below
commit cc0f9f07c17c8118a781fdca55f7fe45f7de916a
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 14 15:16:08 2026 -0400
[REFACTOR][IR] Unify structural mutation modes and native traversal
entrypoints (#20338)
Consolidate structural mutation on explicit `InplaceMode` and make the
native visitor and mutator `AnyView` entrypoints virtual. Remove the
redundant ObjectRef and typed Expr entrypoints so qualified parent calls
bypass the current override while recursive descendants retain virtual
dispatch. `ObjectFunctor::CanDispatch` supports borrowed node arguments
and `AnyView` through inline type-index helpers.
Update tvm-ffi to commit `4ffc1f35f3c9853b53f2e854768a9acc6cbf3b85`,
including the merged structural mutation API (apache/tvm-ffi#786), and
migrate native hooks and IR/TIRX/Relax child calls together. Preserve
inherited permission, per-child uniqueness checks, trusted default
descent, and the raw structural signatures and slot layout. Native C++
consumers must rebuild for the virtual interfaces and enum callback
signatures.
---
3rdparty/tvm-ffi | 2 +-
include/tvm/ir/expr_functor.h | 125 +++++++++++------------
include/tvm/ir/object_functor.h | 175 +++++++++++++-------------------
src/ir/expr.cc | 47 +++++----
src/ir/expr_functor.cc | 215 +++++++++++++++++++---------------------
src/ir/prim/expr.cc | 30 +++---
src/ir/prim/vector_expr.cc | 23 +++--
src/ir/type.cc | 17 ++--
src/relax/distributed/type.cc | 15 +--
src/relax/ir/dependent_type.cc | 19 ++--
src/relax/ir/expr.cc | 52 +++++-----
src/tirx/ir/buffer.cc | 21 ++--
src/tirx/ir/function.cc | 15 +--
src/tirx/ir/iter_var.cc | 5 +-
src/tirx/ir/layout/tile_core.cc | 24 +++--
src/tirx/ir/stmt.cc | 169 +++++++++++++++++--------------
src/tirx/ir/tirx_stmt.cc | 12 ++-
tests/cpp/expr_functor_test.cc | 3 +-
18 files changed, 491 insertions(+), 478 deletions(-)
diff --git a/3rdparty/tvm-ffi b/3rdparty/tvm-ffi
index be35ec1297..4ffc1f35f3 160000
--- a/3rdparty/tvm-ffi
+++ b/3rdparty/tvm-ffi
@@ -1 +1 @@
-Subproject commit be35ec1297d6ddf388b7feb4f2e89d2cf030dda8
+Subproject commit 4ffc1f35f3c9853b53f2e854768a9acc6cbf3b85
diff --git a/include/tvm/ir/expr_functor.h b/include/tvm/ir/expr_functor.h
index b7115ed6c5..f90a5f0d7d 100644
--- a/include/tvm/ir/expr_functor.h
+++ b/include/tvm/ir/expr_functor.h
@@ -367,7 +367,8 @@ class TVM_DLL ExprVisitor : public ObjectVisitor {
* public:
* TVM_DEFINE_OBJECT_FUNCTOR_DEFAULT_CONSTRUCTOR(MyExprMutator, ExprMutator)
* using ExprMutator::Mutate_;
- * virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const MyExprNode* node,
bool allow_inplace);
+ * virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const MyExprNode* node,
InplaceMode
+ * inplace_mode);
*
* protected:
* static void InitVTable(VTable* vtable) {
@@ -383,77 +384,77 @@ class TVM_DLL ExprMutator : public ObjectMutator {
/*! \brief Construct a mutator with the core expression hooks. */
TVM_DEFINE_OBJECT_FUNCTOR_DEFAULT_CONSTRUCTOR(ExprMutator, ObjectMutator)
- using ObjectMutator::MaybeInplaceMutateIfUniqueExpected;
using ObjectMutator::MutateExpected;
- /*!
- * \brief Mutate a borrowed expression without allowing in-place changes.
- * \param value The borrowed expression to mutate.
- * \return An Expr replacement, Unchanged, or an Error if mutation fails.
- * \note This entry trusts the Expr replacement contract of native hooks and
extensions.
- */
- TVM_FFI_INLINE Expected<UnchangedOr<Expr>> MutateExpected(const Expr& value)
noexcept {
- return ffi::details::ExpectedUnsafe::MoveFromTVMFFIAny<UnchangedOr<Expr>>(
-
ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(ObjectMutator::MutateExpected(value)));
- }
- /*!
- * \brief Forward inherited permission and check the expression's uniqueness.
- * \param value The borrowed expression to mutate.
- * \param allow_inplace Whether the path to this expression is already
uniquely owned.
- * \return An Expr replacement, Unchanged, or an Error if mutation fails.
- * \note This entry trusts the Expr replacement contract of native hooks and
extensions.
- */
- TVM_FFI_INLINE Expected<UnchangedOr<Expr>>
MaybeInplaceMutateIfUniqueExpected(
- const Expr& value, bool allow_inplace = true) noexcept {
- return ffi::details::ExpectedUnsafe::MoveFromTVMFFIAny<UnchangedOr<Expr>>(
- ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(
- ObjectMutator::MaybeInplaceMutateIfUniqueExpected(value,
allow_inplace)));
- }
-
// A downstream class overrides any existing hook without rebuilding the
table.
// Extra node types use a fresh inherited table and SetDispatch<Self,
ExtraNode>.
// Hooks borrow the node and return an owning Expr replacement in Any,
Unchanged, or Error.
- // Narrower field types are checked separately. Forward allow_inplace on
every child edge.
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const OpaqueExprNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TupleNode* node, bool
allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TupleGetItemNode*
node, bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TensorLoadNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const VarNode* node, bool
allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const GlobalVarNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const CallNode* node, bool
allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const IntImmNode* node, bool
allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const FloatImmNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const OpNode* node, bool
allow_inplace);
+ // Narrower field types are checked separately. Forward inplace_mode on
every child edge.
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const OpaqueExprNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TupleNode* node,
InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TupleGetItemNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const TensorLoadNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const VarNode* node,
InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const GlobalVarNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const CallNode* node,
InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const IntImmNode* node,
InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const FloatImmNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const OpNode* node,
InplaceMode inplace_mode);
virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::StringImmNode*
node,
- bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::CastNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::AddNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::SubNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MulNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::DivNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::ModNode* node,
bool allow_inplace);
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::CastNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::AddNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::SubNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MulNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::DivNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::ModNode* node,
+ InplaceMode inplace_mode);
virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::FloorDivNode*
node,
- bool allow_inplace);
+ InplaceMode inplace_mode);
virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::FloorModNode*
node,
- bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MinNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MaxNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::EQNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::NENode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LTNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LENode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::GTNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::GENode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::AndNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::OrNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::NotNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::SelectNode*
node, bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LetNode* node,
bool allow_inplace);
- virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::RampNode* node,
bool allow_inplace);
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MinNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::MaxNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::EQNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::NENode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LTNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LENode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::GTNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::GENode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::AndNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::OrNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::NotNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::SelectNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::LetNode* node,
+ InplaceMode inplace_mode);
+ virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::RampNode* node,
+ InplaceMode inplace_mode);
virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::BroadcastNode*
node,
- bool allow_inplace);
+ InplaceMode inplace_mode);
virtual Expected<UnchangedOr<ffi::Any>> Mutate_(const prim::ShuffleNode*
node,
- bool allow_inplace);
+ InplaceMode inplace_mode);
protected:
/*!
diff --git a/include/tvm/ir/object_functor.h b/include/tvm/ir/object_functor.h
index 037a1d0673..a45251b42e 100644
--- a/include/tvm/ir/object_functor.h
+++ b/include/tvm/ir/object_functor.h
@@ -95,11 +95,13 @@ class ObjectFunctor<R(NodeArg, Args...)> {
using result_type = R;
/*!
* \brief Whether a dispatch function is registered for the exact runtime
type.
- * \param n The object to be dispatched.
+ * \param n The borrowed NodeArg or AnyView value to be dispatched.
+ * \tparam T The input representation, convertible to NodeArg or AnyView.
* \return Whether a dispatch function is registered for n's type, excluding
ancestors.
*/
- TVM_FFI_INLINE bool CanDispatch(NodeArg n) const {
- uint32_t type_index = n->type_index();
+ template <typename T>
+ TVM_FFI_INLINE bool CanDispatch(const T& n) const {
+ uint32_t type_index = GetTypeIndex(n);
if (type_index < begin_type_index_) return false;
type_index -= begin_type_index_;
return type_index < func_.size() && func_[type_index] != nullptr;
@@ -139,6 +141,8 @@ class ObjectFunctor<R(NodeArg, Args...)> {
*/
template <typename TNode>
ObjectFunctor& SetDispatch(R (*f)(NodeArg n, Args...)) {
+ static_assert(std::is_base_of_v<ffi::Object, TNode>,
+ "ObjectFunctor dispatch requires an object node type");
uint32_t tindex = TNode::RuntimeTypeIndex();
if (func_.size() <= tindex) {
func_.resize(tindex + 1, nullptr);
@@ -184,6 +188,17 @@ class ObjectFunctor<R(NodeArg, Args...)> {
}
private:
+ // Map null pointers and undefined handles to None so CanDispatch rejects
them without
+ // dereferencing.
+ TVM_FFI_INLINE static uint32_t GetTypeIndex(NodeArg n) {
+ if constexpr (std::is_pointer_v<NodeArg>) {
+ return n != nullptr ? n->type_index() : ffi::TypeIndex::kTVMFFINone;
+ } else {
+ return n.defined() ? n->type_index() : ffi::TypeIndex::kTVMFFINone;
+ }
+ }
+ TVM_FFI_INLINE static uint32_t GetTypeIndex(ffi::AnyView n) { return
n.type_index(); }
+
[[noreturn]] TVM_FFI_COLD_CODE static void ThrowUnregistered(NodeArg n) {
TVM_FFI_THROW(InternalError) << "ObjectFunctor calls un-registered
function on type "
<< n->GetTypeKey();
@@ -199,6 +214,7 @@ class ObjectFunctor<R(NodeArg, Args...)> {
};
using ffi::Expected;
+using ffi::InplaceMode;
using ffi::UnchangedOr;
using ffi::VisitInterrupt;
@@ -218,23 +234,26 @@ class TVM_DLL ObjectVisitor : public
ffi::StructuralVisitorObj {
ObjectVisitor(const ObjectVisitor& other) = delete;
ObjectVisitor& operator=(const ObjectVisitor& other) = delete;
- /*!
- * \brief Visit a borrowed object and propagate an interrupt or error.
- * \param value The borrowed object to visit.
- * \return None on completion, an owning VisitInterrupt, or an Error.
- */
- TVM_FFI_INLINE Expected<ffi::Optional<VisitInterrupt>> VisitExpected(
- const ffi::ObjectRef& value) noexcept {
- return VisitExpected(ffi::AnyView(value));
- }
/*!
* \brief Visit a borrowed object or inline value.
* \param value The borrowed value to visit.
* \return None on completion, an owning VisitInterrupt, or an Error.
+ * \note Overrides must convert thrown errors and propagate errors and
interrupts from their
+ * own work. A qualified Parent::VisitExpected(value) bypasses the
current entry override;
+ * descendant calls still use virtual dispatch.
*/
- TVM_FFI_INLINE Expected<ffi::Optional<VisitInterrupt>> VisitExpected(
- ffi::AnyView value) noexcept {
- if (const auto* object = value.as<ffi::Object>()) return Dispatch(object);
+ virtual Expected<ffi::Optional<VisitInterrupt>> VisitExpected(ffi::AnyView
value) noexcept {
+ if (native_vtable_->CanDispatch(value)) {
+ // Exact registrations contain only object node types, so this
extraction is borrowed
+ // and needs no additional object-type check or owning reference.
+ const auto* object =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ffi::Object>(value);
+ try {
+ return DispatchNative(object);
+ } catch (ffi::Error& error) {
+ return AttachVisitErrorContext(error, object);
+ }
+ }
return StructuralVisitDefault(value);
}
@@ -275,18 +294,6 @@ class TVM_DLL ObjectVisitor : public
ffi::StructuralVisitorObj {
return static_cast<Self*>(self)->Visit_(static_cast<const Node*>(node));
}
- // The AnyView entry establishes that value is a non-null object.
- TVM_FFI_INLINE Expected<ffi::Optional<VisitInterrupt>> Dispatch(
- const ffi::Object* value) noexcept {
- if (native_vtable_->CanDispatch(value)) {
- try {
- return DispatchNative(value);
- } catch (ffi::Error& error) {
- return AttachVisitErrorContext(error, value);
- }
- }
- return StructuralVisitDefault(value);
- }
// Keep one named return value in this scope so native results can be
constructed in place.
TVM_FFI_INLINE Expected<ffi::Optional<VisitInterrupt>> DispatchNative(const
ffi::Object* value) {
Expected<ffi::Optional<VisitInterrupt>> result = (*native_vtable_)(value,
this);
@@ -339,7 +346,7 @@ class TVM_DLL ObjectVisitor : public
ffi::StructuralVisitorObj {
* \brief Native mutation with exact dispatch and structural fallback.
*
* Allocate mutators with ffi::make_object<Derived>(). Inputs are borrowed;
- * replacements and errors own their values. Overrides forward allow_inplace
+ * replacements and errors own their values. Overrides forward inplace_mode
* to child calls, following the structural mutation ownership contract.
*/
class TVM_DLL ObjectMutator : public ffi::StructuralMapEngineBase {
@@ -352,45 +359,38 @@ class TVM_DLL ObjectMutator : public
ffi::StructuralMapEngineBase {
ObjectMutator& operator=(const ObjectMutator& other) = delete;
/*!
- * \brief Mutate a borrowed value without allowing in-place changes.
- * \param value The borrowed object to mutate.
- * \return A replacement, Unchanged, or an Error if mutation fails.
- */
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> MutateExpected(
- const ffi::ObjectRef& value) noexcept {
- return Dispatch(value.get(), false);
- }
- /*!
- * \brief Mutate a borrowed value without allowing in-place changes.
+ * \brief Mutate a borrowed value, checking current uniqueness before
in-place dispatch.
* \param value The borrowed object or inline value to mutate.
+ * \param inplace_mode Inherited permission along the path to this value.
* \return A replacement, Unchanged, or an Error if mutation fails.
+ * \note The base implementation establishes current uniqueness before
invoking typed hooks.
+ * Entry overrides must check uniqueness before writing or forwarding
permission,
+ * never promote kDisallow, and convert thrown errors from their own
work into Expected.
+ * Expr inputs require Expr replacements. A qualified
Parent::MutateExpected(value, mode)
+ * bypasses the current entry override while descendants remain
virtual.
+ * DefaultMutateExpected trusts its caller's established mode and does
not recheck it.
*/
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> MutateExpected(ffi::AnyView
value) noexcept {
- if (const auto* object = value.as<ffi::Object>()) return Dispatch(object,
false);
- return StructuralMutateDefault(value, false);
- }
-
- /*!
- * \brief Forward inherited permission and check the value's uniqueness.
- * \param value The borrowed object to mutate.
- * \param allow_inplace Whether the path to this value is already uniquely
owned.
- * \return A replacement, Unchanged, or an Error if mutation fails.
- */
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>>
MaybeInplaceMutateIfUniqueExpected(
- const ffi::ObjectRef& value, bool allow_inplace = true) noexcept {
- return Dispatch(value.get(), allow_inplace && value.defined() &&
value->unique());
- }
- /*!
- * \brief Forward inherited permission and check the value's uniqueness.
- * \param value The borrowed object or inline value to mutate.
- * \param allow_inplace Whether the path to this value is already uniquely
owned.
- * \return A replacement, Unchanged, or an Error if mutation fails.
- */
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>>
MaybeInplaceMutateIfUniqueExpected(
- ffi::AnyView value, bool allow_inplace = true) noexcept {
+ virtual Expected<UnchangedOr<ffi::Any>> MutateExpected(
+ ffi::AnyView value, InplaceMode inplace_mode = InplaceMode::kDisallow)
noexcept {
+ if (native_vtable_->CanDispatch(value)) {
+ // Only object node types can be registered. Exact dispatch proves this
borrowed extraction
+ // is valid without an additional object-type check or owning reference.
+ const auto* object =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ffi::Object>(value);
+ // Modes encode disallow/allow as 0/1: preserve inherited permission
only for a unique object.
+ inplace_mode =
+ static_cast<InplaceMode>(inplace_mode == InplaceMode::kAllow &&
object->unique());
+ try {
+ return DispatchNative(object, inplace_mode);
+ } catch (ffi::Error& error) {
+ return AttachVisitErrorContext(error, object);
+ }
+ }
const auto* object = value.as<ffi::Object>();
- if (object) return Dispatch(object, allow_inplace && object->unique());
- return StructuralMutateDefault(value, false);
+ // The same 0/1 permission rule also excludes inline values and None.
+ inplace_mode = static_cast<InplaceMode>(inplace_mode ==
InplaceMode::kAllow &&
+ object != nullptr &&
object->unique());
+ return ffi::StructuralMutatorObj::DefaultMutateExpected(value,
inplace_mode);
}
/*!
@@ -414,8 +414,8 @@ class TVM_DLL ObjectMutator : public
ffi::StructuralMapEngineBase {
protected:
/*! \brief Exact native dispatch table with owning replacement and error
results. */
- using VTable =
- ObjectFunctor<Expected<UnchangedOr<ffi::Any>>(const ffi::Object*,
ObjectMutator*, bool)>;
+ using VTable = ObjectFunctor<Expected<UnchangedOr<ffi::Any>>(const
ffi::Object*, ObjectMutator*,
+ InplaceMode)>;
/*!
* \brief Construct a mutator with a finalized native dispatch table.
@@ -445,55 +445,22 @@ class TVM_DLL ObjectMutator : public
ffi::StructuralMapEngineBase {
// Native table callback; core instantiations may be shared by the library.
template <typename Self, typename Node>
static Expected<UnchangedOr<ffi::Any>> DispatchNode(const ffi::Object* node,
ObjectMutator* self,
- bool allow_inplace) {
- return static_cast<Self*>(self)->Mutate_(static_cast<const Node*>(node),
allow_inplace);
+ InplaceMode
inplace_mode) {
+ return static_cast<Self*>(self)->Mutate_(static_cast<const Node*>(node),
inplace_mode);
}
- using ffi::StructuralMutatorObj::DefaultMaybeInplaceMutateExpected;
using ffi::StructuralMutatorObj::DefaultMutateExpected;
- using ffi::StructuralMutatorObj::MaybeInplaceMutate;
using ffi::StructuralMutatorObj::Mutate;
- // Structural ABI entry: the caller guarantees ownership of the entire path.
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> MaybeInplaceMutateExpected(
- const ffi::ObjectRef& value) noexcept {
- return Dispatch(value.get(), true);
- }
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> MaybeInplaceMutateExpected(
- ffi::AnyView value) noexcept {
- if (const auto* object = value.as<ffi::Object>()) return Dispatch(object,
true);
- return StructuralMutateDefault(value, true);
- }
-
- TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> Dispatch(const ffi::Object*
value,
- bool allow_inplace)
noexcept {
- if (value == nullptr) return ffi::Unchanged();
- if (native_vtable_->CanDispatch(value)) {
- try {
- return DispatchNative(value, allow_inplace);
- } catch (ffi::Error& error) {
- return AttachVisitErrorContext(error, value);
- }
- }
- return StructuralMutateDefault(value, allow_inplace);
- }
// Keep one named return value in this scope so native results can be
constructed in place.
TVM_FFI_INLINE Expected<UnchangedOr<ffi::Any>> DispatchNative(const
ffi::Object* value,
- bool
allow_inplace) {
- Expected<UnchangedOr<ffi::Any>> result = (*native_vtable_)(value, this,
allow_inplace);
+ InplaceMode
inplace_mode) {
+ Expected<UnchangedOr<ffi::Any>> result = (*native_vtable_)(value, this,
inplace_mode);
if (TVM_FFI_PREDICT_FALSE(result.is_err())) {
UpdateVisitErrorContext(result, value);
}
return result;
}
- Expected<UnchangedOr<ffi::Any>> StructuralMutateDefault(ffi::AnyView value,
- bool allow_inplace)
noexcept {
- if (allow_inplace) {
- return
ffi::StructuralMutatorObj::DefaultMaybeInplaceMutateExpected(value);
- } else {
- return ffi::StructuralMutatorObj::DefaultMutateExpected(value);
- }
- }
TVM_FFI_COLD_CODE static Expected<UnchangedOr<ffi::Any>>
AttachVisitErrorContext(
ffi::Error& error, const ffi::Object* value) {
if (value) ffi::details::UpdateVisitErrorContext(error,
ffi::GetRef<ffi::ObjectRef>(value));
@@ -522,12 +489,12 @@ class TVM_DLL ObjectMutator : public
ffi::StructuralMapEngineBase {
static TVMFFIAny StructuralVTableMutateImpl(ffi::StructuralMutatorObj* self,
ffi::AnyView value) noexcept {
return ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(
- static_cast<ObjectMutator*>(self)->MutateExpected(value));
+ static_cast<ObjectMutator*>(self)->MutateExpected(value,
InplaceMode::kDisallow));
}
static TVMFFIAny
StructuralVTableMaybeInplaceMutateImpl(ffi::StructuralMutatorObj* self,
ffi::AnyView value)
noexcept {
return ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(
- static_cast<ObjectMutator*>(self)->MaybeInplaceMutateExpected(value));
+ static_cast<ObjectMutator*>(self)->MutateExpected(value,
InplaceMode::kAllow));
}
static TVMFFIAny StructuralVTableVarRemapGetImpl(ffi::StructuralMutatorObj*
self,
ffi::AnyView key) noexcept {
diff --git a/src/ir/expr.cc b/src/ir/expr.cc
index e72902bef6..e2fb8b30ef 100644
--- a/src/ir/expr.cc
+++ b/src/ir/expr.cc
@@ -68,7 +68,7 @@ TVMFFIAny
OpaqueExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
OpaqueExprNode* self = const_cast<OpaqueExprNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
OpaqueExprNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
if (mapped_ty.UnchangedOrSameAs(self->ty)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -110,11 +110,13 @@ TVMFFIAny
TensorLoadMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TensorLoadNode* self = const_cast<TensorLoadNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TensorLoadNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_source,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->source));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_indices,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->indices));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<Expr>, mapped_source,
+ mutator->MutateExpected(self->source, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_indices,
+ mutator->MutateExpected(self->indices, ffi::InplaceMode::kAllow));
if (mapped_ty.UnchangedOrSameAs(self->ty) &&
mapped_source.UnchangedOrSameAs(self->source) &&
mapped_indices.UnchangedOrSameAs(self->indices)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -153,9 +155,10 @@ TVMFFIAny
TupleMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyVi
TupleNode* self = const_cast<TupleNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TupleNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Expr>>,
mapped_fields,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->fields));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Expr>>, mapped_fields,
+ mutator->MutateExpected(self->fields, ffi::InplaceMode::kAllow));
if (mapped_ty.UnchangedOrSameAs(self->ty) &&
mapped_fields.UnchangedOrSameAs(self->fields)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -196,9 +199,9 @@ TVMFFIAny
TupleGetItemMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TupleGetItemNode* self = const_cast<TupleGetItemNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TupleGetItemNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_tuple,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->tuple));
+ mutator->MutateExpected(self->tuple,
ffi::InplaceMode::kAllow));
if (mapped_ty.UnchangedOrSameAs(self->ty) &&
mapped_tuple.UnchangedOrSameAs(self->tuple)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -265,9 +268,10 @@ TVMFFIAny
RangeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyVi
RangeNode* self = const_cast<RangeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
RangeNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_min,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->min));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_extent,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->extent));
+ mutator->MutateExpected(self->min,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_extent,
+ mutator->MutateExpected(self->extent, ffi::InplaceMode::kAllow));
if (mapped_min.UnchangedOrSameAs(self->min) &&
mapped_extent.UnchangedOrSameAs(self->extent)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -362,8 +366,8 @@ TVMFFIAny VarMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyView
mutator->def_region_kind() == kTVMFFIDefRegionKindSimple
? mutator->WithDefRegionKind(
kTVMFFIDefRegionKindNone,
- [&]() { return
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty); })
- : mutator->MaybeInplaceMutateIfUniqueExpected(self->ty);
+ [&]() { return mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow); })
+ : mutator->MutateExpected(self->ty, ffi::InplaceMode::kAllow);
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
std::move(mapped_ty_result));
if (!mapped_ty.UnchangedOrSameAs(self->ty)) {
@@ -472,7 +476,7 @@ TVMFFIAny CallMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyVie
ffi::UnchangedOr<Type> mapped_ty = ffi::Unchanged();
if (!self->ty.as<PrimTypeNode>()) {
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, descended_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
mapped_ty = std::move(descended_ty);
}
// An Op is an interned registry singleton, so it has nothing to substitute.
Broad callbacks do
@@ -480,17 +484,18 @@ TVMFFIAny
CallMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyVie
ffi::UnchangedOr<Expr> mapped_op = ffi::Unchanged();
if (!self->op.as<OpNode>()) {
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, descended_op,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->op));
+ mutator->MutateExpected(self->op,
ffi::InplaceMode::kAllow));
mapped_op = std::move(descended_op);
}
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Expr>>,
mapped_args,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->args));
+ mutator->MutateExpected(self->args,
ffi::InplaceMode::kAllow));
// An empty ty_args has no element to substitute. Broad callbacks do not
see the empty
// container; nonempty type arguments retain normal container descent and
callback behavior.
ffi::UnchangedOr<ffi::Array<Type>> mapped_ty_args = ffi::Unchanged();
if (!self->ty_args.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Type>>,
descended_ty_args,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty_args));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Type>>, descended_ty_args,
+ mutator->MutateExpected(self->ty_args, ffi::InplaceMode::kAllow));
mapped_ty_args = std::move(descended_ty_args);
}
if (mapped_ty.UnchangedOrSameAs(self->ty) &&
mapped_op.UnchangedOrSameAs(self->op) &&
diff --git a/src/ir/expr_functor.cc b/src/ir/expr_functor.cc
index bbdd025639..c4f86f4bd2 100644
--- a/src/ir/expr_functor.cc
+++ b/src/ir/expr_functor.cc
@@ -321,11 +321,11 @@ void ExprMutator::InitVTable(VTable* vtable) {
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const OpaqueExprNode*
node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Type>, ty,
this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, ty,
+ this->MutateExpected(node->ty,
inplace_mode));
if (ty.UnchangedOrSameAs(node->ty)) return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<OpaqueExprNode*>(node);
if (!ty.IsUnchanged()) writable->ty = std::move(ty).ValueUnchecked();
return ffi::Unchanged();
@@ -336,15 +336,15 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const OpaqueExprNode* node,
}
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const TupleNode* node,
bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Type>, ty,
this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<Expr>>, fields,
- this->MaybeInplaceMutateIfUniqueExpected(node->fields, allow_inplace));
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const TupleNode* node,
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, ty,
+ this->MutateExpected(node->ty,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<Expr>>, fields,
+ this->MutateExpected(node->fields,
inplace_mode));
if (ty.UnchangedOrSameAs(node->ty) && fields.UnchangedOrSameAs(node->fields))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<TupleNode*>(node);
if (!ty.IsUnchanged()) writable->ty = std::move(ty).ValueUnchecked();
if (!fields.IsUnchanged()) writable->fields =
std::move(fields).ValueUnchecked();
@@ -358,14 +358,14 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const TupleNode* node, bool
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const TupleGetItemNode*
node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Type>, ty,
this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, ty,
+ this->MutateExpected(node->ty,
inplace_mode));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Expr>, tuple,
-
MaybeInplaceMutateIfUniqueExpected(node->tuple, allow_inplace));
+ MutateExpected(node->tuple, inplace_mode));
if (ty.UnchangedOrSameAs(node->ty) && tuple.UnchangedOrSameAs(node->tuple))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<TupleGetItemNode*>(node);
if (!ty.IsUnchanged()) writable->ty = std::move(ty).ValueUnchecked();
if (!tuple.IsUnchanged()) writable->tuple =
std::move(tuple).ValueUnchecked();
@@ -379,18 +379,17 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const TupleGetItemNode* nod
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const TensorLoadNode*
node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Type>, ty,
this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Expr>, source,
MaybeInplaceMutateIfUniqueExpected(node->source, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<PrimExpr>>, indices,
- this->MaybeInplaceMutateIfUniqueExpected(node->indices, allow_inplace));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, ty,
+ this->MutateExpected(node->ty,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Expr>, source,
+ MutateExpected(node->source,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<PrimExpr>>, indices,
+ this->MutateExpected(node->indices,
inplace_mode));
if (ty.UnchangedOrSameAs(node->ty) && source.UnchangedOrSameAs(node->source)
&&
indices.UnchangedOrSameAs(node->indices))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<TensorLoadNode*>(node);
if (!ty.IsUnchanged()) writable->ty = std::move(ty).ValueUnchecked();
if (!source.IsUnchanged()) writable->source =
std::move(source).ValueUnchecked();
@@ -406,39 +405,37 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const TensorLoadNode* node,
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const GlobalVarNode* node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
// Registry atoms and constant leaves do not descend into metadata or types.
return ffi::Unchanged();
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const CallNode* node,
bool allow_inplace) {
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const CallNode* node,
+ InplaceMode inplace_mode)
{
UnchangedOr<Type> ty = ffi::Unchanged();
if (!node->ty.as<PrimTypeNode>()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Type>, mapped_ty,
- this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, mapped_ty,
+ this->MutateExpected(node->ty,
inplace_mode));
ty = std::move(mapped_ty);
}
UnchangedOr<Expr> op = ffi::Unchanged();
if (!node->op.as<OpNode>()) {
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Expr>, mapped_op,
-
MaybeInplaceMutateIfUniqueExpected(node->op, allow_inplace));
+ MutateExpected(node->op, inplace_mode));
op = std::move(mapped_op);
}
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<Expr>>, args,
- this->MaybeInplaceMutateIfUniqueExpected(node->args, allow_inplace));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<Expr>>, args,
+ this->MutateExpected(node->args,
inplace_mode));
UnchangedOr<ffi::Array<Type>> ty_args = ffi::Unchanged();
if (!node->ty_args.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<Type>>, mapped_ty_args,
- this->MaybeInplaceMutateIfUniqueExpected(node->ty_args,
allow_inplace));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<Type>>,
mapped_ty_args,
+ this->MutateExpected(node->ty_args,
inplace_mode));
ty_args = std::move(mapped_ty_args);
}
if (ty.UnchangedOrSameAs(node->ty) && op.UnchangedOrSameAs(node->op) &&
args.UnchangedOrSameAs(node->args) &&
ty_args.UnchangedOrSameAs(node->ty_args))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<CallNode*>(node);
if (!ty.IsUnchanged()) writable->ty = std::move(ty).ValueUnchecked();
if (!op.IsUnchanged()) writable->op = std::move(op).ValueUnchecked();
@@ -455,33 +452,35 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const CallNode* node, bool
}
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const IntImmNode* node,
bool allow_inplace) {
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const IntImmNode* node,
+ InplaceMode inplace_mode)
{
// Registry atoms and constant leaves do not descend into metadata or types.
return ffi::Unchanged();
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const FloatImmNode* node,
bool allow_inplace) {
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const FloatImmNode* node,
+ InplaceMode inplace_mode)
{
// Registry atoms and constant leaves do not descend into metadata or types.
return ffi::Unchanged();
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const OpNode* node, bool
allow_inplace) {
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const OpNode* node,
InplaceMode inplace_mode) {
// Registry atoms and constant leaves do not descend into metadata or types.
return ffi::Unchanged();
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const
prim::StringImmNode* node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
// Registry atoms and constant leaves do not descend into metadata or types.
return ffi::Unchanged();
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::CastNode*
node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, value,
-
MaybeInplaceMutateIfUniqueExpected(node->value, allow_inplace));
+ MutateExpected(node->value, inplace_mode));
if (value.UnchangedOrSameAs(node->value)) return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::CastNode*>(node);
if (!value.IsUnchanged()) writable->value =
std::move(value).ValueUnchecked();
return ffi::Unchanged();
@@ -492,25 +491,25 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::CastNode* node,
}
}
-#define TVM_IR_BINARY_MUTATE_IMPL(Name)
\
- Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::Name##Node*
node, \
- bool allow_inplace) {
\
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, a,
\
-
MaybeInplaceMutateIfUniqueExpected(node->a, allow_inplace)); \
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, b,
\
-
MaybeInplaceMutateIfUniqueExpected(node->b, allow_inplace)); \
- if (a.UnchangedOrSameAs(node->a) && b.UnchangedOrSameAs(node->b)) return
ffi::Unchanged(); \
- if (allow_inplace) {
\
- auto* writable = const_cast<prim::Name##Node*>(node);
\
- if (!a.IsUnchanged()) writable->a = std::move(a).ValueUnchecked();
\
- if (!b.IsUnchanged()) writable->b = std::move(b).ValueUnchecked();
\
- return ffi::Unchanged();
\
- } else {
\
- auto copy = ffi::make_object<prim::Name##Node>(*node);
\
- copy->a = std::move(a).ValueOrUnchanged(std::move(copy->a));
\
- copy->b = std::move(b).ValueOrUnchanged(std::move(copy->b));
\
- return ffi::Any(Expr(std::move(copy)));
\
- }
\
+#define TVM_IR_BINARY_MUTATE_IMPL(Name)
\
+ Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::Name##Node*
node, \
+ InplaceMode
inplace_mode) { \
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, a,
\
+ MutateExpected(node->a, inplace_mode));
\
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, b,
\
+ MutateExpected(node->b, inplace_mode));
\
+ if (a.UnchangedOrSameAs(node->a) && b.UnchangedOrSameAs(node->b)) return
ffi::Unchanged(); \
+ if (inplace_mode == InplaceMode::kAllow) {
\
+ auto* writable = const_cast<prim::Name##Node*>(node);
\
+ if (!a.IsUnchanged()) writable->a = std::move(a).ValueUnchecked();
\
+ if (!b.IsUnchanged()) writable->b = std::move(b).ValueUnchecked();
\
+ return ffi::Unchanged();
\
+ } else {
\
+ auto copy = ffi::make_object<prim::Name##Node>(*node);
\
+ copy->a = std::move(a).ValueOrUnchanged(std::move(copy->a));
\
+ copy->b = std::move(b).ValueOrUnchanged(std::move(copy->b));
\
+ return ffi::Any(Expr(std::move(copy)));
\
+ }
\
}
TVM_IR_BINARY_MUTATE_IMPL(Add)
TVM_IR_BINARY_MUTATE_IMPL(Sub)
@@ -532,11 +531,11 @@ TVM_IR_BINARY_MUTATE_IMPL(Or)
#undef TVM_IR_BINARY_MUTATE_IMPL
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::NotNode* node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, a,
-
MaybeInplaceMutateIfUniqueExpected(node->a, allow_inplace));
+ MutateExpected(node->a, inplace_mode));
if (a.UnchangedOrSameAs(node->a)) return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::NotNode*>(node);
if (!a.IsUnchanged()) writable->a = std::move(a).ValueUnchecked();
return ffi::Unchanged();
@@ -548,21 +547,18 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::NotNode* node,
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::SelectNode*
node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<PrimExpr>, condition,
- MaybeInplaceMutateIfUniqueExpected(node->condition, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<PrimExpr>, true_value,
- MaybeInplaceMutateIfUniqueExpected(node->true_value, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<PrimExpr>, false_value,
- MaybeInplaceMutateIfUniqueExpected(node->false_value, allow_inplace));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, condition,
+ MutateExpected(node->condition,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, true_value,
+ MutateExpected(node->true_value,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, false_value,
+ MutateExpected(node->false_value,
inplace_mode));
if (condition.UnchangedOrSameAs(node->condition) &&
true_value.UnchangedOrSameAs(node->true_value) &&
false_value.UnchangedOrSameAs(node->false_value))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::SelectNode*>(node);
if (!condition.IsUnchanged()) writable->condition =
std::move(condition).ValueUnchecked();
if (!true_value.IsUnchanged()) writable->true_value =
std::move(true_value).ValueUnchecked();
@@ -578,19 +574,19 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::SelectNode* nod
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::LetNode* node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<Var>, var, WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]
{
- return MaybeInplaceMutateIfUniqueExpected(node->var, allow_inplace);
- }));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Var>, var,
+
WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&] {
+ return MutateExpected(node->var,
inplace_mode);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, value,
-
MaybeInplaceMutateIfUniqueExpected(node->value, allow_inplace));
+ MutateExpected(node->value, inplace_mode));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, body,
-
MaybeInplaceMutateIfUniqueExpected(node->body, allow_inplace));
+ MutateExpected(node->body, inplace_mode));
if (var.UnchangedOrSameAs(node->var) && value.UnchangedOrSameAs(node->value)
&&
body.UnchangedOrSameAs(node->body))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::LetNode*>(node);
if (!var.IsUnchanged()) writable->var = std::move(var).ValueUnchecked();
if (!value.IsUnchanged()) writable->value =
std::move(value).ValueUnchecked();
@@ -606,18 +602,17 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::LetNode* node,
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::RampNode*
node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, base,
-
MaybeInplaceMutateIfUniqueExpected(node->base, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<PrimExpr>, stride,
- MaybeInplaceMutateIfUniqueExpected(node->stride, allow_inplace));
+ MutateExpected(node->base, inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, stride,
+ MutateExpected(node->stride,
inplace_mode));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, lanes,
-
MaybeInplaceMutateIfUniqueExpected(node->lanes, allow_inplace));
+ MutateExpected(node->lanes, inplace_mode));
if (base.UnchangedOrSameAs(node->base) &&
stride.UnchangedOrSameAs(node->stride) &&
lanes.UnchangedOrSameAs(node->lanes))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::RampNode*>(node);
if (!base.IsUnchanged()) writable->base = std::move(base).ValueUnchecked();
if (!stride.IsUnchanged()) writable->stride =
std::move(stride).ValueUnchecked();
@@ -633,14 +628,14 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::RampNode* node,
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const
prim::BroadcastNode* node,
- bool allow_inplace) {
+ InplaceMode inplace_mode)
{
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, value,
-
MaybeInplaceMutateIfUniqueExpected(node->value, allow_inplace));
+ MutateExpected(node->value, inplace_mode));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<PrimExpr>, lanes,
-
MaybeInplaceMutateIfUniqueExpected(node->lanes, allow_inplace));
+ MutateExpected(node->lanes, inplace_mode));
if (value.UnchangedOrSameAs(node->value) &&
lanes.UnchangedOrSameAs(node->lanes))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::BroadcastNode*>(node);
if (!value.IsUnchanged()) writable->value =
std::move(value).ValueUnchecked();
if (!lanes.IsUnchanged()) writable->lanes =
std::move(lanes).ValueUnchecked();
@@ -654,16 +649,14 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const prim::BroadcastNode*
}
Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const prim::ShuffleNode*
node,
- bool allow_inplace) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<PrimExpr>>, vectors,
- this->MaybeInplaceMutateIfUniqueExpected(node->vectors, allow_inplace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- UnchangedOr<ffi::Array<PrimExpr>>, indices,
- this->MaybeInplaceMutateIfUniqueExpected(node->indices, allow_inplace));
+ InplaceMode inplace_mode)
{
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<PrimExpr>>, vectors,
+ this->MutateExpected(node->vectors,
inplace_mode));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<ffi::Array<PrimExpr>>, indices,
+ this->MutateExpected(node->indices,
inplace_mode));
if (vectors.UnchangedOrSameAs(node->vectors) &&
indices.UnchangedOrSameAs(node->indices))
return ffi::Unchanged();
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
auto* writable = const_cast<prim::ShuffleNode*>(node);
if (!vectors.IsUnchanged()) writable->vectors =
std::move(vectors).ValueUnchecked();
if (!indices.IsUnchanged()) writable->indices =
std::move(indices).ValueUnchecked();
@@ -676,7 +669,8 @@ Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const
prim::ShuffleNode* no
}
}
-Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const VarNode* node, bool
allow_inplace) {
+Expected<UnchangedOr<ffi::Any>> ExprMutator::Mutate_(const VarNode* node,
+ InplaceMode inplace_mode)
{
if (node->ty.as<PrimTypeNode>()) return ffi::Unchanged();
if (TVM_FFI_PREDICT_TRUE(var_remap_.empty() && def_region_kind() ==
kTVMFFIDefRegionKindNone)) {
return ffi::Unchanged();
@@ -695,13 +689,12 @@ Expected<UnchangedOr<ffi::Any>>
ExprMutator::Mutate_(const VarNode* node, bool a
if (!node->ty.as<PrimTypeNode>()) {
Expected<UnchangedOr<ffi::Any>> mapped_ty_result =
def_region_kind() == kTVMFFIDefRegionKindSimple
- ? WithDefRegionKind(
- kTVMFFIDefRegionKindNone,
- [&] { return
this->MaybeInplaceMutateIfUniqueExpected(node->ty, allow_inplace); })
- : this->MaybeInplaceMutateIfUniqueExpected(node->ty,
allow_inplace);
+ ? WithDefRegionKind(kTVMFFIDefRegionKindNone,
+ [&] { return this->MutateExpected(node->ty,
inplace_mode); })
+ : this->MutateExpected(node->ty, inplace_mode);
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(UnchangedOr<Type>, mapped_ty,
std::move(mapped_ty_result));
if (!mapped_ty.UnchangedOrSameAs(node->ty)) {
- if (allow_inplace) {
+ if (inplace_mode == InplaceMode::kAllow) {
const_cast<VarNode*>(node)->ty = std::move(mapped_ty).ValueUnchecked();
mapped_value = ffi::Any(node);
} else {
diff --git a/src/ir/prim/expr.cc b/src/ir/prim/expr.cc
index 9bb303bda6..0ca93446cd 100644
--- a/src/ir/prim/expr.cc
+++ b/src/ir/prim/expr.cc
@@ -86,9 +86,9 @@ TVMFFIAny BinaryMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator,
TNode* self = const_cast<TNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, a,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->a));
+ mutator->MutateExpected(self->a,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, b,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->b));
+ mutator->MutateExpected(self->b,
ffi::InplaceMode::kAllow));
if (a.UnchangedOrSameAs(self->a) && b.UnchangedOrSameAs(self->b)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -145,7 +145,7 @@ TVMFFIAny CastMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyVie
CastNode* self = const_cast<CastNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
CastNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
if (mapped_value.UnchangedOrSameAs(self->value)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -180,7 +180,7 @@ TVMFFIAny NotMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyView
NotNode* self = const_cast<NotNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
NotNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_a,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->a));
+ mutator->MutateExpected(self->a,
ffi::InplaceMode::kAllow));
if (mapped_a.UnchangedOrSameAs(self->a)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -225,12 +225,15 @@ TVMFFIAny
SelectMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
SelectNode* self = const_cast<SelectNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
SelectNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_condition,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->condition));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_true_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->true_value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_false_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->false_value));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_condition,
+ mutator->MutateExpected(self->condition, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_true_value,
+ mutator->MutateExpected(self->true_value, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_false_value,
+ mutator->MutateExpected(self->false_value, ffi::InplaceMode::kAllow));
if (mapped_condition.UnchangedOrSameAs(self->condition) &&
mapped_true_value.UnchangedOrSameAs(self->true_value) &&
mapped_false_value.UnchangedOrSameAs(self->false_value)) {
@@ -285,12 +288,13 @@ TVMFFIAny
LetMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
LetNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Var>, mapped_var,
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return
mutator->MaybeInplaceMutateIfUniqueExpected(self->var);
+ return mutator->MutateExpected(self->var,
+
ffi::InplaceMode::kAllow);
}));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (mapped_var.UnchangedOrSameAs(self->var) &&
mapped_value.UnchangedOrSameAs(self->value) &&
mapped_body.UnchangedOrSameAs(self->body)) {
return ffi::Unchanged().CopyToTVMFFIAny();
diff --git a/src/ir/prim/vector_expr.cc b/src/ir/prim/vector_expr.cc
index 9183e31703..acf3f04245 100644
--- a/src/ir/prim/vector_expr.cc
+++ b/src/ir/prim/vector_expr.cc
@@ -89,11 +89,12 @@ TVMFFIAny RampMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyVie
RampNode* self = const_cast<RampNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
RampNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_base,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->base));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_stride,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->stride));
+ mutator->MutateExpected(self->base,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_stride,
+ mutator->MutateExpected(self->stride, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_lanes,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->lanes));
+ mutator->MutateExpected(self->lanes,
ffi::InplaceMode::kAllow));
if (mapped_base.UnchangedOrSameAs(self->base) &&
mapped_stride.UnchangedOrSameAs(self->stride) &&
mapped_lanes.UnchangedOrSameAs(self->lanes)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -136,9 +137,9 @@ TVMFFIAny
BroadcastMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
BroadcastNode* self = const_cast<BroadcastNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BroadcastNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_lanes,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->lanes));
+ mutator->MutateExpected(self->lanes,
ffi::InplaceMode::kAllow));
if (mapped_value.UnchangedOrSameAs(self->value) &&
mapped_lanes.UnchangedOrSameAs(self->lanes)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -179,10 +180,12 @@ TVMFFIAny
ShuffleMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
ShuffleNode* self = const_cast<ShuffleNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ShuffleNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_vectors,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->vectors));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_indices,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->indices));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_vectors,
+ mutator->MutateExpected(self->vectors, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_indices,
+ mutator->MutateExpected(self->indices, ffi::InplaceMode::kAllow));
if (mapped_vectors.UnchangedOrSameAs(self->vectors) &&
mapped_indices.UnchangedOrSameAs(self->indices)) {
return ffi::Unchanged().CopyToTVMFFIAny();
diff --git a/src/ir/type.cc b/src/ir/type.cc
index 9c07cd1100..06e71c3b38 100644
--- a/src/ir/type.cc
+++ b/src/ir/type.cc
@@ -143,7 +143,7 @@ TVMFFIAny
PointerTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
PointerTypeNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<Type>, mapped_element_type,
- mutator->MaybeInplaceMutateIfUniqueExpected(self->element_type));
+ mutator->MutateExpected(self->element_type, ffi::InplaceMode::kAllow));
if (!mapped_element_type.UnchangedOrSameAs(self->element_type)) {
self->element_type = std::move(mapped_element_type).ValueUnchecked();
}
@@ -179,10 +179,12 @@ TVMFFIAny
FuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
FuncTypeNode* self = const_cast<FuncTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
FuncTypeNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Type>>,
mapped_arg_types,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->arg_types));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ret_type,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_type));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Type>>, mapped_arg_types,
+ mutator->MutateExpected(self->arg_types, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<Type>, mapped_ret_type,
+ mutator->MutateExpected(self->ret_type, ffi::InplaceMode::kAllow));
if (!mapped_arg_types.UnchangedOrSameAs(self->arg_types)) {
self->arg_types = std::move(mapped_arg_types).ValueUnchecked();
}
@@ -216,8 +218,9 @@ TVMFFIAny
TupleTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
TupleTypeNode* self = const_cast<TupleTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TupleTypeNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Type>>,
mapped_fields,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->fields));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Type>>, mapped_fields,
+ mutator->MutateExpected(self->fields, ffi::InplaceMode::kAllow));
if (!mapped_fields.UnchangedOrSameAs(self->fields)) {
self->fields = std::move(mapped_fields).ValueUnchecked();
}
diff --git a/src/relax/distributed/type.cc b/src/relax/distributed/type.cc
index c4537393f4..91a423b886 100644
--- a/src/relax/distributed/type.cc
+++ b/src/relax/distributed/type.cc
@@ -66,12 +66,15 @@ TVMFFIAny
DTensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
DTensorTypeNode* self = const_cast<DTensorTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
DTensorTypeNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<DeviceMesh>,
mapped_device_mesh,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->device_mesh));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Placement>,
mapped_placement,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->placement));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<TensorType>,
mapped_tensor_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->tensor_ty));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<DeviceMesh>, mapped_device_mesh,
+ mutator->MutateExpected(self->device_mesh, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<Placement>, mapped_placement,
+ mutator->MutateExpected(self->placement, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<TensorType>, mapped_tensor_ty,
+ mutator->MutateExpected(self->tensor_ty, ffi::InplaceMode::kAllow));
if (!mapped_device_mesh.UnchangedOrSameAs(self->device_mesh)) {
self->device_mesh = std::move(mapped_device_mesh).ValueUnchecked();
}
diff --git a/src/relax/ir/dependent_type.cc b/src/relax/ir/dependent_type.cc
index 2c7177c60a..5cee831d74 100644
--- a/src/relax/ir/dependent_type.cc
+++ b/src/relax/ir/dependent_type.cc
@@ -71,9 +71,9 @@ TVMFFIAny
ShapeTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: ndim (scalar)
ShapeTypeNode* self = const_cast<ShapeTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ShapeTypeNode>(value));
-
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<ffi::Array<PrimExpr>>>,
- mapped_values,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->values));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Optional<ffi::Array<PrimExpr>>>, mapped_values,
+ mutator->MutateExpected(self->values, ffi::InplaceMode::kAllow));
if (!mapped_values.UnchangedOrSameAs(self->values)) {
self->values = std::move(mapped_values).ValueUnchecked();
}
@@ -117,11 +117,12 @@ TVMFFIAny
TensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TensorTypeNode* self = const_cast<TensorTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TensorTypeNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Expr>>,
mapped_shape,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->shape));
+ mutator->MutateExpected(self->shape,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<PrimType>>,
mapped_dtype,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->dtype));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<VDevice>>,
mapped_vdevice,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->vdevice));
+ mutator->MutateExpected(self->dtype,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Optional<VDevice>>, mapped_vdevice,
+ mutator->MutateExpected(self->vdevice, ffi::InplaceMode::kAllow));
if (!mapped_shape.UnchangedOrSameAs(self->shape)) {
self->shape = std::move(mapped_shape).ValueUnchecked();
}
@@ -171,10 +172,10 @@ TVMFFIAny
FuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<ffi::Optional<ffi::Array<Type>>>, mapped_params,
mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->params);
+ return mutator->MutateExpected(self->params, ffi::InplaceMode::kAllow);
}));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ret,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ret));
+ mutator->MutateExpected(self->ret,
ffi::InplaceMode::kAllow));
if (!mapped_params.UnchangedOrSameAs(self->params)) {
self->params = std::move(mapped_params).ValueUnchecked();
}
diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc
index 5b29fded21..3ab6670fd5 100644
--- a/src/relax/ir/expr.cc
+++ b/src/relax/ir/expr.cc
@@ -60,7 +60,7 @@ TVMFFIAny
TypeOnlyExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TNode* self = const_cast<TNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
if (!mapped_ty.UnchangedOrSameAs(self->ty)) {
self->ty = std::move(mapped_ty).ValueUnchecked();
}
@@ -106,14 +106,15 @@ TVMFFIAny IfMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyView
IfNode* self = const_cast<IfNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
IfNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_cond,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->cond));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<SeqExpr>,
mapped_true_branch,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->true_branch));
+ mutator->MutateExpected(self->cond,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<SeqExpr>, mapped_true_branch,
+ mutator->MutateExpected(self->true_branch, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<SeqExpr>, mapped_false_branch,
- mutator->MaybeInplaceMutateIfUniqueExpected(self->false_branch));
+ mutator->MutateExpected(self->false_branch, ffi::InplaceMode::kAllow));
if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty =
std::move(mapped_ty).ValueUnchecked();
if (!mapped_cond.UnchangedOrSameAs(self->cond)) {
self->cond = std::move(mapped_cond).ValueUnchecked();
@@ -156,9 +157,10 @@ TVMFFIAny
ShapeExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ShapeExprNode* self = const_cast<ShapeExprNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ShapeExprNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_values,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->values));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_values,
+ mutator->MutateExpected(self->values, ffi::InplaceMode::kAllow));
if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty =
std::move(mapped_ty).ValueUnchecked();
if (!mapped_values.UnchangedOrSameAs(self->values)) {
self->values = std::move(mapped_values).ValueUnchecked();
@@ -241,8 +243,8 @@ TVMFFIAny
DataflowVarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
mutator->def_region_kind() == kTVMFFIDefRegionKindSimple
? mutator->WithDefRegionKind(
kTVMFFIDefRegionKindNone,
- [&]() { return
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty); })
- : mutator->MaybeInplaceMutateIfUniqueExpected(self->ty);
+ [&]() { return mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow); })
+ : mutator->MutateExpected(self->ty, ffi::InplaceMode::kAllow);
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
std::move(mapped_ty_result));
if (!mapped_ty.UnchangedOrSameAs(self->ty)) {
@@ -294,12 +296,13 @@ TVMFFIAny
SeqExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
SeqExprNode* self = const_cast<SeqExprNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
SeqExprNode>(value));
-
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<BindingBlock>>,
mapped_blocks,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->blocks));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<BindingBlock>>, mapped_blocks,
+ mutator->MutateExpected(self->blocks, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (!mapped_blocks.UnchangedOrSameAs(self->blocks)) {
self->blocks = std::move(mapped_blocks).ValueUnchecked();
}
@@ -355,17 +358,18 @@ TVMFFIAny
FunctionMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: attrs (metadata), is_pure (scalar)
FunctionNode* self = const_cast<FunctionNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
FunctionNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Array<Var>>, mapped_params,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->params);
- }));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Var>>,
mapped_params,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() {
+ return
mutator->MutateExpected(self->params,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ty));
+ mutator->MutateExpected(self->ty,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<SeqExpr>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ret_ty,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_ty));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<Type>, mapped_ret_ty,
+ mutator->MutateExpected(self->ret_ty, ffi::InplaceMode::kAllow));
if (!mapped_params.UnchangedOrSameAs(self->params)) {
self->params = std::move(mapped_params).ValueUnchecked();
}
diff --git a/src/tirx/ir/buffer.cc b/src/tirx/ir/buffer.cc
index aef2bc87a4..ffbea28fab 100644
--- a/src/tirx/ir/buffer.cc
+++ b/src/tirx/ir/buffer.cc
@@ -186,28 +186,31 @@ TVMFFIAny
BufferTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
BufferTypeNode* self = const_cast<BufferTypeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimType>, mapped_dtype,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->dtype));
+ mutator->MutateExpected(self->dtype,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_shape,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->shape));
+ mutator->MutateExpected(self->shape,
ffi::InplaceMode::kAllow));
// Empty strides denote the common compact layout. Broad callbacks do not
see the empty
// container; explicit strides retain normal container descent and callback
behavior.
ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_strides = ffi::Unchanged();
if (!self->strides.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
descended_strides,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->strides));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_strides,
+ mutator->MutateExpected(self->strides, ffi::InplaceMode::kAllow));
mapped_strides = std::move(descended_strides);
}
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_elem_offset,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->elem_offset));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Layout>>,
mapped_layout,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->layout));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_elem_offset,
+ mutator->MutateExpected(self->elem_offset, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Optional<Layout>>, mapped_layout,
+ mutator->MutateExpected(self->layout, ffi::InplaceMode::kAllow));
// allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
// empty container; present addresses retain normal container descent and
callback behavior.
ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_allocated_addr =
ffi::Unchanged();
if (!self->allocated_addr.empty()) {
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_allocated_addr,
- mutator->MaybeInplaceMutateIfUniqueExpected(self->allocated_addr));
+ mutator->MutateExpected(self->allocated_addr,
ffi::InplaceMode::kAllow));
mapped_allocated_addr = std::move(descended_allocated_addr);
}
if (mapped_dtype.UnchangedOrSameAs(self->dtype) &&
mapped_shape.UnchangedOrSameAs(self->shape) &&
diff --git a/src/tirx/ir/function.cc b/src/tirx/ir/function.cc
index 131414a596..e37cd42f5a 100644
--- a/src/tirx/ir/function.cc
+++ b/src/tirx/ir/function.cc
@@ -115,15 +115,16 @@ TVMFFIAny
PrimFuncMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: attrs (metadata), ty (derived by InferType)
PrimFuncNode* self = const_cast<PrimFuncNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
PrimFuncNode>(value));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Var>>,
mapped_params,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() {
+ return
mutator->MutateExpected(self->params,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Array<Var>>, mapped_params,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->params);
- }));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Type>, mapped_ret_type,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_type));
+ ffi::UnchangedOr<Type>, mapped_ret_type,
+ mutator->MutateExpected(self->ret_type, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (!mapped_params.UnchangedOrSameAs(self->params)) {
self->params = std::move(mapped_params).ValueUnchecked();
}
diff --git a/src/tirx/ir/iter_var.cc b/src/tirx/ir/iter_var.cc
index d04d4ef399..1ab656d3ee 100644
--- a/src/tirx/ir/iter_var.cc
+++ b/src/tirx/ir/iter_var.cc
@@ -68,10 +68,11 @@ TVMFFIAny
IterVarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
IterVarNode* self = const_cast<IterVarNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
IterVarNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Range>, mapped_dom,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->dom));
+ mutator->MutateExpected(self->dom,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimVar>, mapped_var,
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return
mutator->MaybeInplaceMutateIfUniqueExpected(self->var);
+ return mutator->MutateExpected(self->var,
+
ffi::InplaceMode::kAllow);
}));
if (mapped_dom.UnchangedOrSameAs(self->dom) &&
mapped_var.UnchangedOrSameAs(self->var)) {
return ffi::Unchanged().CopyToTVMFFIAny();
diff --git a/src/tirx/ir/layout/tile_core.cc b/src/tirx/ir/layout/tile_core.cc
index ab397a3d21..d555b9d0b6 100644
--- a/src/tirx/ir/layout/tile_core.cc
+++ b/src/tirx/ir/layout/tile_core.cc
@@ -66,12 +66,14 @@ TVMFFIAny IterMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noe
TVMFFIAny IterMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
IterNode* self = const_cast<IterNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
IterNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_extent,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->extent));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_stride,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->stride));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_extent,
+ mutator->MutateExpected(self->extent, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_stride,
+ mutator->MutateExpected(self->stride, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Axis>, mapped_axis,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->axis));
+ mutator->MutateExpected(self->axis,
ffi::InplaceMode::kAllow));
if (mapped_extent.UnchangedOrSameAs(self->extent) &&
mapped_stride.UnchangedOrSameAs(self->stride) &&
mapped_axis.UnchangedOrSameAs(self->axis)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -118,12 +120,14 @@ TVMFFIAny
TileLayoutMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
TileLayoutNode* self = const_cast<TileLayoutNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TileLayoutNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Iter>>,
mapped_shard,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->shard));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Iter>>,
mapped_replica,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->replica));
+ mutator->MutateExpected(self->shard,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Iter>>, mapped_replica,
+ mutator->MutateExpected(self->replica, ffi::InplaceMode::kAllow));
using OffsetMap = ffi::Map<Axis, PrimExpr>;
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<OffsetMap>, mapped_offset,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->offset));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<OffsetMap>, mapped_offset,
+ mutator->MutateExpected(self->offset, ffi::InplaceMode::kAllow));
if (mapped_shard.UnchangedOrSameAs(self->shard) &&
mapped_replica.UnchangedOrSameAs(self->replica) &&
mapped_offset.UnchangedOrSameAs(self->offset)) {
diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc
index 0d9ab48e72..0ccf045c04 100644
--- a/src/tirx/ir/stmt.cc
+++ b/src/tirx/ir/stmt.cc
@@ -93,10 +93,11 @@ TVMFFIAny BindMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyVie
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BindNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Var>, mapped_var,
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return
mutator->MaybeInplaceMutateIfUniqueExpected(self->var);
+ return mutator->MutateExpected(self->var,
+
ffi::InplaceMode::kAllow);
}));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
if (mapped_var.UnchangedOrSameAs(self->var) &&
mapped_value.UnchangedOrSameAs(self->value)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -142,11 +143,11 @@ TVMFFIAny
AttrStmtMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
AttrStmtNode* self = const_cast<AttrStmtNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
AttrStmtNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Any>, mapped_node,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->node));
+ mutator->MutateExpected(self->node,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (mapped_node.UnchangedOrSameAs(self->node) &&
mapped_value.UnchangedOrSameAs(self->value) &&
mapped_body.UnchangedOrSameAs(self->body)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -184,8 +185,9 @@ TVMFFIAny
AssertStmtMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: error_kind and message_parts, which are constant assertion
metadata.
AssertStmtNode* self = const_cast<AssertStmtNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
AssertStmtNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_condition,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->condition));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_condition,
+ mutator->MutateExpected(self->condition, ffi::InplaceMode::kAllow));
if (mapped_condition.UnchangedOrSameAs(self->condition)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -248,22 +250,23 @@ TVMFFIAny
ForMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView
// skips: kind and constant annotations; unlike SBlock annotations, these
carry no expressions.
ForNode* self = const_cast<ForNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ForNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<PrimVar>, mapped_loop_var,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->loop_var);
- }));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimVar>, mapped_loop_var,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
+ return
mutator->MutateExpected(self->loop_var,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_min,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->min));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_extent,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->extent));
+ mutator->MutateExpected(self->min,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_extent,
+ mutator->MutateExpected(self->extent, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<ffi::Optional<IterVar>>, mapped_thread_binding,
- mutator->MaybeInplaceMutateIfUniqueExpected(self->thread_binding));
+ mutator->MutateExpected(self->thread_binding, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<PrimExpr>>,
mapped_step,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->step));
+ mutator->MutateExpected(self->step,
ffi::InplaceMode::kAllow));
if (mapped_loop_var.UnchangedOrSameAs(self->loop_var) &&
mapped_min.UnchangedOrSameAs(self->min) &&
mapped_extent.UnchangedOrSameAs(self->extent) &&
mapped_body.UnchangedOrSameAs(self->body) &&
@@ -310,10 +313,11 @@ TVMFFIAny WhileMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) no
TVMFFIAny WhileMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
WhileNode* self = const_cast<WhileNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
WhileNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_condition,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->condition));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_condition,
+ mutator->MutateExpected(self->condition, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (mapped_condition.UnchangedOrSameAs(self->condition) &&
mapped_body.UnchangedOrSameAs(self->body)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -349,7 +353,7 @@ TVMFFIAny
ReturnMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ReturnNode* self = const_cast<ReturnNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ReturnNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
if (mapped_value.UnchangedOrSameAs(self->value)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -412,13 +416,13 @@ TVMFFIAny
DeclBufferMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
DeclBufferNode* self = const_cast<DeclBufferNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
DeclBufferNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<BufferVar>, mapped_buffer,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->buffer);
- }));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferVar>, mapped_buffer,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
+ return
mutator->MutateExpected(self->buffer,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_data,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->data));
+ mutator->MutateExpected(self->data,
ffi::InplaceMode::kAllow));
if (mapped_buffer.UnchangedOrSameAs(self->buffer) &&
mapped_data.UnchangedOrSameAs(self->data)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -457,11 +461,11 @@ TVMFFIAny
AllocBufferMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: constant annotations; unlike SBlock annotations, these carry no
expressions.
AllocBufferNode* self = const_cast<AllocBufferNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
AllocBufferNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<BufferVar>, mapped_buffer,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->buffer);
- }));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferVar>, mapped_buffer,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
+ return
mutator->MutateExpected(self->buffer,
+
ffi::InplaceMode::kAllow);
+ }));
if (mapped_buffer.UnchangedOrSameAs(self->buffer)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -583,7 +587,7 @@ TVMFFIAny
MaybeInplaceMutateSeqStmtRaw(ffi::StructuralMutatorObj* mutator,
// Pass 1: rebuild each element in place and classify the final slot
contents.
for (size_t i = 0; i < size; ++i) {
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped,
-
mutator->MaybeInplaceMutateIfUniqueExpected(slots[i]));
+ mutator->MutateExpected(slots[i],
ffi::InplaceMode::kAllow));
if (!mapped.UnchangedOrSameAs(slots[i].cast<Stmt>())) {
slots[i] = ffi::Any(std::move(mapped).ValueUnchecked());
}
@@ -686,12 +690,15 @@ TVMFFIAny
IfThenElseMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
IfThenElseNode* self = const_cast<IfThenElseNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
IfThenElseNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_condition,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->condition));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_then_case,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->then_case));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Stmt>>,
mapped_else_case,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->else_case));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_condition,
+ mutator->MutateExpected(self->condition, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<Stmt>, mapped_then_case,
+ mutator->MutateExpected(self->then_case, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Optional<Stmt>>, mapped_else_case,
+ mutator->MutateExpected(self->else_case, ffi::InplaceMode::kAllow));
if (mapped_condition.UnchangedOrSameAs(self->condition) &&
mapped_then_case.UnchangedOrSameAs(self->then_case) &&
mapped_else_case.UnchangedOrSameAs(self->else_case)) {
@@ -731,7 +738,7 @@ TVMFFIAny
EvaluateMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
EvaluateNode* self = const_cast<EvaluateNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
EvaluateNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Expr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
if (mapped_value.UnchangedOrSameAs(self->value)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -773,12 +780,14 @@ TVMFFIAny
BufferStoreMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
BufferStoreNode* self = const_cast<BufferStoreNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferStoreNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferVar>, mapped_buffer,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->buffer));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<BufferVar>, mapped_buffer,
+ mutator->MutateExpected(self->buffer, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_indices,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->indices));
+ mutator->MutateExpected(self->value,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_indices,
+ mutator->MutateExpected(self->indices, ffi::InplaceMode::kAllow));
if (mapped_buffer.UnchangedOrSameAs(self->buffer) &&
mapped_value.UnchangedOrSameAs(self->value) &&
mapped_indices.UnchangedOrSameAs(self->indices)) {
@@ -881,10 +890,12 @@ TVMFFIAny
BufferRegionMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
BufferRegionNode* self = const_cast<BufferRegionNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferRegionNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferVar>, mapped_buffer,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->buffer));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<Range>>,
mapped_region,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->region));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<BufferVar>, mapped_buffer,
+ mutator->MutateExpected(self->buffer, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<Range>>, mapped_region,
+ mutator->MutateExpected(self->region, ffi::InplaceMode::kAllow));
if (mapped_buffer.UnchangedOrSameAs(self->buffer) &&
mapped_region.UnchangedOrSameAs(self->region)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -929,13 +940,14 @@ TVMFFIAny
MatchBufferRegionMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator
MatchBufferRegionNode* self = const_cast<MatchBufferRegionNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
MatchBufferRegionNode>(
value));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferVar>, mapped_buffer,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
+ return
mutator->MutateExpected(self->buffer,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<BufferVar>, mapped_buffer,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return mutator->MaybeInplaceMutateIfUniqueExpected(self->buffer);
- }));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<BufferRegion>,
mapped_source,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->source));
+ ffi::UnchangedOr<BufferRegion>, mapped_source,
+ mutator->MutateExpected(self->source, ffi::InplaceMode::kAllow));
if (mapped_buffer.UnchangedOrSameAs(self->buffer) &&
mapped_source.UnchangedOrSameAs(self->source)) {
return ffi::Unchanged().CopyToTVMFFIAny();
@@ -1013,27 +1025,30 @@ TVMFFIAny
SBlockMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
// skips: name_hint
SBlockNode* self = const_cast<SBlockNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
SBlockNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<IterVar>>,
mapped_iter_vars,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->iter_vars));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<IterVar>>, mapped_iter_vars,
+ mutator->MutateExpected(self->iter_vars, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<BufferRegion>>,
mapped_reads,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->reads));
-
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<BufferRegion>>,
mapped_writes,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->writes));
+ mutator->MutateExpected(self->reads,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Array<BufferVar>>, mapped_alloc_buffers,
- mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
- return
mutator->MaybeInplaceMutateIfUniqueExpected(self->alloc_buffers);
- }));
+ ffi::UnchangedOr<ffi::Array<BufferRegion>>, mapped_writes,
+ mutator->MutateExpected(self->writes, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<BufferVar>>,
mapped_alloc_buffers,
+
mutator->WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() {
+ return
mutator->MutateExpected(self->alloc_buffers,
+
ffi::InplaceMode::kAllow);
+ }));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
ffi::UnchangedOr<ffi::Array<MatchBufferRegion>>, mapped_match_buffers,
- mutator->MaybeInplaceMutateIfUniqueExpected(self->match_buffers));
+ mutator->MutateExpected(self->match_buffers, ffi::InplaceMode::kAllow));
using AnnotationMap = ffi::Map<ffi::String, ffi::Any>;
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<AnnotationMap>,
mapped_annotations,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->annotations));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<AnnotationMap>, mapped_annotations,
+ mutator->MutateExpected(self->annotations, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Stmt>>,
mapped_init,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->init));
+ mutator->MutateExpected(self->init,
ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<Stmt>, mapped_body,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->body));
+ mutator->MutateExpected(self->body,
ffi::InplaceMode::kAllow));
if (mapped_iter_vars.UnchangedOrSameAs(self->iter_vars) &&
mapped_reads.UnchangedOrSameAs(self->reads) &&
mapped_writes.UnchangedOrSameAs(self->writes) &&
@@ -1086,7 +1101,7 @@ TVMFFIAny
ScopeIdDefStmtMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ScopeIdDefStmtNode* self = const_cast<ScopeIdDefStmtNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ScopeIdDefStmtNode>(value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ScopeIdDef>, mapped_def,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->def));
+ mutator->MutateExpected(self->def,
ffi::InplaceMode::kAllow));
if (mapped_def.UnchangedOrSameAs(self->def)) {
return ffi::Unchanged().CopyToTVMFFIAny();
}
@@ -1128,12 +1143,14 @@ TVMFFIAny
SBlockRealizeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
ffi::AnyView value) noexcept {
SBlockRealizeNode* self = const_cast<SBlockRealizeNode*>(
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
SBlockRealizeNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_iter_values,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->iter_values));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_predicate,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->predicate));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, mapped_iter_values,
+ mutator->MutateExpected(self->iter_values, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_predicate,
+ mutator->MutateExpected(self->predicate, ffi::InplaceMode::kAllow));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<SBlock>, mapped_block,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->block));
+ mutator->MutateExpected(self->block,
ffi::InplaceMode::kAllow));
if (mapped_iter_values.UnchangedOrSameAs(self->iter_values) &&
mapped_predicate.UnchangedOrSameAs(self->predicate) &&
mapped_block.UnchangedOrSameAs(self->block)) {
diff --git a/src/tirx/ir/tirx_stmt.cc b/src/tirx/ir/tirx_stmt.cc
index 57fa8064e1..ef92f11639 100644
--- a/src/tirx/ir/tirx_stmt.cc
+++ b/src/tirx/ir/tirx_stmt.cc
@@ -96,14 +96,16 @@ TVMFFIAny
TilePrimitiveCallMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
TilePrimitiveCallNode>(
value));
TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<ffi::Any>>,
mapped_args,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->args));
+ mutator->MutateExpected(self->args,
ffi::InplaceMode::kAllow));
using WorkspaceMap = ffi::Map<ffi::String, BufferVar>;
using ConfigMap = ffi::Map<ffi::String, ffi::Any>;
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<WorkspaceMap>,
mapped_workspace,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->workspace));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ConfigMap>, mapped_config,
-
mutator->MaybeInplaceMutateIfUniqueExpected(self->config));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<WorkspaceMap>, mapped_workspace,
+ mutator->MutateExpected(self->workspace, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ConfigMap>, mapped_config,
+ mutator->MutateExpected(self->config, ffi::InplaceMode::kAllow));
if (mapped_args.UnchangedOrSameAs(self->args) &&
mapped_workspace.UnchangedOrSameAs(self->workspace) &&
diff --git a/tests/cpp/expr_functor_test.cc b/tests/cpp/expr_functor_test.cc
index 06cb2c962e..b62cc9dd93 100644
--- a/tests/cpp/expr_functor_test.cc
+++ b/tests/cpp/expr_functor_test.cc
@@ -72,7 +72,8 @@ TEST(ExprVisitor, StructuralFallback) {
class Rewrite : public ExprMutator {
public:
using ExprMutator::Mutate_;
- Expected<UnchangedOr<ffi::Any>> Mutate_(const IntImmNode* node, bool
allow_inplace) override {
+ Expected<UnchangedOr<ffi::Any>> Mutate_(const IntImmNode* node,
+ InplaceMode inplace_mode) override {
return ffi::Any(IntImm::Int32(node->value + 1));
}
};