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
The following commit(s) were added to refs/heads/tvmscript-ast-only-transpiler
by this push:
new a72f369da1 Register parameter declarations in the parser protocol
a72f369da1 is described below
commit a72f369da16aee1403292260549f219f464b42ec
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 17:22:39 2026 +0000
Register parameter declarations in the parser protocol
Expose register_type_var_decl for declaration-capable annotation
constructors while retaining explicit-type-first signature predeclaration and
builder-owned symbol compatibility.
---
python/tvm/script/ir_builder/protocol.py | 33 ++++--------------------------
python/tvm/script/parser/protocol.py | 30 +++++++++++++++++++++++++++
python/tvm/script/parser/transpile.py | 5 ++++-
python/tvm/tirx/script/builder/__init__.py | 4 ++--
tests/python/tvmscript/test_parser.py | 3 ++-
5 files changed, 42 insertions(+), 33 deletions(-)
diff --git a/python/tvm/script/ir_builder/protocol.py
b/python/tvm/script/ir_builder/protocol.py
index a22153c176..6fab759059 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -26,7 +26,7 @@ from builtins import locals as locals
from builtins import slice as slice
from contextlib import contextmanager, nullcontext
from functools import wraps
-from typing import Any, NamedTuple, TypeVar
+from typing import TypeVar
from tvm import ir
@@ -116,33 +116,6 @@ def wrap_expression_constructor(constructor,
call_signature, policy, *, as_type=
return result
-class DeclarationArguments(NamedTuple):
- """Immutable scalar-constructor metadata used by signature predeclaration.
-
- value_parameter names the optional value argument whose absence denotes a
- declaration; dtype is the explicit primitive dtype. Stored on the callable
- by register_declaration and shared across aliases and transpilation passes.
- It contains no symbols or construction results.
- """
-
- value_parameter: str
- dtype: Any = None
-
-
-def register_declaration(constructor, *, value_parameter="expr", dtype=None):
- """Register metadata for a constructor's zero-argument declaration form.
-
- ``constructor`` is a callable supporting attribute assignment;
- ``value_parameter`` defaults to "expr" and names its optional value
argument.
- ``dtype`` is the explicit dtype spelling, default None. Returns the same
- callable after attaching immutable DeclarationArguments; calls are never
- wrapped/evaluated and no frame is entered. Unsupported attribute assignment
- raises the ordinary Python error. Registration persists with the callable.
- """
- constructor.__tvm_declaration_args__ =
DeclarationArguments(value_parameter, dtype)
- return constructor
-
-
def source_span(location):
"""Materialize an IR span from a source location, or preserve an existing
span.
@@ -327,7 +300,7 @@ def callee(builder, value):
def __getattr__(name):
# Compatibility exports forward to the single parser-owned registry. Lazy
# lookup avoids importing parser entry points during builder
initialization.
- aliases = {"function_kind": "function_info"}
+ aliases = {"function_kind": "function_info", "register_declaration":
"register_type_var_decl"}
exported = {
"ExprStrPolicy",
"ExpressionArguments",
@@ -335,6 +308,8 @@ def __getattr__(name):
"expression_args",
"expr_str_policy",
"register_function",
+ "register_declaration",
+ "DeclarationArguments",
"function_kind",
}
if name in exported:
diff --git a/python/tvm/script/parser/protocol.py
b/python/tvm/script/parser/protocol.py
index 82f38d9eb9..e8082cc41d 100644
--- a/python/tvm/script/parser/protocol.py
+++ b/python/tvm/script/parser/protocol.py
@@ -144,6 +144,36 @@ ExpressionArguments = ExprStrPolicy
expression_args = expr_str_args
+class DeclarationArguments(NamedTuple):
+ """Immutable scalar-constructor metadata used by signature predeclaration.
+
+ value_parameter names the optional value argument whose absence denotes a
+ declaration; dtype is the explicit primitive dtype. Stored on the callable
+ by register_type_var_decl and shared across aliases and transpilation
passes.
+ It contains no symbols or construction results.
+ """
+
+ value_parameter: str
+ dtype: Any = None
+
+
+def register_type_var_decl(constructor, *, value_parameter="expr", dtype=None):
+ """Register a declaration-capable constructor for parameter annotations.
+
+ ``constructor`` is a callable supporting attribute assignment;
+ ``value_parameter`` defaults to "expr" and names its optional value
argument.
+ ``dtype`` is the explicit dtype spelling, default None. Returns the same
+ callable after attaching immutable DeclarationArguments. The signature
+ prepass recognizes this marker so an explicit later parameter type is
reserved
+ before an earlier shape string refers to it. This is not assignment-pattern
+ registration and never hoists body declarations. Calls are never
+ wrapped/evaluated and no frame is entered. Unsupported attribute assignment
+ raises the ordinary Python error. Registration persists with the callable.
+ """
+ constructor.__tvm_type_var_decl__ = DeclarationArguments(value_parameter,
dtype)
+ return constructor
+
+
class FunctionDecoratorInfo(NamedTuple):
"""Flat syntax registration for a source function decorator.
diff --git a/python/tvm/script/parser/transpile.py
b/python/tvm/script/parser/transpile.py
index b90adb965c..27d1926e72 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -990,6 +990,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
parameter,
)
)
+ # Pattern: f(A: X.Buffer(("n",)), n: X.int64) -> reserve the explicit
+ # scalar annotation before resolving signature strings. This is solely
a
+ # parameter-annotation prepass; body assignments stay in execution
order.
for parameter in node.args.args:
annotation = parameter.annotation
if annotation is None:
@@ -998,7 +1001,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
annotation.func if isinstance(annotation, ast.Call) else
annotation
)
if (
- getattr(constructor, "__tvm_declaration_args__", None) is not
None
+ getattr(constructor, "__tvm_type_var_decl__", None) is not None
or getattr(constructor, "__tvm_parameter_dtype__", None) is
not None
):
declaration.append(
diff --git a/python/tvm/tirx/script/builder/__init__.py
b/python/tvm/tirx/script/builder/__init__.py
index ae63a3535d..5780e10cd8 100644
--- a/python/tvm/tirx/script/builder/__init__.py
+++ b/python/tvm/tirx/script/builder/__init__.py
@@ -33,13 +33,13 @@ from tvm.script.ir_builder.base import IRBuilderFrame as
_NativeFrame
from tvm.script.ir_builder.protocol import MISSING as _MISSING
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 register_declaration as
_register_declaration
from tvm.script.ir_builder.protocol import source_span as _source_span
from tvm.script.ir_builder.protocol import span_context as _span_context
from tvm.script.ir_builder.type_var_frame import TypeVarDecl as _TypeVarDecl
from tvm.script.ir_builder.type_var_frame import TypeVarFrame as _TypeVarFrame
from tvm.script.ir_builder.type_var_frame import resolve_type_var
from tvm.script.parser.protocol import expr_str_args as _expression_args
+from tvm.script.parser.protocol import register_type_var_decl as
_register_type_var_decl
from tvm.tirx.lang.alloc_pool import SMEMPool, TMEMPool
from . import _ffi_api
@@ -697,7 +697,7 @@ def match_buffer(*args, **kwargs):
# Constructor identities carry syntax policy; aliases share it without
wrappers.
for _constructor in vars(_native).values():
if isinstance(_constructor, _native.DtypeConstructor):
- _register_declaration(_constructor, dtype=_constructor._dtype_str)
+ _register_type_var_decl(_constructor, dtype=_constructor._dtype_str)
del _constructor
diff --git a/tests/python/tvmscript/test_parser.py
b/tests/python/tvmscript/test_parser.py
index c1ca3e4207..1e814a7c3e 100644
--- a/tests/python/tvmscript/test_parser.py
+++ b/tests/python/tvmscript/test_parser.py
@@ -195,7 +195,7 @@ def test_callable_metadata_and_missing_initializer():
def scalar(expr=None):
return object() if expr is None else expr
- protocol.register_declaration(scalar)
+ syntax_protocol.register_type_var_decl(scalar)
previous = object()
compiler, function, builder = _registered(
"""\
@@ -465,6 +465,7 @@ def test_shared_allocator_preserves_source_names():
def test_parser_metadata_registration_boundary():
+ assert protocol.register_declaration is
syntax_protocol.register_type_var_decl
options, defaults = {"private": "local"}, {"local": False}
def decorator():