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 3ed81254901bea97d9a26d505d9f07264318aa24
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 02:45:52 2026 +0000

    Distinguish implicit shape defaults from explicit type variables
---
 python/tvm/script/ir_builder/protocol.py      | 10 +++++++---
 python/tvm/script/parser_v2/annotations.py    | 10 +++++++++-
 python/tvm/tirx/script/builder_v2/__init__.py |  6 ++++--
 3 files changed, 20 insertions(+), 6 deletions(-)

diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index 740091eb2d..98975d3bd7 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -48,16 +48,19 @@ class ExpressionArguments(NamedTuple):
     introduce: bool = False
     dtype: Any = None
     scalar_strings: bool = True
+    implicit_dtype: Any = None
 
 
-def expression_args(*fields, introduce=False, dtype=None, scalar_strings=True):
+def expression_args(*fields, introduce=False, dtype=None, scalar_strings=True, 
implicit_dtype=None):
     """Mark constructor fields whose nested strings are source expressions.
 
     ``introduce`` permits signature/match scopes to introduce otherwise unknown
     names. ``dtype`` optionally selects the dialect's symbol-construction 
policy;
     it is metadata, never an evaluator or a replacement constructor. Set
     ``scalar_strings=False`` when a bare string is literal shorthand while
-    strings nested in tuples/lists remain expressions.
+    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.
     """
 
     def decorate(constructor):
@@ -66,7 +69,7 @@ def expression_args(*fields, introduce=False, dtype=None, 
scalar_strings=True):
         if unknown:
             raise ValueError(f"Unknown expression argument fields: 
{sorted(unknown)}")
         constructor.__tvm_expression_args__ = ExpressionArguments(
-            tuple(fields), bool(introduce), dtype, bool(scalar_strings)
+            tuple(fields), bool(introduce), dtype, bool(scalar_strings), 
implicit_dtype
         )
         return constructor
 
@@ -175,3 +178,4 @@ def select_lazy(operation, condition, true_value, 
false_value):
     if isinstance(condition, bool):
         return true_value() if condition else false_value()
     return operation(condition, true_value(), false_value())
+
diff --git a/python/tvm/script/parser_v2/annotations.py 
b/python/tvm/script/parser_v2/annotations.py
index 9195d8a5db..a659458a2c 100644
--- a/python/tvm/script/parser_v2/annotations.py
+++ b/python/tvm/script/parser_v2/annotations.py
@@ -327,6 +327,8 @@ class AnnotationScope:
                     scope._canonical_type_var(current.id, current, self.dtype)
                     if self.allow_names and current.id in scope.env:
                         scope._symbol(current.id, current, self.dtype)
+                        if self.in_string:
+                            scope._unbound_parameters.discard(current.id)
                     elif (
                         self.allow_names
                         and self.in_string
@@ -360,7 +362,13 @@ class AnnotationScope:
                             self.in_string = True
                             if self.allow_names and isinstance(current, 
ast.Name):
                                 scope._shape_declarations.setdefault(
-                                    current.id, (current, self.dtype)
+                                    current.id,
+                                    (
+                                        current,
+                                        self.dtype
+                                        if self.dtype is not None
+                                        else metadata.implicit_dtype,
+                                    ),
                                 )
                     return self.visit(current)
                 finally:
diff --git a/python/tvm/tirx/script/builder_v2/__init__.py 
b/python/tvm/tirx/script/builder_v2/__init__.py
index 4d9d7cb6aa..9e28d89d70 100644
--- a/python/tvm/tirx/script/builder_v2/__init__.py
+++ b/python/tvm/tirx/script/builder_v2/__init__.py
@@ -45,7 +45,9 @@ 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, dtype="int32")
+@_expression_args(
+    "shape", "strides", "elem_offset", "byte_offset", introduce=True, 
implicit_dtype="int32"
+)
 def Buffer(
     shape,
     dtype="float32",
@@ -470,7 +472,7 @@ def shared_scalar(dtype="float32"):
     return alloc_scalar(dtype, "shared")
 
 
-@_expression_args("shape", "strides", "elem_offset", introduce=True, 
dtype="int32")
+@_expression_args("shape", "strides", "elem_offset", introduce=True, 
implicit_dtype="int32")
 @_wraps(_T.match_buffer)
 def match_buffer(*args, **kwargs):
     """Construct a native buffer match with resolved symbolic shape fields."""

Reply via email to