From 91cee9a51100edfe0f20a61ca42f23f893d30999 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 00:39:49 +0000 Subject: [PATCH 01/10] Add structural hooks for remaining type nodes --- src/ir/type.cc | 140 ++++++++++++++++++++ src/relax/distributed/type.cc | 64 +++++++++ src/relax/ir/dependent_type.cc | 167 ++++++++++++++++++++++++ src/relax/ir/type.cc | 28 +++- tests/cpp/type_structural_hooks_test.cc | 106 +++++++++++++++ 5 files changed, 504 insertions(+), 1 deletion(-) create mode 100644 tests/cpp/type_structural_hooks_test.cc diff --git a/src/ir/type.cc b/src/ir/type.cc index d4a56cb4765e..7d2e1050ad16 100644 --- a/src/ir/type.cc +++ b/src/ir/type.cc @@ -87,6 +87,126 @@ TVMFFIAny PrimTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) n return ffi::Unchanged().CopyToTVMFFIAny(); } +TVMFFIAny PointerTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const PointerTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->element_type)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny PointerTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const PointerTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_element_type, + mutator->MutateExpected(self->element_type)); + if (mapped_element_type.UnchangedOrSameAs(self->element_type)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->element_type = + std::move(mapped_element_type).ValueOrUnchanged(std::move(copy->element_type)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny PointerTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + PointerTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr, mapped_element_type, + mutator->MaybeInplaceMutateIfUniqueExpected(self->element_type)); + if (!mapped_element_type.UnchangedOrSameAs(self->element_type)) { + self->element_type = std::move(mapped_element_type).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny FuncTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const FuncTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->arg_types)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ret_type)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny FuncTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const FuncTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_arg_types, + mutator->MutateExpected(self->arg_types)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_type, + mutator->MutateExpected(self->ret_type)); + if (mapped_arg_types.UnchangedOrSameAs(self->arg_types) && + mapped_ret_type.UnchangedOrSameAs(self->ret_type)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->arg_types = std::move(mapped_arg_types).ValueOrUnchanged(std::move(copy->arg_types)); + copy->ret_type = std::move(mapped_ret_type).ValueOrUnchanged(std::move(copy->ret_type)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny FuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + FuncTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_arg_types, + mutator->MaybeInplaceMutateIfUniqueExpected(self->arg_types)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_type, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_type)); + if (!mapped_arg_types.UnchangedOrSameAs(self->arg_types)) { + self->arg_types = std::move(mapped_arg_types).ValueUnchecked(); + } + if (!mapped_ret_type.UnchangedOrSameAs(self->ret_type)) { + self->ret_type = std::move(mapped_ret_type).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny TupleTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const TupleTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->fields)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny TupleTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const TupleTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_fields, + mutator->MutateExpected(self->fields)); + if (mapped_fields.UnchangedOrSameAs(self->fields)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->fields = std::move(mapped_fields).ValueOrUnchanged(std::move(copy->fields)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny TupleTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + TupleTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_fields, + mutator->MaybeInplaceMutateIfUniqueExpected(self->fields)); + if (!mapped_fields.UnchangedOrSameAs(self->fields)) { + self->fields = std::move(mapped_fields).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny TensorMapTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny TensorMapTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny TensorMapTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + } // namespace Type Type::Missing() { @@ -227,6 +347,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { refl::GlobalDef().def("ir.PointerType", [](Type element_type, ffi::String storage_scope = "") { return PointerType(element_type, storage_scope); }); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&PointerTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&PointerTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&PointerTypeMaybeInplaceMutate)); } TVM_FFI_STATIC_INIT_BLOCK() { @@ -235,6 +360,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { refl::GlobalDef().def("ir.FuncType", [](tvm::ffi::Array arg_types, Type ret_type) { return FuncType(arg_types, ret_type); }); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&FuncTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&FuncTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&FuncTypeMaybeInplaceMutate)); } TVM_FFI_STATIC_INIT_BLOCK() { @@ -245,6 +375,16 @@ TVM_FFI_STATIC_INIT_BLOCK() { .def("ir.TupleType", [](ffi::Array fields, Span span) { return TupleType(fields, span); }) .def("ir.TensorMapType", [](Span span) { return TensorMapType(span); }); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&TupleTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&TupleTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TupleTypeMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&TensorMapTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&TensorMapTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TensorMapTypeMaybeInplaceMutate)); } } // namespace tvm diff --git a/src/relax/distributed/type.cc b/src/relax/distributed/type.cc index 8fb164c2a86e..0a48faad2c72 100644 --- a/src/relax/distributed/type.cc +++ b/src/relax/distributed/type.cc @@ -22,16 +22,80 @@ * \brief Relax DTensor type. */ +#include +#include #include #include namespace tvm { namespace relax { namespace distributed { +namespace { + +TVMFFIAny DTensorTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const DTensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->device_mesh)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->placement)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->tensor_ty)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny DTensorTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const DTensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_device_mesh, + mutator->MutateExpected(self->device_mesh)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_placement, + mutator->MutateExpected(self->placement)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_tensor_ty, + mutator->MutateExpected(self->tensor_ty)); + if (mapped_device_mesh.UnchangedOrSameAs(self->device_mesh) && + mapped_placement.UnchangedOrSameAs(self->placement) && + mapped_tensor_ty.UnchangedOrSameAs(self->tensor_ty)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->device_mesh = std::move(mapped_device_mesh).ValueOrUnchanged(std::move(copy->device_mesh)); + copy->placement = std::move(mapped_placement).ValueOrUnchanged(std::move(copy->placement)); + copy->tensor_ty = std::move(mapped_tensor_ty).ValueOrUnchanged(std::move(copy->tensor_ty)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny DTensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + DTensorTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_device_mesh, + mutator->MaybeInplaceMutateIfUniqueExpected(self->device_mesh)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_placement, + mutator->MaybeInplaceMutateIfUniqueExpected(self->placement)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_tensor_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->tensor_ty)); + if (!mapped_device_mesh.UnchangedOrSameAs(self->device_mesh)) { + self->device_mesh = std::move(mapped_device_mesh).ValueUnchecked(); + } + if (!mapped_placement.UnchangedOrSameAs(self->placement)) { + self->placement = std::move(mapped_placement).ValueUnchecked(); + } + if (!mapped_tensor_ty.UnchangedOrSameAs(self->tensor_ty)) { + self->tensor_ty = std::move(mapped_tensor_ty).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; DTensorTypeNode::RegisterReflection(); PlacementNode::RegisterReflection(); PlacementSpecNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&DTensorTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&DTensorTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&DTensorTypeMaybeInplaceMutate)); } PlacementSpec PlacementSpec::Sharding(int axis) { diff --git a/src/relax/ir/dependent_type.cc b/src/relax/ir/dependent_type.cc index ff0bce1fb6d9..b12901fa045f 100644 --- a/src/relax/ir/dependent_type.cc +++ b/src/relax/ir/dependent_type.cc @@ -21,6 +21,8 @@ * \file src/relax/ir/dependent_type.cc * \brief Relax type nodes. */ +#include +#include #include #include #include @@ -30,12 +32,177 @@ namespace tvm { namespace relax { +namespace { + +TVMFFIAny AnyTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny AnyTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny AnyTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny ShapeTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const ShapeTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->values)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny ShapeTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const ShapeTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>>, + mapped_values, mutator->MutateExpected(self->values)); + if (mapped_values.UnchangedOrSameAs(self->values)) return ffi::Unchanged().CopyToTVMFFIAny(); + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->values = std::move(mapped_values).ValueOrUnchanged(std::move(copy->values)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny ShapeTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + ShapeTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>>, + mapped_values, + mutator->MaybeInplaceMutateIfUniqueExpected(self->values)); + if (!mapped_values.UnchangedOrSameAs(self->values)) { + self->values = std::move(mapped_values).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny TensorTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const TensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->shape)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->dtype)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->vdevice)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny TensorTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const TensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, + mutator->MutateExpected(self->shape)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_dtype, + mutator->MutateExpected(self->dtype)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_vdevice, + mutator->MutateExpected(self->vdevice)); + if (mapped_shape.UnchangedOrSameAs(self->shape) && mapped_dtype.UnchangedOrSameAs(self->dtype) && + mapped_vdevice.UnchangedOrSameAs(self->vdevice)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->shape = std::move(mapped_shape).ValueOrUnchanged(std::move(copy->shape)); + copy->dtype = std::move(mapped_dtype).ValueOrUnchanged(std::move(copy->dtype)); + copy->vdevice = std::move(mapped_vdevice).ValueOrUnchanged(std::move(copy->vdevice)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny TensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + TensorTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, + mutator->MaybeInplaceMutateIfUniqueExpected(self->shape)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_dtype, + mutator->MaybeInplaceMutateIfUniqueExpected(self->dtype)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_vdevice, + mutator->MaybeInplaceMutateIfUniqueExpected(self->vdevice)); + if (!mapped_shape.UnchangedOrSameAs(self->shape)) { + self->shape = std::move(mapped_shape).ValueUnchecked(); + } + if (!mapped_dtype.UnchangedOrSameAs(self->dtype)) { + self->dtype = std::move(mapped_dtype).ValueUnchecked(); + } + if (!mapped_vdevice.UnchangedOrSameAs(self->vdevice)) { + self->vdevice = std::move(mapped_vdevice).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny FuncTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const FuncTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( + kTVMFFIDefRegionKindPattern, [&]() { return visitor->VisitExpected(self->params); })); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ret)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny FuncTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const FuncTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr>>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, + [&]() { return mutator->MutateExpected(self->params); })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret, + mutator->MutateExpected(self->ret)); + if (mapped_params.UnchangedOrSameAs(self->params) && mapped_ret.UnchangedOrSameAs(self->ret)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->params = std::move(mapped_params).ValueOrUnchanged(std::move(copy->params)); + copy->ret = std::move(mapped_ret).ValueOrUnchanged(std::move(copy->ret)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny FuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + FuncTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr>>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() { + return mutator->MaybeInplaceMutateIfUniqueExpected(self->params); + })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ret)); + if (!mapped_params.UnchangedOrSameAs(self->params)) { + self->params = std::move(mapped_params).ValueUnchecked(); + } + if (!mapped_ret.UnchangedOrSameAs(self->ret)) { + self->ret = std::move(mapped_ret).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; AnyTypeNode::RegisterReflection(); ShapeTypeNode::RegisterReflection(); TensorTypeNode::RegisterReflection(); FuncTypeNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&AnyTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&AnyTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&AnyTypeMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&ShapeTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&ShapeTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&ShapeTypeMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&TensorTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&TensorTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TensorTypeMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&FuncTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&FuncTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&FuncTypeMaybeInplaceMutate)); refl::TypeAttrDef().def( "__subscript_expr_realize__", [](Expr value, diff --git a/src/relax/ir/type.cc b/src/relax/ir/type.cc index d6fa7ada9cc4..6664474ce978 100644 --- a/src/relax/ir/type.cc +++ b/src/relax/ir/type.cc @@ -21,6 +21,8 @@ * \file src/relax/ir/type.cc * \brief Relax type system. */ +#include +#include #include #include #include @@ -28,7 +30,31 @@ namespace tvm { namespace relax { -TVM_FFI_STATIC_INIT_BLOCK() { PackedFuncTypeNode::RegisterReflection(); } +namespace { + +TVMFFIAny PackedFuncTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny PackedFuncTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny PackedFuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + PackedFuncTypeNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&PackedFuncTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&PackedFuncTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&PackedFuncTypeMaybeInplaceMutate)); +} PackedFuncType::PackedFuncType(Span span) : Type(ffi::UnsafeInit{}) { ffi::ObjectPtr n = ffi::make_object(); diff --git a/tests/cpp/type_structural_hooks_test.cc b/tests/cpp/type_structural_hooks_test.cc new file mode 100644 index 000000000000..f3d3d600c3a6 --- /dev/null +++ b/tests/cpp/type_structural_hooks_test.cc @@ -0,0 +1,106 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +template +void ExpectStructuralHooks() { + namespace refl = tvm::ffi::reflection; + for (const char* attr_name : + {refl::type_attr::kStructuralVisit, refl::type_attr::kStructuralMutate, + refl::type_attr::kStructuralMaybeInplaceMutate}) { + refl::TypeAttrColumn column(attr_name); + EXPECT_EQ(column[TNode::RuntimeTypeIndex()].type_index(), tvm::ffi::TypeIndex::kTVMFFIOpaquePtr) + << TNode::_type_key << " is missing " << attr_name; + } +} + +TEST(TypeStructuralHooks, EveryConcreteOpenTypeHasExplicitHooks) { + using namespace tvm; + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); +} + +TEST(TypeStructuralHooks, StructuralMapDescendsThroughTypeFields) { + using namespace tvm; + Type input = TupleType( + {PointerType(PrimType::Float(32), "global"), relax::TensorType(PrimType::Float(32), 2)}); + Type mapped = ffi::StructuralMap( + input, + [](const PrimType& type) -> ffi::Expected> { + if (!type.MatchesElementType(DLDataTypeCode::kDLFloat, 32)) { + return ffi::Unchanged(); + } + return ffi::Any(PrimType::Float(64)); + }) + .cast(); + + const auto* tuple = mapped.as(); + ASSERT_NE(tuple, nullptr); + EXPECT_TRUE(tuple->fields[0].as()->element_type.as()->dtype == + PrimType::Float(64)->dtype); + EXPECT_TRUE(tuple->fields[1].as()->dtype.value()->dtype == + PrimType::Float(64)->dtype); +} + +TEST(TypeStructuralHooks, StructuralEqualAndHashStillUseAllReflectedFields) { + using namespace tvm; + ffi::StructuralEqual equal; + ffi::StructuralHash hash; + + PointerType global(PrimType::Float(32), "global"); + PointerType shared(PrimType::Float(32), "shared"); + EXPECT_FALSE(equal(global, shared)); + EXPECT_NE(hash(global), hash(shared)); + + relax::ShapeType rank_one(1); + relax::ShapeType rank_two(2); + EXPECT_FALSE(equal(rank_one, rank_two)); + EXPECT_NE(hash(rank_one), hash(rank_two)); + + relax::TensorType f32(PrimType::Float(32), 2); + relax::TensorType f64(PrimType::Float(64), 2); + EXPECT_FALSE(equal(f32, f64)); + EXPECT_NE(hash(f32), hash(f64)); + + relax::FuncType pure({}, relax::AnyType(), true); + relax::FuncType impure({}, relax::AnyType(), false); + EXPECT_FALSE(equal(pure, impure)); + EXPECT_NE(hash(pure), hash(impure)); +} + +} // namespace From bfe1d4783829f3705ed34d2dd5f0cc5eff380f2c Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 00:40:47 +0000 Subject: [PATCH 02/10] Add structural hooks for remaining Relax expressions --- src/relax/ir/expr.cc | 398 ++++++++++++++++++ tests/cpp/relax_expr_structural_hooks_test.cc | 107 +++++ 2 files changed, 505 insertions(+) create mode 100644 tests/cpp/relax_expr_structural_hooks_test.cc diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc index b21cdbad879c..ba03c090f4da 100644 --- a/src/relax/ir/expr.cc +++ b/src/relax/ir/expr.cc @@ -16,6 +16,8 @@ * specific language governing permissions and limitations * under the License. */ +#include +#include #include #include #include @@ -27,7 +29,350 @@ namespace tvm { namespace relax { +namespace { + +template +TVMFFIAny TypeOnlyExprVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const TNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +template +TVMFFIAny TypeOnlyExprMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const TNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MutateExpected(self->ty)); + if (mapped_ty.UnchangedOrSameAs(self->ty)) return ffi::Unchanged().CopyToTVMFFIAny(); + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->ty = std::move(mapped_ty).ValueOrUnchanged(std::move(copy->ty)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +template +TVMFFIAny TypeOnlyExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + TNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ty)); + if (!mapped_ty.UnchangedOrSameAs(self->ty)) { + self->ty = std::move(mapped_ty).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny ShapeExprVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const ShapeExprNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->values)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny ShapeExprMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const ShapeExprNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MutateExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_values, + mutator->MutateExpected(self->values)); + if (mapped_ty.UnchangedOrSameAs(self->ty) && mapped_values.UnchangedOrSameAs(self->values)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->ty = std::move(mapped_ty).ValueOrUnchanged(std::move(copy->ty)); + copy->values = std::move(mapped_values).ValueOrUnchanged(std::move(copy->values)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny ShapeExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + ShapeExprNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_values, + mutator->MaybeInplaceMutateIfUniqueExpected(self->values)); + if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty = std::move(mapped_ty).ValueUnchecked(); + if (!mapped_values.UnchangedOrSameAs(self->values)) { + self->values = std::move(mapped_values).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny DataflowVarVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const DataflowVarNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + if (!self->ty.as()) { + if (visitor->def_region_kind() == kTVMFFIDefRegionKindSimple) { + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( + kTVMFFIDefRegionKindNone, [&]() { return visitor->VisitExpected(self->ty); })); + } else { + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + } + } + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny DataflowVarMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const DataflowVarNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + ffi::Expected remap_result = mutator->VarRemapGetExpected(value); + TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(remap_result); + if (ffi::details::ExpectedUnsafe::GetData(remap_result).type_index() != + ffi::TypeIndex::kTVMFFINone) { + return ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(std::move(remap_result)); + } + if (mutator->def_region_kind() == kTVMFFIDefRegionKindNone) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::UnchangedOr result = ffi::Unchanged(); + ffi::Any mapped_value; + if (!self->ty.as()) { + ffi::Expected> mapped_ty_result = + mutator->def_region_kind() == kTVMFFIDefRegionKindSimple + ? mutator->WithDefRegionKind(kTVMFFIDefRegionKindNone, + [&]() { return mutator->MutateExpected(self->ty); }) + : mutator->MutateExpected(self->ty); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + std::move(mapped_ty_result)); + if (!mapped_ty.UnchangedOrSameAs(self->ty)) { + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->ty = std::move(mapped_ty).ValueUnchecked(); + mapped_value = ffi::Any(std::move(copy)); + result = mapped_value; + } + } + if (!result.IsUnchanged() || mutator->def_region_kind() == kTVMFFIDefRegionKindPattern) { + ffi::AnyView value_to_store = result.IsUnchanged() ? value : ffi::AnyView(mapped_value); + auto set_result = mutator->VarRemapSetExpected(value, value_to_store); + if (TVM_FFI_PREDICT_FALSE(set_result.is_err())) { + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(set_result).error())); + } + } + return ffi::details::UnchangedOrUnsafe::MoveToTVMFFIAny(std::move(result)); +} + +TVMFFIAny DataflowVarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + DataflowVarNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + ffi::Expected remap_result = mutator->VarRemapGetExpected(value); + TVM_FFI_S_MUTATE_MAYBE_EARLY_RETURN(remap_result); + if (ffi::details::ExpectedUnsafe::GetData(remap_result).type_index() != + ffi::TypeIndex::kTVMFFINone) { + return ffi::details::ExpectedUnsafe::MoveToTVMFFIAny(std::move(remap_result)); + } + if (mutator->def_region_kind() == kTVMFFIDefRegionKindNone) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::UnchangedOr result = ffi::Unchanged(); + ffi::Any mapped_value; + if (!self->ty.as()) { + ffi::Expected> mapped_ty_result = + mutator->def_region_kind() == kTVMFFIDefRegionKindSimple + ? mutator->WithDefRegionKind( + kTVMFFIDefRegionKindNone, + [&]() { return mutator->MaybeInplaceMutateIfUniqueExpected(self->ty); }) + : mutator->MaybeInplaceMutateIfUniqueExpected(self->ty); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + std::move(mapped_ty_result)); + if (!mapped_ty.UnchangedOrSameAs(self->ty)) { + self->ty = std::move(mapped_ty).ValueUnchecked(); + mapped_value = ffi::Any(self); + result = mapped_value; + } + } + if (!result.IsUnchanged() || mutator->def_region_kind() == kTVMFFIDefRegionKindPattern) { + ffi::AnyView value_to_store = result.IsUnchanged() ? value : ffi::AnyView(mapped_value); + auto set_result = mutator->VarRemapSetExpected(value, value_to_store); + if (TVM_FFI_PREDICT_FALSE(set_result.is_err())) { + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(set_result).error())); + } + } + return ffi::details::UnchangedOrUnsafe::MoveToTVMFFIAny(std::move(result)); +} + +TVMFFIAny SeqExprVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const SeqExprNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->blocks)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->body)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny SeqExprMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const SeqExprNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_blocks, + mutator->MutateExpected(self->blocks)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MutateExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MutateExpected(self->body)); + if (mapped_blocks.UnchangedOrSameAs(self->blocks) && mapped_ty.UnchangedOrSameAs(self->ty) && + mapped_body.UnchangedOrSameAs(self->body)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->blocks = std::move(mapped_blocks).ValueOrUnchanged(std::move(copy->blocks)); + copy->ty = std::move(mapped_ty).ValueOrUnchanged(std::move(copy->ty)); + copy->body = std::move(mapped_body).ValueOrUnchanged(std::move(copy->body)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny SeqExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + SeqExprNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_blocks, + mutator->MaybeInplaceMutateIfUniqueExpected(self->blocks)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MaybeInplaceMutateIfUniqueExpected(self->body)); + if (!mapped_blocks.UnchangedOrSameAs(self->blocks)) { + self->blocks = std::move(mapped_blocks).ValueUnchecked(); + } + if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty = std::move(mapped_ty).ValueUnchecked(); + if (!mapped_body.UnchangedOrSameAs(self->body)) { + self->body = std::move(mapped_body).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny IfVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const IfNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->cond)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->true_branch)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->false_branch)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny IfMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const IfNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MutateExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_cond, + mutator->MutateExpected(self->cond)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_true_branch, + mutator->MutateExpected(self->true_branch)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_false_branch, + mutator->MutateExpected(self->false_branch)); + if (mapped_ty.UnchangedOrSameAs(self->ty) && mapped_cond.UnchangedOrSameAs(self->cond) && + mapped_true_branch.UnchangedOrSameAs(self->true_branch) && + mapped_false_branch.UnchangedOrSameAs(self->false_branch)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->ty = std::move(mapped_ty).ValueOrUnchanged(std::move(copy->ty)); + copy->cond = std::move(mapped_cond).ValueOrUnchanged(std::move(copy->cond)); + copy->true_branch = std::move(mapped_true_branch).ValueOrUnchanged(std::move(copy->true_branch)); + copy->false_branch = + std::move(mapped_false_branch).ValueOrUnchanged(std::move(copy->false_branch)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny IfMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + IfNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_cond, + mutator->MaybeInplaceMutateIfUniqueExpected(self->cond)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_true_branch, + mutator->MaybeInplaceMutateIfUniqueExpected(self->true_branch)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr, mapped_false_branch, + mutator->MaybeInplaceMutateIfUniqueExpected(self->false_branch)); + if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty = std::move(mapped_ty).ValueUnchecked(); + if (!mapped_cond.UnchangedOrSameAs(self->cond)) { + self->cond = std::move(mapped_cond).ValueUnchecked(); + } + if (!mapped_true_branch.UnchangedOrSameAs(self->true_branch)) { + self->true_branch = std::move(mapped_true_branch).ValueUnchecked(); + } + if (!mapped_false_branch.UnchangedOrSameAs(self->false_branch)) { + self->false_branch = std::move(mapped_false_branch).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny FunctionVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + const FunctionNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( + kTVMFFIDefRegionKindPattern, [&]() { return visitor->VisitExpected(self->params); })); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ty)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->body)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ret_ty)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny FunctionMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + const FunctionNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() { + return mutator->MutateExpected(self->params); + })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MutateExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MutateExpected(self->body)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_ty, + mutator->MutateExpected(self->ret_ty)); + if (mapped_params.UnchangedOrSameAs(self->params) && mapped_ty.UnchangedOrSameAs(self->ty) && + mapped_body.UnchangedOrSameAs(self->body) && mapped_ret_ty.UnchangedOrSameAs(self->ret_ty)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->params = std::move(mapped_params).ValueOrUnchanged(std::move(copy->params)); + copy->ty = std::move(mapped_ty).ValueOrUnchanged(std::move(copy->ty)); + copy->body = std::move(mapped_body).ValueOrUnchanged(std::move(copy->body)); + copy->ret_ty = std::move(mapped_ret_ty).ValueOrUnchanged(std::move(copy->ret_ty)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny FunctionMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + FunctionNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() { + return mutator->MaybeInplaceMutateIfUniqueExpected(self->params); + })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ty)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MaybeInplaceMutateIfUniqueExpected(self->body)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_ty, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_ty)); + if (!mapped_params.UnchangedOrSameAs(self->params)) { + self->params = std::move(mapped_params).ValueUnchecked(); + } + if (!mapped_ty.UnchangedOrSameAs(self->ty)) self->ty = std::move(mapped_ty).ValueUnchecked(); + if (!mapped_body.UnchangedOrSameAs(self->body)) { + self->body = std::move(mapped_body).ValueUnchecked(); + } + if (!mapped_ret_ty.UnchangedOrSameAs(self->ret_ty)) { + self->ret_ty = std::move(mapped_ret_ty).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; ShapeExprNode::RegisterReflection(); BindingNode::RegisterReflection(); DataflowVarNode::RegisterReflection(); @@ -42,6 +387,59 @@ TVM_FFI_STATIC_INIT_BLOCK() { IfNode::RegisterReflection(); FunctionNode::RegisterReflection(); ExternFuncNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&ShapeExprVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&ShapeExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&ShapeExprMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&DataflowVarVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&DataflowVarMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&DataflowVarMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, + reinterpret_cast(&TypeOnlyExprVisit)) + .attr(refl::type_attr::kStructuralMutate, + reinterpret_cast(&TypeOnlyExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TypeOnlyExprMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, + reinterpret_cast(&TypeOnlyExprVisit)) + .attr(refl::type_attr::kStructuralMutate, + reinterpret_cast(&TypeOnlyExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TypeOnlyExprMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, + reinterpret_cast(&TypeOnlyExprVisit)) + .attr(refl::type_attr::kStructuralMutate, + reinterpret_cast(&TypeOnlyExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TypeOnlyExprMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&SeqExprVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&SeqExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&SeqExprMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&IfVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&IfMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&IfMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&FunctionVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&FunctionMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&FunctionMaybeInplaceMutate)); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, + reinterpret_cast(&TypeOnlyExprVisit)) + .attr(refl::type_attr::kStructuralMutate, + reinterpret_cast(&TypeOnlyExprMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TypeOnlyExprMaybeInplaceMutate)); } If::If(Expr cond, Expr true_branch, Expr false_branch, Span span) { diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc new file mode 100644 index 000000000000..bcb26949f298 --- /dev/null +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -0,0 +1,107 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +#include +#include +#include +#include +#include +#include +#include + +namespace { + +template +void ExpectStructuralHooks() { + namespace refl = tvm::ffi::reflection; + for (const char* attr_name : + {refl::type_attr::kStructuralVisit, refl::type_attr::kStructuralMutate, + refl::type_attr::kStructuralMaybeInplaceMutate}) { + refl::TypeAttrColumn column(attr_name); + EXPECT_EQ(column[TNode::RuntimeTypeIndex()].type_index(), tvm::ffi::TypeIndex::kTVMFFIOpaquePtr) + << TNode::_type_key << " is missing " << attr_name; + } +} + +TEST(RelaxExprStructuralHooks, EveryConcreteExprHasExplicitHooks) { + using namespace tvm::relax; + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); + ExpectStructuralHooks(); +} + +TEST(RelaxExprStructuralHooks, UnchangedCallbackPreservesAncestorIdentity) { + using namespace tvm; + using namespace tvm::relax; + DataflowVar var("x", AnyType()); + Expr input = SeqExpr({}, var); + const auto* original = input.get(); + + auto miss = [](const DataflowVar&) -> ffi::Expected> { + return ffi::Unchanged(); + }; + Expr mapped = ffi::StructuralMap(input, miss).cast(); + EXPECT_EQ(mapped.get(), original); +} + +TEST(RelaxExprStructuralHooks, DataflowVarMapsItsInheritedTypeField) { + using namespace tvm; + using namespace tvm::relax; + DataflowVar input_var("x", AnyType()); + Function input({input_var}, SeqExpr({}, input_var), AnyType()); + auto replace_any_type = [](const AnyType&) -> ffi::Expected> { + return ffi::Any(TensorMapType()); + }; + Function mapped = + ffi::StructuralMap(input, replace_any_type).cast(); + + const auto* var = mapped->params[0].as(); + ASSERT_NE(var, nullptr); + EXPECT_NE(var->ty.as(), nullptr); +} + +TEST(RelaxExprStructuralHooks, StructuralEqualAndHashStillUseReflectedConstants) { + using namespace tvm; + using namespace tvm::relax; + ffi::StructuralEqual equal; + ffi::StructuralHash hash; + + StringImm lhs("lhs"); + StringImm rhs("rhs"); + EXPECT_FALSE(equal(lhs, rhs)); + EXPECT_NE(hash(lhs), hash(rhs)); + + ExternFunc first("first"); + ExternFunc second("second"); + EXPECT_FALSE(equal(first, second)); + EXPECT_NE(hash(first), hash(second)); + + Function pure({}, SeqExpr({}, Tuple(ffi::Array{})), TupleType(ffi::Array{}), true); + Function impure({}, SeqExpr({}, Tuple(ffi::Array{})), TupleType(ffi::Array{}), false); + EXPECT_FALSE(equal(pure, impure)); + EXPECT_NE(hash(pure), hash(impure)); +} + +} // namespace From 2a4a4a36f0f2f09aa96d3d1fff9dcbd4d1347d43 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 01:14:45 +0000 Subject: [PATCH 03/10] Complete structural hooks for concrete Type nodes --- src/tirx/ir/stmt.cc | 20 ++++++++++++++++++-- tests/cpp/type_structural_hooks_test.cc | 25 +++++++++++++++++++++++++ 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc index bd774356b2e5..644cdd1775fa 100644 --- a/src/tirx/ir/stmt.cc +++ b/src/tirx/ir/stmt.cc @@ -46,6 +46,18 @@ using SubscriptSlice = ffi::Array, ffi::Optional, ffi::Optional>, PrimExpr>>; +TVMFFIAny BufferRegionTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny BufferRegionTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny BufferRegionTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + ffi::ObjectRef RealizeBufferRegionSubscript(Expr value, SubscriptSlice slice, Span span) { BufferRegion source = value.as_or_throw(); TVM_FFI_CHECK_LE(slice.size(), source->region.size(), IndexError) @@ -1648,8 +1660,12 @@ BufferRegionType::BufferRegionType() : Type(ffi::UnsafeInit{}) { TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; BufferRegionTypeNode::RegisterReflection(); - refl::TypeAttrDef().def("__subscript_expr_realize__", - RealizeBufferRegionSubscript); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&BufferRegionTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&BufferRegionTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&BufferRegionTypeMaybeInplaceMutate)) + .def("__subscript_expr_realize__", RealizeBufferRegionSubscript); } BufferRegion::BufferRegion(BufferVar buffer, ffi::Array region, Span span) { diff --git a/tests/cpp/type_structural_hooks_test.cc b/tests/cpp/type_structural_hooks_test.cc index f3d3d600c3a6..f120495f4c17 100644 --- a/tests/cpp/type_structural_hooks_test.cc +++ b/tests/cpp/type_structural_hooks_test.cc @@ -25,7 +25,10 @@ #include #include #include +#include #include +#include +#include namespace { @@ -53,6 +56,28 @@ TEST(TypeStructuralHooks, EveryConcreteOpenTypeHasExplicitHooks) { ExpectStructuralHooks(); ExpectStructuralHooks(); ExpectStructuralHooks(); + ExpectStructuralHooks(); +} + +TEST(TypeStructuralHooks, RelaxFuncTypeParametersUsePatternDefinitionRegion) { + using namespace tvm; + tirx::PrimVar symbolic_extent("n", PrimType::Int(64)); + relax::TensorType tensor_type(relax::ShapeExpr(ffi::Array{symbolic_extent}), + PrimType::Float(32)); + relax::FuncType input({tensor_type}, tensor_type, true); + std::vector observed_regions; + auto observe_var = [&](const VarNode*, + TVMFFIDefRegionKind region) -> ffi::Expected> { + observed_regions.push_back(region); + return ffi::Unchanged(); + }; + + relax::FuncType mapped = + ffi::StructuralMap(input, observe_var).cast(); + + EXPECT_TRUE(mapped.same_as(input)); + ASSERT_FALSE(observed_regions.empty()); + EXPECT_EQ(observed_regions.front(), kTVMFFIDefRegionKindPattern); } TEST(TypeStructuralHooks, StructuralMapDescendsThroughTypeFields) { From 63053db7036527195a24fa032f0ed0cc2682f319 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 01:16:13 +0000 Subject: [PATCH 04/10] Give GlobalVar explicit no-descend structural hooks --- src/ir/expr.cc | 24 ++++++++++++++++++- tests/cpp/relax_expr_structural_hooks_test.cc | 4 ++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/src/ir/expr.cc b/src/ir/expr.cc index 1af350c70d23..c9e0e4683c4e 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -381,6 +381,20 @@ TVMFFIAny VarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView return ffi::details::UnchangedOrUnsafe::MoveToTVMFFIAny(std::move(result)); } +TVMFFIAny GlobalVarVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + // GlobalVar is a module-level symbol. name_hint is scalar identity and ty is derived from the + // referenced function, matching GlobalVarNode's custom structural equality/hash definition. + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny GlobalVarMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny GlobalVarMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + TVMFFIAny CallVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { // skips: attrs, constant metadata left untouched like the classic Expr functors. const CallNode* self = @@ -859,7 +873,15 @@ GlobalVar::GlobalVar(ffi::String name_hint, Span span) { data_ = std::move(n); } -TVM_FFI_STATIC_INIT_BLOCK() { GlobalVarNode::RegisterReflection(); } +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + GlobalVarNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&GlobalVarVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&GlobalVarMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&GlobalVarMaybeInplaceMutate)); +} // Call Call::Call(Type ret_ty, Expr op, ffi::Array args, Attrs attrs, ffi::Array ty_args, diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc index bcb26949f298..6c21c02e61ee 100644 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -52,6 +52,10 @@ TEST(RelaxExprStructuralHooks, EveryConcreteExprHasExplicitHooks) { ExpectStructuralHooks(); } +TEST(RelaxExprStructuralHooks, ReviewedCoreExprNodesHaveExplicitHooks) { + ExpectStructuralHooks(); +} + TEST(RelaxExprStructuralHooks, UnchangedCallbackPreservesAncestorIdentity) { using namespace tvm; using namespace tvm::relax; From 2145e896ef4618401af87687edb4e3dd084ac2fd Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 01:18:09 +0000 Subject: [PATCH 05/10] Traverse structural fields of PrimFunc nodes --- src/tirx/ir/function.cc | 73 +++++++++++++++++++ tests/cpp/relax_expr_structural_hooks_test.cc | 19 +++++ 2 files changed, 92 insertions(+) diff --git a/src/tirx/ir/function.cc b/src/tirx/ir/function.cc index 463a15084f7d..e0f295c93118 100644 --- a/src/tirx/ir/function.cc +++ b/src/tirx/ir/function.cc @@ -21,6 +21,8 @@ * \file src/tirx/ir/function.cc * \brief The function data structure. */ +#include +#include #include #include #include @@ -32,9 +34,80 @@ namespace tvm { namespace tirx { +namespace { + +TVMFFIAny PrimFuncVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: attrs (metadata), ty (derived by InferType) + const PrimFuncNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( + kTVMFFIDefRegionKindPattern, [&]() { return visitor->VisitExpected(self->params); })); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->ret_type)); + TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->body)); + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny PrimFuncMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: attrs (metadata), ty (derived by InferType) + const PrimFuncNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() { + return mutator->MutateExpected(self->params); + })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_type, + mutator->MutateExpected(self->ret_type)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MutateExpected(self->body)); + if (mapped_params.UnchangedOrSameAs(self->params) && + mapped_ret_type.UnchangedOrSameAs(self->ret_type) && + mapped_body.UnchangedOrSameAs(self->body)) { + return ffi::Unchanged().CopyToTVMFFIAny(); + } + ffi::ObjectPtr copy = ffi::make_object(*self); + copy->params = std::move(mapped_params).ValueOrUnchanged(std::move(copy->params)); + copy->ret_type = std::move(mapped_ret_type).ValueOrUnchanged(std::move(copy->ret_type)); + copy->body = std::move(mapped_body).ValueOrUnchanged(std::move(copy->body)); + return ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy))); +} + +TVMFFIAny PrimFuncMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, + ffi::AnyView value) noexcept { + // skips: attrs (metadata), ty (derived by InferType) + PrimFuncNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( + ffi::UnchangedOr>, mapped_params, + mutator->WithDefRegionKind(kTVMFFIDefRegionKindPattern, [&]() { + return mutator->MaybeInplaceMutateIfUniqueExpected(self->params); + })); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_ret_type, + mutator->MaybeInplaceMutateIfUniqueExpected(self->ret_type)); + TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_body, + mutator->MaybeInplaceMutateIfUniqueExpected(self->body)); + if (!mapped_params.UnchangedOrSameAs(self->params)) { + self->params = std::move(mapped_params).ValueUnchecked(); + } + if (!mapped_ret_type.UnchangedOrSameAs(self->ret_type)) { + self->ret_type = std::move(mapped_ret_type).ValueUnchecked(); + } + if (!mapped_body.UnchangedOrSameAs(self->body)) { + self->body = std::move(mapped_body).ValueUnchecked(); + } + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; PrimFuncNode::RegisterReflection(); TensorIntrinNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&PrimFuncVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&PrimFuncMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&PrimFuncMaybeInplaceMutate)); } namespace { diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc index 6c21c02e61ee..64b9ebbb8cd1 100644 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -24,6 +24,8 @@ #include #include #include +#include +#include namespace { @@ -54,6 +56,23 @@ TEST(RelaxExprStructuralHooks, EveryConcreteExprHasExplicitHooks) { TEST(RelaxExprStructuralHooks, ReviewedCoreExprNodesHaveExplicitHooks) { ExpectStructuralHooks(); + ExpectStructuralHooks(); +} + +TEST(RelaxExprStructuralHooks, PrimFuncDescendsIntoBody) { + using namespace tvm; + tirx::PrimFunc input({}, tirx::Evaluate(IntImm(PrimType::Int(32), 1))); + auto replace_one = [](const IntImm& value) -> ffi::Expected> { + if (value->value != 1) return ffi::Unchanged(); + return ffi::Any(IntImm(value.ty().as_or_throw(), 2)); + }; + + tirx::PrimFunc mapped = + ffi::StructuralMap(input, replace_one).cast(); + + const auto* evaluate = mapped->body.as(); + ASSERT_NE(evaluate, nullptr); + EXPECT_EQ(evaluate->value.as()->value, 2); } TEST(RelaxExprStructuralHooks, UnchangedCallbackPreservesAncestorIdentity) { From b7cf65b0ddf8c655b132d114e1653ec40b2f2c02 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 01:42:53 +0000 Subject: [PATCH 06/10] Avoid structural descent through Op registry metadata --- src/ir/op.cc | 25 +++++++++++++++++++ tests/cpp/relax_expr_structural_hooks_test.cc | 24 ++++++++++++++++++ 2 files changed, 49 insertions(+) diff --git a/src/ir/op.cc b/src/ir/op.cc index 560cee6c9d89..c300806b74a1 100644 --- a/src/ir/op.cc +++ b/src/ir/op.cc @@ -21,6 +21,8 @@ * \file src/ir/op.cc * \brief Primitive operators and intrinsics. */ +#include +#include #include #include #include @@ -33,9 +35,32 @@ namespace tvm { +namespace { + +TVMFFIAny OpVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + // Ops are unique registry atoms. Avoid reflecting through their registry metadata. + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny OpMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny OpMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +} // namespace + TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; ArgumentInfoNode::RegisterReflection(); OpNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&OpVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&OpMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&OpMaybeInplaceMutate)); } using ffi::Any; diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc index 64b9ebbb8cd1..ff560fcd5477 100644 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -23,6 +23,7 @@ #include #include #include +#include #include #include #include @@ -56,9 +57,32 @@ TEST(RelaxExprStructuralHooks, EveryConcreteExprHasExplicitHooks) { TEST(RelaxExprStructuralHooks, ReviewedCoreExprNodesHaveExplicitHooks) { ExpectStructuralHooks(); + ExpectStructuralHooks(); ExpectStructuralHooks(); } +TEST(RelaxExprStructuralHooks, OpHookPreservesIdentityWithoutDescendingIntoMetadata) { + using namespace tvm; + Op input = Op::Get("ir.prim.likely"); + int op_callbacks = 0; + int metadata_callbacks = 0; + auto observe_op = [&](const Op&) -> ffi::Expected> { + ++op_callbacks; + return ffi::Unchanged(); + }; + auto observe_metadata = [&](const ffi::String&) -> ffi::Expected> { + ++metadata_callbacks; + return ffi::Unchanged(); + }; + + Op mapped = ffi::StructuralMap(input, observe_op, observe_metadata) + .cast(); + + EXPECT_TRUE(mapped.same_as(input)); + EXPECT_EQ(op_callbacks, 1); + EXPECT_EQ(metadata_callbacks, 0); +} + TEST(RelaxExprStructuralHooks, PrimFuncDescendsIntoBody) { using namespace tvm; tirx::PrimFunc input({}, tirx::Evaluate(IntImm(PrimType::Int(32), 1))); From 733bd8e9c54a1d3a08a5ffe19ed36d364e9fd87d Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 02:30:06 +0000 Subject: [PATCH 07/10] Cover fieldless type sentinels in structural traversal --- src/ir/type.cc | 49 ++++++++++++++++++++++++- tests/cpp/type_structural_hooks_test.cc | 19 +++++++++- 2 files changed, 65 insertions(+), 3 deletions(-) diff --git a/src/ir/type.cc b/src/ir/type.cc index 7d2e1050ad16..1947f489856c 100644 --- a/src/ir/type.cc +++ b/src/ir/type.cc @@ -69,6 +69,32 @@ ffi::ObjectPtr GetCachedPrimTypeNode(DLDataType dtype) { // Structural traversal hooks +TVMFFIAny TypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + // Type::Missing() is the only concrete TypeNode value; span is ignored debug metadata. + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny TypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny TypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny OpaqueTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { + // OpaqueType is a field-less construction-time marker; span is ignored debug metadata. + return ffi::AnyView(nullptr).CopyToTVMFFIAny(); +} + +TVMFFIAny OpaqueTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + +TVMFFIAny OpaqueTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) noexcept { + return ffi::Unchanged().CopyToTVMFFIAny(); +} + TVMFFIAny PrimTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { // dtype is a constant: reflected for StructuralEqual/Hash, // not traversed by the visitor/mutator contract. @@ -88,6 +114,7 @@ TVMFFIAny PrimTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) n } TVMFFIAny PointerTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: storage_scope (scalar) const PointerTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->element_type)); @@ -95,6 +122,7 @@ TVMFFIAny PointerTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView valu } TVMFFIAny PointerTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: storage_scope (scalar) const PointerTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_element_type, @@ -110,6 +138,7 @@ TVMFFIAny PointerTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView val TVMFFIAny PointerTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: storage_scope (scalar) PointerTypeNode* self = const_cast( ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( @@ -222,9 +251,25 @@ bool Type::IsMissing() const { return this->same_as(Type::Missing()); } OpaqueType::OpaqueType() : Type(ffi::UnsafeInit{}) { data_ = ffi::make_object(); } -TVM_FFI_STATIC_INIT_BLOCK() { TypeNode::RegisterReflection(); } +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + TypeNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&TypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&TypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&TypeMaybeInplaceMutate)); +} -TVM_FFI_STATIC_INIT_BLOCK() { OpaqueTypeNode::RegisterReflection(); } +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + OpaqueTypeNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&OpaqueTypeVisit)) + .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&OpaqueTypeMutate)) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + reinterpret_cast(&OpaqueTypeMaybeInplaceMutate)); +} // PrimType PrimType::PrimType(DLDataType dtype) : Type(ffi::UnsafeInit{}) { diff --git a/tests/cpp/type_structural_hooks_test.cc b/tests/cpp/type_structural_hooks_test.cc index f120495f4c17..1500bf8b04ee 100644 --- a/tests/cpp/type_structural_hooks_test.cc +++ b/tests/cpp/type_structural_hooks_test.cc @@ -44,8 +44,10 @@ void ExpectStructuralHooks() { } } -TEST(TypeStructuralHooks, EveryConcreteOpenTypeHasExplicitHooks) { +TEST(TypeStructuralHooks, EveryConcreteTypeHasExplicitHooks) { using namespace tvm; + ExpectStructuralHooks(); + ExpectStructuralHooks(); ExpectStructuralHooks(); ExpectStructuralHooks(); ExpectStructuralHooks(); @@ -59,6 +61,21 @@ TEST(TypeStructuralHooks, EveryConcreteOpenTypeHasExplicitHooks) { ExpectStructuralHooks(); } +TEST(TypeStructuralHooks, FieldlessSentinelsPreserveIdentity) { + using namespace tvm; + auto miss = [](const Type&) -> ffi::Expected> { + return ffi::Unchanged(); + }; + + Type missing = Type::Missing(); + Type mapped_missing = ffi::StructuralMap(missing, miss).cast(); + EXPECT_TRUE(mapped_missing.same_as(missing)); + + OpaqueType opaque; + Type mapped_opaque = ffi::StructuralMap(opaque, miss).cast(); + EXPECT_TRUE(mapped_opaque.same_as(opaque)); +} + TEST(TypeStructuralHooks, RelaxFuncTypeParametersUsePatternDefinitionRegion) { using namespace tvm; tirx::PrimVar symbolic_extent("n", PrimType::Int(64)); From 37ccf0de177145bd26a133252b63c9facf58e8fe Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 02:31:13 +0000 Subject: [PATCH 08/10] Document structural traversal field boundaries --- src/ir/expr.cc | 4 ++ src/relax/ir/dependent_type.cc | 45 +++++++++++-------- src/relax/ir/expr.cc | 10 +++++ tests/cpp/relax_expr_structural_hooks_test.cc | 19 ++++++++ 4 files changed, 60 insertions(+), 18 deletions(-) diff --git a/src/ir/expr.cc b/src/ir/expr.cc index c9e0e4683c4e..0c6c76734fe3 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -278,6 +278,8 @@ TVMFFIAny RangeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyVi return ffi::Unchanged().CopyToTVMFFIAny(); } +// DataflowVarNode duplicates this protocol because structural hooks do not inherit. Keep the two +// hook triples in lockstep when changing remap, PrimType-skip, or definition-region behavior. TVMFFIAny VarVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { // skips: name const VarNode* self = @@ -384,6 +386,8 @@ TVMFFIAny VarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView TVMFFIAny GlobalVarVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { // GlobalVar is a module-level symbol. name_hint is scalar identity and ty is derived from the // referenced function, matching GlobalVarNode's custom structural equality/hash definition. + // StructuralMap's callback layer memoizes FreeVar results before and after this hook, so repeated + // occurrences reuse one replacement without a definition-site VarRemap operation here. return ffi::AnyView(nullptr).CopyToTVMFFIAny(); } diff --git a/src/relax/ir/dependent_type.cc b/src/relax/ir/dependent_type.cc index b12901fa045f..14755ca573e9 100644 --- a/src/relax/ir/dependent_type.cc +++ b/src/relax/ir/dependent_type.cc @@ -47,6 +47,7 @@ TVMFFIAny AnyTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyView) no } TVMFFIAny ShapeTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: ndim (scalar) const ShapeTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->values)); @@ -54,6 +55,7 @@ TVMFFIAny ShapeTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) } TVMFFIAny ShapeTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: ndim (scalar) const ShapeTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>>, @@ -66,6 +68,7 @@ TVMFFIAny ShapeTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value TVMFFIAny ShapeTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: ndim (scalar) ShapeTypeNode* self = const_cast( ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>>, @@ -78,6 +81,7 @@ TVMFFIAny ShapeTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, } TVMFFIAny TensorTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: ndim (scalar) const TensorTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->shape)); @@ -87,6 +91,7 @@ TVMFFIAny TensorTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value } TVMFFIAny TensorTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: ndim (scalar) const TensorTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, @@ -108,6 +113,7 @@ TVMFFIAny TensorTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView valu TVMFFIAny TensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: ndim (scalar) TensorTypeNode* self = const_cast( ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, @@ -129,6 +135,7 @@ TVMFFIAny TensorTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, } TVMFFIAny FuncTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: derive_func (environment-backed callable metadata), purity (scalar) const FuncTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( @@ -138,6 +145,7 @@ TVMFFIAny FuncTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) } TVMFFIAny FuncTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: derive_func (environment-backed callable metadata), purity (scalar) const FuncTypeNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( @@ -157,6 +165,7 @@ TVMFFIAny FuncTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) TVMFFIAny FuncTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: derive_func (environment-backed callable metadata), purity (scalar) FuncTypeNode* self = const_cast( ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( @@ -197,29 +206,29 @@ TVM_FFI_STATIC_INIT_BLOCK() { .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&TensorTypeVisit)) .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&TensorTypeMutate)) .attr(refl::type_attr::kStructuralMaybeInplaceMutate, - reinterpret_cast(&TensorTypeMaybeInplaceMutate)); + reinterpret_cast(&TensorTypeMaybeInplaceMutate)) + .def("__subscript_expr_realize__", + [](Expr value, + ffi::Array, ffi::Optional, + ffi::Optional>, + PrimExpr>> + slice, + Span span) -> ffi::ObjectRef { + TVM_FFI_CHECK_EQ(slice.size(), 1, IndexError) + << "A Relax expression requires exactly one index"; + auto index = slice[0].as(); + TVM_FFI_CHECK(index.has_value(), TypeError) + << "A Relax expression requires a point index"; + const auto* imm = index.value().as(); + TVM_FFI_CHECK(imm != nullptr, TypeError) + << "A Relax expression requires a constant integer index"; + return TupleGetItem(value, static_cast(imm->value), span); + }); refl::TypeAttrDef() .attr(refl::type_attr::kStructuralVisit, reinterpret_cast(&FuncTypeVisit)) .attr(refl::type_attr::kStructuralMutate, reinterpret_cast(&FuncTypeMutate)) .attr(refl::type_attr::kStructuralMaybeInplaceMutate, reinterpret_cast(&FuncTypeMaybeInplaceMutate)); - refl::TypeAttrDef().def( - "__subscript_expr_realize__", - [](Expr value, - ffi::Array, ffi::Optional, ffi::Optional>, - PrimExpr>> - slice, - Span span) -> ffi::ObjectRef { - TVM_FFI_CHECK_EQ(slice.size(), 1, IndexError) - << "A Relax expression requires exactly one index"; - auto index = slice[0].as(); - TVM_FFI_CHECK(index.has_value(), TypeError) << "A Relax expression requires a point index"; - const auto* imm = index.value().as(); - TVM_FFI_CHECK(imm != nullptr, TypeError) - << "A Relax expression requires a constant integer index"; - return TupleGetItem(value, static_cast(imm->value), span); - }); } AnyType::AnyType(Span span) : Type(ffi::UnsafeInit{}) { diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc index ba03c090f4da..d841d7817a6c 100644 --- a/src/relax/ir/expr.cc +++ b/src/relax/ir/expr.cc @@ -31,6 +31,9 @@ namespace relax { namespace { +// Traverses only ExprNode::ty. The per-node payload is an intentional constant leaf: +// ConstantNode::data is tensor data; StringImmNode::value and DataTypeImmNode::value are scalars; +// ExternFuncNode::global_symbol is scalar and BaseFuncNode::attrs is metadata. template TVMFFIAny TypeOnlyExprVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { const TNode* self = @@ -103,6 +106,8 @@ TVMFFIAny ShapeExprMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, return ffi::Unchanged().CopyToTVMFFIAny(); } +// Hooks do not inherit, so DataflowVar must mirror the base VarNode remap, PrimType-skip, and +// Simple-to-None definition-region protocol. Keep this hook triple in lockstep with VarNode. TVMFFIAny DataflowVarVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { const DataflowVarNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); @@ -305,7 +310,10 @@ TVMFFIAny IfMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView return ffi::Unchanged().CopyToTVMFFIAny(); } +// Parameters precede the reflected ty field in this hook triple so their pattern region establishes +// the remap before the derived function type can refer to those symbolic definitions. TVMFFIAny FunctionVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { + // skips: attrs (metadata), is_pure (scalar) const FunctionNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( @@ -317,6 +325,7 @@ TVMFFIAny FunctionVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value) } TVMFFIAny FunctionMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: attrs (metadata), is_pure (scalar) const FunctionNode* self = ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_params, @@ -343,6 +352,7 @@ TVMFFIAny FunctionMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) TVMFFIAny FunctionMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { + // skips: attrs (metadata), is_pure (scalar) FunctionNode* self = const_cast( ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN( diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc index ff560fcd5477..5a1886ba0a52 100644 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -83,6 +83,25 @@ TEST(RelaxExprStructuralHooks, OpHookPreservesIdentityWithoutDescendingIntoMetad EXPECT_EQ(metadata_callbacks, 0); } +TEST(RelaxExprStructuralHooks, GlobalVarCallbackReplacementIsMemoized) { + using namespace tvm; + GlobalVar symbol("f"); + Expr input = Tuple({symbol, symbol}); + int callbacks = 0; + auto replace_symbol = [&](const GlobalVar&) -> ffi::Expected> { + ++callbacks; + return ffi::Any(GlobalVar("g")); + }; + + Expr mapped = ffi::StructuralMap(input, replace_symbol).cast(); + const auto* tuple = mapped.as(); + ASSERT_NE(tuple, nullptr); + ASSERT_EQ(tuple->fields.size(), 2U); + EXPECT_EQ(callbacks, 1); + EXPECT_TRUE(tuple->fields[0].same_as(tuple->fields[1])); + EXPECT_FALSE(tuple->fields[0].same_as(symbol)); +} + TEST(RelaxExprStructuralHooks, PrimFuncDescendsIntoBody) { using namespace tvm; tirx::PrimFunc input({}, tirx::Evaluate(IntImm(PrimType::Int(32), 1))); From 272f2018faa711779aa2832f9594406c07dc8d15 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 03:03:37 +0000 Subject: [PATCH 09/10] Make GlobalVar callback identity caller-owned --- src/ir/expr.cc | 4 ++-- tests/cpp/relax_expr_structural_hooks_test.cc | 8 +++----- 2 files changed, 5 insertions(+), 7 deletions(-) diff --git a/src/ir/expr.cc b/src/ir/expr.cc index 0c6c76734fe3..b9b51da08963 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -386,8 +386,8 @@ TVMFFIAny VarMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView TVMFFIAny GlobalVarVisit(ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { // GlobalVar is a module-level symbol. name_hint is scalar identity and ty is derived from the // referenced function, matching GlobalVarNode's custom structural equality/hash definition. - // StructuralMap's callback layer memoizes FreeVar results before and after this hook, so repeated - // occurrences reuse one replacement without a definition-site VarRemap operation here. + // It has no definition site where this hook could establish a VarRemap. A callback that renames + // GlobalVars is therefore responsible for returning one stable replacement per module symbol. return ffi::AnyView(nullptr).CopyToTVMFFIAny(); } diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc index 5a1886ba0a52..4445d489b060 100644 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ b/tests/cpp/relax_expr_structural_hooks_test.cc @@ -83,21 +83,19 @@ TEST(RelaxExprStructuralHooks, OpHookPreservesIdentityWithoutDescendingIntoMetad EXPECT_EQ(metadata_callbacks, 0); } -TEST(RelaxExprStructuralHooks, GlobalVarCallbackReplacementIsMemoized) { +TEST(RelaxExprStructuralHooks, GlobalVarCallbackCanReturnStableReplacement) { using namespace tvm; GlobalVar symbol("f"); Expr input = Tuple({symbol, symbol}); - int callbacks = 0; + GlobalVar replacement("g"); auto replace_symbol = [&](const GlobalVar&) -> ffi::Expected> { - ++callbacks; - return ffi::Any(GlobalVar("g")); + return ffi::Any(replacement); }; Expr mapped = ffi::StructuralMap(input, replace_symbol).cast(); const auto* tuple = mapped.as(); ASSERT_NE(tuple, nullptr); ASSERT_EQ(tuple->fields.size(), 2U); - EXPECT_EQ(callbacks, 1); EXPECT_TRUE(tuple->fields[0].same_as(tuple->fields[1])); EXPECT_FALSE(tuple->fields[0].same_as(symbol)); } From 87d4e68459d450d0a2a4fbe914af7ad739f00d70 Mon Sep 17 00:00:00 2001 From: tqchen Date: Thu, 10 Sep 2026 12:33:15 +0000 Subject: [PATCH 10/10] Remove dedicated structural hook tests --- tests/cpp/relax_expr_structural_hooks_test.cc | 171 ------------------ tests/cpp/type_structural_hooks_test.cc | 148 --------------- 2 files changed, 319 deletions(-) delete mode 100644 tests/cpp/relax_expr_structural_hooks_test.cc delete mode 100644 tests/cpp/type_structural_hooks_test.cc diff --git a/tests/cpp/relax_expr_structural_hooks_test.cc b/tests/cpp/relax_expr_structural_hooks_test.cc deleted file mode 100644 index 4445d489b060..000000000000 --- a/tests/cpp/relax_expr_structural_hooks_test.cc +++ /dev/null @@ -1,171 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace { - -template -void ExpectStructuralHooks() { - namespace refl = tvm::ffi::reflection; - for (const char* attr_name : - {refl::type_attr::kStructuralVisit, refl::type_attr::kStructuralMutate, - refl::type_attr::kStructuralMaybeInplaceMutate}) { - refl::TypeAttrColumn column(attr_name); - EXPECT_EQ(column[TNode::RuntimeTypeIndex()].type_index(), tvm::ffi::TypeIndex::kTVMFFIOpaquePtr) - << TNode::_type_key << " is missing " << attr_name; - } -} - -TEST(RelaxExprStructuralHooks, EveryConcreteExprHasExplicitHooks) { - using namespace tvm::relax; - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); -} - -TEST(RelaxExprStructuralHooks, ReviewedCoreExprNodesHaveExplicitHooks) { - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); -} - -TEST(RelaxExprStructuralHooks, OpHookPreservesIdentityWithoutDescendingIntoMetadata) { - using namespace tvm; - Op input = Op::Get("ir.prim.likely"); - int op_callbacks = 0; - int metadata_callbacks = 0; - auto observe_op = [&](const Op&) -> ffi::Expected> { - ++op_callbacks; - return ffi::Unchanged(); - }; - auto observe_metadata = [&](const ffi::String&) -> ffi::Expected> { - ++metadata_callbacks; - return ffi::Unchanged(); - }; - - Op mapped = ffi::StructuralMap(input, observe_op, observe_metadata) - .cast(); - - EXPECT_TRUE(mapped.same_as(input)); - EXPECT_EQ(op_callbacks, 1); - EXPECT_EQ(metadata_callbacks, 0); -} - -TEST(RelaxExprStructuralHooks, GlobalVarCallbackCanReturnStableReplacement) { - using namespace tvm; - GlobalVar symbol("f"); - Expr input = Tuple({symbol, symbol}); - GlobalVar replacement("g"); - auto replace_symbol = [&](const GlobalVar&) -> ffi::Expected> { - return ffi::Any(replacement); - }; - - Expr mapped = ffi::StructuralMap(input, replace_symbol).cast(); - const auto* tuple = mapped.as(); - ASSERT_NE(tuple, nullptr); - ASSERT_EQ(tuple->fields.size(), 2U); - EXPECT_TRUE(tuple->fields[0].same_as(tuple->fields[1])); - EXPECT_FALSE(tuple->fields[0].same_as(symbol)); -} - -TEST(RelaxExprStructuralHooks, PrimFuncDescendsIntoBody) { - using namespace tvm; - tirx::PrimFunc input({}, tirx::Evaluate(IntImm(PrimType::Int(32), 1))); - auto replace_one = [](const IntImm& value) -> ffi::Expected> { - if (value->value != 1) return ffi::Unchanged(); - return ffi::Any(IntImm(value.ty().as_or_throw(), 2)); - }; - - tirx::PrimFunc mapped = - ffi::StructuralMap(input, replace_one).cast(); - - const auto* evaluate = mapped->body.as(); - ASSERT_NE(evaluate, nullptr); - EXPECT_EQ(evaluate->value.as()->value, 2); -} - -TEST(RelaxExprStructuralHooks, UnchangedCallbackPreservesAncestorIdentity) { - using namespace tvm; - using namespace tvm::relax; - DataflowVar var("x", AnyType()); - Expr input = SeqExpr({}, var); - const auto* original = input.get(); - - auto miss = [](const DataflowVar&) -> ffi::Expected> { - return ffi::Unchanged(); - }; - Expr mapped = ffi::StructuralMap(input, miss).cast(); - EXPECT_EQ(mapped.get(), original); -} - -TEST(RelaxExprStructuralHooks, DataflowVarMapsItsInheritedTypeField) { - using namespace tvm; - using namespace tvm::relax; - DataflowVar input_var("x", AnyType()); - Function input({input_var}, SeqExpr({}, input_var), AnyType()); - auto replace_any_type = [](const AnyType&) -> ffi::Expected> { - return ffi::Any(TensorMapType()); - }; - Function mapped = - ffi::StructuralMap(input, replace_any_type).cast(); - - const auto* var = mapped->params[0].as(); - ASSERT_NE(var, nullptr); - EXPECT_NE(var->ty.as(), nullptr); -} - -TEST(RelaxExprStructuralHooks, StructuralEqualAndHashStillUseReflectedConstants) { - using namespace tvm; - using namespace tvm::relax; - ffi::StructuralEqual equal; - ffi::StructuralHash hash; - - StringImm lhs("lhs"); - StringImm rhs("rhs"); - EXPECT_FALSE(equal(lhs, rhs)); - EXPECT_NE(hash(lhs), hash(rhs)); - - ExternFunc first("first"); - ExternFunc second("second"); - EXPECT_FALSE(equal(first, second)); - EXPECT_NE(hash(first), hash(second)); - - Function pure({}, SeqExpr({}, Tuple(ffi::Array{})), TupleType(ffi::Array{}), true); - Function impure({}, SeqExpr({}, Tuple(ffi::Array{})), TupleType(ffi::Array{}), false); - EXPECT_FALSE(equal(pure, impure)); - EXPECT_NE(hash(pure), hash(impure)); -} - -} // namespace diff --git a/tests/cpp/type_structural_hooks_test.cc b/tests/cpp/type_structural_hooks_test.cc deleted file mode 100644 index 1500bf8b04ee..000000000000 --- a/tests/cpp/type_structural_hooks_test.cc +++ /dev/null @@ -1,148 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace { - -template -void ExpectStructuralHooks() { - namespace refl = tvm::ffi::reflection; - for (const char* attr_name : - {refl::type_attr::kStructuralVisit, refl::type_attr::kStructuralMutate, - refl::type_attr::kStructuralMaybeInplaceMutate}) { - refl::TypeAttrColumn column(attr_name); - EXPECT_EQ(column[TNode::RuntimeTypeIndex()].type_index(), tvm::ffi::TypeIndex::kTVMFFIOpaquePtr) - << TNode::_type_key << " is missing " << attr_name; - } -} - -TEST(TypeStructuralHooks, EveryConcreteTypeHasExplicitHooks) { - using namespace tvm; - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); - ExpectStructuralHooks(); -} - -TEST(TypeStructuralHooks, FieldlessSentinelsPreserveIdentity) { - using namespace tvm; - auto miss = [](const Type&) -> ffi::Expected> { - return ffi::Unchanged(); - }; - - Type missing = Type::Missing(); - Type mapped_missing = ffi::StructuralMap(missing, miss).cast(); - EXPECT_TRUE(mapped_missing.same_as(missing)); - - OpaqueType opaque; - Type mapped_opaque = ffi::StructuralMap(opaque, miss).cast(); - EXPECT_TRUE(mapped_opaque.same_as(opaque)); -} - -TEST(TypeStructuralHooks, RelaxFuncTypeParametersUsePatternDefinitionRegion) { - using namespace tvm; - tirx::PrimVar symbolic_extent("n", PrimType::Int(64)); - relax::TensorType tensor_type(relax::ShapeExpr(ffi::Array{symbolic_extent}), - PrimType::Float(32)); - relax::FuncType input({tensor_type}, tensor_type, true); - std::vector observed_regions; - auto observe_var = [&](const VarNode*, - TVMFFIDefRegionKind region) -> ffi::Expected> { - observed_regions.push_back(region); - return ffi::Unchanged(); - }; - - relax::FuncType mapped = - ffi::StructuralMap(input, observe_var).cast(); - - EXPECT_TRUE(mapped.same_as(input)); - ASSERT_FALSE(observed_regions.empty()); - EXPECT_EQ(observed_regions.front(), kTVMFFIDefRegionKindPattern); -} - -TEST(TypeStructuralHooks, StructuralMapDescendsThroughTypeFields) { - using namespace tvm; - Type input = TupleType( - {PointerType(PrimType::Float(32), "global"), relax::TensorType(PrimType::Float(32), 2)}); - Type mapped = ffi::StructuralMap( - input, - [](const PrimType& type) -> ffi::Expected> { - if (!type.MatchesElementType(DLDataTypeCode::kDLFloat, 32)) { - return ffi::Unchanged(); - } - return ffi::Any(PrimType::Float(64)); - }) - .cast(); - - const auto* tuple = mapped.as(); - ASSERT_NE(tuple, nullptr); - EXPECT_TRUE(tuple->fields[0].as()->element_type.as()->dtype == - PrimType::Float(64)->dtype); - EXPECT_TRUE(tuple->fields[1].as()->dtype.value()->dtype == - PrimType::Float(64)->dtype); -} - -TEST(TypeStructuralHooks, StructuralEqualAndHashStillUseAllReflectedFields) { - using namespace tvm; - ffi::StructuralEqual equal; - ffi::StructuralHash hash; - - PointerType global(PrimType::Float(32), "global"); - PointerType shared(PrimType::Float(32), "shared"); - EXPECT_FALSE(equal(global, shared)); - EXPECT_NE(hash(global), hash(shared)); - - relax::ShapeType rank_one(1); - relax::ShapeType rank_two(2); - EXPECT_FALSE(equal(rank_one, rank_two)); - EXPECT_NE(hash(rank_one), hash(rank_two)); - - relax::TensorType f32(PrimType::Float(32), 2); - relax::TensorType f64(PrimType::Float(64), 2); - EXPECT_FALSE(equal(f32, f64)); - EXPECT_NE(hash(f32), hash(f64)); - - relax::FuncType pure({}, relax::AnyType(), true); - relax::FuncType impure({}, relax::AnyType(), false); - EXPECT_FALSE(equal(pure, impure)); - EXPECT_NE(hash(pure), hash(impure)); -} - -} // namespace