28 const TypeInfo boolT = {mlir::IntegerType::getTypeID(), 1};
29 const TypeInfo i4T = {mlir::IntegerType::getTypeID(), 4};
30 const TypeInfo i8T = {mlir::IntegerType::getTypeID(), 8};
31 const TypeInfo i16T = {mlir::IntegerType::getTypeID(), 16};
32 const TypeInfo i32T = {mlir::IntegerType::getTypeID(), 32};
33 const TypeInfo i48T = {mlir::IntegerType::getTypeID(), 48};
34 const TypeInfo i64T = {mlir::IntegerType::getTypeID(), 64};
35 const TypeInfo bf16T = {mlir::BFloat16Type::getTypeID(), 16};
36 const TypeInfo fp16T = {mlir::Float16Type::getTypeID(), 16};
37 const TypeInfo fp32T = {mlir::Float32Type::getTypeID(), 32};
38 const TypeInfo fp8e4m3T = {mlir::Float8E4M3FNType::getTypeID(), 8};
39 const TypeInfo fp8e5m2T = {mlir::Float8E5M2Type::getTypeID(), 8};
44 const TypeInfo fp6e2m3T = {mlir::Float6E2M3FNType::getTypeID(), 6};
45 const TypeInfo fp6e3m2T = {mlir::Float6E3M2FNType::getTypeID(), 6};
46 const TypeInfo fp4e2m1T = {mlir::Float4E2M1FNType::getTypeID(), 4};
47 const TypeInfo fp8ue8m0T = {mlir::Float8E8M0FNUType::getTypeID(), 8};
48 const TypeInfo mxint8T = {mlir::tosa::mxint8Type::getTypeID(), 8};
51 const TypeID blockScaledID = mlir::tosa::BlockScaledType::getTypeID();
52 const TypeID fp4e2m1ID = mlir::Float4E2M1FNType::getTypeID();
53 const TypeID fp6e2m3ID = mlir::Float6E2M3FNType::getTypeID();
54 const TypeID fp6e3m2ID = mlir::Float6E3M2FNType::getTypeID();
55 const TypeID fp8e4m3ID = mlir::Float8E4M3FNType::getTypeID();
56 const TypeID fp8e5m2ID = mlir::Float8E5M2Type::getTypeID();
57 const TypeID fp8ue8m0ID = mlir::Float8E8M0FNUType::getTypeID();
58 const TypeID mxint8ID = mlir::tosa::mxint8Type::getTypeID();
60 const TypeInfo bs32_fp8ue8m0_fp4e2m1T = {blockScaledID, 4, fp4e2m1ID,
62 tosa::BlockShape::BLOCK_SHAPE_32};
63 const TypeInfo bs32_fp8ue8m0_fp6e2m3T = {blockScaledID, 6, fp6e2m3ID,
65 tosa::BlockShape::BLOCK_SHAPE_32};
66 const TypeInfo bs32_fp8ue8m0_fp6e3m2T = {blockScaledID, 6, fp6e3m2ID,
68 tosa::BlockShape::BLOCK_SHAPE_32};
69 const TypeInfo bs32_fp8ue8m0_fp8e4m3T = {blockScaledID, 8, fp8e4m3ID,
71 tosa::BlockShape::BLOCK_SHAPE_32};
72 const TypeInfo bs32_fp8ue8m0_fp8e5m2T = {blockScaledID, 8, fp8e5m2ID,
74 tosa::BlockShape::BLOCK_SHAPE_32};
75 const TypeInfo bs32_fp8ue8m0_mxint8T = {
76 blockScaledID, 8, mxint8ID, fp8ue8m0ID, tosa::BlockShape::BLOCK_SHAPE_32};
478 const auto maybeOpEntries = getOperatorMatchedEntries<T>(op);
479 if (failed(maybeOpEntries))
482 const auto opEntries = maybeOpEntries.value();
483 if (opEntries.size() == 0) {
499 const auto isVersionCompatible =
502 if (info.operandTypeInfoSet.empty())
505 info.operandTypeInfoSet.front().second};
509 for (
const auto &info : opEntries) {
513 if (isModeAllowed(info) && isVersionCompatible(info))
520 llvm::raw_string_ostream os(message);
523 const size_t numOpEntries = opEntries.size();
524 for (
const auto &[
index, info] : llvm::enumerate(opEntries)) {
525 bool mismatchedVersion =
false;
526 if (!isVersionCompatible(info)) {
527 mismatchedVersion =
true;
528 os <<
"requires specification version compatible with "
533 if (!isModeAllowed(info)) {
534 if (mismatchedVersion)
539 <<
"] profiles/extensions ";
542 if (
index != numOpEntries - 1)
545 os <<
"to be specified in the target environment";
563 const auto maybeProfEntries = getOperatorMatchedEntries<Profile>(op);
564 const auto maybeExtEntries = getOperatorMatchedEntries<Extension>(op);
565 if (failed(maybeProfEntries) && failed(maybeExtEntries))
568 const bool hasEntry =
569 (succeeded(maybeProfEntries) && !maybeProfEntries.value().empty()) ||
570 (succeeded(maybeExtEntries) && !maybeExtEntries.value().empty());
574 llvm::raw_string_ostream os(message);
575 os <<
"illegal: operation operand/result data types did not align with any "
576 "profile or extension, got (";
580 for (
const auto &typeInfo : llvm::drop_end(current))
589 const auto searchBestMatch = [&](
auto map) {
590 for (
const auto &complianceInfos : map[opName]) {
591 for (
const auto &versionedTypeInfos :
592 complianceInfos.operandTypeInfoSet) {
594 if (current.size() != typeInfos.size())
596 const int matches = llvm::count_if(
597 llvm::zip_equal(current, typeInfos), [&](
const auto zipType) {
599 std::get<1>(zipType));
601 if (matches > maxMatches) {
602 maxMatches = matches;
603 bestTypeInfo = typeInfos;
611 os <<
", did you mean (";
612 for (
const auto &typeInfo : llvm::drop_end(bestTypeInfo))
615 os <<
"Otherwise, please refer to the 'supported data types' for '"
616 << opName <<
"' in the specification.";
714 const auto stringifyScalarTypeInfo =
716 if (typeInfo.
typeID == mlir::IntegerType::getTypeID()) {
717 return {
"i" + llvm::utostr(typeInfo.
bitWidth)};
719 if (typeInfo.
typeID == mlir::Float16Type::getTypeID()) {
721 }
else if (typeInfo.
typeID == mlir::Float32Type::getTypeID()) {
723 }
else if (typeInfo.
typeID == mlir::BFloat16Type::getTypeID()) {
725 }
else if (typeInfo.
typeID == mlir::Float8E4M3FNType::getTypeID()) {
727 }
else if (typeInfo.
typeID == mlir::Float8E5M2Type::getTypeID()) {
729 }
else if (typeInfo.
typeID == mlir::Float6E2M3FNType::getTypeID()) {
731 }
else if (typeInfo.
typeID == mlir::Float6E3M2FNType::getTypeID()) {
733 }
else if (typeInfo.
typeID == mlir::Float4E2M1FNType::getTypeID()) {
735 }
else if (typeInfo.
typeID == mlir::Float8E8M0FNUType::getTypeID()) {
737 }
else if (typeInfo.
typeID == tosa::mxint8Type::getTypeID()) {
740 llvm_unreachable(
"unknown type");
743 if (typeInfo.
typeID == tosa::BlockScaledType::getTypeID()) {
747 llvm::raw_svector_ostream os(
result);
749 << tosa::BlockShapeAttr::getBlockShapeValue(typeInfo.
blockShape.value())
750 <<
"_" << stringifyScalarTypeInfo(scaleInfo) <<
"_"
751 << stringifyScalarTypeInfo(valueInfo);
755 return stringifyScalarTypeInfo(typeInfo);