From a6f068dc1ed291693b7680290edc40072848a44a Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 20:50:21 +0000 Subject: [PATCH 1/2] [FIX][IR] Share typed expression mismatch diagnostics --- include/tvm/ir/base_expr.h | 38 +++++++++------------ tests/cpp/expr_test.cc | 17 +++++++++ tests/python/tirx-base/test_tir_op_types.py | 2 +- 3 files changed, 35 insertions(+), 22 deletions(-) diff --git a/include/tvm/ir/base_expr.h b/include/tvm/ir/base_expr.h index 4302a570b2e5..c61a09c4e579 100644 --- a/include/tvm/ir/base_expr.h +++ b/include/tvm/ir/base_expr.h @@ -489,6 +489,20 @@ struct TypeTraits> const auto* expr = details::ObjectUnsafe::RawObjectPtrFromUnowned(src->v_obj); return details::AnyUnsafe::CheckAnyViewStrict(expr->ty); } + + TVM_FFI_INLINE static std::string GetMismatchTypeInfo(const TVMFFIAny* src) { + if (src->type_index >= TypeIndex::kTVMFFIStaticObjectBegin && src->v_obj != nullptr && + details::IsObjectInstance(src->type_index)) { + const auto* expr = details::ObjectUnsafe::RawObjectPtrFromUnowned(src->v_obj); + if (expr->ty.defined()) { + const auto* prim_ty = expr->ty.as(); + std::string ty = + prim_ty ? std::string(ffi::DLDataTypeToString(prim_ty->dtype)) : expr->ty.GetTypeKey(); + return TypeIndexToTypeKey(src->type_index) + "[ty=" + ty + "]"; + } + } + return TypeTraitsBase::GetMismatchTypeInfo(src); + } }; template <> @@ -499,9 +513,6 @@ inline constexpr bool use_default_type_traits_v = false; template <> struct TypeTraits : public ObjectRefWithFallbackTraitsBase { - using Base = - ObjectRefWithFallbackTraitsBase; - TVM_FFI_INLINE static bool CheckAnyStrict(const TVMFFIAny* src) { if (src->type_index == TypeIndex::kTVMFFINone) return PrimExpr::_type_is_nullable; return TypeTraits>::CheckAnyStrict(src); @@ -509,15 +520,8 @@ struct TypeTraits : public ObjectRefWithFallbackTraitsBasetype_index >= TypeIndex::kTVMFFIStaticObjectBegin && source->v_obj != nullptr && - details::IsObjectInstance(source->type_index)) { - const auto* expr = details::ObjectUnsafe::RawObjectPtrFromUnowned(source->v_obj); - if (expr->ty.defined() && !details::AnyUnsafe::CheckAnyViewStrict(expr->ty)) { - return TypeIndexToTypeKey(source->type_index) + "[ty=" + expr->ty.GetTypeKey() + "]"; - } - } - return Base::GetMismatchTypeInfo(source); + TVM_FFI_INLINE static std::string GetMismatchTypeInfo(const TVMFFIAny* src) { + return TypeTraits>::GetMismatchTypeInfo(src); } TVM_DLL static PrimExpr ConvertFallbackValue(StrictBool value); @@ -551,15 +555,7 @@ struct TypeTraits : public ObjectRefWithFallbackTraitsBasetype_index >= TypeIndex::kTVMFFIStaticObjectBegin && - details::IsObjectInstance(src->type_index)) { - const auto* expr = details::ObjectUnsafe::RawObjectPtrFromUnowned(src->v_obj); - if (const auto* ty = expr->ty.as()) { - return TypeIndexToTypeKey(src->type_index) + - "[dtype=" + ffi::DLDataTypeToString(ty->dtype) + "]"; - } - } - return TypeTraits::GetMismatchTypeInfo(src); + return TypeTraits>::GetMismatchTypeInfo(src); } }; diff --git a/tests/cpp/expr_test.cc b/tests/cpp/expr_test.cc index 067962a6e391..19de1e949252 100644 --- a/tests/cpp/expr_test.cc +++ b/tests/cpp/expr_test.cc @@ -111,6 +111,23 @@ TEST(Expr, IntegerView) { EXPECT_EQ(ffi::Any(int64_t{1} << 40).cast().ty().bits(), 64); } +TEST(Expr, TypedMismatchDiagnostics) { + using namespace tvm; + auto check = [](ffi::AnyView value, const std::string& expected) { + TVMFFIAny raw = value.CopyToTVMFFIAny(); + EXPECT_EQ(ffi::TypeTraits>::GetMismatchTypeInfo(&raw), expected); + EXPECT_EQ(ffi::TypeTraits::GetMismatchTypeInfo(&raw), expected); + EXPECT_EQ(ffi::TypeTraits::GetMismatchTypeInfo(&raw), expected); + }; + check(Var("x", PrimType::Float(32)), "ir.Var[ty=float32]"); + check(Var("x", AnyType()), "ir.Var[ty=ir.AnyType]"); + check(Var("x", Type::Missing()), "ir.Var[ty=ir.Type]"); + check(ffi::String("text"), "ffi.SmallStr"); + check(ffi::String("a string stored as an object"), "ffi.String"); + check(1.0, "float"); + check(nullptr, "None"); +} + TEST(ExprNodeRef, Basic) { using namespace tvm; using namespace tvm::tirx; diff --git a/tests/python/tirx-base/test_tir_op_types.py b/tests/python/tirx-base/test_tir_op_types.py index 085b0c69f619..ae23694622f9 100644 --- a/tests/python/tirx-base/test_tir_op_types.py +++ b/tests/python/tirx-base/test_tir_op_types.py @@ -36,7 +36,7 @@ def test_scalar_integer_signature(): assert tirx.tvm_struct_get(ptr, index, 2, dtype="int32").args[1].same_as(index) assert tirx.tvm_struct_get(ptr, 1 << 40, 2, dtype="int32").args[1].ty.dtype == "int64" for dtype in ("float32", "bool", "int32x4", "int32xvscalex4"): - with pytest.raises(TypeError, match=r"index.*expected `ir.IntExpr`.*dtype="): + with pytest.raises(TypeError, match=rf"index.*expected `ir.IntExpr`.*\[ty={dtype}\]"): tirx.tvm_struct_get(ptr, tirx.Var("index", dtype), 2, dtype="int32") From fd280af30242d87ccbef6e100dcbb22f5809f661 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 2 Oct 2026 21:49:38 +0000 Subject: [PATCH 2/2] [FIX][IR] Remove redundant diagnostic test --- tests/cpp/expr_test.cc | 17 ----------------- 1 file changed, 17 deletions(-) diff --git a/tests/cpp/expr_test.cc b/tests/cpp/expr_test.cc index 19de1e949252..067962a6e391 100644 --- a/tests/cpp/expr_test.cc +++ b/tests/cpp/expr_test.cc @@ -111,23 +111,6 @@ TEST(Expr, IntegerView) { EXPECT_EQ(ffi::Any(int64_t{1} << 40).cast().ty().bits(), 64); } -TEST(Expr, TypedMismatchDiagnostics) { - using namespace tvm; - auto check = [](ffi::AnyView value, const std::string& expected) { - TVMFFIAny raw = value.CopyToTVMFFIAny(); - EXPECT_EQ(ffi::TypeTraits>::GetMismatchTypeInfo(&raw), expected); - EXPECT_EQ(ffi::TypeTraits::GetMismatchTypeInfo(&raw), expected); - EXPECT_EQ(ffi::TypeTraits::GetMismatchTypeInfo(&raw), expected); - }; - check(Var("x", PrimType::Float(32)), "ir.Var[ty=float32]"); - check(Var("x", AnyType()), "ir.Var[ty=ir.AnyType]"); - check(Var("x", Type::Missing()), "ir.Var[ty=ir.Type]"); - check(ffi::String("text"), "ffi.SmallStr"); - check(ffi::String("a string stored as an object"), "ffi.String"); - check(1.0, "float"); - check(nullptr, "None"); -} - TEST(ExprNodeRef, Basic) { using namespace tvm; using namespace tvm::tirx;