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 4be5dab6511d22fec8dcff837c57ac7b3efc3269
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 03:17:51 2026 +0000

    Declare first-use symbols in registered compound shape expressions
---
 python/tvm/script/ir_builder/protocol.py      | 20 +++++++++++++++++---
 python/tvm/script/parser_v2/annotations.py    | 25 +++++++++++++++++++++++++
 python/tvm/tirx/script/builder_v2/__init__.py | 17 +++++++++++++++--
 3 files changed, 57 insertions(+), 5 deletions(-)

diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index 35e19104ce..5a971e24a2 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -49,9 +49,17 @@ class ExpressionArguments(NamedTuple):
     dtype: Any = None
     scalar_strings: bool = True
     implicit_dtype: Any = None
+    compound_declarations: bool = False
 
 
-def expression_args(*fields, introduce=False, dtype=None, scalar_strings=True, 
implicit_dtype=None):
+def expression_args(
+    *fields,
+    introduce=False,
+    dtype=None,
+    scalar_strings=True,
+    implicit_dtype=None,
+    compound_declarations=False,
+):
     """Mark constructor fields whose nested strings are source expressions.
 
     ``introduce`` permits signature/match scopes to introduce otherwise unknown
@@ -60,7 +68,8 @@ def expression_args(*fields, introduce=False, dtype=None, 
scalar_strings=True, i
     ``scalar_strings=False`` when a bare string is literal shorthand while
     strings nested in tuples/lists remain expressions. ``implicit_dtype`` 
selects
     only the default for names introduced by strings; explicit TypeVars retain
-    the dialect default unless ``dtype`` is supplied.
+    the dialect default unless ``dtype`` is supplied. ``compound_declarations``
+    also permits first-use declarations inside quoted compound expressions.
     """
 
     def decorate(constructor):
@@ -69,7 +78,12 @@ def expression_args(*fields, introduce=False, dtype=None, 
scalar_strings=True, i
         if unknown:
             raise ValueError(f"Unknown expression argument fields: 
{sorted(unknown)}")
         constructor.__tvm_expression_args__ = ExpressionArguments(
-            tuple(fields), bool(introduce), dtype, bool(scalar_strings), 
implicit_dtype
+            tuple(fields),
+            bool(introduce),
+            dtype,
+            bool(scalar_strings),
+            implicit_dtype,
+            bool(compound_declarations),
         )
         return constructor
 
diff --git a/python/tvm/script/parser_v2/annotations.py 
b/python/tvm/script/parser_v2/annotations.py
index a659458a2c..c9db73be11 100644
--- a/python/tvm/script/parser_v2/annotations.py
+++ b/python/tvm/script/parser_v2/annotations.py
@@ -306,8 +306,21 @@ class AnnotationScope:
                 self.allow_names = False
                 self.in_string = False
                 self.dtype = None
+                self.declaration_dtype = None
+                self.compound_declarations = False
 
             def visit_Name(self, current):
+                if (
+                    self.allow_names
+                    and self.in_string
+                    and self.compound_declarations
+                    and isinstance(current.ctx, ast.Load)
+                    and current.id not in scope.env
+                    and not hasattr(builtins, current.id)
+                ):
+                    scope._shape_declarations.setdefault(
+                        current.id, (current, self.declaration_dtype)
+                    )
                 if collect_declarations:
                     return current
                 if isinstance(current.ctx, ast.Load):
@@ -347,8 +360,16 @@ class AnnotationScope:
 
             def expression_field(self, current, metadata, *, nested=False):
                 old_allow, old_dtype, old_string = self.allow_names, 
self.dtype, self.in_string
+                old_compound, old_declaration_dtype = (
+                    self.compound_declarations,
+                    self.declaration_dtype,
+                )
                 self.allow_names = introduce and metadata.introduce
                 self.dtype = metadata.dtype
+                self.compound_declarations = metadata.compound_declarations
+                self.declaration_dtype = (
+                    metadata.dtype if metadata.dtype is not None else 
metadata.implicit_dtype
+                )
                 try:
                     if isinstance(current, ast.List | ast.Tuple):
                         current.elts = [
@@ -373,6 +394,10 @@ class AnnotationScope:
                     return self.visit(current)
                 finally:
                     self.allow_names, self.dtype, self.in_string = old_allow, 
old_dtype, old_string
+                    self.compound_declarations, self.declaration_dtype = (
+                        old_compound,
+                        old_declaration_dtype,
+                    )
 
             def visit_Call(self, current):
                 constructor = scope._resolve(current.func)
diff --git a/python/tvm/tirx/script/builder_v2/__init__.py 
b/python/tvm/tirx/script/builder_v2/__init__.py
index f12acf0198..1384cab01c 100644
--- a/python/tvm/tirx/script/builder_v2/__init__.py
+++ b/python/tvm/tirx/script/builder_v2/__init__.py
@@ -48,7 +48,13 @@ def type_var(name, *, dtype=None, span=None):
 
 
 @_expression_args(
-    "shape", "strides", "elem_offset", "byte_offset", introduce=True, 
implicit_dtype="int32"
+    "shape",
+    "strides",
+    "elem_offset",
+    "byte_offset",
+    introduce=True,
+    implicit_dtype="int32",
+    compound_declarations=True,
 )
 def Buffer(
     shape,
@@ -474,7 +480,14 @@ def shared_scalar(dtype="float32"):
     return alloc_scalar(dtype, "shared")
 
 
-@_expression_args("shape", "strides", "elem_offset", introduce=True, 
implicit_dtype="int32")
+@_expression_args(
+    "shape",
+    "strides",
+    "elem_offset",
+    introduce=True,
+    implicit_dtype="int32",
+    compound_declarations=True,
+)
 @_wraps(_T.match_buffer)
 def match_buffer(*args, **kwargs):
     """Construct a native buffer match with resolved symbolic shape fields."""

Reply via email to