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 8e7f35a335b5351e2e7c05464f5d901f62332ddc
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 02:10:29 2026 +0000

    Respect declaration order in symbolic annotation resolution
---
 python/tvm/script/parser_v2/annotations.py    | 66 +++++++++++++++++++++------
 python/tvm/tirx/script/builder_v2/__init__.py | 31 ++++++++++---
 2 files changed, 78 insertions(+), 19 deletions(-)

diff --git a/python/tvm/script/parser_v2/annotations.py 
b/python/tvm/script/parser_v2/annotations.py
index 06b5045cc3..9195d8a5db 100644
--- a/python/tvm/script/parser_v2/annotations.py
+++ b/python/tvm/script/parser_v2/annotations.py
@@ -40,6 +40,10 @@ class AnnotationScope:
         self.symbols = {}
         self._evaluated = {}
         self._prepared = {}
+        self._pending_parameters = {}
+        self._parameter_annotations = {}
+        self._unbound_parameters = set()
+        self._shape_declarations = {}
         self._type_vars = {}
         self._counter = 0
         self._used_names = set(env)
@@ -143,13 +147,12 @@ class AnnotationScope:
         return ast.copy_location(ast.Name(name, ast.Load()), node)
 
     def prepare_parameters(self, arguments):
-        """Preallocate scalar parameter identities before dependent 
annotations.
+        """Prepare declared symbols and sequential scalar parameter 
annotations.
 
-        Scalar constructors own dtype metadata. Cached dtype expressions are
-        substituted in the annotation, so even a dynamic dtype is evaluated 
once.
-        ``symbols`` contains scalar parameter objects for the signature 
builder.
-        The caller evaluates annotations and registers parameters sequentially,
-        making each actual parameter available to subsequent annotations.
+        Declaration constructors reserve identities before dependent 
annotations.
+        Type constructors introduce their parameters in signature order. Cached
+        dtype expressions are evaluated once. Bare symbolic strings declare 
shape
+        names across the signature; compound expressions only reference them.
         """
         parameters = [*arguments.posonlyargs, *arguments.args, 
*arguments.kwonlyargs]
         self._used_names.update(
@@ -188,7 +191,14 @@ class AnnotationScope:
                             keyword.value = cached
             else:
                 continue
-            self._symbol(parameter.arg, parameter, dtype, shadow=True)
+            if declaration is not None:
+                self._symbol(parameter.arg, parameter, dtype, shadow=True)
+                self._parameter_annotations[id(annotation)] = parameter.arg
+                self._unbound_parameters.add(parameter.arg)
+            else:
+                self._pending_parameters[id(annotation)] = (parameter, dtype)
+        for prepared in self._prepared.values():
+            self.rewrite(prepared, introduce=True, collect_declarations=True)
         return self.symbols
 
     def _eval(self, node):
@@ -200,6 +210,14 @@ class AnnotationScope:
         """Evaluate an annotation once in the prepared signature scope."""
         key = id(node)
         if key not in self._evaluated:
+            
self._unbound_parameters.discard(self._parameter_annotations.get(key))
+            if key in self._pending_parameters:
+                parameter, dtype = self._pending_parameters[key]
+                if parameter.arg in self.symbols:
+                    self._error(
+                        parameter, "A later parameter cannot adopt an existing 
shape symbol"
+                    )
+                self._symbol(parameter.arg, parameter, dtype, shadow=True)
             prepared = self._prepared.get(key, node)
             # A quoted whole annotation is ordinary Python annotation syntax.
             if isinstance(prepared, ast.Constant) and 
isinstance(prepared.value, str):
@@ -275,7 +293,7 @@ class AnnotationScope:
             index = end
         return positions if decoded == node.value else None
 
-    def rewrite(self, node, *, introduce=False):
+    def rewrite(self, node, *, introduce=False, collect_declarations=False):
         """Return a copied expression AST, registering new symbols in ``env``.
 
         Construction code must execute with this scope's updated environment.
@@ -286,19 +304,36 @@ class AnnotationScope:
         class Rewrite(ast.NodeTransformer):
             def __init__(self):
                 self.allow_names = False
+                self.in_string = False
                 self.dtype = None
 
             def visit_Name(self, current):
+                if collect_declarations:
+                    return current
                 if isinstance(current.ctx, ast.Load):
+                    if not self.in_string and current.id in 
scope._unbound_parameters:
+                        scope._error(current, f"Parameter {current.id!r} is 
not yet bound")
+                    if (
+                        self.allow_names
+                        and not self.in_string
+                        and current.id not in scope.env
+                        and not hasattr(builtins, current.id)
+                    ):
+                        scope._error(current, f"Name {current.id!r} is not 
defined")
                     if isinstance(scope.env.get(current.id), TypeVar) and not 
introduce:
                         scope._error(
                             current, "A TypeVar must be introduced in a 
signature or match scope"
                         )
                     scope._canonical_type_var(current.id, current, self.dtype)
-                    if self.allow_names and (
-                        current.id in scope.env or not hasattr(builtins, 
current.id)
-                    ):
+                    if self.allow_names and current.id in scope.env:
                         scope._symbol(current.id, current, self.dtype)
+                    elif (
+                        self.allow_names
+                        and self.in_string
+                        and current.id in scope._shape_declarations
+                    ):
+                        declaration, dtype = 
scope._shape_declarations[current.id]
+                        scope._symbol(current.id, declaration, dtype)
                 return current
 
             def visit_Attribute(self, current):
@@ -309,7 +344,7 @@ class AnnotationScope:
                 return current
 
             def expression_field(self, current, metadata, *, nested=False):
-                old_allow, old_dtype = self.allow_names, self.dtype
+                old_allow, old_dtype, old_string = self.allow_names, 
self.dtype, self.in_string
                 self.allow_names = introduce and metadata.introduce
                 self.dtype = metadata.dtype
                 try:
@@ -322,9 +357,14 @@ class AnnotationScope:
                     if isinstance(current, ast.Constant) and 
isinstance(current.value, str):
                         if nested or metadata.scalar_strings:
                             current = scope._string_expression(current)
+                            self.in_string = True
+                            if self.allow_names and isinstance(current, 
ast.Name):
+                                scope._shape_declarations.setdefault(
+                                    current.id, (current, self.dtype)
+                                )
                     return self.visit(current)
                 finally:
-                    self.allow_names, self.dtype = old_allow, old_dtype
+                    self.allow_names, self.dtype, self.in_string = old_allow, 
old_dtype, old_string
 
             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 c327ab543e..4b4746f0c0 100644
--- a/python/tvm/tirx/script/builder_v2/__init__.py
+++ b/python/tvm/tirx/script/builder_v2/__init__.py
@@ -45,7 +45,7 @@ def type_var(name, *, dtype=None, span=None):
     return _ir.Var(name, "int64" if dtype is None else dtype, span)
 
 
-@_expression_args("shape", "strides", "elem_offset", "byte_offset", 
introduce=True)
+@_expression_args("shape", "strides", "elem_offset", "byte_offset", 
introduce=True, dtype="int32")
 def Buffer(
     shape,
     dtype="float32",
@@ -187,6 +187,8 @@ def _enter_concise(frame):
 
 
 def _as_expr(value):
+    if isinstance(value, _ffi.ObjectConvertible):
+        value = value.asobject()
     if isinstance(value, _ir.Expr):
         return value
     if isinstance(value, str):
@@ -203,7 +205,12 @@ def _check_unterminated():
     statements = frames[-1].stmts
     while statements:
         last = statements[-1]
-        if isinstance(last, _tir.Return | _tir.Break | _tir.Continue):
+        if isinstance(last, _tir.Return | _tir.Break | _tir.Continue) or (
+            isinstance(last, _tir.Evaluate)
+            and isinstance(last.value, _tir.Call)
+            and isinstance(last.value.op, _ir.Op)
+            and last.value.op.name in ("tirx.break_loop", "tirx.continue_loop")
+        ):
             raise ValueError("An operation cannot follow an unconditional 
terminator")
         if not isinstance(last, _tir.SeqStmt):
             break
@@ -226,6 +233,8 @@ def bind_(
     _check_unterminated()
     with _span_context(span):
         if frame_value:
+            if isinstance(value, _frame.SBlockFrame):
+                raise TypeError("A block does not introduce an as-target 
value")
             if isinstance(value, _python.list | _python.tuple | _ir.Array):
                 for index, item in enumerate(value):
                     bind_(
@@ -240,6 +249,16 @@ def bind_(
             elif isinstance(value, _ir.TensorLoad) and 
_tir.is_buffer_var(value.source):
                 _name(value.source, name, name_span)
             return value
+        if previous is not _MISSING and (
+            _tir.is_buffer_var(previous)
+            or isinstance(previous, _tir.IterVar)
+            or _python.any(
+                isinstance(frame, _frame.SBlockFrame)
+                and _python.any(axis.var.same_as(previous) for axis in 
frame.iter_vars)
+                for frame in _IRBuilder.current().frames
+            )
+        ):
+            raise ValueError(f"Cannot rebind buffer or block axis {name!r}")
         if declaration:
             if not _ir.is_prim_var(value):
                 raise TypeError("A symbol declaration requires a concrete 
primitive variable")
@@ -371,7 +390,7 @@ def break_(*, span=None):
     _require_loop()
     _check_unterminated()
     with _span_context(span):
-        _T.Break()
+        _T.evaluate(_T.break_loop())
 
 
 def continue_(*, span=None):
@@ -379,7 +398,7 @@ def continue_(*, span=None):
     _require_loop()
     _check_unterminated()
     with _span_context(span):
-        _T.Continue()
+        _T.evaluate(_T.continue_loop())
 
 
 def assert_(condition, message="", *, span=None):
@@ -451,7 +470,7 @@ def shared_scalar(dtype="float32"):
     return alloc_scalar(dtype, "shared")
 
 
-@_expression_args("shape", "strides", "elem_offset", introduce=True)
+@_expression_args("shape", "strides", "elem_offset", introduce=True, 
dtype="int32")
 @_wraps(_T.match_buffer)
 def match_buffer(*args, **kwargs):
     """Construct a native buffer match with resolved symbolic shape fields."""
@@ -517,4 +536,4 @@ def select(condition, true_value, false_value):
     """Construct a conditional expression or select an ordinary Python 
value."""
     if not isinstance(condition, _ir.Expr):
         return true_value if condition else false_value
-    return _tir.Select(condition, true_value, false_value)
+    return _tir.if_then_else(condition, true_value, false_value)

Reply via email to