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."""
