This is an automated email from the ASF dual-hosted git repository.
tqchen pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm-ffi.git
The following commit(s) were added to refs/heads/main by this push:
new 54382cf1 [PERF][Rust] Reduce structural traversal and callback
overhead (#806)
54382cf1 is described below
commit 54382cf123ae1b41221d52c9d36b2a0609767d61
Author: Shushi Hong <[email protected]>
AuthorDate: Mon Sep 21 19:19:42 2026 -0400
[PERF][Rust] Reduce structural traversal and callback overhead (#806)
This PR reduces callback dispatch and result-conversion overhead in Rust
structural traversal. It uses concrete-driver mutation ABI entrypoints,
specializes map traversal order, reuses reflection metadata, and avoids
repeated context handling. Containers use their registered hooks.
---
rust/tvm-ffi-macros/src/dispatch.rs | 2 +-
rust/tvm-ffi/src/extra/dispatch.rs | 2 +-
rust/tvm-ffi/src/extra/structural_common.rs | 7 +
rust/tvm-ffi/src/extra/structural_mutate.rs | 427 +++++++++++++--------
rust/tvm-ffi/src/extra/structural_mutate/policy.rs | 56 ++-
rust/tvm-ffi/src/extra/structural_visit.rs | 186 ++++-----
rust/tvm-ffi/src/extra/structural_visit/policy.rs | 29 +-
7 files changed, 423 insertions(+), 286 deletions(-)
diff --git a/rust/tvm-ffi-macros/src/dispatch.rs
b/rust/tvm-ffi-macros/src/dispatch.rs
index 94aeda86..75a04eb4 100644
--- a/rust/tvm-ffi-macros/src/dispatch.rs
+++ b/rust/tvm-ffi-macros/src/dispatch.rs
@@ -264,7 +264,7 @@ fn expand(item_impl: &ItemImpl, args: DispatchArgs) ->
syn::Result<TokenStream2>
{
#policy_method
- #[inline]
+ #[inline(always)]
#[allow(unreachable_code, unused_variables)]
fn visit(
&mut self,
diff --git a/rust/tvm-ffi/src/extra/dispatch.rs
b/rust/tvm-ffi/src/extra/dispatch.rs
index 42b263b4..fdbd0a81 100644
--- a/rust/tvm-ffi/src/extra/dispatch.rs
+++ b/rust/tvm-ffi/src/extra/dispatch.rs
@@ -70,6 +70,6 @@ impl<V: WalkDispatch> NativeVisit for DispatchWalker<V> {
) -> Result<WalkResult> {
self.walker
.dispatch_walk(value, def_region_kind)
- .unwrap_or(Ok(WalkResult::Advance))
+ .unwrap_or_else(|| Ok(WalkResult::Advance))
}
}
diff --git a/rust/tvm-ffi/src/extra/structural_common.rs
b/rust/tvm-ffi/src/extra/structural_common.rs
index 604841c2..070a209e 100644
--- a/rust/tvm-ffi/src/extra/structural_common.rs
+++ b/rust/tvm-ffi/src/extra/structural_common.rs
@@ -170,6 +170,13 @@ impl StructuralView {
Self(raw)
}
+ #[inline]
+ pub(crate) fn from_any(value: &Any) -> &Self {
+ // SAFETY: StructuralView is transparent over TVMFFIAny. The owning
+ // Any keeps its contents live for the lifetime of this shared borrow.
+ unsafe { &*std::ptr::from_ref(value.as_raw_ffi_any()).cast::<Self>() }
+ }
+
#[inline]
pub(crate) fn raw(&self) -> TVMFFIAny {
self.0
diff --git a/rust/tvm-ffi/src/extra/structural_mutate.rs
b/rust/tvm-ffi/src/extra/structural_mutate.rs
index ed972af4..4b3d83f3 100644
--- a/rust/tvm-ffi/src/extra/structural_mutate.rs
+++ b/rust/tvm-ffi/src/extra/structural_mutate.rs
@@ -234,11 +234,12 @@ impl<'a, T> MutateValue<'a, T> {
}
}
+ // Combine permissions here; policy entry and built-in descent recheck
ownership.
fn permit(&self, requested: InplaceMode) -> Permit {
if self.mode == InplaceMode::Disallow {
Permit::Copy
} else {
- requested.permit_if_unique(self.value.raw())
+ requested.permit()
}
}
}
@@ -1099,15 +1100,7 @@ pub trait NativeMap: Sized {
impl<D: MapDispatch> NativeMap for D {
fn map_root(&mut self, root: Any, order: WalkOrder) -> Result<Any> {
- run_structural_mutator(
- root,
- &mut NativeMapper::<_, DefaultMutContextPolicy> {
- dispatch: self,
- order,
- policy: None,
- remap: StructuralVarRemap::default(),
- },
- )
+ run_native_mapper::<_, DefaultMutContextPolicy>(root, self, None,
order)
}
}
@@ -2016,14 +2009,39 @@ struct MemoEntry {
result: Any,
}
-struct NativeMapper<'a, D, Policy> {
+struct NativeMapper<'a, D, Policy, const PRE_ORDER: bool> {
dispatch: &'a mut D,
policy: Option<Rc<Policy>>,
- order: WalkOrder,
remap: StructuralVarRemap,
}
-impl<D: MapDispatch, Policy: MutContextPolicy<D>> NativeMapper<'_, D, Policy> {
+fn run_native_mapper<D: MapDispatch, Policy: MutContextPolicy<D>>(
+ root: Any,
+ dispatch: &mut D,
+ policy: Option<Rc<Policy>>,
+ order: WalkOrder,
+) -> Result<Any> {
+ match order {
+ WalkOrder::PreOrder => NativeMapper::<_, _, true>::run(root, dispatch,
policy),
+ WalkOrder::PostOrder => NativeMapper::<_, _, false>::run(root,
dispatch, policy),
+ }
+}
+
+impl<D: MapDispatch, Policy: MutContextPolicy<D>, const PRE_ORDER: bool>
+ NativeMapper<'_, D, Policy, PRE_ORDER>
+{
+ fn run(root: Any, dispatch: &mut D, policy: Option<Rc<Policy>>) ->
Result<Any> {
+ run_structural_mutator(
+ root,
+ &mut NativeMapper::<_, _, PRE_ORDER> {
+ dispatch,
+ policy,
+ remap: StructuralVarRemap::default(),
+ },
+ )
+ }
+
+ #[inline(always)]
fn map_raw(
&mut self,
raw: TVMFFIAny,
@@ -2042,7 +2060,7 @@ impl<D: MapDispatch, Policy: MutContextPolicy<D>>
NativeMapper<'_, D, Policy> {
Some(result) => {
let mapped = result?;
// A pre-order callback may replace an inline leaf with a
subtree.
- if self.order == WalkOrder::PreOrder &&
!is_plain_inline(mapped.type_index()) {
+ if PRE_ORDER && !is_plain_inline(mapped.type_index()) {
let descended =
self.map_default_root(&mapped, def_region_kind,
Permit::Copy)?;
Ok(if is_unchanged(&descended) {
@@ -2071,8 +2089,8 @@ impl<D: MapDispatch, Policy: MutContextPolicy<D>>
NativeMapper<'_, D, Policy> {
def_region_kind: DefRegionKind,
permit: Permit,
) -> Result<Any> {
- match self.order {
- WalkOrder::PreOrder => {
+ match PRE_ORDER {
+ true => {
// A replacement inherits the input's permission, established
// before the callback can acquire or release ownership.
let permit = permit.inplace_mode(raw).permit();
@@ -2097,18 +2115,18 @@ impl<D: MapDispatch, Policy: MutContextPolicy<D>>
NativeMapper<'_, D, Policy> {
})
}
}
- WalkOrder::PostOrder => {
+ false => {
let mapped = self.default_map_current_raw(raw,
def_region_kind, permit)?;
- let mapped_raw = if is_unchanged(&mapped) {
- raw
+ let original = StructuralView::from_raw(raw);
+ let value = if is_unchanged(&mapped) {
+ &original
} else {
- *mapped.as_raw_ffi_any()
+ StructuralView::from_any(&mapped)
};
- let value = StructuralView::from_raw(mapped_raw);
// Keep child rewrites when the callback leaves its input
unchanged.
- match self.dispatch.dispatch_map(&value, def_region_kind) {
+ match self.dispatch.dispatch_map(value, def_region_kind) {
Some(Ok(result)) if is_unchanged(&result) => Ok(mapped),
- Some(result) => result.map_err(|error|
with_value_context(error, mapped_raw)),
+ Some(result) => result.map_err(|error|
with_value_context(error, value.raw())),
None => Ok(mapped),
}
}
@@ -2137,6 +2155,21 @@ trait MutationDriver: Sized {
permit: Permit,
) -> Result<Any>;
+ // The ABI entry has already installed and validated the active region.
+ #[inline(always)]
+ fn dispatch_abi_raw<const INPLACE: bool>(
+ &mut self,
+ raw: TVMFFIAny,
+ kind: DefRegionKind,
+ ) -> TVMFFIAny {
+ let permit = if INPLACE {
+ Permit::MaybeInPlace
+ } else {
+ Permit::Copy
+ };
+ result_into_raw(self.dispatch_raw(raw, kind, permit))
+ }
+
fn var_remap_get_raw(&mut self, raw: TVMFFIAny) -> Result<Option<Any>>;
fn var_remap_set_raw(&mut self, raw: TVMFFIAny, replacement: &Any) ->
Result<()>;
@@ -2147,9 +2180,13 @@ trait MutationDriver: Sized {
def_region_kind: DefRegionKind,
permit: Permit,
) -> Result<Option<Any>> {
- let mutator = active_mutator()?;
- with_current_driver_context(mutator, self, || {
- call_registered_structural_mutate(mutator, raw, def_region_kind,
permit)
+ let (mutator, context) = checked_driver_context(self)?;
+ let Some(attr) = structural_mutate_hook(raw, permit) else {
+ // No foreign call needs access to the driver on a hook miss.
+ return Ok(None);
+ };
+ with_current_driver_context(mutator, context, || {
+ call_structural_mutate_hook(mutator, raw, def_region_kind,
attr).map(Some)
})
}
@@ -2162,8 +2199,12 @@ trait MutationDriver: Sized {
default_mutate_driver(self, raw, def_region_kind, permit)
}
- fn map_reflected(&mut self, raw: TVMFFIAny, def_region_kind:
DefRegionKind) -> Result<Any> {
- let type_info = checked_type_info(raw.type_index)?;
+ fn map_reflected(
+ &mut self,
+ raw: TVMFFIAny,
+ def_region_kind: DefRegionKind,
+ type_info: *const crate::tvm_ffi_sys::TVMFFITypeInfo,
+ ) -> Result<Any> {
let seq_hash_kind = unsafe {
if (*type_info).metadata.is_null() {
TVMFFISEqHashKind::kTVMFFISEqHashKindUnsupported as i32
@@ -2289,13 +2330,10 @@ struct StructuralMutatorVTable {
var_remap_set: FStructuralVarRemapSet,
}
-type RuntimeDispatchMutateCallback =
- unsafe fn(*mut c_void, TVMFFIAny, DefRegionKind, Permit) -> Result<Any>;
type RuntimeVarRemapGetCallback = unsafe fn(*mut c_void, TVMFFIAny) ->
Result<Option<Any>>;
type RuntimeVarRemapSetCallback = unsafe fn(*mut c_void, TVMFFIAny, &Any) ->
Result<()>;
struct RuntimeMutatorCallbacks {
- dispatch_mutate: RuntimeDispatchMutateCallback,
var_remap_get: RuntimeVarRemapGetCallback,
var_remap_set: RuntimeVarRemapSetCallback,
}
@@ -2361,12 +2399,16 @@ unsafe impl ObjectCore for RuntimeStructuralMutatorObj {
}
}
-static RUST_STRUCTURAL_MUTATOR_VTABLE: StructuralMutatorVTable =
StructuralMutatorVTable {
- mutate: rust_vtable_mutate,
- maybe_inplace_mutate: rust_vtable_maybe_inplace_mutate,
- var_remap_get: rust_vtable_var_remap_get,
- var_remap_set: rust_vtable_var_remap_set,
-};
+struct RuntimeMutatorVTable<D>(PhantomData<D>);
+
+impl<D: MutationDriver> RuntimeMutatorVTable<D> {
+ const VTABLE: StructuralMutatorVTable = StructuralMutatorVTable {
+ mutate: rust_vtable_mutate::<D>,
+ maybe_inplace_mutate: rust_vtable_maybe_inplace_mutate::<D>,
+ var_remap_get: rust_vtable_var_remap_get,
+ var_remap_set: rust_vtable_var_remap_set,
+ };
+}
struct RuntimeContextGuard {
mutator: StructuralMutatorHandle,
@@ -2409,55 +2451,67 @@ unsafe fn take_runtime_context(mutator:
StructuralMutatorHandle) -> Result<Runti
Ok(RuntimeContextGuard { mutator, context })
}
-unsafe extern "C" fn rust_vtable_mutate(
+unsafe extern "C" fn rust_vtable_mutate<D: MutationDriver>(
mutator: StructuralMutatorHandle,
value: AnyView<'static>,
) -> TVMFFIAny {
// SAFETY: this function is installed only in the vtable of a live
- // RuntimeStructuralMutatorObj; `value` is borrowed for this call.
- rust_vtable_mutate_impl(mutator, value, Permit::Copy)
+ // RuntimeStructuralMutatorObj built for D; `value` is borrowed for this
call.
+ rust_vtable_mutate_impl::<D, false>(mutator, value)
}
-unsafe extern "C" fn rust_vtable_maybe_inplace_mutate(
+unsafe extern "C" fn rust_vtable_maybe_inplace_mutate<D: MutationDriver>(
mutator: StructuralMutatorHandle,
value: AnyView<'static>,
) -> TVMFFIAny {
// SAFETY: same vtable and borrowed-value contract as
// `rust_vtable_mutate`.
- rust_vtable_mutate_impl(mutator, value, Permit::MaybeInPlace)
+ rust_vtable_mutate_impl::<D, true>(mutator, value)
}
-/// Run one erased vtable mutation callback and convert its result to ABI form.
+/// Run one typed vtable mutation callback and convert its result to ABI form.
///
/// # Safety
///
-/// `mutator` must be a live runtime mutator handle, and `value` must remain
-/// valid for this call.
-unsafe fn rust_vtable_mutate_impl(
+/// `mutator` must be a live runtime mutator created for `D`, and `value`
+/// must remain valid for this call.
+#[inline(always)]
+unsafe fn rust_vtable_mutate_impl<D: MutationDriver, const INPLACE: bool>(
mutator: StructuralMutatorHandle,
value: AnyView<'static>,
- permit: Permit,
) -> TVMFFIAny {
let context_guard = match take_runtime_context(mutator) {
Ok(guard) => guard,
Err(error) => return result_into_raw(Err(error)),
};
let context = context_guard.context;
- let callback = (*mutator).callbacks.dispatch_mutate;
let raw = *value.as_raw_ffi_any();
- let outcome = catch_unwind(AssertUnwindSafe(|| {
- let kind = def_region_from_raw((*mutator).def_region_mode)?;
- with_active_mutator(mutator, || callback(context, raw, kind, permit))
- }));
+ let outcome = catch_unwind(AssertUnwindSafe(
+ #[inline]
+ || {
+ let kind = match def_region_from_raw((*mutator).def_region_mode) {
+ Ok(kind) => kind,
+ Err(error) => return result_into_raw(Err(error)),
+ };
+ // take_runtime_context verified that this mutator is already
active.
+ runtime_dispatch_mutate::<D, INPLACE>(context, raw, kind)
+ },
+ ));
match outcome {
- Ok(result) => result_into_raw(result),
- Err(payload) => {
- (*mutator).panic = Some(payload);
- result_into_raw(Err(runtime_error("panic in structural mutator
callback")))
- }
+ Ok(result) => result,
+ Err(payload) => mutation_panic_result(mutator, payload),
}
}
+#[cold]
+unsafe fn mutation_panic_result(
+ mutator: StructuralMutatorHandle,
+ payload: Box<dyn std::any::Any + Send>,
+) -> TVMFFIAny {
+ (*mutator).panic = Some(payload);
+ result_into_raw(Err(runtime_error("panic in structural mutator callback")))
+}
+
unsafe extern "C" fn rust_vtable_var_remap_get(
mutator: StructuralMutatorHandle,
var: AnyView<'static>,
@@ -2608,27 +2662,14 @@ fn inactive_mutator_error(mutator:
StructuralMutatorHandle, operation: &str) ->
/// Expose the current mutable reborrow only for the duration of one registered
/// type hook. Nested vtable calls then reborrow from this pointer, and
/// [`take_runtime_context`] hides it again while Rust is executing.
-fn with_current_driver_context<D, T>(
+fn with_current_driver_context<T>(
mutator: StructuralMutatorHandle,
- driver: &mut D,
+ context: *mut c_void,
callback: impl FnOnce() -> Result<T>,
) -> Result<T> {
- let context = std::ptr::from_mut(driver).cast::<c_void>();
- if !is_active_mutator(mutator) {
- return Err(inactive_mutator_error(mutator, "helper"));
- }
+ // checked_driver_context validated this reborrow; no user code runs
+ // between that check and exposing the pointer to the registered hook.
unsafe {
- if (*mutator).context_identity != context {
- return Err(runtime_error(
- "structural mutator helper called on a non-active mutator",
- ));
- }
- if !(*mutator).context.is_null() {
- return Err(runtime_error(
- "structural mutator driver context is already exposed",
- ));
- }
-
(*mutator).context = context;
struct HideContext {
mutator: StructuralMutatorHandle,
@@ -2645,6 +2686,25 @@ fn with_current_driver_context<D, T>(
}
}
+fn checked_driver_context<D>(driver: &mut D) ->
Result<(StructuralMutatorHandle, *mut c_void)> {
+ let mutator = active_mutator()?;
+ let context = std::ptr::from_mut(driver).cast::<c_void>();
+ unsafe {
+ if (*mutator).context_identity != context {
+ return Err(runtime_error(
+ "structural mutator helper called on a non-active mutator",
+ ));
+ }
+ if !(*mutator).context.is_null() {
+ return Err(runtime_error(
+ "structural mutator driver context is already exposed",
+ ));
+ }
+
+ Ok((mutator, context))
+ }
+}
+
fn def_region_from_raw(kind: i32) -> Result<DefRegionKind> {
match kind {
x if x == DefRegionKind::None as i32 => Ok(DefRegionKind::None),
@@ -2654,7 +2714,9 @@ fn def_region_from_raw(kind: i32) ->
Result<DefRegionKind> {
}
}
-impl<D: MapDispatch, Policy: MutContextPolicy<D>> MutationDriver for
NativeMapper<'_, D, Policy> {
+impl<D: MapDispatch, Policy: MutContextPolicy<D>, const PRE_ORDER: bool>
MutationDriver
+ for NativeMapper<'_, D, Policy, PRE_ORDER>
+{
fn dispatch_raw(
&mut self,
raw: TVMFFIAny,
@@ -2664,6 +2726,21 @@ impl<D: MapDispatch, Policy: MutContextPolicy<D>>
MutationDriver for NativeMappe
with_mutation_region(def_region_kind, |kind| self.map_raw(raw, kind,
permit))
}
+ #[inline(always)]
+ fn dispatch_abi_raw<const INPLACE: bool>(
+ &mut self,
+ raw: TVMFFIAny,
+ kind: DefRegionKind,
+ ) -> TVMFFIAny {
+ let permit = if INPLACE {
+ Permit::MaybeInPlace
+ } else {
+ Permit::Copy
+ };
+ result_into_raw(self.map_raw(raw, kind, permit))
+ }
+
+ #[inline]
fn default_map_current_raw(
&mut self,
raw: TVMFFIAny,
@@ -2694,6 +2771,7 @@ impl<D: MapDispatch, Policy: MutContextPolicy<D>>
MutationDriver for NativeMappe
}
impl<U: StructuralMutator> MutationDriver for U {
+ #[inline]
fn dispatch_raw(
&mut self,
raw: TVMFFIAny,
@@ -2703,6 +2781,26 @@ impl<U: StructuralMutator> MutationDriver for U {
dispatch_user_raw(self, raw, def_region_kind, permit)
}
+ #[inline(always)]
+ fn dispatch_abi_raw<const INPLACE: bool>(
+ &mut self,
+ raw: TVMFFIAny,
+ kind: DefRegionKind,
+ ) -> TVMFFIAny {
+ if INPLACE && raw.type_index <
TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
+ return result_into_raw(
+ self.dispatch_mutate(&StructuralView::from_raw(raw), kind)
+ .map_err(|error| with_value_context(error, raw)),
+ );
+ }
+ let permit = if INPLACE {
+ Permit::MaybeInPlace
+ } else {
+ Permit::Copy
+ };
+ result_into_raw(dispatch_user_current_raw(self, raw, kind, permit))
+ }
+
fn var_remap_get_raw(&mut self, raw: TVMFFIAny) -> Result<Option<Any>> {
self.var_remap_get(&StructuralView::from_raw(raw))
}
@@ -2718,13 +2816,13 @@ impl<U: StructuralMutator> MutationDriver for U {
///
/// `context` must come from the current mutable reborrow of a live `D`; the
/// runtime object hides that pointer until this call returns.
-unsafe fn runtime_dispatch_mutate<D: MutationDriver>(
+#[inline(always)]
+unsafe fn runtime_dispatch_mutate<D: MutationDriver, const INPLACE: bool>(
context: *mut c_void,
raw: TVMFFIAny,
def_region_kind: DefRegionKind,
- permit: Permit,
-) -> Result<Any> {
- (&mut *context.cast::<D>()).dispatch_raw(raw, def_region_kind, permit)
+) -> TVMFFIAny {
+ (&mut *context.cast::<D>()).dispatch_abi_raw::<INPLACE>(raw,
def_region_kind)
}
/// Dispatch a variable-remap lookup through the erased driver context.
@@ -2756,21 +2854,26 @@ unsafe fn runtime_var_remap_set<D: MutationDriver>(
fn run_structural_mutator<D: MutationDriver>(root: Any, driver: &mut D) ->
Result<Any> {
let context = std::ptr::from_mut(driver).cast::<c_void>();
let callbacks = RuntimeMutatorCallbacks {
- dispatch_mutate: runtime_dispatch_mutate::<D>,
var_remap_get: runtime_var_remap_get::<D>,
var_remap_set: runtime_var_remap_set::<D>,
};
- run_structural_mutator_with_context(root, context, callbacks)
+ run_structural_mutator_with_context(
+ root,
+ context,
+ callbacks,
+ &RuntimeMutatorVTable::<D>::VTABLE,
+ )
}
fn run_structural_mutator_with_context(
root: Any,
context: *mut c_void,
callbacks: RuntimeMutatorCallbacks,
+ vtable: &'static StructuralMutatorVTable,
) -> Result<Any> {
let mut active = ObjectArc::new(RuntimeStructuralMutatorObj {
base: Object::new(),
- vtable: &RUST_STRUCTURAL_MUTATOR_VTABLE,
+ vtable,
def_region_mode: DefRegionKind::None as i32,
context,
context_identity: context,
@@ -2822,18 +2925,14 @@ fn call_mutator(
(*(*mutator).vtable).mutate
}
};
- with_mutator_def_region(mutator, def_region_kind, || unsafe {
+ with_mutator_def_region(mutator, def_region_kind, |_| unsafe {
let view = AnyView::from_raw_ffi_any(raw);
result_from_raw(callback(mutator, view))
})
}
-fn call_registered_structural_mutate(
- mutator: StructuralMutatorHandle,
- raw: TVMFFIAny,
- def_region_kind: DefRegionKind,
- permit: Permit,
-) -> Result<Option<Any>> {
+#[inline]
+fn structural_mutate_hook(raw: TVMFFIAny, permit: Permit) -> Option<TVMFFIAny>
{
let use_inplace = permit == Permit::MaybeInPlace && object_is_unique(raw);
if use_inplace {
if let Some(attr) = structural_maybe_inplace_mutate_column()
@@ -2842,19 +2941,12 @@ fn call_registered_structural_mutate(
if attr.type_index == TVMFFITypeIndex::kTVMFFIOpaquePtr as i32
|| attr.type_index == TVMFFITypeIndex::kTVMFFIFunction as i32
{
- return call_structural_mutate_hook(mutator, raw,
def_region_kind, attr).map(Some);
+ return Some(attr);
}
}
}
-
- let Some(attr) = structural_mutate_column().and_then(|column|
column.get_raw(raw.type_index))
- else {
- return Ok(None);
- };
- if attr.type_index == TVMFFITypeIndex::kTVMFFINone as i32 {
- return Ok(None);
- }
- call_structural_mutate_hook(mutator, raw, def_region_kind, attr).map(Some)
+ let attr = structural_mutate_column().and_then(|column|
column.get_raw(raw.type_index))?;
+ (attr.type_index != TVMFFITypeIndex::kTVMFFINone as i32).then_some(attr)
}
fn call_structural_mutate_hook(
@@ -2863,7 +2955,7 @@ fn call_structural_mutate_hook(
def_region_kind: DefRegionKind,
attr: TVMFFIAny,
) -> Result<Any> {
- with_mutator_def_region(mutator, def_region_kind, || unsafe {
+ with_mutator_def_region(mutator, def_region_kind, |_| unsafe {
match attr.type_index {
x if x == TVMFFITypeIndex::kTVMFFIOpaquePtr as i32 => {
let pointer = attr.data_union.v_ptr;
@@ -2907,6 +2999,7 @@ unsafe fn borrowed_mutator_view<'a>(mutator:
StructuralMutatorHandle) -> AnyView
AnyView::from_raw_ffi_any(raw)
}
+#[inline(always)]
fn result_into_raw(result: Result<Any>) -> TVMFFIAny {
unsafe {
match result {
@@ -2918,6 +3011,7 @@ fn result_into_raw(result: Result<Any>) -> TVMFFIAny {
/// Resolve only at an owning-value API boundary; internal Any carriers and
/// native hooks keep the unchanged tag to avoid acquiring the original.
+#[inline]
fn resolve_result(result: Any, original: TVMFFIAny) -> Result<Any> {
if is_unchanged(&result) {
owned_from_raw(original)
@@ -2932,6 +3026,7 @@ fn resolve_result(result: Any, original: TVMFFIAny) ->
Result<Any> {
///
/// `raw` must contain one owning TVMFFIAny result that has not already been
/// consumed. An Error object is converted into the Rust error channel.
+#[inline(always)]
unsafe fn result_from_raw(raw: TVMFFIAny) -> Result<Any> {
let value = Any::from_raw_ffi_any(raw);
if value.type_index() != TVMFFITypeIndex::kTVMFFIError as i32 {
@@ -2942,42 +3037,45 @@ unsafe fn result_from_raw(raw: TVMFFIAny) ->
Result<Any> {
}
}
+#[inline(always)]
fn with_mutator_def_region<T>(
mutator: StructuralMutatorHandle,
kind: DefRegionKind,
- callback: impl FnOnce() -> T,
+ callback: impl FnOnce(DefRegionKind) -> T,
) -> T {
unsafe {
let previous = (*mutator).def_region_mode;
// Precedence: a pattern region propagates; entering any kind inside
it has no effect.
- if previous == DefRegionKind::Pattern as i32 {
- return callback();
- }
- (*mutator).def_region_mode = kind as i32;
+ let effective = if previous == DefRegionKind::Pattern as i32 {
+ DefRegionKind::Pattern
+ } else {
+ (*mutator).def_region_mode = kind as i32;
+ kind
+ };
struct Restore {
mutator: StructuralMutatorHandle,
previous: i32,
}
impl Drop for Restore {
fn drop(&mut self) {
- unsafe { (*self.mutator).def_region_mode = self.previous };
+ if self.previous != DefRegionKind::Pattern as i32 {
+ unsafe { (*self.mutator).def_region_mode = self.previous };
+ }
}
}
let _restore = Restore { mutator, previous };
- callback()
+ // One call site keeps the continuation visible to the inliner.
+ callback(effective)
}
}
+#[inline(always)]
fn with_mutation_region<T>(
kind: DefRegionKind,
callback: impl FnOnce(DefRegionKind) -> Result<T>,
) -> Result<T> {
let mutator = active_mutator()?;
- with_mutator_def_region(mutator, kind, || {
- // SAFETY: the active invocation keeps this thread's ABI mutator alive.
- let effective = def_region_from_raw(unsafe {
(*mutator).def_region_mode })?;
- callback(effective)
- })
+ with_mutator_def_region(mutator, kind, callback)
}
#[inline(always)]
@@ -2987,15 +3085,27 @@ fn dispatch_user_raw<U: StructuralMutator>(
def_region_kind: DefRegionKind,
permit: Permit,
) -> Result<Any> {
- with_mutation_region(def_region_kind, |kind| {
- let result = if permit == Permit::MaybeInPlace &&
object_is_unique(raw) {
- let mut scoped_raw = raw;
- mutator.dispatch_maybe_inplace_mutate(InplaceValue::from_raw(&mut
scoped_raw), kind)
- } else {
- mutator.dispatch_mutate(&StructuralView::from_raw(raw), kind)
- };
- result.map_err(|error| with_value_context(error, raw))
- })
+ with_mutation_region(
+ def_region_kind,
+ #[inline(always)]
+ |kind| dispatch_user_current_raw(mutator, raw, kind, permit),
+ )
+}
+
+#[inline(always)]
+fn dispatch_user_current_raw<U: StructuralMutator>(
+ mutator: &mut U,
+ raw: TVMFFIAny,
+ kind: DefRegionKind,
+ permit: Permit,
+) -> Result<Any> {
+ let result = if permit == Permit::MaybeInPlace && object_is_unique(raw) {
+ let mut scoped_raw = raw;
+ mutator.dispatch_maybe_inplace_mutate(InplaceValue::from_raw(&mut
scoped_raw), kind)
+ } else {
+ mutator.dispatch_mutate(&StructuralView::from_raw(raw), kind)
+ };
+ result.map_err(|error| with_value_context(error, raw))
}
fn user_default_mutate<U: StructuralMutator>(
@@ -3037,7 +3147,14 @@ fn default_mutate_driver_impl<D: MutationDriver>(
return owned_from_raw(raw);
}
- let kind = structural_hash_kind(raw)?;
+ let type_info = checked_type_info(raw.type_index)?;
+ let kind = unsafe {
+ if (*type_info).metadata.is_null() {
+ None
+ } else {
+ Some((*(*type_info).metadata).structural_eq_hash_kind)
+ }
+ };
let is_free_var = kind ==
Some(TVMFFISEqHashKind::kTVMFFISEqHashKindFreeVar as i32);
let is_dag_node = kind ==
Some(TVMFFISEqHashKind::kTVMFFISEqHashKindDAGNode as i32);
if is_free_var || is_dag_node {
@@ -3053,7 +3170,7 @@ fn default_mutate_driver_impl<D: MutationDriver>(
return Ok(Unchanged.into());
}
- let result = driver.map_reflected(raw, def_region_kind)?;
+ let result = driver.map_reflected(raw, def_region_kind, type_info)?;
if is_dag_node
|| (is_free_var && (def_region_kind == DefRegionKind::Pattern ||
!is_unchanged(&result)))
{
@@ -3116,12 +3233,18 @@ fn shallow_copy(raw: TVMFFIAny) -> Result<Any> {
"",
));
}
- let function = Function::try_from(unsafe { AnyView::from_raw_ffi_any(attr)
})?;
- // `raw` is borrowed from the active mutation call and remains valid for
- // this synchronous packed call. Avoid an unnecessary object refcount
- // increment/decrement just to pass another borrowed view.
- let source = unsafe { AnyView::from_raw_ffi_any(raw) };
- let result = function.call_packed(&[source])?;
+ let function = unsafe { attr.data_union.v_obj };
+ if function.is_null() {
+ return Err(runtime_error("shallow copy function pointer is null"));
+ }
+ // The registry owns this function throughout the call; borrowing it avoids
+ // an atomic retain/release for every reflected node.
+ let mut result = Any::new();
+ let status =
+ unsafe { TVMFFIFunctionCall(function.cast(), &raw, 1,
Any::as_data_ptr(&mut result)) };
+ if status != 0 {
+ return Err(Error::from_raised());
+ }
let result_raw = *result.as_raw_ffi_any();
let result_pointer = unsafe { result_raw.data_union.v_obj };
let source_pointer = unsafe { raw.data_union.v_obj };
@@ -3181,44 +3304,41 @@ fn call_field_setter(
}
}
-fn structural_hash_kind(raw: TVMFFIAny) -> Result<Option<i32>> {
- if raw.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
- return Ok(None);
- }
- let type_info = checked_type_info(raw.type_index)?;
- unsafe {
- if (*type_info).metadata.is_null() {
- Ok(None)
- } else {
- Ok(Some((*(*type_info).metadata).structural_eq_hash_kind))
- }
- }
-}
-
fn object_identity_key(raw: TVMFFIAny) -> Result<NonNull<TVMFFIObject>> {
if raw.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
- return Err(Error::new(
- TYPE_ERROR,
- "variable-remap keys must be object-backed values",
- "",
- ));
+ return Err(var_remap_key_error());
}
let pointer = unsafe { raw.data_union.v_obj };
NonNull::new(pointer)
.ok_or_else(|| runtime_error("native structural map: identity object
has a null pointer"))
}
+#[cold]
+fn var_remap_key_error() -> Error {
+ Error::new(
+ TYPE_ERROR,
+ "variable-remap keys must be object-backed values",
+ "",
+ )
+}
+
+#[inline]
fn checked_type_info(type_index: i32) -> Result<*const
crate::tvm_ffi_sys::TVMFFITypeInfo> {
let info = unsafe { TVMFFIGetTypeInfo(type_index) };
if info.is_null() {
- Err(runtime_error(&format!(
- "native structural map: unregistered type index {type_index}"
- )))
+ Err(unregistered_type_error(type_index))
} else {
Ok(info)
}
}
+#[cold]
+fn unregistered_type_error(type_index: i32) -> Error {
+ runtime_error(&format!(
+ "native structural map: unregistered type index {type_index}"
+ ))
+}
+
#[inline]
fn object_is_unique(raw: TVMFFIAny) -> bool {
if raw.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
@@ -3228,6 +3348,7 @@ fn object_is_unique(raw: TVMFFIAny) -> bool {
!pointer.is_null() && unsafe { object::unsafe_::strong_count(pointer) == 1
}
}
+#[inline]
fn owned_from_raw(raw: TVMFFIAny) -> Result<Any> {
if let Some(owned) = try_to_owned_without_normalization(raw) {
return Ok(owned);
@@ -3250,6 +3371,7 @@ fn owned_from_raw(raw: TVMFFIAny) -> Result<Any> {
}
}
+#[cold]
fn with_value_context(error: Error, raw: TVMFFIAny) -> Error {
if raw.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
error
@@ -3261,10 +3383,13 @@ fn with_value_context(error: Error, raw: TVMFFIAny) ->
Error {
}
}
+#[cold]
fn with_error_context(error: Error, frame: &str) -> Error {
with_structural_error_context(error, "map", frame)
}
+#[cold]
+#[inline(never)]
fn runtime_error(message: &str) -> Error {
Error::new(RUNTIME_ERROR, message, "")
}
diff --git a/rust/tvm-ffi/src/extra/structural_mutate/policy.rs
b/rust/tvm-ffi/src/extra/structural_mutate/policy.rs
index 4dc2e092..42ac3827 100644
--- a/rust/tvm-ffi/src/extra/structural_mutate/policy.rs
+++ b/rust/tvm-ffi/src/extra/structural_mutate/policy.rs
@@ -72,6 +72,7 @@ impl<State, Outer: MutContextPolicy<State>, Inner:
MutContextPolicy<State>> MutC
}
}
+#[inline(always)]
pub(super) fn mutate_with_policy<State>(
driver: &mut dyn MutateContextDriver<State>,
policy: &impl MutContextPolicy<State>,
@@ -79,15 +80,19 @@ pub(super) fn mutate_with_policy<State>(
kind: DefRegionKind,
) -> Result<Any> {
let raw = value.value.raw();
- with_mutation_region(kind, |kind| {
- let mut ctx = MutateContext {
- driver,
- def_region_kind: kind,
- inplace_mode: value.inplace_mode(),
- _not_send_sync: PhantomData,
- };
- policy.default_mutate(value, &mut ctx).map(Any::from)
- })
+ with_mutation_region(
+ kind,
+ #[inline(always)]
+ |kind| {
+ let mut ctx = MutateContext {
+ driver,
+ def_region_kind: kind,
+ inplace_mode: value.inplace_mode(),
+ _not_send_sync: PhantomData,
+ };
+ policy.default_mutate(value, &mut ctx).map(Any::from)
+ },
+ )
.map_err(|error| with_value_context(error, raw))
}
@@ -146,22 +151,17 @@ impl<State, Driver: MutationDriver +
MutateCallbackState<State>> MutateContextDr
self.driver.callback_state_mut()
}
fn mutate_borrowed(&mut self, value: AnyView<'_>, kind: DefRegionKind) ->
Result<Any> {
- with_mutation_region(kind, |kind| {
- self.driver
- .dispatch_raw(*value.as_raw_ffi_any(), kind, Permit::Copy)
- })
+ self.driver
+ .dispatch_raw(*value.as_raw_ffi_any(), kind, Permit::Copy)
}
fn mutate_owned(&mut self, value: Any, kind: DefRegionKind, mode:
InplaceMode) -> Result<Any> {
let raw = *value.as_raw_ffi_any();
- let result = with_mutation_region(kind, |kind| {
- self.driver.dispatch_raw(raw, kind, mode.permit())
- })?;
+ let result = self.driver.dispatch_raw(raw, kind, mode.permit())?;
Ok(if is_unchanged(&result) { value } else { result })
}
fn default_mutate_borrowed(&mut self, value: AnyView<'_>, kind:
DefRegionKind) -> Result<Any> {
- with_mutation_region(kind, |kind| {
- default_mutate_driver(self.driver, *value.as_raw_ffi_any(), kind,
Permit::Copy)
- })
+ // The policy context has already installed this definition region.
+ default_mutate_driver(self.driver, *value.as_raw_ffi_any(), kind,
Permit::Copy)
}
fn default_mutate_value(
&mut self,
@@ -170,9 +170,7 @@ impl<State, Driver: MutationDriver +
MutateCallbackState<State>> MutateContextDr
mode: InplaceMode,
) -> Result<Any> {
let permit = value.permit(mode);
- with_mutation_region(kind, |kind| {
- default_mutate_driver(self.driver, value.value.raw(), kind, permit)
- })
+ default_mutate_driver(self.driver, value.value.raw(), kind, permit)
}
fn var_remap_get(&mut self, var: &StructuralView) -> Result<Option<Any>> {
self.driver.var_remap_get_raw(var.raw())
@@ -182,7 +180,9 @@ impl<State, Driver: MutationDriver +
MutateCallbackState<State>> MutateContextDr
}
}
-impl<D, Policy> MutateCallbackState<D> for NativeMapper<'_, D, Policy> {
+impl<D, Policy, const PRE_ORDER: bool> MutateCallbackState<D>
+ for NativeMapper<'_, D, Policy, PRE_ORDER>
+{
fn callback_state(&self) -> &D {
self.dispatch
}
@@ -231,15 +231,7 @@ impl<Mapper: MapDispatch, Policy:
MutContextPolicy<Mapper>> NativeMap
for MapWithContextPolicy<Mapper, Policy>
{
fn map_root(&mut self, root: Any, order: WalkOrder) -> Result<Any> {
- run_structural_mutator(
- root,
- &mut NativeMapper {
- dispatch: &mut self.mapper,
- order,
- policy: Some(self.policy.clone()),
- remap: StructuralVarRemap::default(),
- },
- )
+ run_native_mapper(root, &mut self.mapper, Some(self.policy.clone()),
order)
}
}
diff --git a/rust/tvm-ffi/src/extra/structural_visit.rs
b/rust/tvm-ffi/src/extra/structural_visit.rs
index 06ab7896..603b6001 100644
--- a/rust/tvm-ffi/src/extra/structural_visit.rs
+++ b/rust/tvm-ffi/src/extra/structural_visit.rs
@@ -887,7 +887,7 @@ where
) -> Result<WalkResult> {
self.link
.try_call(value, def_region_kind)
- .unwrap_or(Ok(WalkResult::Advance))
+ .unwrap_or_else(|| Ok(WalkResult::Advance))
}
}
@@ -1207,7 +1207,11 @@ fn default_walk_children<V: NativeVisit, const
PRE_ORDER: bool>(
/// Action applied to each child found by the shared traversal.
trait ChildVisit {
- fn visit_child(&mut self, child: TVMFFIAny, def_region_kind:
DefRegionKind) -> NativeResult;
+ fn visit_child(
+ &mut self,
+ child: &StructuralView,
+ def_region_kind: DefRegionKind,
+ ) -> NativeResult;
}
struct WalkChildren<'a, V, const PRE_ORDER: bool> {
@@ -1215,8 +1219,13 @@ struct WalkChildren<'a, V, const PRE_ORDER: bool> {
}
impl<V: NativeVisit, const PRE_ORDER: bool> ChildVisit for WalkChildren<'_, V,
PRE_ORDER> {
- fn visit_child(&mut self, child: TVMFFIAny, def_region_kind:
DefRegionKind) -> NativeResult {
- visit_raw::<V, PRE_ORDER>(child, self.visitor, def_region_kind)
+ #[inline(always)]
+ fn visit_child(
+ &mut self,
+ child: &StructuralView,
+ def_region_kind: DefRegionKind,
+ ) -> NativeResult {
+ visit_raw::<V, PRE_ORDER>(child.raw(), self.visitor, def_region_kind)
}
}
@@ -1226,16 +1235,18 @@ struct UserChildren<'a, V> {
impl<V: StructuralVisitor> ChildVisit for UserChildren<'_, V> {
#[inline]
- fn visit_child(&mut self, child: TVMFFIAny, def_region_kind:
DefRegionKind) -> NativeResult {
- if child.type_index == TVMFFITypeIndex::kTVMFFINone as i32 {
+ fn visit_child(
+ &mut self,
+ child: &StructuralView,
+ def_region_kind: DefRegionKind,
+ ) -> NativeResult {
+ if child.type_index() == TVMFFITypeIndex::kTVMFFINone as i32 {
return Ok(());
}
- match with_visit_region(def_region_kind, |kind| {
- self.visitor.visit(&StructuralView::from_raw(child), kind)
- }) {
+ match with_visit_region(def_region_kind, |kind|
self.visitor.visit(child, kind)) {
Ok(None) => Ok(()),
Ok(Some(interrupt)) => Err(NativeHalt::Interrupt(interrupt.value)),
- Err(error) => Err(with_value_context(NativeHalt::Error(error),
child)),
+ Err(error) => Err(with_value_context(NativeHalt::Error(error),
child.raw())),
}
}
}
@@ -1250,15 +1261,16 @@ fn visit_raw<V: NativeVisit, const PRE_ORDER: bool>(
if value.type_index == TVMFFITypeIndex::kTVMFFINone as i32 {
return Ok(());
}
-
let visit_value = StructuralView::from_raw(value);
if PRE_ORDER {
- match visitor.visit(&visit_value, def_region_kind) {
- Ok(WalkResult::Advance) => {}
- Ok(WalkResult::Skip) => return Ok(()),
- Ok(WalkResult::Interrupt) => return
Err(NativeHalt::Interrupt(Any::new())),
- Ok(WalkResult::InterruptWith(payload)) => return
Err(NativeHalt::Interrupt(payload)),
- Err(error) => return Err(with_value_context(error.into(), value)),
+ let action = visitor
+ .visit(&visit_value, def_region_kind)
+ .map_err(|error| with_value_context(error.into(), value))?;
+ match action {
+ WalkResult::Advance => {}
+ WalkResult::Skip => return Ok(()),
+ WalkResult::Interrupt => return
Err(NativeHalt::Interrupt(Any::new())),
+ WalkResult::InterruptWith(payload) => return
Err(NativeHalt::Interrupt(payload)),
}
}
@@ -1268,7 +1280,9 @@ fn visit_raw<V: NativeVisit, const PRE_ORDER: bool>(
Ok(Some(interrupt)) => return
Err(NativeHalt::Interrupt(interrupt.value)),
Err(error) => return Err(with_value_context(error.into(), value)),
}
- } else {
+ } else if value.type_index >= TVMFFITypeIndex::kTVMFFIStaticObjectBegin as
i32
+ || has_registered_visit_hook(value.type_index)
+ {
// Preserve the raw-result path for ordinary walkers.
let context = std::ptr::from_mut(&mut *visitor).cast::<c_void>();
let children = &mut WalkChildren::<V, PRE_ORDER> { visitor };
@@ -1281,9 +1295,9 @@ fn visit_raw<V: NativeVisit, const PRE_ORDER: bool>(
Ok(())
} else {
match visitor.visit(&visit_value, def_region_kind) {
+ Ok(WalkResult::Advance | WalkResult::Skip) => Ok(()),
Ok(WalkResult::Interrupt) =>
Err(NativeHalt::Interrupt(Any::new())),
Ok(WalkResult::InterruptWith(payload)) =>
Err(NativeHalt::Interrupt(payload)),
- Ok(WalkResult::Advance | WalkResult::Skip) => Ok(()),
Err(error) => Err(with_value_context(error.into(), value)),
}
}
@@ -1378,7 +1392,7 @@ unsafe fn visit_reflected_field<C: ChildVisit>(
));
}
- let borrowed = raw_of_owned(&child);
+ let borrowed = StructuralView::from_any(&child);
let child_region = field_def_region(field, inherited_region);
visitor
.visit_child(borrowed, child_region)
@@ -1555,21 +1569,32 @@ unsafe fn rust_vtable_visit_impl(
};
let context = context_guard.context;
let raw = *value.as_raw_ffi_any();
- let outcome = catch_unwind(AssertUnwindSafe(|| {
- let kind = def_region_from_raw((*visitor).def_region_mode)?;
- callback(context, raw, kind)
- }));
+ let outcome = catch_unwind(AssertUnwindSafe(
+ #[inline(always)]
+ || {
+ let result = def_region_from_raw((*visitor).def_region_mode)
+ .map_err(NativeHalt::Error)
+ .and_then(|kind| callback(context, raw, kind));
+ native_result_into_raw(result)
+ },
+ ));
match outcome {
- Ok(result) => native_result_into_raw(result),
- Err(payload) => {
- (*visitor).panic = Some(payload);
- native_result_into_raw(Err(NativeHalt::Error(runtime_error(
- "panic in structural visitor callback",
- ))))
- }
+ Ok(result) => result,
+ Err(payload) => visit_panic_result(visitor, payload),
}
}
+#[cold]
+unsafe fn visit_panic_result(
+ visitor: StructuralVisitorHandle,
+ payload: Box<dyn std::any::Any + Send>,
+) -> TVMFFIAny {
+ (*visitor).panic = Some(payload);
+ native_result_into_raw(Err(NativeHalt::Error(runtime_error(
+ "panic in structural visitor callback",
+ ))))
+}
+
thread_local! {
static ACTIVE_STRUCTURAL_VISITOR: Cell<*mut ActiveStructuralVisitor> =
const {
Cell::new(std::ptr::null_mut())
@@ -1654,18 +1679,8 @@ fn with_current_visitor_context(
context: *mut c_void,
callback: impl FnOnce() -> NativeResult,
) -> NativeResult {
- let active = active_structural_visitor_state(visitor)
- .ok_or_else(|| inactive_structural_visitor_error(visitor, "helper"))?;
+ let active = checked_visitor_context(visitor, context)?;
unsafe {
- if (*active).context_identity != context {
- return Err(
- runtime_error("structural visitor helper called on a
non-active visitor").into(),
- );
- }
- if !(*active).context.is_null() {
- return Err(runtime_error("structural visitor context is already
exposed").into());
- }
-
(*active).context = context;
struct HideContext {
active: *mut ActiveStructuralVisitor,
@@ -1680,45 +1695,33 @@ fn with_current_visitor_context(
}
}
+fn checked_visitor_context(
+ visitor: StructuralVisitorHandle,
+ context: *mut c_void,
+) -> Result<*mut ActiveStructuralVisitor> {
+ let active = active_structural_visitor_state(visitor)
+ .ok_or_else(|| inactive_structural_visitor_error(visitor, "helper"))?;
+ unsafe {
+ if (*active).context_identity != context {
+ return Err(runtime_error(
+ "structural visitor helper called on a non-active visitor",
+ ));
+ }
+ if !(*active).context.is_null() {
+ return Err(runtime_error(
+ "structural visitor context is already exposed",
+ ));
+ }
+ Ok(active)
+ }
+}
+
#[inline(always)]
unsafe fn runtime_walk<V: NativeVisit, const PRE_ORDER: bool>(
context: *mut c_void,
raw: TVMFFIAny,
def_region_kind: DefRegionKind,
) -> NativeResult {
- if raw.type_index == TVMFFITypeIndex::kTVMFFINone as i32 {
- return Ok(());
- }
- if !V::CUSTOM_DESCENT && raw.type_index <
TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
- let visitor = &mut *context.cast::<V>();
- if PRE_ORDER {
- match visitor.visit(&StructuralView::from_raw(raw),
def_region_kind) {
- Ok(WalkResult::Advance) => {}
- Ok(WalkResult::Skip) => return Ok(()),
- Ok(WalkResult::Interrupt) => return
Err(NativeHalt::Interrupt(Any::new())),
- Ok(WalkResult::InterruptWith(payload)) => {
- return Err(NativeHalt::Interrupt(payload));
- }
- Err(error) => return Err(with_value_context(error.into(),
raw)),
- }
- if !has_registered_visit_hook(raw.type_index) {
- return Ok(());
- }
- let children = &mut WalkChildren::<V, PRE_ORDER> { visitor };
- return visit_children_raw(raw, children, context, def_region_kind)
- .map_err(|halt| with_value_context(halt, raw));
- }
- // Post-order inline values have no children unless their type
- // registered a visit hook. Handle the common case directly here.
- if !has_registered_visit_hook(raw.type_index) {
- return match visitor.visit(&StructuralView::from_raw(raw),
def_region_kind) {
- Ok(WalkResult::Advance | WalkResult::Skip) => Ok(()),
- Ok(WalkResult::Interrupt) =>
Err(NativeHalt::Interrupt(Any::new())),
- Ok(WalkResult::InterruptWith(payload)) =>
Err(NativeHalt::Interrupt(payload)),
- Err(error) => Err(with_value_context(error.into(), raw)),
- };
- }
- }
visit_raw::<V, PRE_ORDER>(raw, &mut *context.cast::<V>(), def_region_kind)
}
@@ -1890,18 +1893,25 @@ fn visit_result_from_any(value: Any) -> NativeResult {
}
}
+#[inline(always)]
fn with_visit_region<T>(
kind: DefRegionKind,
callback: impl FnOnce(DefRegionKind) -> Result<T>,
) -> Result<T> {
let active = active_structural_visitor()?;
- with_visitor_def_region(active, kind, || {
- // SAFETY: the active invocation keeps this thread's ABI visitor alive.
- let kind = def_region_from_raw(unsafe { (*active).def_region_mode })?;
- callback(kind)
- })
+ with_visitor_def_region(
+ active,
+ kind,
+ #[inline(always)]
+ || {
+ // SAFETY: the active invocation keeps this thread's ABI visitor
alive.
+ let kind = def_region_from_raw(unsafe { (*active).def_region_mode
})?;
+ callback(kind)
+ },
+ )
}
+#[inline(always)]
fn with_visitor_def_region<T>(
visitor: StructuralVisitorHandle,
kind: DefRegionKind,
@@ -1910,17 +1920,18 @@ fn with_visitor_def_region<T>(
unsafe {
let previous = (*visitor).def_region_mode;
// Precedence: a pattern region propagates; entering any kind inside
it has no effect.
- if previous == DefRegionKind::Pattern as i32 {
- return callback();
+ if previous != DefRegionKind::Pattern as i32 {
+ (*visitor).def_region_mode = kind as i32;
}
- (*visitor).def_region_mode = kind as i32;
struct Restore {
visitor: StructuralVisitorHandle,
previous: i32,
}
impl Drop for Restore {
fn drop(&mut self) {
- unsafe { (*self.visitor).def_region_mode = self.previous };
+ if self.previous != DefRegionKind::Pattern as i32 {
+ unsafe { (*self.visitor).def_region_mode = self.previous };
+ }
}
}
let _restore = Restore { visitor, previous };
@@ -1938,6 +1949,7 @@ fn def_region_from_raw(kind: i32) ->
Result<DefRegionKind> {
}
}
+#[cold]
fn with_value_context(halt: NativeHalt, value: TVMFFIAny) -> NativeHalt {
if value.type_index < TVMFFITypeIndex::kTVMFFIStaticObjectBegin as i32 {
halt
@@ -2028,6 +2040,7 @@ pub(crate) fn free_var_child_region(
}
}
+#[cold]
fn with_error_context(halt: NativeHalt, frame: &str) -> NativeHalt {
match halt {
NativeHalt::Error(error) => {
@@ -2130,8 +2143,3 @@ unsafe fn visit_field_level<B>(
fn raw_of(view: AnyView<'_>) -> TVMFFIAny {
*view.as_raw_ffi_any()
}
-
-#[inline]
-fn raw_of_owned(any: &Any) -> TVMFFIAny {
- *any.as_raw_ffi_any()
-}
diff --git a/rust/tvm-ffi/src/extra/structural_visit/policy.rs
b/rust/tvm-ffi/src/extra/structural_visit/policy.rs
index 7ffcd2aa..e94615d2 100644
--- a/rust/tvm-ffi/src/extra/structural_visit/policy.rs
+++ b/rust/tvm-ffi/src/extra/structural_visit/policy.rs
@@ -76,23 +76,28 @@ impl<State, Outer: ContextPolicy<State>, Inner:
ContextPolicy<State>> ContextPol
}
}
+#[inline(always)]
pub(super) fn visit_with_policy<State>(
driver: &mut dyn VisitContextDriver<State>,
policy: &impl ContextPolicy<State>,
value: &StructuralView,
def_region_kind: DefRegionKind,
) -> Result<Option<VisitInterrupt>> {
- with_visit_region(def_region_kind, |def_region_kind| {
- policy.default_visit(
- value,
- &mut VisitContext {
- driver,
- current: StructuralView::from_raw(value.raw()),
- def_region_kind,
- _not_send_sync: PhantomData,
- },
- )
- })
+ with_visit_region(
+ def_region_kind,
+ #[inline(always)]
+ |def_region_kind| {
+ policy.default_visit(
+ value,
+ &mut VisitContext {
+ driver,
+ current: StructuralView::from_raw(value.raw()),
+ def_region_kind,
+ _not_send_sync: PhantomData,
+ },
+ )
+ },
+ )
}
struct NextPolicy<'a, State, Policy> {
@@ -216,7 +221,7 @@ impl<Walker: WalkDispatch, Policy: ContextPolicy<Walker>>
NativeVisit
fn visit(&mut self, value: &StructuralView, kind: DefRegionKind) ->
Result<WalkResult> {
self.walker
.dispatch_walk(value, kind)
- .unwrap_or(Ok(WalkResult::Advance))
+ .unwrap_or_else(|| Ok(WalkResult::Advance))
}
fn default_visit_children<const PRE_ORDER: bool>(