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 35f4f3a0 [REFACTOR][EXTRA] Streamline container mutation hooks (#739)
35f4f3a0 is described below
commit 35f4f3a0143e51b2ff8a55a885100fd28020c1ce
Author: Tianqi Chen <[email protected]>
AuthorDate: Fri Sep 4 14:37:07 2026 -0400
[REFACTOR][EXTRA] Streamline container mutation hooks (#739)
Move structural container mutation onto the same exception-free contract
as structural visiting: child failures remain explicit `Expected`
values, while allocation and invariant failures stay fatal.
- share prefix range initialization between Array and List
- avoid throwing dispatch casts and repeated bounds checks after
invariants are established
- cover raised and returned callback errors across Array, List, Map, and
Dict
---
include/tvm/ffi/any.h | 11 ++-
include/tvm/ffi/container/array.h | 18 ----
include/tvm/ffi/container/seq_base.h | 29 ++++++
include/tvm/ffi/container/variant.h | 2 +-
include/tvm/ffi/extra/structural_mutate.h | 5 ++
src/ffi/extra/structural_mutate.cc | 141 ++++++++++++++----------------
tests/cpp/extra/test_structural_mutate.cc | 37 ++++++++
tests/cpp/test_any.cc | 5 ++
8 files changed, 150 insertions(+), 98 deletions(-)
diff --git a/include/tvm/ffi/any.h b/include/tvm/ffi/any.h
index 9f243efb..3d203df3 100644
--- a/include/tvm/ffi/any.h
+++ b/include/tvm/ffi/any.h
@@ -51,6 +51,7 @@ class AnyView {
TVMFFIAny data_;
// Any can see AnyView
friend class Any;
+ friend struct details::AnyUnsafe;
public:
// NOTE: the following functions use style
@@ -629,8 +630,14 @@ struct AnyUnsafe : public ObjectUnsafe {
}
}
- TVM_FFI_INLINE static Object* ObjectPtrFromAnyAfterCheck(const Any& ref) {
- return reinterpret_cast<Object*>(ref.data_.v_obj);
+ template <typename TObject = Object>
+ TVM_FFI_INLINE static TObject* RawObjectPtrFromAnyAfterCheck(const Any& ref)
{
+ return ObjectUnsafe::RawObjectPtrFromUnowned<TObject>(ref.data_.v_obj);
+ }
+
+ template <typename TObject = Object>
+ TVM_FFI_INLINE static TObject* RawObjectPtrFromAnyViewAfterCheck(const
AnyView& ref) {
+ return ObjectUnsafe::RawObjectPtrFromUnowned<TObject>(ref.data_.v_obj);
}
TVM_FFI_INLINE static const TVMFFIAny* TVMFFIAnyPtrFromAny(const Any& ref) {
diff --git a/include/tvm/ffi/container/array.h
b/include/tvm/ffi/container/array.h
index c2da576c..f5421df5 100644
--- a/include/tvm/ffi/container/array.h
+++ b/include/tvm/ffi/container/array.h
@@ -127,24 +127,6 @@ class ArrayObj : public SeqBaseObj {
return p;
}
- /*!
- * \brief Inplace-initialize the elements starting idx from [first, last)
- * \param idx The starting point
- * \param first Begin of iterator
- * \param last End of iterator
- * \tparam IterType The type of iterator
- * \return Self
- */
- template <typename IterType>
- ArrayObj* InitRange(int64_t idx, IterType first, IterType last) {
- Any* itr = MutableBegin() + idx;
- for (; first != last; ++first) {
- Any ref = *first;
- new (itr++) Any(std::move(ref));
- }
- return this;
- }
-
/*! \brief Initial size of ArrayObj */
static constexpr int64_t kInitSize = 4;
diff --git a/include/tvm/ffi/container/seq_base.h
b/include/tvm/ffi/container/seq_base.h
index 4d5a3618..43191734 100644
--- a/include/tvm/ffi/container/seq_base.h
+++ b/include/tvm/ffi/container/seq_base.h
@@ -128,9 +128,38 @@ class SeqBaseObj : public Object, protected TVMFFISeqCell {
if (i < 0 || i >= TVMFFISeqCell::size) {
TVM_FFI_THROW(IndexError) << "Index " << i << " out of bounds " <<
TVMFFISeqCell::size;
}
+ SetItemAfterCheck(i, std::move(item));
+ }
+
+ /*!
+ * \brief Set i-th element after the caller has established that i is in
bounds.
+ * \param i The index.
+ * \param item The value to be set.
+ */
+ TVM_FFI_INLINE void SetItemAfterCheck(int64_t i, Any item) {
static_cast<Any*>(data)[i] = std::move(item);
}
+ /*!
+ * \brief Inplace-initialize the elements starting at idx from [first, last).
+ * \param idx The starting point.
+ * \param first Begin of iterator.
+ * \param last End of iterator.
+ * \tparam IterType The iterator type.
+ * \return Self.
+ */
+ template <typename IterType>
+ SeqBaseObj* InitRange(int64_t idx, IterType first, IterType last) {
+ TVM_FFI_DCHECK_GE(idx, 0);
+ TVM_FFI_DCHECK_LE(idx + std::distance(first, last), TVMFFISeqCell::size);
+ Any* itr = MutableBegin() + idx;
+ for (; first != last; ++first) {
+ Any ref = *first;
+ new (itr++) Any(std::move(ref));
+ }
+ return this;
+ }
+
/*! \brief Remove the last element */
void pop_back() {
if (TVMFFISeqCell::size == 0) {
diff --git a/include/tvm/ffi/container/variant.h
b/include/tvm/ffi/container/variant.h
index 2a027f56..1823da42 100644
--- a/include/tvm/ffi/container/variant.h
+++ b/include/tvm/ffi/container/variant.h
@@ -171,7 +171,7 @@ class Variant {
static_assert(all_object_v,
"All types used in Variant<...> must be derived from
ObjectRef "
"to enable ObjectPtrHash/ObjectPtrEqual");
- return details::AnyUnsafe::ObjectPtrFromAnyAfterCheck(this->data_);
+ return details::AnyUnsafe::RawObjectPtrFromAnyAfterCheck(this->data_);
}
TVM_FFI_INLINE AnyView ToAnyView() const { return data_.operator AnyView(); }
TVM_FFI_INLINE Any MoveToAny() && { return std::move(data_); }
diff --git a/include/tvm/ffi/extra/structural_mutate.h
b/include/tvm/ffi/extra/structural_mutate.h
index 56eff51c..d2d04aa1 100644
--- a/include/tvm/ffi/extra/structural_mutate.h
+++ b/include/tvm/ffi/extra/structural_mutate.h
@@ -57,6 +57,11 @@ class StructuralMutatorObj;
* \param mutator The active structural mutator.
* \param value The borrowed value to mutate.
* \return Raw ``TVMFFIAny`` containing the mutated value or an Error.
+ *
+ * \note The hook is exception-free like \ref FStructuralVisit. Representable
failures must be
+ * returned as an Error. Hook implementations should use non-throwing
accessors when the
+ * engine's type dispatch has already established the type; allocation
failure and violated
+ * container invariants remain fatal.
*/
using FStructuralMutate = TVMFFIAny (*)(StructuralMutatorObj* mutator, AnyView
value) noexcept;
diff --git a/src/ffi/extra/structural_mutate.cc
b/src/ffi/extra/structural_mutate.cc
index b1b74c0f..b8e5a072 100644
--- a/src/ffi/extra/structural_mutate.cc
+++ b/src/ffi/extra/structural_mutate.cc
@@ -101,33 +101,28 @@ Expected<Any> StructuralMapExpected(
template <typename SeqObj>
Expected<Any> MutateSeqContainerExpected(StructuralMutatorObj* mutator,
AnyView value,
const SeqObj* self) noexcept {
- try {
- int64_t size = static_cast<int64_t>(self->size());
- ObjectPtr<SeqObj> output = nullptr;
+ int64_t size = static_cast<int64_t>(self->size());
+ const Any* items = self->begin();
+ ObjectPtr<SeqObj> output = nullptr;
- for (int64_t i = 0; i < size; ++i) {
- const Any& item = self->at(i);
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, mapped_value,
mutator->MutateExpected(item), self);
+ for (int64_t i = 0; i < size; ++i) {
+ const Any& item = items[i];
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, mapped_value,
mutator->MutateExpected(item), self);
- if (output == nullptr) {
- if (item.same_as(mapped_value)) {
- continue;
- }
- output = SeqObj::CreateRepeated(size, Any());
- for (int64_t j = 0; j < i; ++j) {
- output->SetItem(j, self->at(j));
- }
+ if (output == nullptr) {
+ if (item.same_as(mapped_value)) {
+ continue;
}
- output->SetItem(i, std::move(mapped_value));
+ output = SeqObj::CreateRepeated(size, Any());
+ output->InitRange(0, items, items + i);
}
+ output->SetItemAfterCheck(i, std::move(mapped_value));
+ }
- if (output == nullptr) {
- return Any(value);
- }
- return Any(ObjectRef(std::move(output)));
- } catch (const Error& err) {
- return Unexpected(err);
+ if (output == nullptr) {
+ return Any(value);
}
+ return Any(ObjectRef(std::move(output)));
}
/*!
@@ -142,20 +137,16 @@ Expected<Any>
MutateSeqContainerExpected(StructuralMutatorObj* mutator, AnyView
template <typename SeqObj>
Expected<Any> MaybeInplaceMutateSeqContainerExpected(StructuralMutatorObj*
mutator, AnyView value,
SeqObj* self) noexcept {
- try {
- for (int64_t i = 0; i < static_cast<int64_t>(self->size()); ++i) {
- const Any& item = self->at(i);
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, mapped_value,
-
mutator->MaybeInplaceMutateIfUniqueExpected(item), self);
+ for (int64_t i = 0; i < static_cast<int64_t>(self->size()); ++i) {
+ const Any& item = self->begin()[i];
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, mapped_value,
+
mutator->MaybeInplaceMutateIfUniqueExpected(item), self);
- if (!item.same_as(mapped_value)) {
- self->SetItem(i, std::move(mapped_value));
- }
+ if (!item.same_as(mapped_value)) {
+ self->SetItemAfterCheck(i, std::move(mapped_value));
}
- return Any(value);
- } catch (const Error& err) {
- return Unexpected(err);
}
+ return Any(value);
}
/*!
@@ -170,38 +161,34 @@ Expected<Any>
MaybeInplaceMutateSeqContainerExpected(StructuralMutatorObj* mutat
template <typename MapObjType>
Expected<Any> MutateMapValuesExpected(StructuralMutatorObj* mutator, AnyView
value,
const MapObjType* self) noexcept {
- try {
- ObjectPtr<Object> output = nullptr;
- MapBaseObj::iterator output_it;
- size_t index = 0;
+ ObjectPtr<Object> output = nullptr;
+ MapBaseObj::iterator output_it;
+ size_t index = 0;
- for (auto source_it = self->begin(); source_it != self->end();
++source_it, ++index) {
- const Any& old_value = source_it->second;
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, new_value,
mutator->MutateExpected(old_value), self);
- bool changed = !old_value.same_as(new_value);
- if (output == nullptr) {
- if (!changed) {
- continue;
- }
- output = MapObjType::ShallowCopy(self);
- output_it = static_cast<MapBaseObj*>(output.get())->begin();
- for (size_t i = 0; i < index; ++i) {
- ++output_it;
- }
+ for (auto source_it = self->begin(); source_it != self->end(); ++source_it,
++index) {
+ const Any& old_value = source_it->second;
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, new_value,
mutator->MutateExpected(old_value), self);
+ bool changed = !old_value.same_as(new_value);
+ if (output == nullptr) {
+ if (!changed) {
+ continue;
}
- if (changed) {
- output_it->second = std::move(new_value);
+ output = MapObjType::ShallowCopy(self);
+ output_it = static_cast<MapBaseObj*>(output.get())->begin();
+ for (size_t i = 0; i < index; ++i) {
+ ++output_it;
}
- ++output_it;
}
-
- if (output == nullptr) {
- return Any(value);
+ if (changed) {
+ output_it->second = std::move(new_value);
}
- return Any(ObjectRef(std::move(output)));
- } catch (const Error& err) {
- return Unexpected(err);
+ ++output_it;
+ }
+
+ if (output == nullptr) {
+ return Any(value);
}
+ return Any(ObjectRef(std::move(output)));
}
/*!
@@ -216,20 +203,16 @@ Expected<Any>
MutateMapValuesExpected(StructuralMutatorObj* mutator, AnyView val
template <typename MapObjType>
Expected<Any> MaybeInplaceMutateMapValuesExpected(StructuralMutatorObj*
mutator, AnyView value,
MapObjType* self) noexcept {
- try {
- for (auto it = self->begin(); it != self->end(); ++it) {
- const Any& old_value = it->second;
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- Any, new_value,
mutator->MaybeInplaceMutateIfUniqueExpected(old_value), self);
+ for (auto it = self->begin(); it != self->end(); ++it) {
+ const Any& old_value = it->second;
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(Any, new_value,
+
mutator->MaybeInplaceMutateIfUniqueExpected(old_value), self);
- if (!old_value.same_as(new_value)) {
- it->second = std::move(new_value);
- }
+ if (!old_value.same_as(new_value)) {
+ it->second = std::move(new_value);
}
- return Any(value);
- } catch (const Error& err) {
- return Unexpected(err);
}
+ return Any(value);
}
/*! \brief Identity structural mutation hook for immutable String and Bytes
leaves. */
@@ -240,53 +223,57 @@ TVMFFIAny MutateImmutableLeaf(StructuralMutatorObj*,
AnyView value) noexcept {
/*! \brief Structural mutation hook for ArrayObj. */
TVMFFIAny MutateArray(StructuralMutatorObj* mutator, AnyView value) noexcept {
- Expected<Any> result = MutateSeqContainerExpected(mutator, value,
value.cast<const ArrayObj*>());
+ Expected<Any> result = MutateSeqContainerExpected(
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const ArrayObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Maybe-in-place structural mutation hook for ArrayObj. */
TVMFFIAny MaybeInplaceMutateArray(StructuralMutatorObj* mutator, AnyView
value) noexcept {
Expected<Any> result = MaybeInplaceMutateSeqContainerExpected(
- mutator, value, const_cast<ArrayObj*>(value.cast<const ArrayObj*>()));
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<ArrayObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Structural mutation hook for ListObj. */
TVMFFIAny MutateList(StructuralMutatorObj* mutator, AnyView value) noexcept {
- Expected<Any> result = MutateSeqContainerExpected(mutator, value,
value.cast<const ListObj*>());
+ Expected<Any> result = MutateSeqContainerExpected(
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const ListObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Maybe-in-place structural mutation hook for ListObj. */
TVMFFIAny MaybeInplaceMutateList(StructuralMutatorObj* mutator, AnyView value)
noexcept {
Expected<Any> result = MaybeInplaceMutateSeqContainerExpected(
- mutator, value, const_cast<ListObj*>(value.cast<const ListObj*>()));
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<ListObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Structural mutation hook for MapObj. */
TVMFFIAny MutateMap(StructuralMutatorObj* mutator, AnyView value) noexcept {
- Expected<Any> result = MutateMapValuesExpected(mutator, value,
value.cast<const MapObj*>());
+ Expected<Any> result = MutateMapValuesExpected(
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const MapObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Maybe-in-place structural mutation hook for MapObj. */
TVMFFIAny MaybeInplaceMutateMap(StructuralMutatorObj* mutator, AnyView value)
noexcept {
Expected<Any> result = MaybeInplaceMutateMapValuesExpected(
- mutator, value, const_cast<MapObj*>(value.cast<const MapObj*>()));
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<MapObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Structural mutation hook for DictObj. */
TVMFFIAny MutateDict(StructuralMutatorObj* mutator, AnyView value) noexcept {
- Expected<Any> result = MutateMapValuesExpected(mutator, value,
value.cast<const DictObj*>());
+ Expected<Any> result = MutateMapValuesExpected(
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const DictObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
/*! \brief Maybe-in-place structural mutation hook for DictObj. */
TVMFFIAny MaybeInplaceMutateDict(StructuralMutatorObj* mutator, AnyView value)
noexcept {
Expected<Any> result = MaybeInplaceMutateMapValuesExpected(
- mutator, value, const_cast<DictObj*>(value.cast<const DictObj*>()));
+ mutator, value,
details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<DictObj>(value));
return ExpectedUnsafe::MoveToTVMFFIAny(std::move(result));
}
} // namespace details
diff --git a/tests/cpp/extra/test_structural_mutate.cc
b/tests/cpp/extra/test_structural_mutate.cc
index 8d8f87a9..c5ae8599 100644
--- a/tests/cpp/extra/test_structural_mutate.cc
+++ b/tests/cpp/extra/test_structural_mutate.cc
@@ -18,6 +18,8 @@
*/
#include <gtest/gtest.h>
#include <tvm/ffi/container/array.h>
+#include <tvm/ffi/container/dict.h>
+#include <tvm/ffi/container/list.h>
#include <tvm/ffi/container/map.h>
#include <tvm/ffi/extra/structural_mutate.h>
#include <tvm/ffi/string.h>
@@ -265,6 +267,41 @@ TEST(StructuralMap, AcceptsExpectedCallbackReturnTypes) {
"thrown error");
}
+template <WalkOrder order>
+void CheckContainerCallbackErrorsStayExpected() {
+ auto check_error = [](AnyView root, auto callback, const char*
expected_message) {
+ Expected<Any> result = StructuralMapExpected<order>(root,
std::move(callback));
+ ASSERT_TRUE(result.is_err());
+ EXPECT_EQ(result.error().kind(), "ValueError");
+ EXPECT_EQ(result.error().message(), expected_message);
+ };
+ auto raise = [](const TVar&) -> Expected<Any> { throw Error("ValueError",
"raised", ""); };
+ auto returns_error = [](const TVar&) -> Expected<Any> {
+ return Unexpected(Error("ValueError", "returned", ""));
+ };
+
+ AnyArray array_root{Any(TVar("n"))};
+ check_error(array_root, raise, "raised");
+ check_error(array_root, returns_error, "returned");
+
+ List<Any> list_root{Any(TVar("n"))};
+ check_error(list_root, raise, "raised");
+ check_error(list_root, returns_error, "returned");
+
+ StringMap map_root{{"value", TVar("n")}};
+ check_error(map_root, raise, "raised");
+ check_error(map_root, returns_error, "returned");
+
+ Dict<Any, Any> dict_root{{String("value"), Any(TVar("n"))}};
+ check_error(dict_root, raise, "raised");
+ check_error(dict_root, returns_error, "returned");
+}
+
+TEST(StructuralMap, ContainerCallbackErrorsStayExpected) {
+ CheckContainerCallbackErrorsStayExpected<WalkOrder::kPreOrder>();
+ CheckContainerCallbackErrorsStayExpected<WalkOrder::kPostOrder>();
+}
+
template <WalkOrder order>
void CheckRepeatedVarRemap() {
TVar var("n");
diff --git a/tests/cpp/test_any.cc b/tests/cpp/test_any.cc
index a82ad8b6..b6d6bc8b 100644
--- a/tests/cpp/test_any.cc
+++ b/tests/cpp/test_any.cc
@@ -275,6 +275,11 @@ TEST(Any, Object) {
AnyView view2 = any1;
EXPECT_EQ(v1.use_count(), 2);
+ EXPECT_EQ(details::AnyUnsafe::RawObjectPtrFromAnyAfterCheck(any1), v1.get());
+ EXPECT_EQ(details::AnyUnsafe::RawObjectPtrFromAnyAfterCheck<TIntObj>(any1),
v1.get());
+ EXPECT_EQ(details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(view2),
v1.get());
+
EXPECT_EQ(details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<TIntObj>(view2),
v1.get());
+
// convert to weak raw object ptr
const TIntObj* v1_ptr = view2.cast<const TIntObj*>();
EXPECT_EQ(v1.use_count(), 2);