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 78291ea37b5cb9c0c86329c2ba8f09560010eb62
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 18:11:53 2026 +0000

    Complete missing Relax value annotations before normalization
---
 include/tvm/relax/script/builder/ir.h          |  18 ++-
 src/relax/ir/block_builder.cc                  |   6 +-
 src/relax/script/builder/ir.cc                 |  49 +++++-
 tests/python/relax/test_builder_annotations.py | 199 +++++++++++++++++++++++++
 4 files changed, 261 insertions(+), 11 deletions(-)

diff --git a/include/tvm/relax/script/builder/ir.h 
b/include/tvm/relax/script/builder/ir.h
index 86c6fd3dad..c24afc5fff 100644
--- a/include/tvm/relax/script/builder/ir.h
+++ b/include/tvm/relax/script/builder/ir.h
@@ -105,7 +105,11 @@ TVM_DLL void DataflowBlockOutput(const 
ffi::Array<tvm::Var>& vars);
 /*!
  * \brief Emit a binding to the last binding block frame.
  * \param value The right side value of the bindings to be emitted.
- * \param annotate_ty The optional type annotation for the emitted value.
+ * \param annotate_ty Optional output type. Fills missing types on the original
+ * RHS and matching tuple-literal fields before normalization; concrete types
+ * retain their identities and must satisfy the existing compatibility check.
+ * All annotation checks precede completion, so a rejected annotation leaves
+ * the original value types unchanged. Call arguments are not annotated.
  * \return The left side var of the emitted binding.
  */
 TVM_DLL tvm::Var Emit(const tvm::relax::Expr& value,
@@ -126,7 +130,17 @@ TVM_DLL tvm::Var EmitMatchCast(const tvm::relax::Expr& 
value, const tvm::Type& t
  */
 TVM_DLL tvm::Var EmitVarBinding(const tvm::relax::VarBinding& binding);
 
-/*! \brief Emit a binding with separate statement and variable-name ranges. */
+/*!
+ * \brief Emit a binding with separate statement and variable-name ranges.
+ * \param value The original RHS value, completed in place only where its type
+ * is missing, according to Emit's annotation rules.
+ * \param annotate_ty Optional output annotation; concrete type mismatches 
raise
+ * an error before missing types are completed.
+ * \param name_span Optional variable-name range; defaults to the active source
+ * span, which is also recorded for the emitted statement.
+ * \return The emitted variable in the current binding block. No new frame is
+ * entered; annotation and native emission errors propagate.
+ */
 TVM_DLL tvm::Var EmitV2(const tvm::relax::Expr& value, const 
ffi::Optional<tvm::Type>& annotate_ty,
                         const ffi::Optional<Span>& name_span);
 
diff --git a/src/relax/ir/block_builder.cc b/src/relax/ir/block_builder.cc
index 413f41e5bd..54327b843b 100644
--- a/src/relax/ir/block_builder.cc
+++ b/src/relax/ir/block_builder.cc
@@ -658,7 +658,11 @@ class Normalizer : public BlockBuilderImpl, private 
ExprFunctor<Expr(const Expr&
     if (new_op.same_as(op->op) && new_args.same_as(op->args)) {
       call = ffi::GetRef<Call>(op);
     } else {
-      call = Call(Type::Missing(), new_op, new_args, op->attrs, op->ty_args);
+      // Normalizing arguments only names equivalent input values. Preserve an
+      // existing output annotation on the rebuilt call; resetting it to 
Missing
+      // would discard builder-completed types for opaque calls. Unannotated
+      // calls remain Missing and follow the ordinary inference path below.
+      call = Call(op->ty, new_op, new_args, op->attrs, op->ty_args);
     }
 
     if (call->ty.IsMissing()) {
diff --git a/src/relax/script/builder/ir.cc b/src/relax/script/builder/ir.cc
index df6729ed72..87a717b624 100644
--- a/src/relax/script/builder/ir.cc
+++ b/src/relax/script/builder/ir.cc
@@ -232,18 +232,51 @@ TVM_FFI_STATIC_INIT_BLOCK() {
 
 /////////////////////////////// Bindings ///////////////////////////////
 
+namespace {
+
+void CollectMissingValueTypes(const tvm::Expr& value, const tvm::Type& 
annotation,
+                              ffi::Map<tvm::Expr, tvm::Type>* missing_types) {
+  // Record expected types by original value identity, without mutating values
+  // until every concrete field is validated. Shared tuple leaves are checked
+  // against the first annotation collected for that same object.
+  tvm::Type value_type = missing_types->Get(value).value_or(value->ty);
+  if (value_type.IsMissing()) {
+    missing_types->Set(value, annotation);
+  } else {
+    TVM_FFI_ICHECK(tvm::relax::TypeBaseCheck(annotation, value_type) !=
+                   tvm::relax::BaseCheckResult::kFailL0)
+        << "Invalid annotation. Got rhs value type: " << value_type
+        << ", given type: " << annotation;
+  }
+
+  // Tuple literal fields are constituent output values. Complete them before
+  // normalization lifts their bindings and rebuilds the containing tuple;
+  // otherwise a typed tuple can lose its annotation to inferred Any fields.
+  // An annotation never describes arbitrary Call inputs, so do not recurse
+  // through calls or other expression children.
+  if (const auto* tuple = value.as<tvm::TupleNode>()) {
+    if (const auto* tuple_type = annotation.as<tvm::TupleTypeNode>()) {
+      TVM_FFI_ICHECK_EQ(tuple->fields.size(), tuple_type->fields.size())
+          << "Invalid annotation: tuple value and type have different arity";
+      for (size_t i = 0; i < tuple->fields.size(); ++i) {
+        CollectMissingValueTypes(tuple->fields[i], tuple_type->fields[i], 
missing_types);
+      }
+    }
+  }
+}
+
+}  // namespace
+
 tvm::Var Emit(const tvm::relax::Expr& expr, const ffi::Optional<tvm::Type>& 
annotate_ty) {
-  using tvm::relax::GetType;
   BindingBlockFrame block_frame = CheckBindingBlockFrameExistAndUnended();
   const tvm::relax::BlockBuilder& block_builder = GetBlockBuilder();
   if (annotate_ty.has_value()) {
-    const auto& ty = annotate_ty.value();
-    if (expr->ty.IsMissing()) {
-      tvm::relax::UpdateType(expr, ty);
-    } else {
-      TVM_FFI_ICHECK(tvm::relax::TypeBaseCheck(ty, GetType(expr)) !=
-                     tvm::relax::BaseCheckResult::kFailL0)
-          << "Invalid annotation. Got rhs value type: " << GetType(expr) << ", 
given type: " << ty;
+    // This binding-local map lives only until validation/completion finishes.
+    // Update the original RHS objects, not only the newly emitted variable.
+    ffi::Map<tvm::Expr, tvm::Type> missing_types;
+    CollectMissingValueTypes(expr, annotate_ty.value(), &missing_types);
+    for (const auto& [value, type] : missing_types) {
+      tvm::relax::UpdateType(value, type);
     }
   }
   tvm::Var var = block_builder->Emit(expr);
diff --git a/tests/python/relax/test_builder_annotations.py 
b/tests/python/relax/test_builder_annotations.py
new file mode 100644
index 0000000000..2d0216fbe6
--- /dev/null
+++ b/tests/python/relax/test_builder_annotations.py
@@ -0,0 +1,199 @@
+# 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.
+"""Relax annotations complete original missing RHS types before 
normalization."""
+
+import numpy as np
+import pytest
+
+import tvm
+import tvm.testing
+from tvm import ir, relax
+from tvm.relax.script import builder as R
+from tvm.script import relax as SR
+from tvm.script.ir_builder import IRBuilder
+
+
+def _emit(value, annotation):
+    with IRBuilder() as builder:
+        with R.function(is_pure=False):
+            R.func_name("main")
+            result = R.bind_(value, ty=annotation, name="result")
+            R.return_(result)
+    return builder.get(), result
+
+
+def _missing_call():
+    return R.call_packed("test.builder.annotation")
+
+
+def test_missing_rhs_type_and_identity():
+    value = _missing_call()
+    assert value.ty.is_missing()
+    annotation = relax.TensorType([2], "float32")
+    function, result = _emit(value, annotation)
+    assert value.ty.same_as(annotation)
+    assert result.ty.same_as(annotation)
+    assert function.body.blocks[0].bindings[0].value.same_as(value)
+
+
+def test_concrete_rhs_type_is_retained():
+    value = relax.const(np.ones((2,), dtype="float32"))
+    original_type = value.ty
+    function, result = _emit(value, relax.TensorType(dtype="float32"))
+    assert value.ty.same_as(original_type)
+    assert result.ty.same_as(original_type)
+    assert function.body.blocks[0].bindings[0].value.same_as(value)
+
+
+def test_concrete_rhs_mismatch_is_rejected():
+    value = relax.const(np.ones((2,), dtype="float32"))
+    original_type = value.ty
+    with pytest.raises(tvm.error.InternalError, match="Invalid annotation"):
+        _emit(value, relax.TensorType([2], "int32"))
+    assert value.ty.same_as(original_type)
+
+
+def test_nested_tuple_missing_values_keep_identity():
+    first, second = _missing_call(), _missing_call()
+    inner = relax.Tuple([second])
+    value = relax.Tuple([first, inner])
+    first_type = relax.TensorType([2], "float32")
+    second_type = relax.TensorType([3], "int32")
+    inner_type = ir.TupleType([second_type])
+    annotation = ir.TupleType([first_type, inner_type])
+    function, result = _emit(value, annotation)
+    assert first.ty.same_as(first_type)
+    assert second.ty.same_as(second_type)
+    assert inner.ty.same_as(inner_type)
+    assert value.ty.same_as(annotation)
+    tvm.ir.assert_structural_equal(result.ty, annotation)
+    values = [binding.value for block in function.body.blocks for binding in 
block.bindings]
+    assert any(bound.same_as(first) for bound in values)
+    assert any(bound.same_as(second) for bound in values)
+
+
+def test_tuple_rejection_does_not_partially_annotate_values():
+    first = _missing_call()
+    second = relax.const(np.ones((2,), dtype="float32"))
+    original_type = second.ty
+    value = relax.Tuple([first, second])
+    annotation = ir.TupleType([relax.TensorType([2], "float32"), 
ir.PrimType("int32")])
+    with pytest.raises(tvm.error.InternalError, match="Invalid annotation"):
+        _emit(value, annotation)
+    assert first.ty.is_missing()
+    assert value.ty.is_missing()
+    assert second.ty.same_as(original_type)
+
+
+def test_shared_missing_leaf_rejects_conflicting_annotations():
+    leaf = _missing_call()
+    value = relax.Tuple([leaf, leaf])
+    annotation = ir.TupleType([relax.TensorType([2], "float32"), 
ir.PrimType("int32")])
+    with pytest.raises(tvm.error.InternalError, match="Invalid annotation"):
+        _emit(value, annotation)
+    assert leaf.ty.is_missing()
+    assert value.ty.is_missing()
+
+
+def test_tuple_annotation_arity_is_checked_before_completion():
+    leaf = _missing_call()
+    value = relax.Tuple([leaf])
+    with pytest.raises(tvm.error.InternalError, match="different arity"):
+        _emit(value, ir.TupleType([]))
+    assert leaf.ty.is_missing()
+    assert value.ty.is_missing()
+
+
+def test_parser_annotation_completes_original_nested_rhs():
+    captured = []
+
+    def make_value():
+        leaf = _missing_call()
+        value = relax.Tuple([relax.Tuple([leaf])])
+        captured.append((value, leaf))
+        return value
+
+    @SR.function(pure=False)
+    def annotated():
+        result: SR.Tuple(SR.Tuple(SR.Tensor([2], "float32"))) = make_value()
+        return result
+
+    value, leaf = captured[0]
+    expected_leaf = relax.TensorType([2], "float32")
+    expected = ir.TupleType([ir.TupleType([expected_leaf])])
+    tvm.ir.assert_structural_equal(value.ty, expected)
+    tvm.ir.assert_structural_equal(leaf.ty, expected_leaf)
+    tvm.ir.assert_structural_equal(annotated.ret_ty, expected)
+
+
+def test_parser_retains_concrete_rhs_type():
+    value = relax.const(np.ones((2,), dtype="float32"))
+    original_type = value.ty
+
+    def make_value():
+        return value
+
+    @SR.function
+    def annotated():
+        result: SR.Tensor(dtype="float32") = make_value()
+        return result
+
+    assert value.ty.same_as(original_type)
+    tvm.ir.assert_structural_equal(annotated.ret_ty, original_type)
+
+
+def test_parser_rejects_concrete_rhs_mismatch():
+    value = relax.const(np.ones((2,), dtype="float32"))
+    original_type = value.ty
+
+    def make_value():
+        return value
+
+    with pytest.raises(tvm.error.DiagnosticError):
+
+        @SR.function
+        def annotated():
+            result: SR.Tensor([2], "int32") = make_value()
+            return result
+
+    assert value.ty.same_as(original_type)
+
+
[email protected]("annotated", [False, True])
+def 
test_normalized_call_keeps_output_annotation_without_typing_inputs(annotated):
+    argument = _missing_call()
+    value = R.call_packed("test.builder.outer", argument)
+    annotation = relax.TensorType([2], "float32") if annotated else None
+    function, result = _emit(value, annotation)
+    expected = annotation if annotated else relax.AnyType()
+    if annotated:
+        tvm.ir.assert_structural_equal(value.ty, annotation)
+    else:
+        # Ordinary inference acts on the rebuilt call, leaving the original
+        # unannotated input object untouched.
+        assert value.ty.is_missing()
+    tvm.ir.assert_structural_equal(result.ty, expected)
+    # The nested input is inferred independently. A result annotation must not
+    # assign that output tensor type to the opaque call's inputs.
+    assert isinstance(argument.ty, relax.AnyType)
+    binding = function.body.blocks[-1].bindings[-1]
+    assert not binding.value.same_as(value)
+    tvm.ir.assert_structural_equal(binding.value.ty, expected)
+
+
+if __name__ == "__main__":
+    tvm.testing.main()

Reply via email to