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 d1c62b5179ce0a879599f365fa976cc915da55e7 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 14:52:39 2026 +0000 Start fresh annotation symbols after interrupted definition attempts --- python/tvm/script/parser_v2/annotations.py | 9 +++++++++ tests/python/tvmscript/test_parser_v2.py | 22 ++++++++++++++++------ 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/python/tvm/script/parser_v2/annotations.py b/python/tvm/script/parser_v2/annotations.py index af4983c499..ce63a781b3 100644 --- a/python/tvm/script/parser_v2/annotations.py +++ b/python/tvm/script/parser_v2/annotations.py @@ -507,6 +507,14 @@ def enable_eager_constructors(builder, *, classes=()): states = caller.f_locals.setdefault("__tvm_eager_annotations__", {}) key = (node.name, caller.f_code.co_filename) state = states.get(key) + site = (node.lineno, caller.f_lasti) + # A host argument may fail before entering any constructor. + # Re-entering a definition therefore starts a fresh signature, + # even when its decorator never got a chance to consume it. + if state is not None and ( + site[0] != state._eager_site[0] or site[1] <= state._eager_site[1] + ): + state = None if state is None: state = states[key] = fresh_scope for parameter in node.args.args: @@ -514,6 +522,7 @@ def enable_eager_constructors(builder, *, classes=()): state.rewrite( parameter.annotation, introduce=True, collect_declarations=True ) + state._eager_site = site scope = state else: scope = AnnotationScope( diff --git a/tests/python/tvmscript/test_parser_v2.py b/tests/python/tvmscript/test_parser_v2.py index 1307ca79f6..f9747c1a06 100644 --- a/tests/python/tvmscript/test_parser_v2.py +++ b/tests/python/tvmscript/test_parser_v2.py @@ -322,15 +322,22 @@ def test_annotation_identity_effects_and_recovery(postponed, monkeypatch): from typing import TypeVar from tvm.script import relax as R tensor = R.Tensor - calls, errors, functions = [], [], [] + calls, errors, functions, captures = [], [], [], [] def note(label): calls.append(label) return "float32" - for dtype in ["definitely_invalid_dtype", "float32", "float32"]: + def keep(value): + captures.append(value) + return value + def dtypeval(value): + if value == "argument_error": + raise ValueError("annotation argument failed") + return value + for dtype in ["argument_error", "definitely_invalid_dtype", "float32", "float32"]: M = TypeVar("M") try: @R.function - def f(x: tensor(("n", M), note("x")), y: tensor((8,), dtype)) -> \ + def f(x: keep(tensor(("n", M), note("x"))), y: tensor((8,), dtypeval(dtype))) -> \ 'tensor(("n", M), note("return"))': return x functions.append(f) except Exception as error: @@ -346,9 +353,12 @@ def test_annotation_identity_effects_and_recovery(postponed, monkeypatch): postponed, monkeypatch, ) - assert len(module.errors) == 1 - assert "unknown dtype" in str(module.errors[0]).lower() - assert module.calls == ["x", "x", "return", "x", "return", "object", "object"] + assert len(module.errors) == 2 + assert "annotation argument failed" in str(module.errors[0]) + assert "unknown dtype" in str(module.errors[1]).lower() + assert module.calls == ["x", "x", "x", "return", "x", "return", "object", "object"] + for first, second in zip(module.captures, module.captures[1:]): + assert not first.shape[0].same_as(second.shape[0]) first, second = module.functions for function in module.functions: for argument_dim, return_dim in zip(function.params[0].ty.shape, function.ret_ty.shape):
