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 c4ede00b363b531e1bfcfa8d613147556c70426f
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 18:11:52 2026 +0000

    Route global calls through ordinary callable overloads
---
 python/tvm/ir/expr.py                       | 33 ++++++++++++++++-------
 python/tvm/relax/script/builder/__init__.py | 21 +++++----------
 python/tvm/script/ir_builder/protocol.py    | 42 -----------------------------
 python/tvm/script/parser/transpile.py       | 15 +++--------
 python/tvm/tirx/script/builder/__init__.py  |  9 -------
 tests/python/relax/test_tvmscript_parser.py | 15 +++++++----
 tests/python/tvmscript/test_parser.py       | 34 +++++++++++++++++++++++
 7 files changed, 77 insertions(+), 92 deletions(-)

diff --git a/python/tvm/ir/expr.py b/python/tvm/ir/expr.py
index 7cd7680ea8..b2d8b34727 100644
--- a/python/tvm/ir/expr.py
+++ b/python/tvm/ir/expr.py
@@ -93,18 +93,31 @@ class GlobalVar(Expr):
         self.__init_handle_by_constructor__(_ffi_api.GlobalVar, name_hint)
 
     def __call__(self, *args: Expr) -> Expr:
-        """Call the global variable.
+        """Construct an ordinary callable global reference.
+
+        args are source call arguments. Within a primitive builder they retain
+        primitive conversion and the declared function's exact return type,
+        including pointers erased by a Relax-facing signature. Within a Relax
+        function, Python scalars/strings/tuples are converted by its normal 
value
+        rules. Outside construction, preserve the generic Call with missing
+        result type for later normalization. Returns an IR Call without 
entering
+        a frame or retaining state; argument/type errors propagate unchanged.
+        """
+        from tvm.script.ir_builder import IRBuilder
 
-        Parameters
-        ----------
-        args: List[Expr]
-            The arguments to the call.
+        if IRBuilder.is_in_scope():
+            from tvm.relax.script.builder.frame import FunctionFrame
+            from tvm.tirx.script.builder.frame import PrimFuncFrame
 
-        Returns
-        -------
-        call: Expr
-            A call taking the variable as a function.
-        """
+            for frame in reversed(list(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.relax.utils import convert_to_expr
+
+                    return Call(self, [convert_to_expr(value) for value in 
args])
         return Call(self, args)
 
 
diff --git a/python/tvm/relax/script/builder/__init__.py 
b/python/tvm/relax/script/builder/__init__.py
index 7607bebffa..d97526bcba 100644
--- a/python/tvm/relax/script/builder/__init__.py
+++ b/python/tvm/relax/script/builder/__init__.py
@@ -19,8 +19,6 @@
 # pylint: disable=wildcard-import,redefined-builtin,invalid-name
 import builtins as _python
 import numbers as _numbers
-import sys as _sys
-from functools import partial as _partial
 
 import tvm_ffi as _ffi
 
@@ -334,7 +332,13 @@ def bind_(
 
     Returns a named emitted Var, canonical symbol, or unchanged Python value.
     Expression inputs emit bindings in the active native function/block; a
-    TypeVarDecl only updates the symbol frame. Missing initializers, 
incompatible
+    TypeVarDecl only updates the symbol frame. For an annotation, native 
emission
+    fills MissingType on the original RHS value itself before normalization. A
+    tuple annotation also fills its tuple-literal fields recursively, 
preserving
+    shared value identities. Existing concrete types are checked and retained;
+    all compatibility checks finish before any missing value types are changed.
+    An output annotation does not assign types to arbitrary call arguments.
+    Missing initializers, incompatible
     declaration/match-cast types, or operations after an unconditional return
     raise ValueError/TypeError; native construction errors propagate.
     """
@@ -540,17 +544,6 @@ def select(condition, true_value, false_value):
 __all__ += ["logical_and", "logical_not", "logical_or", "select"]
 
 
-def _call_global(function, *args):
-    return _relax.Call(function, [_relax.utils.convert_to_expr(value) for 
value in args])
-
-
-def _global_callee(function):
-    return _partial(_call_global, function)
-
-
-_protocol.register_call_kind(_sys.modules[__name__], _ir.GlobalVar, 
_global_callee)
-
-
 for_ = For
 
 
diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index ebfd9d110f..6468464cbf 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -297,48 +297,6 @@ def _comparison_chain(comparisons, operands, conjunction, 
bind):
     return result
 
 
-def register_call_kind(builder, value_type, adapter):
-    """Register runtime adaptation of callable values for one builder 
namespace.
-
-    ``builder`` is a namespace supporting attributes, ``value_type`` a Python
-    type suitable for isinstance, and ``adapter`` maps matching values to
-    callable construction operations. Returns None and replaces that type's
-    entry in the namespace-owned __tvm_call_kinds__ dict. The dict maps types
-    to adapters, starts empty, and persists across functions; no values or
-    call results are cached. No frames are entered. Invalid attribute writes
-    raise Python errors; invalid types/adapters fail when consumed by callee.
-    """
-    policies = dict(getattr(builder, "__tvm_call_kinds__", {}))
-    policies[value_type] = adapter
-    builder.__tvm_call_kinds__ = policies
-
-
-def callee(builder, value, *, span=None):
-    """Adapt a runtime callable according to one builder's registered policy.
-
-    ``builder`` is the construction namespace and ``value`` the already
-    evaluated call target. Returns the first matching adapter's result, or
-    value unchanged when no type matches. Optional span accepts source_span
-    input forms and wraps the actual construction call in the private builder
-    span context, including operations which emit statements and return None.
-    None adds no wrapper or span instrumentation. Registration is read-only;
-    adapter and call errors propagate with original types and source metadata.
-    """
-    for value_type, adapter in getattr(builder, "__tvm_call_kinds__", 
{}).items():
-        if isinstance(value, value_type):
-            value = adapter(value)
-            break
-    if span is None:
-        return value
-
-    @wraps(value)
-    def construct(*args, **kwargs):
-        with _construction_span(span):
-            return value(*args, **kwargs)
-
-    return construct
-
-
 def __getattr__(name):
     # Compatibility exports forward to the single parser-owned registry. Lazy
     # lookup avoids importing parser entry points during builder 
initialization.
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 01394c8c59..23ef1252bf 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -355,18 +355,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                         child.format_spec = 
self._format_spec(child.format_spec)
         else:
             self._expression_children(node)
-        # Pattern: f(args) -> at(loc, callee(X, f, span=loc)(args)). The
-        # builder adapter owns construction contexts for calls which emit a
-        # statement and return None, including inline helpers. No parser-side
-        # context manager or deferred source-expression callback is introduced.
-        if isinstance(node, ast.Call):
-            node.func = self._call(
-                self.infrastructure_name,
-                "callee",
-                [self._name(self.dialect_prefix, original), node.func],
-                original,
-                span=self.span(original),
-            )
+        # Pattern: f(args, keyword=value) remains an ordinary Python call.
+        # Recursive argument translation above preserves order and locations;
+        # callable overloads own concrete IR argument/result semantics.
         if isinstance(node, ast.Starred) or (
             isinstance(node, ast.Name) and isinstance(node.ctx, ast.Store)
         ):
diff --git a/python/tvm/tirx/script/builder/__init__.py 
b/python/tvm/tirx/script/builder/__init__.py
index fb332d95c1..c7a31aec21 100644
--- a/python/tvm/tirx/script/builder/__init__.py
+++ b/python/tvm/tirx/script/builder/__init__.py
@@ -17,7 +17,6 @@
 """Concrete TIRx construction operations over the shared native IRBuilder 
stack."""
 
 import builtins as _python
-import sys as _sys
 from dataclasses import dataclass as _dataclass
 from dataclasses import field as _field
 from functools import partial as _partial
@@ -34,7 +33,6 @@ from tvm.script.ir_builder.protocol import MISSING as _MISSING
 from tvm.script.ir_builder.protocol import _construction_span
 from tvm.script.ir_builder.protocol import _frame_result as _named_frame_result
 from tvm.script.ir_builder.protocol import at as _at
-from tvm.script.ir_builder.protocol import register_call_kind as 
_register_call_kind
 from tvm.script.ir_builder.protocol import source_span as _source_span
 from tvm.script.ir_builder.type_var_frame import TypeVarDecl as _TypeVarDecl
 from tvm.script.ir_builder.type_var_frame import TypeVarFrame as _TypeVarFrame
@@ -792,13 +790,6 @@ def select(condition, true_value, false_value):
     return _tir.if_then_else(condition, true_value, false_value)
 
 
-def _global_callee(function):
-    return _partial(_native._call_global, function)
-
-
-_register_call_kind(_sys.modules[__name__], _ir.GlobalVar, _global_callee)
-
-
 def if_then_else_(condition, true_value, false_value):
     """Construct a scalar conditional whose compiled code evaluates one arm.
 
diff --git a/tests/python/relax/test_tvmscript_parser.py 
b/tests/python/relax/test_tvmscript_parser.py
index cf52a916b1..6a1109da35 100644
--- a/tests/python/relax/test_tvmscript_parser.py
+++ b/tests/python/relax/test_tvmscript_parser.py
@@ -1883,7 +1883,7 @@ def test_class_normalize():
     _check(InputModule, OutputModule)
 
 
-def test_context_aware_parsing(monkeypatch):
+def test_global_calls_use_callable_overload(monkeypatch):
     @tvm.script.ir_module
     class Module:
         @T.prim_func(s_tir=True)
@@ -1904,13 +1904,18 @@ def test_context_aware_parsing(monkeypatch):
 
     _check(Module)
 
-    # Break the env settings, but context-aware parsing can still handle it
-    def _break_env(self, *args):
-        raise RuntimeError("Fail to pass context-aware parsing")
+    # Generated source uses the same callable overload as ordinary Python.
+    # Preserve the module roundtrip while proving no parser adapter bypasses 
it.
+    calls = []
+    original = tvm.ir.GlobalVar.__call__
 
-    monkeypatch.setattr(tvm.ir.GlobalVar, "__call__", _break_env)
+    def record(self, *args):
+        calls.append(self.name_hint)
+        return original(self, *args)
 
+    monkeypatch.setattr(tvm.ir.GlobalVar, "__call__", record)
     _check(Module)
+    assert calls == ["add"]
 
 
 def test_unit_tuple_on_rhs_of_assign():
diff --git a/tests/python/tvmscript/test_parser.py 
b/tests/python/tvmscript/test_parser.py
index 46667cac02..2420a14080 100644
--- a/tests/python/tvmscript/test_parser.py
+++ b/tests/python/tvmscript/test_parser.py
@@ -691,3 +691,37 @@ def test_named_results_keep_nested_frame_identity():
     _run(compiler, function, builder)
     assert builder.frame_requests == [(inner, "value"), (outer, "value")]
     assert builder.returned == [30]
+
+
+def test_calls_use_ordinary_python_callable_overloads():
+    events = []
+
+    class Callable:
+        def __call__(self, first, *, second):
+            events.append(("call", first, second))
+            return first + second
+
+    def argument(value):
+        events.append(value)
+        return value
+
+    compiler, function, builder = _registered(
+        """
+        @D.function
+        def f():
+            value = target(argument(1), second=argument(2))
+            return value
+        """,
+        {"target": Callable(), "argument": argument},
+    )
+    generated = compiler.transformer().transform_statements(function.body)
+    assert not any(
+        isinstance(node, ast.Attribute) and node.attr == "callee"
+        for statement in generated
+        for node in ast.walk(statement)
+    )
+    assert not hasattr(protocol, "callee")
+    assert not hasattr(protocol, "register_call_kind")
+    _run(compiler, function, builder)
+    assert events == [1, 2, ("call", 1, 2)]
+    assert builder.returned == [3]

Reply via email to