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.
 // ---------------------------------------------------------------------------

Reply via email to