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/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")