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 8065c99f7aac6ed2524582fe50761e160592fb93 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 01:39:53 2026 +0000 Preserve declaration metadata and source ranges in transformed expressions --- python/tvm/script/parser_v2/transform.py | 55 ++++++++++++++++++++++++++------ 1 file changed, 46 insertions(+), 9 deletions(-) diff --git a/python/tvm/script/parser_v2/transform.py b/python/tvm/script/parser_v2/transform.py index ab8729c0df..75ba52708f 100644 --- a/python/tvm/script/parser_v2/transform.py +++ b/python/tvm/script/parser_v2/transform.py @@ -150,7 +150,12 @@ class Transformer(ast.NodeTransformer): value = arguments.arguments.get(policy.value_parameter) return value is None or (isinstance(value, ast.Constant) and value.value is None) - def _expression(self, original, *, attach_span=True): + def _declaration_pattern(self, node): + if isinstance(node, ast.Tuple | ast.List): + return [self._declaration_pattern(item) for item in node.elts] + return self._is_declaration(node) + + def _expression(self, original, *, attach_span=True, declarations=False): node = copy.deepcopy(original) if isinstance(getattr(node, "ctx", None), ast.Store): return node @@ -166,6 +171,11 @@ class Transformer(ast.NodeTransformer): ) if self.expression_rewriter is not None: node = self.expression_rewriter(node) + # Policy-generated scaffolding inherits the replaced expression range; + # original descendants retain their more precise source coordinates. + if not hasattr(node, "lineno"): + ast.copy_location(node, original) + ast.fix_missing_locations(node) if isinstance(node, ast.Await | ast.Yield | ast.YieldFrom | ast.NamedExpr): self._error(original, f"Unsupported expression: {type(node).__name__}") if isinstance(node, ast.JoinedStr): @@ -176,7 +186,7 @@ class Transformer(ast.NodeTransformer): if child.format_spec is not None: child.format_spec = self._format_spec(child.format_spec) else: - self._expression_children(node) + self._expression_children(node, declarations) if isinstance(node, ast.Starred) or ( isinstance(node, ast.Name) and isinstance(node.ctx, ast.Store) ): @@ -193,14 +203,21 @@ class Transformer(ast.NodeTransformer): child.format_spec = self._format_spec(child.format_spec) return node - def _expression_children(self, node): + def _expression_children(self, node, declarations=False): for field, value in ast.iter_fields(node): if isinstance(value, ast.expr): setattr(node, field, self._expression(value)) elif isinstance(value, list): for index, item in enumerate(value): if isinstance(item, ast.expr): - value[index] = self._expression(item) + policy = ( + declarations[index] + if field == "elts" and isinstance(declarations, list) + else False + ) + value[index] = self._expression( + item, attach_span=policy is not True, declarations=policy + ) elif isinstance(item, ast.AST): self._expression_children(item) elif isinstance(value, ast.AST): @@ -226,7 +243,7 @@ class Transformer(ast.NodeTransformer): keywords["frame_value"] = ast.Constant(True) if ty is not None: keywords["ty"] = ty - if declaration: + if declaration is True: keywords["declaration"] = ast.Constant(True) if target.id in self.signature_names: keywords["previous"] = copy.deepcopy( @@ -271,10 +288,27 @@ class Transformer(ast.NodeTransformer): ast.Assign([ast.Tuple(pattern, ast.Store())], unpack), target ) result = [assignment] - for item, name in zip(target.elts, names): + star = next( + (index for index, item in enumerate(target.elts) if isinstance(item, ast.Starred)), + None, + ) + for index, (item, name) in enumerate(zip(target.elts, names)): + policy = False + if isinstance(declaration, list): + source_index = index + if star is not None and index > star: + source_index += len(declaration) - len(target.elts) + if index != star and 0 <= source_index < len(declaration): + policy = declaration[source_index] item = item.value if isinstance(item, ast.Starred) else item result.extend( - self._bind(item, self._name(name, item), statement, frame_value=frame_value) + self._bind( + item, + self._name(name, item), + statement, + declaration=policy, + frame_value=frame_value, + ) ) return result self._error(target, f"Unsupported assignment target: {type(target).__name__}") @@ -282,9 +316,12 @@ class Transformer(ast.NodeTransformer): def visit_Assign(self, node): # Cache first: stores evaluate RHS before target base/index, and chained # assignments share exactly one RHS evaluation. - declaration = self._is_declaration(node.value) + declaration = self._declaration_pattern(node.value) cache, value = self._cache( - self._expression(node.value, attach_span=not declaration), node.value + self._expression( + node.value, attach_span=declaration is not True, declarations=declaration + ), + node.value, ) result = [cache] resolved = self._resolve(node.value)
