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-ffi.git
The following commit(s) were added to refs/heads/main by this push:
new 6b1b19d9 [REFACTOR][EXTRA] Add StructuralMutate, a callback-driven
mutation primitive (#751)
6b1b19d9 is described below
commit 6b1b19d9f4d9f10acc87af63c7dd0cf1ab75e9e9
Author: Tianqi Chen <[email protected]>
AuthorDate: Sun Sep 6 08:06:19 2026 -0400
[REFACTOR][EXTRA] Add StructuralMutate, a callback-driven mutation
primitive (#751)
Structural mutation needs a callback-owned form where a matched callback
decides whether and how to recurse, and its replacement remains final.
This adds StructuralMutate and StructuralMutateExpected with
declaration-ordered typed and runtime dispatch. Both engines preserve
distinct ordinary and maybe-in-place fallbacks, use the exact
Parent::MutatorObjType callback view, and let matched variable callbacks
own remap consistency.
Map and mutate entry points now take the root as Any by value, so
in-place reuse requires explicit ownership transfer with std::move or
Python _move(). Existing StructuralMap callers therefore lose
opportunistic in-place reuse until they move an owned root.
StructuralMutate callbacks may optionally take allow_inplace as a third
argument; two-argument callbacks retain the copy-on-write default.
Object metadata can now promise that subclasses retain one structural
equality/hash kind. Registration rejects violations, and mutation skips
runtime remap metadata lookup for final or kind-fixed non-remappable
callback types.
The runtime and Python interfaces expose callback dispatch and
default_mutate but intentionally omit a direct Python maybe-in-place
mutator method. Shared-subvalue, moved-root identity, nested-hook
ownership, callback-path, and paired visit/mutate recursion tests keep
the interfaces aligned.
Validation: all 493 C++ tests and all 2,495 Python tests pass; the
generated FFI stub and full pre-commit suite, including ty,
clang-format, and CMake lint/format, are clean.
---
docs/concepts/structural_eq_hash.rst | 21 +--
include/tvm/ffi/extra/structural_mutate.h | 224 +++++++++++++++++++++--
include/tvm/ffi/object.h | 13 ++
python/tvm_ffi/__init__.py | 2 +
python/tvm_ffi/_ffi_api.py | 6 +-
python/tvm_ffi/structural.py | 99 ++++++++--
src/ffi/extra/structural_mutate.cc | 106 ++++++++++-
tests/cpp/extra/test_structural_mutate.cc | 294 +++++++++++++++++++++++++++++-
tests/cpp/extra/test_structural_visit.cc | 21 +++
tests/python/test_structural.py | 122 ++++++++++++-
10 files changed, 853 insertions(+), 55 deletions(-)
diff --git a/docs/concepts/structural_eq_hash.rst
b/docs/concepts/structural_eq_hash.rst
index b7182560..51335645 100644
--- a/docs/concepts/structural_eq_hash.rst
+++ b/docs/concepts/structural_eq_hash.rst
@@ -1068,8 +1068,6 @@ A :class:`~tvm_ffi.StructuralMutator` adds ownership and
replacement semantics.
Its main operations are:
- ``mutator.mutate(value)`` maps without intentionally modifying ``value``.
-- ``mutator.maybe_inplace_mutate(value)`` permits a type-specific
implementation
- to reuse a safely mutable value and otherwise falls back to ``mutate``.
- ``mutator.var_remap_get(var)`` and ``mutator.var_remap_set(var, mapped)``
access the current identity-substitution environment.
- ``def_region_kind`` and ``with_def_region_kind`` have the same role as on the
@@ -1081,9 +1079,11 @@ recursively maps each structural field, and installs
mapped fields in that copy.
If no field changes, it returns the original object instead. A nested change
therefore copies only the objects along the changed path; unchanged children
remain shared.
-``maybe_inplace_mutate`` is an explicit optimization path. A type-specific
-``__s_maybe_inplace_mutate__`` hook owns the safety policy and may reuse its
-input. Without that hook, the default implementation calls ``mutate``.
+A type-specific ``__s_maybe_inplace_mutate__`` hook is an internal optimization
+path. The structural-map engine invokes it only for a uniquely owned value and
+otherwise uses ``__s_mutate__``. Python does not expose this dispatch as a
+direct mutator method; move a root with ``root._move()`` to transfer ownership
+to :func:`~tvm_ffi.structural_map`.
.. note::
@@ -1235,12 +1235,11 @@ structural child, and returns an interrupt if one
occurs:
A custom ``__s_mutate__`` hook similarly receives the active mutator. It
should
recursively call ``mutator.mutate`` and return a new value only when needed.
An optional ``__s_maybe_inplace_mutate__`` hook may implement an in-place
-optimization. Callers use ``mutate`` for shared objects and call
-``maybe_inplace_mutate`` only when the input is safe to mutate, so the optional
-hook may rely on that ownership guarantee. A type defining it must also define
-``__s_mutate__``. If the optional hook is absent, ``maybe_inplace_mutate``
uses
-the default non-in-place mutation; generic reflected fields are never mutated
-in place automatically.
+optimization. The structural-map engine dispatches it only when the input is
+safe to mutate, so the optional hook may rely on that ownership guarantee. A
+type defining it must also define ``__s_mutate__``. If the optional hook is
+absent, the engine uses the default non-in-place mutation; generic reflected
+fields are never mutated in place automatically.
When an object marked ``structural_eq="var"`` or ``structural_eq="dag"``
registers
either ``__s_mutate__`` or ``__s_maybe_inplace_mutate__`` hooks, it should:
diff --git a/include/tvm/ffi/extra/structural_mutate.h
b/include/tvm/ffi/extra/structural_mutate.h
index a1b9c8ea..f53b678d 100644
--- a/include/tvm/ffi/extra/structural_mutate.h
+++ b/include/tvm/ffi/extra/structural_mutate.h
@@ -185,6 +185,9 @@ class StructuralMutatorObj : public Object {
*
* \param value The borrowed value to mutate.
* \return The mutated owning value, or an Error if mutation failed.
+ *
+ * \note Call only from a ``__s_maybe_inplace_mutate__`` hook, which is
dispatched
+ * only for a value whose entire path from the root is uniquely owned.
*/
TVM_FFI_INLINE Expected<Any> MaybeInplaceMutateExpected(AnyView value)
noexcept {
return details::ExpectedUnsafe::MoveFromTVMFFIAny<Any>(
@@ -196,6 +199,11 @@ class StructuralMutatorObj : public Object {
*
* \param value The borrowed value to mutate.
* \return The mutated owning value, or an Error if mutation failed.
+ *
+ * \note The caller must already know the entire path from the root is
uniquely
+ * owned, either through an owning moved-in root or while handling a
+ * ``__s_maybe_inplace_mutate__`` hook. This method checks only \p
value
+ * itself, not its ancestors.
*/
TVM_FFI_INLINE Expected<Any> MaybeInplaceMutateIfUniqueExpected(AnyView
value) noexcept {
const Object* obj = value.as<Object>();
@@ -943,13 +951,13 @@ class StructuralMapEngine : public Parent {
if (!matched.has_value()) return false;
}
- // A final statically non-remappable type discards the remap path at
optimization time.
- // Every other case uses runtime metadata: nullable refs may match None,
non-final subclasses
- // may redeclare the kind, and metadata may be absent.
+ // A statically non-remappable type whose subclasses cannot change kind
discards the remap
+ // path at optimization time. Every other case uses runtime metadata:
nullable refs may match
+ // None, non-final subclasses may redeclare the kind, and metadata may be
absent.
const bool remappable = [&]() {
if constexpr (std::is_base_of_v<ObjectRef, TSub>) {
using TNode = typename TSub::ContainerType;
- if constexpr (TNode::_type_final &&
+ if constexpr ((TNode::_type_final ||
TNode::_type_s_eq_hash_subclass_kind_fixed) &&
TNode::_type_s_eq_hash_kind != kTVMFFISEqHashKindFreeVar
&&
TNode::_type_s_eq_hash_kind !=
kTVMFFISEqHashKindDAGNode) {
return false;
@@ -958,10 +966,11 @@ class StructuralMapEngine : public Parent {
if constexpr (std::is_pointer_v<TSub> &&
std::is_base_of_v<Object,
std::remove_cv_t<std::remove_pointer_t<TSub>>>) {
using TNode = std::remove_cv_t<std::remove_pointer_t<TSub>>;
- constexpr bool kFinalNonRemappable =
- TNode::_type_final && TNode::_type_s_eq_hash_kind !=
kTVMFFISEqHashKindFreeVar &&
- TNode::_type_s_eq_hash_kind != kTVMFFISEqHashKindDAGNode;
- if constexpr (kFinalNonRemappable) return false;
+ if constexpr ((TNode::_type_final ||
TNode::_type_s_eq_hash_subclass_kind_fixed) &&
+ TNode::_type_s_eq_hash_kind != kTVMFFISEqHashKindFreeVar
&&
+ TNode::_type_s_eq_hash_kind !=
kTVMFFISEqHashKindDAGNode) {
+ return false;
+ }
}
return this->IsRemappableIdentity(value.type_index());
}();
@@ -1325,6 +1334,138 @@ class StructuralMapDynEngine : public Parent {
Array<Tuple<int32_t, Function>> callbacks_with_def_region_kind_;
};
+/*!
+ * \brief Engine of the callback-dispatched \ref tvm::ffi::StructuralMutate.
+ *
+ * A matched callback owns mutation of its value, so the engine returns the
+ * callback result without descending into it. An unmatched value keeps the
+ * Parent's default mutation. ``Parent::MutatorObjType`` pins the exact
+ * callback-facing mutator view across layer composition.
+ *
+ * \tparam Parent Mutator layer extended by the engine.
+ * \tparam Callbacks Callable types whose first parameter selects the value
type.
+ */
+template <typename Parent, typename... Callbacks>
+class StructuralMutateEngine : public Parent {
+ public:
+ static_assert(std::is_base_of_v<StructuralMapEngineBase, Parent>,
+ "StructuralMutate Parent must derive from
StructuralMapEngineBase");
+
+ /*! \brief Construct a mutate engine over callbacks tested in declaration
order. */
+ explicit StructuralMutateEngine(Callbacks... callbacks)
+ : Parent(VTable()), callbacks_(std::move(callbacks)...) {}
+
+ private:
+ /*! \brief Return this engine's immutable callback-aware mutator vtable. */
+ static const StructuralMutatorVTable* VTable() {
+ static const StructuralMutatorVTable vtable{
+ &StructuralMutateEngine::DispatchMutate,
+ &StructuralMutateEngine::DispatchMaybeInplaceMutate,
+ &StructuralMutateEngine::DispatchVarRemapGet,
+ &StructuralMutateEngine::DispatchVarRemapSet,
+ };
+ return &vtable;
+ }
+
+ /*! \brief Dispatch ordinary mutation from the erased mutator pointer. */
+ static TVMFFIAny DispatchMutate(StructuralMutatorObj* mutator, AnyView
value) noexcept {
+ return static_cast<StructuralMutateEngine*>(mutator)->MutateImplRaw(value);
+ }
+
+ /*! \brief Dispatch maybe-in-place mutation from the erased mutator pointer.
*/
+ static TVMFFIAny DispatchMaybeInplaceMutate(StructuralMutatorObj* mutator,
+ AnyView value) noexcept {
+ return
static_cast<StructuralMutateEngine*>(mutator)->MaybeInplaceMutateImplRaw(value);
+ }
+
+ /*! \brief Mutate one value, handing a matched callback ownership of
descent. */
+ TVMFFIAny MutateImplRaw(AnyView value) noexcept {
+ if (std::optional<Expected<Any>> matched = DispatchCallbacks(value,
false)) {
+ Expected<Any> result = *std::move(matched);
+ if (TVM_FFI_PREDICT_FALSE(result.is_err())) {
+ // Keep callback-boundary context in addition to the default-descent
+ // context: a callback may return a rebuilt value, so the two nodes
can differ.
+ Parent::UpdateVisitErrorContext(result, value);
+ }
+ return details::ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
+ }
+ return Parent::DefaultMutateRaw(value);
+ }
+
+ /*! \brief Maybe mutate one value in place, with callback-owned descent. */
+ TVMFFIAny MaybeInplaceMutateImplRaw(AnyView value) noexcept {
+ if (std::optional<Expected<Any>> matched = DispatchCallbacks(value, true))
{
+ Expected<Any> result = *std::move(matched);
+ if (TVM_FFI_PREDICT_FALSE(result.is_err())) {
+ // Keep callback-boundary context in addition to the default-descent
+ // context: a callback may return a rebuilt value, so the two nodes
can differ.
+ Parent::UpdateVisitErrorContext(result, value);
+ }
+ return details::ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
+ }
+ return Parent::DefaultMaybeInplaceMutateRaw(value);
+ }
+
+ /*! \brief Try one typed callback and preserve Error as an expected result.
*/
+ template <typename Callback>
+ TVM_FFI_INLINE std::optional<Expected<Any>> TryLink(Callback& callback,
AnyView value,
+ bool allow_inplace)
noexcept {
+ using FuncInfo = details::FunctionInfo<std::decay_t<Callback>>;
+ static_assert(FuncInfo::num_args == 2 || FuncInfo::num_args == 3,
+ "StructuralMutate callback must take (value, mutator) or "
+ "(value, mutator, allow_inplace)");
+ using FirstArg = std::tuple_element_t<0, typename FuncInfo::ArgType>;
+ using TSub = std::remove_cv_t<std::remove_reference_t<FirstArg>>;
+ using SecondArg = std::decay_t<std::tuple_element_t<1, typename
FuncInfo::ArgType>>;
+ using Second = std::remove_pointer_t<SecondArg>;
+ static_assert(std::is_same_v<Second, typename Parent::MutatorObjType>,
+ "second StructuralMutate callback argument must be exactly "
+ "Parent::MutatorObjType*");
+ if constexpr (FuncInfo::num_args == 3) {
+ using ThirdArg = std::decay_t<std::tuple_element_t<2, typename
FuncInfo::ArgType>>;
+ static_assert(std::is_same_v<ThirdArg, bool>,
+ "third StructuralMutate callback argument must be bool");
+ }
+ auto* mutator = static_cast<typename Parent::MutatorObjType*>(this);
+ auto invoke = [&](auto&& matched) -> Expected<Any> {
+ try {
+ if constexpr (FuncInfo::num_args == 3) {
+ return callback(std::forward<decltype(matched)>(matched), mutator,
allow_inplace);
+ } else {
+ return callback(std::forward<decltype(matched)>(matched), mutator);
+ }
+ } catch (Error& err) {
+ return Unexpected(std::move(err));
+ }
+ };
+ if constexpr (std::is_same_v<TSub, AnyView>) {
+ return invoke(value);
+ } else if constexpr (std::is_same_v<TSub, Any>) {
+ return invoke(Any(value));
+ } else if (auto matched = value.template as<TSub>()) {
+ return invoke(*std::move(matched));
+ }
+ return std::nullopt;
+ }
+
+ /*! \brief Fold callbacks in declaration order, stopping at the first match.
*/
+ template <size_t... Is>
+ TVM_FFI_INLINE std::optional<Expected<Any>> TryLinks(AnyView value, bool
allow_inplace,
+
std::index_sequence<Is...>) noexcept {
+ std::optional<Expected<Any>> result;
+ (... || (result = TryLink(std::get<Is>(callbacks_), value,
allow_inplace)).has_value());
+ return result;
+ }
+
+ /*! \brief Run the callback chain, or return empty when no callback matched.
*/
+ std::optional<Expected<Any>> DispatchCallbacks(AnyView value, bool
allow_inplace) noexcept {
+ return TryLinks(value, allow_inplace,
std::index_sequence_for<Callbacks...>{});
+ }
+
+ /*! \brief Typed callbacks tested in declaration order, first match wins. */
+ std::tuple<Callbacks...> callbacks_;
+};
+
/*!
* \brief Map a structured value graph and invoke typed replacement callbacks.
*
@@ -1368,7 +1509,7 @@ class StructuralMapDynEngine : public Parent {
*
* \tparam order Whether callbacks run before or after recursively mapping
children.
* \tparam Callbacks Callback types whose first parameters select matching
values.
- * \param root The borrowed root value to map.
+ * \param root The owning root value to map.
* \param callbacks Callbacks tested in declaration order. Each accepts
``(value)`` or
* ``(value, def_region_kind)`` and returns a bare Any-convertible
replacement,
* ``Expected<U>`` where ``U`` is Any-convertible, or an error value.
@@ -1376,9 +1517,14 @@ class StructuralMapDynEngine : public Parent {
*
* \note Returning ``Expected<U>`` expresses errors as values; throwing
``Error`` is also
* supported and is converted to the error state.
+ * \note Pass an owned root with ``std::move(root)`` to permit root reuse. In a
+ * ``__s_maybe_inplace_mutate__`` hook, a nested owned field follows the
idiom
+ * ``self->field = StructuralMap(std::move(self->field), callback)``.
*/
template <WalkOrder order, typename... Callbacks>
-Expected<Any> StructuralMapExpected(AnyView root, Callbacks&&... callbacks)
noexcept {
+// The owning parameter makes caller ownership visible to the uniqueness check.
+Expected<Any> StructuralMapExpected(
+ Any root, Callbacks&&... callbacks) noexcept { //
NOLINT(performance-unnecessary-value-param)
static_assert(sizeof...(Callbacks) != 0, "StructuralMap requires at least
one callback");
using Mutator = StructuralMapEngine<StructuralMapEngineBase, order,
std::decay_t<Callbacks>...>;
StructuralMutator
mutator(make_object<Mutator>(std::forward<Callbacks>(callbacks)...));
@@ -1393,7 +1539,7 @@ Expected<Any> StructuralMapExpected(AnyView root,
Callbacks&&... callbacks) noex
*
* \tparam order Whether callbacks run before or after recursively mapping
children.
* \tparam Callbacks Callback types whose first parameters select matching
values.
- * \param root The borrowed root value to map.
+ * \param root The owning root value to map.
* \param callbacks Callbacks tested in declaration order. Each accepts
``(value)`` or
* ``(value, def_region_kind)`` and returns a bare Any-convertible
replacement,
* ``Expected<U>`` where ``U`` is Any-convertible, or an error value.
@@ -1402,10 +1548,62 @@ Expected<Any> StructuralMapExpected(AnyView root,
Callbacks&&... callbacks) noex
*
* \note Returning ``Expected<U>`` expresses errors as values; throwing
``Error`` is also
* supported and is rethrown by this interface.
+ * \note Pass an owned root with ``std::move(root)`` to permit root reuse.
*/
template <WalkOrder order, typename... Callbacks>
-Any StructuralMap(AnyView root, Callbacks&&... callbacks) {
- return StructuralMapExpected<order>(root,
std::forward<Callbacks>(callbacks)...).value();
+// The owning parameter makes caller ownership visible to the uniqueness check.
+Any StructuralMap(Any root,
+ Callbacks&&... callbacks) { //
NOLINT(performance-unnecessary-value-param)
+ return StructuralMapExpected<order>(std::move(root),
std::forward<Callbacks>(callbacks)...)
+ .value();
+}
+
+/*!
+ * \brief Mutate a structured value with callbacks that own recursion.
+ *
+ * A callback takes one of two signatures:
+ *
+ * - ``Expected<Any>(const T& value, StructuralMutatorObj* mutator)``
+ * - ``Expected<Any>(const T& value, StructuralMutatorObj* mutator, bool
allow_inplace)``
+ *
+ * The returned ``Any`` is the replacement for ``value``; an ``Error`` fails
the
+ * mutation. The first argument selects by FFI type; callbacks are tried in
+ * declaration order and the first match owns mutation -- it drives its own
+ * recursion through the mutator and sets any variable remapping. An unmatched
+ * value takes registered or reflected default mutation.
+ *
+ * \param root The owning root value to mutate.
+ * \param callbacks Callbacks tested in declaration order.
+ * \return The mutated owning value, or an Error if mutation or a callback
fails.
+ *
+ * \note A two-argument callback descends with ``MutateExpected`` and remains
copy-on-write.
+ * A three-argument callback receives ``allow_inplace=true`` only when
its value is on a
+ * uniquely owned path and may then explicitly use the maybe-in-place
mutator operation.
+ * \note Pass an owned root with ``std::move(root)`` to permit root reuse. In a
+ * ``__s_maybe_inplace_mutate__`` hook, the corresponding nested idiom is
+ * ``self->field = StructuralMap(std::move(self->field), callback)``;
const-correctness
+ * rejects that ownership transfer outside a mutable maybe-in-place hook.
+ */
+template <typename... Callbacks>
+// The owning parameter makes caller ownership visible to the uniqueness check.
+Expected<Any> StructuralMutateExpected(
+ Any root, Callbacks&&... callbacks) noexcept { //
NOLINT(performance-unnecessary-value-param)
+ static_assert(sizeof...(Callbacks) != 0, "StructuralMutate requires at least
one callback");
+ using Mutator = StructuralMutateEngine<StructuralMapEngineBase,
std::decay_t<Callbacks>...>;
+ StructuralMutator
mutator(make_object<Mutator>(std::forward<Callbacks>(callbacks)...));
+ return mutator->MaybeInplaceMutateIfUniqueExpected(root);
+}
+
+/*!
+ * \brief Throwing form of \ref tvm::ffi::StructuralMutateExpected.
+ *
+ * \note Pass an owned root with ``std::move(root)`` to permit root reuse.
+ */
+template <typename... Callbacks>
+// The owning parameter makes caller ownership visible to the uniqueness check.
+Any StructuralMutate(Any root,
+ Callbacks&&... callbacks) { //
NOLINT(performance-unnecessary-value-param)
+ return StructuralMutateExpected(std::move(root),
std::forward<Callbacks>(callbacks)...).value();
}
} // namespace ffi
diff --git a/include/tvm/ffi/object.h b/include/tvm/ffi/object.h
index e66d0dee..d8a80efc 100644
--- a/include/tvm/ffi/object.h
+++ b/include/tvm/ffi/object.h
@@ -101,6 +101,9 @@ TVM_FFI_INLINE bool IsObjectInstance(int32_t
object_type_index);
* - _type_mutable:
* Whether we would like to expose cast to non-constant pointer
* ObjectType* from Any/AnyView. By default, we set to false so it is not
exposed.
+ * - _type_s_eq_hash_subclass_kind_fixed:
+ * Whether every subclass must retain this type's structural equality and
hash kind.
+ * By default, this is false so downstream subclasses may select their
own kind.
*
* The following two fields are necessary for base classes that can be
sub-classed.
*
@@ -232,6 +235,8 @@ class Object {
static constexpr int32_t _type_depth = 0;
/*! \brief The structural equality and hash kind of the type */
static constexpr TVMFFISEqHashKind _type_s_eq_hash_kind =
kTVMFFISEqHashKindUnsupported;
+ /*! \brief Whether subclasses must retain this type's structural equality
and hash kind */
+ static constexpr bool _type_s_eq_hash_subclass_kind_fixed = false;
// The following functions are provided by macro
// TVM_FFI_DECLARE_OBJECT_INFO and TVM_FFI_DECLARE_OBJECT_INFO_FINAL
/*!
@@ -1069,6 +1074,10 @@ struct ObjectPtrEqual {
static constexpr int32_t _type_depth = ParentType::_type_depth + 1;
\
TVM_FFI_COLD_CODE static int32_t _GetOrAllocRuntimeTypeIndex() {
\
static_assert(!ParentType::_type_final, "ParentType marked as final");
\
+ static_assert(!ParentType::_type_s_eq_hash_subclass_kind_fixed ||
\
+ TypeName::_type_s_eq_hash_kind ==
ParentType::_type_s_eq_hash_kind, \
+ "Subclass must retain the structural equality and hash kind
of its fixed " \
+ "ancestor");
\
static_assert(TypeName::_type_child_slots == 0 ||
ParentType::_type_child_slots == 0 || \
TypeName::_type_child_slots <
ParentType::_type_child_slots, \
"Need to set _type_child_slots when parent specifies it.");
\
@@ -1092,6 +1101,10 @@ struct ObjectPtrEqual {
static constexpr int32_t _type_depth = ParentType::_type_depth + 1;
\
TVM_FFI_COLD_CODE static int32_t _GetOrAllocRuntimeTypeIndex() {
\
static_assert(!ParentType::_type_final, "ParentType marked as final");
\
+ static_assert(!ParentType::_type_s_eq_hash_subclass_kind_fixed ||
\
+ TypeName::_type_s_eq_hash_kind ==
ParentType::_type_s_eq_hash_kind, \
+ "Subclass must retain the structural equality and hash kind
of its fixed " \
+ "ancestor");
\
static_assert(TypeName::_type_child_slots == 0 ||
ParentType::_type_child_slots == 0 || \
TypeName::_type_child_slots <
ParentType::_type_child_slots, \
"Need to set _type_child_slots when parent specifies it.");
\
diff --git a/python/tvm_ffi/__init__.py b/python/tvm_ffi/__init__.py
index 1f4e812d..f3ee8c86 100644
--- a/python/tvm_ffi/__init__.py
+++ b/python/tvm_ffi/__init__.py
@@ -87,6 +87,7 @@ if TYPE_CHECKING or not _is_config_mode():
structural_equal,
structural_hash,
structural_map,
+ structural_mutate,
structural_visit,
structural_walk,
)
@@ -189,6 +190,7 @@ __all__ = [
"structural_equal",
"structural_hash",
"structural_map",
+ "structural_mutate",
"structural_visit",
"structural_walk",
"system_lib",
diff --git a/python/tvm_ffi/_ffi_api.py b/python/tvm_ffi/_ffi_api.py
index 09fa770d..e5e1842b 100644
--- a/python/tvm_ffi/_ffi_api.py
+++ b/python/tvm_ffi/_ffi_api.py
@@ -112,8 +112,9 @@ if TYPE_CHECKING:
def StructuralKey(_0: Any, /) -> _StructuralKey: ...
def StructuralKeyEqual(_0: Any, _1: Any, /) -> bool: ...
def StructuralMap(_0: Any, _1: Sequence[tuple[int, Callable[..., Any]]],
_2: Sequence[tuple[int, Callable[..., Any]]], _3: int, /) -> Any: ...
+ def StructuralMutate(_0: Any, _1: Sequence[tuple[int, Callable[..., Any],
bool]], /) -> Any: ...
def StructuralMutatorDefRegionKind(_0: _StructuralMutator, /) -> int: ...
- def StructuralMutatorMaybeInplaceMutate(_0: _StructuralMutator, _1: Any,
/) -> Any: ...
+ def StructuralMutatorDefaultMutate(_0: _StructuralMutator, _1: Any, /) ->
Any: ...
def StructuralMutatorMutate(_0: _StructuralMutator, _1: Any, /) -> Any: ...
def StructuralMutatorVarRemapGet(_0: _StructuralMutator, _1: Any, /) ->
Any: ...
def StructuralMutatorVarRemapSet(_0: _StructuralMutator, _1: Any, _2: Any,
/) -> None: ...
@@ -212,8 +213,9 @@ __all__ = [
"StructuralKey",
"StructuralKeyEqual",
"StructuralMap",
+ "StructuralMutate",
"StructuralMutatorDefRegionKind",
- "StructuralMutatorMaybeInplaceMutate",
+ "StructuralMutatorDefaultMutate",
"StructuralMutatorMutate",
"StructuralMutatorVarRemapGet",
"StructuralMutatorVarRemapSet",
diff --git a/python/tvm_ffi/structural.py b/python/tvm_ffi/structural.py
index 60d3b229..2618d0eb 100644
--- a/python/tvm_ffi/structural.py
+++ b/python/tvm_ffi/structural.py
@@ -19,6 +19,7 @@
from __future__ import annotations
+import inspect
from collections.abc import Callable, Sequence
from enum import IntEnum
from typing import TYPE_CHECKING, Any
@@ -42,6 +43,7 @@ __all__ = [
"structural_equal",
"structural_hash",
"structural_map",
+ "structural_mutate",
"structural_visit",
"structural_walk",
]
@@ -450,11 +452,11 @@ class StructuralMutator(Object):
mutation hooks.
"""
- def maybe_inplace_mutate(self, value: Any) -> Any:
- """Mutate ``value``, permitting an in-place implementation when safe.
+ def mutate(self, value: Any) -> Any:
+ """Mutate ``value`` without modifying it in place.
- The caller must ensure that an object-backed ``value`` is safe to
mutate
- in place and use :meth:`mutate` for a shared object.
+ The original value is returned when none of its structural fields
+ change; otherwise, the result is a mutated copy.
Parameters
----------
@@ -464,21 +466,22 @@ class StructuralMutator(Object):
Returns
-------
result
- The mutated owning value. It may refer to the same object as
``value``.
+ The mutated owning value.
"""
- return _ffi_api.StructuralMutatorMaybeInplaceMutate(self, value)
+ return _ffi_api.StructuralMutatorMutate(self, value)
- def mutate(self, value: Any) -> Any:
- """Mutate ``value`` without modifying it in place.
+ def default_mutate(self, value: Any) -> Any:
+ """Mutate ``value`` using its registered or reflected default behavior.
- The original value is returned when none of its structural fields
- change; otherwise, the result is a mutated copy.
+ This bypasses the active engine callback for ``value`` itself while
+ recursive children re-enter the same mutator. A ``structural_mutate``
+ callback may use it on its matched value to request default descent.
Parameters
----------
value
- Value to mutate.
+ Value whose default mutation should run.
Returns
-------
@@ -486,7 +489,9 @@ class StructuralMutator(Object):
The mutated owning value.
"""
- return _ffi_api.StructuralMutatorMutate(self, value)
+ return _ffi_api.StructuralMutatorDefaultMutate( # ty:
ignore[unresolved-attribute]
+ self, value
+ )
def var_remap_get(self, var: Object) -> Any | None:
"""Return the replacement recorded for a variable identity.
@@ -681,6 +686,49 @@ def structural_visit(
return _ffi_api.StructuralVisit(root, entries)
+def structural_mutate(
+ root: Any,
+ callbacks: tuple | Sequence | Callable = (),
+) -> Any:
+ """Mutate a value with callbacks that own recursive mutation.
+
+ Each callback receives ``(value, mutator)`` and may optionally receive a
+ third ``allow_inplace`` boolean, which is true only when the callback's
+ value is on a uniquely owned path. The flag lets a callback choose an
+ ownership-aware implementation; recursive descent still uses
+ :meth:`StructuralMutator.mutate` for selected children or
+ :meth:`StructuralMutator.default_mutate` for the matched value's default
+ mutation. Its returned value is final and is not traversed again. Entries
+ use ``structural_map`` matching rules, and an unmatched value follows
+ registered/default mutation.
+
+ Parameters
+ ----------
+ root
+ Root value to mutate. Passing a regular Python reference preserves it
+ through copy-on-write; passing ``root._move()`` transfers ownership and
+ permits in-place mutation along unique paths.
+ callbacks
+ Callback entries tried in order; the first match owns mutation.
+
+ Returns
+ -------
+ result
+ The mutated owning value.
+
+ """
+ callback_entries = _normalize_callbacks(callbacks,
api_name="structural_mutate")
+ entries: list[tuple[int, Callable[..., Any], bool]] = [
+ (
+ _callback_type_to_type_index(t, api_name="structural_mutate"),
+ fn,
+ _callback_accepts_allow_inplace(fn),
+ )
+ for t, fn in callback_entries
+ ]
+ return _ffi_api.StructuralMutate(root, entries) # ty:
ignore[unresolved-attribute]
+
+
def structural_map(
root: Any,
callbacks: tuple | Sequence | Callable = (),
@@ -695,7 +743,9 @@ def structural_map(
Parameters
----------
root
- Root value to map.
+ Root value to map. Passing a regular Python reference preserves it
+ through copy-on-write; passing ``root._move()`` transfers ownership and
+ permits in-place mutation along unique paths.
callbacks
Normal callbacks. These callbacks receive one argument, ``value``, and
@@ -803,6 +853,29 @@ def _normalize_callbacks(
return callback_entries
+def _callback_accepts_allow_inplace(callback: Callable[..., Any]) -> bool:
+ """Return whether a StructuralMutate callback accepts its optional flag."""
+ try:
+ signature = inspect.signature(callback)
+ except (TypeError, ValueError):
+ # Some extension callables do not expose a signature. Keep their
+ # existing two-argument, copy-safe behavior.
+ return False
+
+ try:
+ signature.bind(None, None, False)
+ except TypeError:
+ try:
+ signature.bind(None, None)
+ except TypeError as err:
+ raise TypeError(
+ "structural_mutate callback must accept (value, mutator) or "
+ "(value, mutator, allow_inplace)"
+ ) from err
+ return False
+ return True
+
+
def _callback_type_to_type_index(callback_type: type[Any] | Any, *, api_name:
str) -> int:
"""Convert a callback arg type to a type index."""
annotation = Any if callback_type is object else callback_type
diff --git a/src/ffi/extra/structural_mutate.cc
b/src/ffi/extra/structural_mutate.cc
index f094cafc..71511a94 100644
--- a/src/ffi/extra/structural_mutate.cc
+++ b/src/ffi/extra/structural_mutate.cc
@@ -47,8 +47,10 @@ namespace details {
* \param order Integer value of \ref WalkOrder.
* \return The mapped owning value, or an Error.
*/
+// The owning parameter makes caller ownership visible to the uniqueness check.
Expected<Any> StructuralMapExpected(
- AnyView root, const Array<Tuple<int32_t, Function>>& callbacks,
+ Any root, // NOLINT(performance-unnecessary-value-param)
+ const Array<Tuple<int32_t, Function>>& callbacks,
const Array<Tuple<int32_t, Function>>& callbacks_with_def_region_kind, int
order) noexcept {
if (order == static_cast<int>(WalkOrder::kPreOrder)) {
using Mutator = StructuralMapDynEngine<StructuralMapEngineBase,
WalkOrder::kPreOrder>;
@@ -61,6 +63,88 @@ Expected<Any> StructuralMapExpected(
}
}
+/*! \brief Runtime counterpart of the typed callback-owned mutate engine. */
+template <typename Parent>
+class StructuralMutateDynEngine : public Parent {
+ public:
+ explicit StructuralMutateDynEngine(Array<Tuple<int32_t, Function, bool>>
callbacks)
+ : Parent(VTable()), callbacks_(std::move(callbacks)) {}
+
+ private:
+ static const StructuralMutatorVTable* VTable() {
+ static const StructuralMutatorVTable vtable{
+ &StructuralMutateDynEngine::DispatchMutate,
+ &StructuralMutateDynEngine::DispatchMaybeInplaceMutate,
+ &StructuralMutateDynEngine::DispatchVarRemapGet,
+ &StructuralMutateDynEngine::DispatchVarRemapSet,
+ };
+ return &vtable;
+ }
+
+ static TVMFFIAny DispatchMutate(StructuralMutatorObj* mutator, AnyView
value) noexcept {
+ return
static_cast<StructuralMutateDynEngine*>(mutator)->MutateImplRaw(value);
+ }
+
+ static TVMFFIAny DispatchMaybeInplaceMutate(StructuralMutatorObj* mutator,
+ AnyView value) noexcept {
+ return
static_cast<StructuralMutateDynEngine*>(mutator)->MaybeInplaceMutateImplRaw(value);
+ }
+
+ std::optional<Expected<Any>> DispatchCallback(AnyView value, bool
allow_inplace) noexcept {
+ for (const auto& entry : callbacks_) {
+ if (!RuntimeTypeIndexMatch(value.type_index(), entry.template get<0>()))
continue;
+ if (entry.template get<2>()) {
+ return entry.template get<1>().template CallExpected<Any>(
+ value, GetRef<StructuralMutator>(this), allow_inplace);
+ }
+ return entry.template get<1>().template CallExpected<Any>(value,
+
GetRef<StructuralMutator>(this));
+ }
+ return std::nullopt;
+ }
+
+ TVMFFIAny MutateImplRaw(AnyView value) noexcept {
+ if (std::optional<Expected<Any>> matched = DispatchCallback(value, false))
{
+ Expected<Any> result = *std::move(matched);
+ if (TVM_FFI_PREDICT_FALSE(result.is_err())) {
+ Parent::UpdateVisitErrorContext(result, value);
+ }
+ return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
+ }
+ return Parent::DefaultMutateRaw(value);
+ }
+
+ TVMFFIAny MaybeInplaceMutateImplRaw(AnyView value) noexcept {
+ if (std::optional<Expected<Any>> matched = DispatchCallback(value, true)) {
+ Expected<Any> result = *std::move(matched);
+ if (TVM_FFI_PREDICT_FALSE(result.is_err())) {
+ Parent::UpdateVisitErrorContext(result, value);
+ }
+ return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
+ }
+ return Parent::DefaultMaybeInplaceMutateRaw(value);
+ }
+
+ Array<Tuple<int32_t, Function, bool>> callbacks_;
+};
+
+/*!
+ * \brief Runtime callback-driven structural mutation.
+ * \param root The root value to mutate.
+ * \param callbacks Runtime ``(type_index, callback, accepts_allow_inplace)``
entries. A callback
+ * whose marker is true is invoked as ``callback(value,
mutator, allow_inplace)``;
+ * otherwise it is invoked as ``callback(value, mutator)``.
+ * \return The mutated owning value, or an Error.
+ */
+// The owning parameter makes caller ownership visible to the uniqueness check.
+Expected<Any> StructuralMutateExpected(
+ Any root, // NOLINT(performance-unnecessary-value-param)
+ const Array<Tuple<int32_t, Function, bool>>& callbacks) noexcept {
+ using Mutator = StructuralMutateDynEngine<StructuralMapEngineBase>;
+ StructuralMutator mutator(make_object<Mutator>(callbacks));
+ return mutator->MaybeInplaceMutateIfUniqueExpected(root);
+}
+
// ---------------------------------------------------------------------------
// Built-in container structural mutation.
// ---------------------------------------------------------------------------
@@ -254,9 +338,11 @@ TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::ObjectDef<StructuralMutatorObj>(); // NOLINT(bugprone-unused-raii)
refl::GlobalDef()
- .def_method("ffi.StructuralMutatorMaybeInplaceMutate",
- &StructuralMutatorObj::MaybeInplaceMutate)
.def_method("ffi.StructuralMutatorMutate", &StructuralMutatorObj::Mutate)
+ .def_method("ffi.StructuralMutatorDefaultMutate",
+ [](const StructuralMutator& mutator, AnyView value) {
+ return mutator->DefaultMutateExpected(value).value();
+ })
.def_method("ffi.StructuralMutatorVarRemapGet",
[](const StructuralMutator& mutator, AnyView var) {
return mutator->VarRemapGetExpected(var).value();
@@ -272,12 +358,20 @@ TVM_FFI_STATIC_INIT_BLOCK() {
return mutator->WithDefRegionKind(kind, callback);
})
.def("ffi.StructuralMap",
- [](AnyView root, const Array<Tuple<int32_t, Function>>& callbacks,
+ // The owning parameter makes caller ownership visible to the
uniqueness check.
+ [](Any root, // NOLINT(performance-unnecessary-value-param)
+ const Array<Tuple<int32_t, Function>>& callbacks,
const Array<Tuple<int32_t, Function>>&
callbacks_with_def_region_kind,
int32_t order) -> Any {
- return details::StructuralMapExpected(root, callbacks,
callbacks_with_def_region_kind,
- order)
+ return details::StructuralMapExpected(std::move(root), callbacks,
+
callbacks_with_def_region_kind, order)
.value();
+ })
+ .def("ffi.StructuralMutate",
+ // The owning parameter makes caller ownership visible to the
uniqueness check.
+ [](Any root, // NOLINT(performance-unnecessary-value-param)
+ const Array<Tuple<int32_t, Function, bool>>& callbacks) -> Any {
+ return details::StructuralMutateExpected(std::move(root),
callbacks).value();
});
refl::EnsureTypeAttrColumn(refl::type_attr::kStructuralMutate);
refl::EnsureTypeAttrColumn(refl::type_attr::kStructuralMaybeInplaceMutate);
diff --git a/tests/cpp/extra/test_structural_mutate.cc
b/tests/cpp/extra/test_structural_mutate.cc
index 8e77a8d9..0170c45d 100644
--- a/tests/cpp/extra/test_structural_mutate.cc
+++ b/tests/cpp/extra/test_structural_mutate.cc
@@ -40,7 +40,64 @@ using namespace tvm::ffi::testing;
using AnyArray = Array<Any>;
using StringMap = Map<String, Any>;
-TVM_FFI_STATIC_INIT_BLOCK() { TMutatePairObj::RegisterReflection(); }
+class TNestedMapHookObj : public Object {
+ public:
+ AnyArray field;
+
+ explicit TNestedMapHookObj(AnyArray field) : field(std::move(field)) {}
+
+ static TVMFFIAny StructuralMutate(StructuralMutatorObj* mutator, AnyView
value) noexcept {
+ const auto* self = value.cast<const TNestedMapHookObj*>();
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, mapped,
mutator->MutateExpected(self->field));
+ AnyArray mapped_field = mapped.cast<AnyArray>();
+ if (mapped_field.same_as(self->field)) {
+ return details::AnyUnsafe::MoveAnyToTVMFFIAny(Any(value));
+ }
+ return details::AnyUnsafe::MoveAnyToTVMFFIAny(
+ Any(make_object<TNestedMapHookObj>(std::move(mapped_field))));
+ }
+
+ static TVMFFIAny MaybeInplaceMutate(StructuralMutatorObj*, AnyView value)
noexcept {
+ auto* self = value.cast<TNestedMapHookObj*>();
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ Any, mapped,
+ StructuralMapExpected<WalkOrder::kPostOrder>(
+ Any(std::move(self->field)),
+ [](int64_t item) -> Expected<Any> { return Any(item + 1); }));
+ self->field = mapped.cast<AnyArray>();
+ return details::AnyUnsafe::MoveAnyToTVMFFIAny(Any(value));
+ }
+
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<TNestedMapHookObj>().def_rw("field",
&TNestedMapHookObj::field);
+ refl::EnsureTypeAttrColumn(refl::type_attr::kStructuralMutate);
+ refl::EnsureTypeAttrColumn(refl::type_attr::kStructuralMaybeInplaceMutate);
+ refl::TypeAttrDef<TNestedMapHookObj>()
+ .attr(refl::type_attr::kStructuralMutate,
+
reinterpret_cast<void*>(static_cast<FStructuralMutate>(&StructuralMutate)))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+
reinterpret_cast<void*>(static_cast<FStructuralMutate>(&MaybeInplaceMutate)));
+ }
+
+ static constexpr bool _type_mutable = true;
+ static constexpr TVMFFISEqHashKind _type_s_eq_hash_kind =
kTVMFFISEqHashKindTreeNode;
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("test.NestedMapHook", TNestedMapHookObj,
Object);
+};
+
+class TNestedMapHook : public ObjectRef {
+ public:
+ explicit TNestedMapHook(AnyArray field) {
+ data_ = make_object<TNestedMapHookObj>(std::move(field));
+ }
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(TNestedMapHook, ObjectRef,
TNestedMapHookObj);
+};
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ TMutatePairObj::RegisterReflection();
+ TNestedMapHookObj::RegisterReflection();
+}
Expected<Any> Increment(int64_t value) { return Any(value + 1); }
@@ -91,6 +148,16 @@ class StructuralMapWithMutateCount : public
StructuralMapEngineBase {
int marker_ = 17;
};
+class StructuralMutateLayer : public StructuralMapEngineBase {
+ public:
+ using MutatorObjType = StructuralMutateLayer;
+
+ explicit StructuralMutateLayer(const StructuralMutatorVTable* vtable)
+ : StructuralMapEngineBase(vtable) {}
+
+ int callback_tag() const { return 23; }
+};
+
TEST(StructuralMap, ParentLayerOwnsBothDescentsAndProvidesState) {
std::vector<int> callback_counts;
auto identity = [&](const AnyArray& value, const MutateCount& live_count,
const int& live_marker,
@@ -136,6 +203,175 @@ TEST(StructuralMap,
ParentLayerOwnsBothDescentsAndProvidesState) {
EXPECT_TRUE(mapped[0].cast<TVar>().same_as(mapped[1].cast<TVar>()));
}
+TEST(StructuralMutate, CallbackOwnsMutationAndErrorsStayExpected) {
+ std::vector<int64_t> trace;
+ auto mutate_array = [&](const AnyArray& value, StructuralMutateLayer*
mutator) -> Expected<Any> {
+ EXPECT_EQ(mutator->callback_tag(), 23);
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, first,
mutator->MutateExpected(value[0]));
+ return Any(AnyArray{std::move(first), int64_t{10}});
+ };
+ auto mutate_int = [&](int64_t value, StructuralMutateLayer*) ->
Expected<Any> {
+ trace.push_back(value);
+ return Any(value + 1);
+ };
+ using Mutator =
+ StructuralMutateEngine<StructuralMutateLayer, decltype(mutate_array),
decltype(mutate_int)>;
+ StructuralMutator mutator(make_object<Mutator>(std::move(mutate_array),
std::move(mutate_int)));
+
+ AnyArray mapped =
+ mutator->MutateExpected(AnyArray{int64_t{1},
int64_t{2}}).value().cast<AnyArray>();
+ ASSERT_EQ(mapped.size(), 2U);
+ EXPECT_EQ(mapped[0].cast<int64_t>(), 2);
+ EXPECT_EQ(mapped[1].cast<int64_t>(), 10);
+ EXPECT_EQ(trace, std::vector<int64_t>{1});
+
+ AnyArray default_mapped =
+ StructuralMutate(
+ AnyArray{int64_t{3}, int64_t{4}},
+ [](int64_t value, StructuralMutatorObj*) -> Expected<Any> { return
Any(value + 1); })
+ .cast<AnyArray>();
+ EXPECT_EQ(default_mapped[0].cast<int64_t>(), 4);
+ EXPECT_EQ(default_mapped[1].cast<int64_t>(), 5);
+
+ Expected<Any> returned_error =
+ StructuralMutateExpected(int64_t{1}, [](int64_t, StructuralMutatorObj*)
-> Expected<Any> {
+ return Unexpected(Error("ValueError", "returned mutate error", ""));
+ });
+ ASSERT_TRUE(returned_error.is_err());
+ EXPECT_EQ(returned_error.error().message(), "returned mutate error");
+
+ Expected<Any> thrown_error =
+ StructuralMutateExpected(int64_t{1}, [](int64_t, StructuralMutatorObj*)
-> Expected<Any> {
+ TVM_FFI_THROW(ValueError) << "thrown mutate error";
+ return Any(nullptr);
+ });
+ ASSERT_TRUE(thrown_error.is_err());
+ EXPECT_EQ(thrown_error.error().message(), "thrown mutate error");
+}
+
+TEST(StructuralMutate, CallbackControlsRecursion) {
+ TPair root(TPair(TInt(1), TInt(2)), TPair(TInt(3), TInt(4)));
+ ObjectRef original_rhs = root->rhs;
+
+ TPair mapped =
+ StructuralMutate(
+ root,
+ [](const TPair& pair, StructuralMutatorObj* mutator) ->
Expected<Any> {
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, lhs,
mutator->MutateExpected(pair->lhs));
+ return Any(TPair(lhs.cast<ObjectRef>(), pair->rhs));
+ },
+ [](const TInt& value, StructuralMutatorObj*) -> Expected<Any> {
+ return Any(TInt(value->value + 100));
+ })
+ .cast<TPair>();
+
+ TPair mapped_lhs = mapped->lhs.as_or_throw<TPair>();
+ TPair mapped_rhs = mapped->rhs.as_or_throw<TPair>();
+ EXPECT_EQ(mapped_lhs->lhs.as_or_throw<TInt>()->value, 101);
+ EXPECT_EQ(mapped_lhs->rhs.as_or_throw<TInt>()->value, 2);
+ EXPECT_EQ(mapped_rhs->lhs.as_or_throw<TInt>()->value, 3);
+ EXPECT_EQ(mapped_rhs->rhs.as_or_throw<TInt>()->value, 4);
+ EXPECT_TRUE(mapped->rhs.same_as(original_rhs));
+}
+
+TEST(StructuralMutate, PreservesUniqueContainerIdentity) {
+ AnyArray inner{int64_t{1}};
+ const Object* inner_address = inner.get();
+ AnyArray root{Any(std::move(inner))};
+ const Object* root_address = root.get();
+
+ AnyArray mapped =
+ StructuralMutate(std::move(root), [](int64_t value,
StructuralMutatorObj*) -> Expected<Any> {
+ return Any(value + 1);
+ }).cast<AnyArray>();
+
+ AnyArray mapped_inner = mapped[0].cast<AnyArray>();
+ EXPECT_EQ(mapped.get(), root_address);
+ EXPECT_EQ(mapped_inner.get(), inner_address);
+ EXPECT_EQ(mapped_inner[0].cast<int64_t>(), 2);
+}
+
+TEST(StructuralMutate, RootByValueProtectsSharedParentSubvalue) {
+ AnyArray child{int64_t{1}};
+ AnyArray outer{Any(std::move(child))};
+ const Object* child_address = outer[0].cast<AnyArray>().get();
+
+ AnyArray mapped =
+ StructuralMutate(outer[0], [](int64_t value, StructuralMutatorObj*) ->
Expected<Any> {
+ return Any(value + 1);
+ }).cast<AnyArray>();
+
+ EXPECT_NE(mapped.get(), child_address);
+ EXPECT_EQ(outer[0].cast<AnyArray>()[0].cast<int64_t>(), 1);
+ EXPECT_EQ(mapped[0].cast<int64_t>(), 2);
+}
+
+TEST(StructuralMutate, CallbackArityControlsInplaceMutation) {
+ AnyArray inplace_root{int64_t{1}};
+ AnyArray copy_on_write_root{int64_t{1}};
+ const Object* inplace_root_address = inplace_root.get();
+ const Object* copy_on_write_root_address = copy_on_write_root.get();
+ std::vector<bool> allow_inplace_trace;
+
+ AnyArray inplace_mapped =
+ StructuralMutate(
+ std::move(inplace_root),
+ [&](const AnyArray& value, StructuralMutatorObj* mutator,
+ bool allow_inplace) -> Expected<Any> {
+ allow_inplace_trace.push_back(allow_inplace);
+ return allow_inplace ?
mutator->DefaultMaybeInplaceMutateExpected(value)
+ : mutator->DefaultMutateExpected(value);
+ },
+ [&](int64_t value, StructuralMutatorObj*, bool allow_inplace) ->
Expected<Any> {
+ allow_inplace_trace.push_back(allow_inplace);
+ return Any(value + 1);
+ })
+ .cast<AnyArray>();
+
+ AnyArray copy_on_write_mapped =
+ StructuralMutate(
+ std::move(copy_on_write_root),
+ [](const AnyArray& value, StructuralMutatorObj* mutator) ->
Expected<Any> {
+ return mutator->DefaultMutateExpected(value);
+ },
+ [](int64_t value, StructuralMutatorObj*) -> Expected<Any> { return
Any(value + 1); })
+ .cast<AnyArray>();
+
+ EXPECT_EQ(inplace_mapped.get(), inplace_root_address);
+ EXPECT_NE(copy_on_write_mapped.get(), copy_on_write_root_address);
+ EXPECT_EQ(inplace_mapped[0].cast<int64_t>(), 2);
+ EXPECT_EQ(copy_on_write_mapped[0].cast<int64_t>(), 2);
+ EXPECT_EQ(allow_inplace_trace, (std::vector<bool>{true, false}));
+}
+
+TEST(StructuralMutate, MatchedVarOwnsRemapConsistency) {
+ TVar var("n");
+ AnyArray root{var, var};
+ int callback_count = 0;
+
+ AnyArray mapped =
+ StructuralMutate(
+ root,
+ [&](const TVar& value, StructuralMutatorObj* mutator) ->
Expected<Any> {
+ ++callback_count;
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, remapped,
mutator->VarRemapGetExpected(value));
+ if (remapped.type_index() != TypeIndex::kTVMFFINone) {
+ return remapped;
+ }
+ Any replacement(TVar(value->name + "-mapped"));
+ Expected<void> set_result = mutator->VarRemapSetExpected(value,
replacement);
+ if (set_result.is_err()) {
+ return Unexpected(std::move(set_result).error());
+ }
+ return replacement;
+ })
+ .cast<AnyArray>();
+
+ EXPECT_EQ(callback_count, 2);
+ EXPECT_TRUE(mapped[0].cast<TVar>().same_as(mapped[1].cast<TVar>()));
+ EXPECT_EQ(mapped[0].cast<TVar>()->name, "n-mapped");
+}
+
template <WalkOrder order>
void CheckNestedArrayMapOrder(const std::vector<std::string>& expected_trace) {
AnyArray inner_array{int64_t{1}};
@@ -148,7 +384,7 @@ void CheckNestedArrayMapOrder(const
std::vector<std::string>& expected_trace) {
AnyArray mapped =
StructuralMap<order>(
- root,
+ std::move(root),
[&](const AnyArray& array) -> Expected<Any> {
trace.emplace_back(array.get() == root_address ? "outer-array" :
"inner-array");
return Any(array);
@@ -184,6 +420,33 @@ TEST(StructuralMap,
MapsNestedArrayAndMapInConfiguredOrder) {
CheckNestedArrayMapOrder<WalkOrder::kPostOrder>({"int", "inner-array",
"map", "outer-array"});
}
+TEST(StructuralMap, RootByValueProtectsSharedParentSubvalue) {
+ AnyArray child{int64_t{1}};
+ AnyArray outer{Any(std::move(child))};
+ const Object* child_address = outer[0].cast<AnyArray>().get();
+
+ AnyArray mapped = StructuralMap<WalkOrder::kPostOrder>(outer[0],
Increment).cast<AnyArray>();
+
+ EXPECT_NE(mapped.get(), child_address);
+ EXPECT_EQ(outer[0].cast<AnyArray>()[0].cast<int64_t>(), 1);
+ EXPECT_EQ(mapped[0].cast<int64_t>(), 2);
+}
+
+TEST(StructuralMap, MaybeInplaceHookMovesNestedFieldIntoStructuralMap) {
+ TNestedMapHook root(AnyArray{int64_t{1}});
+ const Object* root_address = root.get();
+ const Object* field_address = root->field.get();
+
+ TNestedMapHook mapped =
+ StructuralMap<WalkOrder::kPostOrder>(
+ Any(std::move(root)), [](const String& value) -> Expected<Any> {
return Any(value); })
+ .cast<TNestedMapHook>();
+
+ EXPECT_EQ(mapped.get(), root_address);
+ EXPECT_EQ(mapped->field.get(), field_address);
+ EXPECT_EQ(mapped->field[0].cast<int64_t>(), 2);
+}
+
TEST(StructuralMap, RegisteredMutateHookUsesAssignOrReturn) {
TVar lhs("lhs");
TVar rhs("rhs");
@@ -552,4 +815,31 @@ TEST(StructuralMapDyn,
ParentLayerRunsThroughHeaderDefinedEngine) {
CheckDynamicParentLayer<WalkOrder::kPostOrder>();
}
+Any CallDynStructuralMutate(Any root, //
NOLINT(performance-unnecessary-value-param)
+ const Array<Tuple<int32_t, Function, bool>>&
callbacks) {
+ Function fn = Function::GetGlobalRequired("ffi.StructuralMutate");
+ return fn(std::move(root), callbacks);
+}
+
+TEST(StructuralMutateDyn, PreservesDistinctDefaultDescentPaths) {
+ Function increment = Function::FromTyped(
+ [](int64_t value, const StructuralMutator&) -> Any { return Any(value +
1); });
+ Array<Tuple<int32_t, Function, bool>> callbacks{
+ Tuple<int32_t, Function, bool>(TypeIndex::kTVMFFIInt, increment, false)};
+
+ AnyArray unique_root{int64_t{1}};
+ AnyArray unique_mapped = CallDynStructuralMutate(unique_root,
callbacks).cast<AnyArray>();
+ EXPECT_FALSE(unique_mapped.same_as(unique_root));
+ EXPECT_EQ(unique_root[0].cast<int64_t>(), 1);
+ EXPECT_EQ(unique_mapped[0].cast<int64_t>(), 2);
+
+ AnyArray shared_root{int64_t{1}};
+ AnyArray extra_owner = shared_root; //
NOLINT(performance-unnecessary-copy-initialization)
+ AnyArray shared_mapped = CallDynStructuralMutate(shared_root,
callbacks).cast<AnyArray>();
+ EXPECT_FALSE(shared_mapped.same_as(shared_root));
+ EXPECT_TRUE(extra_owner.same_as(shared_root));
+ EXPECT_EQ(shared_root[0].cast<int64_t>(), 1);
+ EXPECT_EQ(shared_mapped[0].cast<int64_t>(), 2);
+}
+
} // namespace
diff --git a/tests/cpp/extra/test_structural_visit.cc
b/tests/cpp/extra/test_structural_visit.cc
index 82a6bccf..75513128 100644
--- a/tests/cpp/extra/test_structural_visit.cc
+++ b/tests/cpp/extra/test_structural_visit.cc
@@ -750,4 +750,25 @@ TEST(StructuralVisit, CallbackDrivenTraversal) {
ExpectTrace(error_trace, {"throw"});
}
+TEST(StructuralVisit, CallbackVisitsLhsOnly) {
+ TPair root(TPair(TVar("lhs"), TVar("inner-rhs")),
+ TPair(TVar("root-rhs-lhs"), TVar("root-rhs-rhs")));
+ std::vector<std::string> trace;
+
+ Expected<Optional<VisitInterrupt>> result = StructuralVisitExpected(
+ root,
+ [](const TPair& pair, StructuralVisitorObj* visitor) ->
Expected<Optional<VisitInterrupt>> {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(pair->lhs));
+ return Optional<VisitInterrupt>(std::nullopt);
+ },
+ [&](const TVar& var, StructuralVisitorObj*) ->
Expected<Optional<VisitInterrupt>> {
+ trace.emplace_back(var->name);
+ return Optional<VisitInterrupt>(std::nullopt);
+ });
+
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
+ ExpectTrace(trace, {"lhs"});
+}
+
} // namespace
diff --git a/tests/python/test_structural.py b/tests/python/test_structural.py
index 9ad0f3ae..96873830 100644
--- a/tests/python/test_structural.py
+++ b/tests/python/test_structural.py
@@ -329,6 +329,90 @@ def test_structural_visit_default_visit_binding() -> None:
assert nested_trace == ["array", 1, 2]
+def test_structural_mutate_callback_owned_recursion_and_errors() -> None:
+ trace: list[int | str] = []
+
+ def mutate_array(value: tvm_ffi.Array, mutator: tvm_ffi.StructuralMutator)
-> object:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ trace.append("array")
+ return tvm_ffi.Array([mutator.mutate(value[0]), 10])
+
+ def mutate_int(value: int, mutator: tvm_ffi.StructuralMutator) -> int:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ trace.append(value)
+ return value + 1
+
+ mapped = tvm_ffi.structural_mutate(
+ tvm_ffi.Array([1, 2]),
+ [(tvm_ffi.Array, mutate_array), (int, mutate_int)],
+ )
+ assert list(mapped) == [2, 10]
+ assert trace == ["array", 1]
+
+ default_trace: list[int] = []
+
+ def default_mutate_array(value: tvm_ffi.Array, mutator:
tvm_ffi.StructuralMutator) -> object:
+ return mutator.default_mutate(value)
+
+ def default_mutate_int(value: int, mutator: tvm_ffi.StructuralMutator) ->
int:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ default_trace.append(value)
+ return value + 1
+
+ default_owned = tvm_ffi.structural_mutate(
+ tvm_ffi.Array([1, 2]),
+ [(tvm_ffi.Array, default_mutate_array), (int, default_mutate_int)],
+ )
+ assert list(default_owned) == [2, 3]
+ assert default_trace == [1, 2]
+
+ default_root = tvm_ffi.Array([3, 4])
+ default_mapped = tvm_ffi.structural_mutate(default_root, (int, mutate_int))
+ assert not default_mapped.same_as(default_root)
+ assert list(default_root) == [3, 4]
+ assert list(default_mapped) == [4, 5]
+
+ inplace_trace: list[bool] = []
+
+ def mutate_with_flag(
+ value: int, mutator: tvm_ffi.StructuralMutator, allow_inplace: bool
+ ) -> int:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ inplace_trace.append(allow_inplace)
+ return value + 1
+
+ flagged_root = tvm_ffi.Array([1])
+ flagged_mapped = tvm_ffi.structural_mutate(flagged_root, (int,
mutate_with_flag))
+ assert inplace_trace == [False]
+ assert not flagged_mapped.same_as(flagged_root)
+ assert list(flagged_root) == [1]
+ assert list(flagged_mapped) == [2]
+
+ direct_trace: list[int] = []
+
+ def fail_directly(value: int, mutator: tvm_ffi.StructuralMutator) ->
object:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ direct_trace.append(value)
+ raise ValueError("direct structural mutate failure")
+
+ with pytest.raises(ValueError, match="direct structural mutate failure"):
+ tvm_ffi.structural_mutate(1, (int, fail_directly))
+ assert direct_trace == [1]
+
+ nested_trace: list[int] = []
+
+ def fail_nested(value: int, mutator: tvm_ffi.StructuralMutator) -> int:
+ assert isinstance(mutator, tvm_ffi.StructuralMutator)
+ nested_trace.append(value)
+ if value == 2:
+ raise ValueError("nested structural mutate failure")
+ return value
+
+ with pytest.raises(ValueError, match="nested structural mutate failure"):
+ tvm_ffi.structural_mutate(tvm_ffi.Array([1, 2, 3]), (int, fail_nested))
+ assert nested_trace == [1, 2]
+
+
def test_structural_walk_nested_containers_and_skips_map_keys() -> None:
root = tvm_ffi.Array(
[
@@ -478,7 +562,7 @@ def test_structural_map_nested_array_map_order_and_keys()
-> None:
trace: list[str] = []
def map_array(value: tvm_ffi.Array) -> tvm_ffi.Array:
- trace.append("outer-array" if value.same_as(root) else
"inner-array")
+ trace.append("outer-array" if isinstance(value[0], tvm_ffi.Map)
else "inner-array")
return value
def map_map(value: tvm_ffi.Map) -> tvm_ffi.Map:
@@ -504,9 +588,10 @@ def test_structural_map_nested_array_map_order_and_keys()
-> None:
else:
mapped = tvm_ffi.structural_map(root, callbacks, order=order)
- assert mapped.__chandle__() == root_handle
- assert mapped[0].__chandle__() == map_handle
- assert mapped[0]["value"].__chandle__() == inner_array_handle
+ assert mapped.__chandle__() != root_handle
+ assert mapped[0].__chandle__() != map_handle
+ assert mapped[0]["value"].__chandle__() != inner_array_handle
+ assert list(root[0]["value"]) == [1]
assert list(mapped[0]["value"]) == [2]
assert "value" in mapped[0]
assert "renamed" not in mapped[0]
@@ -522,14 +607,15 @@ def test_structural_map_nested_array_map_order_and_keys()
-> None:
def test_structural_map_array_ownership() -> None:
- # A unique outer Array is reused, but its externally shared child is
copied.
+ # Python retains the outer Array, so changed paths are copied.
shared_child = tvm_ffi.Array([1])
root = tvm_ffi.Array([shared_child])
root_handle = root.__chandle__()
mapped = tvm_ffi.structural_map(root, (int, lambda value: value + 1))
- assert mapped.__chandle__() == root_handle
+ assert mapped.__chandle__() != root_handle
assert not mapped[0].same_as(shared_child)
+ assert list(root[0]) == [1]
assert list(shared_child) == [1]
assert list(mapped[0]) == [2]
@@ -545,16 +631,36 @@ def test_structural_map_array_ownership() -> None:
assert list(shared_root[0]) == [1]
assert list(mapped[0]) == [2]
+ # Moving the root transfers its only Python-owned reference to the engine.
+ moved_root = tvm_ffi.Array([1])
+ moved_handle = moved_root.__chandle__()
+ mapped = tvm_ffi.structural_map(moved_root._move(), (int, lambda value:
value + 1))
+
+ assert mapped.__chandle__() == moved_handle
+ assert list(mapped) == [2]
+
+ # Moving a wrapper obtained from a container does not transfer the
+ # container's reference, so its retained value remains unchanged.
+ owner = tvm_ffi.Array([tvm_ffi.Array([1])])
+ borrowed = owner[0]
+ borrowed_handle = borrowed.__chandle__()
+ mapped = tvm_ffi.structural_map(borrowed._move(), (int, lambda value:
value + 1))
+
+ assert mapped.__chandle__() != borrowed_handle
+ assert list(owner[0]) == [1]
+ assert list(mapped) == [2]
+
def test_structural_map_map_value_ownership() -> None:
- # A unique Map is reused, but its externally shared value is copied.
+ # Python retains the Map, so changed paths are copied.
shared_value = tvm_ffi.Array([1])
root = tvm_ffi.Map({"value": shared_value})
root_handle = root.__chandle__()
mapped = tvm_ffi.structural_map(root, (int, lambda value: value + 1))
- assert mapped.__chandle__() == root_handle
+ assert mapped.__chandle__() != root_handle
assert not mapped["value"].same_as(shared_value)
+ assert list(root["value"]) == [1]
assert list(shared_value) == [1]
assert list(mapped["value"]) == [2]