Skip to content

Commit 60c250b

Browse files
authored
[Fix][Relax] Scope BasePyModule's Python function registry entries to the module (#20417)
`BasePyModule.__del__` calls `vm.builtin.clear_py_func_registry`, which empties the process-wide registry rather than removing the module's own entries. Collecting any module therefore drops every registration in the process, and a later `R.call_py_func` fails with `Python function '<name>' not found in registry`. Whether it happens depends on when the collector runs, so it shows up as an unrelated test failing. The module built in `tests/python/relax/test_pytorch_integration.py` registers nothing itself, but its finalizer wipes the names registered by `tests/python/relax/test_relax_operators.py`: ```python tvm.get_global_func("vm.builtin.register_py_func")("victim", tvm.runtime.convert(lambda x: x)) bystander = BasePyModule.__new__(BasePyModule) # registers nothing del bystander gc.collect() tvm.get_global_func("vm.builtin.get_py_func")("victim") # InternalError: Python function 'victim' not found in registry ``` This adds `vm.builtin.unregister_py_func` and records the names each module registers, so the finalizer removes only its own entries. Registering a name again transfers ownership to the newer module, which keeps a stale module from dropping a live one's function. The per-module cleanup was also unreachable before this change. The wrapper stored in the registry captured `self`, so the registry held a strong reference to every module that registered a function and none of them were ever collected — the finalizer only ever ran for modules that had registered nothing, which are exactly the ones that then cleared everyone else's entries. The wrapper now holds a weak reference to its module, and reports a clear error if it is called after the module is gone: ``` RuntimeError: Python function 'f' belongs to a BasePyModule that has been destroyed; keep the module alive while calling into it. ``` `vm.builtin.clear_py_func_registry` is left in place; it simply has no callers now.
1 parent 0db97a5 commit 60c250b

3 files changed

Lines changed: 85 additions & 10 deletions

File tree

‎python/tvm/relax/base_py_module.py‎

Lines changed: 49 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,8 @@
1919

2020
import inspect
2121
import os
22+
import sys
23+
import weakref
2224
from typing import Any, Optional, Union
2325

2426
import numpy as np
@@ -43,6 +45,11 @@
4345
_FASTER_DLPACK_EXTENSION = None
4446

4547

48+
# Maps a registered Python function name to the id of the BasePyModule that registered it last.
49+
# A module only unregisters names it still owns, so re-registering a name transfers ownership.
50+
_PY_FUNC_OWNERS: dict[str, int] = {}
51+
52+
4653
class BasePyModule:
4754
"""Base class that allows Python functions in IRModule with DLPack conversion.
4855
@@ -56,12 +63,31 @@ class BasePyModule:
5663
subclass with ``R.py_module`` adds this executable runtime interface.
5764
"""
5865

66+
# Names this instance registered with the VM's Python function registry. The registry is
67+
# global, so a module must only remove entries it still owns.
68+
_registered_py_funcs: tuple[str, ...] = ()
69+
5970
def __del__(self):
60-
"""Clean up registered Python functions on module destruction."""
71+
"""Unregister the Python functions this module still owns."""
72+
registered = getattr(self, "_registered_py_funcs", None)
73+
if not registered:
74+
return
6175
try:
62-
clear_func = tvm.get_global_func("vm.builtin.clear_py_func_registry")
63-
clear_func()
64-
except (ValueError, AttributeError):
76+
# Release ownership first: if the call below fails, the map must not keep pointing at
77+
# a destroyed module, whose id() a later module could reuse.
78+
owned = [name for name in registered if _PY_FUNC_OWNERS.get(name) == id(self)]
79+
for func_name in owned:
80+
del _PY_FUNC_OWNERS[func_name]
81+
# Once finalization starts the registry goes away with the process, and the module
82+
# globals this needs may already be cleared.
83+
if not owned or sys.is_finalizing():
84+
return
85+
unregister_py_func = tvm.get_global_func("vm.builtin.unregister_py_func")
86+
for func_name in owned:
87+
unregister_py_func(func_name)
88+
except Exception: # pylint: disable=broad-except
89+
# A finalizer must not raise: either the interpreter is shutting down or the runtime
90+
# is already unusable, and in both cases the registry no longer matters.
6591
pass
6692

6793
def __init__(
@@ -203,21 +229,36 @@ def _register_python_functions(self):
203229
for func_name, py_func in self.ir_mod.__pyfuncs__.items():
204230

205231
def create_py_func_wrapper(name, original_func):
232+
# The registry owns the wrapper, so capture the module weakly. A strong capture
233+
# would make every registering module immortal, keeping its entries registered
234+
# for the lifetime of the process.
235+
module_ref = weakref.ref(self)
236+
206237
def wrapper(*args, **kwargs):
207-
converted_args = [self._convert_tvm_to_pytorch(arg) for arg in args]
238+
module = module_ref()
239+
if module is None:
240+
raise RuntimeError(
241+
f"Python function '{name}' belongs to a BasePyModule that has been "
242+
"destroyed; keep the module alive while calling into it."
243+
)
244+
245+
converted_args = [module._convert_tvm_to_pytorch(arg) for arg in args]
208246
converted_kwargs = {
209-
k: self._convert_tvm_to_pytorch(v) for k, v in kwargs.items()
247+
k: module._convert_tvm_to_pytorch(v) for k, v in kwargs.items()
210248
}
211249

212-
result = original_func(self, *converted_args, **converted_kwargs)
250+
result = original_func(module, *converted_args, **converted_kwargs)
213251

214-
return self._convert_pytorch_to_tvm(result)
252+
return module._convert_pytorch_to_tvm(result)
215253

216254
wrapper.__name__ = name
217255
return wrapper
218256

219257
wrapped_func = create_py_func_wrapper(func_name, py_func)
220258
register_py_func(func_name, wrapped_func)
259+
_PY_FUNC_OWNERS[func_name] = id(self)
260+
if func_name not in self._registered_py_funcs:
261+
self._registered_py_funcs = (*self._registered_py_funcs, func_name)
221262

222263
def call_tir(self, tir_func, args, out_ty):
223264
"""Call a TIR function with PyTorch tensors."""

‎src/runtime/vm/builtin.cc‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -472,6 +472,12 @@ void ClearPyFuncRegistry() { py_func_registry.clear(); }
472472
*/
473473
void RegisterPyFunc(const std::string& name, ffi::Function func) { py_func_registry[name] = func; }
474474

475+
/*!
476+
* \brief Unregister a Python function registered with RegisterPyFunc
477+
* \param name The function name
478+
*/
479+
void UnregisterPyFunc(const std::string& name) { py_func_registry.erase(name); }
480+
475481
/*!
476482
* \brief Get a registered Python function
477483
* \param name The function name
@@ -522,6 +528,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
522528
refl::GlobalDef()
523529
.def_packed("vm.builtin.call_py_func", CallPyFunc)
524530
.def("vm.builtin.register_py_func", RegisterPyFunc)
531+
.def("vm.builtin.unregister_py_func", UnregisterPyFunc)
525532
.def("vm.builtin.get_py_func", GetPyFunc)
526533
.def("vm.builtin.clear_py_func_registry", ClearPyFuncRegistry);
527534
}

‎tests/python/relax/test_relax_operators.py‎

Lines changed: 29 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,11 @@
1616
# under the License.
1717
# ruff: noqa: E501, F841
1818

19+
import gc
1920
import sys
2021
import tempfile
22+
import weakref
23+
from types import SimpleNamespace
2124

2225
import numpy as np
2326
import pytest
@@ -478,8 +481,32 @@ def multiple_calls(x: R.Tensor((2,), "float32")):
478481
expected2 = 1.0 / (1.0 + np.exp(-np.maximum(y_data, 0.0)))
479482
assert (result2.numpy() == expected2).all()
480483

481-
clear_func = tvm.get_global_func("vm.builtin.clear_py_func_registry")
482-
clear_func()
484+
unregister_func = tvm.get_global_func("vm.builtin.unregister_py_func")
485+
unregister_func("torch_relu")
486+
unregister_func("torch_sigmoid")
487+
488+
489+
def test_py_func_registry_is_scoped_to_its_module():
490+
"""A module's finalizer must drop its own registrations and nothing else."""
491+
from tvm.relax.base_py_module import BasePyModule
492+
493+
get_func = tvm.get_global_func("vm.builtin.get_py_func")
494+
tvm.get_global_func("vm.builtin.register_py_func")("registry_probe", lambda x: x)
495+
496+
# __new__ skips __init__'s JIT compilation; only the registration matters here.
497+
module = BasePyModule.__new__(BasePyModule)
498+
module.ir_mod = SimpleNamespace(pyfuncs={"registry_owned": lambda self, x: x})
499+
module._register_python_functions()
500+
module_ref = weakref.ref(module)
501+
502+
del module
503+
gc.collect()
504+
505+
assert module_ref() is None, "the registry must not keep the module alive"
506+
assert get_func("registry_probe") is not None, "another owner's function was dropped"
507+
with pytest.raises(tvm.error.InternalError, match="not found in registry"):
508+
get_func("registry_owned")
509+
tvm.get_global_func("vm.builtin.unregister_py_func")("registry_probe")
483510

484511

485512
def test_op_to_device(exec_mode):

0 commit comments

Comments
 (0)