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)
