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

Reply via email to