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");
1175 if (!innerCastOp.getInputUnsigned() && castOp.getInputUnsigned()) {
1177 castOp,
"avoid rewriting cast(input_unsigned=false) -> "
1178 "cast(input_unsigned=true)");
1183 innerCastOp.getInputUnsigned());
1189 return semantics.nonFiniteBehavior !=
1190 llvm::fltNonfiniteBehavior::FiniteOnly;
1194 return semantics.nonFiniteBehavior == llvm::fltNonfiniteBehavior::IEEE754;
1198 const ShapedType outType)
const {
1200 if (inType.getElementType().isInteger() &&
1201 outType.getElementType().isInteger()) {
1203 const auto inTypeSignedness =
1204 cast<IntegerType>(inType.getElementType()).getSignedness();
1205 const auto outTypeSignedness =
1206 cast<IntegerType>(outType.getElementType()).getSignedness();
1208 return (inTypeSignedness != outTypeSignedness ||
1209 inType.getElementTypeBitWidth() >
1210 outType.getElementTypeBitWidth());
1213 if (inType.getElementType().isFloat() &&
1214 outType.getElementType().isFloat()) {
1216 FloatType inElemTy = cast<FloatType>(inType.getElementType());
1217 FloatType outElemTy = cast<FloatType>(outType.getElementType());
1218 llvm::fltSemantics inTypeSemantics = inElemTy.getFloatSemantics();
1219 llvm::fltSemantics outTypeSemantics = outElemTy.getFloatSemantics();
1225 [[maybe_unused]]
const auto isSupported = [](
Type elemType) {
1226 return llvm::isa<Float8E4M3FNType, Float8E5M2Type, BFloat16Type,
1227 Float16Type, Float32Type>(elemType);
1230 assert(isSupported(inElemTy) &&
1231 "unsupported input element type in isNarrowingCast");
1232 assert(isSupported(outElemTy) &&
1233 "unsupported output element type in isNarrowingCast");
1236 inTypeSemantics.maxExponent > outTypeSemantics.maxExponent ||
1237 inTypeSemantics.minExponent < outTypeSemantics.minExponent ||
1238 inTypeSemantics.precision > outTypeSemantics.precision ||
1255 const Value outerInput = castOp.getInput();
1256 auto innerCastOp = outerInput.
getDefiningOp<tosa::CastOp>();
1259 "input must be a cast operation");
1261 const Value innerInput = innerCastOp.getInput();
1262 const auto innerInputTy = llvm::cast<ShapedType>(innerInput.
getType());
1263 const auto innerOutputTy = llvm::cast<ShapedType>(innerCastOp.getType());
1264 const auto outerOutputTy = llvm::cast<ShapedType>(castOp.getType());
1266 if (!llvm::isa<tosa::BlockScaledType>(innerInputTy.getElementType()))
1268 castOp,
"inner cast input must have block scaled element type");
1270 if (innerInputTy != outerOutputTy)
1272 castOp,
"inner input type must match outer output type");
1274 const Type innerOutputElemType = innerOutputTy.getElementType();
1275 const bool isLosslessCast =
1276 isa<Float32Type, BFloat16Type>(innerOutputElemType);
1277 if (!isLosslessCast)
1279 castOp,
"avoid cancelling casts that should be lossy");
1287void CastOp::getCanonicalizationPatterns(RewritePatternSet &results,
1288 MLIRContext *context) {
1289 results.
add<NonNarrowingCastsOptimization,
1290 CancellingBlockScaledCastsOptimization>(context);
1299 const Value castToBlockScaledInput = castToBlockScaledOp.getInputData();
1300 auto castFromBlockScaledOp =
1301 castToBlockScaledInput.
getDefiningOp<tosa::CastFromBlockScaledOp>();
1302 if (!castFromBlockScaledOp)
1304 castToBlockScaledOp,
1305 "input must be cast_from_block_scaled operation");
1307 const Value innerData = castFromBlockScaledOp.getInputData();
1308 const Value innerScale = castFromBlockScaledOp.getInputScale();
1309 const auto innerDataTy = llvm::cast<ShapedType>(innerData.
getType());
1310 const auto innerScaleTy = llvm::cast<ShapedType>(innerScale.
getType());
1312 const Value outerData = castToBlockScaledOp.getOutputData();
1313 const Value outerScale = castToBlockScaledOp.getOutputScale();
1314 const auto outerDataTy = llvm::cast<ShapedType>(outerData.
getType());
1315 const auto outerScaleTy = llvm::cast<ShapedType>(outerScale.
getType());
1317 if (innerDataTy != outerDataTy || innerScaleTy != outerScaleTy) {
1319 castToBlockScaledOp,
1320 "inputs types to cast_from_block_scaled operation must match output "
1321 "types to cast_to_block_scaled");
1324 if (castFromBlockScaledOp.getBlockSize() !=
1325 castToBlockScaledOp.getBlockSize()) {
1327 castToBlockScaledOp,
"block sizes for cast_from_block_scaled and "
1328 "cast_to_block_scaled must match");
1331 rewriter.
replaceOp(castToBlockScaledOp, {innerData, innerScale});
1337void CastToBlockScaledOp::getCanonicalizationPatterns(
1338 RewritePatternSet &results, MLIRContext *context) {
1339 results.
add<CancellingCastToFromBlockScaledOptimization>(context);
1347 const FailureOr<int32_t> rowCount =
1349 if (failed(rowCount) || rowCount.value() != 1)
1353 op, op.getOutput().
getType(), op.getValues(), op.getIndices());
1358void RowGatherOp::getCanonicalizationPatterns(RewritePatternSet &results,
1359 MLIRContext *context) {
1360 results.
add<RowGatherToGather>(context);
1367template <
typename Folder>
1368static DenseElementsAttr
1370 bool foldDenseValues =
false) {
1374 if (!returnTy.hasRank() || !returnTy.hasStaticShape())
1378 const auto rETy = llvm::cast<ShapedType>(
rhs.getType()).getElementType();
1382 if (
lhs.isSplat() &&
rhs.isSplat()) {
1383 if (isa<FloatType>(lETy)) {
1384 const APFloat l =
lhs.getSplatValue<APFloat>();
1385 const APFloat r =
rhs.getSplatValue<APFloat>();
1386 const auto maybeResult = Folder::fold(l, r);
1387 if (failed(maybeResult))
1392 if (
const auto lIntTy = llvm::dyn_cast<IntegerType>(lETy)) {
1393 const APInt l =
lhs.getSplatValue<APInt>();
1394 const APInt r =
rhs.getSplatValue<APInt>();
1395 const auto maybeResult = Folder::fold(l, r, lIntTy.isUnsigned());
1396 if (failed(maybeResult))
1402 if (foldDenseValues) {
1403 assert(lETy.isIntOrIndex() &&
1404 "Only integer types are currently supported.");
1407 llvm::zip(
lhs.getValues<APInt>(),
rhs.getValues<APInt>())) {
1408 const auto maybeResult = Folder::fold(l, r,
false);
1409 if (failed(maybeResult))
1411 resultValues.push_back(maybeResult.value());
1419template <
typename Folder>
1421 bool foldDenseValues =
false) {
1425 if (!returnTy.hasRank() || !returnTy.hasStaticShape())
1431 if (
const auto vIntTy = llvm::dyn_cast<IntegerType>(vETy)) {
1433 const auto maybeResult = Folder::fold(v, vIntTy.isUnsigned());
1434 if (failed(maybeResult))
1440 if (foldDenseValues) {
1444 for (
auto const &v : val.
getValues<APInt>()) {
1445 const auto maybeResult = Folder::fold(v,
false);
1446 if (failed(maybeResult))
1448 resultValues.push_back(maybeResult.value());
1463 assert(dense.isSplat());
1464 APInt a = dense.getSplatValue<APInt>();
1465 return a.getSExtValue();
1470 const bool isUnsigned) {
1473 isUnsigned ?
lhs.uadd_ov(
rhs, overflow) :
lhs.sadd_ov(
rhs, overflow);
1479 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1486 const bool isUnsigned) {
1489 isUnsigned ?
lhs.usub_ov(
rhs, overflow) :
lhs.ssub_ov(
rhs, overflow);
1495 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1502 const bool isUnsigned) {
1504 const unsigned originalWidth =
lhs.getBitWidth();
1507 if (
lhs.getBitWidth() !=
rhs.getBitWidth()) {
1512 if (
lhs == 0 ||
rhs == 0)
1513 return APInt::getZero(originalWidth);
1515 bool overflow =
false;
1517 isUnsigned ?
lhs.umul_ov(
rhs, overflow) :
lhs.smul_ov(
rhs, overflow);
1522 return result.trunc(originalWidth);
1525 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1531 return a.isNegative() !=
b.isNegative();
1538 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1546 APInt::udivrem(
lhs,
rhs, q, r);
1547 if (!r.isZero() && Ceil) {
1554 bool overflow{
false};
1555 APInt
const q =
lhs.sdiv_ov(
rhs, overflow);
1558 APInt
const r =
lhs.srem(
rhs);
1568 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1576 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1578 if (
lhs.isNegative() || (!
rhs.isStrictlyPositive()))
1588 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1590 auto const r = t.mod(
rhs);
1591 if (llvm::APFloatBase::opStatus::opOK == r) {
1601 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1603 return lhs.getSExtValue() >=
rhs.getSExtValue() ?
lhs :
rhs;
1606 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1614 if (
lhs.getBitWidth() !=
rhs.getBitWidth())
1616 return lhs.getSExtValue() <=
rhs.getSExtValue() ?
lhs :
rhs;
1619 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1625 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1626 auto const numBits = value.getBitWidth();
1628 auto const zextv = value.getZExtValue();
1629 if (zextv >= numBits)
1631 return APInt::getOneBitSet(numBits, zextv);
1633 auto const sextv = value.getSExtValue();
1634 if (sextv < 0 || sextv >= numBits || (value.isNegative()))
1636 return APInt::getOneBitSet(numBits, sextv);
1646 assert(!isUnsigned &&
1647 "unsigned values are not supported for shape div folders");
1648 if (
lhs.isNegative() || !
rhs.isStrictlyPositive())
1653 static FailureOr<APFloat>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1659 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1660 if (!value.isStrictlyPositive())
1662 return APInt(value.getBitWidth(), value.ceilLogBase2());
1667 static FailureOr<APInt>
fold(
const APInt &value,
bool isUnsigned) {
1668 if (!value.isStrictlyPositive())
1670 return APInt(value.getBitWidth(), value.logBase2());
1676 const bool isUnsigned) {
1677 return isUnsigned ? APInt(1,
lhs.ugt(
rhs)) : APInt(1,
lhs.sgt(
rhs));
1680 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1681 return APInt(1,
lhs >
rhs);
1687 const bool isUnsigned) {
1688 return isUnsigned ? APInt(1,
lhs.uge(
rhs)) : APInt(1,
lhs.sge(
rhs));
1691 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1692 return APInt(1,
lhs >=
rhs);
1698 const bool isUnsigned) {
1699 return APInt(1,
lhs ==
rhs);
1702 static FailureOr<APInt>
fold(
const APFloat &
lhs,
const APFloat &
rhs) {
1703 return APInt(1,
lhs ==
rhs);
1708 if (llvm::isa<FloatType>(elemType))
1710 if (llvm::isa<IntegerType>(elemType))
1716 if (llvm::isa<FloatType>(elemType))
1717 return val && val.
isSplat() &&
1719 if (llvm::isa<IntegerType>(elemType)) {
1720 const int64_t shifted = 1LL << shift;
1721 return val && val.
isSplat() &&
1727OpFoldResult AddOp::fold(FoldAdaptor adaptor) {
1728 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1729 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1730 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1731 if (!lhsTy || !rhsTy || !resultTy)
1735 if (!lhsTy.getElementType().isIntOrIndexOrFloat() ||
1736 !rhsTy.getElementType().isIntOrIndexOrFloat())
1739 auto resultETy = resultTy.getElementType();
1741 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1743 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1746 lhsTy.getShape(), rhsTy.getShape());
1747 if (isBroadcastable && lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr))
1749 if (isBroadcastable && rhsTy == resultTy &&
isSplatZero(resultETy, lhsAttr))
1752 if (!lhsAttr || !rhsAttr)
1758OpFoldResult ArgMaxOp::fold(FoldAdaptor adaptor) {
1759 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput().
getType());
1760 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
1761 if (!inputTy || !outputTy || !inputTy.hasStaticShape() ||
1762 !outputTy.hasStaticShape())
1766 if (inputTy.getDimSize(getAxis()) == 1 && outputElementTy.
isInteger()) {
1767 const auto outputElemIntTy = cast<IntegerType>(outputElementTy);
1768 const APInt zero = APInt::getZero(outputElemIntTy.getWidth());
1775OpFoldResult IntDivOp::fold(FoldAdaptor adaptor) {
1776 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1777 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1778 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1779 if (!lhsTy || !rhsTy || !resultTy)
1781 if (lhsTy.getElementType() != rhsTy.getElementType())
1786 auto resultETy = resultTy.getElementType();
1788 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1790 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1791 if (lhsAttr && lhsAttr.isSplat() && rhsAttr && rhsAttr.isSplat()) {
1792 if (llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape() &&
1793 lhsAttr.getSplatValue<APInt>().isZero() &&
1794 !rhsAttr.getSplatValue<APInt>().isZero()) {
1795 return lhsAttr.resizeSplat(resultTy);
1799 if (rhsAttr && rhsAttr.isSplat()) {
1801 lhsTy.getShape(), rhsTy.getShape());
1802 if (isBroadcastable && lhsTy == resultTy &&
1803 llvm::isa<IntegerType>(resultETy) &&
1804 rhsAttr.getSplatValue<APInt>().isOne())
1808 if (rhsAttr && lhsAttr && rhsAttr.isSplat() && lhsAttr.isSplat() &&
1809 llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape()) {
1810 APInt l = lhsAttr.getSplatValue<APInt>();
1811 APInt r = rhsAttr.getSplatValue<APInt>();
1813 auto intTy = dyn_cast<mlir::IntegerType>(resultETy);
1815 DivFoldAdaptor<
false>::fold(l, r, intTy.isUnsigned());
1828std::optional<APInt> mulInt(APInt
lhs, APInt
rhs, int32_t shift,
1829 unsigned bitwidth) {
1830 bool overflow =
false;
1831 APInt
result =
lhs.sext(64).smul_ov(
rhs.sext(64), overflow);
1834 return std::nullopt;
1837 auto round = APInt(64, 1) << (shift - 1);
1839 result.ashrInPlace(shift);
1842 if (!(
result.getSExtValue() >= INT32_MIN &&
1843 result.getSExtValue() <= INT32_MAX)) {
1845 return std::nullopt;
1849 return result.trunc(bitwidth);
1852DenseElementsAttr mulBinaryFolder(DenseElementsAttr
lhs, DenseElementsAttr
rhs,
1853 RankedTensorType ty, int32_t shift) {
1855 if (!ty.hasStaticShape())
1858 if (llvm::isa<IntegerType>(ty.getElementType())) {
1859 APInt l =
lhs.getSplatValue<APInt>();
1860 APInt r =
rhs.getSplatValue<APInt>();
1866 auto bitwidth = ty.getElementType().getIntOrFloatBitWidth();
1867 const std::optional<APInt>
result = mulInt(l, r, shift, bitwidth);
1873 if (llvm::isa<FloatType>(ty.getElementType())) {
1874 APFloat l =
lhs.getSplatValue<APFloat>();
1875 APFloat r =
rhs.getSplatValue<APFloat>();
1885OpFoldResult MulOp::fold(FoldAdaptor adaptor) {
1886 auto lhs = getInput1();
1887 auto rhs = getInput2();
1888 auto lhsTy = llvm::dyn_cast<RankedTensorType>(
lhs.getType());
1889 auto rhsTy = llvm::dyn_cast<RankedTensorType>(
rhs.getType());
1890 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1891 if (!lhsTy || !rhsTy || !resultTy)
1894 auto resultETy = resultTy.getElementType();
1896 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1898 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1903 if (resultETy.isInteger(32)) {
1904 ElementsAttr shift_elem;
1905 if (getShift().getImpl()) {
1909 shift = shift_elem.getValues<IntegerAttr>()[0].getInt();
1913 if (rhsTy == resultTy &&
isSplatZero(resultETy, lhsAttr) &&
1914 resultTy.hasStaticShape())
1916 return lhsAttr.resizeSplat(resultTy);
1917 if (lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr) &&
1918 resultTy.hasStaticShape())
1919 return rhsAttr.resizeSplat(resultTy);
1922 lhsTy.getShape(), rhsTy.getShape());
1923 if (isBroadcastable && rhsTy == resultTy &&
1926 if (isBroadcastable && lhsTy == resultTy &&
1930 return mulBinaryFolder(lhsAttr, rhsAttr, resultTy, shift);
1933OpFoldResult SubOp::fold(FoldAdaptor adaptor) {
1934 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1935 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1936 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1937 if (!lhsTy || !rhsTy || !resultTy)
1941 if (!lhsTy.getElementType().isIntOrIndexOrFloat() ||
1942 !rhsTy.getElementType().isIntOrIndexOrFloat())
1945 auto resultETy = resultTy.getElementType();
1947 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1949 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1952 lhsTy.getShape(), rhsTy.getShape());
1953 if (isBroadcastable && lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr))
1956 if (!lhsAttr || !rhsAttr)
1962OpFoldResult GreaterOp::fold(FoldAdaptor adaptor) {
1963 auto resultTy = llvm::cast<ShapedType>(
getType());
1965 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1967 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1969 if (!lhsAttr || !rhsAttr)
1975OpFoldResult GreaterEqualOp::fold(FoldAdaptor adaptor) {
1976 auto resultTy = llvm::cast<ShapedType>(
getType());
1978 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1980 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1982 if (!lhsAttr || !rhsAttr)
1988OpFoldResult EqualOp::fold(FoldAdaptor adaptor) {
1989 auto resultTy = llvm::cast<ShapedType>(
getType());
1991 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1993 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1994 Value
lhs = getInput1();
1995 Value
rhs = getInput2();
1996 auto lhsTy = llvm::cast<ShapedType>(
lhs.getType());
2000 if (llvm::isa<IntegerType>(lhsTy.getElementType()) && resultTy.hasRank() &&
2001 resultTy.hasStaticShape() &&
lhs ==
rhs) {
2005 if (!lhsAttr || !rhsAttr)
2011OpFoldResult CastOp::fold(FoldAdaptor adaptor) {
2015 auto operand = llvm::dyn_cast_if_present<ElementsAttr>(adaptor.getInput());
2019 auto inTy = llvm::cast<ShapedType>(getInput().
getType());
2020 auto outTy = llvm::cast<ShapedType>(
getType());
2021 if (!outTy.hasRank() || !outTy.hasStaticShape())
2023 auto inETy = inTy.getElementType();
2024 auto outETy = outTy.getElementType();
2026 if (operand.isSplat()) {
2027 if (llvm::isa<FloatType>(inETy) && llvm::isa<FloatType>(outETy)) {
2029 auto splatVal = operand.getSplatValue<APFloat>();
2030 auto &semantics = llvm::cast<FloatType>(outETy).getFloatSemantics();
2031 splatVal.convert(semantics, llvm::RoundingMode::NearestTiesToEven,
2036 if (llvm::isa<IntegerType>(inETy) && llvm::isa<FloatType>(outETy)) {
2037 const bool unsign = llvm::cast<IntegerType>(inETy).isUnsignedInteger() ||
2038 adaptor.getInputUnsigned();
2039 APFloat splatVal(llvm::cast<FloatType>(outETy).getFloatSemantics());
2040 splatVal.convertFromAPInt(operand.getSplatValue<APInt>(), !unsign,
2041 llvm::RoundingMode::NearestTiesToEven);
2045 if (llvm::isa<FloatType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2046 auto unsign = llvm::cast<IntegerType>(outETy).isUnsignedInteger();
2047 auto intVal = APSInt(
2048 llvm::cast<IntegerType>(outETy).getIntOrFloatBitWidth(), unsign);
2049 auto floatVal = operand.getSplatValue<APFloat>();
2051 floatVal.convertToInteger(intVal, llvm::RoundingMode::NearestTiesToEven,
2056 if (llvm::isa<IntegerType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2057 const auto inIntType = llvm::cast<IntegerType>(inETy);
2058 const bool unsignIn =
2059 inIntType.isUnsignedInteger() || adaptor.getInputUnsigned();
2061 inETy.getIntOrFloatBitWidth() > outETy.getIntOrFloatBitWidth();
2062 auto intVal = operand.getSplatValue<APInt>();
2063 auto bitwidth = outETy.getIntOrFloatBitWidth();
2066 if (outETy.isInteger(1)) {
2067 intVal = APInt(bitwidth, intVal.isZero() ? 0 : 1);
2069 intVal = intVal.trunc(bitwidth);
2070 }
else if (unsignIn || inIntType.isInteger(1)) {
2071 intVal = intVal.zext(bitwidth);
2073 intVal = intVal.sext(bitwidth);
2083OpFoldResult ConstOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2085OpFoldResult ConstShapeOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2087#define REDUCE_FOLDER(OP) \
2088 OpFoldResult OP::fold(FoldAdaptor adaptor) { \
2089 ShapedType inputTy = llvm::cast<ShapedType>(getInput().getType()); \
2090 if (!inputTy.hasRank()) \
2092 if (inputTy != getType()) \
2094 if (inputTy.getRank() == 0 || inputTy.getDimSize(getAxis()) == 1) \
2095 return getInput(); \
2108 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2109 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2111 if (!inputTy || !outputTy)
2117 if (inputTy == outputTy && inputTy.getNumDynamicDims() < 2)
2121 if (
auto reshapeOp = llvm::dyn_cast_if_present<tosa::ReshapeOp>(
2122 getInput1().getDefiningOp())) {
2123 getInput1Mutable().assign(reshapeOp.getInput1());
2128 if (!inputTy.getElementType().isIntOrIndexOrFloat())
2132 if (!outputTy.hasStaticShape())
2136 if (
auto operand = llvm::dyn_cast_if_present<DenseResourceElementsAttr>(
2137 adaptor.getInput1()))
2138 return DenseResourceElementsAttr::get(outputTy, operand.getRawHandle());
2142 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2144 if (operand.isSplat())
2149 if (!getInput1().hasOneUse())
2156 return operand.reshape(
2157 llvm::cast<ShapedType>(operand.getType()).clone(shapeVec));
2163OpFoldResult PadOp::fold(FoldAdaptor adaptor) {
2165 if (adaptor.getPadding() && getInput1().
getType() ==
getType()) {
2166 auto densePad = llvm::dyn_cast<DenseElementsAttr>(adaptor.getPadding());
2167 if (densePad && densePad.isSplat() &&
2168 densePad.getSplatValue<APInt>().isZero()) {
2178OpFoldResult ResizeOp::fold(FoldAdaptor adaptor) {
2180 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getScale());
2182 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getOffset());
2184 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getBorder());
2185 if (!scaleAttr || !offsetAttr || !borderAttr) {
2192 if (scale.size() != 4 || offset.size() != 2 || border.size() != 2) {
2197 if (scale[0] != scale[1] || scale[2] != scale[3]) {
2202 if (offset[0] != 0 || offset[1] != 0) {
2207 if (border[0] != 0 || border[1] != 0) {
2211 return foldToInputIfTypeMatches(
getType(), getInput());
2214OpFoldResult ReverseOp::fold(FoldAdaptor adaptor) {
2215 auto operand = getInput1();
2216 auto operandTy = llvm::cast<ShapedType>(operand.getType());
2217 auto axis = getAxis();
2219 const bool isSplatInput =
2220 llvm::isa_and_nonnull<SplatElementsAttr>(adaptor.getInput1());
2221 if (!operandTy.hasRank() ||
2222 (!isSplatInput && operandTy.getDimSize(axis) != 1))
2224 return foldToInputIfTypeMatches(
getType(), operand);
2227OpFoldResult SliceOp::fold(FoldAdaptor adaptor) {
2228 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2229 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2231 if (!inputTy || !outputTy)
2234 if (inputTy == outputTy && inputTy.hasStaticShape())
2239 DenseElementsAttr startElems;
2245 llvm::all_of(startElems.
getValues<APInt>(),
2246 [](
const APInt &val) { return val.isZero(); });
2251 DenseElementsAttr sizeElems;
2255 auto inputShape = inputTy.getShape();
2256 auto sizeValues = sizeElems.
getValues<APInt>();
2258 bool sizeMatchesInput =
true;
2259 for (
const auto &[i, sizeVal] : llvm::enumerate(sizeValues)) {
2260 int64_t size = sizeVal.getSExtValue();
2262 if (inputTy.isDynamicDim(i)) {
2266 sizeMatchesInput =
false;
2273 sizeMatchesInput =
false;
2279 if (sizeMatchesInput)
2284 if (!adaptor.getInput1())
2288 if (!inputTy.getElementType().isIntOrIndexOrFloat() ||
2289 !outputTy.getElementType().isIntOrIndexOrFloat())
2292 auto operand = llvm::cast<ElementsAttr>(adaptor.getInput1());
2293 if (operand.isSplat() && outputTy.hasStaticShape()) {
2297 if (inputTy.hasStaticShape() && outputTy.hasStaticShape() &&
2298 outputTy.getNumElements() == 1) {
2299 llvm::SmallVector<uint64_t>
indices =
2300 llvm::to_vector(startElems.
getValues<uint64_t>());
2301 if (
auto values = operand.tryGetValues<Attribute>())
2308OpFoldResult tosa::SelectOp::fold(FoldAdaptor adaptor) {
2309 const Value pred = getPred();
2310 const Value onTrue = getOnTrue();
2311 const Value onFalse = getOnFalse();
2313 const auto predTy = llvm::dyn_cast<RankedTensorType>(pred.
getType());
2314 const auto onTrueTy = llvm::dyn_cast<RankedTensorType>(onTrue.
getType());
2315 const auto onFalseTy = llvm::dyn_cast<RankedTensorType>(onFalse.
getType());
2316 if (!predTy || !onTrueTy || !onFalseTy)
2319 const Type resultTy =
getType();
2321 const ArrayRef<int64_t> predShape = predTy.getShape();
2322 const ArrayRef<int64_t> onTrueShape = onTrueTy.getShape();
2324 if (onTrue == onFalse && onTrueTy == resultTy &&
2329 llvm::dyn_cast_if_present<DenseIntElementsAttr>(adaptor.getInput1());
2332 if (!predicate.isSplat())
2335 const bool predicateValue = predicate.getSplatValue<APInt>().getBoolValue();
2337 SmallVector<SmallVector<int64_t>, 3> shapes;
2338 shapes.emplace_back(predShape);
2339 shapes.emplace_back(onTrueShape);
2340 shapes.emplace_back(onFalseTy.getShape());
2341 const bool isBroadcastable =
2344 if (predicateValue ==
true && onTrueTy == resultTy && isBroadcastable)
2346 if (predicateValue ==
false && onFalseTy == resultTy && isBroadcastable)
2352 const auto inputType =
2353 dyn_cast<RankedTensorType>(tileOp.getInput1().getType());
2354 const auto outputType = dyn_cast<RankedTensorType>(tileOp.getType());
2355 if (!inputType || !outputType)
2359 if (failed(tileOp.getConstantMultiples(multiples)))
2362 for (
const auto [
index, multiple] : llvm::enumerate(multiples)) {
2365 if (outputType.isDynamicDim(
index))
2367 if (inputType.getDimSize(
index) != 1)
2380 Value tileOutput = tileOp.getOutput();
2383 "tile output must have one use");
2386 const bool isBinaryElementwise =
2389 if (!isBinaryElementwise && !isa<tosa::MulOp>(user))
2391 tileOp,
"consumer must be binary broadcastable");
2395 tileOp,
"tile must only expand statically-known singleton dims");
2399 Value otherOperand = lhsOperand == tileOutput ? rhsOperand : lhsOperand;
2400 Value tileInput = tileOp.getInput1();
2402 const ShapedType newOtherType = cast<ShapedType>(otherOperand.
getType());
2403 const ShapedType newTileType = cast<ShapedType>(tileInput.
getType());
2406 newOtherType.getShape(), newTileType.getShape(), broadcastedShape);
2408 const ShapedType outputType = cast<ShapedType>(user->
getResultTypes()[0]);
2409 if (!llvm::equal(broadcastedShape, outputType.getShape()))
2411 tileOp,
"tile output must be broadcastable to consumer operands");
2415 mapper.
map(tileOutput, tileOp.getInput1());
2422void TileOp::getCanonicalizationPatterns(RewritePatternSet &results,
2423 MLIRContext *context) {
2424 results.
add<RemoveBroadcastTileFromBinaryElementwise>(context);
2427OpFoldResult TileOp::fold(FoldAdaptor adaptor) {
2429 if (
auto multiples = llvm::dyn_cast_if_present<DenseElementsAttr>(
2430 adaptor.getMultiples())) {
2431 if (multiples.isSplat() &&
2432 multiples.getSplatValue<APInt>().getSExtValue() == 1)
2434 if (
auto int_array_attr =
2435 llvm::dyn_cast<DenseIntElementsAttr>(multiples)) {
2436 if (llvm::all_of(int_array_attr.getValues<APInt>(),
2437 [](APInt v) { return v.getSExtValue() == 1; }))
2445OpFoldResult TransposeOp::fold(FoldAdaptor adaptor) {
2446 auto resultTy = llvm::cast<ShapedType>(
getType());
2450 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2451 if (input.isSplat() && resultTy.hasRank() && resultTy.hasStaticShape() &&
2452 input.
getType().getElementType() == resultTy.getElementType())
2453 return input.reshape(resultTy);
2457 const llvm::ArrayRef<int32_t> perms = getPerms();
2459 if (!llvm::equal(llvm::seq<int32_t>(0, perms.size()), perms))
2462 return foldToInputIfTypeMatches(
getType(), getInput1());
2465OpFoldResult tosa::NegateOp::fold(FoldAdaptor adaptor) {
2468 auto definingOp = getInput1().getDefiningOp<tosa::NegateOp>();
2474 if (FailureOr<int64_t> maybeIZp = getInput1ZeroPoint();
2475 failed(maybeIZp) || *maybeIZp != 0) {
2479 if (FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
2480 failed(maybeOZp) || *maybeOZp != 0) {
2484 if (FailureOr<int64_t> maybeIZp = definingOp.getInput1ZeroPoint();
2485 failed(maybeIZp) || *maybeIZp != 0) {
2489 if (FailureOr<int64_t> maybeOZp = definingOp.getOutputZeroPoint();
2490 failed(maybeOZp) || *maybeOZp != 0) {
2495 return foldToInputIfTypeMatches(
getType(), definingOp.getInput1());
2498OpFoldResult tosa::AbsOp::fold(FoldAdaptor adaptor) {
2499 auto input = getInput1();
2502 return foldToInputIfTypeMatches(
getType(), input);
2507OpFoldResult tosa::ReciprocalOp::fold(FoldAdaptor adaptor) {
2508 auto input = adaptor.getInput1();
2510 auto inputAttr = llvm::dyn_cast_if_present<DenseElementsAttr>(input);
2512 if (!inputAttr || !inputAttr.isSplat())
2515 auto shapeType = llvm::cast<ShapedType>(
getType());
2516 if (!shapeType.hasRank() || !shapeType.hasStaticShape())
2518 if (
auto floatType = llvm::dyn_cast<FloatType>(inputAttr.getElementType())) {
2519 auto floatVal = inputAttr.getSplatValue<APFloat>();
2521 ReciprocalOp::calcOneElement(floatVal));
2527template <
typename Op,
typename OpFoldAdaptor>
2529 auto input1ConstShape =
2530 dyn_cast<tosa::ConstShapeOp>(op->getInput().getDefiningOp());
2531 if (!input1ConstShape)
2534 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2540template <
typename Op,
typename OpFoldAdaptor>
2542 auto input1ConstShape =
2543 dyn_cast<tosa::ConstShapeOp>(op->getInput1().getDefiningOp());
2544 auto input2ConstShape =
2545 dyn_cast<tosa::ConstShapeOp>(op->getInput2().getDefiningOp());
2546 if (!input1ConstShape || !input2ConstShape)
2549 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2550 const auto input2Attr = cast<DenseElementsAttr>(input2ConstShape.getValues());
2553 input1Attr.getType(),
2557OpFoldResult tosa::DimOp::fold(FoldAdaptor adaptor) {
2558 const auto inputTy = llvm::dyn_cast<ShapedType>(getInput1().
getType());
2559 if (!inputTy || !inputTy.hasRank())
2561 const int32_t axis = getAxis();
2562 const int64_t dimSize = inputTy.getDimSize(axis);
2563 if (ShapedType::isDynamic(dimSize))
2567 const auto resultAttrTy =
2568 RankedTensorType::get(1, builder.getIndexType());
2573 auto const inputs = op->getInput();
2579 concatDims.reserve( 64);
2580 for (
auto const &v : inputs) {
2581 auto vConstShape = dyn_cast<tosa::ConstShapeOp>(v.getDefiningOp());
2585 const auto vAttr = cast<DenseElementsAttr>(vConstShape.getValues());
2588 auto const vAttrVals = vAttr.getValues<APInt>();
2589 for (
auto const &v : vAttrVals) {
2590 concatDims.push_back(v);
2594 auto *ctx = op->getContext();
2595 assert(ctx !=
nullptr &&
"ctx is nullptr");
2596 auto const rankedTy = RankedTensorType::get(
2597 {
static_cast<int64_t>(concatDims.size())}, IndexType::get(ctx));
2603 auto const input1 = op->getInput();
2604 auto const input2 = op->getStart();
2605 auto const input3 = op->getSize();
2607 auto input1ConstShape = dyn_cast<tosa::ConstShapeOp>(input1.getDefiningOp());
2609 if (!input1ConstShape)
2612 auto const input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2616 auto const input1Vals = input1Attr.getValues<APInt>();
2617 auto const totalInput1 = input1Vals.size();
2622 if (failed(start) || failed(size))
2625 auto const startV =
static_cast<int32_t
>(start.value());
2626 auto const sizeV =
static_cast<int32_t
>(size.value());
2628 if ((sizeV <= 0) || (startV < 0) ||
2629 (
static_cast<size_t>(startV + sizeV) > totalInput1))
2633 sliceOfInput.reserve(totalInput1);
2635 for (
auto i = startV; i < (startV + sizeV); i++) {
2636 sliceOfInput.push_back(input1Vals[i]);
2639 auto *ctx = op->getContext();
2640 assert(ctx !=
nullptr &&
"ctx is nullptr");
2642 auto const rankedTy = RankedTensorType::get(
2643 {
static_cast<int64_t>(sliceOfInput.size())}, IndexType::get(ctx));
2648OpFoldResult tosa::AddShapeOp::fold(FoldAdaptor adaptor) {
2652OpFoldResult tosa::SubShapeOp::fold(FoldAdaptor adaptor) {
2656OpFoldResult tosa::MulShapeOp::fold(FoldAdaptor adaptor) {
2660OpFoldResult tosa::DivCeilShapeOp::fold(FoldAdaptor adaptor) {
2661 return binaryFold<DivCeilShapeOp, ShapeDivFoldAdaptor<
true>>(
this);
2664OpFoldResult tosa::DivFloorShapeOp::fold(FoldAdaptor adaptor) {
2665 return binaryFold<DivFloorShapeOp, ShapeDivFoldAdaptor<
false>>(
this);
2668OpFoldResult tosa::ModShapeOp::fold(FoldAdaptor adaptor) {
2672OpFoldResult tosa::MaxShapeOp::fold(FoldAdaptor adaptor) {
2676OpFoldResult tosa::MinShapeOp::fold(FoldAdaptor adaptor) {
2680OpFoldResult tosa::Exp2ShapeOp::fold(FoldAdaptor adaptor) {
2684OpFoldResult tosa::Log2CeilShapeOp::fold(FoldAdaptor adaptor) {
2688OpFoldResult tosa::Log2FloorShapeOp::fold(FoldAdaptor adaptor) {
2692OpFoldResult tosa::ConcatShapeOp::fold(FoldAdaptor adaptor) {
2696OpFoldResult 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