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

Reply via email to