This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch tvmscript-generic-parser-builder in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 5d0bd83a44ab97a8c689c4671af787b58f296c57 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 03:37:55 2026 +0000 Release Python callbacks safely during registry cleanup --- python/tvm/runtime/__init__.py | 9 +++++++++ src/runtime/vm/builtin.cc | 25 ++++++++++++++++++++----- tests/python/relax/test_relax_operators.py | 21 +++++++++++++++++++++ 3 files changed, 50 insertions(+), 5 deletions(-) diff --git a/python/tvm/runtime/__init__.py b/python/tvm/runtime/__init__.py index 56b904a519..6b220aaeb1 100644 --- a/python/tvm/runtime/__init__.py +++ b/python/tvm/runtime/__init__.py @@ -17,7 +17,10 @@ # under the License. """TVM runtime namespace.""" +import atexit as _atexit + from tvm_ffi import convert, Object +from tvm_ffi import get_global_func as _get_global_func from tvm_ffi._dtype import dtype as DataType, DataTypeCode # Import _ffi_node_api for its side effect of installing AsRepr as @@ -52,3 +55,9 @@ except (ImportError, ValueError): disco = None # type: ignore[assignment] from tvm_ffi import Shape as ShapeTuple + +# Release Python callbacks while their interpreter is still alive, before native +# static destruction. The VM extension is optional in runtime-only builds. +_clear_py_func_registry = _get_global_func("vm.builtin.clear_py_func_registry", allow_missing=True) +if _clear_py_func_registry is not None: + _atexit.register(_clear_py_func_registry) diff --git a/src/runtime/vm/builtin.cc b/src/runtime/vm/builtin.cc index fa5f5ce7ce..fc0544f868 100644 --- a/src/runtime/vm/builtin.cc +++ b/src/runtime/vm/builtin.cc @@ -458,19 +458,34 @@ TVM_FFI_STATIC_INIT_BLOCK() { //------------------------------------- // Global registry for Python functions -static std::unordered_map<std::string, ffi::Function> py_func_registry; +struct PyFuncRegistry { + std::unordered_map<std::string, ffi::Function> functions; + + void Clear() { + // Releasing a callback can run an owner's destructor, which clears the + // registry again. Detach the entries before releasing any Python objects. + decltype(functions) removed; + removed.swap(functions); + } + + ~PyFuncRegistry() { Clear(); } +}; + +static PyFuncRegistry py_func_registry; /*! * \brief Clear the Python function registry on shutdown */ -void ClearPyFuncRegistry() { py_func_registry.clear(); } +void ClearPyFuncRegistry() { py_func_registry.Clear(); } /*! * \brief Register a Python function for call_py_func * \param name The function name * \param func The Python function wrapped as ffi::Function */ -void RegisterPyFunc(const std::string& name, ffi::Function func) { py_func_registry[name] = func; } +void RegisterPyFunc(const std::string& name, ffi::Function func) { + py_func_registry.functions[name] = func; +} /*! * \brief Get a registered Python function @@ -478,8 +493,8 @@ void RegisterPyFunc(const std::string& name, ffi::Function func) { py_func_regis * \return The Python function */ ffi::Function GetPyFunc(const std::string& name) { - auto it = py_func_registry.find(name); - if (it == py_func_registry.end()) { + auto it = py_func_registry.functions.find(name); + if (it == py_func_registry.functions.end()) { TVM_FFI_THROW(InternalError) << "Python function '" << name << "' not found in registry"; } return it->second; diff --git a/tests/python/relax/test_relax_operators.py b/tests/python/relax/test_relax_operators.py index 3a93ac674b..4bb50f9077 100644 --- a/tests/python/relax/test_relax_operators.py +++ b/tests/python/relax/test_relax_operators.py @@ -18,6 +18,7 @@ from __future__ import annotations +import subprocess import sys import tempfile @@ -408,6 +409,26 @@ def test_op_call_inplace_packed(exec_mode): assert (result[1].numpy() == sum).all() [email protected]("clear_explicitly", [False, True]) +def test_py_func_registry_reentrant_cleanup(clear_explicitly): + # A callback can retain an owner whose destructor clears the registry again. + # Exercise process teardown separately so shutdown errors fail the test. + source = """ +import sys +import tvm +class Owner: + def __del__(self): + tvm.get_global_func("vm.builtin.clear_py_func_registry")() +owner = Owner() +tvm.get_global_func("vm.builtin.register_py_func")("cleanup", lambda retained=owner: 7) +assert tvm.get_global_func("vm.builtin.get_py_func")("cleanup")() == 7 +del owner +if sys.argv[1] == "True": + tvm.get_global_func("vm.builtin.clear_py_func_registry")() +""" + subprocess.run([sys.executable, "-c", source, str(clear_explicitly)], check=True, timeout=60) + + def test_op_call_py_func(exec_mode): """Test R.call_py_func operator functionality.""" import torch
