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())
1377 const auto lETy = llvm::cast<ShapedType>(lhs.getType()).
getElementType();
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();
1469 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
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) {
1485 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &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) {
1501 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &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();
1536 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
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);
1560 if (Ceil && !r.isZero() && !
signsDiffer(lhs, rhs)) {
1568 static FailureOr<APFloat>
fold(
const APFloat &lhs,
const APFloat &rhs) {
1574 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
1576 if (lhs.getBitWidth() != rhs.getBitWidth())
1578 if (lhs.isNegative() || (!rhs.isStrictlyPositive()))
1582 return lhs.urem(rhs);
1585 return lhs.srem(rhs);
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) {
1599 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
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) {
1607 return lhs >= rhs ? lhs : rhs;
1612 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &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) {
1620 return lhs <= rhs ? lhs : 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);
1644 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
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());
1675 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &rhs,
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);
1686 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &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);
1697 static FailureOr<APInt>
fold(
const APInt &lhs,
const APInt &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 ArgMinOp::fold(FoldAdaptor adaptor) {
1776 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput().
getType());
1777 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
1778 if (!inputTy || !outputTy || !inputTy.hasStaticShape() ||
1779 !outputTy.hasStaticShape())
1783 if (inputTy.getDimSize(getAxis()) == 1 && outputElementTy.
isInteger()) {
1784 const auto outputElemIntTy = cast<IntegerType>(outputElementTy);
1785 const APInt zero = APInt::getZero(outputElemIntTy.getWidth());
1792OpFoldResult IntDivOp::fold(FoldAdaptor adaptor) {
1793 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1794 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1795 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1796 if (!lhsTy || !rhsTy || !resultTy)
1798 if (lhsTy.getElementType() != rhsTy.getElementType())
1803 auto resultETy = resultTy.getElementType();
1805 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1807 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1808 if (lhsAttr && lhsAttr.isSplat() && rhsAttr && rhsAttr.isSplat()) {
1809 if (llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape() &&
1810 lhsAttr.getSplatValue<APInt>().isZero() &&
1811 !rhsAttr.getSplatValue<APInt>().isZero()) {
1812 return lhsAttr.resizeSplat(resultTy);
1816 if (rhsAttr && rhsAttr.isSplat()) {
1818 lhsTy.getShape(), rhsTy.getShape());
1819 if (isBroadcastable && lhsTy == resultTy &&
1820 llvm::isa<IntegerType>(resultETy) &&
1821 rhsAttr.getSplatValue<APInt>().isOne())
1825 if (rhsAttr && lhsAttr && rhsAttr.isSplat() && lhsAttr.isSplat() &&
1826 llvm::isa<IntegerType>(resultETy) && resultTy.hasStaticShape()) {
1827 APInt l = lhsAttr.getSplatValue<APInt>();
1828 APInt r = rhsAttr.getSplatValue<APInt>();
1830 auto intTy = dyn_cast<mlir::IntegerType>(resultETy);
1832 DivFoldAdaptor<
false>::fold(l, r, intTy.isUnsigned());
1845std::optional<APInt> mulInt(APInt
lhs, APInt
rhs, int32_t shift,
1846 unsigned bitwidth) {
1847 bool overflow =
false;
1848 APInt
result =
lhs.sext(64).smul_ov(
rhs.sext(64), overflow);
1851 return std::nullopt;
1854 auto round = APInt(64, 1) << (shift - 1);
1856 result.ashrInPlace(shift);
1859 if (!(
result.getSExtValue() >= INT32_MIN &&
1860 result.getSExtValue() <= INT32_MAX)) {
1862 return std::nullopt;
1866 return result.trunc(bitwidth);
1869DenseElementsAttr mulBinaryFolder(DenseElementsAttr
lhs, DenseElementsAttr
rhs,
1870 RankedTensorType ty, int32_t shift) {
1872 if (!ty.hasStaticShape())
1875 if (llvm::isa<IntegerType>(ty.getElementType())) {
1876 APInt l =
lhs.getSplatValue<APInt>();
1877 APInt r =
rhs.getSplatValue<APInt>();
1883 auto bitwidth = ty.getElementType().getIntOrFloatBitWidth();
1884 const std::optional<APInt>
result = mulInt(l, r, shift, bitwidth);
1890 if (llvm::isa<FloatType>(ty.getElementType())) {
1891 APFloat l =
lhs.getSplatValue<APFloat>();
1892 APFloat r =
rhs.getSplatValue<APFloat>();
1902OpFoldResult MulOp::fold(FoldAdaptor adaptor) {
1903 auto lhs = getInput1();
1904 auto rhs = getInput2();
1905 auto lhsTy = llvm::dyn_cast<RankedTensorType>(
lhs.getType());
1906 auto rhsTy = llvm::dyn_cast<RankedTensorType>(
rhs.getType());
1907 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1908 if (!lhsTy || !rhsTy || !resultTy)
1911 auto resultETy = resultTy.getElementType();
1913 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1915 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1920 if (resultETy.isInteger(32)) {
1921 ElementsAttr shift_elem;
1922 if (getShift().getImpl()) {
1926 shift = shift_elem.getValues<IntegerAttr>()[0].getInt();
1930 if (rhsTy == resultTy &&
isSplatZero(resultETy, lhsAttr) &&
1931 resultTy.hasStaticShape())
1933 return lhsAttr.resizeSplat(resultTy);
1934 if (lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr) &&
1935 resultTy.hasStaticShape())
1936 return rhsAttr.resizeSplat(resultTy);
1939 lhsTy.getShape(), rhsTy.getShape());
1940 if (isBroadcastable && rhsTy == resultTy &&
1943 if (isBroadcastable && lhsTy == resultTy &&
1947 return mulBinaryFolder(lhsAttr, rhsAttr, resultTy, shift);
1950OpFoldResult SubOp::fold(FoldAdaptor adaptor) {
1951 auto lhsTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
1952 auto rhsTy = llvm::dyn_cast<RankedTensorType>(getInput2().
getType());
1953 auto resultTy = llvm::dyn_cast<RankedTensorType>(
getType());
1954 if (!lhsTy || !rhsTy || !resultTy)
1958 if (!lhsTy.getElementType().isIntOrIndexOrFloat() ||
1959 !rhsTy.getElementType().isIntOrIndexOrFloat())
1962 auto resultETy = resultTy.getElementType();
1964 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1966 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1969 lhsTy.getShape(), rhsTy.getShape());
1970 if (isBroadcastable && lhsTy == resultTy &&
isSplatZero(resultETy, rhsAttr))
1973 if (!lhsAttr || !rhsAttr)
1979OpFoldResult GreaterOp::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());
1986 if (!lhsAttr || !rhsAttr)
1992OpFoldResult GreaterEqualOp::fold(FoldAdaptor adaptor) {
1993 auto resultTy = llvm::cast<ShapedType>(
getType());
1995 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
1997 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
1999 if (!lhsAttr || !rhsAttr)
2005OpFoldResult EqualOp::fold(FoldAdaptor adaptor) {
2006 auto resultTy = llvm::cast<ShapedType>(
getType());
2008 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1());
2010 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput2());
2011 Value
lhs = getInput1();
2012 Value
rhs = getInput2();
2013 auto lhsTy = llvm::cast<ShapedType>(
lhs.getType());
2017 if (llvm::isa<IntegerType>(lhsTy.getElementType()) && resultTy.hasRank() &&
2018 resultTy.hasStaticShape() &&
lhs ==
rhs) {
2022 if (!lhsAttr || !rhsAttr)
2028OpFoldResult CastOp::fold(FoldAdaptor adaptor) {
2032 auto operand = llvm::dyn_cast_if_present<ElementsAttr>(adaptor.getInput());
2036 auto inTy = llvm::cast<ShapedType>(getInput().
getType());
2037 auto outTy = llvm::cast<ShapedType>(
getType());
2038 if (!outTy.hasRank() || !outTy.hasStaticShape())
2040 auto inETy = inTy.getElementType();
2041 auto outETy = outTy.getElementType();
2043 if (operand.isSplat()) {
2044 if (llvm::isa<FloatType>(inETy) && llvm::isa<FloatType>(outETy)) {
2046 auto splatVal = operand.getSplatValue<APFloat>();
2047 auto &semantics = llvm::cast<FloatType>(outETy).getFloatSemantics();
2048 splatVal.convert(semantics, llvm::RoundingMode::NearestTiesToEven,
2053 if (llvm::isa<IntegerType>(inETy) && llvm::isa<FloatType>(outETy)) {
2054 const bool unsign = llvm::cast<IntegerType>(inETy).isUnsignedInteger() ||
2055 adaptor.getInputUnsigned();
2057 splatVal.convertFromAPInt(operand.getSplatValue<APInt>(), !unsign,
2058 llvm::RoundingMode::NearestTiesToEven);
2062 if (llvm::isa<FloatType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2063 auto unsign = llvm::cast<IntegerType>(outETy).isUnsignedInteger();
2064 auto intVal = APSInt(
2065 llvm::cast<IntegerType>(outETy).getIntOrFloatBitWidth(), unsign);
2066 auto floatVal = operand.getSplatValue<APFloat>();
2068 floatVal.convertToInteger(intVal, llvm::RoundingMode::NearestTiesToEven,
2073 if (llvm::isa<IntegerType>(inETy) && llvm::isa<IntegerType>(outETy)) {
2074 const auto inIntType = llvm::cast<IntegerType>(inETy);
2075 const bool unsignIn =
2076 inIntType.isUnsignedInteger() || adaptor.getInputUnsigned();
2078 inETy.getIntOrFloatBitWidth() > outETy.getIntOrFloatBitWidth();
2079 auto intVal = operand.getSplatValue<APInt>();
2080 auto bitwidth = outETy.getIntOrFloatBitWidth();
2083 if (outETy.isInteger(1)) {
2084 intVal = APInt(bitwidth, intVal.isZero() ? 0 : 1);
2086 intVal = intVal.trunc(bitwidth);
2087 }
else if (unsignIn || inIntType.isInteger(1)) {
2088 intVal = intVal.zext(bitwidth);
2090 intVal = intVal.sext(bitwidth);
2100OpFoldResult ConstOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2102OpFoldResult ConstShapeOp::fold(FoldAdaptor adaptor) {
return getValuesAttr(); }
2104#define REDUCE_FOLDER(OP) \
2105 OpFoldResult OP::fold(FoldAdaptor adaptor) { \
2106 ShapedType inputTy = llvm::cast<ShapedType>(getInput().getType()); \
2107 if (!inputTy.hasRank()) \
2109 if (inputTy != getType()) \
2111 if (inputTy.getRank() == 0 || inputTy.getDimSize(getAxis()) == 1) \
2112 return getInput(); \
2125 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2126 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2128 if (!inputTy || !outputTy)
2134 if (inputTy == outputTy && inputTy.getNumDynamicDims() < 2)
2138 if (
auto reshapeOp = llvm::dyn_cast_if_present<tosa::ReshapeOp>(
2139 getInput1().getDefiningOp())) {
2140 getInput1Mutable().assign(reshapeOp.getInput1());
2145 if (!inputTy.getElementType().isIntOrIndexOrFloat())
2149 if (!outputTy.hasStaticShape())
2153 if (
auto operand = llvm::dyn_cast_if_present<DenseResourceElementsAttr>(
2154 adaptor.getInput1()))
2155 return DenseResourceElementsAttr::get(outputTy, operand.getRawHandle());
2159 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2161 if (operand.isSplat())
2166 if (!getInput1().hasOneUse())
2173 return operand.reshape(
2174 llvm::cast<ShapedType>(operand.getType()).clone(shapeVec));
2180OpFoldResult PadOp::fold(FoldAdaptor adaptor) {
2182 if (adaptor.getPadding() && getInput1().
getType() ==
getType()) {
2183 auto densePad = llvm::dyn_cast<DenseElementsAttr>(adaptor.getPadding());
2184 if (densePad && densePad.isSplat() &&
2185 densePad.getSplatValue<APInt>().isZero()) {
2195OpFoldResult ResizeOp::fold(FoldAdaptor adaptor) {
2197 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getScale());
2199 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getOffset());
2201 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getBorder());
2202 if (!scaleAttr || !offsetAttr || !borderAttr) {
2209 if (scale.size() != 4 || offset.size() != 2 || border.size() != 2) {
2214 if (scale[0] != scale[1] || scale[2] != scale[3]) {
2219 if (offset[0] != 0 || offset[1] != 0) {
2224 if (border[0] != 0 || border[1] != 0) {
2228 return foldToInputIfTypeMatches(
getType(), getInput());
2231OpFoldResult ReverseOp::fold(FoldAdaptor adaptor) {
2232 auto operand = getInput1();
2233 auto operandTy = llvm::cast<ShapedType>(operand.getType());
2234 auto axis = getAxis();
2237 const bool isSplatInput =
2238 !isa<BlockScaledType>(operandTy.getElementType()) &&
2239 llvm::isa_and_nonnull<SplatElementsAttr>(adaptor.getInput1());
2240 if (!operandTy.hasRank() ||
2241 (!isSplatInput && operandTy.getDimSize(axis) != 1))
2243 return foldToInputIfTypeMatches(
getType(), operand);
2246OpFoldResult SliceOp::fold(FoldAdaptor adaptor) {
2247 auto inputTy = llvm::dyn_cast<RankedTensorType>(getInput1().
getType());
2248 auto outputTy = llvm::dyn_cast<RankedTensorType>(
getType());
2250 if (!inputTy || !outputTy)
2253 if (inputTy == outputTy && inputTy.hasStaticShape())
2258 DenseElementsAttr startElems;
2264 llvm::all_of(startElems.
getValues<APInt>(),
2265 [](
const APInt &val) { return val.isZero(); });
2270 DenseElementsAttr sizeElems;
2274 auto inputShape = inputTy.getShape();
2275 auto sizeValues = sizeElems.
getValues<APInt>();
2277 bool sizeMatchesInput =
true;
2278 for (
const auto &[i, sizeVal] : llvm::enumerate(sizeValues)) {
2279 int64_t size = sizeVal.getSExtValue();
2281 if (inputTy.isDynamicDim(i)) {
2285 sizeMatchesInput =
false;
2292 sizeMatchesInput =
false;
2298 if (sizeMatchesInput)
2303 if (!adaptor.getInput1())
2307 if (!inputTy.getElementType().isIntOrIndexOrFloat() ||
2308 !outputTy.getElementType().isIntOrIndexOrFloat())
2311 auto operand = llvm::cast<ElementsAttr>(adaptor.getInput1());
2312 if (operand.isSplat() && outputTy.hasStaticShape()) {
2316 if (inputTy.hasStaticShape() && outputTy.hasStaticShape() &&
2317 outputTy.getNumElements() == 1) {
2318 llvm::SmallVector<uint64_t>
indices =
2319 llvm::to_vector(startElems.
getValues<uint64_t>());
2320 if (
auto values = operand.tryGetValues<Attribute>())
2327OpFoldResult tosa::SelectOp::fold(FoldAdaptor adaptor) {
2328 const Value pred = getPred();
2329 const Value onTrue = getOnTrue();
2330 const Value onFalse = getOnFalse();
2332 const auto predTy = llvm::dyn_cast<RankedTensorType>(pred.
getType());
2333 const auto onTrueTy = llvm::dyn_cast<RankedTensorType>(onTrue.
getType());
2334 const auto onFalseTy = llvm::dyn_cast<RankedTensorType>(onFalse.
getType());
2335 if (!predTy || !onTrueTy || !onFalseTy)
2338 const Type resultTy =
getType();
2340 const ArrayRef<int64_t> predShape = predTy.getShape();
2341 const ArrayRef<int64_t> onTrueShape = onTrueTy.getShape();
2343 if (onTrue == onFalse && onTrueTy == resultTy &&
2348 llvm::dyn_cast_if_present<DenseIntElementsAttr>(adaptor.getInput1());
2351 if (!predicate.isSplat())
2354 const bool predicateValue = predicate.getSplatValue<APInt>().getBoolValue();
2356 SmallVector<SmallVector<int64_t>, 3> shapes;
2357 shapes.emplace_back(predShape);
2358 shapes.emplace_back(onTrueShape);
2359 shapes.emplace_back(onFalseTy.getShape());
2360 const bool isBroadcastable =
2363 if (predicateValue ==
true && onTrueTy == resultTy && isBroadcastable)
2365 if (predicateValue ==
false && onFalseTy == resultTy && isBroadcastable)
2371 const auto inputType =
2372 dyn_cast<RankedTensorType>(tileOp.getInput1().getType());
2373 const auto outputType = dyn_cast<RankedTensorType>(tileOp.getType());
2374 if (!inputType || !outputType)
2378 if (failed(tileOp.getConstantMultiples(multiples)))
2381 for (
const auto [
index, multiple] : llvm::enumerate(multiples)) {
2384 if (outputType.isDynamicDim(
index))
2386 if (inputType.getDimSize(
index) != 1)
2399 Value tileOutput = tileOp.getOutput();
2402 "tile output must have one use");
2405 const bool isBinaryElementwise =
2408 if (!isBinaryElementwise && !isa<tosa::MulOp>(user))
2410 tileOp,
"consumer must be binary broadcastable");
2414 tileOp,
"tile must only expand statically-known singleton dims");
2418 Value otherOperand = lhsOperand == tileOutput ? rhsOperand : lhsOperand;
2419 Value tileInput = tileOp.getInput1();
2421 const ShapedType newOtherType = cast<ShapedType>(otherOperand.
getType());
2422 const ShapedType newTileType = cast<ShapedType>(tileInput.
getType());
2425 newOtherType.getShape(), newTileType.getShape(), broadcastedShape);
2427 const ShapedType outputType = cast<ShapedType>(user->
getResultTypes()[0]);
2428 if (!llvm::equal(broadcastedShape, outputType.getShape()))
2430 tileOp,
"tile output must be broadcastable to consumer operands");
2434 mapper.
map(tileOutput, tileOp.getInput1());
2441void TileOp::getCanonicalizationPatterns(RewritePatternSet &results,
2442 MLIRContext *context) {
2443 results.
add<RemoveBroadcastTileFromBinaryElementwise>(context);
2446OpFoldResult TileOp::fold(FoldAdaptor adaptor) {
2448 if (
auto multiples = llvm::dyn_cast_if_present<DenseElementsAttr>(
2449 adaptor.getMultiples())) {
2450 if (multiples.isSplat() &&
2451 multiples.getSplatValue<APInt>().getSExtValue() == 1)
2453 if (
auto int_array_attr =
2454 llvm::dyn_cast<DenseIntElementsAttr>(multiples)) {
2455 if (llvm::all_of(int_array_attr.getValues<APInt>(),
2456 [](APInt v) { return v.getSExtValue() == 1; }))
2464OpFoldResult TransposeOp::fold(FoldAdaptor adaptor) {
2465 auto resultTy = llvm::cast<ShapedType>(
getType());
2469 llvm::dyn_cast_if_present<DenseElementsAttr>(adaptor.getInput1())) {
2470 if (input.isSplat() && resultTy.hasRank() && resultTy.hasStaticShape() &&
2471 input.
getType().getElementType() == resultTy.getElementType())
2472 return input.reshape(resultTy);
2476 const llvm::ArrayRef<int32_t> perms = getPerms();
2478 if (!llvm::equal(llvm::seq<int32_t>(0, perms.size()), perms))
2481 return foldToInputIfTypeMatches(
getType(), getInput1());
2484OpFoldResult tosa::NegateOp::fold(FoldAdaptor adaptor) {
2487 auto definingOp = getInput1().getDefiningOp<tosa::NegateOp>();
2493 if (FailureOr<int64_t> maybeIZp = getInput1ZeroPoint();
2494 failed(maybeIZp) || *maybeIZp != 0) {
2498 if (FailureOr<int64_t> maybeOZp = getOutputZeroPoint();
2499 failed(maybeOZp) || *maybeOZp != 0) {
2503 if (FailureOr<int64_t> maybeIZp = definingOp.getInput1ZeroPoint();
2504 failed(maybeIZp) || *maybeIZp != 0) {
2508 if (FailureOr<int64_t> maybeOZp = definingOp.getOutputZeroPoint();
2509 failed(maybeOZp) || *maybeOZp != 0) {
2514 return foldToInputIfTypeMatches(
getType(), definingOp.getInput1());
2517OpFoldResult tosa::AbsOp::fold(FoldAdaptor adaptor) {
2518 auto input = getInput1();
2521 return foldToInputIfTypeMatches(
getType(), input);
2526OpFoldResult tosa::ReciprocalOp::fold(FoldAdaptor adaptor) {
2527 auto input = adaptor.getInput1();
2529 auto inputAttr = llvm::dyn_cast_if_present<DenseElementsAttr>(input);
2531 if (!inputAttr || !inputAttr.isSplat())
2534 auto shapeType = llvm::cast<ShapedType>(
getType());
2535 if (!shapeType.hasRank() || !shapeType.hasStaticShape())
2537 if (
auto floatType = llvm::dyn_cast<FloatType>(inputAttr.getElementType())) {
2538 auto floatVal = inputAttr.getSplatValue<APFloat>();
2540 ReciprocalOp::calcOneElement(floatVal));
2546template <
typename Op,
typename OpFoldAdaptor>
2548 auto input1ConstShape =
2549 dyn_cast<tosa::ConstShapeOp>(op->getInput().getDefiningOp());
2550 if (!input1ConstShape)
2553 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2559template <
typename Op,
typename OpFoldAdaptor>
2561 auto input1ConstShape =
2562 dyn_cast<tosa::ConstShapeOp>(op->getInput1().getDefiningOp());
2563 auto input2ConstShape =
2564 dyn_cast<tosa::ConstShapeOp>(op->getInput2().getDefiningOp());
2565 if (!input1ConstShape || !input2ConstShape)
2568 const auto input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2569 const auto input2Attr = cast<DenseElementsAttr>(input2ConstShape.getValues());
2572 input1Attr.getType(),
2576OpFoldResult tosa::DimOp::fold(FoldAdaptor adaptor) {
2577 const auto inputTy = llvm::dyn_cast<ShapedType>(getInput1().
getType());
2578 if (!inputTy || !inputTy.hasRank())
2580 const int32_t axis = getAxis();
2581 const int64_t dimSize = inputTy.getDimSize(axis);
2582 if (ShapedType::isDynamic(dimSize))
2586 const auto resultAttrTy =
2587 RankedTensorType::get(1, builder.getIndexType());
2592 auto const inputs = op->getInput();
2598 concatDims.reserve( 64);
2599 for (
auto const &v : inputs) {
2600 auto vConstShape = dyn_cast<tosa::ConstShapeOp>(v.getDefiningOp());
2604 const auto vAttr = cast<DenseElementsAttr>(vConstShape.getValues());
2607 auto const vAttrVals = vAttr.getValues<APInt>();
2608 for (
auto const &v : vAttrVals) {
2609 concatDims.push_back(v);
2613 auto *ctx = op->getContext();
2614 assert(ctx !=
nullptr &&
"ctx is nullptr");
2615 auto const rankedTy = RankedTensorType::get(
2616 {
static_cast<int64_t>(concatDims.size())}, IndexType::get(ctx));
2622 auto const input1 = op->getInput();
2623 auto const input2 = op->getStart();
2624 auto const input3 = op->getSize();
2626 auto input1ConstShape = dyn_cast<tosa::ConstShapeOp>(input1.getDefiningOp());
2628 if (!input1ConstShape)
2631 auto const input1Attr = cast<DenseElementsAttr>(input1ConstShape.getValues());
2635 auto const input1Vals = input1Attr.getValues<APInt>();
2636 auto const totalInput1 = input1Vals.size();
2641 if (failed(start) || failed(size))
2644 auto const startV =
static_cast<int32_t
>(start.value());
2645 auto const sizeV =
static_cast<int32_t
>(size.value());
2647 if ((sizeV <= 0) || (startV < 0) ||
2648 (
static_cast<size_t>(startV + sizeV) > totalInput1))
2652 sliceOfInput.reserve(totalInput1);
2654 for (
auto i = startV; i < (startV + sizeV); i++) {
2655 sliceOfInput.push_back(input1Vals[i]);
2658 auto *ctx = op->getContext();
2659 assert(ctx !=
nullptr &&
"ctx is nullptr");
2661 auto const rankedTy = RankedTensorType::get(
2662 {
static_cast<int64_t>(sliceOfInput.size())}, IndexType::get(ctx));
2667OpFoldResult tosa::AddShapeOp::fold(FoldAdaptor adaptor) {
2671OpFoldResult tosa::SubShapeOp::fold(FoldAdaptor adaptor) {
2675OpFoldResult tosa::MulShapeOp::fold(FoldAdaptor adaptor) {
2679OpFoldResult tosa::DivCeilShapeOp::fold(FoldAdaptor adaptor) {
2680 return binaryFold<DivCeilShapeOp, ShapeDivFoldAdaptor<
true>>(
this);
2683OpFoldResult tosa::DivFloorShapeOp::fold(FoldAdaptor adaptor) {
2684 return binaryFold<DivFloorShapeOp, ShapeDivFoldAdaptor<
false>>(
this);
2687OpFoldResult tosa::ModShapeOp::fold(FoldAdaptor adaptor) {
2691OpFoldResult tosa::MaxShapeOp::fold(FoldAdaptor adaptor) {
2695OpFoldResult tosa::MinShapeOp::fold(FoldAdaptor adaptor) {
2699OpFoldResult tosa::Exp2ShapeOp::fold(FoldAdaptor adaptor) {
2703OpFoldResult tosa::Log2CeilShapeOp::fold(FoldAdaptor adaptor) {
2707OpFoldResult tosa::Log2FloorShapeOp::fold(FoldAdaptor adaptor) {
2711OpFoldResult tosa::ConcatShapeOp::fold(FoldAdaptor adaptor) {
2715OpFoldResult 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.
static const llvm::fltSemantics * getFloatSemantics(TruncfSrcElemTypes etype)
Float semantics the element type attributes of xevm.truncf and xevm.extf stand for.
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