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