29#include "llvm/ADT/STLExtras.h"
30#include "llvm/ADT/Sequence.h"
31#include "llvm/ADT/SmallVectorExtras.h"
38template <
typename OpTy>
42 typename OpTy::Properties properties{};
43 OpTy::populateDefaultProperties(
46 return OpTy::create(builder, loc, resultTypes, operands, properties,
72template <
typename OpTy>
80 auto nanMode = op.getNanMode();
81 if (nanMode == NanPropagationMode::PROPAGATE)
85 Value lhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
86 arith::CmpFPredicate::UNO, lhs, lhs);
87 Value rhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
88 arith::CmpFPredicate::UNO, rhs, rhs);
90 arith::SelectOp::create(rewriter, op.getLoc(), lhsIsNaN, rhs,
result);
91 return arith::SelectOp::create(rewriter, op.getLoc(), rhsIsNaN, lhs,
97 ConversionPatternRewriter &rewriter) {
103 if (isa<tosa::AbsOp>(op) && isa<FloatType>(elementTy))
107 if (isa<tosa::AbsOp>(op) && isa<IntegerType>(elementTy)) {
108 auto zero = arith::ConstantOp::create(rewriter, loc,
109 rewriter.getZeroAttr(elementTy));
110 auto neg = arith::SubIOp::create(rewriter, loc, zero, args[0]);
111 return arith::MaxSIOp::create(rewriter, loc, args[0], neg);
115 if (isa<tosa::AddOp>(op) && isa<FloatType>(elementTy))
119 if (isa<tosa::AddOp>(op) && isa<IntegerType>(elementTy))
124 if (isa<tosa::SubOp>(op) && isa<FloatType>(elementTy))
128 if (isa<tosa::SubOp>(op) && isa<IntegerType>(elementTy))
133 if (isa<tosa::IntDivOp>(op) && isa<IntegerType>(elementTy))
138 if (isa<tosa::ReciprocalOp>(op) && isa<FloatType>(elementTy)) {
140 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
141 return arith::DivFOp::create(rewriter, loc, one, args[0]);
145 if (isa<tosa::MulOp>(op)) {
146 auto shiftVal = cast<tosa::MulOp>(op).getShift();
148 bool shiftIsConstant =
true;
151 shift = shiftElem.
getValues<IntegerAttr>()[0].getInt();
153 shiftIsConstant =
false;
155 if (isa<FloatType>(elementTy)) {
157 (
void)rewriter.notifyMatchFailure(op,
158 "Cannot have shift value for float");
161 return arith::MulFOp::create(rewriter, loc, args[0], args[1]);
164 if (isa<IntegerType>(elementTy)) {
168 if (shift > 0 || !shiftIsConstant) {
175 a = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), a);
177 if (!
b.getType().isInteger(32))
178 b = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(),
b);
180 auto shiftAmount = shiftIsConstant ? shiftConst : args[2];
181 auto roundingAttr = RoundingModeAttr::get(rewriter.getContext(),
182 RoundingMode::SINGLE_ROUND);
184 tosa::ApplyScaleOp::create(rewriter, loc, rewriter.getI32Type(), a,
185 b, shiftAmount, roundingAttr);
191 int bWidth =
b.getType().getIntOrFloatBitWidth();
192 int cWidth = resultTypes[0].getIntOrFloatBitWidth();
195 a = arith::ExtSIOp::create(rewriter, loc, resultTypes[0], a);
197 b = arith::ExtSIOp::create(rewriter, loc, resultTypes[0],
b);
199 return arith::MulIOp::create(rewriter, loc, resultTypes, a,
b);
204 if (isa<tosa::NegateOp>(op)) {
205 auto negate = cast<tosa::NegateOp>(op);
208 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
209 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
210 bool hasInZp = !failed(maybeInZp);
211 bool hasOutZp = !failed(maybeOutZp);
217 if (isa<FloatType>(elementTy))
218 return arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
220 if (isa<IntegerType>(elementTy)) {
222 Type intermediateType;
225 int intermediateBitWidth = 64;
227 if (hasInZp && hasOutZp) {
229 const int64_t zpAdd = inZp + outZp;
231 APInt::getSignedMaxValue(inputBitWidth).getSExtValue() +
236 if (maxValue <= APInt::getSignedMaxValue(16).getSExtValue()) {
237 intermediateBitWidth = 16;
238 }
else if (maxValue <= APInt::getSignedMaxValue(32).getSExtValue()) {
239 intermediateBitWidth = 32;
242 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
243 zpAddValue = arith::ConstantOp::create(
244 rewriter, loc, rewriter.getIntegerAttr(intermediateType, zpAdd));
246 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
247 Value arg1 = args[1];
248 Value arg2 = args[2];
250 if (arg1.
getType() != intermediateType)
251 arg1 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg1);
252 if (arg2.
getType() != intermediateType)
253 arg2 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg2);
255 arith::AddIOp::create(rewriter, loc, intermediateType, arg1, arg2);
261 if (ext.
getType() != intermediateType)
262 ext = arith::ExtSIOp::create(rewriter, loc, intermediateType, ext);
263 auto sub = arith::SubIOp::create(rewriter, loc, zpAddValue, ext);
267 rewriter, loc, intermediateType,
268 APInt::getSignedMinValue(inputBitWidth).getSExtValue());
270 rewriter, loc, intermediateType,
271 APInt::getSignedMaxValue(inputBitWidth).getSExtValue());
275 if (
clamp.getType() == elementTy)
277 return arith::TruncIOp::create(rewriter, loc, elementTy,
clamp);
282 if (isa<tosa::BitwiseAndOp>(op) && isa<IntegerType>(elementTy))
283 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
286 if (isa<tosa::BitwiseOrOp>(op) && isa<IntegerType>(elementTy))
287 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
290 if (isa<tosa::BitwiseNotOp>(op) && isa<IntegerType>(elementTy)) {
291 auto allOnesAttr = rewriter.getIntegerAttr(
292 elementTy, APInt::getAllOnes(elementTy.getIntOrFloatBitWidth()));
293 auto allOnes = arith::ConstantOp::create(rewriter, loc, allOnesAttr);
294 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], allOnes);
298 if (isa<tosa::BitwiseXorOp>(op) && isa<IntegerType>(elementTy))
299 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
302 if (isa<tosa::LogicalLeftShiftOp>(op) && isa<IntegerType>(elementTy))
307 if (isa<tosa::LogicalRightShiftOp>(op) && isa<IntegerType>(elementTy))
312 if (isa<tosa::ArithmeticRightShiftOp>(op) && isa<IntegerType>(elementTy)) {
314 rewriter, loc, resultTypes, args);
315 bool round = cast<tosa::ArithmeticRightShiftOp>(op).getRound();
320 Type i1Ty = IntegerType::get(rewriter.getContext(), 1);
321 auto one = arith::ConstantOp::create(rewriter, loc,
322 IntegerAttr::get(elementTy, 1));
323 auto zero = arith::ConstantOp::create(rewriter, loc,
324 IntegerAttr::get(elementTy, 0));
326 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 0));
328 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 1));
331 auto shiftValueGreaterThanZero = arith::CmpIOp::create(
332 rewriter, loc, arith::CmpIPredicate::sgt, args[1], zero);
336 arith::SubIOp::create(rewriter, loc, resultTypes, args[1], one);
338 arith::ShRSIOp::create(rewriter, loc, resultTypes, args[0], subtract)
341 rewriter, loc,
TypeRange{i1Ty}, shifted);
343 arith::AndIOp::create(rewriter, loc, i1Ty, truncated, i1one);
345 auto shouldRound = arith::SelectOp::create(
346 rewriter, loc, i1Ty, shiftValueGreaterThanZero, isInputOdd, i1zero);
348 arith::ExtUIOp::create(rewriter, loc, resultTypes, shouldRound);
349 return arith::AddIOp::create(rewriter, loc, resultTypes,
result, extended);
353 if (isa<tosa::ClzOp>(op) && isa<IntegerType>(elementTy)) {
354 return math::CountLeadingZerosOp::create(rewriter, loc, elementTy, args[0]);
358 if (isa<tosa::LogicalAndOp>(op) && elementTy.isInteger(1))
359 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
362 if (isa<tosa::LogicalNotOp>(op) && elementTy.isInteger(1)) {
363 auto one = arith::ConstantOp::create(rewriter, loc,
364 rewriter.getIntegerAttr(elementTy, 1));
365 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], one);
369 if (isa<tosa::LogicalOrOp>(op) && elementTy.isInteger(1))
370 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
373 if (isa<tosa::LogicalXorOp>(op) && elementTy.isInteger(1))
374 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
377 if (isa<tosa::PowOp>(op) && isa<FloatType>(elementTy))
382 if (isa<tosa::RsqrtOp>(op) && isa<FloatType>(elementTy))
387 if (isa<tosa::LogOp>(op) && isa<FloatType>(elementTy))
392 if (isa<tosa::ExpOp>(op) && isa<FloatType>(elementTy))
397 if (isa<tosa::SinOp>(op) && isa<FloatType>(elementTy))
402 if (isa<tosa::CosOp>(op) && isa<FloatType>(elementTy))
407 if (isa<tosa::TanhOp>(op) && isa<FloatType>(elementTy))
412 if (isa<tosa::ErfOp>(op) && llvm::isa<FloatType>(elementTy))
417 if (isa<tosa::GreaterOp>(op) && isa<FloatType>(elementTy))
418 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGT,
421 if (isa<tosa::GreaterOp>(op) && elementTy.isSignlessInteger())
422 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sgt,
426 if (isa<tosa::GreaterEqualOp>(op) && isa<FloatType>(elementTy))
427 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGE,
430 if (isa<tosa::GreaterEqualOp>(op) && elementTy.isSignlessInteger())
431 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sge,
435 if (isa<tosa::EqualOp>(op) && isa<FloatType>(elementTy))
436 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OEQ,
439 if (isa<tosa::EqualOp>(op) && elementTy.isSignlessInteger())
440 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
444 if (isa<tosa::SelectOp>(op)) {
446 if (isa<FloatType>(elementTy) || isa<IntegerType>(elementTy))
447 return arith::SelectOp::create(rewriter, loc, args[0], args[1], args[2]);
451 if (isa<tosa::MaximumOp>(op) && isa<FloatType>(elementTy)) {
452 auto max = arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
454 rewriter, args[0], args[1],
max);
457 if (isa<tosa::MaximumOp>(op) && elementTy.isSignlessInteger()) {
458 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
462 if (isa<tosa::MinimumOp>(op) && isa<FloatType>(elementTy)) {
463 auto min = arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
465 rewriter, args[0], args[1],
min);
468 if (isa<tosa::MinimumOp>(op) && elementTy.isSignlessInteger()) {
469 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
473 if (isa<tosa::CeilOp>(op) && isa<FloatType>(elementTy))
478 if (isa<tosa::FloorOp>(op) && isa<FloatType>(elementTy))
483 if (isa<tosa::ClampOp>(op) && isa<FloatType>(elementTy)) {
484 bool losesInfo =
false;
485 auto clampOp = cast<tosa::ClampOp>(op);
486 APFloat minApf = cast<FloatAttr>(clampOp.getMinValAttr()).getValue();
487 APFloat maxApf = cast<FloatAttr>(clampOp.getMaxValAttr()).getValue();
489 APFloat::rmNearestTiesToEven, &losesInfo);
491 APFloat::rmNearestTiesToEven, &losesInfo);
492 auto min = arith::ConstantOp::create(
493 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, minApf));
494 auto max = arith::ConstantOp::create(
495 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, maxApf));
498 const auto nanMode = clampOp.getNanMode();
501 if (!isa<FloatType>(elementTy))
506 if (nanMode == NanPropagationMode::PROPAGATE)
520 Value isNaN = arith::CmpFOp::create(
521 rewriter, op->
getLoc(), arith::CmpFPredicate::UNO, args[0], args[0]);
524 return arith::SelectOp::create(rewriter, op->
getLoc(), isNaN,
min,
result);
527 if (isa<tosa::ClampOp>(op) && isa<IntegerType>(elementTy)) {
528 auto intTy = cast<IntegerType>(elementTy);
529 auto clampOp = cast<tosa::ClampOp>(op);
531 cast<IntegerAttr>(clampOp.getMinValAttr()).getValue().getSExtValue();
533 cast<IntegerAttr>(clampOp.getMaxValAttr()).getValue().getSExtValue();
535 int64_t minRepresentable = std::numeric_limits<int64_t>::min();
536 int64_t maxRepresentable = std::numeric_limits<int64_t>::max();
537 if (intTy.isUnsignedInteger()) {
538 minRepresentable = 0;
539 if (intTy.getIntOrFloatBitWidth() <= 63) {
541 (
int64_t)APInt::getMaxValue(intTy.getIntOrFloatBitWidth())
544 }
else if (intTy.getIntOrFloatBitWidth() <= 64) {
546 minRepresentable = APInt::getSignedMinValue(intTy.getIntOrFloatBitWidth())
548 maxRepresentable = APInt::getSignedMaxValue(intTy.getIntOrFloatBitWidth())
553 min = std::max(
min, minRepresentable);
554 max = std::max(
max, minRepresentable);
555 min = std::min(
min, maxRepresentable);
556 max = std::min(
max, maxRepresentable);
559 intTy.getIntOrFloatBitWidth());
561 intTy.getIntOrFloatBitWidth());
563 intTy.isUnsignedInteger());
567 if (isa<tosa::SigmoidOp>(op) && isa<FloatType>(elementTy)) {
569 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
570 auto negate = arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
571 auto exp = mlir::math::ExpOp::create(rewriter, loc, resultTypes, negate);
572 auto added = arith::AddFOp::create(rewriter, loc, exp, one);
573 return arith::DivFOp::create(rewriter, loc, one, added);
577 if (isa<tosa::CastOp>(op)) {
578 Type srcTy = elementTy;
579 Type dstTy = resultTypes.front();
581 (
void)rewriter.notifyMatchFailure(op,
"unsupported type");
591 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && bitExtend)
595 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && !bitExtend)
600 if (srcTy.
isInteger(1) && arith::UIToFPOp::areCastCompatible(srcTy, dstTy))
604 if (srcTy.
isInteger(1) && isa<IntegerType>(dstTy) && bitExtend)
611 auto unrealizedCast =
612 UnrealizedConversionCastOp::create(
616 return arith::UIToFPOp::create(rewriter, loc, resultTypes[0],
621 if (arith::SIToFPOp::areCastCompatible(srcTy, dstTy))
622 return arith::SIToFPOp::create(rewriter, loc, resultTypes, args,
626 if (isa<FloatType>(srcTy) && dstTy.
isInteger(1)) {
627 Value zero = arith::ConstantOp::create(rewriter, loc,
628 rewriter.getFloatAttr(srcTy, 0.0));
629 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::UNE,
633 if (arith::FPToSIOp::areCastCompatible(srcTy, dstTy)) {
634 auto rounded = math::RoundEvenOp::create(rewriter, loc, args[0]);
636 const auto &fltSemantics = cast<FloatType>(srcTy).getFloatSemantics();
640 APFloat::semanticsMaxExponent(fltSemantics)) {
643 auto conv = arith::FPToSIOp::create(rewriter, loc, dstTy, rounded);
644 auto posInf = arith::ConstantOp::create(
647 APFloat::getInf(fltSemantics)));
648 auto negInf = arith::ConstantOp::create(
650 rewriter.getFloatAttr(
652 APFloat::getInf(fltSemantics,
true)));
653 auto overflow = arith::CmpFOp::create(
654 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, posInf);
655 auto underflow = arith::CmpFOp::create(
656 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, negInf);
657 auto intMin = arith::ConstantOp::create(
659 rewriter.getIntegerAttr(
662 auto intMax = arith::ConstantOp::create(
664 rewriter.getIntegerAttr(
668 arith::SelectOp::create(rewriter, loc, overflow, intMax, conv);
669 return arith::SelectOp::create(rewriter, loc, underflow, intMin,
673 auto intMinFP = arith::ConstantOp::create(
675 rewriter.getFloatAttr(
681 if (cast<FloatType>(srcTy).getFPMantissaWidth() >=
687 auto intMaxFP = arith::ConstantOp::create(
689 rewriter.getFloatAttr(
696 return arith::FPToSIOp::create(rewriter, loc, dstTy, clamped);
703 auto intMaxPlusOneFP = arith::ConstantOp::create(
705 rewriter.getFloatAttr(
712 auto intMax = arith::ConstantOp::create(
714 rewriter.getIntegerAttr(
718 arith::MaximumFOp::create(rewriter, loc, rounded, intMinFP);
720 arith::FPToSIOp::create(rewriter, loc, dstTy, minClampedFP);
721 auto overflow = arith::CmpFOp::create(
722 rewriter, loc, arith::CmpFPredicate::UGE, rounded, intMaxPlusOneFP);
723 return arith::SelectOp::create(rewriter, loc, overflow, intMax,
729 if (isa<IntegerType>(srcTy) && dstTy.
isInteger(1)) {
732 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ne,
736 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && bitExtend)
737 return arith::ExtSIOp::create(rewriter, loc, resultTypes, args,
740 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && !bitExtend) {
741 return arith::TruncIOp::create(rewriter, loc, dstTy, args[0]);
745 (
void)rewriter.notifyMatchFailure(
746 op,
"unhandled op for linalg body calculation for elementwise op");
767 return tensor::DimOp::create(rewriter, loc,
tensor, indexValue).getResult();
773 auto shapedType = dyn_cast<ShapedType>(
tensor.getType());
774 assert(shapedType && shapedType.hasRank() &&
"expected a ranked shaped type");
775 assert(
index >= 0 &&
index < shapedType.getRank() &&
"index out of bounds");
776 if (shapedType.isDynamicDim(
index))
782 auto isRanked = [](
Value value) {
783 return isa<RankedTensorType>(value.getType());
785 return llvm::all_of(operation->
getOperands(), isRanked) &&
786 llvm::all_of(operation->
getResults(), isRanked);
799static std::pair<OpFoldResult, Value>
805 for (
auto operand : operands) {
806 auto size = cast<RankedTensorType>(operand.getType()).getDimSize(dim);
807 if (ShapedType::isStatic(size) && size > 1)
812 auto operandsWithDynamicDim =
813 llvm::filter_to_vector(operands, [&](
Value operand) {
814 return cast<RankedTensorType>(operand.
getType()).isDynamicDim(dim);
818 if (operandsWithDynamicDim.empty())
825 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[0], dim);
826 if (operandsWithDynamicDim.size() == 1)
827 return {targetSize, operandsWithDynamicDim[0]};
830 for (
size_t i = 1; i < operandsWithDynamicDim.size(); i++) {
832 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[i], dim);
833 targetSize = arith::MaxUIOp::create(rewriter, loc, targetSize, nextSize);
835 return {targetSize,
nullptr};
843 assert(!operands.empty());
844 auto rank = cast<RankedTensorType>(operands.front().
getType()).getRank();
847 for (
auto dim : llvm::seq<int64_t>(0, rank)) {
848 auto [targetSize, masterOperand] =
850 targetShape.push_back(targetSize);
851 masterOperands.push_back(masterOperand);
853 return {targetShape, masterOperands};
859 Value masterOperand) {
861 auto rankedTensorType = cast<RankedTensorType>(operand.
getType());
862 if (!rankedTensorType.isDynamicDim(dim))
869 if (operand == masterOperand)
873 auto rank = rankedTensorType.getRank();
875 for (
auto index : llvm::seq<int64_t>(0, rank)) {
878 affineExprs.push_back(affineExpr);
880 auto broadcastAffineMap =
886 auto one =
createIndex(rewriter, loc, indexPool, 1);
887 auto runtimeSize =
getTensorDim(rewriter, loc, indexPool, operand, dim);
888 auto broadcastNecessary = arith::CmpIOp::create(
889 rewriter, loc, arith::CmpIPredicate::eq, runtimeSize, one);
899 for (
auto index : llvm::seq<int64_t>(0, rank)) {
900 auto size =
index == dim ? targetSize
903 outputTensorShape.push_back(size);
905 Value outputTensor = tensor::EmptyOp::create(
906 opBuilder, loc, outputTensorShape, rankedTensorType.getElementType());
910 linalg::GenericOp::create(
911 opBuilder, loc, outputTensor.
getType(), operand, outputTensor,
915 linalg::YieldOp::create(opBuilder, loc, blockArgs.front());
920 auto castResultTensor = rewriter.
createOrFold<tensor::CastOp>(
921 loc, operand.
getType(), resultTensor);
924 scf::YieldOp::create(opBuilder, loc, castResultTensor);
929 scf::YieldOp::create(opBuilder, loc, operand);
933 auto ifOp = scf::IfOp::create(rewriter, loc, broadcastNecessary,
934 emitThenRegion, emitElseRegion);
935 return ifOp.getResult(0);
942 int64_t rank = cast<RankedTensorType>(operand.
getType()).getRank();
943 assert((
int64_t)targetShape.size() == rank);
944 assert((
int64_t)masterOperands.size() == rank);
945 for (
auto index : llvm::seq<int64_t>(0, rank))
958 if (operands.size() == 1)
962 bool hasDynamic =
false;
963 for (
auto op : operands) {
964 const auto tType = dyn_cast<RankedTensorType>(op.getType());
965 if (tType && !tType.hasStaticShape()) {
974 return llvm::map_to_vector(operands, [&](
Value operand) {
976 targetShape, masterOperands);
986 auto resultType = cast_or_null<RankedTensorType>(
989 return rewriter.notifyMatchFailure(operation,
"failed to convert type");
991 Value outputTensor = tensor::EmptyOp::create(rewriter, loc, targetShape,
992 resultType.getElementType());
997 auto rank = resultType.getRank();
998 auto affineMaps = llvm::map_to_vector(operands, [&](
Value operand) {
999 auto shape = cast<ShapedType>(operand.
getType()).getShape();
1001 for (
auto it : llvm::enumerate(
shape)) {
1005 bool requiresBroadcast =
1006 (it.value() == 1 && resultType.getDimSize(it.index()) != 1);
1007 auto affineExpr = requiresBroadcast
1008 ? rewriter.getAffineConstantExpr(0)
1009 : rewriter.getAffineDimExpr(it.index());
1010 affineExprs.push_back(affineExpr);
1012 return AffineMap::get(rank, 0, affineExprs, rewriter.getContext());
1014 affineMaps.push_back(rewriter.getMultiDimIdentityMap(rank));
1017 bool encounteredError =
false;
1018 auto linalgOp = linalg::GenericOp::create(
1019 rewriter, loc, outputTensor.
getType(), operands, outputTensor, affineMaps,
1024 {resultType.getElementType()}, rewriter);
1026 encounteredError =
true;
1029 linalg::YieldOp::create(opBuilder, loc, opResult);
1031 if (encounteredError)
1032 return rewriter.notifyMatchFailure(
1033 operation,
"unable to create linalg.generic body for elementwise op");
1036 auto castResult = rewriter.createOrFold<tensor::CastOp>(
1037 loc, resultType, linalgOp->getResult(0));
1038 rewriter.replaceOp(operation, castResult);
1045 if (isa<tosa::MulOp>(operation)) {
1049 return operands.take_front(2);
1051 return operands.take_front(3);
1053 if (
auto negate = dyn_cast<tosa::NegateOp>(operation)) {
1054 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
1055 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
1056 if (failed(maybeOutZp) && failed(maybeInZp))
1059 return operands.take_front(1);
1066 ConversionPatternRewriter &rewriter,
1070 assert(operation->
getNumResults() == 1 &&
"elementwise op expects 1 result");
1072 "elementwise op expects at least 1 operand");
1074 return rewriter.notifyMatchFailure(operation,
1075 "Unranked tensors not supported");
1079 auto loc = operation->
getLoc();
1081 auto [targetShape, masterOperands] =
1083 auto broadcastOperands =
1085 targetShape, masterOperands);
1087 targetShape, converter);
1099 bool negative,
bool allowNonFinites) {
1100 if (allowNonFinites && APFloat::semanticsHasInf(semantics))
1101 return APFloat::getInf(semantics, negative);
1102 return APFloat::getLargest(semantics, negative);
1109 bool allowNonFinites) {
1110 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy))
1113 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy))
1116 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy))
1119 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy))
1122 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy))
1126 false, allowNonFinites));
1128 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy))
1132 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy))
1136 true, allowNonFinites));
1138 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy))
1142 if (isa<tosa::ReduceAllOp>(op) && elementTy.
isInteger(1))
1145 if (isa<tosa::ReduceAnyOp>(op) && elementTy.
isInteger(1))
1148 if (isa<tosa::ArgMaxOp>(op) && isa<FloatType>(elementTy))
1152 true, allowNonFinites));
1154 if (isa<tosa::ArgMaxOp>(op) && isa<IntegerType>(elementTy))
1168 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy)) {
1170 rewriter, loc,
TypeRange{elementTy}, args);
1173 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy)) {
1175 rewriter, loc,
TypeRange{elementTy}, args);
1178 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy)) {
1180 rewriter, loc,
TypeRange{elementTy}, args);
1183 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy)) {
1185 rewriter, loc,
TypeRange{elementTy}, args);
1188 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy)) {
1189 return arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
1192 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy)) {
1193 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
1196 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy)) {
1197 return arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
1200 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy)) {
1201 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
1204 if (isa<tosa::ReduceAllOp>(op) && elementTy.
isInteger(1))
1205 return arith::AndIOp::create(rewriter, loc, args);
1207 if (isa<tosa::ReduceAnyOp>(op) && elementTy.
isInteger(1))
1208 return arith::OrIOp::create(rewriter, loc, args);
1216template <
typename OpTy>
1219 bool allowNonFinites) {
1220 auto loc = op->getLoc();
1221 auto inputTy = dyn_cast<RankedTensorType>(op->getOperand(0).getType());
1222 auto resultTy = dyn_cast<RankedTensorType>(op->getResult(0).getType());
1223 if (!inputTy || !resultTy)
1226 auto elementTy = resultTy.getElementType();
1227 Value input = op->getOperand(0);
1230 bool widenAccTy = std::is_same_v<OpTy, tosa::ReduceSumOp> &&
1231 isa<FloatType>(elementTy) &&
1232 cast<FloatType>(elementTy).isBF16();
1237 for (
unsigned i = 0; i < inputTy.getRank(); i++) {
1239 reduceShape.push_back(inputTy.getDimSize(i));
1240 if (inputTy.isDynamicDim(i))
1241 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1246 inputs.push_back(input);
1250 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1253 auto fillValueAttr =
1257 op,
"No initial value found for reduction operation");
1259 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
1261 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValue},
1264 outputs.push_back(filledTensor);
1266 bool isNanIgnoreMode =
false;
1267 if constexpr (std::is_same_v<OpTy, tosa::ReduceMinOp> ||
1268 std::is_same_v<OpTy, tosa::ReduceMaxOp>) {
1270 if (isa<FloatType>(elementTy) &&
1271 op.getNanMode() == NanPropagationMode::IGNORE) {
1272 isNanIgnoreMode =
true;
1278 auto trueValue = arith::ConstantOp::create(rewriter, loc, trueAttr);
1279 auto emptyBoolTensor =
1280 tensor::EmptyOp::create(rewriter, loc, reduceShape,
1281 trueValue.getType(), dynDims)
1283 auto allResultsNaNTensor =
1284 linalg::FillOp::create(rewriter, loc,
ValueRange{trueValue},
1296 inputs.push_back(input);
1297 outputs.push_back(allResultsNaNTensor);
1301 bool didEncounterError =
false;
1302 linalg::LinalgOp linalgOp = linalg::ReduceOp::create(
1303 rewriter, loc, inputs, outputs, axis,
1305 std::array<Value, 2> binaryArgs{
1306 blockArgs[0], isNanIgnoreMode ? blockArgs[2] : blockArgs[1]};
1309 if (binaryArgs[0].
getType() != accTy)
1310 binaryArgs[0] = arith::ExtFOp::create(
1311 nestedBuilder, nestedLoc,
TypeRange{accTy},
1312 ValueRange{binaryArgs[0]}, arith::ExtFOp::Properties{});
1317 didEncounterError =
true;
1320 if (isNanIgnoreMode) {
1321 auto inputValue = blockArgs[0];
1322 auto initialValue = blockArgs[2];
1323 auto oldAllResultsNanFlagValue = blockArgs[3];
1326 Value isNaN = arith::CmpFOp::create(nestedBuilder, op->getLoc(),
1327 arith::CmpFPredicate::UNO,
1328 inputValue, inputValue);
1330 auto selectOp = arith::SelectOp::create(nestedBuilder, op->getLoc(),
1331 isNaN, initialValue,
result);
1334 auto newAllResultsNanFlagValue = arith::AndIOp::create(
1335 nestedBuilder, op->getLoc(), oldAllResultsNanFlagValue, isNaN);
1336 resultsToYield.push_back(selectOp);
1337 resultsToYield.push_back(newAllResultsNanFlagValue);
1339 resultsToYield.push_back(
result);
1341 linalg::YieldOp::create(nestedBuilder, loc, resultsToYield);
1344 if (!didEncounterError)
1346 op,
"unable to create linalg.generic body for reduce op");
1348 if (isNanIgnoreMode) {
1358 auto nanValue = arith::ConstantOp::create(rewriter, loc, nanValueAttr);
1359 auto emptyNanTensor =
1360 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1362 auto nanFilledTensor =
1363 linalg::FillOp::create(rewriter, loc,
ValueRange{nanValue},
1369 auto finalEmptyTensor =
1370 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1376 ins.push_back(linalgOp->getOpResult(1));
1377 ins.push_back(nanFilledTensor);
1378 ins.push_back(linalgOp->getResult(0));
1379 outs.push_back(finalEmptyTensor);
1381 linalg::ElementwiseOp::create(rewriter, op->getLoc(), ins, outs,
1382 mlir::linalg::ElementwiseKind::select);
1383 linalgOp = linalgSelect;
1387 Value reducedRes = linalgOp->getResult(0);
1390 tensor::EmptyOp::create(rewriter, loc, reduceShape, elementTy, dynDims)
1393 const unsigned reducedRank =
1394 cast<ShapedType>(reducedRes.
getType()).getRank();
1397 linalg::GenericOp::create(
1403 Value truncf = arith::TruncFOp::create(nestedBuilder, nestedLoc,
1404 elementTy, args[0]);
1405 linalg::YieldOp::create(nestedBuilder, nestedLoc, truncf);
1411 uint64_t expandInputRank = cast<ShapedType>(reducedRes.
getType()).getRank();
1412 reassociationMap.resize(expandInputRank);
1414 for (uint64_t i = 0; i < expandInputRank; i++) {
1415 int32_t dimToPush = i > axis ? i + 1 : i;
1419 if (expandInputRank != 0) {
1420 int32_t expandedDim = axis < expandInputRank ? axis : expandInputRank - 1;
1421 reassociationMap[expandedDim].push_back(
1436template <
typename SrcOp>
1437class PointwiseConverter :
public OpConversionPattern<SrcOp> {
1439 using OpConversionPattern<SrcOp>::OpConversionPattern;
1440 using typename OpConversionPattern<SrcOp>::OpAdaptor;
1443 matchAndRewrite(SrcOp op, OpAdaptor operands,
1444 ConversionPatternRewriter &rewriter)
const final {
1446 op, operands.getOperands(), rewriter, *this->getTypeConverter());
1456 auto inputType = cast<RankedTensorType>(input.
getType());
1457 auto elemType = inputType.getElementType();
1458 auto collapsedType = RankedTensorType::get({}, elemType);
1460 return tensor::CollapseShapeOp::create(rewriter, loc, collapsedType, input,
1467 output.reserve(input.size());
1469 for (
auto v : llvm::map_range(
1470 input, [](int32_t val) {
return static_cast<int8_t
>(val); })) {
1471 output.push_back(v);
1483static void setupLinalgGenericOpInputAndIndexingMap(
1486 bool isConstant, tosa::RescaleOp op,
Value &constant,
int64_t &arg,
1487 bool isShift =
false) {
1489 auto loc = op.getLoc();
1490 auto inputTy = cast<ShapedType>(op.getInput().getType());
1491 unsigned rank = inputTy.getRank();
1497 if (values.size() == 1) {
1498 IntegerAttr intAttr = isShift
1501 constant = arith::ConstantOp::create(rewriter, loc, intAttr);
1505 auto tensorType = RankedTensorType::get(
1506 {
static_cast<int64_t>(values.size())}, elementType);
1512 genericInputs.push_back(
1513 arith::ConstantOp::create(rewriter, loc, EltAttr));
1521 auto operand = isShift ? op.getShift() : op.getMultiplier();
1522 auto tensorType = dyn_cast<RankedTensorType>(operand.getType());
1523 if (tensorType && tensorType.hasStaticShape() &&
1524 tensorType.getShape()[0] == 1) {
1529 genericInputs.push_back(collapse1xNTensorToN(rewriter, operand, loc));
1530 indexingMaps.push_back(broadcastMap);
1532 genericInputs.push_back(operand);
1538 arg = indexingMaps.size() - 1;
1543 FailureOr<int64_t> maybeZp,
Location loc,
1545 bool isOutputZp =
false) {
1548 const uint32_t attrBitwidth =
1549 isOutputZp ? 32 : (bitwidth > 32 ? bitwidth : 32);
1556 result = blockArgs[zpArg];
1557 auto zpTy =
result.getType();
1558 if (zpTy.getIntOrFloatBitWidth() < attrBitwidth) {
1561 if (zpTy.isUnsignedInteger()) {
1563 UnrealizedConversionCastOp::create(
1568 if (zpTy.isUnsignedInteger()) {
1569 return arith::ExtUIOp::create(builder, loc, extendType,
result);
1571 return arith::ExtSIOp::create(builder, loc, extendType,
result);
1575 return arith::ConstantOp::create(builder, loc,
1576 IntegerAttr::get(extendType, *maybeZp));
1583 using OpRewritePattern<tosa::RescaleOp>::OpRewritePattern;
1585 LogicalResult matchAndRewrite(tosa::RescaleOp op,
1586 PatternRewriter &rewriter)
const final {
1587 auto loc = op.getLoc();
1588 auto input = op.getInput();
1589 auto inputTy = cast<ShapedType>(op.getInput().getType());
1590 auto outputTy = cast<ShapedType>(op.getOutput().getType());
1591 unsigned rank = inputTy.getRank();
1594 if (op.getRoundingMode() == RoundingMode::INEXACT_ROUND)
1596 op,
"tosa.rescale with rounding mode = 'INEXACT_ROUND' is not "
1597 "currently supported");
1598 if (op.getRoundingMode() == RoundingMode::DOUBLE_ROUND && !op.getScale32())
1600 op,
"tosa.rescale requires scale32 for double_round to be true");
1602 if (!isa<IntegerType>(inputTy.getElementType()))
1605 SmallVector<Value> dynDims;
1606 for (
int i = 0; i < outputTy.getRank(); i++) {
1607 if (outputTy.isDynamicDim(i)) {
1608 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1612 DenseElementsAttr shiftElems;
1613 bool isShiftConstant =
false;
1615 isShiftConstant =
true;
1617 DenseElementsAttr multiplierElems;
1618 bool isMultiplierConstant =
false;
1620 isMultiplierConstant =
true;
1622 llvm::SmallVector<int32_t> shiftValues;
1623 llvm::SmallVector<int32_t> multiplierValues;
1626 if (isMultiplierConstant && isShiftConstant) {
1628 shiftValues = llvm::map_to_vector(
1629 shiftElems.
getValues<IntegerAttr>(), [](IntegerAttr attr) -> int32_t {
1630 return static_cast<int32_t>(attr.getInt());
1633 llvm::map_to_vector(multiplierElems.
getValues<IntegerAttr>(),
1634 [](IntegerAttr attr) -> int32_t {
1635 return static_cast<int32_t>(attr.getInt());
1639 for (
int i = 0, s = multiplierValues.size(); i < s; i++) {
1640 if (shiftValues[i] > 63) {
1642 multiplierValues[i] = 0;
1647 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND &&
1648 llvm::any_of(shiftValues, [](int32_t v) {
return v > 31; });
1650 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND;
1652 RoundingMode roundingMode =
1653 doubleRound ? RoundingMode::DOUBLE_ROUND : RoundingMode::SINGLE_ROUND;
1655 SmallVector<AffineMap> indexingMaps = {
1657 SmallVector<Value, 4> genericInputs = {input};
1661 Value multiplierConstant;
1662 int64_t multiplierArg = 0;
1663 setupLinalgGenericOpInputAndIndexingMap(
1664 rewriter, multiplierValues, genericInputs, indexingMaps,
1665 isMultiplierConstant, op, multiplierConstant, multiplierArg);
1669 Value shiftConstant;
1670 int64_t shiftArg = 0;
1671 setupLinalgGenericOpInputAndIndexingMap(
1672 rewriter, shiftValues, genericInputs, indexingMaps, isShiftConstant, op,
1673 shiftConstant, shiftArg,
true);
1678 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1679 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1689 genericInputs.push_back(
1690 collapse1xNTensorToN(rewriter, op->getOperand(3), loc));
1691 indexingMaps.push_back(broadcastMap);
1692 iZpArg = indexingMaps.size() - 1;
1696 genericInputs.push_back(
1697 collapse1xNTensorToN(rewriter, op->getOperand(4), loc));
1698 indexingMaps.push_back(broadcastMap);
1699 oZpArg = indexingMaps.size() - 1;
1706 Value emptyTensor = tensor::EmptyOp::create(
1707 rewriter, loc, outputTy.getShape(), outputTy.getElementType(),
1708 ArrayRef<Value>({dynDims}));
1710 auto linalgOp = linalg::GenericOp::create(
1711 rewriter, loc, outputTy, genericInputs,
ValueRange{emptyTensor},
1713 [&](OpBuilder &nestedBuilder, Location nestedLoc,
1715 Value value = blockArgs[0];
1716 Type valueTy = value.
getType();
1718 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1719 auto inputZp = getExtendZp(nestedBuilder, valueTy, maybeIZp,
1720 nestedLoc, blockArgs, iZpArg);
1722 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1723 auto outputZp = getExtendZp(nestedBuilder, valueTy, maybeOZp,
1724 nestedLoc, blockArgs, oZpArg,
true);
1726 IntegerType outIntType =
1727 cast<IntegerType>(blockArgs.back().
getType());
1728 unsigned outBitWidth = outIntType.getWidth();
1729 assert(outBitWidth <= 32 &&
"Unexpected output zeropoint bitwidth");
1731 Value multiplier = multiplierConstant ? multiplierConstant
1732 : blockArgs[multiplierArg];
1733 Value shift = shiftConstant ? shiftConstant : blockArgs[shiftArg];
1736 value = UnrealizedConversionCastOp::create(
1737 nestedBuilder, nestedLoc,
1738 nestedBuilder.getIntegerType(
1744 if (op.getInputUnsigned()) {
1745 value = arith::ExtUIOp::create(nestedBuilder, nestedLoc,
1746 nestedBuilder.getI32Type(), value);
1748 value = arith::ExtSIOp::create(nestedBuilder, nestedLoc,
1749 nestedBuilder.getI32Type(), value);
1754 arith::SubIOp::create(nestedBuilder, nestedLoc, value, inputZp);
1756 value = tosa::ApplyScaleOp::create(nestedBuilder, loc,
1757 nestedBuilder.getI32Type(), value,
1758 multiplier, shift, roundingMode);
1762 arith::AddIOp::create(nestedBuilder, nestedLoc, value, outputZp);
1765 int32_t intMin = APInt::getSignedMinValue(outBitWidth).getSExtValue();
1766 int32_t intMax = APInt::getSignedMaxValue(outBitWidth).getSExtValue();
1769 if (op.getOutputUnsigned()) {
1771 intMax = APInt::getMaxValue(outBitWidth).getZExtValue();
1774 auto intMinVal = arith::ConstantOp::create(
1775 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMin));
1776 auto intMaxVal = arith::ConstantOp::create(
1777 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMax));
1780 nestedBuilder,
false);
1782 if (outIntType.getWidth() < 32) {
1783 value = arith::TruncIOp::create(
1784 nestedBuilder, nestedLoc,
1788 if (outIntType.isUnsignedInteger()) {
1789 value = UnrealizedConversionCastOp::create(nestedBuilder, nestedLoc,
1793 linalg::YieldOp::create(nestedBuilder, loc, value);
1796 rewriter.
replaceOp(op, linalgOp->getResults());
1806 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
1808 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1809 PatternRewriter &rewriter)
const final {
1810 Location loc = op.getLoc();
1811 ImplicitLocOpBuilder builder(loc, rewriter);
1812 auto input = op.getInput();
1813 auto inputTy = cast<RankedTensorType>(input.getType());
1814 auto resultTy = cast<RankedTensorType>(op.getType());
1815 const bool isBilinear = op.getMode() == ResizeMode::BILINEAR;
1817 auto inputH = inputTy.getDimSize(1);
1818 auto inputW = inputTy.getDimSize(2);
1819 auto outputH = resultTy.getDimSize(1);
1820 auto outputW = resultTy.getDimSize(2);
1822 if (inputH != 1 || inputW != 1 || outputH != 1 || outputW != 1)
1824 op,
"tosa.resize is not a pure 1x1->1x1 image operation");
1826 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
1827 op.getMode() != ResizeMode::BILINEAR)
1829 op,
"tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
1831 if (inputTy == resultTy) {
1836 SmallVector<int64_t> scale;
1842 SmallVector<ReassociationExprs, 4> reassociationMap(2);
1849 RankedTensorType::get({inputTy.getDimSize(0), inputTy.getDimSize(3)},
1850 inputTy.getElementType());
1851 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy, input,
1855 llvm::SmallVector<Value> outputDynSize;
1856 if (inputTy.isDynamicDim(0))
1857 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1858 if (inputTy.isDynamicDim(3))
1859 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1862 auto genericTy = collapseTy.clone(resultTy.getElementType());
1864 tensor::EmptyOp::create(builder, genericTy.getShape(),
1865 resultTy.getElementType(), outputDynSize);
1867 SmallVector<utils::IteratorType> iterators(genericTy.getRank(),
1868 utils::IteratorType::parallel);
1870 auto generic = linalg::GenericOp::create(
1872 ArrayRef<AffineMap>{genericMap, genericMap}, iterators,
1873 [=](OpBuilder &
b, Location loc,
ValueRange args) {
1874 Value value = args[0];
1876 if (inputTy.getElementType() != resultTy.getElementType()) {
1877 value = arith::ExtSIOp::create(
b, loc, resultTy.getElementType(),
1880 if (isBilinear && scale[0] != 0) {
1881 Value scaleY = arith::ConstantOp::create(
1882 b, loc,
b.getI32IntegerAttr(scale[0]));
1883 value = arith::MulIOp::create(
b, loc, value, scaleY);
1886 if (isBilinear && scale[2] != 0) {
1887 Value scaleX = arith::ConstantOp::create(
1888 b, loc,
b.getI32IntegerAttr(scale[2]));
1889 value = arith::MulIOp::create(
b, loc, value, scaleX);
1893 linalg::YieldOp::create(
b, loc, value);
1897 op, resultTy,
generic.getResults()[0], reassociationMap);
1909 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1913 auto input = op.getInput();
1914 auto inputTy = dyn_cast<RankedTensorType>(input.getType());
1915 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
1917 if (!inputTy || !resultTy)
1919 "requires ranked input/output types");
1921 auto batch = inputTy.getDimSize(0);
1922 auto channels = inputTy.getDimSize(3);
1923 auto inputH = inputTy.getDimSize(1);
1924 auto inputW = inputTy.getDimSize(2);
1925 auto outputH = resultTy.getDimSize(1);
1926 auto outputW = resultTy.getDimSize(2);
1928 if ((inputH != 1 || outputH == 1) && (inputW != 1 || outputW == 1))
1930 op,
"tosa.resize has no broadcasting behavior");
1935 resizeShape.push_back(batch);
1936 resizeShape.push_back(inputH == 1 ? 1 : outputH);
1937 resizeShape.push_back(inputW == 1 ? 1 : outputW);
1938 resizeShape.push_back(channels);
1940 auto resizeTy = resultTy.clone(resizeShape);
1942 tosa::ResizeOp::create(builder, resizeTy, input, op.getScale(),
1943 op.getOffset(), op.getBorder(), op.getMode());
1950 reassociationMap.push_back({});
1953 reassociationMap.push_back({});
1958 collapseShape.push_back(outputH);
1960 collapseShape.push_back(outputW);
1961 collapseShape.push_back(channels);
1963 auto collapseTy = resultTy.clone(collapseShape);
1964 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy,
1965 resize, reassociationMap);
1969 if (inputTy.isDynamicDim(0))
1970 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1971 if (inputTy.isDynamicDim(3))
1972 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1975 utils::IteratorType::parallel);
1976 Value empty = tensor::EmptyOp::create(
1977 builder, resultTy.getShape(), resultTy.getElementType(), outputDynSize);
1994 Value value = args[0];
1995 linalg::YieldOp::create(
b, loc, value);
2004 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
2006 LogicalResult matchAndRewrite(tosa::ResizeOp op,
2007 PatternRewriter &rewriter)
const final {
2008 Location loc = op.getLoc();
2009 ImplicitLocOpBuilder
b(loc, rewriter);
2010 auto input = op.getInput();
2011 auto inputTy = cast<ShapedType>(input.getType());
2012 auto resultTy = cast<ShapedType>(op.getType());
2013 auto resultETy = resultTy.getElementType();
2015 bool floatingPointMode = isa<FloatType>(resultETy);
2016 auto floatTy = resultETy;
2018 auto imageH = inputTy.getShape()[1];
2019 auto imageW = inputTy.getShape()[2];
2021 auto dynamicDimsOr =
2023 if (!dynamicDimsOr.has_value())
2025 op,
"unable to get dynamic dimensions of tosa.resize");
2027 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
2028 op.getMode() != ResizeMode::BILINEAR)
2030 op,
"tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
2032 SmallVector<AffineMap, 2> affineMaps = {
2034 auto emptyTensor = tensor::EmptyOp::create(
b, resultTy.getShape(),
2035 resultETy, *dynamicDimsOr);
2036 auto genericOp = linalg::GenericOp::create(
2039 Value resize = genericOp.getResult(0);
2042 OpBuilder::InsertionGuard regionGuard(
b);
2043 b.createBlock(&genericOp.getRegion(), genericOp.getRegion().end(),
2045 Value batch = linalg::IndexOp::create(
b, 0);
2046 Value y = linalg::IndexOp::create(
b, 1);
2047 Value x = linalg::IndexOp::create(
b, 2);
2048 Value channel = linalg::IndexOp::create(
b, 3);
2051 arith::ConstantOp::create(
b,
b.getZeroAttr(
b.getI32Type()));
2052 Value zeroFp = arith::ConstantOp::create(
b,
b.getZeroAttr(floatTy));
2054 arith::ConstantOp::create(
b,
b.getI32IntegerAttr(imageH - 1));
2056 arith::ConstantOp::create(
b,
b.getI32IntegerAttr(imageW - 1));
2058 Value inY = arith::IndexCastOp::create(
b,
b.getI32Type(), y);
2059 Value inX = arith::IndexCastOp::create(
b,
b.getI32Type(), x);
2061 SmallVector<int64_t> scale, offset, border;
2066 op,
"tosa.resize scale/offset/border should have compile time "
2067 "constant values.");
2070 Value yScaleN, yScaleD, xScaleN, xScaleD;
2071 yScaleN = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[0]));
2072 yScaleD = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[1]));
2073 xScaleN = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[2]));
2074 xScaleD = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[3]));
2076 Value yOffset, xOffset, yBorder, xBorder;
2077 yOffset = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(offset[0]));
2078 xOffset = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(offset[1]));
2079 yBorder = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(border[0]));
2080 xBorder = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(border[1]));
2083 auto getIndexAndDeltaFp = [&](Value &index, Value &delta, Value in,
2084 Value scaleN, Value scaleD, Value offset,
2085 int size, ImplicitLocOpBuilder &
b) {
2093 Value val = arith::MulIOp::create(
b, in, scaleD);
2094 val = arith::AddIOp::create(
b, val, offset);
2095 index = arith::FloorDivSIOp::create(
b, val, scaleN);
2098 Value scaledIndex = arith::MulIOp::create(
b, index, scaleN);
2099 Value r = arith::SubIOp::create(
b, val, scaledIndex);
2100 Value rFp = arith::SIToFPOp::create(
b, floatTy, r);
2103 Value scaleNfp = arith::UIToFPOp::create(
b, floatTy, scaleN);
2104 delta = arith::DivFOp::create(
b, rFp, scaleNfp);
2108 auto getIndexAndDeltaInt = [&](Value &index, Value &delta, Value in,
2109 Value scaleN, Value scaleD, Value offset,
2110 int size, ImplicitLocOpBuilder &
b) {
2119 Value val = arith::MulIOp::create(
b, in, scaleD);
2120 val = arith::AddIOp::create(
b, val, offset);
2121 index = arith::FloorDivSIOp::create(
b, val, scaleN);
2122 delta = arith::MulIOp::create(
b, index, scaleN);
2123 delta = arith::SubIOp::create(
b, val, delta);
2126 Value ix, iy, dx, dy;
2127 if (floatingPointMode) {
2128 getIndexAndDeltaFp(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH,
b);
2129 getIndexAndDeltaFp(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW,
b);
2131 getIndexAndDeltaInt(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH,
b);
2132 getIndexAndDeltaInt(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW,
b);
2135 if (op.getMode() == ResizeMode::NEAREST_NEIGHBOR) {
2136 auto one = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(1));
2138 auto getNearestIndexAndClamp = [&](Value val, Value dval, Value scale,
2139 Value
max,
int size,
2140 ImplicitLocOpBuilder &
b) -> Value {
2146 if (floatingPointMode) {
2148 arith::ConstantOp::create(
b,
b.getFloatAttr(floatTy, 0.5f));
2149 pred = arith::CmpFOp::create(
b, arith::CmpFPredicate::OGE, dval, h);
2151 Value dvalDouble = arith::ShLIOp::create(
b, dval, one);
2152 pred = arith::CmpIOp::create(
b, arith::CmpIPredicate::sge,
2156 auto offset = arith::SelectOp::create(
b, pred, one, zeroI32);
2157 val = arith::AddIOp::create(
b, val, offset);
2159 return arith::IndexCastOp::create(
b,
b.getIndexType(), val);
2162 iy = getNearestIndexAndClamp(iy, dy, yScaleN, hMax, imageH,
b);
2163 ix = getNearestIndexAndClamp(ix, dx, xScaleN, wMax, imageW,
b);
2165 Value
result = tensor::ExtractOp::create(
2168 linalg::YieldOp::create(
b,
result);
2171 assert(op.getMode() == ResizeMode::BILINEAR);
2173 auto oneVal = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(1));
2175 auto getClampedIdxs = [&](Value &val0, Value &val1,
int size, Value in,
2176 Value
max, ImplicitLocOpBuilder &
b) {
2178 val1 = arith::AddIOp::create(
b, val0, oneVal);
2183 val0 = arith::IndexCastOp::create(
b,
b.getIndexType(), val0);
2184 val1 = arith::IndexCastOp::create(
b,
b.getIndexType(), val1);
2192 Value x0, x1, y0, y1;
2193 getClampedIdxs(y0, y1, imageH, iy, hMax,
b);
2194 getClampedIdxs(x0, x1, imageW, ix, wMax,
b);
2196 Value y0x0 = tensor::ExtractOp::create(
2198 Value y0x1 = tensor::ExtractOp::create(
2200 Value y1x0 = tensor::ExtractOp::create(
2202 Value y1x1 = tensor::ExtractOp::create(
2205 if (floatingPointMode) {
2207 arith::ConstantOp::create(
b,
b.getFloatAttr(floatTy, 1.0f));
2208 auto interpolate = [&](Value val0, Value val1, Value delta,
2210 ImplicitLocOpBuilder &
b) -> Value {
2213 Value oneMinusDelta = arith::SubFOp::create(
b, oneVal, delta);
2214 Value mul0 = arith::MulFOp::create(
b, val0, oneMinusDelta);
2215 Value mul1 = arith::MulFOp::create(
b, val1, delta);
2216 return arith::AddFOp::create(
b, mul0, mul1);
2222 Value topAcc = interpolate(y0x0, y0x1, dx, imageW,
b);
2227 Value bottomAcc = interpolate(y1x0, y1x1, dx, imageW,
b);
2231 Value
result = interpolate(topAcc, bottomAcc, dy, imageH,
b);
2232 linalg::YieldOp::create(
b,
result);
2235 y0x0 = arith::ExtSIOp::create(
b, resultETy, y0x0);
2236 y0x1 = arith::ExtSIOp::create(
b, resultETy, y0x1);
2237 y1x0 = arith::ExtSIOp::create(
b, resultETy, y1x0);
2238 y1x1 = arith::ExtSIOp::create(
b, resultETy, y1x1);
2241 if (resultETy.getIntOrFloatBitWidth() > deltaBitwidth) {
2242 dx = arith::ExtSIOp::create(
b, resultETy, dx);
2243 dy = arith::ExtSIOp::create(
b, resultETy, dy);
2246 Value yScaleNExt = yScaleN;
2247 Value xScaleNExt = xScaleN;
2249 const int64_t scaleBitwidth =
2251 if (resultETy.getIntOrFloatBitWidth() > scaleBitwidth) {
2252 yScaleNExt = arith::ExtSIOp::create(
b, resultETy, yScaleN);
2253 xScaleNExt = arith::ExtSIOp::create(
b, resultETy, xScaleN);
2256 auto interpolate = [](Value val0, Value val1, Value weight1,
2257 Value scale,
int inputSize,
2258 ImplicitLocOpBuilder &
b) -> Value {
2260 return arith::MulIOp::create(
b, val0, scale);
2261 Value weight0 = arith::SubIOp::create(
b, scale, weight1);
2262 Value mul0 = arith::MulIOp::create(
b, val0, weight0);
2263 Value mul1 = arith::MulIOp::create(
b, val1, weight1);
2264 return arith::AddIOp::create(
b, mul0, mul1);
2267 Value topAcc = interpolate(y0x0, y0x1, dx, xScaleNExt, imageW,
b);
2268 Value bottomAcc = interpolate(y1x0, y1x1, dx, xScaleNExt, imageW,
b);
2270 interpolate(topAcc, bottomAcc, dy, yScaleNExt, imageH,
b);
2271 linalg::YieldOp::create(
b,
result);
2284template <
typename SrcOp>
2287 using OpRewritePattern<SrcOp>::OpRewritePattern;
2289 LogicalResult matchAndRewrite(SrcOp op,
2290 PatternRewriter &rewriter)
const final {
2291 rewriter.
replaceOp(op, op.getOperation()->getOperands());
2296template <
typename SrcOp>
2299 ReduceConverter(MLIRContext *context,
bool allowNonFinites)
2300 : OpRewritePattern<SrcOp>(context), allowNonFinites(allowNonFinites) {}
2302 LogicalResult matchAndRewrite(SrcOp reduceOp,
2303 PatternRewriter &rewriter)
const final {
2309 bool allowNonFinites;
2314 using OpRewritePattern<tosa::ReverseOp>::OpRewritePattern;
2316 LogicalResult matchAndRewrite(tosa::ReverseOp op,
2317 PatternRewriter &rewriter)
const final {
2318 auto loc = op.getLoc();
2319 Value input = op.getInput1();
2320 auto inputTy = cast<ShapedType>(input.
getType());
2321 auto resultTy = cast<ShapedType>(op.getType());
2322 auto axis = op.getAxis();
2324 SmallVector<Value> dynDims;
2325 for (
int i = 0; i < inputTy.getRank(); i++) {
2326 if (inputTy.isDynamicDim(i)) {
2327 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2331 Value axisDimSize = tensor::DimOp::create(rewriter, loc, input, axis);
2334 auto emptyTensor = tensor::EmptyOp::create(
2335 rewriter, loc, inputTy.getShape(),
2336 inputTy.getElementType(), ArrayRef<Value>({dynDims}))
2338 SmallVector<AffineMap, 2> affineMaps = {
2342 op, resultTy, ArrayRef<Value>({}),
ValueRange{emptyTensor}, affineMaps,
2344 [&](OpBuilder &nestedBuilder, Location nestedLoc,
ValueRange args) {
2345 llvm::SmallVector<Value>
indices;
2346 for (
unsigned int i = 0; i < inputTy.getRank(); i++) {
2348 linalg::IndexOp::create(rewriter, nestedLoc, i).getResult();
2352 arith::SubIOp::create(rewriter, nestedLoc, axisDimSize, one);
2353 index = arith::SubIOp::create(rewriter, nestedLoc, sizeMinusOne,
2360 auto extract = tensor::ExtractOp::create(nestedBuilder, nestedLoc,
2362 linalg::YieldOp::create(nestedBuilder, op.getLoc(),
2363 extract.getResult());
2373struct TileConverter :
public OpConversionPattern<tosa::TileOp> {
2374 using OpConversionPattern<tosa::TileOp>::OpConversionPattern;
2377 matchAndRewrite(tosa::TileOp op, OpAdaptor adaptor,
2378 ConversionPatternRewriter &rewriter)
const override {
2379 auto loc = op.getLoc();
2380 auto input = op.getInput1();
2381 auto inputTy = cast<ShapedType>(input.
getType());
2382 auto inputShape = inputTy.getShape();
2383 auto resultTy = cast<ShapedType>(op.getType());
2384 auto elementTy = inputTy.getElementType();
2385 int64_t rank = inputTy.getRank();
2387 SmallVector<int64_t> multiples;
2388 if (
failed(op.getConstantMultiples(multiples)))
2392 SmallVector<int64_t, 2> genericShape;
2393 for (
int i = 0; i < rank; i++) {
2394 int64_t dim = multiples[i];
2395 genericShape.push_back(dim == -1 ? ShapedType::kDynamic : dim);
2396 genericShape.push_back(inputShape[i]);
2399 SmallVector<Value> dynDims;
2400 for (
int i = 0; i < inputTy.getRank(); i++) {
2401 if (inputTy.isDynamicDim(i) || multiples[i] == -1) {
2402 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2406 auto emptyTensor = tensor::EmptyOp::create(
2407 rewriter, op.getLoc(), genericShape, elementTy, dynDims);
2410 SmallVector<AffineExpr, 4> dimExprs;
2411 dimExprs.reserve(rank);
2412 for (
unsigned i = 0; i < rank; ++i)
2413 dimExprs.push_back(rewriter.getAffineDimExpr(i * 2 + 1));
2415 auto readAffineMap =
2417 rewriter.getContext());
2419 SmallVector<AffineMap, 2> affineMaps = {
2420 readAffineMap, rewriter.getMultiDimIdentityMap(genericShape.size())};
2422 auto genericOp = linalg::GenericOp::create(
2423 rewriter, loc, RankedTensorType::get(genericShape, elementTy), input,
2426 [&](OpBuilder &nestedBuilder, Location nestedLoc,
ValueRange args) {
2427 linalg::YieldOp::create(nestedBuilder, op.getLoc(), *args.begin());
2432 rewriter.replaceOpWithNewOp<tosa::ReshapeOp>(
2433 op, resultTy, genericOp.getResult(0), shapeValue);
2453 ArgMaxConverter(MLIRContext *context,
bool allowNonFinites)
2454 : OpRewritePattern<tosa::ArgMaxOp>(context),
2455 allowNonFinites(allowNonFinites) {}
2457 LogicalResult matchAndRewrite(tosa::ArgMaxOp argmaxOp,
2458 PatternRewriter &rewriter)
const final {
2459 auto loc = argmaxOp.getLoc();
2460 Value input = argmaxOp.getInput();
2461 auto inputTy = cast<ShapedType>(input.
getType());
2462 auto resultTy = cast<ShapedType>(argmaxOp.getOutput().getType());
2463 auto inElementTy = inputTy.getElementType();
2464 auto outElementTy = resultTy.getElementType();
2465 int axis = argmaxOp.getAxis();
2466 auto resultMaxTy = RankedTensorType::get(resultTy.getShape(), inElementTy);
2468 if (!isa<IntegerType>(outElementTy))
2469 return rewriter.notifyMatchFailure(
2471 "tosa.arg_max to linalg.* requires integer-like result type");
2473 SmallVector<Value> dynDims;
2474 for (
int i = 0; i < inputTy.getRank(); i++) {
2475 if (inputTy.isDynamicDim(i) && i != axis) {
2476 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2481 auto emptyTensorIdx =
2482 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2483 outElementTy, dynDims)
2485 auto fillValueIdx = arith::ConstantOp::create(
2486 rewriter, loc, rewriter.getIntegerAttr(outElementTy, 0));
2487 auto filledTensorIdx =
2488 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValueIdx},
2493 auto emptyTensorMax =
2494 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(), inElementTy,
2498 argmaxOp, inElementTy, rewriter, allowNonFinites);
2500 if (!fillValueMaxAttr)
2501 return rewriter.notifyMatchFailure(
2502 argmaxOp,
"unsupported tosa.argmax element type");
2505 arith::ConstantOp::create(rewriter, loc, fillValueMaxAttr);
2506 auto filledTensorMax =
2507 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValueMax},
2513 SmallVector<utils::IteratorType, 4> iteratorTypes;
2514 iteratorTypes.resize(inputTy.getRank(), utils::IteratorType::parallel);
2515 iteratorTypes[axis] = utils::IteratorType::reduction;
2517 SmallVector<AffineExpr, 2> srcExprs;
2518 SmallVector<AffineExpr, 2> dstExprs;
2519 for (
int i = 0, rank = inputTy.getRank(); i != rank; ++i) {
2525 bool didEncounterError =
false;
2527 rewriter.getContext());
2528 auto linalgOp = linalg::GenericOp::create(
2529 rewriter, loc, ArrayRef<Type>({resultTy, resultMaxTy}), input,
2530 ValueRange({filledTensorIdx, filledTensorMax}), maps, iteratorTypes,
2531 [&](OpBuilder &nestedBuilder, Location nestedLoc,
2533 auto newValue = blockArgs[0];
2534 auto oldIndex = blockArgs[1];
2535 auto oldValue = blockArgs[2];
2537 Value newIndex = arith::IndexCastOp::create(
2538 rewriter, nestedLoc, oldIndex.getType(),
2539 linalg::IndexOp::create(rewriter, loc, axis));
2542 if (isa<FloatType>(inElementTy)) {
2543 if (argmaxOp.getNanMode() == NanPropagationMode::IGNORE) {
2546 predicate = arith::CmpFOp::create(rewriter, nestedLoc,
2547 arith::CmpFPredicate::OGT,
2548 newValue, oldValue);
2553 Value gt = arith::CmpFOp::create(rewriter, nestedLoc,
2554 arith::CmpFPredicate::UGT,
2555 newValue, oldValue);
2556 Value oldNonNaN = arith::CmpFOp::create(rewriter, nestedLoc,
2557 arith::CmpFPredicate::ORD,
2558 oldValue, oldValue);
2559 predicate = arith::AndIOp::create(
2560 rewriter, nestedLoc, rewriter.getI1Type(), gt, oldNonNaN);
2562 }
else if (isa<IntegerType>(inElementTy)) {
2563 predicate = arith::CmpIOp::create(rewriter, nestedLoc,
2564 arith::CmpIPredicate::sgt,
2565 newValue, oldValue);
2567 didEncounterError =
true;
2571 auto resultMax = arith::SelectOp::create(
2572 rewriter, nestedLoc, predicate, newValue, oldValue);
2573 auto resultIndex = arith::SelectOp::create(
2574 rewriter, nestedLoc, predicate, newIndex, oldIndex);
2575 linalg::YieldOp::create(nestedBuilder, nestedLoc,
2579 if (didEncounterError)
2580 return rewriter.notifyMatchFailure(
2581 argmaxOp,
"unsupported tosa.argmax element type");
2583 rewriter.replaceOp(argmaxOp, linalgOp.getResult(0));
2588 bool allowNonFinites;
2591class GatherConverter :
public OpConversionPattern<tosa::GatherOp> {
2593 using OpConversionPattern<tosa::GatherOp>::OpConversionPattern;
2595 matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
2596 ConversionPatternRewriter &rewriter)
const final {
2597 auto input = adaptor.getOperands()[0];
2598 auto indices = adaptor.getOperands()[1];
2600 auto valuesTy = dyn_cast<RankedTensorType>(op.getValues().getType());
2601 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
2602 if (!valuesTy || !resultTy)
2603 return rewriter.notifyMatchFailure(op,
"unranked tensors not supported");
2605 auto dynamicDims = inferDynamicDimsForGather(
2606 rewriter, op.getLoc(), adaptor.getValues(), adaptor.getIndices());
2608 auto resultElementTy = resultTy.getElementType();
2610 auto loc = op.getLoc();
2612 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2613 resultElementTy, dynamicDims)
2616 SmallVector<AffineMap, 2> affineMaps = {
2618 resultTy.getRank(), 0,
2619 {rewriter.getAffineDimExpr(0), rewriter.getAffineDimExpr(1)},
2620 rewriter.getContext()),
2621 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2623 auto genericOp = linalg::GenericOp::create(
2627 [&](OpBuilder &
b, Location loc,
ValueRange args) {
2628 auto indexValue = args[0];
2629 auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
2630 Value index1 = arith::IndexCastOp::create(
2631 rewriter, loc, rewriter.getIndexType(), indexValue);
2632 auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
2633 Value extract = tensor::ExtractOp::create(
2634 rewriter, loc, input,
ValueRange{index0, index1, index2});
2635 linalg::YieldOp::create(rewriter, loc, extract);
2637 rewriter.replaceOp(op, genericOp.getResult(0));
2641 static llvm::SmallVector<Value> inferDynamicDimsForGather(OpBuilder &builder,
2645 llvm::SmallVector<Value> results;
2647 auto addDynamicDimension = [&](Value source, int64_t dim) {
2649 if (
auto dimValue = llvm::dyn_cast_if_present<Value>(sz))
2650 results.push_back(dimValue);
2653 addDynamicDimension(values, 0);
2654 addDynamicDimension(
indices, 1);
2655 addDynamicDimension(values, 2);
2665 using OpRewritePattern<tosa::TableOp>::OpRewritePattern;
2667 LogicalResult matchAndRewrite(tosa::TableOp op,
2668 PatternRewriter &rewriter)
const final {
2669 auto loc = op.getLoc();
2670 Value input = op.getInput1();
2671 Value table = op.getTable();
2672 auto inputTy = cast<ShapedType>(input.
getType());
2673 auto tableTy = cast<ShapedType>(table.
getType());
2674 auto resultTy = cast<ShapedType>(op.getType());
2676 auto inputElementTy = inputTy.getElementType();
2677 auto tableElementTy = tableTy.getElementType();
2678 auto resultElementTy = resultTy.getElementType();
2680 SmallVector<Value> dynDims;
2681 for (
int i = 0; i < resultTy.getRank(); ++i) {
2682 if (inputTy.isDynamicDim(i)) {
2684 tensor::DimOp::create(rewriter, loc, op.getOperand(0), i));
2689 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2690 resultElementTy, dynDims)
2693 SmallVector<AffineMap, 2> affineMaps = {
2694 rewriter.getMultiDimIdentityMap(resultTy.getRank()),
2695 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2697 auto genericOp = linalg::GenericOp::create(
2700 rewriter.replaceOp(op, genericOp.getResult(0));
2703 OpBuilder::InsertionGuard regionGuard(rewriter);
2704 Block *block = rewriter.createBlock(
2705 &genericOp.getRegion(), genericOp.getRegion().end(),
2706 TypeRange({inputElementTy, resultElementTy}), {loc, loc});
2709 rewriter.setInsertionPointToStart(block);
2710 if (inputElementTy.isInteger(8) && tableElementTy.isInteger(8) &&
2711 resultElementTy.isInteger(8)) {
2712 Value index = arith::IndexCastOp::create(
2713 rewriter, loc, rewriter.getIndexType(), inputValue);
2715 index = arith::AddIOp::create(rewriter, loc, rewriter.getIndexType(),
2718 tensor::ExtractOp::create(rewriter, loc, table,
ValueRange{index});
2719 linalg::YieldOp::create(rewriter, loc, extract);
2723 if (inputElementTy.isInteger(16) && tableElementTy.isInteger(16) &&
2724 resultElementTy.isInteger(32)) {
2725 Value extend = arith::ExtSIOp::create(
2726 rewriter, loc, rewriter.getI32Type(), inputValue);
2728 auto offset = arith::ConstantOp::create(
2729 rewriter, loc, rewriter.getI32IntegerAttr(32768));
2730 auto seven = arith::ConstantOp::create(rewriter, loc,
2731 rewriter.getI32IntegerAttr(7));
2732 auto one = arith::ConstantOp::create(rewriter, loc,
2733 rewriter.getI32IntegerAttr(1));
2734 auto b1111111 = arith::ConstantOp::create(
2735 rewriter, loc, rewriter.getI32IntegerAttr(127));
2741 auto extendAdd = arith::AddIOp::create(rewriter, loc, extend, offset);
2742 Value index = arith::ShRUIOp::create(rewriter, loc, extendAdd, seven);
2744 arith::AndIOp::create(rewriter, loc, extendAdd, b1111111);
2749 Value indexPlusOne = arith::AddIOp::create(rewriter, loc, index, one);
2751 index = arith::IndexCastOp::create(rewriter, loc,
2752 rewriter.getIndexType(), index);
2753 indexPlusOne = arith::IndexCastOp::create(
2754 rewriter, loc, rewriter.getIndexType(), indexPlusOne);
2757 tensor::ExtractOp::create(rewriter, loc, table,
ValueRange{index});
2758 Value
next = tensor::ExtractOp::create(rewriter, loc, table,
2762 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), base);
2764 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), next);
2768 Value baseScaled = arith::ShLIOp::create(rewriter, loc, base, seven);
2769 Value diff = arith::SubIOp::create(rewriter, loc, next, base);
2770 Value diffScaled = arith::MulIOp::create(rewriter, loc, diff, fraction);
2772 arith::AddIOp::create(rewriter, loc, baseScaled, diffScaled);
2774 linalg::YieldOp::create(rewriter, loc,
result);
2780 return rewriter.notifyMatchFailure(
2781 op,
"unable to create body for tosa.table op");
2786 using OpRewritePattern<RFFT2dOp>::OpRewritePattern;
2788 static bool isRankedTensor(Type type) {
return isa<RankedTensorType>(type); }
2790 static OpFoldResult halfPlusOne(OpBuilder &builder, Location loc,
2796 auto divBy2 = builder.
createOrFold<arith::DivUIOp>(loc, value, two);
2797 auto plusOne = builder.
createOrFold<arith::AddIOp>(loc, divBy2, one);
2801 static RankedTensorType
2802 computeOutputShape(OpBuilder &builder, Location loc, Value input,
2803 llvm::SmallVectorImpl<Value> &dynamicSizes) {
2809 dims[2] = halfPlusOne(builder, loc, dims[2]);
2811 llvm::SmallVector<int64_t, 3> staticSizes;
2814 auto elementType = cast<RankedTensorType>(input.
getType()).getElementType();
2815 return RankedTensorType::get(staticSizes, elementType);
2818 static Value createZeroTensor(PatternRewriter &rewriter, Location loc,
2819 RankedTensorType type,
2820 llvm::ArrayRef<Value> dynamicSizes) {
2822 tensor::EmptyOp::create(rewriter, loc, type, dynamicSizes);
2823 auto fillValueAttr = rewriter.
getZeroAttr(type.getElementType());
2824 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
2826 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValue},
2829 return filledTensor;
2832 static Value castIndexToFloat(OpBuilder &builder, Location loc,
2833 FloatType type, Value value) {
2834 auto integerVal = arith::IndexCastUIOp::create(
2836 type.getIntOrFloatBitWidth() > 32 ? builder.
getI64Type()
2840 return arith::UIToFPOp::create(builder, loc, type, integerVal);
2843 static Value createLinalgIndex(OpBuilder &builder, Location loc,
2844 FloatType type, int64_t index) {
2845 auto indexVal = linalg::IndexOp::create(builder, loc, index);
2846 return castIndexToFloat(builder, loc, type, indexVal);
2849 template <
typename... Args>
2850 static llvm::SmallVector<AffineExpr, 4> affineDimsExpr(OpBuilder &builder,
2855 LogicalResult matchAndRewrite(RFFT2dOp rfft2d,
2856 PatternRewriter &rewriter)
const override {
2857 if (!llvm::all_of(rfft2d->getOperandTypes(), isRankedTensor) ||
2858 !llvm::all_of(rfft2d->getResultTypes(), isRankedTensor)) {
2860 "only supports ranked tensors");
2863 auto loc = rfft2d.getLoc();
2864 auto input = rfft2d.getInputReal();
2866 dyn_cast<FloatType>(cast<ShapedType>(input.
getType()).getElementType());
2869 "only supports float element types");
2872 llvm::SmallVector<Value> dynamicSizes;
2873 auto outputType = computeOutputShape(rewriter, loc, input, dynamicSizes);
2876 llvm::SmallVector<utils::IteratorType, 5> iteratorTypes = {
2877 utils::IteratorType::parallel, utils::IteratorType::parallel,
2878 utils::IteratorType::parallel, utils::IteratorType::reduction,
2879 utils::IteratorType::reduction};
2882 llvm::SmallVector<Value> genericOpInputs = {input};
2883 llvm::SmallVector<Value> genericOpOutputs = {
2884 createZeroTensor(rewriter, loc, outputType, dynamicSizes),
2885 createZeroTensor(rewriter, loc, outputType, dynamicSizes)};
2889 llvm::ArrayRef{affineDimsExpr(rewriter, 0, 3, 4),
2890 affineDimsExpr(rewriter, 0, 1, 2),
2891 affineDimsExpr(rewriter, 0, 1, 2)},
2895 auto dimH = rewriter.
createOrFold<tensor::DimOp>(loc, input, 1);
2896 auto dimW = rewriter.
createOrFold<tensor::DimOp>(loc, input, 2);
2899 auto zeroFloat = arith::ConstantOp::create(
2900 rewriter, loc, rewriter.
getZeroAttr(elementType));
2901 auto twoPiAttr = rewriter.
getFloatAttr(elementType, 6.283185307179586);
2902 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
2907 auto constH = castIndexToFloat(rewriter, loc, elementType, dimH);
2908 auto constW = castIndexToFloat(rewriter, loc, elementType, dimW);
2909 auto halfH = index::DivUOp::create(rewriter, loc, dimH, twoIndex);
2910 auto halfW = index::DivUOp::create(rewriter, loc, dimW, twoIndex);
2912 auto buildBody = [&](OpBuilder &builder, Location loc,
ValueRange args) {
2913 Value valReal = args[0];
2914 Value sumReal = args[1];
2915 Value sumImag = args[2];
2918 Value oy = linalg::IndexOp::create(builder, loc, 1);
2919 Value ox = linalg::IndexOp::create(builder, loc, 2);
2920 Value iy = linalg::IndexOp::create(builder, loc, 3);
2921 Value ix = linalg::IndexOp::create(builder, loc, 4);
2926 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
2927 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
2929 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
2930 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
2932 auto iyRemFloat = castIndexToFloat(builder, loc, elementType, iyRem);
2933 auto ixRemFloat = castIndexToFloat(builder, loc, elementType, ixRem);
2935 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
2936 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
2937 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
2938 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
2945 auto iyIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2947 auto iyIsHalfH = arith::CmpIOp::create(
2948 builder, loc, arith::CmpIPredicate::eq, iyRem, halfH);
2949 auto ixIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2951 auto ixIsHalfW = arith::CmpIOp::create(
2952 builder, loc, arith::CmpIPredicate::eq, ixRem, halfW);
2954 auto iyIsSinSkippable =
2955 arith::OrIOp::create(builder, loc, iyIs0, iyIsHalfH);
2956 auto ixIsSinSkippable =
2957 arith::OrIOp::create(builder, loc, ixIs0, ixIsHalfW);
2958 auto shouldSkipSin = arith::AndIOp::create(builder, loc, iyIsSinSkippable,
2963 auto cosAngle = math::CosOp::create(builder, loc, angle);
2964 auto sinAngle = math::SinOp::create(builder, loc, angle);
2965 auto imagWeight = arith::SelectOp::create(builder, loc, shouldSkipSin,
2966 zeroFloat, sinAngle);
2967 auto realComponent =
2968 arith::MulFOp::create(builder, loc, valReal, cosAngle);
2969 auto imagComponent =
2970 arith::MulFOp::create(builder, loc, valReal, imagWeight);
2975 arith::AddFOp::create(builder, loc, sumReal, realComponent);
2977 arith::SubFOp::create(builder, loc, sumImag, imagComponent);
2979 linalg::YieldOp::create(builder, loc,
ValueRange{outReal, outImag});
2983 rfft2d, rfft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
2984 indexingMaps, iteratorTypes, buildBody);
2993 LogicalResult matchAndRewrite(FFT2dOp fft2d,
2994 PatternRewriter &rewriter)
const override {
2995 if (!llvm::all_of(fft2d->getOperandTypes(),
2996 RFFT2dConverter::isRankedTensor) ||
2997 !llvm::all_of(fft2d->getResultTypes(),
2998 RFFT2dConverter::isRankedTensor)) {
3002 Location loc = fft2d.getLoc();
3003 Value input_real = fft2d.getInputReal();
3004 Value input_imag = fft2d.getInputImag();
3005 BoolAttr inverse = fft2d.getInverseAttr();
3007 auto real_el_ty = cast<FloatType>(
3008 cast<ShapedType>(input_real.
getType()).getElementType());
3009 [[maybe_unused]]
auto imag_el_ty = cast<FloatType>(
3010 cast<ShapedType>(input_imag.
getType()).getElementType());
3012 assert(real_el_ty == imag_el_ty);
3015 SmallVector<Value> dynamicSizes;
3020 SmallVector<int64_t, 3> staticSizes;
3023 auto outputType = RankedTensorType::get(staticSizes, real_el_ty);
3026 SmallVector<utils::IteratorType, 5> iteratorTypes = {
3027 utils::IteratorType::parallel, utils::IteratorType::parallel,
3028 utils::IteratorType::parallel, utils::IteratorType::reduction,
3029 utils::IteratorType::reduction};
3032 SmallVector<Value> genericOpInputs = {input_real, input_imag};
3033 SmallVector<Value> genericOpOutputs = {
3034 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
3036 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
3041 ArrayRef{RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
3042 RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
3043 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2),
3044 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2)},
3048 auto dimH = rewriter.
createOrFold<tensor::DimOp>(loc, input_real, 1);
3049 auto dimW = rewriter.
createOrFold<tensor::DimOp>(loc, input_real, 2);
3052 auto twoPiAttr = rewriter.
getFloatAttr(real_el_ty, 6.283185307179586);
3053 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
3055 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimH);
3057 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimW);
3059 auto buildBody = [&](OpBuilder &builder, Location loc,
ValueRange args) {
3060 Value valReal = args[0];
3061 Value valImag = args[1];
3062 Value sumReal = args[2];
3063 Value sumImag = args[3];
3066 Value oy = linalg::IndexOp::create(builder, loc, 1);
3067 Value ox = linalg::IndexOp::create(builder, loc, 2);
3068 Value iy = linalg::IndexOp::create(builder, loc, 3);
3069 Value ix = linalg::IndexOp::create(builder, loc, 4);
3073 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
3074 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
3076 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
3077 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
3080 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, iyRem);
3082 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, ixRem);
3084 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
3085 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
3087 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
3088 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
3091 angle = arith::MulFOp::create(
3092 builder, loc, angle,
3093 arith::ConstantOp::create(rewriter, loc,
3099 auto cosAngle = math::CosOp::create(builder, loc, angle);
3100 auto sinAngle = math::SinOp::create(builder, loc, angle);
3102 auto rcos = arith::MulFOp::create(builder, loc, valReal, cosAngle);
3103 auto rsin = arith::MulFOp::create(builder, loc, valImag, sinAngle);
3104 auto realComponent = arith::AddFOp::create(builder, loc, rcos, rsin);
3106 auto icos = arith::MulFOp::create(builder, loc, valImag, cosAngle);
3107 auto isin = arith::MulFOp::create(builder, loc, valReal, sinAngle);
3109 auto imagComponent = arith::SubFOp::create(builder, loc, icos, isin);
3114 arith::AddFOp::create(builder, loc, sumReal, realComponent);
3116 arith::AddFOp::create(builder, loc, sumImag, imagComponent);
3118 linalg::YieldOp::create(builder, loc,
ValueRange{outReal, outImag});
3122 fft2d, fft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
3123 indexingMaps, iteratorTypes, buildBody);
3133 const TosaToLinalgOptions &
options) {
3136 patterns->
add<GenericResizeConverter>(patterns->
getContext(),
3140 patterns->
add<MaterializeResizeBroadcast>(patterns->
getContext(),
3145 PointwiseConverter<tosa::AddOp>,
3146 PointwiseConverter<tosa::SubOp>,
3147 PointwiseConverter<tosa::MulOp>,
3148 PointwiseConverter<tosa::IntDivOp>,
3149 PointwiseConverter<tosa::NegateOp>,
3150 PointwiseConverter<tosa::PowOp>,
3151 PointwiseConverter<tosa::ReciprocalOp>,
3152 PointwiseConverter<tosa::RsqrtOp>,
3153 PointwiseConverter<tosa::LogOp>,
3154 PointwiseConverter<tosa::ExpOp>,
3155 PointwiseConverter<tosa::AbsOp>,
3156 PointwiseConverter<tosa::SinOp>,
3157 PointwiseConverter<tosa::CosOp>,
3158 PointwiseConverter<tosa::TanhOp>,
3159 PointwiseConverter<tosa::ErfOp>,
3160 PointwiseConverter<tosa::BitwiseAndOp>,
3161 PointwiseConverter<tosa::BitwiseOrOp>,
3162 PointwiseConverter<tosa::BitwiseNotOp>,
3163 PointwiseConverter<tosa::BitwiseXorOp>,
3164 PointwiseConverter<tosa::LogicalAndOp>,
3165 PointwiseConverter<tosa::LogicalNotOp>,
3166 PointwiseConverter<tosa::LogicalOrOp>,
3167 PointwiseConverter<tosa::LogicalXorOp>,
3168 PointwiseConverter<tosa::CastOp>,
3169 PointwiseConverter<tosa::LogicalLeftShiftOp>,
3170 PointwiseConverter<tosa::LogicalRightShiftOp>,
3171 PointwiseConverter<tosa::ArithmeticRightShiftOp>,
3172 PointwiseConverter<tosa::ClzOp>,
3173 PointwiseConverter<tosa::SelectOp>,
3174 PointwiseConverter<tosa::GreaterOp>,
3175 PointwiseConverter<tosa::GreaterEqualOp>,
3176 PointwiseConverter<tosa::EqualOp>,
3177 PointwiseConverter<tosa::MaximumOp>,
3178 PointwiseConverter<tosa::MinimumOp>,
3179 PointwiseConverter<tosa::CeilOp>,
3180 PointwiseConverter<tosa::FloorOp>,
3181 PointwiseConverter<tosa::ClampOp>,
3182 PointwiseConverter<tosa::SigmoidOp>
3186 IdentityNConverter<tosa::IdentityOp>,
3198 ReduceConverter<tosa::ReduceAllOp>,
3199 ReduceConverter<tosa::ReduceAnyOp>,
3200 ReduceConverter<tosa::ReduceMinOp>,
3201 ReduceConverter<tosa::ReduceMaxOp>,
3202 ReduceConverter<tosa::ReduceSumOp>,
3203 ReduceConverter<tosa::ReduceProductOp>,
*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 inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
static llvm::ManagedStatic< PassManagerOptions > options
static Value clamp(ImplicitLocOpBuilder &builder, Value value, Value lowerBound, Value upperBound)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
static TypedAttr createInitialValueForReduceOp(Operation *op, Type elementTy, PatternRewriter &rewriter, bool allowNonFinites)
static OpFoldResult getOrFoldTensorDim(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value tensor, int64_t index)
static LogicalResult emitElementwiseComputation(ConversionPatternRewriter &rewriter, Location loc, Operation *operation, ValueRange operands, ArrayRef< OpFoldResult > targetShape, const TypeConverter &converter)
static Value createLinalgBodyCalculationForReduceOp(Operation *op, ValueRange args, Type elementTy, PatternRewriter &rewriter)
static OpTy createWithDefaultProperties(OpBuilder &builder, Location loc, TypeRange resultTypes, ValueRange operands)
static Value getTensorDim(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value tensor, int64_t index)
static Value createIndex(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, int64_t index)
static std::pair< OpFoldResult, Value > computeTargetSize(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, ValueRange operands, int64_t dim)
DenseMap< int64_t, Value > IndexPool
static LogicalResult reduceMatchAndRewriteHelper(OpTy op, uint64_t axis, PatternRewriter &rewriter, bool allowNonFinites)
static Value broadcastDynamicDimensions(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value operand, ArrayRef< OpFoldResult > targetShape, ArrayRef< Value > masterOperands)
static Value broadcastDynamicDimension(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, Value operand, int64_t dim, OpFoldResult targetSize, Value masterOperand)
static LogicalResult elementwiseMatchAndRewriteHelper(Operation *operation, ValueRange operands, ConversionPatternRewriter &rewriter, const TypeConverter &converter)
static APFloat getFloatMinMaxIdentity(const llvm::fltSemantics &semantics, bool negative, bool allowNonFinites)
static Value createLinalgBodyCalculationForElementwiseOp(Operation *op, ValueRange args, ArrayRef< Type > resultTypes, ConversionPatternRewriter &rewriter)
static ValueRange getBroadcastableOperands(Operation *operation, ValueRange operands)
static Value materializeBinaryNanCheckIfRequired(OpTy op, PatternRewriter &rewriter, Value lhs, Value rhs, Value result)
static std::pair< SmallVector< OpFoldResult >, SmallVector< Value > > computeTargetShape(PatternRewriter &rewriter, Location loc, IndexPool &indexPool, ValueRange operands)
static bool operandsAndResultsRanked(Operation *operation)
static const llvm::fltSemantics * getFloatSemantics(TruncfSrcElemTypes etype)
Float semantics the element type attributes of xevm.truncf and xevm.extf stand for.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
BlockArgument getArgument(unsigned i)
bool getValue() const
Return the boolean value of this attribute.
IntegerAttr getIndexAttr(int64_t value)
IntegerAttr getI32IntegerAttr(int32_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
AffineMap getMultiDimIdentityMap(unsigned rank)
FloatAttr getFloatAttr(Type type, double value)
AffineExpr getAffineConstantExpr(int64_t constant)
IntegerType getIntegerType(unsigned width)
Ty getType(Args &&...args)
Get or construct an instance of the type Ty with provided arguments.
BoolAttr getBoolAttr(bool value)
TypedAttr getZeroAttr(Type type)
AffineExpr getAffineDimExpr(unsigned position)
MLIRContext * getContext() const
IntegerAttr getI8IntegerAttr(int8_t value)
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.
An attribute that represents a reference to a dense integer vector or tensor object.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
Location getLoc()
The source location the operation was defined or derived from.
unsigned getNumOperands()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
unsigned getNumResults()
Return the number of results held by this operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
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,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class provides an abstraction over the different types of ranges over Values.
type_range getType() const
Type front()
Return first type in the range.
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.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
InFlightDiagnostic & next(InFlightDiagnostic &diag)
Starts a new message part in an in-flight diagnostic.
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given tensor value.
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Value clampFloatHelper(Location loc, Value arg, Value min, Value max, OpBuilder &rewriter)
SmallVector< utils::IteratorType > getNParallelLoopsAttrs(unsigned nParallelLoops)
std::optional< SmallVector< Value > > checkHasDynamicBatchDims(PatternRewriter &rewriter, Op op, ArrayRef< Value > params)
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter, RewritePatternSet *patterns, const TosaToLinalgOptions &options=TosaToLinalgOptions())
Populates conversion passes from TOSA dialect to Linalg dialect.
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
SmallVector< int64_t > convertFromMlirShape(ArrayRef< int64_t > shape)
Value clampIntHelper(Location loc, Value arg, Value min, Value max, OpBuilder &rewriter, bool isUnsigned)
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.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
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...