29#include "llvm/ADT/APFloat.h"
30#include "llvm/ADT/APInt.h"
56 (padConstAttr.
size() != 1)) {
61 if (
auto padConstFpAttr = mlir::dyn_cast<DenseFPElementsAttr>(padConstAttr)) {
62 float padConstVal = (*padConstFpAttr.begin()).convertToFloat();
63 return padConstVal == 0.0f;
67 if (
auto padConstIntAttr =
68 mlir::dyn_cast<DenseIntElementsAttr>(padConstAttr)) {
77 int64_t padConstVal = (*padConstIntAttr.begin()).getSExtValue();
78 return zpVal == padConstVal;
86template <
typename OpTy>
87struct PoolPadFoldAdaptor;
90struct PoolPadFoldAdaptor<
tosa::MaxPool2dOp> {
91 using OpTy = tosa::MaxPool2dOp;
92 static bool checkKernelCompliance(OpTy op,
const ArrayRef<int64_t> newPad) {
93 const llvm::ArrayRef<int64_t> kernel = op.getKernel();
94 if (newPad[2] >= kernel[1] || newPad[3] >= kernel[1] ||
95 newPad[0] >= kernel[0] || newPad[1] >= kernel[0])
99 static bool checkPadConstCompliance(OpTy, Value padConst) {
101 DenseElementsAttr padConstAttr;
103 padConstAttr.
size() != 1) {
108 if (
auto padConstFpAttr =
109 mlir::dyn_cast<DenseFPElementsAttr>(padConstAttr)) {
110 const APFloat padConstVal = *padConstFpAttr.begin();
111 const APFloat lowestVal =
112 APFloat::getLargest(padConstVal.getSemantics(),
true);
113 return padConstVal == lowestVal;
115 if (
auto padConstIntAttr =
116 mlir::dyn_cast<DenseIntElementsAttr>(padConstAttr)) {
117 const APInt padConstVal = *padConstIntAttr.begin();
118 const unsigned int bitWidth = padConstVal.getBitWidth();
119 const APInt lowestVal =
120 padConstIntAttr.getElementType().isUnsignedInteger()
121 ? APInt::getZero(bitWidth)
122 : APInt::getSignedMinValue(bitWidth);
123 return padConstVal == lowestVal;
129 static void replaceOpWithNewPad(PatternRewriter &rewriter, OpTy op,
130 Value padInput, ArrayRef<int64_t> newPad) {
132 op, op.getType(), padInput, op.getKernel(), op.getStride(),
137template <
typename OpTy>
138struct ConvPadFoldAdaptor {
139 static bool checkKernelCompliance(OpTy,
const ArrayRef<int64_t>) {
142 static bool checkPadConstCompliance(OpTy op, Value padConst) {
145 static void replaceOpWithNewPad(PatternRewriter &rewriter, OpTy op,
146 Value padInput, ArrayRef<int64_t> newPad) {
148 op, op.getResult().
getType(), padInput, op.getWeight(), op.getBias(),
149 op.getInputZp(), op.getWeightZp(), newPad, op.getStrideAttr(),
150 op.getDilationAttr(), op.getAccType(), op.getLocalBound());
158template <
typename OpTy,
typename AdaptorTy>
160 using OpRewritePattern<OpTy>::OpRewritePattern;
162 LogicalResult matchAndRewrite(OpTy tensorOp,
163 PatternRewriter &rewriter)
const override {
165 auto padOp = tensorOp.getInput().template getDefiningOp<tosa::PadOp>();
168 "Producer must be a tosa::PadOp.");
171 const std::vector<int64_t> &tensorOpPad = tensorOp.getPad().vec();
172 if (tensorOpPad.size() != 4)
174 tensorOp,
"Tensor operation padding shall have 4 elements.");
177 DenseIntElementsAttr padOpPadding;
181 "The `padding` input specified on the tosa::PadOp must be constant.");
185 if (padOpPadding.size() != 8)
187 "Pad padding should have 8 elements.");
188 int64_t padNBefore = (*(padOpPadding.
begin() + 0)).getLimitedValue();
189 int64_t padNAfter = (*(padOpPadding.
begin() + 1)).getLimitedValue();
190 int64_t padHBefore = (*(padOpPadding.
begin() + 2)).getLimitedValue();
191 int64_t padHAfter = (*(padOpPadding.
begin() + 3)).getLimitedValue();
192 int64_t padWBefore = (*(padOpPadding.
begin() + 4)).getLimitedValue();
193 int64_t padWAfter = (*(padOpPadding.
begin() + 5)).getLimitedValue();
194 int64_t padCBefore = (*(padOpPadding.
begin() + 6)).getLimitedValue();
195 int64_t padCAfter = (*(padOpPadding.
begin() + 7)).getLimitedValue();
197 if (padNBefore != 0 || padNAfter != 0 || padCBefore != 0 || padCAfter != 0)
199 tensorOp,
"Folding padding in N or C dimensions is not supported.");
203 SmallVector<int64_t> foldedPad(tensorOpPad.size());
204 foldedPad[0] = padHBefore + tensorOpPad[0];
205 foldedPad[1] = padHAfter + tensorOpPad[1];
206 foldedPad[2] = padWBefore + tensorOpPad[2];
207 foldedPad[3] = padWAfter + tensorOpPad[3];
210 if (!AdaptorTy::checkKernelCompliance(tensorOp, foldedPad)) {
212 tensorOp,
"Padding size not aligned with kernel restrictions.");
216 if (!AdaptorTy::checkPadConstCompliance(tensorOp, padOp.getPadConst())) {
219 "Padding constant is not aligned with operator zero-point.");
223 if (llvm::any_of(foldedPad, [](int64_t padVal) {
return padVal > 8192; })) {
225 tensorOp,
"Padding size more than the 8K level limit.");
229 AdaptorTy::replaceOpWithNewPad(rewriter, tensorOp, padOp.getInput1(),
240 FoldPadToTensorOp<tosa::Conv2DOp, ConvPadFoldAdaptor<tosa::Conv2DOp>>>(
246 results.
add<FoldPadToTensorOp<tosa::DepthwiseConv2DOp,
247 ConvPadFoldAdaptor<tosa::DepthwiseConv2DOp>>>(
264 op,
"expected constant kernel, stride, and pad operands");
267 rewriter, op.getLoc(), op.
getType(), op.getInput(), op.getInputZp(),
282 if (op.getInput().getType() != op.getOutput().getType())
284 op,
"expected input and output types to match");
286 const auto inputType = llvm::cast<ShapedType>(op.getInput().getType());
287 if (!llvm::isa<FloatType>(inputType.getElementType()))
289 "expected floating-point input type");
297 "expected input and output zero points to be statically verifiable");
299 if (!llvm::all_of(op.getKernel(), [](
int64_t val) { return val == 1; }))
302 if (!llvm::all_of(op.getStride(), [](
int64_t val) { return val == 1; }))
305 if (!llvm::all_of(op.getPad(), [](
int64_t val) { return val == 0; }))
318void AvgPool2dAdaptiveOp::getCanonicalizationPatterns(
328 Value input = op.getInput();
329 Value output = op.getOutput();
330 ShapedType inputType = llvm::cast<ShapedType>(input.
getType());
331 ShapedType outputType = llvm::cast<ShapedType>(output.
getType());
334 llvm::all_of(op.getKernel(), [](
int64_t val) { return val == 1; }) &&
335 llvm::all_of(op.getStride(), [](
int64_t val) { return val == 1; }) &&
336 llvm::all_of(op.getPad(), [](
int64_t val) { return val == 0; }) &&
337 op.getNanMode() == tosa::NanPropagationMode::PROPAGATE) {
342 if (!inputType.hasStaticShape() || !outputType.hasStaticShape()) {
348 if (outputShape[1] != 1 || outputShape[2] != 1) {
353 if (inputShape[1] != 1 || inputShape[2] != 1) {
365 FoldPadToTensorOp<tosa::MaxPool2dOp,
366 PoolPadFoldAdaptor<tosa::MaxPool2dOp>>>(
383 op,
"expected constant kernel, stride, and pad operands");
386 rewriter, op.getLoc(), op.
getType(), op.getInput(),
395void MaxPool2dAdaptiveOp::getCanonicalizationPatterns(
409 if (op.getInput1().size() != 1)
411 if (op.getInput1().front().getType() != op.getType()) {
414 op.getInput1().front())
419 rewriter.
replaceOp(op, op.getInput1().front());
434 concatOperands.reserve(2 * op.getNumOperands());
436 int32_t maxNumOperands = 0;
442 bool foundRewritableConcat =
false;
443 for (
Value operand : op.getOperands()) {
444 concatOperands.emplace_back(operand);
446 auto producer = operand.getDefiningOp<tosa::ConcatOp>();
451 if (op.getAxis() != producer.getAxis())
455 foundRewritableConcat =
true;
456 concatOperands.pop_back();
457 llvm::append_range(concatOperands, producer->getOperands());
460 if (!foundRewritableConcat)
462 "No rewritable concat operand found.");
464 if (maxNumOperands > 0 &&
465 concatOperands.size() >
static_cast<size_t>(maxNumOperands))
467 op,
"Rewriting would exceed the maximum number of operands for the "
468 "target environment level.");
471 op, op.getType(), concatOperands, op.getAxisAttr());
481LogicalResult SelectOp::canonicalize(SelectOp op,
PatternRewriter &rewriter) {
482 auto notOp = op.getInput1().getDefiningOp<tosa::LogicalNotOp>();
486 op.getOperation()->setOperands(
487 {notOp.getInput1(), op.getOnFalse(), op.getOnTrue()});
499 auto innerTranspose =
500 transposeOp.getInput1().getDefiningOp<tosa::TransposeOp>();
503 "input must be transpose operation");
507 innerTranspose.getPerms();
509 if (transposePerms.size() != innerTransposePerms.size())
512 "transpose and inner transpose perms sizes must be equal");
513 if (transposePerms.empty())
515 transposeOp,
"transpose perms sizes must be positive");
519 for (
int i = 0, s = transposePerms.size(); i < s; ++i)
520 perms[i] = innerTransposePerms[transposePerms[i]];
523 transposeOp, transposeOp.getResult().
getType(),
536 if (op.getInput1().getDefiningOp<tosa::TransposeOp>())
538 op,
"Src is from transpose, can compose transposes");
542 if (isa_and_nonnull<tosa::TransposeOp>(subop))
544 op,
"Dest is used by transpose, can compose transposes");
547 auto input = op.getInput1();
548 auto inputTy = llvm::cast<ShapedType>(input.
getType());
549 if (!inputTy.hasRank())
553 for (
int i = 0; i < inputTy.getRank(); ++i)
554 if (inputTy.isDynamicDim(i))
563 nonZeroPerms.reserve(permValues.size());
564 for (
auto idx : permValues) {
565 auto sz = inputTy.getDimSize(idx);
567 nonZeroPerms.push_back(idx);
570 for (
int i = 1, s = nonZeroPerms.size(); i < s; ++i)
571 if (nonZeroPerms[i - 1] > nonZeroPerms[i])
573 "Transpose changes memory layout.");
576 newShape.reserve(inputTy.getRank());
577 for (
int i = 0, s = inputTy.getRank(); i < s; ++i)
578 newShape.push_back(inputTy.getDimSize(permValues[i]));
581 op, op.getType(), op.getInput1(),
589 results.
add<ConsolidateTransposeOptimization, TransposeIsReshape>(context);
597 Value input = op.getInput();
598 auto inputType = llvm::cast<ShapedType>(op.getInput().getType());
599 auto inputElementType = inputType.getElementType();
601 if (isa<FloatType>(inputElementType)) {
603 const auto minClamp =
604 llvm::cast<mlir::FloatAttr>(op.getMinValAttr()).getValue();
605 const auto maxClamp =
606 llvm::cast<mlir::FloatAttr>(op.getMaxValAttr()).getValue();
607 const bool isMin = minClamp.isNegInfinity();
608 const bool isMax = maxClamp.isInfinity();
610 if (isMin && isMax) {
618 const bool isBoolean = inputElementType.isInteger(1);
619 if (inputElementType.isUnsignedInteger() || isBoolean) {
620 const int64_t minClamp = llvm::cast<mlir::IntegerAttr>(op.getMinValAttr())
623 const int64_t maxClamp = llvm::cast<mlir::IntegerAttr>(op.getMaxValAttr())
627 const unsigned bitWidth = inputElementType.getIntOrFloatBitWidth();
628 const int64_t intMin = APInt::getMinValue(bitWidth).getZExtValue();
629 const int64_t intMax = APInt::getMaxValue(bitWidth).getZExtValue();
631 if (minClamp <= intMin && maxClamp >= intMax) {
638 if (llvm::isa<IntegerType>(inputElementType)) {
640 llvm::cast<mlir::IntegerAttr>(op.getMinValAttr()).getInt();
642 llvm::cast<mlir::IntegerAttr>(op.getMaxValAttr()).getInt();
644 const unsigned bitWidth = inputElementType.getIntOrFloatBitWidth();
645 const int64_t intMin = APInt::getSignedMinValue(bitWidth).getSExtValue();
646 const int64_t intMax = APInt::getSignedMaxValue(bitWidth).getSExtValue();
648 if (minClamp <= intMin && maxClamp >= intMax) {
680 template <
typename T>
694 Value input = op.getInput();
702 const auto opNanMode = op.getNanMode();
703 const auto clampNanMode = clampOp.getNanMode();
704 if (opNanMode == NanPropagationMode::IGNORE &&
705 clampNanMode == NanPropagationMode::PROPAGATE)
708 auto maxValAttr = op.getMaxValAttr();
709 auto minValAttr = op.getMinValAttr();
710 auto clampOpMaxValAttr = clampOp.getMaxValAttr();
711 auto clampOpMinValAttr = clampOp.getMinValAttr();
713 auto inputEType = llvm::cast<ShapedType>(input.
getType()).getElementType();
715 llvm::dyn_cast<mlir::quant::UniformQuantizedType>(inputEType)) {
720 if (mlir::isa<FloatType>(inputEType)) {
721 auto floatMaxValAttr = cast<mlir::FloatAttr>(maxValAttr);
722 auto floatMinValAttr = cast<mlir::FloatAttr>(minValAttr);
723 auto clampOpFloatMaxValAttr = cast<mlir::FloatAttr>(clampOpMaxValAttr);
724 auto clampOpFloatMinValAttr = cast<mlir::FloatAttr>(clampOpMinValAttr);
727 const auto opMinFloat = floatMinValAttr.getValue();
728 const auto opMaxFloat = floatMaxValAttr.getValue();
729 const auto clampOpMinFloat = clampOpFloatMinValAttr.getValue();
730 const auto clampOpMaxFloat = clampOpFloatMaxValAttr.getValue();
734 if (!opRangeFloatRange.
intersects(clampRangeFloatRange))
738 auto newMinVal = std::max(opMinFloat, clampOpMinFloat);
739 auto newMaxVal = std::min(opMaxFloat, clampOpMaxFloat);
740 newMinValAttr = rewriter.
getFloatAttr(inputEType, newMinVal);
741 newMaxValAttr = rewriter.
getFloatAttr(inputEType, newMaxVal);
743 assert(mlir::isa<IntegerType>(inputEType));
744 auto intMaxValAttr = cast<mlir::IntegerAttr>(maxValAttr);
745 auto intMinValAttr = cast<mlir::IntegerAttr>(minValAttr);
746 auto clampOpIntMaxValAttr = cast<mlir::IntegerAttr>(clampOpMaxValAttr);
747 auto clampOpIntMinValAttr = cast<mlir::IntegerAttr>(clampOpMinValAttr);
749 if (inputEType.isUnsignedInteger()) {
751 const auto opMinInt = intMinValAttr.getUInt();
752 const auto opMaxInt = intMaxValAttr.getUInt();
753 const auto clampOpMinInt = clampOpIntMinValAttr.getUInt();
754 const auto clampOpMaxInt = clampOpIntMaxValAttr.getUInt();
758 if (!opRangeIntRange.
intersects(clampRangeIntRange))
762 auto newMinVal = std::max(opMinInt, clampOpMinInt);
763 auto newMaxVal = std::min(opMaxInt, clampOpMaxInt);
768 const auto opMinInt = intMinValAttr.getInt();
769 const auto opMaxInt = intMaxValAttr.getInt();
770 const auto clampOpMinInt = clampOpIntMinValAttr.getInt();
771 const auto clampOpMaxInt = clampOpIntMaxValAttr.getInt();
775 if (!opRangeIntRange.
intersects(clampRangeIntRange))
779 auto newMinVal = std::max(opMinInt, clampOpMinInt);
780 auto newMaxVal = std::min(opMaxInt, clampOpMaxInt);
786 auto newMode = (opNanMode != clampNanMode)
787 ? tosa::NanPropagationMode::IGNORE
791 NanPropagationModeAttr::get(rewriter.
getContext(), newMode);
794 op, op.getType(), clampOp.getInput(), newMinValAttr, newMaxValAttr,
800void ClampOp::getCanonicalizationPatterns(RewritePatternSet &results,
801 MLIRContext *context) {
802 results.
add<ClampIsNoOp>(context);
803 results.
add<ClampClampOptimization>(context);
811 Value sliceInput = sliceOp.getInput1();
815 sliceOp,
"slice input must be concat operation");
818 auto concatType = dyn_cast<RankedTensorType>(concatOp.getType());
819 if (!concatType || !concatType.hasStaticShape())
821 sliceOp,
"slice input must be a static ranked tensor");
822 int32_t axis = concatOp.getAxis();
829 sliceOp,
"start of slice must be a static ranked shape");
833 sliceOp,
"size of slice must be a static ranked shape");
843 std::optional<Value> replaceWithSlice;
844 for (
auto input : inputs) {
845 auto inputType = dyn_cast<RankedTensorType>(input.
getType());
846 if (!inputType || !inputType.hasStaticShape())
848 sliceOp,
"concat input must be a static ranked tensor");
850 if (sliceStarts[axis] >= 0 && (sliceStarts[axis] + sliceSizes[axis]) <=
851 inputType.getDimSize(axis)) {
857 tosa::SliceOp::create(rewriter, sliceOp.getLoc(), sliceOp.
getType(),
858 input, start_op, size_op)
862 sliceStarts[axis] -= inputType.getDimSize(axis);
865 if (!replaceWithSlice)
867 sliceOp,
"corresponding concat input not found for slice");
869 rewriter.
replaceOp(sliceOp, replaceWithSlice.value());
879 Value sliceInput = sliceOp.getInput1();
885 "slice input must be a pad operation");
888 if (!padOp->hasOneUse())
890 "pad shall have a single consumer");
893 auto inputTy = dyn_cast<RankedTensorType>(padOp.getInput1().getType());
894 auto padTy = dyn_cast<RankedTensorType>(padOp.getType());
895 if (!inputTy || !padTy || !inputTy.hasRank())
897 "slice input must be a ranked tensor");
904 "`padding` input specified on the tosa::PadOp must be constant.");
907 llvm::to_vector(paddingElems.getValues<
int64_t>());
913 sliceOp,
"start of slice must be a static ranked shape");
920 sliceOp,
"size of slice must be a static ranked shape");
925 const int64_t rank = inputTy.getRank();
926 if (llvm::any_of(llvm::seq<int64_t>(0, rank), [&](
int64_t i) {
927 const bool isDimDynamic = inputTy.isDynamicDim(i);
928 const bool isDimSliced =
931 return isDimDynamic && isDimSliced;
934 sliceOp,
"axis that are sliced shall be statically known.");
941 bool updated =
false;
943 for (
int64_t i = 0; i < rank; ++i) {
944 const int64_t padLo = padPaddings[i * 2];
945 const int64_t padHi = padPaddings[i * 2 + 1];
946 const int64_t sliceStart = sliceStarts[i];
947 const int64_t sliceSize = sliceSizes[i];
948 const int64_t sliceEnd = sliceStart + sliceSize;
951 if (inputTy.isDynamicDim(i)) {
952 newPadPaddings[i * 2] = padLo;
953 newPadPaddings[i * 2 + 1] = padHi;
954 newSliceStarts[i] = sliceStart;
959 const int64_t dimSize = inputTy.getShape()[i];
960 const int64_t dimTotal = padLo + dimSize + padHi;
963 if (sliceStart < 0 || sliceEnd > dimTotal)
967 const int64_t newSliceStart = std::max<int64_t>(sliceStart - padLo, 0);
968 newSliceStarts[i] = newSliceStart;
969 updated |= newSliceStart != sliceStart;
972 const int64_t newPadLo = std::max<int64_t>(padLo - sliceStart, 0);
974 std::max<int64_t>(sliceEnd - (padLo + dimSize), 0);
975 newPadPaddings[i * 2] = newPadLo;
976 newPadPaddings[i * 2 + 1] = newPadHi;
977 updated |= (newPadLo != padLo) || (newPadHi != padHi);
981 newPadPaddings[i * 2] + dimSize + newPadPaddings[i * 2 + 1];
987 sliceOp,
"terminate condition; nothing to rewrite");
993 RankedTensorType::get(newPadShape, inputTy.getElementType());
994 auto newPadOp = tosa::PadOp::create(rewriter, padOp.getLoc(), newPadTy,
995 padOp.getInput1(), newPaddingsOp,
996 padOp.getPadConst());
1002 newPadOp.getResult(), newStartOp,
1017 ShapedType resultType = cast<ShapedType>(sliceOp.getType());
1018 if (!resultType.hasRank())
1021 ElementsAttr sizeElems;
1024 sliceOp,
"size of slice must be a static ranked shape");
1028 llvm::to_vector(sizeElems.getValues<
int64_t>());
1030 bool replaceSliceSize{
false};
1034 for (
const auto &[
index, size] : llvm::enumerate(sliceSizes)) {
1036 sliceSizes[
index] = resultType.getDimSize(
index);
1037 replaceSliceSize =
true;
1041 if (!replaceSliceSize) {
1043 sliceOp,
"no dimension of size of slice is dynamic that resolves "
1044 "to static output shape");
1049 tosa::SliceOp::create(rewriter, sliceOp.getLoc(), sliceOp.
getType(),
1050 sliceOp.getInput1(), sliceOp.getStart(), size_op);
1052 rewriter.
replaceOp(sliceOp, newSliceOp.getResult());
1057void SliceOp::getCanonicalizationPatterns(RewritePatternSet &results,
1058 MLIRContext *context) {
1059 results.
add<ConcatSliceOptimization, PadSliceOptimization,
1060 SliceDynamicSizeCanonicalization>(context);
1068 const Value castInput = castOp.getInput();
1072 "input must be cast operation");
1074 const Value innerCastInput = innerCastOp.getInput();
1076 const ShapedType innerInputType =
1077 llvm::cast<ShapedType>(innerCastInput.
getType());
1078 const ShapedType innerOutputType =
1079 llvm::cast<ShapedType>(innerCastOp.getType());
1080 const ShapedType outerOutputType = llvm::cast<ShapedType>(castOp.getType());
1082 const Type innerInputElemType = innerInputType.getElementType();
1083 const Type innerOutputElemType = innerOutputType.getElementType();
1084 const Type outerOutputElemType = outerOutputType.getElementType();
1087 outerOutputElemType};
1089 if (llvm::any_of(types, [](
const Type type) {
1093 llvm::isa<Float8E4M3FNType, Float8E5M2Type, BFloat16Type,
1094 Float16Type, Float32Type>(type));
1097 castOp,
"only integer and f32, f16, bf16, f8E4M3FN, f8E5M2 types are "
1100 if (llvm::isa<Float8E5M2Type>(innerInputElemType) &&
1101 llvm::isa<Float8E4M3FNType>(outerOutputElemType)) {
1103 castOp,
"avoid introducing f8E5M2 -> f8E4M3FN casts which are not "
1107 if (llvm::isa<Float8E4M3FNType>(innerInputElemType) &&
1108 llvm::isa<Float8E5M2Type>(outerOutputElemType)) {
1110 castOp,
"avoid introducing f8E4M3FN -> f8E5M2 casts which are not "
1114 if (llvm::isa<Float8E5M2Type, Float8E4M3FNType>(innerInputElemType) &&
1117 castOp,
"avoid introducing fp8 -> integer casts which are not "
1122 llvm::isa<Float8E5M2Type, Float8E4M3FNType>(outerOutputElemType)) {
1124 castOp,
"avoid introducing integer -> fp8 casts which are not "
1128 if (llvm::isa<Float16Type>(innerInputElemType) &&
1129 llvm::isa<BFloat16Type>(outerOutputElemType)) {
1131 castOp,
"avoid introducing fp16 -> bf16 casts which are not "
1135 if (llvm::isa<BFloat16Type>(innerInputElemType) &&
1136 llvm::isa<Float16Type>(outerOutputElemType)) {
1138 castOp,
"avoid introducing bf16 -> fp16 casts which are not "
1142 const auto isIntegerOneOfWidth = [](
Type type,
size_t bitwidth1,
1147 if (isIntegerOneOfWidth(innerInputElemType, 8, 16) &&
1150 castOp,
"avoid introducing i8/i16 -> i64 casts which are not "
1154 if (isIntegerOneOfWidth(innerInputElemType, 1, 64) &&
1157 castOp,
"avoid introducing bool/i64 to float casts which are not "
1158 "supported in all versions of TOSA");
1162 isIntegerOneOfWidth(outerOutputElemType, 1, 64)) {
1164 castOp,
"avoid introducing float to bool/i64 casts which are not "
1165 "supported in all versions of TOSA");
1171 "inner cast operation is narrowing");
1180 return semantics.nonFiniteBehavior !=
1181 llvm::fltNonfiniteBehavior::FiniteOnly;
1185 return semantics.nonFiniteBehavior == llvm::fltNonfiniteBehavior::IEEE754;
1189 const ShapedType outType)
const {
1191 if (inType.getElementType().isInteger() &&
1192 outType.getElementType().isInteger()) {
1194 const auto inTypeSignedness =
1195 cast<IntegerType>(inType.getElementType()).getSignedness();
1196 const auto outTypeSignedness =
1197 cast<IntegerType>(outType.getElementType()).getSignedness();
1199 return (inTypeSignedness != outTypeSignedness ||
1200 inType.getElementTypeBitWidth() >
1201 outType.getElementTypeBitWidth());
1204 if (inType.getElementType().isFloat() &&
1205 outType.getElementType().isFloat()) {
1207 FloatType inElemTy = cast<FloatType>(inType.getElementType());
1208 FloatType outElemTy = cast<FloatType>(outType.getElementType());
1209 llvm::fltSemantics inTypeSemantics = inElemTy.getFloatSemantics();
1210 llvm::fltSemantics outTypeSemantics = outElemTy.getFloatSemantics();
1216 [[maybe_unused]]
const auto isSupported = [](
Type elemType) {
1217 return llvm::isa<Float8E4M3FNType, Float8E5M2Type, BFloat16Type,
1218 Float16Type, Float32Type>(elemType);
1221 assert(isSupported(inElemTy) &&
1222 "unsupported input element type in isNarrowingCast");
1223 assert(isSupported(outElemTy) &&
1224 "unsupported output element type in isNarrowingCast");
1227 inTypeSemantics.maxExponent > outTypeSemantics.maxExponent ||
1228 inTypeSemantics.minExponent < outTypeSemantics.minExponent ||
1229 inTypeSemantics.precision > outTypeSemantics.precision ||
1246 const Value outerInput = castOp.getInput();
1247 auto innerCastOp = outerInput.
getDefiningOp<tosa::CastOp>();
1250 "input must be a cast operation");
1252 const Value innerInput = innerCastOp.getInput();
1253 const auto innerInputTy = llvm::cast<ShapedType>(innerInput.
getType());
1254 const auto innerOutputTy = llvm::cast<ShapedType>(innerCastOp.getType());
1255 const auto outerOutputTy = llvm::cast<ShapedType>(castOp.getType());
1257 if (!llvm::isa<tosa::BlockScaledType>(innerInputTy.getElementType()))
1259 castOp,
"inner cast input must have block scaled element type");
1261 if (innerInputTy != outerOutputTy)
1263 castOp,
"inner input type must match outer output type");
1265 const Type innerOutputElemType = innerOutputTy.getElementType();
1266 const bool isLosslessCast =
1267 isa<Float32Type, BFloat16Type>(innerOutputElemType);
1268 if (!isLosslessCast)
1270 castOp,
"avoid cancelling casts that should be lossy");
1278void CastOp::getCanonicalizationPatterns(RewritePatternSet &results,
1279 MLIRContext *context) {
1280 results.
add<NonNarrowingCastsOptimization,
1281 CancellingBlockScaledCastsOptimization>(context);
1290 const Value castToBlockScaledInput = castToBlockScaledOp.getInputData();
1291 auto castFromBlockScaledOp =
1292 castToBlockScaledInput.
getDefiningOp<tosa::CastFromBlockScaledOp>();
1293 if (!castFromBlockScaledOp)
1295 castToBlockScaledOp,
1296 "input must be cast_from_block_scaled operation");
1298 const Value innerData = castFromBlockScaledOp.getInputData();
1299 const Value innerScale = castFromBlockScaledOp.getInputScale();
1300 const auto innerDataTy = llvm::cast<ShapedType>(innerData.
getType());
1301 const auto innerScaleTy = llvm::cast<ShapedType>(innerScale.
getType());
1303 const Value outerData = castToBlockScaledOp.getOutputData();
1304 const Value outerScale = castToBlockScaledOp.getOutputScale();
1305 const auto outerDataTy = llvm::cast<ShapedType>(outerData.
getType());
1306 const auto outerScaleTy = llvm::cast<ShapedType>(outerScale.
getType());
1308 if (innerDataTy != outerDataTy || innerScaleTy != outerScaleTy) {
1310 castToBlockScaledOp,
1311 "inputs types to cast_from_block_scaled operation must match output "
1312 "types to cast_to_block_scaled");
1315 if (castFromBlockScaledOp.getBlockSize() !=
1316 castToBlockScaledOp.getBlockSize()) {
1318 castToBlockScaledOp,
"block sizes for cast_from_block_scaled and "
1319 "cast_to_block_scaled must match");
1322 rewriter.
replaceOp(castToBlockScaledOp, {innerData, innerScale});
1328void CastToBlockScaledOp::getCanonicalizationPatterns(
1329 RewritePatternSet &results, MLIRContext *context) {
1330 results.
add<CancellingCastToFromBlockScaledOptimization>(context);
1338 const FailureOr<int32_t> rowCount =
1340 if (failed(rowCount) || rowCount.value() != 1)
1344 op, op.getOutput().
getType(), op.getValues(), op.getIndices());
1349void RowGatherOp::getCanonicalizationPatterns(RewritePatternSet &results,
1350 MLIRContext *context) {
1351 results.
add<RowGatherToGather>(context);
1358template <
typename Folder>
1359static DenseElementsAttr
1361 bool foldDenseValues =
false) {
1365 if (!returnTy.hasRank() || !returnTy.hasStaticShape())
1369 const auto rETy = llvm::cast<ShapedType>(
rhs.getType()).getElementType();
1373 if (
lhs.isSplat() &&
rhs.isSplat()) {
1374 if (isa<FloatType>(lETy)) {
1375 const APFloat l =
lhs.getSplatValue<APFloat>();
1376 const APFloat r =
rhs.getSplatValue<APFloat>();
1377 const auto maybeResult = Folder::fold(l, r);
1378 if (failed(maybeResult))
1383 if (
const auto lIntTy = llvm::dyn_cast<IntegerType>(lETy)) {
1384 const APInt l =
lhs.getSplatValue<APInt>();
1385 const APInt r =
rhs.getSplatValue<APInt>();
1386 const auto maybeResult = Folder::fold(l, r, lIntTy.isUnsigned());
1387 if (failed(maybeResult))
1393 if (foldDenseValues) {
1394 assert(lETy.isIntOrIndex() &&
1395 "Only integer types are currently supported.");
1398 llvm::zip(
lhs.getValues<APInt>(),
rhs.getValues<APInt>())) {
1399 const auto maybeResult = Folder::fold(l, r,
false);
1400 if (failed(maybeResult))
1402 resultValues.push_back(maybeResult.value());
1410template <
typename Folder>
1412 bool foldDenseValues =
false) {
1416 if (!returnTy.hasRank() || !returnTy.hasStaticShape())
1422 if (
const auto vIntTy = llvm::dyn_cast<IntegerType>(vETy)) {
1424 const auto maybeResult = Folder::fold(v, vIntTy.isUnsigned());
1425 if (failed(maybeResult))
1431 if (foldDenseValues) {
1435 for (
auto const &v : val.
getValues<APInt>()) {
1436 const auto maybeResult = Folder::fold(v,
false);
1437 if (failed(maybeResult))
1439 resultValues.push_back(maybeResult.value());
1454 assert(dense.isSplat());
1455 APInt a = dense.getSplatValue<APInt>();
1456 return a.getSExtValue();
1461 const bool isUnsigned) {
1464 isUnsigned ?
lhs.uadd_ov(
rhs, overflow) :
lhs.sadd_ov(
rhs, overflow);
1470 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1477 const bool isUnsigned) {
1480 isUnsigned ?
lhs.usub_ov(
rhs, overflow) :
lhs.ssub_ov(
rhs, overflow);
1486 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1493 const bool isUnsigned) {
1495 const unsigned originalWidth =
lhs.getBitWidth();
1498 if (
lhs.getBitWidth() !=
rhs.getBitWidth()) {
1503 if (
lhs == 0 ||
rhs == 0)
1504 return APInt::getZero(originalWidth);
1506 bool overflow =
false;
1508 isUnsigned ?
lhs.umul_ov(
rhs, overflow) :
lhs.smul_ov(
rhs, overflow);
1513 return result.trunc(originalWidth);
1516 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1522 return a.isNegative() !=
b.isNegative();
1529 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1537 APInt::udivrem(
lhs,
rhs, q, r);
1538 if (!r.isZero() && Ceil) {
1545 bool overflow{
false};
1546 APInt
const q =
lhs.sdiv_ov(
rhs, overflow);
1549 APInt
const r =
lhs.srem(
rhs);
1559 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1567 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1569 if (
lhs.isNegative() || (!
rhs.isStrictlyPositive()))
1579 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1581 auto const r = t.mod(
rhs);
1582 if (llvm::APFloatBase::opStatus::opOK == r) {
1592 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1594 return lhs.getSExtValue() >=
rhs.getSExtValue() ?
lhs :
rhs;
1597 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1605 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1607 return lhs.getSExtValue() <=
rhs.getSExtValue() ?
lhs :
rhs;
1610 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1616 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1617 auto const numBits = value.getBitWidth();
1619 auto const zextv = value.getZExtValue();
1620 if (zextv >= numBits)
1622 return APInt::getOneBitSet(numBits, zextv);
1624 auto const sextv = value.getSExtValue();
1625 if (sextv < 0 || sextv >= numBits || (value.isNegative()))
1627 return APInt::getOneBitSet(numBits, sextv);
1637 assert(!isUnsigned &&
1638 "unsigned values are not supported for shape div folders");
1639 if (
lhs.isNegative() || !
rhs.isStrictlyPositive())
1644 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1650 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1651 if (!value.isStrictlyPositive())
1653 return APInt(value.getBitWidth(), value.ceilLogBase2());
1658 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1659 if (!value.isStrictlyPositive())
1661 return APInt(value.getBitWidth(), value.logBase2());
1667 const bool isUnsigned) {
1668 return isUnsigned ? APInt(1,
lhs.ugt(
rhs)) : APInt(1,
lhs.sgt(
rhs));
1671 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1672 return APInt(1,
lhs >
rhs);
1678 const bool isUnsigned) {
1679 return isUnsigned ? APInt(1,
lhs.uge(
rhs)) : APInt(1,
lhs.sge(
rhs));
1682 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1683 return APInt(1,
lhs >=
rhs);
1689 const bool isUnsigned) {
1690 return APInt(1,
lhs ==
rhs);
1693 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1694 return APInt(1,
lhs ==
rhs);
1699 if (llvm::isa<FloatType>(elemType))
1701 if (llvm::isa<IntegerType>(elemType))
1707 if (llvm::isa<FloatType>(elemType))
1708 return val && val.
isSplat() &&
1710 if (llvm::isa<IntegerType>(elemType)) {
1711 const int64_t shifted = 1LL << shift;
1712 return val && val.
isSplat() &&
1718OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
1719 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1720 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1721 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1722 if (!lhsTy || !rhsTy || !resultTy)
1726 if (!lhsTy.getElementType().isIntOrIndexOrFloat() ||
1727 !rhsTy.getElementType().isIntOrIndexOrFloat())
1730 auto resultETy = resultTy.getElementType();
1732 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1734 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1737 lhsTy.getShape(), rhsTy.getShape());
1738 if (isBroadcastable && lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr))
1740 if (isBroadcastable && rhsTy == resultTy &&
isSplatZero(resultETy, lhsAttr))
1743 if (!lhsAttr || !rhsAttr)
1749OpFoldResult ArgMaxOp::fold(FoldAdaptor adaptor) {
1750 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput().
getType());
1751 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
1752 if (!inputTy || !outputTy || !inputTy.hasStaticShape() ||
1753 !outputTy.hasStaticShape())
1757 if (inputTy.getDimSize(getAxis()) == 1 && outputElementTy.
isInteger()) {
1758 const auto outputElemIntTy = cast<IntegerType>(outputElementTy);
1759 const APInt zero = APInt::getZero(outputElemIntTy.getWidth());
1766OpFoldResult IntDivOp::fold(FoldAdaptor adaptor) {
1767 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1768 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1769 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1770 if (!lhsTy || !rhsTy || !resultTy)
1772 if (lhsTy.getElementType() != rhsTy.getElementType())
1777 auto resultETy = resultTy.getElementType();
1779 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1781 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1782 if (lhsAttr && lhsAttr.isSplat() && rhsAttr && rhsAttr.isSplat()) {
1783 if (llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape() &&
1784 lhsAttr.getSplatValue<APInt>().isZero() &&
1785 !rhsAttr.getSplatValue<APInt>().isZero()) {
1786 return lhsAttr.resizeSplat(resultTy);
1790 if (rhsAttr && rhsAttr.isSplat()) {
1792 lhsTy.getShape(), rhsTy.getShape());
1793 if (isBroadcastable && lhsTy == resultTy &&
1794 llvm::isa<IntegerType>(resultETy) &&
1795 rhsAttr.getSplatValue<APInt>().isOne())
1799 if (rhsAttr && lhsAttr && rhsAttr.isSplat() && lhsAttr.isSplat() &&
1800 llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape()) {
1801 APInt l = lhsAttr.getSplatValue<APInt>();
1802 APInt r = rhsAttr.getSplatValue<APInt>();
1804 auto intTy = dyn_cast<mlir::IntegerType>(resultETy);
1806 DivFoldAdaptor<
false>::fold(l, r, intTy.isUnsigned());
1819std::optional<APInt> mulInt(APInt
lhs, APInt
rhs, int32_t shift,
1820 unsigned bitwidth) {
1821 bool overflow =
false;
1822 APInt
result =
lhs.sext(64).smul_ov(
rhs.sext(64), overflow);
1825 return std::nullopt;
1828 auto round = APInt(64, 1) << (shift - 1);
1830 result.ashrInPlace(shift);
1833 if (!(
result.getSExtValue() >= INT32_MIN &&
1834 result.getSExtValue() <= INT32_MAX)) {
1836 return std::nullopt;
1840 return result.trunc(bitwidth);
1843DenseElementsAttr mulBinaryFolder(DenseElementsAttr
lhs, DenseElementsAttr
rhs,
1844 RankedTensorType ty, int32_t shift) {
1846 if (!ty.hasStaticShape())
1849 if (llvm::isa<IntegerType>(ty.getElementType())) {
1850 APInt l =
lhs.getSplatValue<APInt>();
1851 APInt r =
rhs.getSplatValue<APInt>();
1857 auto bitwidth = ty.getElementType().getIntOrFloatBitWidth();
1858 const std::optional<APInt>
result = mulInt(l, r, shift, bitwidth);
1864 if (llvm::isa<FloatType>(ty.getElementType())) {
1865 APFloat l =
lhs.getSplatValue<APFloat>();
1866 APFloat r =
rhs.getSplatValue<APFloat>();
1876OpFoldResult MulOp::fold(FoldAdaptor adaptor) {
1877 auto lhs = getInput1();
1878 auto rhs = getInput2();
1879 auto lhsTy = llvm::dyn_cast<RankedTensorType>(
lhs.getType());
1880 auto rhsTy = llvm::dyn_cast<RankedTensorType>(
rhs.getType());
1881 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1882 if (!lhsTy || !rhsTy || !resultTy)
1885 auto resultETy = resultTy.getElementType();
1887 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1889 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1894 if (resultETy.isInteger(32)) {
1895 ElementsAttr shift_elem;
1896 if (getShift().getImpl()) {
1900 shift = shift_elem.getValues<IntegerAttr>()[0].getInt();
1904 if (rhsTy == resultTy &&
isSplatZero(resultETy, lhsAttr) &&
1905 resultTy.hasStaticShape())
1907 return lhsAttr.resizeSplat(resultTy);
1908 if (lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr) &&
1909 resultTy.hasStaticShape())
1910 return rhsAttr.resizeSplat(resultTy);
1913 lhsTy.getShape(), rhsTy.getShape());
1914 if (isBroadcastable && rhsTy == resultTy &&
1917 if (isBroadcastable && lhsTy == resultTy &&
1921 return mulBinaryFolder(lhsAttr, rhsAttr, resultTy, shift);
1924OpFoldResult SubOp::fold(FoldAdaptor adaptor) {
1925 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1926 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1927 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1928 if (!lhsTy || !rhsTy || !resultTy)
1932 if (!lhsTy.getElementType().isIntOrIndexOrFloat() ||
1933 !rhsTy.getElementType().isIntOrIndexOrFloat())
1936 auto resultETy = resultTy.getElementType();
1938 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1940 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1943 lhsTy.getShape(), rhsTy.getShape());
1944 if (isBroadcastable && lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr))
1947 if (!lhsAttr || !rhsAttr)
1953OpFoldResult GreaterOp::fold(FoldAdaptor adaptor) {
1954 auto resultTy = llvm::cast<ShapedType>(
getType());
1956 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1958 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1960 if (!lhsAttr || !rhsAttr)
1966OpFoldResult GreaterEqualOp::fold(FoldAdaptor adaptor) {
1967 auto resultTy = llvm::cast<ShapedType>(
getType());
1969 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1971 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1973 if (!lhsAttr || !rhsAttr)
1979OpFoldResult EqualOp::fold(FoldAdaptor adaptor) {
1980 auto resultTy = llvm::cast<ShapedType>(
getType());
1982 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1984 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1985 Value
lhs = getInput1();
1986 Value
rhs = getInput2();
1987 auto lhsTy = llvm::cast<ShapedType>(
lhs.getType());
1991 if (llvm::isa<IntegerType>(lhsTy.getElementType()) && resultTy.hasRank() &&
1992 resultTy.hasStaticShape() &&
lhs ==
rhs) {
1996 if (!lhsAttr || !rhsAttr)
2002OpFoldResult CastOp::fold(FoldAdaptor adaptor) {
2006 auto operand = llvm::dyn_cast_if_present<ElementsAttr>(adaptor.getInput());
2010 auto inTy = llvm::cast<ShapedType>(getInput().
getType());
2011 auto outTy = llvm::cast<ShapedType>(
getType());
2012 if (!outTy.hasRank() || !outTy.hasStaticShape())
2014 auto inETy = inTy.getElementType();
2015 auto outETy = outTy.getElementType();
2017 if (operand.isSplat()) {
2018 if (llvm::isa<FloatType>(inETy) && llvm::isa<FloatType>(outETy)) {
2020 auto splatVal = operand.getSplatValue<APFloat>();
2021 auto &semantics = llvm::cast<FloatType>(outETy).getFloatSemantics();
2022 splatVal.convert(semantics, llvm::RoundingMode::NearestTiesToEven,
2027 if (llvm::isa<IntegerType>(inETy) && llvm::isa<FloatType>(outETy)) {
2028 auto unsign = llvm::cast<IntegerType>(inETy).isUnsignedInteger();
2029 APFloat splatVal(llvm::cast<FloatType>(outETy).getFloatSemantics());
2030 splatVal.convertFromAPInt(operand.getSplatValue<APInt>(), !unsign,
2031 llvm::RoundingMode::NearestTiesToEven);
2035 if (llvm::isa<FloatType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2036 auto unsign = llvm::cast<IntegerType>(outETy).isUnsignedInteger();
2037 auto intVal = APSInt(
2038 llvm::cast<IntegerType>(outETy).getIntOrFloatBitWidth(), unsign);
2039 auto floatVal = operand.getSplatValue<APFloat>();
2041 floatVal.convertToInteger(intVal, llvm::RoundingMode::NearestTiesToEven,
2046 if (llvm::isa<IntegerType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2047 const auto inIntType = llvm::cast<IntegerType>(inETy);
2048 auto unsignIn = inIntType.isUnsignedInteger();
2050 inETy.getIntOrFloatBitWidth() > outETy.getIntOrFloatBitWidth();
2051 auto intVal = operand.getSplatValue<APInt>();
2052 auto bitwidth = outETy.getIntOrFloatBitWidth();
2055 if (outETy.isInteger(1)) {
2056 intVal = APInt(bitwidth, intVal.isZero() ? 0 : 1);
2058 intVal = intVal.trunc(bitwidth);
2059 }
else if (unsignIn || inIntType.isInteger(1)) {
2060 intVal = intVal.zext(bitwidth);
2062 intVal = intVal.sext(bitwidth);
2072OpFoldResult ConstOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2074OpFoldResult ConstShapeOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2076#define REDUCE_FOLDER(OP) \
2077 OpFoldResult OP::fold(FoldAdaptor adaptor) { \
2078 ShapedType inputTy = llvm::cast<ShapedType>(getInput().getType()); \
2079 if (!inputTy.hasRank()) \
2081 if (inputTy != getType()) \
2083 if (inputTy.getRank() == 0 || inputTy.getDimSize(getAxis()) == 1) \
2084 return getInput(); \
2097 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2098 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2100 if (!inputTy || !outputTy)
2106 if (inputTy == outputTy && inputTy.getNumDynamicDims() < 2)
2110 if (
auto reshapeOp = llvm::dyn_cast_if_present<tosa::ReshapeOp>(
2111 getInput1().getDefiningOp())) {
2112 getInput1Mutable().assign(reshapeOp.getInput1());
2117 if (!inputTy.getElementType().isIntOrIndexOrFloat())
2122 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2124 if (!outputTy.hasStaticShape())
2128 if (operand.isSplat())
2133 if (!getInput1().hasOneUse())
2140 return operand.reshape(
2141 llvm::cast<ShapedType>(operand.getType()).clone(shapeVec));
2147OpFoldResult PadOp::fold(FoldAdaptor adaptor) {
2149 if (adaptor.getPadding() && getInput1().
getType() ==
getType()) {
2150 auto densePad = llvm::dyn_cast<DenseElementsAttr>(adaptor.getPadding());
2151 if (densePad && densePad.isSplat() &&
2152 densePad.getSplatValue<APInt>().isZero()) {
2162OpFoldResult ResizeOp::fold(FoldAdaptor adaptor) {
2164 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getScale());
2166 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getOffset());
2168 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getBorder());
2169 if (!scaleAttr || !offsetAttr || !borderAttr) {
2176 if (scale.size() != 4 || offset.size() != 2 || border.size() != 2) {
2181 if (scale[0] != scale[1] || scale[2] != scale[3]) {
2186 if (offset[0] != 0 || offset[1] != 0) {
2191 if (border[0] != 0 || border[1] != 0) {
2195 return foldToInputIfTypeMatches(
getType(), getInput());
2198OpFoldResult ReverseOp::fold(FoldAdaptor adaptor) {
2199 auto operand = getInput1();
2200 auto operandTy = llvm::cast<ShapedType>(operand.getType());
2201 auto axis = getAxis();
2203 const bool isSplatInput =
2204 llvm::isa_and_nonnull<SplatElementsAttr>(adaptor.getInput1());
2205 if (!operandTy.hasRank() ||
2206 (!isSplatInput && operandTy.getDimSize(axis) != 1))
2208 return foldToInputIfTypeMatches(
getType(), operand);
2211OpFoldResult SliceOp::fold(FoldAdaptor adaptor) {
2212 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2213 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2215 if (!inputTy || !outputTy)
2218 if (inputTy == outputTy && inputTy.hasStaticShape())
2223 DenseElementsAttr startElems;
2229 llvm::all_of(startElems.
getValues<APInt>(),
2230 [](
const APInt &val) { return val.isZero(); });
2235 DenseElementsAttr sizeElems;
2239 auto inputShape = inputTy.getShape();
2240 auto sizeValues = sizeElems.
getValues<APInt>();
2242 bool sizeMatchesInput =
true;
2243 for (
const auto &[i, sizeVal] : llvm::enumerate(sizeValues)) {
2244 int64_t size = sizeVal.getSExtValue();
2246 if (inputTy.isDynamicDim(i)) {
2250 sizeMatchesInput =
false;
2257 sizeMatchesInput =
false;
2263 if (sizeMatchesInput)
2268 if (!adaptor.getInput1())
2272 if (!inputTy.getElementType().isIntOrIndexOrFloat() ||
2273 !outputTy.getElementType().isIntOrIndexOrFloat())
2276 auto operand = llvm::cast<ElementsAttr>(adaptor.getInput1());
2277 if (operand.isSplat() && outputTy.hasStaticShape()) {
2281 if (inputTy.hasStaticShape() && outputTy.hasStaticShape() &&
2282 outputTy.getNumElements() == 1) {
2283 llvm::SmallVector<uint64_t>
indices =
2284 llvm::to_vector(startElems.
getValues<uint64_t>());
2285 if (
auto values = operand.tryGetValues<Attribute>())
2292OpFoldResult tosa::SelectOp::fold(FoldAdaptor adaptor) {
2293 const Value pred = getPred();
2294 const Value onTrue = getOnTrue();
2295 const Value onFalse = getOnFalse();
2297 const auto predTy = llvm::dyn_cast<RankedTensorType>(pred.
getType());
2298 const auto onTrueTy = llvm::dyn_cast<RankedTensorType>(onTrue.
getType());
2299 const auto onFalseTy = llvm::dyn_cast<RankedTensorType>(onFalse.
getType());
2300 if (!predTy || !onTrueTy || !onFalseTy)
2303 const Type resultTy =
getType();
2305 const ArrayRef<int64_t> predShape = predTy.getShape();
2306 const ArrayRef<int64_t> onTrueShape = onTrueTy.getShape();
2308 if (onTrue == onFalse && onTrueTy == resultTy &&
2313 llvm::dyn_cast_if_present<DenseIntElementsAttr>(adaptor.getInput1());
2316 if (!predicate.isSplat())
2319 const bool predicateValue = predicate.getSplatValue<APInt>().getBoolValue();
2321 SmallVector<SmallVector<int64_t>, 3> shapes;
2322 shapes.emplace_back(predShape);
2323 shapes.emplace_back(onTrueShape);
2324 shapes.emplace_back(onFalseTy.getShape());
2325 const bool isBroadcastable =
2328 if (predicateValue ==
true && onTrueTy == resultTy && isBroadcastable)
2330 if (predicateValue ==
false && onFalseTy == resultTy && isBroadcastable)
2336 const auto inputType =
2337 dyn_cast<RankedTensorType>(tileOp.getInput1().getType());
2338 const auto outputType = dyn_cast<RankedTensorType>(tileOp.getType());
2339 if (!inputType || !outputType)
2343 if (failed(tileOp.getConstantMultiples(multiples)))
2346 for (
const auto [
index, multiple] : llvm::enumerate(multiples)) {
2349 if (outputType.isDynamicDim(
index))
2351 if (inputType.getDimSize(
index) != 1)
2364 Value tileOutput = tileOp.getOutput();
2367 "tile output must have one use");
2370 const bool isBinaryElementwise =
2373 if (!isBinaryElementwise && !isa<tosa::MulOp>(user))
2375 tileOp,
"consumer must be binary broadcastable");
2379 tileOp,
"tile must only expand statically-known singleton dims");
2383 Value otherOperand = lhsOperand == tileOutput ? rhsOperand : lhsOperand;
2384 Value tileInput = tileOp.getInput1();
2386 const ShapedType newOtherType = cast<ShapedType>(otherOperand.
getType());
2387 const ShapedType newTileType = cast<ShapedType>(tileInput.
getType());
2390 newOtherType.getShape(), newTileType.getShape(), broadcastedShape);
2392 const ShapedType outputType = cast<ShapedType>(user->
getResultTypes()[0]);
2393 if (!llvm::equal(broadcastedShape, outputType.getShape()))
2395 tileOp,
"tile output must be broadcastable to consumer operands");
2399 mapper.
map(tileOutput, tileOp.getInput1());
2406void TileOp::getCanonicalizationPatterns(RewritePatternSet &results,
2407 MLIRContext *context) {
2408 results.
add<RemoveBroadcastTileFromBinaryElementwise>(context);
2411OpFoldResult TileOp::fold(FoldAdaptor adaptor) {
2413 if (
auto multiples = llvm::dyn_cast_if_present<DenseElementsAttr>(
2414 adaptor.getMultiples())) {
2415 if (multiples.isSplat() &&
2416 multiples.getSplatValue<APInt>().getSExtValue() == 1)
2418 if (
auto int_array_attr =
2419 llvm::dyn_cast<DenseIntElementsAttr>(multiples)) {
2420 if (llvm::all_of(int_array_attr.getValues<APInt>(),
2421 [](APInt v) { return v.getSExtValue() == 1; }))
2429OpFoldResult TransposeOp::fold(FoldAdaptor adaptor) {
2430 auto resultTy = llvm::cast<ShapedType>(
getType());
2434 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2435 if (input.isSplat() && resultTy.hasRank() && resultTy.hasStaticShape() &&
2436 input.
getType().getElementType() == resultTy.getElementType())
2437 return input.reshape(resultTy);
2441 const llvm::ArrayRef<int32_t> perms = getPerms();
2443 if (!llvm::equal(llvm::seq<int32_t>(0, perms.size()), perms))
2446 return foldToInputIfTypeMatches(
getType(), getInput1());
2449OpFoldResult tosa::NegateOp::fold(FoldAdaptor adaptor) {
2452 auto definingOp = getInput1().getDefiningOp<tosa::NegateOp>();
2458 if (FailureOr<int64_t> maybeIZp = getInput1ZeroPoint();
2459 failed(maybeIZp) || *maybeIZp != 0) {
2463 if (FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
2464 failed(maybeOZp) || *maybeOZp != 0) {
2468 if (FailureOr<int64_t> maybeIZp = definingOp.getInput1ZeroPoint();
2469 failed(maybeIZp) || *maybeIZp != 0) {
2473 if (FailureOr<int64_t> maybeOZp = definingOp.getOutputZeroPoint();
2474 failed(maybeOZp) || *maybeOZp != 0) {
2479 return foldToInputIfTypeMatches(
getType(), definingOp.getInput1());
2482OpFoldResult tosa::AbsOp::fold(FoldAdaptor adaptor) {
2483 auto input = getInput1();
2486 return foldToInputIfTypeMatches(
getType(), input);
2491OpFoldResult tosa::ReciprocalOp::fold(FoldAdaptor adaptor) {
2492 auto input = adaptor.getInput1();
2494 auto inputAttr = llvm::dyn_cast_if_present<DenseElementsAttr>(input);
2496 if (!inputAttr || !inputAttr.isSplat())
2499 auto shapeType = llvm::cast<ShapedType>(
getType());
2500 if (!shapeType.hasRank() || !shapeType.hasStaticShape())
2502 if (
auto floatType = llvm::dyn_cast<FloatType>(inputAttr.getElementType())) {
2503 auto floatVal = inputAttr.getSplatValue<APFloat>();
2505 ReciprocalOp::calcOneElement(floatVal));
2511template <
typename Op,
typename OpFoldAdaptor>
2513 auto input1ConstShape =
2514 dyn_cast<tosa::ConstShapeOp>(op->getInput().getDefiningOp());
2515 if (!input1ConstShape)
2518 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2524template <
typename Op,
typename OpFoldAdaptor>
2526 auto input1ConstShape =
2527 dyn_cast<tosa::ConstShapeOp>(op->getInput1().getDefiningOp());
2528 auto input2ConstShape =
2529 dyn_cast<tosa::ConstShapeOp>(op->getInput2().getDefiningOp());
2530 if (!input1ConstShape || !input2ConstShape)
2533 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2534 const auto input2Attr = cast<DenseElementsAttr>(input2ConstShape.getValues());
2537 input1Attr.getType(),
2541OpFoldResult tosa::DimOp::fold(FoldAdaptor adaptor) {
2542 const auto inputTy = llvm::dyn_cast<ShapedType>(getInput1().
getType());
2543 if (!inputTy || !inputTy.hasRank())
2545 const int32_t axis = getAxis();
2546 const int64_t dimSize = inputTy.getDimSize(axis);
2547 if (ShapedType::isDynamic(dimSize))
2551 const auto resultAttrTy =
2552 RankedTensorType::get(1, builder.getIndexType());
2557 auto const inputs = op->getInput();
2563 concatDims.reserve( 64);
2564 for (
auto const &v : inputs) {
2565 auto vConstShape = dyn_cast<tosa::ConstShapeOp>(v.getDefiningOp());
2569 const auto vAttr = cast<DenseElementsAttr>(vConstShape.getValues());
2572 auto const vAttrVals = vAttr.getValues<APInt>();
2573 for (
auto const &v : vAttrVals) {
2574 concatDims.push_back(v);
2578 auto *ctx = op->getContext();
2579 assert(ctx !=
nullptr &&
"ctx is nullptr");
2580 auto const rankedTy = RankedTensorType::get(
2581 {
static_cast<int64_t>(concatDims.size())}, IndexType::get(ctx));
2587 auto const input1 = op->getInput();
2588 auto const input2 = op->getStart();
2589 auto const input3 = op->getSize();
2591 auto input1ConstShape = dyn_cast<tosa::ConstShapeOp>(input1.getDefiningOp());
2593 if (!input1ConstShape)
2596 auto const input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2600 auto const input1Vals = input1Attr.getValues<APInt>();
2601 auto const totalInput1 = input1Vals.size();
2606 if (failed(start) || failed(size))
2609 auto const startV =
static_cast<int32_t
>(start.value());
2610 auto const sizeV =
static_cast<int32_t
>(size.value());
2612 if ((sizeV <= 0) || (startV < 0) ||
2613 (
static_cast<size_t>(startV + sizeV) > totalInput1))
2617 sliceOfInput.reserve(totalInput1);
2619 for (
auto i = startV; i < (startV + sizeV); i++) {
2620 sliceOfInput.push_back(input1Vals[i]);
2623 auto *ctx = op->getContext();
2624 assert(ctx !=
nullptr &&
"ctx is nullptr");
2626 auto const rankedTy = RankedTensorType::get(
2627 {
static_cast<int64_t>(sliceOfInput.size())}, IndexType::get(ctx));
2632OpFoldResult tosa::AddShapeOp::fold(FoldAdaptor adaptor) {
2636OpFoldResult tosa::SubShapeOp::fold(FoldAdaptor adaptor) {
2640OpFoldResult tosa::MulShapeOp::fold(FoldAdaptor adaptor) {
2644OpFoldResult tosa::DivCeilShapeOp::fold(FoldAdaptor adaptor) {
2645 return binaryFold<DivCeilShapeOp, ShapeDivFoldAdaptor<
true>>(
this);
2648OpFoldResult tosa::DivFloorShapeOp::fold(FoldAdaptor adaptor) {
2649 return binaryFold<DivFloorShapeOp, ShapeDivFoldAdaptor<
false>>(
this);
2652OpFoldResult tosa::ModShapeOp::fold(FoldAdaptor adaptor) {
2656OpFoldResult tosa::MaxShapeOp::fold(FoldAdaptor adaptor) {
2660OpFoldResult tosa::MinShapeOp::fold(FoldAdaptor adaptor) {
2664OpFoldResult tosa::Exp2ShapeOp::fold(FoldAdaptor adaptor) {
2668OpFoldResult tosa::Log2CeilShapeOp::fold(FoldAdaptor adaptor) {
2672OpFoldResult tosa::Log2FloorShapeOp::fold(FoldAdaptor adaptor) {
2676OpFoldResult tosa::ConcatShapeOp::fold(FoldAdaptor adaptor) {
2680OpFoldResult tosa::SliceShapeOp::fold(FoldAdaptor adaptor) {
static bool isSplatZero(Type elemType, DenseElementsAttr val)
Returns true if 'val' is a splat of zero, false otherwise.
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
#define REDUCE_FOLDER(OP)
OpFoldResult concatShapeFold(tosa::ConcatShapeOp *op)
static DenseElementsAttr binaryFolder(DenseElementsAttr lhs, DenseElementsAttr rhs, ShapedType returnTy, bool foldDenseValues=false)
static DenseElementsAttr unaryFolder(DenseElementsAttr val, ShapedType returnTy, bool foldDenseValues=false)
static LogicalResult verifyTileIsBroadcast(tosa::TileOp tileOp)
OpFoldResult sliceShapeFold(tosa::SliceShapeOp *op)
static FailureOr< int64_t > getSingleI64From1ElementTensor(Value v)
OpFoldResult binaryFold(Op *op)
static bool isSplatOne(Type elemType, DenseElementsAttr val, int64_t shift)
OpFoldResult unaryShapeFold(Op *op)
static bool checkMatchingPadConstAndZp(Value padConst, Value zp)
static bool signsDiffer(const APInt &a, const APInt &b)
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Attributes are known-constant values of operations.
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
IntegerAttr getIntegerAttr(Type type, int64_t value)
DenseI64ArrayAttr getDenseI64ArrayAttr(ArrayRef< int64_t > values)
FloatAttr getFloatAttr(Type type, double value)
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
MLIRContext * getContext() const
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.
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
int64_t size() const
Returns the number of elements held by this attribute.
bool isSplat() const
Returns true if this attribute corresponds to a splat, i.e.
Type getElementType() const
Return the element type of this DenseElementsAttr.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
ShapedType getType() const
Return the type of this ElementsAttr, guaranteed to be a vector or tensor with static shape.
An attribute that represents a reference to a dense integer vector or tensor object.
iterator begin() const
Iterator access to the integer element values.
This is a utility class for mapping one set of IR entities to another.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
MLIRContext is the top-level object for a collection of MLIR operations.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents a single result from folding an operation.
This class indicates that an op is tosa-elementwise (permits broadcasting, unlike Elementwise trait).
This provides public APIs that all operations should have.
This class implements the operand iterators for the Operation class.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
unsigned getNumOperands()
result_type_range getResultTypes()
result_range getResults()
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
bool isInteger() const
Return true if this is an integer type (with the specified width).
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.
user_iterator user_begin() const
bool hasOneUse() const
Returns true if this value has exactly one use.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
bool staticallyKnownBroadcastable(ArrayRef< SmallVector< int64_t, 6 > > shapes)
Returns true if a broadcast between n shapes is guaranteed to be successful and not result in an erro...
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...
DynamicAPInt round(const Fraction &f)
TosaLevel getTosaLevelFromEnum(const Level level)
constexpr int64_t kInferableDimSize
Represents a dimension in the shape of a tensor that can be inferred based on the other provided dime...
SmallVector< int64_t > convertFromIntAttr(const DenseElementsAttr &attr, const int rank)
TargetEnvAttr lookupTargetEnv(Operation *op)
FailureOr< T > getConstantScalarIntValue(Value val)
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
Type getStorageElementTypeFromQuantized(quant::QuantizedType quantizedType)
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.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
LogicalResult matchAndRewrite(tosa::AvgPool2dAdaptiveOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::AvgPool2dOp op, PatternRewriter &rewriter) const override
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
LogicalResult matchAndRewrite(tosa::CastOp castOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::CastToBlockScaledOp castToBlockScaledOp, PatternRewriter &rewriter) const override
bool intersects(const ClampRange< T > &otherRange)
ClampRange(const T &start, const T &end)
LogicalResult matchAndRewrite(tosa::ClampOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::ClampOp op, PatternRewriter &rewriter) const override
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
LogicalResult matchAndRewrite(tosa::ConcatOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::SliceOp sliceOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::ConcatOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::TransposeOp transposeOp, PatternRewriter &rewriter) const override
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
static FailureOr< APInt > fold(const APInt &value, bool isUnsigned)
static FailureOr< APInt > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
static FailureOr< APInt > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &value, bool isUnsigned)
static FailureOr< APInt > fold(const APInt &value, bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, bool isUnsigned)
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
LogicalResult matchAndRewrite(tosa::MaxPool2dAdaptiveOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::MaxPool2dOp op, PatternRewriter &rewriter) const override
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
bool isNarrowingCast(const ShapedType inType, const ShapedType outType) const
LogicalResult matchAndRewrite(tosa::CastOp castOp, PatternRewriter &rewriter) const override
bool supportsInf(const llvm::fltSemantics &semantics) const
bool supportsNaN(const llvm::fltSemantics &semantics) const
LogicalResult matchAndRewrite(tosa::SliceOp sliceOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::TileOp tileOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(tosa::RowGatherOp op, PatternRewriter &rewriter) const override
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, bool isUnsigned)
LogicalResult matchAndRewrite(tosa::SliceOp sliceOp, PatternRewriter &rewriter) const override
static FailureOr< APFloat > fold(const APFloat &lhs, const APFloat &rhs)
static FailureOr< APInt > fold(const APInt &lhs, const APInt &rhs, const bool isUnsigned)
LogicalResult matchAndRewrite(tosa::TransposeOp op, PatternRewriter &rewriter) const override
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...
int32_t MAX_TENSOR_LIST_SIZE