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 1c3f11d5adbce43e4a471586b6e7655f920a3255 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 02:36:30 2026 +0000 Bridge parsed Python modules to the runtime factory --- python/tvm/relax/base_py_module.py | 32 ++++++++++++++++++++++++++++++-- python/tvm/script/parser_v2/functions.py | 19 ++++++++++++++++--- 2 files changed, 46 insertions(+), 5 deletions(-) diff --git a/python/tvm/relax/base_py_module.py b/python/tvm/relax/base_py_module.py index 825b532364..661fe88d01 100644 --- a/python/tvm/relax/base_py_module.py +++ b/python/tvm/relax/base_py_module.py @@ -194,6 +194,9 @@ class BasePyModule: if not hasattr(self.ir_mod, "pyfuncs") or not self.ir_mod.pyfuncs: return + for func_name, py_func in self.ir_mod.pyfuncs.items(): + self.add_python_function(func_name, py_func) + try: register_py_func = tvm.get_global_func("vm.builtin.register_py_func") except ValueError: @@ -208,14 +211,14 @@ class BasePyModule: k: self._convert_tvm_to_pytorch(v) for k, v in kwargs.items() } - result = original_func(self, *converted_args, **converted_kwargs) + result = original_func(*converted_args, **converted_kwargs) return self._convert_pytorch_to_tvm(result) wrapper.__name__ = name return wrapper - wrapped_func = create_py_func_wrapper(func_name, py_func) + wrapped_func = create_py_func_wrapper(func_name, getattr(self, func_name)) register_py_func(func_name, wrapped_func) def call_tir(self, tir_func, args, out_ty): @@ -612,3 +615,28 @@ class BasePyModule: script_content = self.script(**kwargs) cprint(script_content, style=style, black_format=black_format) + + +class PyModuleFactory: + """A parsed module that creates independent executable Python module instances. + + The parsed IR and original Python callables remain available for inspection. + Each invocation clones the IR container while retaining module metadata and + closure identities, then uses the normal runtime compilation and call bridge. + """ + + def __init__(self, ir_module: IRModule, original_class=None): + self.ir_module = ir_module + self.original_class = original_class + self.__name__ = getattr(ir_module, "__name__", "Module") + + def __getattr__(self, name): + return getattr(self.ir_module, name) + + def __getitem__(self, name): + return self.ir_module[name] + + def __call__(self, device=None, target=None): + instance_module = self.ir_module.clone() + instance_module.pyfuncs = dict(getattr(self.ir_module, "pyfuncs", {})) + return BasePyModule(instance_module, tvm.cpu(0) if device is None else device, target) diff --git a/python/tvm/script/parser_v2/functions.py b/python/tvm/script/parser_v2/functions.py index 540fbea6b1..cce0dbfa19 100644 --- a/python/tvm/script/parser_v2/functions.py +++ b/python/tvm/script/parser_v2/functions.py @@ -27,18 +27,24 @@ from types import CodeType from tvm.script.ir_builder import ir as I _OPAQUE_FACTORY = None +_MODULE_ADAPTER = None -def register_opaque_factory(factory): +def register_opaque_factory(factory, *, module_adapter=None): """Register the owner-provided constructor for a module's opaque function slot. The factory receives ``(name, python_callable, source_text, span)`` and must - return a concrete BaseFunc. Registration never executes the Python body. + return a concrete BaseFunc. The optional adapter receives a completed module, + its original class (if available), and its resolved base classes. Registration + never executes the Python body. """ - global _OPAQUE_FACTORY + global _OPAQUE_FACTORY, _MODULE_ADAPTER if not callable(factory): raise TypeError("An opaque function factory must be callable") + if module_adapter is not None and not callable(module_adapter): + raise TypeError("An opaque module adapter must be callable") _OPAQUE_FACTORY = factory + _MODULE_ADAPTER = module_adapter def is_python_function(compiler, node, env): @@ -126,6 +132,13 @@ def attach_python(module, functions): return module +def adapt_python_module(module, *, original=None, bases=()): + """Let the registered owner supply an executable module wrapper if needed.""" + if _MODULE_ADAPTER is None: + return module + return _MODULE_ADAPTER(module, original, bases) + + def _global_loads(node, filename): """Find free sibling uses without confusing attributes or shadowed locals.""" node = copy.deepcopy(node)
