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)

Reply via email to