This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch tvmscript-ast-only-transpiler in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 875f0eb517c5187b83a2d6edf88edd9fe7d9107c Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 22:16:51 2026 +0000 [TVMScript] Preserve native identity in callable builder references --- python/tvm/script/ir_builder/base.py | 5 +- python/tvm/script/ir_builder/construction.py | 84 +++++++++++++++++++++++++--- 2 files changed, 80 insertions(+), 9 deletions(-) diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index bc045e7d22..1ac539ddb8 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -353,7 +353,10 @@ def at(span, value): """ if span is not None and isinstance(value, ir.Expr) and IRBuilder.is_in_scope(): with _construction_span(span): - return IRBuilder.current()._set_current_source_span(value) + # Native code mutates and returns the same node. Keep the original + # Python facade too: FFI rewrapping would erase builder-owned callable + # behavior on GlobalVar subclasses despite preserving native identity. + IRBuilder.current()._set_current_source_span(value) return value diff --git a/python/tvm/script/ir_builder/construction.py b/python/tvm/script/ir_builder/construction.py index f0879a7c7b..6e3e76265c 100644 --- a/python/tvm/script/ir_builder/construction.py +++ b/python/tvm/script/ir_builder/construction.py @@ -25,6 +25,8 @@ from contextlib import contextmanager from types import SimpleNamespace from typing import TypeVar +import tvm_ffi + from tvm import ir from . import ir as I @@ -37,6 +39,65 @@ _OPAQUE_FACTORY = None _MODULE_ADAPTER = None +class _FunctionReference(ir.GlobalVar): + """Keep a module-owned GlobalVar callable without changing its native object. + + The Python subclass shares the original native handle. It remains a real + GlobalVar at FFI boundaries, where a merely callable ObjectConvertible can + otherwise be converted to an ffi.Function before its asobject is consulted. + The current native function frame determines call construction, so the same + module member works from both dialects. No global call hook or parser + call-target adapter is installed. + """ + + def __init__(self, reference): + # The FFI constructor convention retains the returned existing handle; + # it does not create, clone, or replace the underlying native GlobalVar. + # The temporary identity callback owns no registration or persistent state. + self.__init_handle_by_constructor__(tvm_ffi.convert_func(lambda: reference)) + + def asobject(self): + """Return the original native reference without copying or changing it. + + Returns + ------- + ir.GlobalVar + Exact object retained by the module or function construction record. + + Notes + ----- + The facade itself holds the retained original native handle. No frame + is entered, no IR is created, and the reference type is never copied. + """ + return self + + def __call__(self, *args): + # Imports are deferred until execution to avoid the builder namespaces' + # initialization cycle. Only the nearest function owns call semantics; + # module/symbol/block frames do not select a dialect. + if IRBuilder.is_in_scope(): + from tvm.relax.script.builder.frame import FunctionFrame + from tvm.tirx.script.builder.frame import PrimFuncFrame + + for frame in reversed(IRBuilder.current().frames): + if isinstance(frame, PrimFuncFrame): + from tvm.tirx.script.builder.ir import _call_global + + return _call_global(self, *args) + if isinstance(frame, FunctionFrame): + from tvm import relax + + return relax.Call(self, [relax.utils.convert_to_expr(arg) for arg in args]) + return ir.GlobalVar.__call__(self, *args) + + +def _function_reference(value): + """Expose a module GlobalVar as a callable while preserving other values.""" + if isinstance(value, ir.GlobalVar) and not isinstance(value, _FunctionReference): + return _FunctionReference(value) + return value + + def register_opaque_factory(factory, *, module_adapter=None): """Register construction and adaptation callbacks for Python module functions. @@ -111,7 +172,9 @@ class FunctionRecord: signature_symbols retains the names introduced by parameter annotations, excluding return-only free symbols; it is updated after each parameter. return_type starts MISSING and is set by returns(); reference/function start - None and are populated after declaration/definition respectively. All state + None and are populated after declaration/definition respectively. Global + references expose ordinary calls through a wrapper retaining the exact native + GlobalVar; local variable references pass through unchanged. All state belongs to this generated-program invocation and dies with its result graph. Examples @@ -180,7 +243,7 @@ class FunctionRecord: with self.builder.decl_function(**self.options, **mode, span=self.span) as frame: self.builder.func_name(self.name) yield self - self.reference = frame.reference + self.reference = _function_reference(frame.reference) def parameter(self, name, annotation, location=None): """Construct and retain a parameter in the active declaration. @@ -324,13 +387,14 @@ class FunctionRecord: ------- object Existing canonical symbol, a symbol resolved from a host TypeVar, - or any other fallback unchanged. + a callable GlobalVar reference, or any other fallback unchanged. Notes ----- A host TypeVar resolves in this record's retained symbol frame before annotation arithmetic executes. Explicit predeclarations have already - run, so they retain their requested dtype. No frame is entered and + run, so they retain their requested dtype. No frame is entered. + GlobalVars retain their native identity behind callable references; ordinary host values are neither converted nor cached. Examples @@ -339,7 +403,7 @@ class FunctionRecord: """ if isinstance(fallback, TypeVar): return self.symbols.resolve(name) - return self.symbols.symbols.get(name, fallback) + return self.symbols.symbols.get(name, _function_reference(fallback)) def returns(self, annotation): """Set the return annotation on the record and active native frame. @@ -488,7 +552,8 @@ class ModuleProgram: Returns ------- ir.GlobalVar - Native reference, also exposed on namespace. + Callable Python facade sharing the original native GlobalVar handle, + also published on namespace. Raises ------ @@ -505,7 +570,7 @@ class ModuleProgram: >>> with ModuleProgram("Example") as program: ... reference = program.reserve("main") """ - reference = I.reserve_function(name) + reference = _function_reference(I.reserve_function(name)) setattr(self.namespace, name, reference) return reference @@ -544,6 +609,7 @@ class ModuleProgram: reference = I.decl_function(name, value) I.def_function(name, value) value = reference + value = _function_reference(value) setattr(self.namespace, name, value) return value @@ -564,7 +630,8 @@ class ModuleProgram: Returns ------- ir.GlobalVar - Opaque function reference, also exposed on namespace. + Callable opaque-function reference sharing the original GlobalVar + handle, also published on namespace. Raises ------ @@ -588,6 +655,7 @@ class ModuleProgram: reference = I.decl_function(name, opaque) I.def_function(name, opaque) self.python_functions[name] = function + reference = _function_reference(reference) setattr(self.namespace, name, reference) return reference
