17 const TypeInfo boolT = {mlir::IntegerType::getTypeID(), 1};
18 const TypeInfo i4T = {mlir::IntegerType::getTypeID(), 4};
19 const TypeInfo i8T = {mlir::IntegerType::getTypeID(), 8};
20 const TypeInfo i16T = {mlir::IntegerType::getTypeID(), 16};
21 const TypeInfo i32T = {mlir::IntegerType::getTypeID(), 32};
22 const TypeInfo i48T = {mlir::IntegerType::getTypeID(), 48};
23 const TypeInfo i64T = {mlir::IntegerType::getTypeID(), 64};
24 const TypeInfo bf16T = {mlir::BFloat16Type::getTypeID(), 16};
25 const TypeInfo fp16T = {mlir::Float16Type::getTypeID(), 16};
26 const TypeInfo fp32T = {mlir::Float32Type::getTypeID(), 32};
27 const TypeInfo fp8e4m3T = {mlir::Float8E4M3FNType::getTypeID(), 8};
28 const TypeInfo fp8e5m2T = {mlir::Float8E5M2Type::getTypeID(), 8};
33 const TypeInfo fp6e2m3T = {mlir::Float6E2M3FNType::getTypeID(), 6};
34 const TypeInfo fp6e3m2T = {mlir::Float6E3M2FNType::getTypeID(), 6};
35 const TypeInfo fp4e2m1T = {mlir::Float4E2M1FNType::getTypeID(), 4};
36 const TypeInfo fp8ue8m0T = {mlir::Float8E8M0FNUType::getTypeID(), 8};
37 const TypeInfo mxint8T = {mlir::tosa::mxint8Type::getTypeID(), 8};
40 const TypeID blockScaledID = mlir::tosa::BlockScaledType::getTypeID();
41 const TypeID fp4e2m1ID = mlir::Float4E2M1FNType::getTypeID();
42 const TypeID fp6e2m3ID = mlir::Float6E2M3FNType::getTypeID();
43 const TypeID fp6e3m2ID = mlir::Float6E3M2FNType::getTypeID();
44 const TypeID fp8e4m3ID = mlir::Float8E4M3FNType::getTypeID();
45 const TypeID fp8e5m2ID = mlir::Float8E5M2Type::getTypeID();
46 const TypeID fp8ue8m0ID = mlir::Float8E8M0FNUType::getTypeID();
47 const TypeID mxint8ID = mlir::tosa::mxint8Type::getTypeID();
49 const TypeInfo bs32_fp8ue8m0_fp4e2m1T = {blockScaledID, 4, fp4e2m1ID,
51 tosa::BlockShape::BLOCK_SHAPE_32};
52 const TypeInfo bs32_fp8ue8m0_fp6e2m3T = {blockScaledID, 6, fp6e2m3ID,
54 tosa::BlockShape::BLOCK_SHAPE_32};
55 const TypeInfo bs32_fp8ue8m0_fp6e3m2T = {blockScaledID, 6, fp6e3m2ID,
57 tosa::BlockShape::BLOCK_SHAPE_32};
58 const TypeInfo bs32_fp8ue8m0_fp8e4m3T = {blockScaledID, 8, fp8e4m3ID,
60 tosa::BlockShape::BLOCK_SHAPE_32};
61 const TypeInfo bs32_fp8ue8m0_fp8e5m2T = {blockScaledID, 8, fp8e5m2ID,
63 tosa::BlockShape::BLOCK_SHAPE_32};
64 const TypeInfo bs32_fp8ue8m0_mxint8T = {
65 blockScaledID, 8, mxint8ID, fp8ue8m0ID, tosa::BlockShape::BLOCK_SHAPE_32};
466 const auto maybeOpEntries = getOperatorMatchedEntries<T>(op);
467 if (failed(maybeOpEntries))
470 const auto opEntries = maybeOpEntries.value();
471 if (opEntries.size() == 0) {
487 const auto isVersionCompatible =
490 if (info.operandTypeInfoSet.empty())
493 info.operandTypeInfoSet.front().second};
497 for (
const auto &info : opEntries) {
501 if (isModeAllowed(info) && isVersionCompatible(info))
508 llvm::raw_string_ostream os(message);
511 const size_t numOpEntries = opEntries.size();
512 for (
const auto &[
index, info] : llvm::enumerate(opEntries)) {
513 bool mismatchedVersion =
false;
514 if (!isVersionCompatible(info)) {
515 mismatchedVersion =
true;
516 os <<
"requires specification version compatible with "
521 if (!isModeAllowed(info)) {
522 if (mismatchedVersion)
527 <<
"] profiles/extensions ";
530 if (
index != numOpEntries - 1)
533 os <<
"to be specified in the target environment";
551 const auto maybeProfEntries = getOperatorMatchedEntries<Profile>(op);
552 const auto maybeExtEntries = getOperatorMatchedEntries<Extension>(op);
553 if (failed(maybeProfEntries) && failed(maybeExtEntries))
556 const bool hasEntry =
557 (succeeded(maybeProfEntries) && !maybeProfEntries.value().empty()) ||
558 (succeeded(maybeExtEntries) && !maybeExtEntries.value().empty());
562 llvm::raw_string_ostream os(message);
563 os <<
"illegal: operation operand/result data types did not align with any "
564 "profile or extension, got (";
568 for (
const auto &typeInfo : llvm::drop_end(current))
577 const auto searchBestMatch = [&](
auto map) {
578 for (
const auto &complianceInfos : map[opName]) {
579 for (
const auto &versionedTypeInfos :
580 complianceInfos.operandTypeInfoSet) {
582 if (current.size() != typeInfos.size())
584 const int matches = llvm::count_if(
585 llvm::zip_equal(current, typeInfos), [&](
const auto zipType) {
587 std::get<1>(zipType));
589 if (matches > maxMatches) {
590 maxMatches = matches;
591 bestTypeInfo = typeInfos;
599 os <<
", did you mean (";
600 for (
const auto &typeInfo : llvm::drop_end(bestTypeInfo))
603 os <<
"Otherwise, please refer to the 'supported data types' for '"
604 << opName <<
"' in the specification.";
702 const auto stringifyScalarTypeInfo =
704 if (typeInfo.
typeID == mlir::IntegerType::getTypeID()) {
705 return {
"i" + llvm::utostr(typeInfo.
bitWidth)};
707 if (typeInfo.
typeID == mlir::Float16Type::getTypeID()) {
709 }
else if (typeInfo.
typeID == mlir::Float32Type::getTypeID()) {
711 }
else if (typeInfo.
typeID == mlir::BFloat16Type::getTypeID()) {
713 }
else if (typeInfo.
typeID == mlir::Float8E4M3FNType::getTypeID()) {
715 }
else if (typeInfo.
typeID == mlir::Float8E5M2Type::getTypeID()) {
717 }
else if (typeInfo.
typeID == mlir::Float6E2M3FNType::getTypeID()) {
719 }
else if (typeInfo.
typeID == mlir::Float6E3M2FNType::getTypeID()) {
721 }
else if (typeInfo.
typeID == mlir::Float4E2M1FNType::getTypeID()) {
723 }
else if (typeInfo.
typeID == mlir::Float8E8M0FNUType::getTypeID()) {
725 }
else if (typeInfo.
typeID == tosa::mxint8Type::getTypeID()) {
728 llvm_unreachable(
"unknown type");
731 if (typeInfo.
typeID == tosa::BlockScaledType::getTypeID()) {
735 llvm::raw_svector_ostream os(
result);
737 << tosa::BlockShapeAttr::getBlockShapeValue(typeInfo.
blockShape.value())
738 <<
"_" << stringifyScalarTypeInfo(scaleInfo) <<
"_"
739 << stringifyScalarTypeInfo(valueInfo);
743 return stringifyScalarTypeInfo(typeInfo);