30#include "llvm/ADT/APFloat.h"
31#include "llvm/ADT/SmallVectorExtras.h"
32#include "llvm/ADT/TypeSwitch.h"
40#include "mlir/Dialect/Tosa/IR/TosaOpsDialect.cpp.inc"
47#include "mlir/Dialect/Tosa/IR/TosaEnums.cpp.inc"
48#include "mlir/Dialect/Tosa/IR/TosaInterfaces.cpp.inc"
51#include "mlir/Dialect/Tosa/IR/TosaDialectBytecode.cpp.inc"
56struct TosaInlinerInterface :
public DialectInlinerInterface {
57 using DialectInlinerInterface::DialectInlinerInterface;
65 IRMapping &map)
const final {
71 IRMapping &map)
const final {
72 return (isa<tosa::IfOp>(dest->getParentOp()) ||
73 isa<tosa::WhileOp>(dest->getParentOp()));
78struct TosaDialectBytecodeInterface :
public BytecodeDialectInterface {
79 TosaDialectBytecodeInterface(Dialect *dialect)
80 : BytecodeDialectInterface(dialect) {}
85 Attribute readAttribute(DialectBytecodeReader &reader)
const override {
89 LogicalResult writeAttribute(Attribute attr,
90 DialectBytecodeWriter &writer)
const override {
91 return ::writeAttribute(attr, writer);
97 Type readType(DialectBytecodeReader &reader)
const override {
101 LogicalResult writeType(Type type,
102 DialectBytecodeWriter &writer)
const override {
103 return ::writeType(type, writer);
106 void writeVersion(DialectBytecodeWriter &writer)
const final {
110 std::unique_ptr<DialectVersion>
111 readVersion(DialectBytecodeReader &reader)
const final {
113 reader.
emitError(
"Dialect does not support versioning");
117 LogicalResult upgradeFromVersion(Operation *topLevelOp,
118 const DialectVersion &version)
const final {
131 return {&getBodyGraph()};
140 return dim == -1 ? ShapedType::kDynamic : dim;
146 Type elementType = variableOp.getType();
149 return RankedTensorType::get(
shape, elementType);
156void TosaDialect::initialize() {
158#define GET_TYPEDEF_LIST
159#include "mlir/Dialect/Tosa/IR/TosaOpsTypesBase.cpp.inc"
163#include "mlir/Dialect/Tosa/IR/TosaOps.cpp.inc"
166#define GET_ATTRDEF_LIST
167#include "mlir/Dialect/Tosa/IR/TosaAttributes.cpp.inc"
169 addInterfaces<TosaDialectBytecodeInterface, TosaInlinerInterface>();
170 declarePromisedInterfaces<
171 shard::ShardingInterface, ClampOp, SigmoidOp, TanhOp, AddOp,
172 ArithmeticRightShiftOp, BitwiseAndOp, BitwiseOrOp, BitwiseXorOp, IntDivOp,
173 LogicalAndOp, LogicalLeftShiftOp, LogicalRightShiftOp, LogicalOrOp,
174 LogicalXorOp, MaximumOp, MinimumOp, MulOp, PowOp, SubOp, AbsOp,
175 BitwiseNotOp, CeilOp, ClzOp, ExpOp, FloorOp, LogOp, LogicalNotOp,
176 NegateOp, ReciprocalOp, RsqrtOp, SelectOp, EqualOp, GreaterOp,
177 GreaterEqualOp, MatMulOp>();
184 if (llvm::isa<shapeType>(type) && llvm::isa<DenseIntElementsAttr>(value)) {
185 return tosa::ConstShapeOp::create(builder, loc, type,
186 llvm::cast<DenseIntElementsAttr>(value));
188 if (llvm::isa<ElementsAttr>(value))
189 return tosa::ConstOp::create(builder, loc, type,
190 llvm::cast<ElementsAttr>(value));
200ParseResult getShapeAndElementType(
OpAsmParser &parser,
Type parsedType,
202 TypeAttr &typeAttr) {
203 if (
auto shapedType = dyn_cast<ShapedType>(parsedType)) {
204 if (!shapedType.hasRank())
206 <<
"expected ranked type";
208 auto elementType = shapedType.getElementType();
209 typeAttr = TypeAttr::get(elementType);
216 <<
"expected shaped type";
233 <<
"expected attribute";
235 if (
auto typedAttr = dyn_cast<TypedAttr>(initialValueAttr)) {
236 return getShapeAndElementType(parser, typedAttr.getType(), varShapeAttr,
240 <<
"expected Typed attr";
243 initialValueAttr =
nullptr;
247 <<
"expected type after colon";
249 return getShapeAndElementType(parser, parsedType, varShapeAttr, typeAttr);
254 TypeAttr typeAttr,
Attribute initialValueAttr) {
255 bool needsSpace =
false;
256 if (!dyn_cast_or_null<TypedAttr>(initialValueAttr)) {
259 Type elementType = typeAttr.getValue();
260 RankedTensorType tensorType =
262 auto tensorTypeAttr = TypeAttr::get(tensorType);
267 if (initialValueAttr) {
287 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(srcType))
297 Value valZp, StringRef name) {
302 mlir::isa<IntegerType>(eType) && mlir::isa<IntegerType>(eZpType);
306 if (!bothInts || !sameBitWidth) {
308 <<
"expected " << name <<
" and " << name
309 <<
"_zp to both be integer of the same bitwidth, but got " << eType
310 <<
" vs. " << eZpType;
317 Value src, int32_t val) {
320 const auto padConstType = mlir::RankedTensorType::get({1}, srcType);
321 const auto padConstEType = mlir::RankedTensorType::get({1}, srcElemType);
322 const auto padConstAttr{
323 llvm::isa<FloatType>(srcElemType)
328 return tosa::ConstOp::create(builder, loc, padConstType, padConstAttr);
332 if (
auto blockScaledTy = dyn_cast<tosa::BlockScaledType>(type))
334 if (dyn_cast<tosa::mxint8Type>(type))
343 const StringRef operandName,
344 const StringRef dimName) {
345 if (ShapedType::isDynamic(currDim)) {
348 }
else if (ShapedType::isStatic(newDim) && currDim != newDim) {
350 << dimName <<
" of " << operandName <<
" to match size " << currDim
351 <<
", got " << newDim;
358 auto printDim = [&](
int64_t dim) {
359 if (ShapedType::isDynamic(dim))
365 llvm::interleaveComma(
shape,
diag, printDim);
371 StringRef outputName =
"output") {
372 assert(outputType.hasRank() &&
"expected output type to be ranked");
378 diag << outputName <<
" shape ";
380 diag <<
" to be compatible with inferred shape ";
388 const int64_t stride,
const int64_t dilation,
const llvm::StringRef dimName,
389 const llvm::StringRef dimAxis,
const llvm::StringRef padBeforeName,
390 const llvm::StringRef padAfterName) {
391 if (inputSize == ShapedType::kDynamic || kernelSize == ShapedType::kDynamic)
396 const std::optional<int64_t> calculatedOutSizeMinusOne =
idivCheck(
397 inputSize - 1 + padBefore + padAfter - (kernelSize - 1) * dilation,
399 if (!calculatedOutSizeMinusOne.has_value())
401 << dimName <<
" - 1 + pad_" << padBeforeName <<
" + pad_"
402 << padAfterName <<
" - (kernel_" << dimName <<
" - 1) * dilation_"
403 << dimAxis <<
" to be wholly divisible by stride_" << dimAxis
404 <<
", got (" << inputSize <<
" - 1 + " << padBefore <<
" + "
405 << padAfter <<
" - (" << kernelSize <<
" - 1) * " << dilation
408 const int64_t calculatedOutSize = calculatedOutSizeMinusOne.value() + 1;
409 if (outputSize != ShapedType::kDynamic && calculatedOutSize != outputSize)
411 << dimName <<
" did not match expected: "
412 <<
"calculated=" << calculatedOutSize <<
", expected=" << outputSize;
420size_t mlir::tosa::mxint8Type::getDenseElementBitSize()
const {
return 8; }
423mlir::tosa::mxint8Type::convertToAttribute(
ArrayRef<char> rawData)
const {
424 assert(rawData.size() == 1 &&
"expected 1 byte for tosa.mxint8 element");
425 const auto intType = IntegerType::get(
getContext(), 8);
426 return intType.convertToAttribute(rawData);
429LogicalResult mlir::tosa::mxint8Type::convertFromAttribute(
431 const auto intAttr = dyn_cast<IntegerAttr>(attr);
434 const Type attrType = intAttr.getType();
437 return cast<IntegerType>(attrType).convertFromAttribute(attr,
result);
446 bool allowScaleValues) {
447 const auto tensorType = llvm::cast<ShapedType>(type);
448 const BlockScaledType elemType =
449 llvm::dyn_cast<BlockScaledType>(tensorType.getElementType());
453 if (!allowScaleValues && elemType.hasScaleValues()) {
456 <<
"block scaled tensor type with scale values is not allowed";
460 if (!tensorType.hasRank())
463 if (tensorType.getRank() == 0) {
465 emitError() <<
"block scaled tensor type must have rank greater than "
471 const uint32_t blockSize =
472 BlockShapeAttr::getBlockShapeValue(elemType.getBlockShape());
474 if (allowScaleValues && elemType.hasScaleValues() &&
475 tensorType.hasStaticShape()) {
476 const size_t numBlocks = tensorType.getNumElements() / blockSize;
477 if (elemType.getScaleValues().size() != numBlocks) {
479 emitError() <<
"block scaled tensor type with scale values must have "
480 "scale values for each block, expected "
481 << numBlocks <<
", got "
482 << elemType.getScaleValues().size();
487 const int64_t blockedDimension = tensorShape.back();
488 if (ShapedType::isDynamic(blockedDimension))
491 if (blockedDimension % blockSize != 0) {
493 emitError() <<
"last dimension of block scaled tensor type ("
494 << blockedDimension <<
") must be divisible by block size ("
510 type, [ctx] {
return emitError(UnknownLoc::get(ctx)); })) &&
512 return ": " + message;
521 const auto parseScaleValue = [&]() -> ParseResult {
528 if (floatValue < 0.0)
529 return parser.
emitError(loc,
"scale value must be non-negative, got ")
532 Type attrType = scaleType;
536 if (attrType != scaleType)
537 return parser.
emitError(loc,
"parsed attribute type ")
538 << attrType <<
" does not match expected scale type " << scaleType;
540 scaleValues.push_back(FloatAttr::get(attrType, floatValue));
549 llvm::interleaveComma(scaleValues, printer, [&](
Attribute scaleValue) {
554size_t mlir::tosa::BlockScaledType::getDenseElementBitSize()
const {
556 if (isa<tosa::mxint8Type>(valueType))
562mlir::tosa::BlockScaledType::convertToAttribute(
ArrayRef<char> rawData)
const {
566 assert(rawData.size() == 1 &&
"expected 1 byte for block_scaled element");
568 if (
const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
569 return mxint8Value.convertToAttribute(rawData);
570 if (!isa<FloatType>(valueType))
575LogicalResult mlir::tosa::BlockScaledType::convertFromAttribute(
578 if (
const auto mxint8Value = dyn_cast<tosa::mxint8Type>(valueType))
579 return mxint8Value.convertFromAttribute(attr,
result);
581 const auto floatAttr = dyn_cast<FloatAttr>(attr);
582 if (!floatAttr || floatAttr.getType() != valueType)
593template <
typename A, std::enable_if_t<std::is_same_v<A, ArgMaxOp::Adaptor> ||
594 std::is_same_v<A, ArgMinOp::Adaptor>,
597 MLIRContext *context, ::std::optional<Location> location, A adaptor,
600 IntegerAttr axis = adaptor.getProperties().axis;
601 int32_t axisVal = axis.getValue().getSExtValue();
608 const auto inputRank = inputShape.
getRank();
610 outShape.reserve(inputRank - 1);
611 for (
int i = 0, s = inputRank; i < s; i++) {
626 const ShapedType resultType = llvm::cast<ShapedType>(op.getType());
628 if (
const auto resultETy = resultType.getElementType();
629 !resultETy.isIntOrIndex())
630 return op.emitOpError(
"result tensor is not of integer type");
632 const auto inputType = llvm::cast<ShapedType>(op.getInput().getType());
633 if (!inputType.hasRank())
637 const int64_t axis = op.getAxisAttr().getInt();
638 if (((axis < 0) || axis >= inputType.getRank()))
639 return op.emitOpError(
"specified axis is outside the rank of the tensor");
641 if (!resultType.hasRank())
647 expectedOutputShape.erase(expectedOutputShape.begin() + axis);
649 return op.emitOpError(
"expected output shape '")
650 << expectedOutputShape <<
"', got '" << outputShape <<
"'";
657 const auto inputType = llvm::dyn_cast<TensorType>(op.getInput().getType());
658 const auto weightType = llvm::dyn_cast<TensorType>(op.getWeight().getType());
660 auto inputEType = inputType.getElementType();
661 auto weightEType = weightType.getElementType();
663 llvm::cast<ShapedType>(op.getBias().getType()).getElementType();
665 llvm::cast<ShapedType>(op.getResult().getType()).getElementType();
666 bool biasIsFloat = llvm::isa<FloatType>(biasEType);
667 bool resultIsFloat = llvm::isa<FloatType>(resultEType);
669 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(inputEType))
672 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(weightEType))
675 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(biasEType))
678 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(resultEType))
681 if (biasIsFloat && resultIsFloat && (biasEType != resultEType)) {
685 "expect both bias and result to have same element type, got ")
686 << biasEType <<
" and " << resultEType;
690 const bool isInputBlockScaled = llvm::isa<BlockScaledType>(inputEType);
691 const bool isWeightBlockScaled = llvm::isa<BlockScaledType>(weightEType);
692 const bool isInputFloat = llvm::isa<FloatType>(inputEType);
693 const bool isWeightFloat = llvm::isa<FloatType>(weightEType);
695 const bool isInputBSorFloat = isInputBlockScaled || isInputFloat;
696 const bool isWeightBSorFloat = isWeightBlockScaled || isWeightFloat;
699 if (isInputBSorFloat != isWeightBSorFloat) {
701 "expect both input and weight to be float or not together, got ")
702 << inputEType <<
" and " << weightEType;
707 if (!isInputBlockScaled && inputEType != inputZpEType) {
708 return op.emitOpError(
"expect both input and its zero point are the same "
709 "element type, got ")
710 << inputEType <<
" and " << inputZpEType;
712 if (isInputBlockScaled && !llvm::isa<Float32Type>(inputZpEType)) {
713 return op.emitOpError(
714 "expect block scaled input to have fp32 zero point, got ")
715 << inputEType <<
" and " << inputZpEType;
719 if (!isWeightBlockScaled && weightEType != weightZpEType) {
720 return op.emitOpError(
"expect both weight and its zero point are the same "
721 "element type, got ")
722 << weightEType <<
" and " << weightZpEType;
724 if (isWeightBlockScaled && !llvm::isa<Float32Type>(weightZpEType)) {
725 return op.emitOpError(
726 "expect block scaled weight to have fp32 zero point, got ")
727 << weightEType <<
" and " << weightZpEType;
730 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
731 if (succeeded(maybeIZp) && op.verifyInputZeroPoint(*maybeIZp).failed())
734 FailureOr<int64_t> maybeWZp = op.getWeightZeroPoint();
735 if (succeeded(maybeWZp) && op.verifyWeightZeroPoint(*maybeWZp).failed())
741LogicalResult tosa::ConstOp::verify() {
743 auto attrType = llvm::dyn_cast<TensorType>(getValuesAttr().
getType());
744 auto outputType = llvm::dyn_cast<TensorType>(getOutput().
getType());
746 if (!attrType || !outputType) {
747 emitOpError(
"expected tensors for attr/result type");
751 const Type attrElemType = attrType.getElementType();
752 const Type resultElemType = outputType.getElementType();
755 llvm::dyn_cast<mlir::quant::QuantizedType>(resultElemType)) {
760 if (
auto attrBlockScaledType =
761 llvm::dyn_cast<mlir::tosa::BlockScaledType>(attrElemType)) {
762 if (!attrBlockScaledType.hasScaleValues())
764 "attribute block scaled type must have scale values");
766 const auto emitAttributeError = [&op]() {
767 return op.
emitOpError(
"attribute block scaled type is invalid: ");
773 const BlockScaledType resultBlockScaledType =
774 llvm::dyn_cast<mlir::tosa::BlockScaledType>(resultElemType);
775 if (!resultBlockScaledType)
777 "result type must be block scaled type if attribute is block "
780 if (attrBlockScaledType.getValueType() !=
781 resultBlockScaledType.getValueType() ||
782 attrBlockScaledType.getScaleType() !=
783 resultBlockScaledType.getScaleType() ||
784 attrBlockScaledType.getBlockShape() !=
785 resultBlockScaledType.getBlockShape())
787 "expected block scaled element type to be compatible "
788 "between attr and result, got ")
789 << attrBlockScaledType <<
" vs. " << resultBlockScaledType;
794 if (attrElemType != resultElemType)
795 return emitOpError(
"expected same attr/result element types");
803 llvm::cast<ShapedType>(op.getInput().getType()).getElementType();
805 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(inputEType))
809 llvm::cast<ShapedType>(op.getResult().getType()).getElementType();
811 if (
auto quantType = llvm::dyn_cast<mlir::quant::QuantizedType>(resultEType))
825 if (llvm::any_of(padding, [](
int64_t p) {
return p < 0; }))
826 return op.emitOpError(
"expect all padding values to be >= 0, got ")
830 if (llvm::any_of(strides, [](
int64_t s) {
return s < 1; }))
831 return op.emitOpError(
"expect all stride values to be >= 1, got ")
835 if (llvm::any_of(dilations, [](
int64_t d) {
return d < 1; }))
836 return op.emitOpError(
"expect all dilation values to be >= 1, got ")
839 const RankedTensorType outputType =
840 llvm::dyn_cast<RankedTensorType>(op.getOutput().getType());
845 const RankedTensorType inputType =
846 llvm::dyn_cast<RankedTensorType>(op.getInput().getType());
847 const RankedTensorType weightType =
848 llvm::dyn_cast<RankedTensorType>(op.getWeight().getType());
850 if (inputType && weightType) {
852 if constexpr (std::is_same<T, tosa::Conv2DOp>::value) {
854 op, inputType.getDimSize(1), weightType.getDimSize(1),
855 outputType.getDimSize(1), padding[0], padding[1], strides[0],
856 dilations[0],
"height",
"y",
"top",
"bottom")))
860 op, inputType.getDimSize(2), weightType.getDimSize(2),
861 outputType.getDimSize(2), padding[2], padding[3], strides[1],
862 dilations[1],
"width",
"x",
"left",
"right")))
867 if constexpr (std::is_same<T, tosa::DepthwiseConv2DOp>::value) {
869 op, inputType.getDimSize(1), weightType.getDimSize(0),
870 outputType.getDimSize(1), padding[0], padding[1], strides[0],
871 dilations[0],
"height",
"y",
"top",
"bottom")))
875 op, inputType.getDimSize(2), weightType.getDimSize(1),
876 outputType.getDimSize(2), padding[2], padding[3], strides[1],
877 dilations[1],
"width",
"x",
"left",
"right")))
882 if constexpr (std::is_same<T, tosa::Conv3DOp>::value) {
884 op, inputType.getDimSize(1), weightType.getDimSize(1),
885 outputType.getDimSize(1), padding[0], padding[1], strides[0],
886 dilations[0],
"depth",
"d",
"front",
"back")))
890 op, inputType.getDimSize(2), weightType.getDimSize(2),
891 outputType.getDimSize(2), padding[2], padding[3], strides[1],
892 dilations[1],
"height",
"y",
"top",
"bottom")))
896 op, inputType.getDimSize(3), weightType.getDimSize(3),
897 outputType.getDimSize(3), padding[4], padding[5], strides[2],
898 dilations[2],
"width",
"x",
"left",
"right")))
903 const RankedTensorType biasType =
904 llvm::dyn_cast<RankedTensorType>(op.getBias().getType());
909 const int64_t biasChannels = biasType.getDimSize(0);
911 outputType.getDimSize(outputType.getRank() - 1);
912 if (biasChannels == ShapedType::kDynamic ||
913 outputChannels == ShapedType::kDynamic)
917 if (biasChannels != outputChannels && biasChannels != 1)
918 return op.emitOpError(
919 "bias channels expected to be equal to output channels (")
920 << outputChannels <<
") or 1, got " << biasChannels;
927 StringRef name1,
Type type2,
929 auto shapeType1 = dyn_cast<ShapedType>(type1);
930 auto shapeType2 = dyn_cast<ShapedType>(type2);
931 if (!shapeType1 || !shapeType2)
934 auto elemType1 = shapeType1.getElementType();
935 auto elemType2 = shapeType2.getElementType();
936 if (elemType1 != elemType2)
938 <<
"require same element type for " << name1 <<
" (" << elemType1
939 <<
") and " << name2 <<
" (" << elemType2 <<
")";
943 <<
"require same shapes for " << name1 <<
" (" << type1 <<
") and "
944 << name2 <<
" (" << type2 <<
")";
954 if (list1.size() != list2.size())
956 <<
"require same number of values in " << name1 <<
" ("
957 << list1.size() <<
") and " << name2 <<
" (" << list2.size() <<
")";
959 for (
auto [type1, type2] :
979 op->template getParentWithTrait<OpTrait::SymbolTable>();
986 const auto varOp = symTable.
lookup<tosa::VariableOp>(op.getName());
990 return op->emitOpError(
"'")
991 << op.getName() <<
"' has not been declared by 'tosa.variable'";
1005 StringRef aName =
"input",
1006 StringRef bName =
"output") {
1007 auto aTType = llvm::dyn_cast<TensorType>(aType);
1008 auto bTType = llvm::dyn_cast<TensorType>(bType);
1010 op->
emitOpError(
"expect shaped tensor for") << aName <<
", got " << aType;
1014 op->
emitOpError(
"expect shaped tensor for") << bName <<
", got" << bType;
1017 auto aElementType = aTType.getElementType();
1018 auto bElementType = bTType.getElementType();
1020 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(aElementType);
1022 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(bElementType);
1023 if ((aElementType.isIntOrIndexOrFloat() || aQuantType) &&
1024 (bElementType.isIntOrIndexOrFloat() || bQuantType) &&
1025 aElementType != bElementType) {
1031 << aName <<
" and " << bName <<
" to have same element type, got "
1032 << aElementType <<
" and " << bElementType;
1038LogicalResult tosa::ArgMaxOp::verify() {
return argMaxMinVerify(*
this); }
1040LogicalResult tosa::ArgMinOp::verify() {
return argMaxMinVerify(*
this); }
1050 const bool hasKernel = kernel.size() > 0;
1051 const bool hasStrides = strides.size() > 0;
1052 const bool hasPad = padding.size() > 0;
1054 if (hasKernel && llvm::any_of(kernel, [](
int64_t s) {
return s < 1; }))
1055 return op->
emitOpError(
"expect all kernel values to be >= 1, got ")
1058 if (hasStrides && llvm::any_of(strides, [](
int64_t s) {
return s < 1; }))
1059 return op->
emitOpError(
"expect all stride values to be >= 1, got ")
1062 if (hasPad && llvm::any_of(padding, [](
int64_t p) {
return p < 0; }))
1063 return op->
emitOpError(
"expect all padding values to be >= 0, got ")
1066 if (hasKernel && hasPad) {
1068 const int64_t kernelX = kernel[1];
1069 const int64_t padLeft = padding[2];
1070 const int64_t padRight = padding[3];
1071 if (padRight >= kernelX || padLeft >= kernelX)
1072 return op->
emitOpError(
"expected left/right padding to be less than the "
1073 "width of the kernel, got pad_left=")
1074 << padLeft <<
", pad_right=" << padRight
1075 <<
", kernel_x=" << kernelX;
1077 const int64_t kernelY = kernel[0];
1078 const int64_t padTop = padding[0];
1079 const int64_t padBottom = padding[1];
1080 if (padTop >= kernelY || padBottom >= kernelY)
1081 return op->
emitOpError(
"expected top/bottom padding to be less than the "
1082 "height of the kernel, got pad_top=")
1083 << padTop <<
", pad_bottom=" << padBottom
1084 <<
", kernel_y=" << kernelY;
1087 const auto inputType = llvm::dyn_cast<RankedTensorType>(input.
getType());
1088 const auto outputType = llvm::dyn_cast<RankedTensorType>(output.
getType());
1089 if (!inputType || !outputType)
1092 if (hasKernel && hasStrides && hasPad) {
1093 const auto verifyOutputSize =
1097 const llvm::StringRef dimName,
const llvm::StringRef dimAxis,
1098 const llvm::StringRef padBeforeName,
1099 const llvm::StringRef padAfterName) -> LogicalResult {
1100 if (ShapedType::isDynamic(inputSize))
1103 const std::optional<int64_t> calculatedOutSizeMinusOne =
1104 idivCheck(inputSize + padBefore + padAfter - kernelSize, strideSize);
1105 if (!calculatedOutSizeMinusOne.has_value())
1107 << dimName <<
" + pad_" << padBeforeName <<
" + pad_"
1108 << padAfterName <<
" - kernel_" << dimAxis
1109 <<
" to be wholly divisible by stride_" << dimAxis <<
", got ("
1110 << inputSize <<
" + " << padBefore <<
" + " << padAfter <<
" - "
1111 << kernelSize <<
") / " << strideSize;
1113 const int64_t calculatedOutSize = calculatedOutSizeMinusOne.value() + 1;
1114 if (ShapedType::isStatic(outputSize) && calculatedOutSize != outputSize)
1116 << dimName <<
" did not match expected: " <<
"calculated="
1117 << calculatedOutSize <<
", expected=" << outputSize;
1122 if (failed(verifyOutputSize(inputType.getDimSize(1),
1123 outputType.getDimSize(1), kernel[0], strides[0],
1124 padding[0], padding[1],
"height",
"y",
"top",
1128 if (failed(verifyOutputSize(
1129 inputType.getDimSize(2), outputType.getDimSize(2), kernel[1],
1130 strides[1], padding[2], padding[3],
"width",
"x",
"left",
"right")))
1136template <
typename T>
1139 op.getPad(), op.getInput(), op.getOutput());
1142template <
typename T>
1146 const Type inputZpETy =
1148 const Type outputZpETy =
1151 auto accType = op.getAccType();
1152 if (llvm::isa<IntegerType>(inputETy) && !accType.isInteger(32))
1153 return op.emitOpError(
"accumulator type for integer tensor is not i32");
1155 if (inputETy.
isF16() && !(accType.isF16() || accType.isF32()))
1156 return op.emitOpError(
"accumulator type for f16 tensor is not f16/f32");
1158 if (inputETy.
isBF16() && !accType.isF32())
1159 return op.emitOpError(
"accumulator type for bf16 tensor is not f32");
1161 if (inputETy.
isF32() && !accType.isF32())
1162 return op.emitOpError(
"accumulator type for f32 tensor is not f32");
1164 if (inputETy != inputZpETy)
1165 return op.emitOpError(
"expect both input and its zero point are the same "
1166 "element type, got ")
1167 << inputETy <<
" and " << inputZpETy;
1169 if (resultETy != outputZpETy)
1170 return op.emitOpError(
"expect both output and its zero point are the same "
1171 "element type, got ")
1172 << resultETy <<
" and " << outputZpETy;
1174 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1175 if (succeeded(maybeIZp) && op.verifyInputZeroPoint(*maybeIZp).failed())
1178 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1179 if (succeeded(maybeOZp) && op.verifyOutputZeroPoint(*maybeOZp).failed())
1186struct AdaptivePoolingConstShapeValues {
1187 llvm::SmallVector<int64_t> kernel;
1188 llvm::SmallVector<int64_t> stride;
1189 llvm::SmallVector<int64_t> pad;
1193template <
typename T>
1195 std::is_same_v<T, tosa::AvgPool2dAdaptiveOp> ||
1196 std::is_same_v<T, tosa::MaxPool2dAdaptiveOp>;
1198template <
typename T,
1199 typename std::enable_if<IsSupportedAdaptivePoolConstShapeVerifyOp<T>,
1202 T op, AdaptivePoolingConstShapeValues &values) {
1208LogicalResult tosa::AvgPool2dOp::verify() {
1216LogicalResult tosa::AvgPool2dAdaptiveOp::verify() {
1217 AdaptivePoolingConstShapeValues values;
1226 values.pad, getInput(), getOutput())))
1235LogicalResult tosa::ClampOp::verify() {
1237 llvm::cast<ShapedType>(getInput().
getType()).getElementType();
1238 if (
auto quantType =
1239 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(inputETy)) {
1243 llvm::cast<ShapedType>(getOutput().
getType()).getElementType();
1244 if (
auto quantType =
1245 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(outputETy)) {
1248 if (inputETy != outputETy)
1249 return emitOpError(
"input/output element types are incompatible.");
1251 auto maxValAttr = getMaxValAttr();
1252 auto minValAttr = getMinValAttr();
1256 if (inputETy.
isInteger(dataTypeBitWidth)) {
1260 auto intMaxValAttr = mlir::dyn_cast<mlir::IntegerAttr>(maxValAttr);
1261 auto intMinValAttr = mlir::dyn_cast<mlir::IntegerAttr>(minValAttr);
1262 if (!intMaxValAttr || !intMinValAttr ||
1263 (intMaxValAttr.getType() != intMinValAttr.getType()) ||
1264 (intMaxValAttr.getType() != inputETy))
1265 return emitOpError(
"min/max attributes types are incompatible with "
1266 "input/output element types.");
1269 const bool isBoolean = inputETy.
isInteger(1);
1270 const APInt minVal = intMinValAttr.getValue();
1271 const APInt maxVal = intMaxValAttr.getValue();
1272 if ((isUnsigned || isBoolean) ? maxVal.ult(minVal) : maxVal.slt(minVal))
1273 return emitOpError(
"expected min_val <= max_val, got min_val=")
1274 << minValAttr <<
", max_val=" << maxValAttr;
1279 auto floatMaxValAttr = mlir::dyn_cast<mlir::FloatAttr>(maxValAttr);
1280 auto floatMinValAttr = mlir::dyn_cast<mlir::FloatAttr>(minValAttr);
1281 if (!floatMaxValAttr || !floatMinValAttr ||
1282 (floatMaxValAttr.getType() != floatMinValAttr.getType()) ||
1283 (floatMaxValAttr.getType() != inputETy))
1284 return emitOpError(
"min/max attributes types are incompatible with "
1285 "input/output element types.");
1287 const APFloat minVal = floatMinValAttr.getValue();
1288 const APFloat maxVal = floatMaxValAttr.getValue();
1289 if (minVal.isNaN() || maxVal.isNaN())
1290 return emitOpError(
"min/max attributes should not be 'NaN', got min_val=")
1291 << minValAttr <<
", max_val=" << maxValAttr;
1293 if (maxVal < minVal)
1294 return emitOpError(
"expected min_val <= max_val, got min_val=")
1295 << minValAttr <<
", max_val=" << maxValAttr;
1315 result.addOperands({input, weight, bias, zps.first, zps.second});
1316 result.addAttribute(
"pad", pad);
1317 result.addAttribute(
"stride", stride);
1318 result.addAttribute(
"dilation", dilation);
1319 result.addAttribute(
"acc_type", accType);
1320 Type finalOutputType = outputType;
1326 result.addTypes(finalOutputType);
1337 result.addOperands({input, weight, bias, zps.first, zps.second});
1338 result.addAttribute(
"out_pad", outpad);
1339 result.addAttribute(
"stride", stride);
1340 result.addAttribute(
"acc_type", accType);
1341 Type finalOutputType = outputType;
1347 result.addTypes(finalOutputType);
1354 result.addOperands({a,
b, zps.first, zps.second});
1356 Type finalOutputType{outputType};
1359 auto inputBits = eType.getIntOrFloatBitWidth();
1361 auto outputShapedType = llvm::dyn_cast<ShapedType>(outputType);
1362 assert(outputShapedType &&
"Output must be a shaped type");
1364 IntegerType accElementType;
1365 if (inputBits == 16)
1370 finalOutputType = outputShapedType.clone(accElementType);
1372 result.addTypes(finalOutputType);
1393 DenseArrayAttr kernel, DenseArrayAttr stride,
1394 DenseArrayAttr pad, TypeAttr accType) {
1399 if (
auto quantAttr =
1401 inputZp = quantAttr.getInputZp();
1402 outputZp = quantAttr.getOutputZp();
1404 const std::optional<Value> inputZpOp =
1409 "Failed to create input zero point tensor for quantized AVG_POOL2D op");
1411 const std::optional<Value> outputZpOp =
1414 (
void)
emitError(loc,
"Failed to create output zero point tensor for "
1415 "quantized AVG_POOL2D op");
1418 if (inputZpOp && outputZpOp) {
1419 result.addOperands({input, inputZpOp.value(), outputZpOp.value()});
1424 result.addOperands({input});
1426 result.addAttribute(
"kernel", kernel);
1427 result.addAttribute(
"stride", stride);
1428 result.addAttribute(
"pad", pad);
1429 result.addAttribute(
"acc_type", accType);
1430 result.types.push_back(outputType);
1443 if (
auto quantAttr =
1445 inputZp = quantAttr.getInputZp();
1446 outputZp = quantAttr.getOutputZp();
1448 const std::optional<Value> inputZpOp =
1452 "Failed to create input zero point tensor for quantized "
1453 "AVG_POOL2D_ADAPTIVE op");
1455 const std::optional<Value> outputZpOp =
1458 (
void)
emitError(loc,
"Failed to create output zero point tensor for "
1459 "quantized AVG_POOL2D_ADAPTIVE op");
1462 if (inputZpOp && outputZpOp) {
1467 result.addOperands({input, inputZpOp.value(), outputZpOp.value(),
1468 kernelShape, strideShape, padShape});
1473 result.addOperands({input});
1475 result.addAttribute(
"acc_type", accType);
1476 result.types.push_back(outputType);
1490 input1Zp = quantAttr.getInputZp();
1491 outputZp = quantAttr.getOutputZp();
1493 const std::optional<Value> input1ZpOp =
1497 loc,
"Failed to create input1 zero point for quantized NEGATE op");
1500 const std::optional<Value> outputZpOp =
1504 loc,
"Failed to create output zero point for quantized NEGATE op");
1507 if (input1ZpOp && outputZpOp) {
1508 result.addOperands({input, input1ZpOp.value(), outputZpOp.value()});
1513 result.addOperands({input});
1516 result.types.push_back(outputType);
1529 zp =
static_cast<int32_t
>(quantAttr.getInputZp());
1532 result.addOperands({input, paddings, padConstOp});
1533 result.types.push_back(outputType);
1537 StringRef name,
Type variableType,
1542 auto shapedType = dyn_cast<ShapedType>(variableType);
1544 (
void)
emitError(loc,
"variable type must be a shaped type");
1547 if (!shapedType.hasRank()) {
1548 (
void)
emitError(loc,
"variable type must be a ranked type");
1552 auto elementType = shapedType.getElementType();
1553 auto elementTypeAttr = TypeAttr::get(elementType);
1557 result.addAttribute(
"sym_name", nameAttr);
1558 result.addAttribute(
"var_shape", varShapeAttr);
1559 result.addAttribute(
"type", elementTypeAttr);
1560 result.addAttribute(
"initial_value", initialValue);
1573 if (ShapedType::isStatic(dim1) && ShapedType::isStatic(dim2) && dim1 != dim2)
1577 return ShapedType::isDynamic(dim1) ? dim2 : dim1;
1583 for (
int i = 0, e = operands.size(); i != e; ++i) {
1585 if (!
shape.hasRank()) {
1590 outRank = std::max<int64_t>(outRank,
shape.getRank());
1593 outShape.resize(outRank, 1);
1595 for (
int i = 0, e = operands.size(); i != e; ++i) {
1597 auto rankDiff = outShape.size() -
shape.getRank();
1599 for (
size_t i = 0, e =
shape.getRank(); i < e; ++i) {
1600 auto dim1 = outShape[i + rankDiff];
1601 auto dim2 =
shape.getDimSize(i);
1603 const FailureOr<int64_t> maybeResolvedDim =
1605 if (failed(maybeResolvedDim))
1607 const int64_t resolvedDim = *maybeResolvedDim;
1608 outShape[i + rankDiff] = resolvedDim;
1615LogicalResult tosa::ArgMaxOp::inferReturnTypeComponents(
1616 MLIRContext *context, ::std::optional<Location> location,
1617 ArgMaxOp::Adaptor adaptor,
1620 inferredReturnShapes);
1623LogicalResult tosa::ArgMinOp::inferReturnTypeComponents(
1624 MLIRContext *context, ::std::optional<Location> location,
1625 ArgMinOp::Adaptor adaptor,
1628 inferredReturnShapes);
1631LogicalResult tosa::RFFT2dOp::inferReturnTypeComponents(
1632 MLIRContext *context, ::std::optional<Location> location,
1633 RFFT2dOp::Adaptor adaptor,
1635 ShapeAdaptor inputShape(adaptor.getInputReal().getType());
1637 if (!inputShape.hasRank())
1641 outputShape.resize(3, ShapedType::kDynamic);
1642 outputShape[0] = inputShape.getDimSize(0);
1643 outputShape[1] = inputShape.getDimSize(1);
1644 int64_t inWidth = inputShape.getDimSize(2);
1648 if (inWidth != ShapedType::kDynamic)
1649 outputShape[2] = inWidth / 2 + 1;
1658 const llvm::StringRef dimName) {
1659 const bool isPowerOfTwo = (dimSize & (dimSize - 1)) == 0 && dimSize > 0;
1662 << dimName <<
" to be a power of two, got " << dimSize;
1667LogicalResult tosa::RFFT2dOp::verify() {
1668 const auto outputTypes = getResultTypes();
1670 return emitOpError(
"expected output shapes to match, got ") << outputTypes;
1672 const auto inputType =
1673 llvm::dyn_cast<RankedTensorType>(getInputReal().
getType());
1677 const int64_t height = inputType.getDimSize(1);
1678 if (ShapedType::isStatic(height) &&
1682 const int64_t width = inputType.getDimSize(2);
1683 if (ShapedType::isStatic(width) &&
1687 const auto outputType = llvm::dyn_cast<RankedTensorType>(outputTypes[0]);
1693 outputType.getShape().drop_back())))
1694 return emitOpError(
"expected batch and height dimensions of input/output "
1695 "to match, got input=")
1696 << inputType <<
" output=" << outputType;
1699 const int64_t outputWidth = outputType.getDimSize(2);
1700 if (ShapedType::isStatic(width) && ShapedType::isStatic(outputWidth) &&
1701 (outputWidth != (width / 2) + 1))
1703 "expected output width to be equal to input_width / 2 + 1, got ")
1709LogicalResult tosa::FFT2dOp::inferReturnTypeComponents(
1710 MLIRContext *context, ::std::optional<Location> location,
1711 FFT2dOp::Adaptor adaptor,
1713 inferredReturnShapes.push_back(
1715 inferredReturnShapes.push_back(
1720LogicalResult tosa::FFT2dOp::verify() {
1721 const auto inputRealType =
1722 llvm::dyn_cast<RankedTensorType>(getInputReal().
getType());
1723 const auto inputImagType =
1724 llvm::dyn_cast<RankedTensorType>(getInputImag().
getType());
1725 if (!inputRealType || !inputImagType)
1728 const auto trySelectStaticDim = [](
const int64_t a,
const int64_t b) {
1729 return ShapedType::isDynamic(a) ? a :
b;
1732 const int64_t height = trySelectStaticDim(inputRealType.getDimSize(1),
1733 inputImagType.getDimSize(1));
1734 if (ShapedType::isStatic(height) &&
1738 const int64_t width = trySelectStaticDim(inputRealType.getDimSize(2),
1739 inputImagType.getDimSize(2));
1740 if (ShapedType::isStatic(width) &&
1747LogicalResult tosa::ConcatOp::inferReturnTypeComponents(
1748 MLIRContext *context, ::std::optional<Location> location,
1749 ConcatOp::Adaptor adaptor,
1752 const Properties &prop = adaptor.getProperties();
1753 int32_t axis = prop.axis.getValue().getSExtValue();
1755 bool hasRankedInput =
false;
1756 for (
auto operand : adaptor.getOperands()) {
1758 if (!operandShape.hasRank())
1762 if (!hasRankedInput)
1763 outputShape.resize(operandShape.getRank(), ShapedType::kDynamic);
1766 for (
int i = 0, s = operandShape.getRank(); i < s; i++) {
1767 if (i == axis || operandShape.isDynamicDim(i))
1769 if (outputShape[i] == ShapedType::kDynamic)
1770 outputShape[i] = operandShape.getDimSize(i);
1771 if (outputShape[i] != operandShape.getDimSize(i))
1773 "Cannot concat tensors with different sizes"
1774 " on the non-axis dimension ",
1778 hasRankedInput =
true;
1781 if (adaptor.getInput1().empty())
1785 llvm::cast<TensorType>(adaptor.getInput1().getType()[0]).getElementType();
1786 if (!hasRankedInput) {
1793 for (
auto operand : adaptor.getOperands()) {
1798 if (!operandShape.hasRank() || operandShape.isDynamicDim(axis)) {
1799 concatDimSize = ShapedType::kDynamic;
1803 concatDimSize += operandShape.getDimSize(axis);
1806 outputShape[axis] = concatDimSize;
1812LogicalResult tosa::ConcatOp::verify() {
1814 auto outType = getOutput().getType();
1818 if (inputList.empty())
1819 return emitOpError(
"expect at least one input");
1821 if (!llvm::all_of(inputList, [&](
auto input) {
1823 *
this, input.getType(), outType));
1828 const int32_t axis = getAxis();
1830 for (
const auto &input : inputList) {
1831 const Type inputType = input.getType();
1833 if (currShape.hasRank()) {
1834 firstRankedInputShape = currShape;
1836 if (axis < 0 || axis >= firstRankedInputShape.
getRank())
1837 return emitOpError(
"expect axis to be within range 0 < axis < "
1838 "rank(input1[firstRankedTensorIdx]), got ")
1844 const auto allOperandsHasRank = [](
const Value input) {
1847 if (llvm::all_of(inputList, allOperandsHasRank)) {
1850 for (
const auto &[
index, input] : llvm::enumerate(inputList.drop_front())) {
1852 const int64_t inputRank = inputShape.getRank();
1853 const size_t operandNum =
index + 1;
1856 if (inputRank != firstInputRank)
1858 "expect all operands to have the same rank, but got ")
1859 << firstInputRank <<
" vs " << inputRank <<
" on operands 0 and "
1863 for (
int i = 0; i < inputRank; i++) {
1864 const int64_t inputDim = inputShape.getDimSize(i);
1866 if (i == axis || firstRankedInputShape.
isDynamicDim(i) ||
1867 inputShape.isDynamicDim(i))
1869 if (inputDim != firstInputDim)
1870 return emitOpError(
"expect all operand shapes to have the same sizes "
1871 "on non-axis dimensions, but got ")
1872 << inputDim <<
" vs " << firstInputDim <<
" at index " << i
1873 <<
" on operands 0 and " << operandNum;
1878 if (outputShape.hasRank() && outputShape.getRank() != firstInputRank)
1879 return emitOpError(
"expect output rank to match inputs rank, got ")
1880 << outputShape.getRank() <<
" vs " << firstInputRank;
1884 for (
const auto &input : inputList) {
1886 if (inputShape.isDynamicDim(axis)) {
1891 axisSum += inputShape.getDimSize(axis);
1894 if (axisSum >= 0 && outputShape.hasRank() &&
1895 !outputShape.isDynamicDim(axis) &&
1896 axisSum != outputShape.getDimSize(axis))
1897 return emitOpError(
"requires sum of axis dimensions of input1 "
1898 "equal to output axis dimension, got ")
1899 << axisSum <<
" and " << outputShape.getDimSize(axis);
1905LogicalResult tosa::EqualOp::inferReturnTypeComponents(
1906 MLIRContext *context, ::std::optional<Location> location,
1910 auto elementType = IntegerType::get(context, 1);
1923 if (l.size() != r.size() || l.size() != 1)
1933 if (!
shape.hasRank())
1934 return ShapedType::kDynamic;
1935 const int64_t inputAxis = axis - (outputRank -
shape.getRank());
1936 return inputAxis < 0 ? 1 :
shape.getDimSize(inputAxis);
1939static FailureOr<SmallVector<int64_t>>
1941 int64_t outputRank,
bool transposeB) {
1942 if (outputRank < 2 ||
1950 for (
int64_t axis = 0; axis < outputRank - 2; ++axis) {
1954 if (failed(resolvedDim))
1956 outputShape[axis] = *resolvedDim;
1962 outputShape[outputRank - 1] =
1977 inferredReturnShapes.emplace_back();
1984 if (ShapedType::isStatic(aChannels) && ShapedType::isStatic(bChannels) &&
1985 aChannels != bChannels)
1989 FailureOr<SmallVector<int64_t>> outputShape =
1991 if (failed(outputShape))
1994 inferredReturnShapes.emplace_back(*outputShape);
1998LogicalResult tosa::MatMulOp::inferReturnTypeComponents(
1999 MLIRContext *context, ::std::optional<Location> location,
2000 MatMulOp::Adaptor adaptor,
2005 inferredReturnShapes);
2008template <
typename T>
2010 Type bElementType) {
2011 const auto aQuantizedEType =
2012 llvm::dyn_cast<quant::UniformQuantizedType>(aElementType);
2013 const auto bQuantizedEType =
2014 llvm::dyn_cast<quant::UniformQuantizedType>(bElementType);
2016 if (aQuantizedEType || bQuantizedEType) {
2017 if (!aQuantizedEType || !bQuantizedEType) {
2018 return op.emitOpError(
"expect operands to be both quantized or both not "
2020 << aElementType <<
" and " << bElementType;
2023 auto aQuantWidth = aQuantizedEType.getStorageTypeIntegralWidth();
2024 auto bQuantWidth = bQuantizedEType.getStorageTypeIntegralWidth();
2025 if (aQuantWidth != bQuantWidth) {
2026 return op.emitOpError(
"expect quantized operands to have same widths, "
2028 << aQuantWidth <<
" and " << bQuantWidth;
2035template <
typename T>
2037 StringRef inputName,
2042 Type expectedElementType = inputStorageElementType;
2044 if (isa<BlockScaledType>(inputElementType))
2045 expectedElementType = Float32Type::get(op.getContext());
2047 if (expectedElementType == zpElementType)
2051 diag << inputName <<
" and " << zpName;
2052 if (isa<BlockScaledType>(inputElementType))
2053 diag <<
" have compatible element types, got " << inputElementType
2054 <<
" and " << zpElementType;
2056 diag <<
" have the same element type, got " << inputStorageElementType
2057 <<
" and " << zpElementType;
2065 batchShape.reserve(
shape.getRank() - 2);
2066 for (
int64_t i = 0, e =
shape.getRank() - 2; i < e; ++i)
2067 batchShape.push_back(
shape.getDimSize(i));
2071template <
typename T>
2075 const auto outputType = cast<ShapedType>(op.getResult().getType());
2078 : ShapedType::kDynamic;
2086 const int64_t minimumOutputRank =
2089 if (outputType.hasRank() && outputType.getRank() < minimumOutputRank)
2090 return op.emitOpError(
"expected output rank of at least ")
2091 << minimumOutputRank <<
", got " << outputType.getRank();
2093 const bool bothInputsRanked = aShape.
hasRank() && bShape.
hasRank();
2094 if (!bothInputsRanked && !outputType.hasRank())
2099 const int64_t expectedOutputRank =
2100 bothInputsRanked ? minimumOutputRank : outputType.getRank();
2101 FailureOr<SmallVector<int64_t>> expectedOutputShape =
2103 if (failed(expectedOutputShape)) {
2105 "expected batch dimensions of a and b to be broadcast compatible, "
2108 diag <<
"] and b=[";
2114 if (outputType.hasRank())
2116 op.getOperation(), outputType, *expectedOutputShape);
2120LogicalResult MatMulOp::verify() {
2132 FailureOr<int64_t> maybeAZp = getAZeroPoint();
2133 if (succeeded(maybeAZp) && verifyAZeroPoint(*maybeAZp).failed())
2136 FailureOr<int64_t> maybeBZp = getBZeroPoint();
2137 if (succeeded(maybeBZp) && verifyBZeroPoint(*maybeBZp).failed())
2143LogicalResult tosa::MatMulTOp::inferReturnTypeComponents(
2144 MLIRContext *context, ::std::optional<Location> location,
2145 MatMulTOp::Adaptor adaptor,
2150 inferredReturnShapes);
2153LogicalResult MatMulTOp::verify() {
2165 FailureOr<int64_t> maybeAZp = getAZeroPoint();
2166 if (succeeded(maybeAZp) && verifyAZeroPoint(*maybeAZp).failed())
2169 FailureOr<int64_t> maybeBZp = getBZeroPoint();
2170 if (succeeded(maybeBZp) && verifyBZeroPoint(*maybeBZp).failed())
2176LogicalResult tosa::MatmulTBlockScaledOp::inferReturnTypeComponents(
2177 MLIRContext *context, ::std::optional<Location> location,
2178 MatmulTBlockScaledOp::Adaptor adaptor,
2182 const auto aDataShape = cast<ShapedType>(adaptor.getAData().getType());
2183 if (aDataShape.hasRank()) {
2184 outShape[0] = aDataShape.getDimSize(0);
2185 outShape[1] = aDataShape.getDimSize(1);
2188 const auto aScaleShape = cast<ShapedType>(adaptor.getAScale().getType());
2189 if (aScaleShape.hasRank()) {
2190 outShape[0] = ShapedType::isDynamic(outShape[0]) ? aScaleShape.getDimSize(0)
2192 outShape[1] = ShapedType::isDynamic(outShape[1]) ? aScaleShape.getDimSize(1)
2197 const auto bDataShape = cast<ShapedType>(adaptor.getBData().getType());
2198 if (bDataShape.hasRank()) {
2199 const int64_t bDataBatchSize = bDataShape.getDimSize(0);
2200 if (bDataBatchSize != 1)
2202 ShapedType::isDynamic(outShape[0]) ? bDataBatchSize : outShape[0];
2203 outShape[2] = bDataShape.getDimSize(1);
2206 const auto bScaleShape = cast<ShapedType>(adaptor.getBScale().getType());
2207 if (bScaleShape.hasRank()) {
2208 const int64_t bScaleBatchSize = bScaleShape.getDimSize(0);
2209 if (bScaleBatchSize != 1)
2211 ShapedType::isDynamic(outShape[0]) ? bScaleBatchSize : outShape[0];
2212 outShape[2] = ShapedType::isDynamic(outShape[2]) ? bScaleShape.getDimSize(1)
2220LogicalResult MatmulTBlockScaledOp::verify() {
2222 const Type aDataType = getAData().getType();
2223 const Type bDataType = getBData().getType();
2229 int64_t N = ShapedType::kDynamic;
2230 int64_t D = ShapedType::kDynamic;
2231 int64_t H = ShapedType::kDynamic;
2234 int64_t multiplesOfC = ShapedType::kDynamic;
2246 "a_scale",
"batch")) ||
2248 "a_scale",
"height")))
2256 "b_data",
"batch")) ||
2258 "b_data",
"channels")))
2266 "b_scale",
"batch")) ||
2268 "b_scale",
"width")) ||
2276 if (ShapedType::isStatic(N) && ShapedType::isStatic(D) && N != D && D != 1)
2277 return emitOpError(
"expect B matrix batch size to be broadcast compatible "
2279 << D <<
" vs N=" << N;
2282 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(
getBlockSize());
2283 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
2284 return emitOpError(
"expect block size to be 32, got ") << blockSize;
2285 if (ShapedType::isStatic(C) && C % blockSize != 0)
2286 return emitOpError(
"expect C to be a multiple of block size, got C=")
2287 <<
C <<
", block_size=" << blockSize;
2290 if (ShapedType::isStatic(C) && ShapedType::isStatic(multiplesOfC) &&
2291 multiplesOfC != C / blockSize)
2293 "expect scale operands dimension 2 to equal C/block_size (")
2294 <<
C <<
"/" << blockSize <<
")" <<
", got " << multiplesOfC;
2297 N = ShapedType::isDynamic(N) ? D : N;
2299 const auto outputType = cast<ShapedType>(getResult().
getType());
2300 if (outputType.hasRank() &&
2305 opError <<
" to be compatible with expected output shape ";
2313LogicalResult tosa::PadOp::inferReturnTypeComponents(
2314 MLIRContext *context, ::std::optional<Location> location,
2315 PadOp::Adaptor adaptor,
2317 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2319 cast<tosa::shapeType>(adaptor.getPadding().getType()).getRank();
2324 if (!inputShape.hasRank()) {
2325 outputShape.resize(paddingRank / 2, ShapedType::kDynamic);
2334 outputShape.resize(inputShape.getRank(), ShapedType::kDynamic);
2339 outputShape.reserve(inputShape.getRank());
2340 for (
int i = 0, s = inputShape.getRank(); i < s; i++) {
2341 if (inputShape.isDynamicDim(i)) {
2342 outputShape.push_back(ShapedType::kDynamic);
2345 auto padFront = paddingValues[i * 2];
2346 auto padBack = paddingValues[i * 2 + 1];
2347 if (padFront < 0 || padBack < 0) {
2349 outputShape.push_back(ShapedType::kDynamic);
2353 outputShape.push_back(inputShape.getDimSize(i) + padFront + padBack);
2360LogicalResult tosa::PadOp::verify() {
2367 if (
auto padConst = getPadConst()) {
2375 RankedTensorType inputType =
2376 llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2377 RankedTensorType outputType =
2378 llvm::dyn_cast<RankedTensorType>(getOutput().
getType());
2379 if (!inputType || !outputType)
2386 auto inputRank = inputType.getRank();
2391 auto paddingValues = paddingAttr.getValues<APInt>();
2392 if (paddingValues.size() !=
static_cast<size_t>(inputRank * 2))
2393 return emitOpError() <<
"padding tensor must have " << inputRank
2394 <<
" * 2 = " << inputRank * 2 <<
" elements, but got "
2395 << paddingValues.size();
2397 auto inputShape = inputType.getShape();
2398 auto outputShape = outputType.getShape();
2400 for (
int64_t i = 0; i < inputRank; ++i) {
2401 int64_t padStart = paddingValues[i * 2].getSExtValue();
2402 int64_t padEnd = paddingValues[i * 2 + 1].getSExtValue();
2404 if ((padStart < 0 && padStart != -1) || (padEnd < 0 && padEnd != -1)) {
2405 return emitOpError()
2406 <<
"invalid padding values at dimension " << i
2407 <<
": values must be non-negative or -1 for dynamic padding, got ["
2408 << padStart <<
", " << padEnd <<
"]";
2412 if (inputShape[i] == ShapedType::kDynamic ||
2413 outputShape[i] == ShapedType::kDynamic)
2416 if (outputShape[i] != inputShape[i] + padStart + padEnd) {
2417 return emitOpError() <<
"mismatch in output shape at dimension " << i
2418 <<
": expected " << inputShape[i] <<
" + "
2419 << padStart <<
" + " << padEnd <<
" = "
2420 << (inputShape[i] + padStart + padEnd)
2421 <<
", but got " << outputShape[i];
2428LogicalResult tosa::SliceOp::inferReturnTypeComponents(
2429 MLIRContext *context, ::std::optional<Location> location,
2430 SliceOp::Adaptor adaptor,
2439 auto rank = cast<tosa::shapeType>(adaptor.getSize().getType()).getRank();
2447 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2450 if (inputShape.hasRank()) {
2451 for (
size_t i = 0; i < size.size(); i++) {
2452 if (size[i] != 0 && size[i] >= -1 && start[i] >= 0 &&
2453 (ShapedType::isDynamic(inputShape.getDimSize(i)) ||
2454 start[i] < inputShape.getDimSize(i))) {
2456 if (ShapedType::isDynamic(inputShape.getDimSize(i))) {
2459 outputShape[i] = size[i];
2463 if (size[i] == -1) {
2464 outputShape[i] = inputShape.getDimSize(i) - start[i];
2465 }
else if (start[i] + size[i] <= inputShape.getDimSize(i)) {
2467 outputShape[i] = size[i];
2479LogicalResult tosa::SliceOp::verify() {
2480 const Value input = getInput1();
2481 const Value output = getOutput();
2487 const Value start = getStart();
2488 const Value size = getSize();
2492 if (inputShape.hasRank()) {
2493 const auto inputRank = inputShape.getRank();
2494 if (outputShape.hasRank() && inputRank != outputShape.getRank())
2496 "expect input1 and output to have the same ranks, got ")
2497 << inputRank <<
" and " << outputShape.getRank();
2499 const auto startShapeRank =
2500 llvm::cast<tosa::shapeType>(start.
getType()).getRank();
2501 if (inputRank != startShapeRank)
2502 return emitOpError(
"length of start is not equal to rank of input shape");
2504 const auto sizeShapeRank =
2505 llvm::cast<tosa::shapeType>(size.
getType()).getRank();
2506 if (inputRank != sizeShapeRank)
2507 return emitOpError(
"length of size is not equal to rank of input shape");
2512 if (startValues.size()) {
2513 if (llvm::any_of(startValues, [](
const int64_t v) {
2516 return emitOpError(
"start values must be non-negative, got [")
2517 << startValues <<
"]";
2523 if (
const auto blockScaledType = llvm::dyn_cast<BlockScaledType>(elemType)) {
2524 const auto startBlock = startValues.back();
2525 const auto scaleBlock =
2526 BlockShapeAttr::getBlockShapeValue(blockScaledType.getBlockShape());
2527 if (startBlock % scaleBlock != 0) {
2529 "expected start innermost block size to match data type "
2530 "for block scaled input, got start block=")
2531 << startBlock <<
", scale block=" << scaleBlock;
2539 if (llvm::any_of(sizeValues, [](
const int64_t v) {
2542 return emitOpError(
"size values must be > 0, got [") << sizeValues <<
"]";
2543 if (outputShape.hasRank()) {
2545 outputShape.getDims(outputDims);
2546 const bool hasNoInferableDims = llvm::all_of(
2548 if (hasNoInferableDims &&
2550 return emitOpError(
"expected output shape to match size values, got ")
2551 << output.
getType() <<
" vs [" << sizeValues <<
"]";
2554 if (inputShape.hasRank() && startValues.size()) {
2556 inputShape.getDims(inputDims);
2557 for (
const auto &[
index, vals] :
2558 llvm::enumerate(llvm::zip_equal(startValues, sizeValues, inputDims))) {
2559 const auto &[start, size, inputDim] = vals;
2561 ShapedType::isDynamic(inputDim))
2563 if (start + size > inputDim)
2564 return emitOpError(
"start + size must be less than or equal to input "
2565 "dimension size, got start=")
2566 << start <<
", size=" << size
2567 <<
" vs input dim size=" << inputDim <<
" at dimension "
2575LogicalResult tosa::MulOp::inferReturnTypeComponents(
2576 MLIRContext *context, ::std::optional<Location> location,
2591LogicalResult tosa::MulOp::verify() {
2592 const Value output = getOutput();
2597 if (
auto resIntType = dyn_cast<IntegerType>(resElemType)) {
2598 IntegerType lhsIntType =
2600 IntegerType rhsIntType =
2602 if (!lhsIntType || !rhsIntType || lhsIntType != rhsIntType)
2603 return emitOpError(
"requires the same element type for all operands");
2608 if (lhsIntType.getWidth() > resIntType.getWidth())
2609 return emitOpError(
"invalid data type size for operands or result");
2614 for (
int i = 0; i < 2; ++i) {
2617 "requires the same element type for all operands and results");
2621 ElementsAttr shiftElem;
2623 int32_t shift = shiftElem.getValues<IntegerAttr>()[0].getInt();
2625 return emitOpError() <<
"require shift to be 0 for float type";
2633 TypeRange operandTypes = getOperandTypes();
2634 ShapedType aType = cast<ShapedType>(operandTypes[0]);
2635 ShapedType bType = cast<ShapedType>(operandTypes[1]);
2637 const bool aHasRank = aType.hasRank();
2638 const bool bHasRank = bType.hasRank();
2640 bool hasExpectedOutputShape =
false;
2643 if (aHasRank && bHasRank) {
2644 const int64_t aRank = aType.getRank();
2645 const int64_t bRank = bType.getRank();
2647 return emitOpError(
"a and b operands don't have matching ranks, got ")
2648 << aRank <<
" and " << bRank;
2652 aType.getShape(), bType.getShape(), expectedOutputShape))
2653 return emitOpError(
"a and b operands don't have broadcast-compatible "
2655 << aType <<
" and " << bType;
2656 hasExpectedOutputShape =
true;
2659 ShapedType resultType = cast<ShapedType>(output.
getType());
2660 if (!resultType.hasRank())
2663 const int64_t resultRank = resultType.getRank();
2664 if (aHasRank && resultRank != aType.getRank())
2665 return emitOpError(
"result type has different rank than a, got ")
2666 << resultRank <<
" vs " << aType.getRank();
2667 if (bHasRank && resultRank != bType.getRank())
2668 return emitOpError(
"result type has different rank than b, got ")
2669 << resultRank <<
" vs " << bType.getRank();
2671 if (hasExpectedOutputShape &&
2673 expectedOutputShape)))
2679LogicalResult tosa::TableOp::inferReturnTypeComponents(
2680 MLIRContext *context, ::std::optional<Location> location,
2681 TableOp::Adaptor adaptor,
2683 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2685 if (!inputShape.hasRank()) {
2690 inferredReturnShapes.resize(1);
2691 inputShape.getDims(inferredReturnShapes[0]);
2695LogicalResult tosa::TableOp::verify() {
2696 const TensorType inputType = getInput1().getType();
2697 const TensorType outputType = getOutput().getType();
2706 auto inputDims = inputType.
getShape();
2707 auto outputDims = outputType.
getShape();
2708 for (
auto it : llvm::enumerate(llvm::zip(inputDims, outputDims))) {
2710 auto [inputDim, outputDim] = it.value();
2711 if (ShapedType::isStatic(outputDim) && outputDim != inputDim) {
2712 return emitOpError() <<
"dim(result, " << dim <<
") = " << outputDim
2713 <<
" doesn't match dim(input, " << dim
2714 <<
") = " << inputDim;
2727 llvm::map_to_vector(multiplesAttr.getValues<APInt>(),
2728 [](
const APInt &val) { return val.getSExtValue(); });
2732LogicalResult tosa::TileOp::inferReturnTypeComponents(
2733 MLIRContext *context, ::std::optional<Location> location,
2734 TileOp::Adaptor adaptor,
2741 cast<tosa::shapeType>(adaptor.getMultiples().getType()).getRank();
2748 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2750 if (!inputShape.hasRank()) {
2751 outputShape.resize(multiples.size(), ShapedType::kDynamic);
2752 inferredReturnShapes.push_back(
2756 if (
static_cast<size_t>(inputShape.getRank()) != multiples.size())
2760 outputShape.reserve(multiples.size());
2761 for (
int i = 0, s = inputShape.getRank(); i < s; i++) {
2762 if (multiples[i] == ShapedType::kDynamic) {
2763 outputShape.push_back(ShapedType::kDynamic);
2765 int64_t dim = inputShape.getDimSize(i);
2766 if (dim != ShapedType::kDynamic)
2767 dim *= multiples[i];
2768 outputShape.push_back(dim);
2776LogicalResult tosa::TileOp::verify() {
2782 ShapedType inputType = llvm::cast<ShapedType>(getInput1().
getType());
2783 ShapedType outputType = llvm::cast<ShapedType>(
getType());
2785 shapeType multiplesType =
2786 llvm::cast<tosa::shapeType>(getMultiples().
getType());
2788 auto multiplesRank = multiplesType.getRank();
2790 if (inputType.hasRank()) {
2791 if (inputType.getRank() != multiplesRank)
2792 return emitOpError(
"expect 'multiples' to have rank ")
2793 << inputType.getRank() <<
" but got " << multiplesRank <<
".";
2794 if (outputType.hasRank() &&
2798 }
else if (outputType.hasRank() && outputType.getRank() != multiplesRank)
2799 return emitOpError(
"expect 'multiples' array to have length ")
2800 << outputType.getRank() <<
" but got " << multiplesRank <<
".";
2803 if (getConstantMultiples(multiples).succeeded() &&
2804 llvm::any_of(multiples, [](
int64_t v) {
return v <= 0 && v != -1; }))
2806 "expect element of 'multiples' to be positive integer or -1.");
2812 if (l.size() != r.size() || l.size() != 1)
2817LogicalResult tosa::ReshapeOp::inferReturnTypeComponents(
2818 MLIRContext *context, ::std::optional<Location> location,
2819 ReshapeOp::Adaptor adaptor,
2821 ShapeAdaptor inputShape(adaptor.getInput1().getType());
2826 auto rank = cast<tosa::shapeType>(adaptor.getShape().getType()).getRank();
2835 if (!inputShape.hasRank() || !inputShape.hasStaticShape()) {
2836 inferredReturnShapes.push_back(
2844 int64_t numElements = inputShape.getNumElements();
2846 for (
auto val : newShapeValue) {
2847 if (ShapedType::isStatic(val)) {
2853 for (
auto &val : newShapeValue) {
2854 if (ShapedType::isDynamic(val))
2855 val = numElements / staticMul;
2858 inferredReturnShapes.push_back(
2863llvm::LogicalResult tosa::ReshapeOp::verify() {
2869 TensorType inputType = getInput1().getType();
2874 return mlir::success();
2878 if (missingDims > 1)
2879 return emitOpError() <<
"expected at most one target dimension to be "
2882 const auto outputType = dyn_cast<RankedTensorType>(
getType());
2886 if ((
int64_t)shapeValues.size() != outputType.getRank())
2887 return emitOpError() <<
"new shape does not match result rank";
2889 for (
auto [newShapeDim, outputShapeDim] :
2890 zip(shapeValues, outputType.getShape())) {
2892 newShapeDim != ShapedType::kDynamic &&
2893 outputShapeDim != ShapedType::kDynamic && newShapeDim != outputShapeDim)
2894 return emitOpError() <<
"new shape is inconsistent with result shape";
2897 return emitOpError() <<
"new shape has invalid tensor dimension size "
2901 if (inputType.hasStaticShape()) {
2902 int64_t inputElementsNum = inputType.getNumElements();
2903 if (outputType.hasStaticShape()) {
2904 int64_t outputElementsNum = outputType.getNumElements();
2905 if (inputElementsNum != outputElementsNum) {
2906 return emitOpError() <<
"cannot reshape " << inputElementsNum
2907 <<
" elements into " << outputElementsNum;
2913 return (dim > 0) ?
acc * dim :
acc;
2915 bool isStaticNewShape =
2916 llvm::all_of(shapeValues, [](
int64_t s) {
return s > 0; });
2917 if ((isStaticNewShape && inputElementsNum != newShapeElementsNum) ||
2918 (!isStaticNewShape && newShapeElementsNum > inputElementsNum)) {
2919 return emitOpError() <<
"cannot reshape " << inputElementsNum
2920 <<
" elements into " << newShapeElementsNum;
2924 return mlir::success();
2927bool tosa::ReshapeBlockScaledOp::isCompatibleReturnTypes(
TypeRange l,
2929 if (l.size() != r.size() || l.size() < 1 || l.size() > 2)
2937LogicalResult tosa::ReshapeBlockScaledOp::inferReturnTypeComponents(
2938 MLIRContext *context, ::std::optional<Location> location,
2939 ReshapeBlockScaledOp::Adaptor adaptor,
2942 const auto numInputs = adaptor.getInput().size();
2943 ShapeAdaptor inputShape(adaptor.getInput()[0].getType());
2946 const auto newShape = adaptor.getNewValueShape();
2948 auto rank = cast<tosa::shapeType>(newShape.getType()).getRank();
2957 const uint32_t blockSize =
2958 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
2961 if (numInputs == 2) {
2962 newScaleShapeValue.assign(newShapeValue.begin(), newShapeValue.end());
2963 if (!newScaleShapeValue.empty() &&
2964 ShapedType::isStatic(newScaleShapeValue.back()))
2965 newScaleShapeValue.back() /= blockSize;
2968 inferredReturnShapes.push_back(
2970 if (numInputs == 2) {
2972 for (
size_t idx = 0; idx < newShapeValue.size(); idx++) {
2973 if (ShapedType::isDynamic(newScaleShapeValue[idx])) {
2974 newScaleShapeValue[idx] = newShapeValue[idx];
2975 if (idx + 1 == newShapeValue.size())
2976 newScaleShapeValue[idx] /= blockSize;
2987llvm::LogicalResult tosa::ReshapeBlockScaledOp::verify() {
2991 if (inputList.size() == 0)
2992 return emitOpError(
"requires at least one input");
2994 if (inputList.size() > 2)
2995 return emitOpError(
"requires at most two inputs");
2997 if (inputList.size() != outputList.size())
2998 return emitOpError(
"requires number of results to match inputs");
3006 if (inputList.size() == 2 &&
3007 cast<tosa::shapeType>(getNewValueShape().
getType()).getRank() == 0)
3008 return emitOpError(
"requires new shape to have a rank greater than 0");
3010 const auto inputType = llvm::cast<ShapedType>(inputList[0].
getType());
3011 if (!inputType.hasRank())
3013 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(
getBlockSize());
3015 if (inputList.size() == 2) {
3016 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
3017 return emitOpError(
"expect block size to be 32, got ") << blockSize;
3018 if (llvm::any_of(inputList, [](
Value v) {
3019 const auto input = cast<ShapedType>(v.
getType());
3020 return input.hasRank() && input.getRank() == 0;
3023 "requires all input shapes have a rank greater than 0");
3024 if (llvm::any_of(outputList, [](
Value v) {
3025 const auto output = cast<ShapedType>(v.
getType());
3026 return output.hasRank() && output.getRank() == 0;
3029 "requires all result shapes have a rank greater than 0");
3037 const auto inputScaleType = llvm::cast<ShapedType>(inputList[1].
getType());
3038 if (inputScaleType.hasRank()) {
3039 if (inputType.getRank() != inputScaleType.getRank())
3040 return emitOpError(
"input shapes do not have same rank");
3043 for (
auto dimIdx = 0; dimIdx < inputType.getRank() - 1; dimIdx++) {
3044 const int64_t inputValueDim = inputType.getDimSize(dimIdx);
3045 const int64_t inputScaleDim = inputScaleType.getShape()[dimIdx];
3046 if (ShapedType::isStatic(inputValueDim) &&
3047 ShapedType::isStatic(inputScaleDim) &&
3048 inputValueDim != inputScaleDim)
3049 return emitOpError(
"input shapes for data and scale do not match on "
3056 inputType.getDimSize(inputType.getRank() - 1);
3057 if (ShapedType::isStatic(lastValueDim)) {
3058 if (lastValueDim % blockSize != 0)
3059 return emitOpError(
"expect last dimension of input_data (")
3060 << lastValueDim <<
") to be divisible by block_size ("
3061 << blockSize <<
")";
3064 inputScaleType.getDimSize(inputScaleType.getRank() - 1);
3066 if (ShapedType::isStatic(lastScaleDim) &&
3067 lastScaleDim != lastValueDim / blockSize)
3068 return emitOpError(
"expect last dimension of scale_data (")
3069 << lastScaleDim <<
") to be " << lastValueDim <<
"/"
3074 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_1))
3075 return emitOpError(
"expect block size to be 1, got ") << blockSize;
3083 return mlir::success();
3086 if (inputList.size() == 2) {
3087 const int64_t lastShapeDim = shapeValues.back();
3088 if (ShapedType::isStatic(lastShapeDim) && lastShapeDim % blockSize != 0)
3089 return emitOpError(
"expect last dimension of new shape (")
3090 << lastShapeDim <<
") to be divisible by block_size (" << blockSize
3094 const auto outputType = llvm::cast<ShapedType>(outputList[0].
getType());
3095 if (!outputType.hasRank())
3098 if (
static_cast<int64_t>(shapeValues.size()) != outputType.getRank())
3099 return emitOpError() <<
"result does not match new shape rank";
3101 for (
auto [newShapeDim, outputShapeDim] :
3102 zip(shapeValues, outputType.getShape())) {
3103 if (ShapedType::isStatic(newShapeDim) &&
3104 ShapedType::isStatic(outputShapeDim) && newShapeDim != outputShapeDim)
3105 return emitOpError() <<
"result shape is inconsistent with new shape";
3108 if (outputList.size() == 2) {
3112 scaleShapeValues.back() /= blockSize;
3114 const auto outputScaleType =
3115 llvm::cast<ShapedType>(outputList[1].
getType());
3116 if (outputScaleType.hasRank()) {
3117 if ((
int64_t)scaleShapeValues.size() != outputScaleType.getRank())
3118 return emitOpError() <<
"result scale does not match new shape rank";
3120 for (
auto [newScaleShapeDim, outputScaleShapeDim] :
3121 zip(scaleShapeValues, outputScaleType.getShape())) {
3122 if (ShapedType::isStatic(newScaleShapeDim) &&
3123 ShapedType::isStatic(outputScaleShapeDim) &&
3124 newScaleShapeDim != outputScaleShapeDim)
3125 return emitOpError()
3126 <<
"result scale shape is inconsistent with new shape";
3131 if (inputType.hasStaticShape()) {
3132 int64_t inputElementsNum = inputType.getNumElements();
3133 if (outputType.hasStaticShape()) {
3134 int64_t outputElementsNum = outputType.getNumElements();
3135 if (inputElementsNum != outputElementsNum) {
3136 return emitOpError() <<
"cannot reshape " << inputElementsNum
3137 <<
" elements into " << outputElementsNum;
3143 return (dim > 0) ?
acc * dim :
acc;
3145 bool isStaticNewShape =
3146 llvm::all_of(shapeValues, [](
int64_t s) {
return s > 0; });
3147 if ((isStaticNewShape && inputElementsNum != newShapeElementsNum) ||
3148 (!isStaticNewShape && newShapeElementsNum > inputElementsNum)) {
3149 return emitOpError() <<
"cannot reshape " << inputElementsNum
3150 <<
" elements into " << newShapeElementsNum;
3154 return mlir::success();
3161 ElementsAttr zpAttr;
3166 Type zpElemType = zpAttr.getElementType();
3168 if (llvm::isa<FloatType>(zpElemType)) {
3169 if (zpAttr.getValues<APFloat>()[0].isZero()) {
3176 if (llvm::isa<IntegerType>(zpElemType)) {
3178 return zpAttr.getValues<APInt>()[0].getSExtValue();
3179 return zpAttr.getValues<APInt>()[0].getZExtValue();
3186template <
typename T>
3188 const std::string &operand) {
3191 if (!zpElemType.
isInteger(8) && zp != 0) {
3193 std::string lower = operand;
3194 llvm::transform(lower, lower.begin(), ::tolower);
3195 return op.emitOpError()
3196 << lower <<
" zero point must be zero for non-int8 integer types";
3204 const std::string &operand) {
3205 bool isInputZp = (operand ==
"Input");
3207 bool tensorUnsigned =
3208 isInputZp ? op.getInputUnsigned() : op.getOutputUnsigned();
3209 StringRef tensorName = isInputZp ?
"input" :
"output";
3215 !(zpElemType.
isInteger(16) && tensorUnsigned)) {
3216 return op.emitOpError()
3217 <<
"expect " << tensorName <<
"_zp of 0, got " << zp;
3219 if (zpElemType.
isInteger(16) && tensorUnsigned && zp != 32768) {
3220 return op.emitOpError() <<
"expect " << tensorName
3221 <<
"_zp of 0 or 32768 for unsigned int16 "
3222 << tensorName <<
", got " << zp;
3229#define ZERO_POINT_HELPER(OP, OPERAND_NAME, SIGN_EXTEND) \
3230 FailureOr<int64_t> tosa::OP::get##OPERAND_NAME##ZeroPoint() { \
3231 return getZeroPoint(get##OPERAND_NAME##Zp(), SIGN_EXTEND); \
3233 LogicalResult tosa::OP::verify##OPERAND_NAME##ZeroPoint(int64_t zp) { \
3234 return verifyZeroPoint(*this, get##OPERAND_NAME##Zp(), zp, #OPERAND_NAME); \
3257#undef ZERO_POINT_HELPER
3259LogicalResult tosa::TransposeOp::inferReturnTypeComponents(
3260 MLIRContext *context, ::std::optional<Location> location,
3261 TransposeOp::Adaptor adaptor,
3263 ShapeAdaptor inputShape(adaptor.getInput1().getType());
3272 const auto inputRank = inputShape.
getRank();
3276 if (adaptor.getPerms().size() !=
static_cast<size_t>(inputRank)) {
3282 if (inputRank == 0) {
3288 bool allTheSame =
true;
3289 for (
int i = 1, s = inputRank; i < s; i++) {
3299 outputShape.resize(inputRank, inputShape.
getDimSize(0));
3304 outputShape.resize(inputRank, ShapedType::kDynamic);
3307 if (llvm::any_of(adaptor.getPerms(),
3308 [inputRank](
const auto i) { return i >= inputRank; }))
3311 outputShape.reserve(inputRank);
3312 for (
int i = 0, s = inputRank; i < s; i++) {
3313 outputShape[i] = inputShape.
getDimSize(adaptor.getPerms()[i]);
3320LogicalResult tosa::TransposeOp::verify() {
3332 if (inputShape.hasRank() &&
3333 constantPerms.size() !=
static_cast<size_t>(inputShape.getRank()))
3334 return emitOpError() <<
"expected perms attribute to have size "
3335 << inputShape.getRank()
3336 <<
" (input rank) but got size "
3337 << constantPerms.size();
3339 if (inputShape.hasRank() && outputShape.hasRank() &&
3340 inputShape.getRank() != outputShape.getRank())
3341 return emitOpError()
3342 <<
"expected input tensor rank to equal result tensor rank";
3344 if (outputShape.hasRank() &&
3345 constantPerms.size() !=
static_cast<size_t>(outputShape.getRank()))
3346 return emitOpError() <<
"expected perms attribute to have size "
3347 << outputShape.getRank()
3348 <<
" (output rank) but got size "
3349 << constantPerms.size();
3351 if (!llvm::all_of(constantPerms,
3352 [&constantPerms](int32_t s) {
3354 static_cast<size_t>(s) < constantPerms.size();
3357 constantPerms, [](int32_t v) ->
int64_t {
return v; })))
3358 return emitOpError() <<
"expected valid permutation indices";
3361 constantPerms.back() !=
static_cast<int32_t
>(constantPerms.size()) - 1) {
3362 return emitOpError() <<
"expected no-op permutation on innermost dimension "
3363 "for block scaled input";
3367 if (inputShape.hasStaticShape() && outputShape.hasStaticShape() &&
3368 inputShape.getNumElements() != outputShape.getNumElements())
3369 return emitOpError() <<
"expected input1 and output to have same numbers "
3371 << inputShape.getNumElements() <<
" and "
3372 << outputShape.getNumElements();
3376 if (inputShape.hasRank() && outputShape.hasRank()) {
3377 for (
auto i = 0; i < outputShape.getRank(); i++) {
3378 if (inputShape.isDynamicDim(constantPerms[i]) ||
3379 outputShape.isDynamicDim(i))
3382 if (inputShape.getDimSize(constantPerms[i]) != outputShape.getDimSize(i))
3383 return emitOpError()
3384 <<
"expected output tensor dim " << i <<
" to match "
3385 <<
"input dim " << constantPerms[i] <<
" with value of "
3386 << inputShape.getDimSize(constantPerms[i]);
3393LogicalResult TransposeOp::reifyResultShapes(
3396 const llvm::ArrayRef<int32_t> transposePerms = getPerms();
3398 Value input = getInput1();
3399 auto inputType = cast<TensorType>(input.
getType());
3401 SmallVector<OpFoldResult> returnedDims(inputType.getRank());
3402 for (
auto dim : transposePerms) {
3403 int32_t dimInInput = transposePerms[dim];
3404 if (inputType.isDynamicDim(dimInInput))
3406 tensor::DimOp::create(builder, getLoc(), input, dimInInput)
3410 builder.
getIndexAttr(inputType.getDimSize(dimInInput));
3413 reifiedReturnShapes.emplace_back(std::move(returnedDims));
3417LogicalResult tosa::GatherOp::inferReturnTypeComponents(
3418 MLIRContext *context, ::std::optional<Location> location,
3419 GatherOp::Adaptor adaptor,
3420 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3421 llvm::SmallVector<int64_t> outputShape;
3422 outputShape.resize(3, ShapedType::kDynamic);
3424 ShapeAdaptor valuesShape(adaptor.getValues().getType());
3425 if (valuesShape.hasRank()) {
3426 outputShape[0] = valuesShape.getDimSize(0);
3427 outputShape[2] = valuesShape.getDimSize(2);
3430 ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3431 if (indicesShape.hasRank()) {
3432 if (outputShape[0] == ShapedType::kDynamic)
3433 outputShape[0] = indicesShape.getDimSize(0);
3434 if (outputShape[1] == ShapedType::kDynamic)
3435 outputShape[1] = indicesShape.getDimSize(1);
3438 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3442LogicalResult tosa::RowGatherOp::inferReturnTypeComponents(
3443 MLIRContext *context, ::std::optional<Location> location,
3444 RowGatherOp::Adaptor adaptor,
3445 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3446 llvm::SmallVector<int64_t> outputShape;
3447 outputShape.resize(3, ShapedType::kDynamic);
3449 const ShapeAdaptor valuesShape(adaptor.getValues().getType());
3450 if (valuesShape.hasRank()) {
3451 outputShape[0] = valuesShape.getDimSize(0);
3452 outputShape[2] = valuesShape.getDimSize(2);
3455 const ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3456 if (indicesShape.hasRank()) {
3457 if (outputShape[0] == ShapedType::kDynamic)
3458 outputShape[0] = indicesShape.getDimSize(0);
3460 const FailureOr<int32_t> maybeRowCount =
3462 if (succeeded(maybeRowCount)) {
3463 const int64_t indicesW = indicesShape.getDimSize(1);
3464 if (ShapedType::isStatic(indicesW))
3465 outputShape[1] = indicesW * maybeRowCount.value();
3469 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3473LogicalResult tosa::RowGatherBlockScaledOp::inferReturnTypeComponents(
3474 MLIRContext *context, ::std::optional<Location> location,
3475 RowGatherBlockScaledOp::Adaptor adaptor,
3476 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3477 const auto values = adaptor.getValues();
3481 SmallVector<int64_t> dataShape(3, ShapedType::kDynamic);
3482 const ShapeAdaptor valuesShape(values.front().getType());
3483 if (valuesShape.hasRank()) {
3484 dataShape[0] = valuesShape.getDimSize(0);
3485 dataShape[2] = valuesShape.getDimSize(2);
3488 const ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3489 if (indicesShape.hasRank()) {
3490 if (dataShape[0] == ShapedType::kDynamic)
3491 dataShape[0] = indicesShape.getDimSize(0);
3495 succeeded(rowCount) && rowCount.value() > 0) {
3496 const int64_t indicesW = indicesShape.getDimSize(1);
3497 if (ShapedType::isStatic(indicesW))
3498 dataShape[1] = indicesW * rowCount.value();
3502 inferredReturnShapes.push_back(ShapedTypeComponents(dataShape));
3503 if (values.size() == 1)
3506 SmallVector<int64_t> scaleShape = dataShape;
3507 const uint32_t blockSize =
3508 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
3509 if (ShapedType::isStatic(dataShape[2]))
3510 scaleShape[2] = dataShape[2] / blockSize;
3512 inferredReturnShapes.push_back(ShapedTypeComponents(scaleShape));
3516LogicalResult tosa::GatherOp::verify() {
3523 const ShapeAdaptor valuesShape(getValues().
getType());
3525 const ShapeAdaptor outputShape(getOutput().
getType());
3527 int64_t n = ShapedType::kDynamic;
3528 int64_t w = ShapedType::kDynamic;
3529 int64_t c = ShapedType::kDynamic;
3531 if (valuesShape.hasRank()) {
3532 n = valuesShape.getDimSize(0);
3533 c = valuesShape.getDimSize(2);
3535 if (indicesShape.hasRank()) {
3536 const int64_t indicesN = indicesShape.getDimSize(0);
3537 w = indicesShape.getDimSize(1);
3538 if (n == ShapedType::kDynamic)
3540 else if (indicesN != ShapedType::kDynamic && n != indicesN)
3541 return emitOpError() <<
"requires indices dimension 0 to have size " << n
3542 <<
", got " << indicesN;
3544 if (outputShape.hasRank()) {
3545 const int64_t outputN = outputShape.getDimSize(0);
3546 const int64_t outputW = outputShape.getDimSize(1);
3547 const int64_t outputC = outputShape.getDimSize(2);
3548 if (n != ShapedType::kDynamic && outputN != ShapedType::kDynamic &&
3550 return emitOpError() <<
"requires output dimension 0 to have size " << n
3551 <<
", got " << outputN;
3553 if (w != ShapedType::kDynamic && outputW != ShapedType::kDynamic &&
3555 return emitOpError() <<
"requires output dimension 1 to have size " << w
3556 <<
", got " << outputW;
3557 if (c != ShapedType::kDynamic && outputC != ShapedType::kDynamic &&
3559 return emitOpError() <<
"requires output dimension 2 to have size " << c
3560 <<
", got " << outputC;
3565LogicalResult tosa::RowGatherOp::verify() {
3570 const FailureOr<int32_t> maybeRowCount =
3572 if (succeeded(maybeRowCount) && maybeRowCount.value() <= 0)
3573 return emitOpError() <<
"requires row_count to be > 0, got "
3574 << maybeRowCount.value();
3576 int64_t n = ShapedType::kDynamic;
3577 int64_t c = ShapedType::kDynamic;
3578 int64_t w = ShapedType::kDynamic;
3580 const ShapeAdaptor valuesShape(getValues().
getType());
3581 if (valuesShape.hasRank()) {
3582 n = valuesShape.getDimSize(0);
3583 c = valuesShape.getDimSize(2);
3587 if (indicesShape.hasRank()) {
3589 "indices",
"batch")))
3591 w = indicesShape.getDimSize(1);
3594 const ShapeAdaptor outputShape(getOutput().
getType());
3595 if (outputShape.hasRank()) {
3597 "output",
"batch")) ||
3599 "output",
"channels")))
3602 if (succeeded(maybeRowCount) && maybeRowCount.value() > 0 &&
3603 ShapedType::isStatic(w)) {
3604 const int64_t expectedOutputRows = w * maybeRowCount.value();
3605 if (ShapedType::isStatic(outputShape.getDimSize(1)) &&
3606 outputShape.getDimSize(1) != expectedOutputRows)
3607 return emitOpError()
3608 <<
"requires output dimension to be equal to "
3609 "indices[1]*row_count ("
3610 << expectedOutputRows <<
"), got " << outputShape.getDimSize(1);
3617LogicalResult tosa::RowGatherBlockScaledOp::verify() {
3618 const OperandRange values = getValues();
3619 const ResultRange output = getOutput();
3620 if (values.empty() || values.size() > 2)
3621 return emitOpError()
3622 <<
"expects values tensor list length to be 1 or 2, got "
3624 if (output.size() != values.size())
3625 return emitOpError()
3626 <<
"expects output tensor list length to match values tensor list "
3628 << output.size() <<
" results for " << values.size()
3629 <<
" input tensors";
3631 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(
getBlockSize());
3632 if (values.size() == 1 && blockSize != 1)
3633 return emitOpError()
3634 <<
"requires block_size to be BLOCK_SIZE_1 when values tensor list "
3636 if (values.size() == 2 && blockSize == 1)
3637 return emitOpError()
3638 <<
"requires block_size to not be BLOCK_SIZE_1 when values tensor "
3642 output[0].
getType(),
"values[0]",
3647 "values[1]",
"output[1]")))
3651 succeeded(rowCount) && rowCount.value() <= 0)
3652 return emitOpError() <<
"requires row_count to be > 0, got "
3653 << rowCount.value();
3655 int64_t n = ShapedType::kDynamic;
3656 int64_t k = ShapedType::kDynamic;
3657 int64_t c = ShapedType::kDynamic;
3658 int64_t w = ShapedType::kDynamic;
3659 int64_t multiplesOfC = ShapedType::kDynamic;
3661 const ShapeAdaptor valuesDataShape(values[0].
getType());
3662 if (valuesDataShape.hasRank()) {
3663 n = valuesDataShape.getDimSize(0);
3664 k = valuesDataShape.getDimSize(1);
3665 c = valuesDataShape.getDimSize(2);
3668 if (ShapedType::isStatic(c) && c % blockSize != 0)
3669 return emitOpError() <<
"expects channels of values[0] (" << c
3670 <<
") to be divisible by block_size (" << blockSize
3674 if (indicesShape.hasRank()) {
3676 "indices",
"batch")))
3678 w = indicesShape.getDimSize(1);
3681 const ShapeAdaptor outputDataShape(output[0].
getType());
3682 if (outputDataShape.hasRank()) {
3684 "output[0]",
"batch")) ||
3686 "output[0]",
"channels")))
3690 succeeded(rowCount) && rowCount.value() > 0 &&
3691 ShapedType::isStatic(w)) {
3692 const int64_t expectedOutputRows = w * rowCount.value();
3693 if (ShapedType::isStatic(outputDataShape.getDimSize(1)) &&
3694 outputDataShape.getDimSize(1) != expectedOutputRows)
3695 return emitOpError() <<
"requires output[0] dimension 1 to have size "
3696 << expectedOutputRows <<
", got "
3697 << outputDataShape.getDimSize(1);
3701 if (values.size() == 2) {
3702 const ShapeAdaptor valuesScaleShape(values[1].
getType());
3703 if (valuesScaleShape.hasRank()) {
3705 "values[1]",
"batch")) ||
3707 "values[1]",
"rows")))
3709 multiplesOfC = valuesScaleShape.getDimSize(2);
3712 const ShapeAdaptor outputScaleShape(output[1].
getType());
3713 if (outputScaleShape.hasRank()) {
3715 "output[1]",
"batch")))
3719 succeeded(rowCount) && rowCount.value() > 0 &&
3720 ShapedType::isStatic(w)) {
3721 const int64_t expectedOutputRows = w * rowCount.value();
3722 if (ShapedType::isStatic(outputScaleShape.getDimSize(1)) &&
3723 outputScaleShape.getDimSize(1) != expectedOutputRows)
3724 return emitOpError() <<
"requires output[1] dimension 1 to have size "
3725 << expectedOutputRows <<
", got "
3726 << outputScaleShape.getDimSize(1);
3729 if (ShapedType::isDynamic(multiplesOfC))
3730 multiplesOfC = outputScaleShape.getDimSize(2);
3731 else if (ShapedType::isStatic(outputScaleShape.getDimSize(2)) &&
3732 multiplesOfC != outputScaleShape.getDimSize(2))
3733 return emitOpError()
3734 <<
"expected channels of output[1] to match size "
3735 << multiplesOfC <<
", got " << outputScaleShape.getDimSize(2);
3738 if (ShapedType::isStatic(c) && ShapedType::isStatic(multiplesOfC) &&
3739 multiplesOfC != c / blockSize)
3740 return emitOpError()
3741 <<
"expects channels of scale tensors to equal C/block_size (" << c
3742 <<
"/" << blockSize <<
"), got " << multiplesOfC;
3748LogicalResult tosa::ResizeOp::inferReturnTypeComponents(
3749 MLIRContext *context, ::std::optional<Location> location,
3750 ResizeOp::Adaptor adaptor,
3751 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3752 llvm::SmallVector<int64_t, 4> outputShape;
3753 outputShape.resize(4, ShapedType::kDynamic);
3755 ShapeAdaptor inputShape(adaptor.getInput().getType());
3756 if (!inputShape.hasRank())
3759 outputShape[0] = inputShape.getDimSize(0);
3760 outputShape[3] = inputShape.getDimSize(3);
3761 int64_t inputHeight = inputShape.getDimSize(1);
3762 int64_t inputWidth = inputShape.getDimSize(2);
3764 if ((inputHeight == ShapedType::kDynamic) ||
3765 (inputWidth == ShapedType::kDynamic))
3768 SmallVector<int64_t> scaleInt, offsetInt, borderInt;
3779 const int64_t outputHeight =
3780 (((inputHeight - 1) * scaleInt[0] - offsetInt[0] + borderInt[0]) /
3784 const int64_t outputWidth =
3785 (((inputWidth - 1) * scaleInt[2] - offsetInt[1] + borderInt[1]) /
3789 if (outputHeight < 0 || outputWidth < 0) {
3792 "calculated output height and width must be non-negative, "
3794 outputHeight,
", width = ", outputWidth);
3797 outputShape[1] = outputHeight;
3798 outputShape[2] = outputWidth;
3799 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3803LogicalResult tosa::ResizeOp::verify() {
3804 const Value input = getInput();
3805 const Value output = getOutput();
3808 if (isa<BlockScaledType>(inputElementType) &&
3809 getMode() != ResizeMode::NEAREST_NEIGHBOR)
3810 return emitOpError(
"requires NEAREST_NEIGHBOR mode for block scaled input");
3812 const RankedTensorType inputType =
3813 llvm::dyn_cast<RankedTensorType>(input.
getType());
3814 const RankedTensorType outputType =
3815 llvm::dyn_cast<RankedTensorType>(output.
getType());
3817 SmallVector<int64_t> scaleValues;
3818 SmallVector<int64_t> offsetValues;
3819 SmallVector<int64_t> borderValues;
3827 if (llvm::any_of(scaleValues, [](int64_t s) {
return s <= 0; }))
3828 return emitOpError(
"expect all scale values to be > 0, got ")
3831 const int64_t scaleYN = scaleValues[0];
3832 const int64_t scaleYD = scaleValues[1];
3833 const int64_t scaleXN = scaleValues[2];
3834 const int64_t scaleXD = scaleValues[3];
3836 const int64_t offsetY = offsetValues[0];
3837 const int64_t offsetX = offsetValues[1];
3839 const int64_t borderY = borderValues[0];
3840 const int64_t borderX = borderValues[1];
3847 const int64_t oh = outputType.getDimSize(1);
3848 const int64_t ow = outputType.getDimSize(2);
3849 const int64_t ih = inputType.getDimSize(1);
3850 const int64_t iw = inputType.getDimSize(2);
3856 if (ih != ShapedType::kDynamic && ih != 1) {
3857 const std::optional<int64_t> calculatedOutHeightMinusOne =
3858 idivCheck((ih - 1) * scaleYN - offsetY + borderY, scaleYD);
3859 if (!calculatedOutHeightMinusOne.has_value())
3860 return emitOpError(
"expected (input_height - 1) * scale_y_n - offset_y + "
3862 <<
"to be wholly divisible by scale_y_d, got ((" << ih
3863 <<
" - 1) * " << scaleYN <<
" - " << offsetY <<
" + " << borderY
3864 <<
") / " << scaleYD;
3865 const int64_t calculatedOutHeight = calculatedOutHeightMinusOne.value() + 1;
3866 if (oh != ShapedType::kDynamic && calculatedOutHeight != oh)
3867 return emitOpError(
"calculated output height did not match expected: ")
3868 <<
"calculated=" << calculatedOutHeight <<
", expected=" << oh;
3875 if (iw != ShapedType::kDynamic && iw != 1) {
3876 const int64_t scaledInWidth = (iw - 1) * scaleXN - offsetX + borderX;
3877 const std::optional<int64_t> calculatedOutWidthMinusOne =
3879 if (!calculatedOutWidthMinusOne.has_value())
3880 return emitOpError(
"expected (input_width - 1) * scale_x_n - offset_x + "
3882 <<
"to be wholly divisible by scale_x_d, got ((" << iw
3883 <<
" - 1) * " << scaleXN <<
" - " << offsetX <<
" + " << borderX
3884 <<
") / " << scaleXD;
3885 const int64_t calculatedOutWidth = calculatedOutWidthMinusOne.value() + 1;
3886 if (ow != ShapedType::kDynamic && calculatedOutWidth != ow)
3887 return emitOpError(
"calculated output width did not match expected: ")
3888 <<
"calculated=" << calculatedOutWidth <<
", expected=" << ow;
3894LogicalResult tosa::ScatterOp::inferReturnTypeComponents(
3895 MLIRContext *context, ::std::optional<Location> location,
3896 ScatterOp::Adaptor adaptor,
3897 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
3898 llvm::SmallVector<int64_t> outputShape;
3899 outputShape.resize(3, ShapedType::kDynamic);
3901 ShapeAdaptor valuesInShape(adaptor.getValuesIn().getType());
3902 if (valuesInShape.hasRank()) {
3903 outputShape[0] = valuesInShape.getDimSize(0);
3904 outputShape[1] = valuesInShape.getDimSize(1);
3905 outputShape[2] = valuesInShape.getDimSize(2);
3908 ShapeAdaptor indicesShape(adaptor.getIndices().getType());
3909 if (indicesShape.hasRank()) {
3910 if (outputShape[0] == ShapedType::kDynamic)
3911 outputShape[0] = indicesShape.getDimSize(0);
3914 ShapeAdaptor inputShape(adaptor.getInput().getType());
3915 if (inputShape.hasRank()) {
3916 if (outputShape[0] == ShapedType::kDynamic)
3917 outputShape[0] = inputShape.getDimSize(0);
3918 if (outputShape[2] == ShapedType::kDynamic)
3919 outputShape[2] = inputShape.getDimSize(2);
3922 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
3926LogicalResult tosa::ScatterOp::verify() {
3936 const ShapeAdaptor valuesInShape(getValuesIn().
getType());
3938 const ShapeAdaptor inputShape(getInput().
getType());
3939 const ShapeAdaptor outputShape(getValuesOut().
getType());
3941 int64_t n = ShapedType::kDynamic;
3942 int64_t k = ShapedType::kDynamic;
3943 int64_t w = ShapedType::kDynamic;
3944 int64_t c = ShapedType::kDynamic;
3945 if (valuesInShape.hasRank()) {
3946 n = valuesInShape.getDimSize(0);
3947 k = valuesInShape.getDimSize(1);
3948 c = valuesInShape.getDimSize(2);
3950 if (indicesShape.hasRank()) {
3951 const int64_t indicesN = indicesShape.getDimSize(0);
3952 w = indicesShape.getDimSize(1);
3953 if (n == ShapedType::kDynamic)
3955 else if (indicesN != ShapedType::kDynamic && n != indicesN)
3956 return emitOpError() <<
"requires indices dimension 0 to have size " << n
3957 <<
", got " << indicesN;
3959 if (inputShape.hasRank()) {
3960 const int64_t inputN = inputShape.getDimSize(0);
3961 const int64_t inputW = inputShape.getDimSize(1);
3962 const int64_t inputC = inputShape.getDimSize(2);
3963 if (n == ShapedType::kDynamic)
3965 else if (inputN != ShapedType::kDynamic && n != inputN)
3966 return emitOpError() <<
"requires input dimension 0 to have size " << n
3967 <<
", got " << inputN;
3968 if (w == ShapedType::kDynamic)
3970 else if (inputW != ShapedType::kDynamic && w != inputW)
3971 return emitOpError() <<
"requires input dimension 1 to have size " << w
3972 <<
", got " << inputW;
3974 if (c == ShapedType::kDynamic)
3976 else if (inputC != ShapedType::kDynamic && c != inputC)
3977 return emitOpError() <<
"requires input dimension 2 to have size " << c
3978 <<
", got " << inputC;
3980 if (outputShape.hasRank()) {
3981 const int64_t outputN = outputShape.getDimSize(0);
3982 const int64_t outputK = outputShape.getDimSize(1);
3983 const int64_t outputC = outputShape.getDimSize(2);
3984 if (n != ShapedType::kDynamic && outputN != ShapedType::kDynamic &&
3986 return emitOpError() <<
"requires values_out dimension 0 to have size "
3987 << n <<
", got " << outputN;
3988 if (k == ShapedType::kDynamic)
3990 else if (outputK != ShapedType::kDynamic && k != outputK)
3991 return emitOpError() <<
"requires values_out dimension 1 to have size "
3992 << k <<
", got " << outputK;
3993 if (c != ShapedType::kDynamic && outputC != ShapedType::kDynamic &&
3995 return emitOpError() <<
"requires values_out dimension 2 to have size "
3996 << c <<
", got " << outputC;
3998 if (k != ShapedType::kDynamic && w != ShapedType::kDynamic && !(k >= w))
3999 return emitOpError() <<
"requires dimensions K >= W, got K=" << k
4008 int64_t axisVal = axis.getValue().getSExtValue();
4009 if (!operandShape.
hasRank() || operandShape.
getRank() <= axisVal) {
4015 operandShape.
getDims(outputShape);
4016 outputShape[axisVal] = 1;
4021#define COMPATIBLE_RETURN_TYPES(OP) \
4022 bool OP::isCompatibleReturnTypes(TypeRange l, TypeRange r) { \
4023 if (l.size() != r.size() || l.size() != 1) \
4025 if (getElementTypeOrSelf(l[0]) != getElementTypeOrSelf(r[0])) \
4027 return succeeded(verifyCompatibleShape(l[0], r[0])); \
4030#define REDUCE_SHAPE_INFER(OP) \
4031 LogicalResult OP::inferReturnTypeComponents( \
4032 MLIRContext *context, ::std::optional<Location> location, \
4033 OP::Adaptor adaptor, \
4034 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) { \
4036 llvm::cast<TensorType>(adaptor.getInput().getType()).getElementType(); \
4037 ShapeAdaptor inputShape(adaptor.getInput().getType()); \
4038 const Properties &prop = adaptor.getProperties(); \
4039 return ReduceInferReturnTypes(inputShape, inputType, prop.axis, \
4040 inferredReturnShapes); \
4042 COMPATIBLE_RETURN_TYPES(OP)
4050#undef REDUCE_SHAPE_INFER
4052#undef COMPATIBLE_RETURN_TYPES
4054template <
typename T>
4057 TensorType inputType = op.getInput().getType();
4058 TensorType outputType = op.getOutput().getType();
4059 int32_t reduceAxis = op.getAxis();
4061 if (reduceAxis < 0) {
4062 op.emitOpError(
"reduce axis must not be negative");
4066 int64_t inputRank = inputType.getRank();
4069 if (reduceAxis >= inputRank && (reduceAxis != 0 || inputRank != 0)) {
4070 op.emitOpError(
"expect input tensor rank (")
4071 << inputRank <<
") to be larger than reduce axis (" << reduceAxis
4077 int64_t outputRank = outputType.getRank();
4078 if (inputType.
hasRank() && outputRank != inputType.getRank()) {
4080 "expect output tensor rank to be equal to input tensor rank");
4083 if (reduceAxis >= outputRank && (reduceAxis != 0 || outputRank != 0)) {
4084 op.emitOpError(
"expect output tensor rank (")
4085 << outputRank <<
") to be larger than reduce axis (" << reduceAxis
4091 if (outputRank != 0) {
4092 auto outputShape = outputType.
getShape();
4093 if (!outputType.isDynamicDim(reduceAxis) &&
4094 outputShape[reduceAxis] != 1) {
4095 op.emitOpError(
"expect reduced dimension size to be 1, got ")
4096 << outputShape[reduceAxis];
4104LogicalResult tosa::ReduceAllOp::verify() {
return verifyReduceOp(*
this); }
4105LogicalResult tosa::ReduceAnyOp::verify() {
return verifyReduceOp(*
this); }
4106LogicalResult tosa::ReduceMaxOp::verify() {
return verifyReduceOp(*
this); }
4107LogicalResult tosa::ReduceMinOp::verify() {
return verifyReduceOp(*
this); }
4108LogicalResult tosa::ReduceProductOp::verify() {
return verifyReduceOp(*
this); }
4109LogicalResult tosa::ReduceSumOp::verify() {
return verifyReduceOp(*
this); }
4123#define NARY_SHAPE_INFER(OP) \
4124 LogicalResult OP::inferReturnTypeComponents( \
4125 MLIRContext *context, ::std::optional<Location> location, \
4126 ValueShapeRange operands, DictionaryAttr attributes, \
4127 PropertyRef properties, RegionRange regions, \
4128 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) { \
4129 return NAryInferReturnTypes(operands, inferredReturnShapes); \
4169#undef PRED_SHAPE_INFER
4171LogicalResult tosa::NegateOp::inferReturnTypeComponents(
4172 MLIRContext *context, ::std::optional<Location> location,
4173 NegateOp::Adaptor adaptor,
4175 ShapeAdaptor inputShape(adaptor.getInput1().getType());
4180LogicalResult tosa::NegateOp::verify() {
4182 const Type input1Type = getInput1().getType();
4183 const Type outputType = getOutput().getType();
4188 const SmallVector<Type, 2> types = {input1Type, outputType};
4190 return emitOpError() <<
"requires the same shape for input1 and output";
4193 const Type input1ZpEType =
4195 if (input1EType != input1ZpEType) {
4196 return emitOpError(
"expect both input1 and its zero point are the same "
4197 "element type, got ")
4198 << input1EType <<
" and " << input1ZpEType;
4201 const Type outputZpEType =
4203 if (outputEType != outputZpEType) {
4204 return emitOpError(
"expect both output and its zero point are the same "
4205 "element type, got ")
4206 << outputEType <<
" and " << outputZpEType;
4209 FailureOr<int64_t> maybeIZp = getInput1ZeroPoint();
4210 if (succeeded(maybeIZp) && verifyInput1ZeroPoint(*maybeIZp).failed())
4213 FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
4214 if (succeeded(maybeOZp) && verifyOutputZeroPoint(*maybeOZp).failed())
4225 outputShape.resize(4, ShapedType::kDynamic);
4240 if (ShapedType::isStatic(height)) {
4241 int64_t padded = height + pad[0] + pad[1] - kernel[0];
4242 outputShape[1] = padded / stride[0] + 1;
4245 if (ShapedType::isStatic(width)) {
4246 int64_t padded = width + pad[2] + pad[3] - kernel[1];
4247 outputShape[2] = padded / stride[1] + 1;
4254template <
typename AdaptorT>
4260 if (ShapedType::isDynamic(current))
4261 current = candidate;
4270 : adaptor(adaptor) {}
4274 const ShapeAdaptor inputShape(adaptor.getInput().getType());
4282 outputShape[0] = outputBatch;
4283 inputSpatial[0] = inputHeight;
4284 inputSpatial[1] = inputWidth;
4289 const ShapeAdaptor weightShape(adaptor.getWeight().getType());
4297 outputShape[3] = outputChannels;
4298 weightSpatial[0] = kernelHeight;
4299 weightSpatial[1] = kernelWidth;
4308 padValues.assign(adaptor.getPad().begin(), adaptor.getPad().end());
4309 strideValues.assign(adaptor.getStride().begin(), adaptor.getStride().end());
4310 dilationValues.assign(adaptor.getDilation().begin(),
4311 adaptor.getDilation().end());
4316 Conv2DOp::Adaptor adaptor;
4324 : adaptor(adaptor) {}
4328 const ShapeAdaptor inputDataShape(adaptor.getInputData().getType());
4329 if (inputDataShape.
hasRank()) {
4334 outputShape[0] = outputBatch;
4335 inputSpatial[0] = inputHeight;
4336 inputSpatial[1] = inputWidth;
4339 const ShapeAdaptor inputScaleShape(adaptor.getInputScale().getType());
4340 if (!inputScaleShape.
hasRank())
4354 const ShapeAdaptor weightDataShape(adaptor.getWeightData().getType());
4355 if (weightDataShape.
hasRank()) {
4360 outputShape[3] = outputChannels;
4361 weightSpatial[0] = kernelHeight;
4362 weightSpatial[1] = kernelWidth;
4365 const ShapeAdaptor weightScaleShape(adaptor.getWeightScale().getType());
4366 if (!weightScaleShape.
hasRank())
4395 Conv2DBlockScaledOp::Adaptor adaptor;
4403 : adaptor(adaptor) {}
4407 const ShapeAdaptor inputShape(adaptor.getInput().getType());
4416 outputShape[0] = outputBatch;
4417 inputSpatial[0] = inputDepth;
4418 inputSpatial[1] = inputHeight;
4419 inputSpatial[2] = inputWidth;
4424 const ShapeAdaptor weightShape(adaptor.getWeight().getType());
4433 outputShape[4] = outputChannels;
4434 weightSpatial[0] = kernelDepth;
4435 weightSpatial[1] = kernelHeight;
4436 weightSpatial[2] = kernelWidth;
4445 padValues.assign(adaptor.getPad().begin(), adaptor.getPad().end());
4446 strideValues.assign(adaptor.getStride().begin(), adaptor.getStride().end());
4447 dilationValues.assign(adaptor.getDilation().begin(),
4448 adaptor.getDilation().end());
4453 Conv3DOp::Adaptor adaptor;
4456template <
typename AdaptorT>
4462 ShapedType::kDynamic);
4464 ShapedType::kDynamic);
4466 ShapedType::kDynamic);
4468 convShapeAdaptor.inferInputShape(outputShape, inputSpatial);
4469 convShapeAdaptor.inferWeightShape(outputShape, weightSpatial);
4471 const ShapeAdaptor biasShape = adaptor.getBias().getType();
4474 if (biasSize != 1) {
4475 const size_t outputChannelDim = convShapeAdaptor.getOutputRank() - 1;
4476 outputShape[outputChannelDim] =
4477 ShapedType::isDynamic(outputShape[outputChannelDim])
4479 : outputShape[outputChannelDim];
4486 if (failed(convShapeAdaptor.getSpatialParameters(padValues, strideValues,
4492 for (
int64_t dim = 0; dim < convShapeAdaptor.getNumSpatialDims(); ++dim) {
4493 if (!ShapedType::isStatic(inputSpatial[dim]) ||
4494 !ShapedType::isStatic(weightSpatial[dim]))
4497 inputSpatial[dim] + padValues[2 * dim] + padValues[2 * dim + 1];
4499 (weightSpatial[dim] - 1) * dilationValues[dim] + 1;
4500 const int64_t unstridedResult = inputSize - filterSize + 1;
4501 outputShape[dim + 1] = (unstridedResult - 1) / strideValues[dim] + 1;
4508LogicalResult Conv2DOp::inferReturnTypeComponents(
4509 MLIRContext *context, ::std::optional<Location> location,
4510 Conv2DOp::Adaptor adaptor,
4511 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4515LogicalResult Conv2DOp::verify() {
4522LogicalResult Conv2DBlockScaledOp::inferReturnTypeComponents(
4523 MLIRContext *context, ::std::optional<Location> location,
4524 Conv2DBlockScaledOp::Adaptor adaptor,
4525 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4529LogicalResult Conv2DBlockScaledOp::verify() {
4531 getWeightData().
getType(),
"input_data",
4534 getWeightScale().
getType(),
"input_scale",
4537 getOutput().
getType(),
"bias",
"output")))
4541 int64_t N = ShapedType::kDynamic;
4542 int64_t IH = ShapedType::kDynamic;
4543 int64_t IW = ShapedType::kDynamic;
4544 int64_t IC = ShapedType::kDynamic;
4545 int64_t multiplesOfIC = ShapedType::kDynamic;
4546 int64_t OC = ShapedType::kDynamic;
4547 int64_t KH = ShapedType::kDynamic;
4548 int64_t KW = ShapedType::kDynamic;
4550 const ShapeAdaptor inputDataShape(getInputData().
getType());
4551 if (inputDataShape.hasRank()) {
4552 N = inputDataShape.getDimSize(0);
4553 IH = inputDataShape.getDimSize(1);
4554 IW = inputDataShape.getDimSize(2);
4555 IC = inputDataShape.getDimSize(3);
4558 const ShapeAdaptor inputScaleShape(getInputScale().
getType());
4559 if (inputScaleShape.hasRank()) {
4561 "input_scale",
"batch size")) ||
4563 "input_scale",
"input height")) ||
4565 "input_scale",
"input width")))
4567 multiplesOfIC = inputScaleShape.getDimSize(3);
4570 const ShapeAdaptor weightDataShape(getWeightData().
getType());
4571 if (weightDataShape.hasRank()) {
4572 OC = weightDataShape.getDimSize(0);
4573 KH = weightDataShape.getDimSize(1);
4574 KW = weightDataShape.getDimSize(2);
4576 "weight_data",
"input channels")))
4580 const ShapeAdaptor weightScaleShape(getWeightScale().
getType());
4581 if (weightScaleShape.hasRank()) {
4583 "weight_scale",
"output channels")) ||
4585 "weight_scale",
"kernel height")) ||
4587 "weight_scale",
"kernel width")) ||
4589 weightScaleShape.getDimSize(3),
4590 "weight_scale",
"input channel blocks")))
4594 const uint32_t blockSize = BlockSizeAttr::getBlockSizeValue(
getBlockSize());
4595 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
4596 return emitOpError(
"expect block size to be 32, got ") << blockSize;
4598 if (ShapedType::isStatic(IC) && IC % blockSize != 0)
4599 return emitOpError(
"expect IC to be a multiple of block size, got IC=")
4600 << IC <<
", block_size=" << blockSize;
4603 if (ShapedType::isStatic(IC) && ShapedType::isStatic(multiplesOfIC) &&
4604 multiplesOfIC != IC / blockSize)
4606 "expect scale operands dimension 2 to equal IC/block_size (")
4607 << IC <<
"/" << blockSize <<
")"
4608 <<
", got " << multiplesOfIC;
4611 SmallVector<int64_t> padValues;
4613 if (llvm::any_of(padValues, [](int64_t p) {
return p < 0; }))
4614 return emitOpError(
"expect all padding values to be >= 0, got ")
4618 SmallVector<int64_t> strideValues;
4620 if (llvm::any_of(strideValues, [](int64_t s) {
return s < 1; }))
4621 return emitOpError(
"expect all stride values to be >= 1, got ")
4625 SmallVector<int64_t> dilationValues;
4628 if (llvm::any_of(dilationValues, [](int64_t d) {
return d < 1; }))
4629 return emitOpError(
"expect all dilation values to be >= 1, got ")
4634 const ShapeAdaptor outputShape(getOutput().
getType());
4635 if (!padValues.empty() && !strideValues.empty() && !dilationValues.empty() &&
4636 outputShape.hasRank()) {
4638 padValues[0], padValues[1], strideValues[0],
4639 dilationValues[0],
"height",
"y",
"top",
4642 padValues[2], padValues[3], strideValues[1],
4643 dilationValues[1],
"width",
"x",
"left",
4649 const ShapeAdaptor biasShape(getBias().
getType());
4650 if (biasShape.hasRank() && outputShape.hasRank()) {
4651 const int64_t biasChannels = biasShape.getDimSize(0);
4652 const int64_t outputChannels =
4653 outputShape.getDimSize(outputShape.getRank() - 1);
4654 if (biasChannels == ShapedType::kDynamic ||
4655 outputChannels == ShapedType::kDynamic)
4659 if (biasChannels != outputChannels && biasChannels != 1)
4661 "bias channels expected to be equal to output channels (")
4662 << outputChannels <<
") or 1, got " << biasChannels;
4668LogicalResult Conv3DOp::inferReturnTypeComponents(
4669 MLIRContext *context, ::std::optional<Location> location,
4670 Conv3DOp::Adaptor adaptor,
4671 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4675LogicalResult Conv3DOp::verify() {
4682LogicalResult AvgPool2dOp::inferReturnTypeComponents(
4683 MLIRContext *context, ::std::optional<Location> location,
4684 AvgPool2dOp::Adaptor adaptor,
4685 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4686 ShapeAdaptor inputShape(adaptor.getInput().getType());
4687 const Properties &prop = adaptor.getProperties();
4689 inferredReturnShapes);
4692LogicalResult AvgPool2dAdaptiveOp::inferReturnTypeComponents(
4693 MLIRContext *context, ::std::optional<Location> location,
4694 AvgPool2dAdaptiveOp::Adaptor adaptor,
4695 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4696 ShapeAdaptor inputShape(adaptor.getInput().getType());
4698 llvm::SmallVector<int64_t> kernelValues;
4699 llvm::SmallVector<int64_t> strideValues;
4700 llvm::SmallVector<int64_t> padValues;
4707 padValues, inferredReturnShapes);
4710 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4711 if (inputShape.hasRank()) {
4713 outputShape[0] = inputShape.getDimSize(0);
4714 outputShape[3] = inputShape.getDimSize(3);
4717 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4721LogicalResult MaxPool2dOp::inferReturnTypeComponents(
4722 MLIRContext *context, ::std::optional<Location> location,
4723 MaxPool2dOp::Adaptor adaptor,
4724 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4725 ShapeAdaptor inputShape(adaptor.getInput().getType());
4726 const Properties &prop = adaptor.getProperties();
4728 inferredReturnShapes);
4731LogicalResult MaxPool2dAdaptiveOp::inferReturnTypeComponents(
4732 MLIRContext *context, ::std::optional<Location> location,
4733 MaxPool2dAdaptiveOp::Adaptor adaptor,
4734 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4735 ShapeAdaptor inputShape(adaptor.getInput().getType());
4737 llvm::SmallVector<int64_t> kernelValues;
4738 llvm::SmallVector<int64_t> strideValues;
4739 llvm::SmallVector<int64_t> padValues;
4746 padValues, inferredReturnShapes);
4749 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4750 if (inputShape.hasRank()) {
4751 outputShape[0] = inputShape.getDimSize(0);
4752 outputShape[3] = inputShape.getDimSize(3);
4754 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4758LogicalResult MaxPool2dOp::verify() {
4769LogicalResult MaxPool2dAdaptiveOp::verify() {
4774 AdaptivePoolingConstShapeValues values;
4778 values.pad, getInput(), getOutput())))
4784LogicalResult DepthwiseConv2DOp::inferReturnTypeComponents(
4785 MLIRContext *context, ::std::optional<Location> location,
4786 DepthwiseConv2DOp::Adaptor adaptor,
4787 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4788 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4790 int64_t inputWidth = ShapedType::kDynamic;
4791 int64_t inputHeight = ShapedType::kDynamic;
4792 int64_t inputChannels = ShapedType::kDynamic;
4794 int64_t weightWidth = ShapedType::kDynamic;
4795 int64_t weightHeight = ShapedType::kDynamic;
4796 int64_t depthChannels = ShapedType::kDynamic;
4799 ShapeAdaptor inputShape(adaptor.getInput().getType());
4800 if (inputShape.hasRank()) {
4801 outputShape[0] = inputShape.getDimSize(0);
4802 inputHeight = inputShape.getDimSize(1);
4803 inputWidth = inputShape.getDimSize(2);
4804 inputChannels = inputShape.getDimSize(3);
4808 ShapeAdaptor weightShape(adaptor.getWeight().getType());
4809 if (weightShape.hasRank()) {
4810 weightHeight = weightShape.getDimSize(0);
4811 weightWidth = weightShape.getDimSize(1);
4812 inputChannels = ShapedType::isDynamic(inputChannels)
4813 ? weightShape.getDimSize(2)
4815 depthChannels = weightShape.getDimSize(3);
4820 if (ShapedType::isStatic(inputChannels) &&
4821 ShapedType::isStatic(depthChannels)) {
4822 outputShape[3] = inputChannels * depthChannels;
4826 ShapeAdaptor biasShape(adaptor.getBias().getType());
4827 if (biasShape.hasRank() && ShapedType::isDynamic(outputShape[3])) {
4828 int64_t bc = biasShape.getDimSize(0);
4829 if (bc != ShapedType::kDynamic && bc != 1)
4830 outputShape[3] = bc;
4833 llvm::ArrayRef<int64_t> dilation = adaptor.getDilation();
4834 llvm::ArrayRef<int64_t> padding = adaptor.getPad();
4835 llvm::ArrayRef<int64_t> stride = adaptor.getStride();
4837 if (ShapedType::isStatic(inputHeight) && ShapedType::isStatic(weightHeight)) {
4838 int64_t inputSize = inputHeight + padding[0] + padding[1];
4839 int64_t filterSize = (weightHeight - 1) * dilation[0] + 1;
4840 int64_t unstridedResult = inputSize - filterSize + 1;
4841 outputShape[1] = (unstridedResult - 1) / stride[0] + 1;
4844 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(weightWidth)) {
4845 int64_t inputSize = inputWidth + padding[2] + padding[3];
4846 int64_t filterSize = (weightWidth - 1) * dilation[1] + 1;
4847 int64_t unstridedResult = inputSize - filterSize + 1;
4848 outputShape[2] = (unstridedResult - 1) / stride[1] + 1;
4851 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4855LogicalResult DepthwiseConv2DOp::verify() {
4862LogicalResult TransposeConv2DOp::inferReturnTypeComponents(
4863 MLIRContext *context, ::std::optional<Location> location,
4864 TransposeConv2DOp::Adaptor adaptor,
4865 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
4866 llvm::SmallVector<int64_t> outputShape(4, ShapedType::kDynamic);
4868 int64_t inputWidth = ShapedType::kDynamic;
4869 int64_t inputHeight = ShapedType::kDynamic;
4870 int64_t weightWidth = ShapedType::kDynamic;
4871 int64_t weightHeight = ShapedType::kDynamic;
4874 ShapeAdaptor inputShape(adaptor.getInput().getType());
4875 if (inputShape.hasRank()) {
4876 outputShape[0] = ShapedType::isDynamic(outputShape[0])
4877 ? inputShape.getDimSize(0)
4879 inputHeight = inputShape.getDimSize(1);
4880 inputWidth = inputShape.getDimSize(2);
4884 ShapeAdaptor weightShape(adaptor.getWeight().getType());
4885 if (weightShape.hasRank()) {
4886 outputShape[3] = ShapedType::isDynamic(outputShape[3])
4887 ? weightShape.getDimSize(0)
4889 weightHeight = weightShape.getDimSize(1);
4890 weightWidth = weightShape.getDimSize(2);
4894 ShapeAdaptor biasShape(adaptor.getBias().getType());
4895 if (biasShape.hasRank() && ShapedType::isDynamic(outputShape[3])) {
4896 int64_t bc = biasShape.getDimSize(0);
4897 if (bc != ShapedType::kDynamic && bc != 1)
4898 outputShape[3] = bc;
4901 llvm::ArrayRef<int64_t> padding = adaptor.getOutPad();
4902 llvm::ArrayRef<int64_t> stride = adaptor.getStride();
4904 if (ShapedType::isStatic(inputHeight) && ShapedType::isStatic(weightHeight)) {
4905 int64_t calculateSize =
4906 (inputHeight - 1) * stride[0] + padding[0] + padding[1] + weightHeight;
4908 ShapedType::isDynamic(outputShape[1]) ? calculateSize : outputShape[1];
4911 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(weightWidth)) {
4912 int64_t calculateSize =
4913 (inputWidth - 1) * stride[1] + padding[2] + padding[3] + weightWidth;
4915 ShapedType::isDynamic(outputShape[2]) ? calculateSize : outputShape[2];
4918 inferredReturnShapes.push_back(ShapedTypeComponents(outputShape));
4922LogicalResult TransposeConv2DOp::verify() {
4926 const llvm::ArrayRef<int64_t> strides = getStride();
4927 const int64_t strideY = strides[0];
4928 const int64_t strideX = strides[1];
4930 if (strideY < 1 || strideX < 1)
4931 return emitOpError(
"expect all stride values to be >= 1, got [")
4934 const auto checkPadAgainstKernelDim =
4935 [
this](int64_t padValue, int64_t kernelDimSize, llvm::StringRef padName,
4936 llvm::StringRef kernelDimName) -> LogicalResult {
4937 if (padValue <= -kernelDimSize)
4938 return emitOpError(
"expected ")
4939 << padName <<
" > -" << kernelDimName <<
", but got: " << padName
4940 <<
"=" << padValue <<
" and " << kernelDimName <<
"="
4945 const llvm::ArrayRef<int64_t> padding = getOutPad();
4946 const int64_t outPadTop = padding[0];
4947 const int64_t outPadBottom = padding[1];
4948 const int64_t outPadLeft = padding[2];
4949 const int64_t outPadRight = padding[3];
4951 const auto weightType =
4952 llvm::dyn_cast<RankedTensorType>(getWeight().
getType());
4955 const int64_t kernelHeight = weightType.getDimSize(1);
4956 if (ShapedType::isStatic(kernelHeight)) {
4957 if (
failed(checkPadAgainstKernelDim(outPadTop, kernelHeight,
4958 "out_pad_top",
"KH")))
4961 if (
failed(checkPadAgainstKernelDim(outPadBottom, kernelHeight,
4962 "out_pad_bottom",
"KH")))
4966 const int64_t kernelWidth = weightType.getDimSize(2);
4967 if (ShapedType::isStatic(kernelWidth)) {
4968 if (
failed(checkPadAgainstKernelDim(outPadLeft, kernelWidth,
4969 "out_pad_left",
"KW")))
4972 if (
failed(checkPadAgainstKernelDim(outPadRight, kernelWidth,
4973 "out_pad_right",
"KW")))
4979 const auto outputType =
4980 llvm::dyn_cast<RankedTensorType>(getOutput().
getType());
4984 const auto inputType = llvm::dyn_cast<RankedTensorType>(getInput().
getType());
4985 if (inputType && weightType) {
4986 const int64_t inputHeight = inputType.getDimSize(1);
4987 const int64_t kernelHeight = weightType.getDimSize(1);
4988 const int64_t outputHeight = outputType.getDimSize(1);
4990 if (ShapedType::isStatic(inputHeight) &&
4991 ShapedType::isStatic(outputHeight)) {
4993 (inputHeight - 1) * strideY + outPadTop + outPadBottom + kernelHeight)
4995 "dimension mismatch: expected OH == (IH - 1) * stride_y "
4996 "+ out_pad_top + out_pad_bottom + KH, but got ")
4997 << outputHeight <<
" != (" << inputHeight <<
" - 1) * "
4998 << strideY <<
" + " << outPadTop <<
" + " << outPadBottom
4999 <<
" + " << kernelHeight;
5002 const int64_t inputWidth = inputType.getDimSize(2);
5003 const int64_t kernelWidth = weightType.getDimSize(2);
5004 const int64_t outputWidth = outputType.getDimSize(2);
5006 if (ShapedType::isStatic(inputWidth) && ShapedType::isStatic(outputWidth)) {
5008 (inputWidth - 1) * strideX + outPadLeft + outPadRight + kernelWidth)
5010 "dimension mismatch: expected OW == (IW - 1) * stride_x "
5011 "+ out_pad_left + out_pad_right + KW, but got ")
5012 << outputWidth <<
" != (" << inputWidth <<
" - 1) * " << strideX
5013 <<
" + " << outPadLeft <<
" + " << outPadRight <<
" + "
5018 const auto biasType = llvm::dyn_cast<RankedTensorType>(getBias().
getType());
5023 const int64_t biasChannels = biasType.getDimSize(0);
5026 if (biasChannels == ShapedType::kDynamic)
5029 const int64_t outputChannels = outputType.getDimSize(3);
5030 if (!ShapedType::isDynamic(outputChannels) &&
5031 biasChannels != outputChannels && biasChannels != 1)
5033 "bias channels expected to be equal to output channels (")
5034 << outputChannels <<
") or 1, got " << biasChannels;
5039LogicalResult RescaleOp::verify() {
5040 const auto inputType = llvm::cast<ShapedType>(getInput().
getType());
5041 auto inputElementType =
5043 if (!mlir::isa<IntegerType>(inputElementType)) {
5044 emitOpError(
"expect input to have integer element type, got ")
5045 << inputElementType;
5049 const auto outputType = llvm::cast<ShapedType>(getOutput().
getType());
5050 auto outputElementType =
5052 if (!mlir::isa<IntegerType>(outputElementType)) {
5053 emitOpError(
"expect output to have integer element type, got ")
5054 << outputElementType;
5066 FailureOr<int64_t> maybeIZp = getInputZeroPoint();
5067 if (succeeded(maybeIZp) && verifyInputZeroPoint(*maybeIZp).failed())
5070 FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
5071 if (succeeded(maybeOZp) && verifyOutputZeroPoint(*maybeOZp).failed())
5074 const auto multiplierType = llvm::cast<ShapedType>(getMultiplier().
getType());
5076 if (getScale32() && !multiplierType.getElementType().isInteger(32)) {
5077 emitOpError(
"expect i32 element type for multiplier for scale32=true, got ")
5078 << multiplierType.getElementType();
5083 if (!getScale32() && !multiplierType.getElementType().isInteger(16)) {
5085 "expect i16 element type for multiplier for scale32=false, got ")
5086 << multiplierType.getElementType();
5090 if (!inputType.hasRank())
5096 int64_t numChannels = 1;
5097 if (getPerChannel()) {
5098 if (inputType.getRank() < 1) {
5099 emitOpError(
"requires input to be at least rank 1 when per_channel is "
5100 "true, but got rank ")
5101 << inputType.getRank();
5104 numChannels = inputType.getDimSize(inputType.getRank() - 1);
5107 if (outputType.hasRank()) {
5109 getOperation(), outputType, inputType.getShape())))
5113 if (multiplierType.hasRank()) {
5114 ArrayRef<int64_t> multiplierShape = multiplierType.getShape();
5116 if (multiplierShape[0] != ShapedType::kDynamic &&
5117 multiplierShape[0] != numChannels) {
5118 emitOpError(
"expect shape of { ")
5119 << numChannels <<
" } for multiplier input, got { "
5120 << multiplierShape[0] <<
" }";
5125 const auto shiftType = llvm::cast<ShapedType>(getShift().
getType());
5126 if (shiftType.hasRank()) {
5127 ArrayRef<int64_t> shiftShape = shiftType.getShape();
5129 if (shiftShape[0] != ShapedType::kDynamic && shiftShape[0] != numChannels) {
5130 emitOpError(
"expect shape of { ")
5131 << numChannels <<
" } for shift input, got { " << shiftShape[0]
5140LogicalResult RescaleOp::inferReturnTypeComponents(
5141 MLIRContext *context, ::std::optional<Location> location,
5142 RescaleOp::Adaptor adaptor,
5143 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5144 ShapeAdaptor inputShape(adaptor.getInput().getType());
5145 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5149LogicalResult CastOp::verify() {
5150 const ShapedType inputType = llvm::cast<ShapedType>(getInput().
getType());
5151 const ShapedType outputType = llvm::cast<ShapedType>(
getType());
5152 const Type inputElementType = inputType.getElementType();
5153 const Type outputElementType = outputType.getElementType();
5155 const bool inputIsBlockScaled = llvm::isa<BlockScaledType>(inputElementType);
5156 const bool outputIsBlockScaled =
5157 llvm::isa<BlockScaledType>(outputElementType);
5159 const bool isUnsigned = this->getInputUnsigned();
5164 return emitOpError()
5165 <<
"attribute input_unsigned requires integer type inputs. Got: "
5168 if (!inputIsBlockScaled && !outputIsBlockScaled)
5171 if (inputIsBlockScaled && outputIsBlockScaled)
5172 return emitOpError()
5173 <<
"requires exactly one of input or output to have block scaled "
5176 const Type scalarElementType =
5177 inputIsBlockScaled ? outputElementType : inputElementType;
5178 if (!llvm::isa<FloatType>(scalarElementType))
5179 return emitOpError()
5180 <<
"requires non-block-scaled element type to be floating-point "
5181 "when casting to or from block scaled element type, got "
5182 << scalarElementType;
5187LogicalResult CastFromBlockScaledOp::inferReturnTypeComponents(
5188 MLIRContext *context, ::std::optional<Location> location,
5189 CastFromBlockScaledOp::Adaptor adaptor,
5190 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5191 const ShapeAdaptor inputShape(adaptor.getInputData().getType());
5192 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5196LogicalResult CastFromBlockScaledOp::verify() {
5197 const Type inputDataType = getInputData().getType();
5198 const Type outputDataType = getResult().getType();
5200 return emitOpError() <<
"require compatible shapes for input_data ("
5201 << inputDataType <<
") and " <<
"output_data ("
5202 << outputDataType <<
")";
5204 const ShapeAdaptor inputDataShape = ShapeAdaptor(inputDataType);
5206 if (inputDataShape.
hasRank()) {
5207 const unsigned int blockSize =
5209 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
5210 return emitOpError(
"expect block size to be 32, got ") << blockSize;
5211 const int64_t inputDataLastDim =
5213 if (inputDataLastDim % blockSize != 0)
5214 return emitOpError() <<
"expect last dimension of input_data ("
5216 <<
") to be divisible by block_size (" << blockSize
5219 const Type inputScaleType = getInputScale().getType();
5220 const ShapeAdaptor inputScaleShape = ShapeAdaptor(inputScaleType);
5222 if (inputScaleShape.
hasRank()) {
5223 SmallVector<int64_t> inputDataDims, inputScaleDims;
5224 inputDataShape.
getDims(inputDataDims);
5225 inputScaleShape.
getDims(inputScaleDims);
5227 if (inputDataDims.size() != inputScaleDims.size() ||
5229 ArrayRef<int64_t>(inputDataDims).drop_back(1),
5230 ArrayRef<int64_t>(inputScaleDims).drop_back(1))))
5231 return emitOpError()
5232 <<
"require compatible shapes for input_data (" << inputDataType
5233 <<
") and " <<
"input_scale (" << inputScaleType
5234 <<
") except for the last dimension";
5236 const SmallVector<int64_t, 2> dimsToCheck{inputDataLastDim / blockSize,
5237 inputScaleDims.back()};
5238 if (ShapedType::isStatic(inputDataLastDim) &&
5240 return emitOpError()
5241 <<
"expect last dimension of input_scale ("
5242 << inputScaleDims.back()
5243 <<
") to be equal to last dimension of input_data / block_size ("
5244 << inputDataDims.back() / blockSize <<
")";
5251LogicalResult CastToBlockScaledOp::inferReturnTypeComponents(
5252 MLIRContext *context, ::std::optional<Location> location,
5253 CastToBlockScaledOp::Adaptor adaptor,
5254 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5255 const ShapeAdaptor inputShape(adaptor.getInputData().getType());
5256 inferredReturnShapes.push_back(ShapedTypeComponents(inputShape));
5257 if (!inputShape.hasRank())
5261 SmallVector<int64_t> outputScaleShape;
5262 inputShape.getDims(outputScaleShape);
5263 const int64_t lastDimLoc = inputShape.getRank() - 1;
5264 const int64_t lastDimSize = inputShape.getDimSize(lastDimLoc);
5265 if (ShapedType::isStatic(lastDimSize)) {
5266 const unsigned int blockSize =
5267 BlockSizeAttr::getBlockSizeValue(adaptor.getBlockSize());
5268 outputScaleShape[lastDimLoc] = lastDimSize / blockSize;
5270 inferredReturnShapes.push_back(ShapedTypeComponents(outputScaleShape));
5274LogicalResult CastToBlockScaledOp::verify() {
5275 const Type inputDataType = getInputData().getType();
5276 const Type outputDataType = getResult(0).getType();
5278 return emitOpError() <<
"require compatible shapes for input_data ("
5279 << inputDataType <<
") and " <<
"output_data ("
5280 << outputDataType <<
")";
5282 const unsigned int blockSize =
5284 if (blockSize != BlockSizeAttr::getBlockSizeValue(BlockSize::BLOCK_SIZE_32))
5285 return emitOpError(
"expect block size to be 32, got ") << blockSize;
5286 const ShapeAdaptor inputDataShape = ShapeAdaptor(inputDataType);
5287 if (inputDataShape.
hasRank()) {
5288 const int64_t inputDataLastDim =
5290 if (ShapedType::isStatic(inputDataLastDim) &&
5291 inputDataLastDim % blockSize != 0)
5292 return emitOpError() <<
"expect last dimension of input_data ("
5294 <<
") to be divisible by block_size (" << blockSize
5298 const ShapeAdaptor outputDataShape = ShapeAdaptor(outputDataType);
5299 const Type outputScaleType = getResult(1).getType();
5300 const ShapeAdaptor outputScaleShape = ShapeAdaptor(outputScaleType);
5302 SmallVector<int64_t> outputDataDims, outputScaleDims;
5303 outputDataShape.
getDims(outputDataDims);
5304 outputScaleShape.
getDims(outputScaleDims);
5306 if (outputDataDims.size() != outputScaleDims.size() ||
5308 ArrayRef<int64_t>(outputDataDims).drop_back(1),
5309 ArrayRef<int64_t>(outputScaleDims).drop_back(1))))
5310 return emitOpError() <<
"require compatible shapes for output_data ("
5311 << outputDataType <<
") and " <<
"output_scale ("
5313 <<
") except for the last dimension";
5315 const int64_t outputDataLastDim = outputDataDims.back();
5316 const SmallVector<int64_t, 2> dimsToCheck{outputDataLastDim / blockSize,
5317 outputScaleDims.back()};
5318 if (ShapedType::isStatic(outputDataLastDim) &&
5320 return emitOpError()
5321 <<
"expect last dimension of output_scale ("
5322 << outputScaleDims.back()
5323 <<
") to be equal to last dimension of output_data / block_size ("
5324 << outputDataDims.back() / blockSize <<
")";
5330LogicalResult IfOp::inferReturnTypeComponents(
5331 MLIRContext *context, ::std::optional<Location> location,
5332 IfOp::Adaptor adaptor,
5333 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5334 llvm::SmallVector<tosa::YieldOp> yieldOps;
5335 for (Region *region : adaptor.getRegions()) {
5336 for (
auto &block : *region)
5337 if (
auto returnOp = dyn_cast<tosa::YieldOp>(block.getTerminator()))
5338 yieldOps.push_back(returnOp);
5341 if (yieldOps.empty())
5345 llvm::SmallVector<ValueKnowledge> resultKnowledge;
5346 resultKnowledge.reserve(yieldOps.front().getNumOperands());
5347 for (
auto operand : yieldOps.front().getOperands()) {
5348 resultKnowledge.push_back(
5352 for (
auto yieldOp : yieldOps) {
5353 if (resultKnowledge.size() != yieldOp.getNumOperands())
5356 for (
const auto &it : llvm::enumerate(yieldOp.getOperands())) {
5357 int32_t index = it.index();
5359 resultKnowledge[index],
5363 resultKnowledge[index] = meet;
5367 for (
const ValueKnowledge &
result : resultKnowledge) {
5368 inferredReturnShapes.push_back(
result.getShapedTypeComponents());
5374LogicalResult WhileOp::inferReturnTypeComponents(
5375 MLIRContext *context, ::std::optional<Location> location,
5376 WhileOp::Adaptor adaptor,
5377 SmallVectorImpl<ShapedTypeComponents> &inferredReturnShapes) {
5378 llvm::SmallVector<tosa::YieldOp> yieldOps;
5379 for (
auto &block : adaptor.getBodyGraph())
5380 if (
auto returnOp = dyn_cast<tosa::YieldOp>(block.getTerminator()))
5381 yieldOps.push_back(returnOp);
5385 if (yieldOps.empty())
5389 llvm::SmallVector<ValueKnowledge> resultKnowledge;
5390 resultKnowledge.reserve(yieldOps.front().getNumOperands());
5391 for (
auto operand : yieldOps.front().getOperands()) {
5392 resultKnowledge.push_back(
5396 for (
auto yieldOp : yieldOps) {
5397 if (resultKnowledge.size() != yieldOp.getNumOperands())
5400 for (
const auto &it : llvm::enumerate(yieldOp.getOperands())) {
5401 int32_t index = it.index();
5403 resultKnowledge[index],
5405 resultKnowledge[index] = meet;
5410 for (
const ValueKnowledge &
result : resultKnowledge) {
5411 inferredReturnShapes.push_back(
result.getShapedTypeComponents());
5417std::optional<SmallVector<int64_t, 4>> ApplyScaleOp::getShapeForUnroll() {
5418 if (
auto vt = llvm::dyn_cast<VectorType>(
getType()))
5419 return llvm::to_vector<4>(vt.getShape());
5420 return std::nullopt;
5426 StringRef prefix =
"") {
5427 assert(blocksArgs.size() == initializers.size() &&
5428 "expected same length of arguments and initializers");
5429 if (initializers.empty())
5432 parser << prefix <<
'(';
5433 llvm::interleaveComma(
5434 llvm::zip(blocksArgs, initializers), parser,
5435 [&](
auto it) { parser << std::get<0>(it) <<
" = " << std::get<1>(it); });
5440ParseResult IfOp::parse(OpAsmParser &parser, OperationState &
result) {
5442 result.regions.reserve(2);
5443 Region *thenRegion =
result.addRegion();
5444 Region *elseRegion =
result.addRegion();
5446 OpAsmParser::UnresolvedOperand cond;
5451 SmallVector<OpAsmParser::Argument, 4> regionArgs;
5452 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
5455 OptionalParseResult listResult =
5463 "expected type for condition operand");
5469 "expected type for condition operand");
5477 FunctionType functionType;
5481 <<
"expected list of types for block arguments "
5482 <<
"followed by arrow type and list of return types";
5484 result.addTypes(functionType.getResults());
5486 if (functionType.getNumInputs() != operands.size()) {
5488 <<
"expected as many input types as operands " <<
"(expected "
5489 << operands.size() <<
" got " << functionType.getNumInputs()
5520void IfOp::print(OpAsmPrinter &p) {
5521 p <<
" " << getCondition();
5524 getInputList(),
" ");
5526 p << getCondition().getType();
5528 if (!getInputList().empty()) {
5530 llvm::interleaveComma(getInputList().getTypes(), p);
5539 auto &elseRegion = getElseGraph();
5540 if (!elseRegion.
empty()) {
5548LogicalResult IfOp::verify() {
5550 "'then_graph' arguments", getInputList(),
5556 "'else_graph' arguments", getInputList(),
5562 if (getThenGraph().front().mightHaveTerminator()) {
5564 dyn_cast<tosa::YieldOp>(getThenGraph().front().getTerminator());
5566 *
this, thenYield.getInputs(),
"'then_graph' results",
5567 getOutputList(),
"'output_list'")
5573 if (getElseGraph().front().mightHaveTerminator()) {
5575 dyn_cast<tosa::YieldOp>(getElseGraph().front().getTerminator());
5577 *
this, elseYield.getInputs(),
"'else_graph' results",
5578 getOutputList(),
"'output_list'")
5583 auto condType = getCondition().getType();
5585 return emitOpError() <<
"'condition' must be a size 1 tensor, got "
5591LogicalResult WhileOp::verify() {
5593 getOutputList(),
"'output_list'")
5598 "'cond_graph' arguments", getInputList(),
5604 "'body_graph' arguments", getInputList(),
5609 if (getBodyGraph().front().mightHaveTerminator()) {
5611 dyn_cast<tosa::YieldOp>(getBodyGraph().front().getTerminator());
5613 "'body_graph' results",
5614 getInputList(),
"'input_list'")
5621 if (!getCondGraph().front().mightHaveTerminator())
5625 dyn_cast<tosa::YieldOp>(getCondGraph().front().getTerminator());
5629 if (condYield.getInputs().size() != 1)
5630 return emitOpError() <<
"require 'cond_graph' only have one result";
5632 auto condOutType = condYield.getInputs()[0].getType();
5634 return emitOpError() <<
"'cond_graph' result must be a size 1 tensor, got "
5638 return emitOpError() <<
"'cond_graph' result must be a boolean tensor, got "
5644LogicalResult ReverseOp::verify() {
5645 TensorType inputType = getInput1().getType();
5646 int32_t reverseAxis = getAxis();
5648 if (reverseAxis < 0)
5649 return emitOpError(
"expected non-negative reverse axis");
5651 int64_t inputRank = inputType.getRank();
5654 if (reverseAxis >= inputRank && (reverseAxis != 0 || inputRank != 0))
5655 return emitOpError(
"expect input tensor rank (")
5656 << inputRank <<
") to be larger than reverse axis (" << reverseAxis
5663LogicalResult tosa::SelectOp::verify() {
5674 auto predicateType = llvm::dyn_cast<ShapedType>(getPred().
getType());
5675 if (!predicateType) {
5676 return emitOpError(
"expect shaped tensor for input1, got ")
5677 << getInput1().getType();
5679 auto predicateElementType = predicateType.getElementType();
5680 if (!predicateElementType.isInteger(1)) {
5681 return emitOpError(
"expect element type of bool for input1, got ")
5682 << predicateElementType;
5688LogicalResult tosa::VariableReadOp::verify() {
5696LogicalResult tosa::VariableWriteOp::verify() {
5705ParseResult WhileOp::parse(OpAsmParser &parser, OperationState &
result) {
5706 SmallVector<OpAsmParser::Argument, 4> regionArgs;
5707 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
5708 Region *cond =
result.addRegion();
5709 Region *body =
result.addRegion();
5711 OptionalParseResult listResult =
5716 FunctionType functionType;
5721 result.addTypes(functionType.getResults());
5723 if (functionType.getNumInputs() != operands.size()) {
5725 <<
"expected as many input types as operands " <<
"(expected "
5726 << operands.size() <<
" got " << functionType.getNumInputs() <<
")";
5736 for (
size_t i = 0, e = regionArgs.size(); i != e; ++i)
5737 regionArgs[i].type = functionType.getInput(i);
5739 return failure(parser.
parseRegion(*cond, regionArgs) ||
5744void WhileOp::print(OpAsmPrinter &parser) {
5746 getInputList(),
" ");
5749 getResults().getTypes());
5755 (*this)->getDiscardableAttrDictionary().getValue());
5764 auto zpType = mlir::RankedTensorType::get({1}, srcElemType);
5765 if (llvm::isa<FloatType>(srcElemType)) {
5767 zpType, builder.
getFloatAttr(srcElemType,
static_cast<double>(zp)));
5768 return tosa::ConstOp::create(builder, loc, zpType, zpAttr);
5770 if (llvm::isa<IntegerType>(srcElemType)) {
5773 return tosa::ConstOp::create(builder, loc, zpType, zpAttr);
5775 llvm::errs() <<
"zero point is not allowed for unsupported data types\n";
5776 return std::nullopt;
5784 return mlir::isa<tosa::shapeType>(t);
5791 return emitError() <<
"invalid rank (must be >= 0): " << rank;
5797 if (mlir::isa<::mlir::tosa::shapeType>(v.getType())) {
5798 Operation *definingOp = v.getDefiningOp();
5800 return op->
emitOpError(
"shape operand is not compile time resolvable");
5813 auto getRank = [](
const Type type) {
5814 return mlir::cast<mlir::tosa::shapeType>(type).getRank();
5820 for (
auto type : operandTypes) {
5821 if (getRank(type) != rank) {
5822 return op->
emitOpError(
"operands don't have matching ranks");
5825 for (
auto type : resultTypes) {
5826 if (getRank(type) != rank) {
5827 return op->
emitOpError(
"result shape has different rank than operands");
5837LogicalResult tosa::ConstShapeOp::verify() {
5839 auto valuesRank = getValues().getType().getRank();
5840 if (valuesRank != 1)
5841 return emitOpError(
"expect elements in attribute values with rank 1");
5843 auto count = getValues().getNumElements();
5844 auto rank = (cast<tosa::shapeType>(getResult().
getType())).getRank();
5845 if (count != rank && (count != 1 || rank != 0)) {
5846 return emitOpError(
"expect number of elements in attribute values (")
5847 << count <<
") to be equal to the rank (" << rank
5848 <<
") for the result shape type";
5853LogicalResult tosa::DimOp::verify() {
5854 const tosa::shapeType outShapeType =
5855 cast<tosa::shapeType>(getResult().
getType());
5856 if (outShapeType.getRank() != 1)
5857 return emitOpError(
"expect output shape type to contain one element, got ")
5862 const int64_t inputRank = inputType.getRank();
5863 const int64_t axis = getAxisAttr().getInt();
5864 if (axis < 0 || axis >= inputRank)
5865 return emitOpError(
"expect axis to be in the range [0, ")
5866 << inputRank <<
"), got " << axis;
5871LogicalResult tosa::ConcatShapeOp::verify() {
5872 const tosa::shapeType outShapeType =
5873 cast<tosa::shapeType>(getResult().
getType());
5874 const int64_t outputRank = outShapeType.getRank();
5877 if (inputList.size() == 0)
5878 return emitOpError(
"requires at least one input shape");
5880 if (llvm::any_of(inputList, [](Value v) {
5881 return cast<tosa::shapeType>(v.
getType()).getRank() == 0;
5883 return emitOpError(
"requires all inputs shapes have a rank greater than 0");
5885 const int64_t inputsRank =
5886 llvm::accumulate(inputList, 0, [](int64_t acc,
const Value &input) {
5887 const tosa::shapeType inShapeType =
5888 cast<tosa::shapeType>(input.
getType());
5889 return acc + inShapeType.getRank();
5891 if (outputRank != inputsRank)
5892 return emitOpError(
"requires output shape rank to be equal to the sum of "
5893 "the input shape ranks (")
5894 << inputsRank <<
"), got " << outputRank;
5899LogicalResult tosa::SliceShapeOp::verify() {
5900 std::optional<int32_t> start;
5901 DenseIntElementsAttr startAttr;
5903 start = startAttr.getValues<int32_t>()[0];
5904 if (start && start.value() < 0)
5905 return emitOpError(
"expected non-negative start index, got ")
5908 std::optional<int32_t> size;
5909 DenseIntElementsAttr sizeAttr;
5911 size = sizeAttr.getValues<int32_t>()[0];
5912 if (size && size.value() <= 0)
5913 return emitOpError(
"expected positive size, got ") << size.value();
5918 const tosa::shapeType outShapeType =
5919 cast<tosa::shapeType>(getResult().
getType());
5920 const int64_t outputRank = outShapeType.getRank();
5921 if (outputRank != size)
5923 "expected output type size to be equal to size attribute, got ")
5924 << outputRank <<
" vs " << size.value();
5929 const tosa::shapeType inShapeType =
5930 cast<tosa::shapeType>(getInput().
getType());
5931 const int64_t inputRank = inShapeType.getRank();
5932 const int64_t sliceSize = start.value() + size.value();
5933 if (sliceSize > inputRank)
5934 return emitOpError(
"expected start + size to be less than or equal to "
5935 "input shape rank (")
5936 << inputRank <<
"), got " << sliceSize;
5945#define GET_ATTRDEF_CLASSES
5946#include "mlir/Dialect/Tosa/IR/TosaAttributes.cpp.inc"
5951#define GET_TYPEDEF_CLASSES
5952#include "mlir/Dialect/Tosa/IR/TosaOpsTypesBase.cpp.inc"
5973 printer << keyword <<
'(';
5995#define GET_OP_CLASSES
5996#include "mlir/Dialect/Tosa/IR/TosaOps.cpp.inc"
static void printInitializationList(OpAsmPrinter &p, Block::BlockArgListType blocksArgs, ValueRange initializers, StringRef prefix="")
Prints the initialization list in the form of <prefix>(inner = outer, inner2 = outer2,...
true
Given two iterators into the same block, return "true" if a is before `b.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
static Type getValueType(Attribute attr)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
static LogicalResult verifyMatMulShapes(T op, bool transposeB)
static ParseResult parseOptionalBoolClause(OpAsmParser &parser, StringRef keyword, BoolAttr &result)
static void printShapeToDiagnostic(InFlightDiagnostic &diag, ArrayRef< int64_t > shape)
static void buildMatMulOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
static LogicalResult verifySameElementTypes(Operation *op, Type aType, Type bType, StringRef aName="input", StringRef bName="output")
static ParseResult parseLocalBound(OpAsmParser &parser, BoolAttr &result)
LogicalResult inferConvReturnTypeComponents(AdaptorT adaptor, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
static int64_t getMatMulBatchDim(const ShapeAdaptor &shape, int64_t outputRank, int64_t axis)
static SmallVector< int64_t > convertToMlirShape(ArrayRef< int64_t > shape)
static LogicalResult ReduceInferReturnTypes(ShapeAdaptor operandShape, Type inputType, IntegerAttr axis, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
static void printScaleValues(AsmPrinter &printer, ArrayRef< Attribute > scaleValues, Type)
static void buildAvgPool2dAdaptiveOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, DenseI64ArrayAttr kernel, DenseI64ArrayAttr stride, DenseI64ArrayAttr pad, TypeAttr accType)
This builder mirrors avg_pool2d quant-info handling and materializes kernel/stride/pad as const_shape...
static LogicalResult verifyRescaleValueAndZpTypes(Operation *op, Value val, Value valZp, StringRef name)
static void printOptionalBoolClause(OpAsmPrinter &printer, StringRef keyword, BoolAttr attr)
static LogicalResult errorIfShapeNotSizeOne(Operation *op, Type type)
static void printLocalBound(OpAsmPrinter &printer, Operation *, BoolAttr attr)
LogicalResult argMaxMinVerify(T op)
static LogicalResult verifyMatMulZeroPointType(T op, Value input, Value zp, StringRef inputName, StringRef zpName)
static ParseResult parseScaleValues(AsmParser &parser, SmallVector< Attribute > &scaleValues, Type scaleType)
static ParseResult parseInputUnsigned(OpAsmParser &parser, BoolAttr &result)
#define REDUCE_SHAPE_INFER(OP)
static LogicalResult verifyConvOp(T op)
static LogicalResult verifyAvgPoolCommonTypeAndZpChecks(T op)
static LogicalResult verifyVariableOpErrorIf(T op, Type type, StringRef name)
static LogicalResult poolingInferReturnTypes(ShapeAdaptor inputShape, ArrayRef< int64_t > kernel, ArrayRef< int64_t > stride, ArrayRef< int64_t > pad, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
static void buildPadOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value paddings)
This builder is called on TOSA pad operator that needs to create its own OptionalAttr quantization_at...
static LogicalResult verifyPoolingOpImpl(Operation *op, ArrayRef< int64_t > kernel, ArrayRef< int64_t > strides, ArrayRef< int64_t > padding, Value input, Value output)
static std::optional< int64_t > idivCheck(const int64_t lhs, const int64_t rhs)
static void buildVariableOp(OpBuilder &builder, OperationState &result, StringRef name, Type variableType, Attribute initialValue)
static void buildMatMulLikeOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
LogicalResult verifyConvOutputSize(Operation *op, const int64_t inputSize, const int64_t kernelSize, const int64_t outputSize, const int64_t padBefore, const int64_t padAfter, const int64_t stride, const int64_t dilation, const llvm::StringRef dimName, const llvm::StringRef dimAxis, const llvm::StringRef padBeforeName, const llvm::StringRef padAfterName)
static LogicalResult verifyReduceOp(T op)
#define NARY_SHAPE_INFER(OP)
#define ZERO_POINT_HELPER(OP, OPERAND_NAME, SIGN_EXTEND)
static void buildTransConvOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value weight, Value bias, DenseI64ArrayAttr outpad, DenseI64ArrayAttr stride, TypeAttr accType)
Handles tosa.transpose_conv2d which has outpad and output shape attributes.
LogicalResult inferArgMaxMinReturnTypeComponents(MLIRContext *context, ::std::optional< Location > location, A adaptor, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
static void extractAdaptivePoolingConstShapeOperands(T op, AdaptivePoolingConstShapeValues &values)
static LogicalResult verifyConvOpErrorIf(T op)
static FailureOr< int64_t > getZeroPoint(Value val, bool signExtend)
static constexpr bool IsSupportedAdaptivePoolConstShapeVerifyOp
LogicalResult tryUpdateDimOrFailure(Operation *op, int64_t &currDim, const int64_t newDim, const StringRef operandName, const StringRef dimName)
static LogicalResult verifyConvOpModes(T op)
static LogicalResult NAryInferReturnTypes(const ValueShapeRange &operands, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
#define COMPATIBLE_RETURN_TYPES(OP)
static LogicalResult resolveBroadcastShape(const ValueShapeRange &operands, SmallVector< int64_t > &outShape)
static LogicalResult verifyMatMulQuantizedOperandsType(T op, Type aElementType, Type bElementType)
static LogicalResult verifyOutputShapeCompatibleWithExpected(Operation *op, ShapedType outputType, ArrayRef< int64_t > expectedShape, StringRef outputName="output")
static void printInputUnsigned(OpAsmPrinter &printer, Operation *, BoolAttr attr)
static void buildNegateOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input)
This builder is called on single-parameter negate operator to construct input and output zero points ...
static void buildConvOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, Value weight, Value bias, DenseI64ArrayAttr pad, DenseI64ArrayAttr stride, DenseI64ArrayAttr dilation, TypeAttr accType)
This builder is called on all convolution operators except TransposeConv, which has specialized outpu...
static SmallVector< int64_t > getMatMulBatchShape(const ShapeAdaptor &shape)
static void buildAvgPool2dOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value input, DenseArrayAttr kernel, DenseArrayAttr stride, DenseArrayAttr pad, TypeAttr accType)
Both the tosa.avg_pool2d and unary ops use the same UnaryOpQuantizationAttr but avg_pool operator has...
static LogicalResult errorIfTypeOrShapeMismatch(Operation *op, Type type1, StringRef name1, Type type2, StringRef name2)
static LogicalResult inferMatMulReturnTypeComponents(const ShapeAdaptor &aShape, const ShapeAdaptor &bShape, bool transposeB, SmallVectorImpl< ShapedTypeComponents > &inferredReturnShapes)
static void buildMatMulTOpWithQuantInfo(OpBuilder &builder, OperationState &result, Type outputType, Value a, Value b)
static FailureOr< SmallVector< int64_t > > resolveMatMulOutputShape(const ShapeAdaptor &aShape, const ShapeAdaptor &bShape, int64_t outputRank, bool transposeB)
static FailureOr< int64_t > resolveBroadcastDim(const int64_t dim1, const int64_t dim2)
static LogicalResult verifyZeroPoint(T op, Value val, const int64_t &zp, const std::string &operand)
static LogicalResult verifyPoolingOp(T op)
static LogicalResult verifyDimIsPowerOfTwo(Operation *op, const int64_t dimSize, const llvm::StringRef dimName)
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
static void updateIfDynamic(int64_t ¤t, int64_t candidate)
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
ConvInferShapeAdaptor(Conv2DBlockScaledOp::Adaptor adaptor)
int64_t getOutputRank() const
int64_t getNumSpatialDims() const
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
ConvInferShapeAdaptor(Conv2DOp::Adaptor adaptor)
int64_t getNumSpatialDims() const
int64_t getOutputRank() const
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
int64_t getNumSpatialDims() const
void inferWeightShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &weightSpatial)
int64_t getOutputRank() const
ConvInferShapeAdaptor(Conv3DOp::Adaptor adaptor)
void inferInputShape(SmallVectorImpl< int64_t > &outputShape, SmallVectorImpl< int64_t > &inputSpatial)
LogicalResult getSpatialParameters(SmallVector< int64_t > &padValues, SmallVector< int64_t > &strideValues, SmallVector< int64_t > &dilationValues)
This base class exposes generic asm parser hooks, usable across the various derived parsers.
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseOptionalColon()=0
Parse a : token if present.
virtual ParseResult parseOptionalAttrDictWithKeyword(NamedAttrList &result)=0
Parse a named dictionary into 'result' if the attributes keyword is present.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseColon()=0
Parse a : token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
virtual ParseResult parseFloat(double &result)=0
Parse a floating point value from the stream.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
This base class exposes generic asm printer hooks, usable across the various derived printers.
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
virtual void printAttribute(Attribute attr)
void printArrowTypeList(TypeRange &&types)
Attributes are known-constant values of operations.
MutableArrayRef< BlockArgument > BlockArgListType
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getIndexAttr(int64_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
FloatAttr getFloatAttr(Type type, double value)
IntegerType getIntegerType(unsigned width)
StringAttr getStringAttr(const Twine &bytes)
DenseIntElementsAttr getIndexTensorAttr(ArrayRef< int64_t > values)
An attribute that represents a reference to a dense vector or tensor object.
auto getValues() const
Return the held element values as a range of the given type.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
An attribute that represents a reference to a dense integer vector or tensor object.
This class contains all of the information necessary to report a diagnostic to the DiagnosticEngine.
virtual InFlightDiagnostic emitError(const Twine &msg={}) const =0
Emit an error to the reader.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
This class represents a diagnostic that is inflight and set to be reported.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual OptionalParseResult parseOptionalAssignmentList(SmallVectorImpl< Argument > &lhs, SmallVectorImpl< UnresolvedOperand > &rhs)=0
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDictWithKeyword(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary prefixed with 'attribute...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
void printFunctionalType(Operation *op)
Print the complete type of an operation in functional form.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
This class helps build Operations.
This class indicates that op operates on tosa shape types.
Operation is the basic unit of execution within MLIR.
ResultRange result_range
Support result iteration.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
OperandRange operand_range
operand_type_range getOperandTypes()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
ParseResult value() const
Access the internal ParseResult value.
bool has_value() const
Returns true if we contain a valid ParseResult value.
Type-safe wrapper around a void* for passing properties, including the properties structs of operatio...
This class provides an abstraction over the different types of ranges over Regions.
This diagnostic handler is a simple RAII class that registers and erases a diagnostic handler on a gi...
Adaptor class to abstract the differences between whether value is from a ShapedType or ShapedTypeCom...
bool isDynamicDim(int index) const
Returns whether the index'th dimension is dynamic.
int64_t getDimSize(int index) const
Returns the size of the index'th dimension.
int64_t getRank() const
Returns the rank of the shape.
bool hasStaticShape() const
Returns whether the shape is fully static.
int64_t getNumElements() const
Returns the number of elements in the shape.
void getDims(SmallVectorImpl< int64_t > &res) const
Populates the dimensions from shape referenced.
bool hasRank() const
Returns whether the shape has a rank.
ShapedTypeComponents that represents the components of a ShapedType.
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
Operation * lookup(StringRef name) const
Look up a symbol with the specified name, returning null if no such name exists.
Tensor types represent multi-dimensional arrays, and have two variants: RankedTensorType and Unranked...
ArrayRef< int64_t > getShape() const
Returns the shape of this tensor type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
Range of values and shapes (corresponding effectively to Shapes dialect's ValueShape type concept).
ShapeAdaptor getShape(int index) const
Returns the shape of index'th operand.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
ArrayRef< T > asArrayRef() const
LogicalResult verifyAtLeastNOperands(Operation *op, unsigned numOperands)
LogicalResult verifyTosaShapeOperatorWithSameRanks(Operation *op)
LogicalResult verifyTosaResolvableShapeOperands(Operation *op)
bool getBroadcastedShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2, SmallVectorImpl< int64_t > &resultShape)
Returns true and sets resultShape to the broadcasted shape from the two given shapes if they are broa...
LogicalResult convertFloatTypeFromAttribute(Type type, Attribute attr, llvm::SmallVectorImpl< char > &result)
Float type implementation of DenseElementTypeInterface::convertFromAttribute.
Attribute convertFloatTypeToAttribute(Type type, llvm::ArrayRef< char > rawData)
Float type implementation of DenseElementTypeInterface::convertToAttribute.
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
SmallVector< unsigned > getBlockSize(AffineMap dimToLvl)
Given the dimToLvl map, returns the block sizes in a vector.
ConvOpQuantizationAttr buildConvOpQuantizationAttr(OpBuilder &builder, Value input, Value weight)
Method to build ConvOpQuantizationAttr, called from ConvOpQuantInfoBuilder/TransConvOpQuantInfoBuilde...
Type getStorageElementTypeOrSelf(Type type)
RankedTensorType getVariableType(VariableOp variableOp)
Type buildConvOpResultTypeInfo(OpBuilder &builder, Type outputType, Value input, Value weight)
construct ConvOp output type with correct bitwidth based on input/weight width.
ParseResult parseVariableOpTypeOrInitialValue(OpAsmParser &parser, DenseElementsAttr &varShapeAttr, TypeAttr &typeAttr, Attribute &initialValueAttr)
PadOpQuantizationAttr buildPadOpQuantizationAttr(OpBuilder &builder, Value input)
Builds PadOpQuantizationAttr, called from PadOpQuantInfoBuilder: inputZp: input zeropoint.
constexpr int64_t kInferableDimSize
Represents a dimension in the shape of a tensor that can be inferred based on the other provided dime...
std::pair< Value, Value > createZPsAsConst(OpBuilder &builder, Value input, Value weight)
void printVariableOpTypeOrInitialValue(OpAsmPrinter &p, Operation *op, DenseElementsAttr varShapeAttr, TypeAttr typeAttr, Attribute initialValueAttr)
FailureOr< T > getConstantScalarIntValue(Value val)
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
MatMulOpQuantizationAttr buildMatMulOpQuantizationAttr(OpBuilder &builder, Value a, Value b)
Builds MatMulOpQuantizationAttr, called from MatMulOpQuantInfoBuilder: aZp: input a zeropoint bZp: in...
unsigned getBitWidth(Type type)
std::optional< Value > createZeroPointTensor(OpBuilder &builder, Location loc, Type srcElemType, int64_t zp=0)
bool isa_tosa_shape_type(mlir::Type t)
SmallVector< int64_t > convertFromMlirShape(ArrayRef< int64_t > shape)
UnaryOpQuantizationAttr buildUnaryOpQuantizationAttr(OpBuilder &builder, Value input, Type outputRawType)
Builds UnaryOpQuantizationAttr UnaryOpQuantInfoBuilder: inputZp: input zeropoint outputZp: output zer...
Type getStorageElementTypeFromQuantized(quant::QuantizedType quantizedType)
Value createPadConstTensor(OpBuilder &builder, Location loc, Value src, int32_t val=0)
LogicalResult verifyBlockScaledTensorType(mlir::Type type, llvm::function_ref< mlir::InFlightDiagnostic()> emitError=nullptr, bool allowScaleValues=false)
std::string getTosaTensorTypeErrorMessage(mlir::Type type)
bool getConstShapeValues(Operation *op, llvm::SmallVector< int64_t > &result_shape)
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
LogicalResult verifyCompatibleShapes(TypeRange types1, TypeRange types2)
Returns success if the given two arrays have the same number of elements and each pair wise entries h...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
LogicalResult emitOptionalError(std::optional< Location > loc, Args &&...args)
Overloads of the above emission functions that take an optionally null location.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
LogicalResult verifyCompatibleDims(ArrayRef< int64_t > dims)
Dimensions are compatible if all non-dynamic dims are equal.
LogicalResult verifyRanksMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching ranks.
LogicalResult verifyCompatibleShape(ArrayRef< int64_t > shape1, ArrayRef< int64_t > shape2)
Returns success if the given two shapes are compatible.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
llvm::function_ref< Fn > function_ref
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
static ValueKnowledge meet(const ValueKnowledge &lhs, const ValueKnowledge &rhs)
static ValueKnowledge getKnowledgeFromType(Type type)