This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch tvmscript-ast-only-transpiler
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit e073ec33b690ec52980e5383a993c2ad5d17dc02
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 17:38:24 2026 +0000

    Share source metadata and make span instrumentation optional
---
 python/tvm/relax/script/builder/__init__.py        | 12 ++--
 python/tvm/script/ir_builder/protocol.py           | 54 ++++++++++------
 python/tvm/script/parser/diagnostics.py            | 14 ++++
 python/tvm/script/parser/frontend.py               | 50 ++++++++++-----
 python/tvm/script/parser/transpile.py              | 24 +++++--
 python/tvm/tirx/script/builder/__init__.py         | 42 ++++++------
 tests/python/tvmscript/test_parser.py              | 75 +++++++++++++++++++---
 .../python/tvmscript/test_tvmscript_parser_tir.py  |  1 -
 8 files changed, 195 insertions(+), 77 deletions(-)

diff --git a/python/tvm/relax/script/builder/__init__.py 
b/python/tvm/relax/script/builder/__init__.py
index ac22ef94df..30e845fd88 100644
--- a/python/tvm/relax/script/builder/__init__.py
+++ b/python/tvm/relax/script/builder/__init__.py
@@ -204,12 +204,12 @@ class _Frame:
         raise AttributeError("This frame does not declare a function")
 
     def __enter__(self):
-        with _protocol.span_context(self.span):
+        with _protocol._construction_span(self.span):
             self.native.__enter__()
         return self
 
     def __exit__(self, exc_type, exc_value, traceback):
-        with _protocol.span_context(self.span):
+        with _protocol._construction_span(self.span):
             self.native.__exit__(exc_type, exc_value, traceback)
         if exc_type is None:
             if isinstance(self.native, _frame.BindingBlockFrame):
@@ -239,7 +239,7 @@ def decl_function(is_pure=True, is_private=False, *, 
local=False, span=None):
 
 def arg(name, ty, *, span=None):
     """Add a parameter, retaining a cached parameter's identity when 
supplied."""
-    with _protocol.span_context(span):
+    with _protocol._construction_span(span):
         if isinstance(ty, _ir.Var):
             return _ffi_api.ArgVar(name, ty)
         return _protocol.at(span, _native.arg(name, _type(ty)))
@@ -358,7 +358,7 @@ def bind_(
         return value.value
     ty = None if ty is None else _type(ty)
     value = _value(value, ty)
-    with _protocol.span_context(span):
+    with _protocol._construction_span(span):
         if isinstance(value, _relax.MatchCast):
             if ty is not None and not _ffi.structural_equal(ty, value.ty):
                 raise TypeError("The binding annotation differs from the 
match-cast type")
@@ -389,7 +389,7 @@ def emit_(value, *, span=None):
 def return_(value=None, *, span=None):
     """Record the function result without exiting Python construction."""
     _check_unterminated()
-    with _protocol.span_context(span):
+    with _protocol._construction_span(span):
         if value is None:
             value = _relax.Tuple([])
         _native.func_ret_value(_value(value))
@@ -416,7 +416,7 @@ def assert_(condition, message="", *, span=None):
     """Construct a runtime assertion with construction-time diagnostic text."""
     if not isinstance(message, _python.str):
         raise TypeError("An assertion message must be construction-time text")
-    with _protocol.span_context(span):
+    with _protocol._construction_span(span):
         emit_(_protocol.at(span, _native.assert_op(condition, 
format=message)), span=span)
 
 
diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index 6fab759059..5d09a7a683 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -119,31 +119,44 @@ def wrap_expression_constructor(constructor, 
call_signature, policy, *, as_type=
 def source_span(location):
     """Materialize an IR span from a source location, or preserve an existing 
span.
 
-    ``location`` is None, an ir.Span, or ``(filename, lineno, end_lineno,
+    ``location`` is None, an ir.Span, or ``(source_name, lineno, end_lineno,
     col_offset, end_col_offset)`` using Python AST's one-based lines and
-    zero-based UTF-8 byte columns. The returned ir.Span retains those exact
-    columns, matching TVMScript's source-range convention.
-    None returns None. Invalid tuple shape/types propagate TypeError/ValueError
-    from unpacking or span construction. This function does not enter a frame.
+    zero-based UTF-8 byte columns. Generated locations carry the single
+    ir.SourceName created for their source unit; this function reuses it 
without
+    construction or caching. A filename string is also accepted for handwritten
+    callers and is converted to ir.SourceName here. The returned ir.Span 
retains
+    the exact coordinates, matching TVMScript's source-range convention.
+
+    None returns None without creating source metadata or accessing a builder.
+    Existing spans pass through unchanged. Invalid tuple shape/types propagate
+    TypeError/ValueError from unpacking or span construction. This function 
does
+    not enter a frame or retain source-unit state.
     """
     if location is None or isinstance(location, ir.Span):
         return location
-    filename, line, end_line, column, end_column = location
-    return ir.Span(ir.SourceName(filename), line, end_line, column, end_column)
+    source_name, line, end_line, column, end_column = location
+    if isinstance(source_name, str):
+        source_name = ir.SourceName(source_name)
+    return ir.Span(source_name, line, end_line, column, end_column)
 
 
 @contextmanager
-def span_context(span):
-    """Temporarily attach source context to nested builder operations.
-
-    ``span`` accepts the same None/ir.Span/location-tuple forms as source_span.
-    Yields None and pushes/pops the existing builder span stack when a builder
-    is active. Without a builder it still materializes diagnostic locations.
+def _construction_span(span):
+    """Apply a builder operation's span through the existing native span stack.
+
+    This private construction helper is used only by builders, never emitted by
+    the transpiler. ``span`` accepts the source_span input forms. None yields
+    directly without IR instrumentation. Otherwise it pushes/pops the existing
+    builder span stack when active, and retains no context after exit. Without
+    a builder it still materializes diagnostic locations.
     Exceptions retain their original type and receive __tvm_script_location__
     as a plain location tuple only when an inner operation has not already
     supplied a more precise range. Diagnostics need never inspect an IR object.
     Span-construction and nested-operation errors propagate unchanged.
     """
+    if span is None:
+        yield
+        return
     span = source_span(span)
     context = (
         IRBuilder.current().with_source_span(span)
@@ -169,14 +182,17 @@ def span_context(span):
 def at(span, value):
     """Attach source metadata to a concrete expression where possible.
 
-    ``span`` is None, ir.Span, or a location tuple; ``value`` is any Python
-    value. Returns the original value, attaching the normalized current source
-    range to IR expressions without a span when a builder is active. Ordinary
-    values and calls without a builder pass through. Temporarily enters the
-    span stack, never a construction frame; invalid spans propagate errors.
+    ``span`` accepts the source_span input forms; ``value`` is any Python 
value.
+    None returns value immediately without creating or inspecting source 
metadata
+    or accessing the builder. Otherwise the normalized source range is attached
+    to IR expressions without a span when a builder is active. Ordinary values
+    and calls without a builder pass through. The operation temporarily enters
+    the native span stack, never a construction frame; invalid spans propagate
+    errors. It does not retain metadata or provide a context-manager/callback
+    form.
     """
     if span is not None and isinstance(value, ir.Expr) and 
IRBuilder.is_in_scope():
-        with span_context(span):
+        with _construction_span(span):
             return IRBuilder.current()._set_current_source_span(value)
     return value
 
diff --git a/python/tvm/script/parser/diagnostics.py 
b/python/tvm/script/parser/diagnostics.py
index 25e0ab660a..f8563b9142 100644
--- a/python/tvm/script/parser/diagnostics.py
+++ b/python/tvm/script/parser/diagnostics.py
@@ -16,6 +16,7 @@
 # under the License.
 """Source-located construction errors without replacing their original 
traceback."""
 
+import ast
 import linecache
 import traceback
 
@@ -52,6 +53,19 @@ def diagnostic_error(error, compiler):
             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)
+            if getattr(frame, "end_lineno", None) is None:
+                # Python before 3.11 supplies only a line. Recover the original
+                # statement range from source AST, without an IR span context 
or
+                # executing/reparsing a generated program. Prefer the narrowest
+                # statement beginning on that line (a nested body over its 
def).
+                candidates = [
+                    node
+                    for node in ast.walk(compiler.tree)
+                    if isinstance(node, ast.stmt) and node.lineno == start
+                ]
+                if candidates:
+                    node = min(candidates, key=lambda item: item.end_lineno - 
item.lineno)
+                    end, column, end_column = node.end_lineno, 
node.col_offset, node.end_col_offset
         else:
             node = compiler.tree.body[-1]
             start, end = node.lineno, node.lineno
diff --git a/python/tvm/script/parser/frontend.py 
b/python/tvm/script/parser/frontend.py
index ea6516cff7..29bca6d3d9 100644
--- a/python/tvm/script/parser/frontend.py
+++ b/python/tvm/script/parser/frontend.py
@@ -28,6 +28,7 @@ from functools import wraps
 from typing import TypeVar
 
 from tvm.error import DiagnosticError
+from tvm.ir import SourceName
 from tvm.script.ir_builder import construction, protocol
 
 from . import protocol as syntax_protocol
@@ -207,12 +208,17 @@ class Compiler:
     filename and compile_flags retain the source's file and annotation mode.
     name_map holds reserved identifiers and next-prefix counters, shared across
     all generated functions. builder_name/infrastructure_name are reserved 
aliases.
-    This entry owns no IR state: build() passes its execution result opaquely.
+    track_span defaults True and owns exactly one SourceName source-metadata
+    object per unit, held in env under source_name_binding. False skips both
+    source-name creation and generated IR span instrumentation; AST 
coordinates,
+    compilation filename and Python tracebacks remain unchanged. No concrete IR
+    construction state is retained: build() passes its result opaquely.
     """
 
-    def __init__(self, source, env=None, filename=None):
+    def __init__(self, source, env=None, filename=None, *, track_span: bool = 
True):
         self.env = {"TypeVar": TypeVar, "tvm": sys.modules.get("tvm"), 
**_NAMESPACES, **(env or {})}
         self.original = source
+        self.track_span = track_span
         members = vars(source).values() if inspect.isclass(source) else 
(source,)
         self.compile_flags = 0
         for member in members:
@@ -249,6 +255,12 @@ class Compiler:
         NameCollector(self.name_map).visit(self.tree)
         self.builder_name = self.fresh()
         self.infrastructure_name = self.fresh()
+        # One source metadata object is shared by this unit and nested 
factories.
+        # This narrow metadata exception never constructs Expr, Type or Span.
+        # The hygienic binding cannot collide with source or captured names.
+        self.source_name_binding = self.fresh() if track_span else None
+        if track_span:
+            self.env[self.source_name_binding] = SourceName(self.filename)
 
     def fresh(self, prefix="_t"):
         """Allocate a generated identifier without changing any source name.
@@ -269,20 +281,25 @@ class Compiler:
         """Emit location data only; builders materialize source spans at 
execution.
 
         node supplies original line/end-line and UTF-8 column/end-column 
fields.
-        The returned AST tuple keeps those values plus the original filename;
-        generated wrappers inherit the full node range via copy_location.
+        The returned tuple references this unit's shared SourceName binding;
+        generated nodes inherit the full range via copy_location. With tracking
+        disabled, return None data and never materialize source metadata.
         """
+        if not self.track_span:
+            return ast.copy_location(ast.Constant(None), node)
         return ast.copy_location(
             ast.Tuple(
                 [
-                    ast.Constant(value)
-                    for value in (
-                        self.filename,
-                        node.lineno,
-                        node.end_lineno,
-                        node.col_offset,
-                        node.end_col_offset,
-                    )
+                    ast.Name(self.source_name_binding, ast.Load()),
+                    *[
+                        ast.Constant(value)
+                        for value in (
+                            node.lineno,
+                            node.end_lineno,
+                            node.col_offset,
+                            node.end_col_offset,
+                        )
+                    ],
                 ],
                 ast.Load(),
             ),
@@ -312,6 +329,7 @@ class Compiler:
             self.span_ast,
             None,
             name_map=self.name_map,
+            track_span=self.track_span,
             **options,
         )
 
@@ -421,12 +439,14 @@ class Compiler:
         return namespace[result]
 
 
-def parse(source, extra_vars=None, *, filename=None, **options):
+def parse(source, extra_vars=None, *, filename=None, track_span: bool = True, 
**options):
     """Transpile and execute a source string, Python function, or Python class.
 
     source is the original object/text; extra_vars optionally supplies lexical
     bindings overriding captured values. filename optionally overrides the 
source
-    filename. Remaining options are accepted for entry-point compatibility; 
actual
+    filename. track_span (default True) enables shared source metadata and IR
+    location instrumentation; False preserves Python locations only. Remaining
+    options are accepted for entry-point compatibility; actual
     construction policy comes from registered decorators in the source.
     Returns the generated builder program's opaque result. Source/host/builder
     failures become DiagnosticError with original-source ranges; an existing
@@ -435,7 +455,7 @@ def parse(source, extra_vars=None, *, filename=None, 
**options):
     """
     env = {} if isinstance(source, str) else _capture(source)
     env.update(extra_vars or {})
-    compiler = Compiler(source, env, filename)
+    compiler = Compiler(source, env, filename, track_span=track_span)
     try:
         return compiler.build()
     except DiagnosticError:
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index f96a02a396..6f45b83e5e 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -137,6 +137,10 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         Shared name reservations/counters; initialized from environment if 
absent.
         NameCollector reserves all source names before whole-program lowering.
 
+    track_span : bool
+        Emit IR source instrumentation, default True. False retains AST ranges
+        without expression at calls or builder span keywords.
+
     State
     -----
     filename, span/fresh callbacks and infrastructure_name belong to one 
complete
@@ -175,8 +179,10 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         nested_function=None,
         preserve_return=False,
         name_map=None,
+        track_span=True,
     ):
         self.filename = filename
+        self.track_span = track_span
         self.namespace_bindings = dict(environment)
         self.dialect_prefix = builder_name
         self.infrastructure_name = infrastructure_name
@@ -207,9 +213,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
     def transform_statements(self, body):
         """Return translated statements from an original ast.stmt sequence.
 
-        body is copied before visiting; source nodes are never mutated. Each
-        translated statement is wrapped in a location context inheriting its
-        original full range. Syntactic bound/optional-name state is updated for
+        body is copied before visiting; source nodes are never mutated. Builder
+        operations receive their statement location directly through span=;
+        source expressions use value-wrapping at calls. Syntactic name state 
is updated for
         subsequent statements in this body. Unsupported syntax raises 
SyntaxError;
         this method never executes callbacks or enters construction frames.
         """
@@ -218,10 +224,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             translated = self.visit(statement)
             if translated is not None:
                 block = translated if isinstance(translated, list) else 
[translated]
-                context = self._call(
-                    self.infrastructure_name, "span_context", 
[self.span(statement)], statement
-                )
-                result.append(self._with(context, block, statement))
+                result.extend(block)
         return result
 
     def _error(self, node, message):
@@ -240,6 +243,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         )
 
     def _call(self, namespace, member, args, original, **keywords):
+        # Disabled tracking removes instrumentation, while copy_location below
+        # preserves Python diagnostics for the same generated computation.
+        if not self.track_span:
+            if namespace == self.infrastructure_name and member == "_at":
+                return args[1]
+            keywords.pop("span", None)
+            keywords.pop("name_span", None)
         return self._located(
             ast.Call(
                 self._attribute(namespace, member, original),
diff --git a/python/tvm/tirx/script/builder/__init__.py 
b/python/tvm/tirx/script/builder/__init__.py
index 5780e10cd8..27da1890dc 100644
--- a/python/tvm/tirx/script/builder/__init__.py
+++ b/python/tvm/tirx/script/builder/__init__.py
@@ -31,10 +31,10 @@ from tvm.script.ir_builder import IRBuilder as _IRBuilder
 from tvm.script.ir_builder import ir as _I
 from tvm.script.ir_builder.base import IRBuilderFrame as _NativeFrame
 from tvm.script.ir_builder.protocol import MISSING as _MISSING
+from tvm.script.ir_builder.protocol import _construction_span
 from tvm.script.ir_builder.protocol import at as _at
 from tvm.script.ir_builder.protocol import register_call_kind as 
_register_call_kind
 from tvm.script.ir_builder.protocol import source_span as _source_span
-from tvm.script.ir_builder.protocol import span_context as _span_context
 from tvm.script.ir_builder.type_var_frame import TypeVarDecl as _TypeVarDecl
 from tvm.script.ir_builder.type_var_frame import TypeVarFrame as _TypeVarFrame
 from tvm.script.ir_builder.type_var_frame import resolve_type_var
@@ -110,7 +110,7 @@ def Buffer(
     eager annotation expressions return MissingType; dtype/scope/name strings
     remain literals. Handwritten calls in builders require concrete 
expressions.
     """
-    with _span_context(span):
+    with _construction_span(span):
         return _at(
             span,
             _native.buffer(
@@ -141,7 +141,7 @@ def Ptr(dtype, storage_scope="global", *, span=None):
         dtype = dtype.ty
     if isinstance(dtype, _ir.PrimType):
         dtype = dtype.dtype
-    with _span_context(span):
+    with _construction_span(span):
         return _at(span, _native.ptr(dtype, storage_scope))
 
 
@@ -158,12 +158,12 @@ class _Frame:
         self.result = {}
 
     def __enter__(self):
-        with _span_context(self.span):
+        with _construction_span(self.span):
             value = self.native.__enter__()
         return self if value is self.native else value
 
     def __exit__(self, *exc):
-        with _span_context(self.span):
+        with _construction_span(self.span):
             return self.native.__exit__(*exc)
 
     @property
@@ -177,13 +177,13 @@ class _Frame:
 
 def function(*, private=False, s_tir=False, persistent=False, span=None):
     """Enter a native primitive-function definition frame."""
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_native.prim_func(private=private, s_tir=s_tir, 
persistent=persistent), span)
 
 
 def decl_function(*, private=False, s_tir=False, persistent=False, span=None):
     """Declare a bodyless signature using the native function frame."""
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_ffi_api.DeclFunction(private, s_tir, persistent), span)
 
 
@@ -195,7 +195,7 @@ def arg(name, annotation, *, span=None):
         annotation = annotation.ty
     if isinstance(annotation, _ir.Type):
         annotation = _ir.Var(name, annotation)
-    with _span_context(span):
+    with _construction_span(span):
         if _tir.is_buffer_var(annotation) and annotation.ty.layout is not None:
             frames = _IRBuilder.current().frames
             if _python.any(
@@ -223,7 +223,7 @@ def func_ret_type(annotation, *, span=None):
         annotation = annotation()
     if isinstance(annotation, _ir.Expr | _TypeVarDecl):
         annotation = annotation.ty
-    with _span_context(span):
+    with _construction_span(span):
         return _native.func_ret(annotation)
 
 
@@ -293,7 +293,7 @@ def bind_(
             if not _ffi.structural_equal(annotation, value.ty):
                 raise TypeError("The symbol declaration has an incompatible 
type")
         return _TypeVarFrame.current().resolve(name, value.ty, span=name_span)
-    with _span_context(span):
+    with _construction_span(span):
         if frame_value:
             if isinstance(value, _frame.SBlockFrame):
                 raise TypeError("A block does not introduce an as-target 
value")
@@ -423,7 +423,7 @@ def emit_(value, *, span=None):
     """Consume one expression statement, including effect-only calls."""
     if value is None or isinstance(value, str | _ir.Var):
         return
-    with _span_context(span):
+    with _construction_span(span):
         if isinstance(value, _NativeFrame | _Frame):
             _enter_concise(value)
         elif hasattr(value, "frames"):
@@ -437,7 +437,7 @@ def emit_(value, *, span=None):
 
 def setitem(target, key, value, *, span=None):
     """Construct an indexed store after the caller has evaluated its 
operands."""
-    with _span_context(span):
+    with _construction_span(span):
         buffer_store(target, value, key)
 
 
@@ -462,7 +462,7 @@ def return_(value=None, *, span=None):
     """Emit a native return; subsequent unreachable statements remain in the 
IR."""
     if value is None:
         raise TypeError("A primitive function return requires an expression")
-    with _span_context(span):
+    with _construction_span(span):
         _native.Return(_as_expr(value))
 
 
@@ -478,14 +478,14 @@ def _require_loop():
 def break_(*, span=None):
     """Construct a break targeting the nearest primitive loop."""
     _require_loop()
-    with _span_context(span):
+    with _construction_span(span):
         _native.evaluate(_native.break_loop())
 
 
 def continue_(*, span=None):
     """Construct a continue targeting the nearest primitive loop."""
     _require_loop()
-    with _span_context(span):
+    with _construction_span(span):
         _native.evaluate(_native.continue_loop())
 
 
@@ -500,23 +500,23 @@ def assert_(condition, message="", *, span=None):
         message = [str(part) for part in message]
     if not isinstance(message, list | tuple):
         message = [message]
-    with _span_context(span):
+    with _construction_span(span):
         with _native.Assert(condition, message, error_kind=kind):
             pass
 
 
 def If(condition, *, span=None):
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_native.If(condition), span)
 
 
 def Then(*, span=None):
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_native.Then(), span)
 
 
 def Else(*, span=None):
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_native.Else(), span)
 
 
@@ -635,7 +635,7 @@ def for_(iterable, *, names=None, span=None):
             raise TypeError("Loop names must be a source identifier or tuple 
of identifiers")
         if _python.sum(name.startswith("*") for name in names) > 1:
             raise ValueError("Loop targets may contain only one starred group")
-    with _span_context(span):
+    with _construction_span(span):
         if isinstance(iterable, _IterationSpec):
             if iterable.kind == "grid":
                 iterable = _native.grid(*iterable.arguments, 
dtype=iterable.dtype)
@@ -653,7 +653,7 @@ For = for_
 
 
 def While(condition, *, span=None):
-    with _span_context(span):
+    with _construction_span(span):
         return _Frame(_native.While(condition), span)
 
 
diff --git a/tests/python/tvmscript/test_parser.py 
b/tests/python/tvmscript/test_parser.py
index 06f81e5d72..55d6f1877c 100644
--- a/tests/python/tvmscript/test_parser.py
+++ b/tests/python/tvmscript/test_parser.py
@@ -66,10 +66,10 @@ class _Recorder:
         return value
 
 
-def _registered(source, env=None):
+def _registered(source, env=None, **options):
     recorder = _Recorder()
     namespace = {"D": SimpleNamespace(function=make_decorator(recorder)), 
**(env or {})}
-    compiler = Compiler(textwrap.dedent(source), namespace, 
filename="<protocol-test>")
+    compiler = Compiler(textwrap.dedent(source), namespace, 
filename="<protocol-test>", **options)
     function = compiler.tree.body[0]
     kind, _ = compiler.function_kind(function, compiler.env)
     assert kind.builder is recorder
@@ -111,11 +111,9 @@ def test_source_translation():
     actual = ast.unparse(ast.fix_missing_locations(generated))
     expected = """\
 def f(x):
-    with I.span_context(S):
-        _value_0 = I._at(S, I._at(S, I._at(S, x) * I._at(S, x)) + I._at(S, 1))
-        y = X.bind_(_value_0, span=S, name='y', name_span=S)
-    with I.span_context(S):
-        X.return_(I._at(S, y), span=S)"""
+    _value_0 = I._at(S, I._at(S, I._at(S, x) * I._at(S, x)) + I._at(S, 1))
+    y = X.bind_(_value_0, span=S, name='y', name_span=S)
+    X.return_(I._at(S, y), span=S)"""
     assert actual == expected
     assert ast.dump(function, include_attributes=True) == original
 
@@ -293,7 +291,9 @@ def test_shared_parser_dependency_direction():
             elif isinstance(node, ast.Constant) and isinstance(node.value, 
str):
                 values = [node.value]
             if isinstance(node, ast.ImportFrom):
-                assert node.module not in {"tvm.ir", "tvm.ir.prim"}, path
+                if node.module in {"tvm.ir", "tvm.ir.prim"}:
+                    assert path.name == "frontend.py" and node.module == 
"tvm.ir", path
+                    assert [alias.name for alias in node.names] == 
["SourceName"], path
                 if node.module == "tvm":
                     assert all(alias.name != "ir" for alias in node.names), 
path
             if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
@@ -544,3 +544,62 @@ def outer():
     assert [item[0] for item in outer._tvm_type_var_declarations] == ["a"]
     assert [item[0] for item in inner._tvm_type_var_declarations] == ["b"]
     assert {"outer", "inner", "a", "b", "hidden"}.issubset(compiler.name_map)
+
+
[email protected]("track_span", [True, False])
+def test_optional_shared_source_metadata(track_span, monkeypatch):
+    from tvm.script.parser import frontend
+
+    source_names = []
+    original = frontend.SourceName
+
+    def source_name(filename):
+        value = original(filename)
+        source_names.append(value)
+        return value
+
+    monkeypatch.setattr(frontend, "SourceName", source_name)
+    source = """\
+        @D.function
+        def f(x):
+            y = x * x + 1
+            return y
+    """
+    compiler, function, builder = _registered(
+        source, {"x": ir.Var("x", "int32")}, track_span=track_span
+    )
+    generated = compiler.transformer().transform_statements(function.body)
+    calls = [node for statement in generated for node in ast.walk(statement)]
+    assert not any(
+        isinstance(node, ast.Attribute) and node.attr == "span_context" for 
node in calls
+    )
+    if not track_span:
+        assert not any(isinstance(node, ast.Attribute) and node.attr == "_at" 
for node in calls)
+        assert not any(
+            isinstance(node, ast.keyword) and node.arg in {"span", 
"name_span"} for node in calls
+        )
+    _run(compiler, function, builder)
+    addition = builder.returned[0]
+    if track_span:
+        assert len(source_names) == 1
+        assert addition.span.source_name.same_as(source_names[0])
+        assert addition.a.span.source_name.same_as(source_names[0])
+        assert (addition.a.span.line, addition.a.span.column, 
addition.a.span.end_column) == (
+            3,
+            8,
+            13,
+        )
+    else:
+        assert source_names == []
+        assert addition.span is None and addition.a.span is None
+
+    # Disabling IR instrumentation must preserve the original Python exception
+    # location, including multi-line statements on runtimes without column 
data.
+    with pytest.raises(DiagnosticError) as error:
+        parser.parse(
+            "@T.prim_func\ndef fail():\n    T.evaluate(\n        
missing_name\n    )\n",
+            filename="<tracking-error>",
+            track_span=track_span,
+        )
+    assert "<tracking-error>:" in str(error.value)
+    assert "missing_name" in str(error.value)
diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py 
b/tests/python/tvmscript/test_tvmscript_parser_tir.py
index e6cabc423b..ee53bbd9a5 100644
--- a/tests/python/tvmscript/test_tvmscript_parser_tir.py
+++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py
@@ -183,7 +183,6 @@ def 
test_tir_return_annotation_does_not_define_symbolic_var():
             """
 @T.prim_func
 def main() -> T.Buffer(("n",), "float32"):
-    n = T.int32()
     A = T.alloc_buffer((n,), "float32")
     return A
 """

Reply via email to