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):

Reply via email to