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 20cecf090565d36adf368c31d88909bc52a76bd6 Author: Tianqi Chen <[email protected]> AuthorDate: Sun Sep 20 23:20:22 2026 +0000 [TVMScript] Reuse declared signature symbols through binding metadata --- python/tvm/relax/script/builder_v2/__init__.py | 15 ++++++++++- python/tvm/script/ir_builder/protocol.py | 14 ++++++++++ python/tvm/tirx/script/builder_v2/__init__.py | 36 +++++++++++++++++++++++++- 3 files changed, 63 insertions(+), 2 deletions(-) diff --git a/python/tvm/relax/script/builder_v2/__init__.py b/python/tvm/relax/script/builder_v2/__init__.py index fafeb2fd47..7ea3a135da 100644 --- a/python/tvm/relax/script/builder_v2/__init__.py +++ b/python/tvm/relax/script/builder_v2/__init__.py @@ -221,13 +221,26 @@ def bind_( span=None, name_span=None, previous=_protocol.MISSING, + declaration=False, ): """Emit an immutable Relax binding and return the newly bound value.""" + _check_unterminated() + if declaration: + if not _ir.is_prim_var(value): + raise TypeError("A symbol declaration requires a concrete primitive variable") + if ty is not None and not _ffi.structural_equal(_type(ty), value.ty): + raise TypeError("The symbol declaration has an incompatible type") + if previous is not _protocol.MISSING: + if not _ir.is_prim_var(previous) or not _ffi.structural_equal(previous.ty, value.ty): + raise TypeError("The symbol declaration has an incompatible signature dtype") + return previous + if name is not None: + _IRBuilder.name(name, value) + return _protocol.at(name_span if name_span is not None else span, value) if value is _protocol.MISSING: raise ValueError("Relax bindings require an initializer") if isinstance(value, _I.meta_var): return value.value - _check_unterminated() ty = None if ty is None else _type(ty) value = _value(value, ty) with _protocol.span_context(span): diff --git a/python/tvm/script/ir_builder/protocol.py b/python/tvm/script/ir_builder/protocol.py index fbb0249db1..87cb155862 100644 --- a/python/tvm/script/ir_builder/protocol.py +++ b/python/tvm/script/ir_builder/protocol.py @@ -21,6 +21,7 @@ and return concrete values. Dialects register their own function kinds here, so translation consumes construction policies without importing their owners. """ +from builtins import slice as slice from contextlib import nullcontext from dataclasses import dataclass from inspect import signature @@ -71,6 +72,19 @@ def expression_args(*fields, introduce=False, dtype=None, scalar_strings=True): return decorate +class DeclarationArguments(NamedTuple): + """Callable syntax that declares a symbol when its value argument is absent.""" + + value_parameter: str + dtype: Any = None + + +def register_declaration(constructor, *, value_parameter="expr", dtype=None): + """Mark a concrete constructor's declaration form without wrapping the call.""" + constructor.__tvm_declaration_args__ = DeclarationArguments(value_parameter, dtype) + return constructor + + @dataclass(frozen=True) class FunctionKind: """Construction namespace and policies registered by a function decorator.""" diff --git a/python/tvm/tirx/script/builder_v2/__init__.py b/python/tvm/tirx/script/builder_v2/__init__.py index 0d79bd9f0b..1647786a15 100644 --- a/python/tvm/tirx/script/builder_v2/__init__.py +++ b/python/tvm/tirx/script/builder_v2/__init__.py @@ -20,6 +20,8 @@ import builtins as _python from functools import partial as _partial from functools import wraps as _wraps +import tvm_ffi as _ffi + from tvm import ir as _ir from tvm import tirx as _tir from tvm.script.ir_builder import IRBuilder as _IRBuilder @@ -28,6 +30,7 @@ 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 expression_args as _expression_args +from tvm.script.ir_builder.protocol import register_declaration as _register_declaration from tvm.script.ir_builder.protocol import span_context as _span_context from tvm.tirx.script import builder as _T from tvm.tirx.script.builder import * # pylint: disable=wildcard-import,unused-wildcard-import @@ -205,11 +208,35 @@ def _check_unterminated(): statements = last.seq -def bind_(value=_MISSING, *, ty=None, name=None, span=None, name_span=None, previous=_MISSING): +def bind_( + value=_MISSING, + *, + ty=None, + name=None, + span=None, + name_span=None, + previous=_MISSING, + declaration=False, +): """Bind concrete values, preserving existing mutable scalar storage.""" name_span = span if name_span is None else name_span _check_unterminated() with _span_context(span): + if declaration: + if not _ir.is_prim_var(value): + raise TypeError("A symbol declaration requires a concrete primitive variable") + if ty is not None: + annotation = ty() if callable(ty) else ty + annotation = annotation.ty if isinstance(annotation, _ir.Expr) else annotation + if not _ffi.structural_equal(annotation, value.ty): + raise TypeError("The symbol declaration has an incompatible type") + if previous is not _MISSING: + if not _ir.is_prim_var(previous) or not _ffi.structural_equal( + previous.ty, value.ty + ): + raise TypeError("The symbol declaration has an incompatible signature dtype") + return previous + return _name(value, name, name_span) if previous is not _MISSING and isinstance(previous, _ir.TensorLoad): if value is _MISSING: raise ValueError("A reassignment requires an initializer") @@ -411,3 +438,10 @@ def shared_scalar(dtype="float32"): def match_buffer(*args, **kwargs): """Construct a native buffer match with resolved symbolic shape fields.""" return _T.match_buffer(*args, **kwargs) + + +# Constructor identities carry syntax policy; aliases share it without wrappers. +for _constructor in vars(_T).values(): + if isinstance(_constructor, _T.DtypeConstructor): + _register_declaration(_constructor, dtype=_constructor._dtype_str) +del _constructor
