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]
