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 58ac1b8098f1137c799c26e48e23c7dd381de08c Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 01:40:09 2026 +0000 Preserve source environments and construction diagnostics --- python/tvm/script/ir_builder/protocol.py | 14 +++++-- python/tvm/script/parser_v2/diagnostics.py | 63 ++++++++++++++++++++++++++++ python/tvm/script/parser_v2/frontend.py | 66 ++++++++++++++++++++++++++---- 3 files changed, 133 insertions(+), 10 deletions(-) diff --git a/python/tvm/script/ir_builder/protocol.py b/python/tvm/script/ir_builder/protocol.py index fb4a73cb5a..1813bd61b7 100644 --- a/python/tvm/script/ir_builder/protocol.py +++ b/python/tvm/script/ir_builder/protocol.py @@ -23,7 +23,7 @@ translation consumes construction policies without importing their owners. from builtins import locals as locals from builtins import slice as slice -from contextlib import nullcontext +from contextlib import contextmanager, nullcontext from dataclasses import dataclass from inspect import signature from typing import Any, NamedTuple @@ -105,9 +105,17 @@ def function_kind(decorator): return getattr(decorator, "__tvm_function_kind__", None) +@contextmanager def span_context(span): - """Use the existing builder's source-span stack for an eager operation.""" - return IRBuilder.current().with_source_span(span) if span is not None else nullcontext() + """Use the existing span stack, preserving a failing operation's source range.""" + context = IRBuilder.current().with_source_span(span) if span is not None else nullcontext() + try: + with context: + yield + except Exception as error: + if span is not None and not hasattr(error, "__tvm_script_span__"): + error.__tvm_script_span__ = span + raise def at(span, value): diff --git a/python/tvm/script/parser_v2/diagnostics.py b/python/tvm/script/parser_v2/diagnostics.py new file mode 100644 index 0000000000..430cac796a --- /dev/null +++ b/python/tvm/script/parser_v2/diagnostics.py @@ -0,0 +1,63 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Source-located construction errors without replacing their original traceback.""" + +import linecache +import traceback + +from tvm.error import DiagnosticError + + +def diagnostic_error(error, compiler): + """Render the innermost original operation and retain the original cause.""" + filename = compiler.filename + span = getattr(error, "__tvm_script_span__", None) + if span is not None: + start, end = span.line, span.end_line + column, end_column = span.column, span.end_column + elif isinstance(error, SyntaxError) and error.lineno: + start = error.lineno + end = error.end_lineno or start + column = max((error.offset or 1) - 1, 0) + end_column = max((error.end_offset or column + 2) - 1, column + 1) + else: + frames = [ + frame + for frame in traceback.extract_tb(error.__traceback__) + if frame.filename == filename + ] + if frames: + frame = frames[-1] + start, end = frame.lineno, getattr(frame, "end_lineno", None) or frame.lineno + column = getattr(frame, "colno", None) or 0 + end_column = getattr(frame, "end_colno", None) + else: + node = compiler.tree.body[-1] + start, end = node.lineno, node.lineno + column, end_column = node.col_offset, None + lines = [f"{filename}:{start}: {type(error).__name__}: {error}"] + for number in range(start, end + 1): + source = linecache.getline(filename, number).rstrip("\n") + first = column if number == start else len(source) - len(source.lstrip()) + last = end_column if number == end and end_column is not None else len(source) + lines.extend( + ( + f" {number} | {source}", + " " * (len(str(number)) + 4 + first) + "^" * max(last - first, 1), + ) + ) + return DiagnosticError("\n".join(lines)) diff --git a/python/tvm/script/parser_v2/frontend.py b/python/tvm/script/parser_v2/frontend.py index f96fa413f2..a54a407c28 100644 --- a/python/tvm/script/parser_v2/frontend.py +++ b/python/tvm/script/parser_v2/frontend.py @@ -16,6 +16,8 @@ # under the License. """Source acquisition and declaration/body execution for registered builders.""" +import __future__ + import ast import copy import inspect @@ -24,12 +26,15 @@ import textwrap from dataclasses import dataclass, field from functools import wraps from types import SimpleNamespace +from typing import TypeVar from tvm import ir +from tvm.error import DiagnosticError from tvm.script.ir_builder import IRBuilder, protocol from tvm.script.ir_builder import ir as I from .annotations import AnnotationScope +from .diagnostics import diagnostic_error from .functions import FunctionGroup, attach_python, declare_python, is_python_function _NAMESPACES = {} @@ -47,7 +52,9 @@ def _resolve(node, env, filename): def _capture(obj): target = obj if inspect.isfunction(obj) else None - env = dict(getattr(target, "__globals__", {})) + module = inspect.getmodule(obj) + env = dict(vars(module)) if module is not None else {} + env.update(getattr(target, "__globals__", {})) if target is not None: closure = inspect.getclosurevars(target) env.update(closure.globals) @@ -106,11 +113,16 @@ def make_helper(builder, *, preserve_return=True): def decorator(function=None, **options): def apply(function): + definition_env = _capture(function) + @wraps(function) def invoke(*args, **kwargs): bound = inspect.signature(function).bind(*args, **kwargs) bound.apply_defaults() - compiler = Compiler(function, _capture(function)) + environment = ( + definition_env if options.get("hygienic", True) else _capture(function) + ) + compiler = Compiler(function, environment) node = compiler.tree.body[0] return compiler.run_statements( node.body, @@ -152,8 +164,14 @@ class Compiler: """Build from a copied original AST, retaining file and range information.""" def __init__(self, source, env=None, filename=None): - self.env = {**_NAMESPACES, **(env or {})} + self.env = {"TypeVar": TypeVar, **_NAMESPACES, **(env or {})} self.original = source + members = vars(source).values() if inspect.isclass(source) else (source,) + self.compile_flags = 0 + for member in members: + code = getattr(member, "__code__", None) + if code is not None: + self.compile_flags |= code.co_flags & __future__.annotations.compiler_flag if isinstance(source, str): text = source self.filename = filename or "<tvmscript>" @@ -292,6 +310,7 @@ class Compiler: set(spec.params) | set(spec.scope.symbols), scope=spec.scope, ) + frame.function.__name__ = spec.node.name return frame.function def run_statements(self, body, builder, env, bound_names, *, scope=None, preserve_return=False): @@ -411,11 +430,22 @@ class Compiler: ) ast.copy_location(helper, body[0]) module = ast.fix_missing_locations(ast.Module([helper], [])) - exec(compile(module, self.filename, "exec"), namespace) + exec( + compile(module, self.filename, "exec", flags=self.compile_flags, dont_inherit=True), + namespace, + ) return namespace[helper_name](*(namespace[name] for name in names)) def build(self): nodes = self.tree.body + env = dict(self.env) + if nodes and isinstance(nodes[-1], ast.FunctionDef | ast.ClassDef): + prefix, nodes = nodes[:-1], nodes[-1:] + if any(isinstance(node, ast.FunctionDef | ast.ClassDef) for node in prefix): + raise SyntaxError("Source must contain one function or module class") + if prefix: + setup = ast.fix_missing_locations(ast.Module(copy.deepcopy(prefix), [])) + exec(compile(setup, self.filename, "exec"), env) if len(nodes) == 1 and isinstance(nodes[0], ast.ClassDef): root = nodes[0] statements = root.body @@ -424,7 +454,6 @@ class Compiler: statements = nodes else: raise SyntaxError("Source must contain one function or module class") - env = dict(self.env) functions = [node for node in statements if isinstance(node, ast.FunctionDef)] python_functions = [] with IRBuilder() as builder: @@ -445,6 +474,20 @@ class Compiler: ), env, ) + if isinstance(statement, ast.Assign | ast.AnnAssign): + targets = ( + statement.targets + if isinstance(statement, ast.Assign) + else [statement.target] + ) + for target in targets: + if isinstance(target, ast.Name): + value = env[target.id] + if isinstance(value, ir.BaseFunc): + reference = I.decl_function(target.id, value) + I.def_function(target.id, value) + env[target.id] = reference + setattr(env[root.name], target.id, env[target.id]) ir_functions = [] for node in functions: if is_python_function(self, node, env): @@ -464,14 +507,23 @@ class Compiler: module = builder.get() if python_functions: attach_python(module, python_functions) - return module if root is not None else results[functions[0].name] + if root is not None: + module.__name__ = root.name + return module + return results[functions[0].name] def parse(source, extra_vars=None, *, filename=None, **options): """Construct from source using only entry-module registered construction policies.""" env = {} if isinstance(source, str) else _capture(source) env.update(extra_vars or {}) - return Compiler(source, env, filename).build() + compiler = Compiler(source, env, filename) + try: + return compiler.build() + except DiagnosticError: + raise + except Exception as error: + raise diagnostic_error(error, compiler) from error def ir_module(module=None, **options):
