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