18 const TypeInfo boolT = {mlir::IntegerType::getTypeID(), 1};
19 const TypeInfo i4T = {mlir::IntegerType::getTypeID(), 4};
20 const TypeInfo i8T = {mlir::IntegerType::getTypeID(), 8};
21 const TypeInfo i16T = {mlir::IntegerType::getTypeID(), 16};
22 const TypeInfo i32T = {mlir::IntegerType::getTypeID(), 32};
23 const TypeInfo i48T = {mlir::IntegerType::getTypeID(), 48};
24 const TypeInfo i64T = {mlir::IntegerType::getTypeID(), 64};
25 const TypeInfo bf16T = {mlir::BFloat16Type::getTypeID(), 16};
26 const TypeInfo fp16T = {mlir::Float16Type::getTypeID(), 16};
27 const TypeInfo fp32T = {mlir::Float32Type::getTypeID(), 32};
28 const TypeInfo fp8e4m3T = {mlir::Float8E4M3FNType::getTypeID(), 8};
29 const TypeInfo fp8e5m2T = {mlir::Float8E5M2Type::getTypeID(), 8};
34 const TypeInfo fp6e2m3T = {mlir::Float6E2M3FNType::getTypeID(), 6};
35 const TypeInfo fp6e3m2T = {mlir::Float6E3M2FNType::getTypeID(), 6};
36 const TypeInfo fp4e2m1T = {mlir::Float4E2M1FNType::getTypeID(), 4};
37 const TypeInfo fp8ue8m0T = {mlir::Float8E8M0FNUType::getTypeID(), 8};
38 const TypeInfo mxint8T = {mlir::tosa::mxint8Type::getTypeID(), 8};
41 const TypeID blockScaledID = mlir::tosa::BlockScaledType::getTypeID();
42 const TypeID fp4e2m1ID = mlir::Float4E2M1FNType::getTypeID();
43 const TypeID fp6e2m3ID = mlir::Float6E2M3FNType::getTypeID();
44 const TypeID fp6e3m2ID = mlir::Float6E3M2FNType::getTypeID();
45 const TypeID fp8e4m3ID = mlir::Float8E4M3FNType::getTypeID();
46 const TypeID fp8e5m2ID = mlir::Float8E5M2Type::getTypeID();
47 const TypeID fp8ue8m0ID = mlir::Float8E8M0FNUType::getTypeID();
48 const TypeID mxint8ID = mlir::tosa::mxint8Type::getTypeID();
50 const TypeInfo bs32_fp8ue8m0_fp4e2m1T = {blockScaledID, 4, fp4e2m1ID,
52 tosa::BlockShape::BLOCK_SHAPE_32};
53 const TypeInfo bs32_fp8ue8m0_fp6e2m3T = {blockScaledID, 6, fp6e2m3ID,
55 tosa::BlockShape::BLOCK_SHAPE_32};
56 const TypeInfo bs32_fp8ue8m0_fp6e3m2T = {blockScaledID, 6, fp6e3m2ID,
58 tosa::BlockShape::BLOCK_SHAPE_32};
59 const TypeInfo bs32_fp8ue8m0_fp8e4m3T = {blockScaledID, 8, fp8e4m3ID,
61 tosa::BlockShape::BLOCK_SHAPE_32};
62 const TypeInfo bs32_fp8ue8m0_fp8e5m2T = {blockScaledID, 8, fp8e5m2ID,
64 tosa::BlockShape::BLOCK_SHAPE_32};
65 const TypeInfo bs32_fp8ue8m0_mxint8T = {
66 blockScaledID, 8, mxint8ID, fp8ue8m0ID, tosa::BlockShape::BLOCK_SHAPE_32};
467 const auto maybeOpEntries = getOperatorMatchedEntries<T>(op);
468 if (failed(maybeOpEntries))
471 const auto opEntries = maybeOpEntries.value();
472 if (opEntries.size() == 0) {
488 const auto isVersionCompatible =
491 if (info.operandTypeInfoSet.empty())
494 info.operandTypeInfoSet.front().second};
498 for (
const auto &info : opEntries) {
502 if (isModeAllowed(info) && isVersionCompatible(info))
509 llvm::raw_string_ostream os(message);
512 const size_t numOpEntries = opEntries.size();
513 for (
const auto &[
index, info] : llvm::enumerate(opEntries)) {
514 bool mismatchedVersion =
false;
515 if (!isVersionCompatible(info)) {
516 mismatchedVersion =
true;
517 os <<
"requires specification version compatible with "
522 if (!isModeAllowed(info)) {
523 if (mismatchedVersion)
528 <<
"] profiles/extensions ";
531 if (
index != numOpEntries - 1)
534 os <<
"to be specified in the target environment";
552 const auto maybeProfEntries = getOperatorMatchedEntries<Profile>(op);
553 const auto maybeExtEntries = getOperatorMatchedEntries<Extension>(op);
554 if (failed(maybeProfEntries) && failed(maybeExtEntries))
557 const bool hasEntry =
558 (succeeded(maybeProfEntries) && !maybeProfEntries.value().empty()) ||
559 (succeeded(maybeExtEntries) && !maybeExtEntries.value().empty());
563 llvm::raw_string_ostream os(message);
564 os <<
"illegal: operation operand/result data types did not align with any "
565 "profile or extension, got (";
569 for (
const auto &typeInfo : llvm::drop_end(current))
578 const auto searchBestMatch = [&](
auto map) {
579 for (
const auto &complianceInfos : map[opName]) {
580 for (
const auto &versionedTypeInfos :
581 complianceInfos.operandTypeInfoSet) {
583 if (current.size() != typeInfos.size())
585 const int matches = llvm::count_if(
586 llvm::zip_equal(current, typeInfos), [&](
const auto zipType) {
588 std::get<1>(zipType));
590 if (matches > maxMatches) {
591 maxMatches = matches;
592 bestTypeInfo = typeInfos;
600 os <<
", did you mean (";
601 for (
const auto &typeInfo : llvm::drop_end(bestTypeInfo))
604 os <<
"Otherwise, please refer to the 'supported data types' for '"
605 << opName <<
"' in the specification.";
703 const auto stringifyScalarTypeInfo =
705 if (typeInfo.
typeID == mlir::IntegerType::getTypeID()) {
706 return {
"i" + llvm::utostr(typeInfo.
bitWidth)};
708 if (typeInfo.
typeID == mlir::Float16Type::getTypeID()) {
710 }
else if (typeInfo.
typeID == mlir::Float32Type::getTypeID()) {
712 }
else if (typeInfo.
typeID == mlir::BFloat16Type::getTypeID()) {
714 }
else if (typeInfo.
typeID == mlir::Float8E4M3FNType::getTypeID()) {
716 }
else if (typeInfo.
typeID == mlir::Float8E5M2Type::getTypeID()) {
718 }
else if (typeInfo.
typeID == mlir::Float6E2M3FNType::getTypeID()) {
720 }
else if (typeInfo.
typeID == mlir::Float6E3M2FNType::getTypeID()) {
722 }
else if (typeInfo.
typeID == mlir::Float4E2M1FNType::getTypeID()) {
724 }
else if (typeInfo.
typeID == mlir::Float8E8M0FNUType::getTypeID()) {
726 }
else if (typeInfo.
typeID == tosa::mxint8Type::getTypeID()) {
729 llvm_unreachable(
"unknown type");
732 if (typeInfo.
typeID == tosa::BlockScaledType::getTypeID()) {
736 llvm::raw_svector_ostream os(
result);
738 << tosa::BlockShapeAttr::getBlockShapeValue(typeInfo.
blockShape.value())
739 <<
"_" << stringifyScalarTypeInfo(scaleInfo) <<
"_"
740 << stringifyScalarTypeInfo(valueInfo);
744 return stringifyScalarTypeInfo(typeInfo);