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 3e8b2a86 fix(structural): update StructuralVisitor def-region
propagation (#681)
3e8b2a86 is described below
commit 3e8b2a860936fdacc85d65ec07bcebe150d23b5f
Author: Kathryn (Jinqi) Chen <[email protected]>
AuthorDate: Tue Jul 21 20:45:15 2026 -0700
fix(structural): update StructuralVisitor def-region propagation (#681)
## PR summary
### Summary
This PR corrects how `StructuralVisitor` propagates definition-region
state while traversing reflected fields. When a non-recursive definition
is applied to the defined FreeVar identity, its children should be
visited with a "use" instead of "def".
Also make `StructuralVisitor` constructor a protected member with vtable
as input.
---
include/tvm/ffi/extra/structural_visit.h | 147 +++++++++++--------------------
python/tvm_ffi/_ffi_api.py | 2 -
python/tvm_ffi/structural.py | 15 +---
src/ffi/extra/structural_visit.cc | 4 +-
tests/cpp/extra/test_structural_visit.cc | 56 +++++++++---
5 files changed, 98 insertions(+), 126 deletions(-)
diff --git a/include/tvm/ffi/extra/structural_visit.h
b/include/tvm/ffi/extra/structural_visit.h
index 7a147cfa..e4ea5fad 100644
--- a/include/tvm/ffi/extra/structural_visit.h
+++ b/include/tvm/ffi/extra/structural_visit.h
@@ -138,9 +138,6 @@ struct StructuralVisitorVTable {
*/
class StructuralVisitorObj : public Object {
public:
- /*! \brief Construct the default structural visitor. */
- StructuralVisitorObj() : StructuralVisitorObj(VTable()) {}
-
/*!
* \brief Visit a value, dispatching through this visitor's vtable.
*
@@ -164,16 +161,41 @@ class StructuralVisitorObj : public Object {
}
/*!
- * \brief Visit using the structural visit behavior registered by
kStructuralVisit for each type,
- * or reflected structural fields when no custom behavior is registered.
+ * \brief Return the current def-region context.
+ * \return The active def-region kind.
*/
- TVM_FFI_INLINE Optional<VisitInterrupt> DefaultVisit(AnyView value) {
- return DefaultVisitExpected(value).value();
+ TVM_FFI_INLINE TVMFFIDefRegionKind def_region_kind() const { return
def_region_mode_; }
+
+ /*!
+ * \brief Temporarily switch the def-region context while invoking \p
callback.
+ *
+ * \param kind The def-region kind to set during the callback.
+ * \param callback A nullary callable that performs recursive visiting.
+ * \return The value returned by \p callback.
+ */
+ template <typename Callback>
+ TVM_FFI_INLINE auto WithDefRegionKind(TVMFFIDefRegionKind kind, Callback&&
callback) {
+ class Scope {
+ public:
+ Scope(StructuralVisitorObj* visitor, TVMFFIDefRegionKind kind)
+ : visitor_(visitor), old_kind_(visitor->def_region_mode_) {
+ visitor_->def_region_mode_ = kind;
+ }
+ ~Scope() { visitor_->def_region_mode_ = old_kind_; }
+ Scope(const Scope&) = delete;
+ Scope& operator=(const Scope&) = delete;
+
+ private:
+ StructuralVisitorObj* visitor_;
+ TVMFFIDefRegionKind old_kind_;
+ };
+ Scope scope(this, kind);
+ return std::forward<Callback>(callback)();
}
/*!
- * \brief Visit using the registered structural visit behavior by
kStructuralVisit, propagating
- * errors by Expected.
+ * \brief Visit using the structural visit behavior registered by
kStructuralVisit for each type,
+ * or reflected structural fields when no custom behavior is registered.
*
* \param value The value to visit.
* \return Expected interrupt state. An error means traversal failed.
@@ -209,43 +231,6 @@ class StructuralVisitorObj : public Object {
return details::VisitReflectedFieldsExpected(this, value.cast<const
Object*>());
}
- /*!
- * \brief Return the current def-region context.
- * \return The active def-region kind.
- */
- TVM_FFI_INLINE TVMFFIDefRegionKind def_region_kind() const { return
def_region_mode_; }
-
- /*!
- * \brief Temporarily switch the def-region context while invoking \p
callback.
- *
- * This helper scopes updates to the traversal state used by def/use-region
- * aware visitors. The previous state is restored when the callback returns
- * or throws.
- *
- * \param kind The def-region kind to set during the callback.
- * \param callback A nullary callable that performs recursive visiting.
- * \return The value returned by \p callback.
- */
- template <typename Callback>
- TVM_FFI_INLINE auto WithDefRegionKind(TVMFFIDefRegionKind kind, Callback&&
callback) {
- class Scope {
- public:
- Scope(StructuralVisitorObj* visitor, TVMFFIDefRegionKind kind)
- : visitor_(visitor), old_kind_(visitor->def_region_mode_) {
- visitor_->def_region_mode_ = kind;
- }
- ~Scope() { visitor_->def_region_mode_ = old_kind_; }
- Scope(const Scope&) = delete;
- Scope& operator=(const Scope&) = delete;
-
- private:
- StructuralVisitorObj* visitor_;
- TVMFFIDefRegionKind old_kind_;
- };
- Scope scope(this, kind);
- return std::forward<Callback>(callback)();
- }
-
/// \cond Doxygen_Suppress
static constexpr const bool _type_mutable = true;
TVM_FFI_DECLARE_OBJECT_INFO("ffi.StructuralVisitor", StructuralVisitorObj,
Object);
@@ -253,12 +238,8 @@ class StructuralVisitorObj : public Object {
protected:
/*!
- * \brief Construct a structural visitor subclass with a custom dispatch
vtable.
- *
- * \param vtable The non-null dispatch table for this visitor.
- *
- * \note This constructor is for internal subclasses. The vtable and its
- * ``visit`` callback must be valid for the lifetime of the visitor.
+ * \brief Construct a structural visitor from an immutable dispatch vtable.
+ * \param vtable The non-null dispatch table for this visitor. It must
outlive this object.
*/
explicit StructuralVisitorObj(const StructuralVisitorVTable* vtable) :
vtable_(vtable) {}
@@ -276,33 +257,6 @@ class StructuralVisitorObj : public Object {
* to scope temporary changes.
*/
TVMFFIDefRegionKind def_region_mode_ = kTVMFFIDefRegionKindNone;
-
- private:
- /*!
- * \brief Return the vtable used by the default visitor.
- * \return Pointer to the static structural visitor vtable.
- */
- static const StructuralVisitorVTable* VTable() {
- static const StructuralVisitorVTable
vtable{&StructuralVisitorObj::DispatchVisit};
- return &vtable;
- }
-
- /*!
- * \brief Dispatch from the vtable to the default visitor.
- * \param visitor The structural visitor object.
- * \param value The value to visit.
- * \return Interrupt state, or an error if traversal failed.
- */
- static TVMFFIAny DispatchVisit(StructuralVisitorObj* visitor, AnyView value)
noexcept {
- auto interrupt = visitor->DefaultVisitExpected(value);
- if (TVM_FFI_PREDICT_FALSE(interrupt.type_index() ==
TypeIndex::kTVMFFIError)) {
- if (value.type_index() >= TypeIndex::kTVMFFIStaticObjectBegin) {
- Error err = interrupt.error();
- details::UpdateVisitErrorContext(err, value.cast<ObjectRef>());
- }
- }
- return details::ExpectedUnsafe::MoveToTVMFFIAny(std::move(interrupt));
- }
};
/*!
@@ -312,10 +266,6 @@ class StructuralVisitorObj : public Object {
*/
class StructuralVisitor : public ObjectRef {
public:
- /*!
- * \brief Construct the default structural visitor.
- */
- StructuralVisitor() : ObjectRef(make_object<StructuralVisitorObj>()) {}
/*!
* \brief Construct from an existing object pointer.
* \param n The object pointer to wrap.
@@ -355,6 +305,13 @@ TVM_FFI_INLINE static Expected<Optional<VisitInterrupt>>
VisitReflectedFieldsExp
StructuralVisitorObj* visitor, const Object* obj) noexcept {
int32_t type_index = obj->type_index();
const TVMFFITypeInfo* type_info = TVMFFIGetTypeInfo(type_index);
+ // A non-recursive definition applies to a FreeVar itself, but not to its
children. All other
+ // inherited modes propagate until an explicit field annotation overrides
them.
+ TVMFFIDefRegionKind inherited_kind = visitor->def_region_kind();
+ if (inherited_kind == kTVMFFIDefRegionKindNonRecursive &&
type_info->metadata != nullptr &&
+ type_info->metadata->structural_eq_hash_kind ==
kTVMFFISEqHashKindFreeVar) {
+ inherited_kind = kTVMFFIDefRegionKindNone;
+ }
Expected<Optional<VisitInterrupt>> result =
Optional<VisitInterrupt>(std::nullopt);
reflection::ForEachFieldInfoWithEarlyStop(
@@ -372,19 +329,15 @@ TVM_FFI_INLINE static Expected<Optional<VisitInterrupt>>
VisitReflectedFieldsExp
return true;
}
- TVMFFIDefRegionKind kind = kTVMFFIDefRegionKindNone;
+ TVMFFIDefRegionKind kind = inherited_kind;
if (field_info->flags & kTVMFFIFieldFlagBitMaskSEqHashDefNonRecursive)
{
kind = kTVMFFIDefRegionKindNonRecursive;
} else if (field_info->flags &
kTVMFFIFieldFlagBitMaskSEqHashDefRecursive) {
kind = kTVMFFIDefRegionKindRecursive;
}
- if (kind != kTVMFFIDefRegionKindNone) {
- result = visitor->WithDefRegionKind(
- kind, [&]() { return visitor->VisitExpected(field_value); });
- } else {
- result = visitor->VisitExpected(field_value);
- }
+ result =
+ visitor->WithDefRegionKind(kind, [&]() { return
visitor->VisitExpected(field_value); });
return StructuralVisitNeedEarlyReturn(result);
});
return result;
@@ -478,12 +431,12 @@ struct TypeTraits<WalkResult> : public
TypeTraits<WalkResult::Storage> {
/// \endcond
/*!
- * \brief Callback order for \ref tvm::ffi::StructuralWalk.
+ * \brief Callback order for recursive structural traversal.
*/
enum class WalkOrder : int32_t {
- /*! \brief Invoke the callback before visiting children. */
+ /*! \brief Invoke the callback before traversing children. */
kPreOrder = 0,
- /*! \brief Invoke the callback after visiting children. */
+ /*! \brief Invoke the callback after traversing children. */
kPostOrder = 1,
};
@@ -541,11 +494,13 @@ class StructuralWalkCallbackVisitorObj : public
StructuralVisitorObj {
private:
/*!
- * \brief Return the vtable used by this visitor.
- * \return Pointer to the static structural visitor vtable.
+ * \brief Return the shared callback-aware visitor vtable.
+ * \return Pointer to the immutable visitor vtable for this specialization.
*/
static const StructuralVisitorVTable* VTable() {
- static const StructuralVisitorVTable
vtable{&StructuralWalkCallbackVisitorObj::DispatchVisit};
+ static const StructuralVisitorVTable vtable{
+ &StructuralWalkCallbackVisitorObj::DispatchVisit,
+ };
return &vtable;
}
diff --git a/python/tvm_ffi/_ffi_api.py b/python/tvm_ffi/_ffi_api.py
index b1216d8b..38a54fff 100644
--- a/python/tvm_ffi/_ffi_api.py
+++ b/python/tvm_ffi/_ffi_api.py
@@ -111,7 +111,6 @@ if TYPE_CHECKING:
def StructuralHash(_0: Any, _1: bool, _2: bool, /) -> int: ...
def StructuralKey(_0: Any, /) -> _StructuralKey: ...
def StructuralKeyEqual(_0: Any, _1: Any, /) -> bool: ...
- def StructuralVisitor() -> _StructuralVisitor: ...
def StructuralVisitorDefRegionKind(_0: _StructuralVisitor, /) -> int: ...
def StructuralVisitorVisit(_0: _StructuralVisitor, _1: Any, /) ->
_VisitInterrupt | None: ...
def StructuralVisitorWithDefRegionKind(_0: _StructuralVisitor, _1: int,
_2: Callable[..., Any], /) -> Any: ...
@@ -203,7 +202,6 @@ __all__ = [
"StructuralHash",
"StructuralKey",
"StructuralKeyEqual",
- "StructuralVisitor",
"StructuralVisitorDefRegionKind",
"StructuralVisitorVisit",
"StructuralVisitorWithDefRegionKind",
diff --git a/python/tvm_ffi/structural.py b/python/tvm_ffi/structural.py
index ff584f5f..ed87dc7d 100644
--- a/python/tvm_ffi/structural.py
+++ b/python/tvm_ffi/structural.py
@@ -352,21 +352,10 @@ class StructuralVisitor(Object):
"""Low-level structural traversal visitor.
This class exposes the low-level visitor object used by structural
- traversal hooks.
+ traversal hooks. Instances are supplied by an active traversal and cannot
+ be constructed directly.
"""
- # tvm-ffi-stubgen(begin): object/ffi.StructuralVisitor
- # fmt: off
- if TYPE_CHECKING:
- def __init__(self) -> None: ...
- def __ffi_init__(self) -> None: ... # ty:
ignore[invalid-method-override]
- # fmt: on
- # tvm-ffi-stubgen(end)
-
- def __init__(self) -> None:
- """Create a default structural visitor."""
- self.__init_handle_by_constructor__(_ffi_api.StructuralVisitor)
-
def visit(self, value: Any) -> VisitInterrupt | None:
"""Low-level API to visit ``value`` using this visitor's dispatch
behavior.
diff --git a/src/ffi/extra/structural_visit.cc
b/src/ffi/extra/structural_visit.cc
index 34967f5e..913bb4e4 100644
--- a/src/ffi/extra/structural_visit.cc
+++ b/src/ffi/extra/structural_visit.cc
@@ -135,11 +135,9 @@ TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::ObjectDef<VisitInterruptObj>().def_ro("value",
&VisitInterruptObj::value,
refl::default_value(nullptr));
- refl::ObjectDef<StructuralVisitorObj>().def(
- refl::init<>(), "Constructor that creates a default structural visitor");
+ refl::ObjectDef<StructuralVisitorObj>(); // NOLINT(bugprone-unused-raii)
refl::GlobalDef()
.def("ffi.VisitInterrupt", [](Any value) { return
VisitInterrupt(std::move(value)); })
- .def("ffi.StructuralVisitor", []() { return StructuralVisitor(); })
.def_method("ffi.StructuralVisitorVisit", &StructuralVisitorObj::Visit)
.def_method("ffi.StructuralVisitorDefRegionKind",
&StructuralVisitorObj::def_region_kind)
.def_method(
diff --git a/tests/cpp/extra/test_structural_visit.cc
b/tests/cpp/extra/test_structural_visit.cc
index 0c65a298..6bff25ad 100644
--- a/tests/cpp/extra/test_structural_visit.cc
+++ b/tests/cpp/extra/test_structural_visit.cc
@@ -47,7 +47,9 @@ class TestVisitorObj : public StructuralVisitorObj {
private:
static const StructuralVisitorVTable* VTable() {
- static const StructuralVisitorVTable
vtable{&TestVisitorObj::DispatchVisit};
+ static const StructuralVisitorVTable vtable{
+ &TestVisitorObj::DispatchVisit,
+ };
return &vtable;
}
@@ -126,9 +128,10 @@ TEST(StructuralVisitor, TraversesPair) {
ObjectRef root = TPair(lhs, rhs);
StructuralVisitor visitor = MakeTestVisitor();
- Optional<VisitInterrupt> result = visitor->DefaultVisit(root);
+ Expected<Optional<VisitInterrupt>> result =
visitor->DefaultVisitExpected(root);
- EXPECT_FALSE(result.has_value());
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
ASSERT_EQ(AsTestVisitor(visitor)->visited.size(), 2U);
EXPECT_TRUE(AsTestVisitor(visitor)->visited[0].same_as(lhs));
EXPECT_TRUE(AsTestVisitor(visitor)->visited[1].same_as(rhs));
@@ -142,9 +145,10 @@ TEST(StructuralVisitor, TraversesFunction) {
ObjectRef root = TFunc(params, body, String("ignored function comment"));
StructuralVisitor visitor = MakeTestVisitor();
- Optional<VisitInterrupt> result = visitor->DefaultVisit(root);
+ Expected<Optional<VisitInterrupt>> result =
visitor->DefaultVisitExpected(root);
- EXPECT_FALSE(result.has_value());
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
TestVisitorObj* test_visitor = AsTestVisitor(visitor);
ASSERT_EQ(test_visitor->visited.size(), 4U);
EXPECT_TRUE(test_visitor->visited[0].same_as(params));
@@ -164,10 +168,11 @@ TEST(StructuralVisitor, StopsOnInterrupt) {
StructuralVisitor visitor = MakeTestVisitor();
SetInterrupt(visitor, lhs);
- Optional<VisitInterrupt> result = visitor->DefaultVisit(root);
+ Expected<Optional<VisitInterrupt>> result =
visitor->DefaultVisitExpected(root);
- ASSERT_TRUE(result.has_value());
- EXPECT_EQ(result.value()->value.cast<String>(), "stop");
+ ASSERT_TRUE(result.is_ok());
+ ASSERT_TRUE(result.value().has_value());
+ EXPECT_EQ(result.value().value()->value.cast<String>(), "stop");
ASSERT_EQ(AsTestVisitor(visitor)->visited.size(), 1U);
EXPECT_TRUE(AsTestVisitor(visitor)->visited[0].same_as(lhs));
}
@@ -178,9 +183,10 @@ TEST(StructuralVisitor, TraversesArray) {
Array<ObjectRef> root = {lhs, rhs};
StructuralVisitor visitor = MakeTestVisitor();
- Optional<VisitInterrupt> result = visitor->DefaultVisit(root);
+ Expected<Optional<VisitInterrupt>> result =
visitor->DefaultVisitExpected(root);
- EXPECT_FALSE(result.has_value());
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
ASSERT_EQ(AsTestVisitor(visitor)->visited.size(), 2U);
EXPECT_TRUE(AsTestVisitor(visitor)->visited[0].same_as(lhs));
EXPECT_TRUE(AsTestVisitor(visitor)->visited[1].same_as(rhs));
@@ -192,9 +198,10 @@ TEST(StructuralVisitor, TraversesMap) {
Map<Any, Any> root{{key, value}};
StructuralVisitor visitor = MakeTestVisitor();
- Optional<VisitInterrupt> result = visitor->DefaultVisit(root);
+ Expected<Optional<VisitInterrupt>> result =
visitor->DefaultVisitExpected(root);
- EXPECT_FALSE(result.has_value());
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
ASSERT_EQ(AsTestVisitor(visitor)->visited.size(), 2U);
EXPECT_TRUE(AsTestVisitor(visitor)->visited[0].same_as(key));
EXPECT_TRUE(AsTestVisitor(visitor)->visited[1].same_as(value));
@@ -249,6 +256,31 @@ TEST(StructuralVisitor, RestoresFuncDefRegion) {
EXPECT_EQ(test_visitor->def_region_kind(), kTVMFFIDefRegionKindNone);
}
+TEST(StructuralVisitor, ExplicitDefRegionsOverrideFreeVarFieldClamp) {
+ TVarWithDep recursive("recursive");
+ TVarWithDep non_recursive("non-recursive");
+ TDefHolder holder(recursive, non_recursive);
+ TVarWithDep root("outer", holder);
+ StructuralVisitor visitor = MakeTestVisitor();
+
+ Expected<Optional<VisitInterrupt>> result = visitor->WithDefRegionKind(
+ kTVMFFIDefRegionKindNonRecursive, [&]() { return
visitor->VisitExpected(root); });
+
+ ASSERT_TRUE(result.is_ok());
+ EXPECT_FALSE(result.value().has_value());
+ TestVisitorObj* test_visitor = AsTestVisitor(visitor);
+ ASSERT_EQ(test_visitor->visited.size(), 4U);
+ EXPECT_TRUE(test_visitor->visited[0].same_as(root));
+ EXPECT_EQ(test_visitor->modes[0], kTVMFFIDefRegionKindNonRecursive);
+ EXPECT_TRUE(test_visitor->visited[1].same_as(holder));
+ EXPECT_EQ(test_visitor->modes[1], kTVMFFIDefRegionKindNone);
+ EXPECT_TRUE(test_visitor->visited[2].same_as(recursive));
+ EXPECT_EQ(test_visitor->modes[2], kTVMFFIDefRegionKindRecursive);
+ EXPECT_TRUE(test_visitor->visited[3].same_as(non_recursive));
+ EXPECT_EQ(test_visitor->modes[3], kTVMFFIDefRegionKindNonRecursive);
+ EXPECT_EQ(test_visitor->def_region_kind(), kTVMFFIDefRegionKindNone);
+}
+
// ---------------------------------------------------------------------------
// StructuralWalk behavior.
// ---------------------------------------------------------------------------