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

tqchen pushed a commit to branch script/canonical-parser-df
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit c9ce33b0ebf4dff1c9fd9cfb9bfdddb4ec4993f3
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 00:42:06 2026 +0000

    [FR] Recover class source from its decoration location
---
 python/tvm/script/parser/frontend.py     |  49 ++++++++++++--
 tests/script/test_parser_class_source.py | 110 +++++++++++++++++++++++++++++++
 2 files changed, 154 insertions(+), 5 deletions(-)

diff --git a/python/tvm/script/parser/frontend.py 
b/python/tvm/script/parser/frontend.py
index 347148d809..b12ac0ae2f 100644
--- a/python/tvm/script/parser/frontend.py
+++ b/python/tvm/script/parser/frontend.py
@@ -448,6 +448,27 @@ def pyfunc(function):
 syntax_protocol.register_function(pyfunc, None, python=True)
 
 
+def _source_lines(source, definition_source):
+    """Recover a class from its exact decoration site when module inspection 
fails."""
+    try:
+        lines, start = inspect.getsourcelines(source)
+        return lines, start, inspect.getsourcefile(source)
+    except OSError:
+        if not inspect.isclass(source) or definition_source is None:
+            raise
+        filename, lineno = definition_source
+        lines = linecache.getlines(filename)
+        # Gallery runners may execute the class in a temporary __main__ module
+        # without __file__. Its decorator still has the original code location.
+        tree = ast.parse("".join(lines), filename)
+        for node in ast.walk(tree):
+            if isinstance(node, ast.ClassDef) and node.name == source.__name__:
+                start = min([node.lineno, *(item.lineno for item in 
node.decorator_list)])
+                if start <= lineno <= node.lineno:
+                    return lines[start - 1 : node.end_lineno], start, filename
+        raise
+
+
 class Compiler:
     """Acquire source and execute a location-preserving builder program.
 
@@ -490,7 +511,14 @@ class Compiler:
     """
 
     def __init__(
-        self, source, env=None, filename=None, *, track_span: bool = True, 
definition_scope=None
+        self,
+        source,
+        env=None,
+        filename=None,
+        *,
+        track_span: bool = True,
+        definition_scope=None,
+        definition_source=None,
     ):
         self.env = {"TypeVar": TypeVar, "tvm": sys.modules.get("tvm"), 
**_NAMESPACES, **(env or {})}
         self.original = source
@@ -513,9 +541,9 @@ class Compiler:
                 self.filename,
             )
         else:
-            lines, start = inspect.getsourcelines(source)
+            lines, start, source_filename = _source_lines(source, 
definition_source)
             text = "".join(lines)
-            self.filename = filename or inspect.getsourcefile(source)
+            self.filename = filename or source_filename
             indent = len(lines[0]) - len(lines[0].lstrip())
         self.tree = ast.parse(textwrap.dedent(text), self.filename)
         if start != 1:
@@ -921,7 +949,12 @@ def parse(source, extra_vars=None, *, filename=None, 
track_span: bool = True, **
         "_definition_scope", getattr(source, "__tvm_definition_scope__", {})
     )
     compiler = Compiler(
-        source, env, filename, track_span=track_span, 
definition_scope=definition_scope
+        source,
+        env,
+        filename,
+        track_span=track_span,
+        definition_scope=definition_scope,
+        definition_source=options.pop("_definition_source", None),
     )
     try:
         root = compiler.tree.body[-1]
@@ -1038,9 +1071,15 @@ def ir_module(module=None, **options):
             if frame.f_code is ir_module.__code__:
                 frame = frame.f_back
             definition_scope = _definition_scope(frame)
+            definition_source = (frame.f_code.co_filename, frame.f_lineno)
         finally:
             del frame
-        result = parse(module, _definition_scope=definition_scope, **options)
+        result = parse(
+            module,
+            _definition_scope=definition_scope,
+            _definition_source=definition_source,
+            **options,
+        )
         from tvm.relax.base_py_module import BasePyModule
 
         if issubclass(module, BasePyModule):
diff --git a/tests/script/test_parser_class_source.py 
b/tests/script/test_parser_class_source.py
new file mode 100644
index 0000000000..a00a0c4271
--- /dev/null
+++ b/tests/script/test_parser_class_source.py
@@ -0,0 +1,110 @@
+# 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.
+"""Class source recovery uses the exact decorator location in gallery 
runners."""
+
+import importlib.util
+import linecache
+import sys
+
+import pytest
+
+from tvm.script import ir as I
+from tvm.script import tirx as T
+
+
+def _execute(source, filename, monkeypatch):
+    module = 
importlib.util.module_from_spec(importlib.util.spec_from_loader("__main__", 
None))
+    module.__dict__.update(I=I, T=T)
+    with monkeypatch.context() as patch:
+        patch.setitem(sys.modules, "__main__", module)
+        exec(compile(source, filename, "exec"), module.__dict__)
+    return module
+
+
[email protected]("cached", [False, True])
+def test_gallery_classes_retain_distinct_source_locations(tmp_path, 
monkeypatch, cached):
+    source = """
[email protected]_module
+class Repeated:
+    @T.prim_func
+    def main(A: T.Buffer((1,), "int32")):
+        A[0] = 11
+first = Repeated
+
[email protected]_module()
+class Repeated:
+    @T.prim_func
+    def main(A: T.Buffer((1,), "int32")):
+        A[0] = 22
+second = Repeated
+"""
+    filename = "<gallery-cached-classes>" if cached else str(tmp_path / 
"gallery.py")
+    if cached:
+        monkeypatch.setitem(
+            linecache.cache, filename, (len(source), None, 
source.splitlines(True), filename)
+        )
+    else:
+        (tmp_path / "gallery.py").write_text(source)
+    module = _execute(source, filename, monkeypatch)
+    for result, value, line in ((module.first, 11, 6), (module.second, 22, 
13)):
+        store = result["main"].body
+        assert store.value.value == value
+        assert store.span.line == line
+        assert store.span.source_name.name == filename
+
+
+def test_gallery_nested_class_preserves_definition_scope(tmp_path, 
monkeypatch):
+    source = """
+VALUE = 7
+
+def build(n):
+    @I.ir_module
+    class Module:
+        @T.prim_func
+        def main(A: T.Buffer((n,), "int32")):
+            A[0] = VALUE
+    return Module
+
+def caller():
+    VALUE = 99
+    return build(4)
+
+result = caller()
+"""
+    path = tmp_path / "nested_gallery.py"
+    path.write_text(source)
+    module = _execute(source, str(path), monkeypatch)
+    function = module.result["main"]
+    assert function.params[0].ty.shape[0].value == 4
+    assert function.body.value.value == 7
+
+
+def test_gallery_empty_class_needs_no_member_source(tmp_path, monkeypatch):
+    source = """
[email protected]_module
+class Empty:
+    pass
+"""
+    path = tmp_path / "empty_gallery.py"
+    path.write_text(source)
+    module = _execute(source, str(path), monkeypatch)
+    assert len(module.Empty.functions) == 0
+
+
+def test_unavailable_class_source_still_raises(monkeypatch):
+    with pytest.raises(OSError, match="source code not available"):
+        _execute("@I.ir_module\nclass Empty:\n    pass\n", 
"<missing-gallery-source>", monkeypatch)

Reply via email to