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)
