29#include "llvm/ADT/STLExtras.h"
30#include "llvm/ADT/Sequence.h"
31#include "llvm/ADT/SmallVectorExtras.h"
60template <
typename OpTy>
68 auto nanMode = op.getNanMode();
69 if (nanMode == NanPropagationMode::PROPAGATE)
73 Value lhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
74 arith::CmpFPredicate::UNO,
lhs,
lhs);
75 Value rhsIsNaN = arith::CmpFOp::create(rewriter, op.getLoc(),
76 arith::CmpFPredicate::UNO,
rhs,
rhs);
78 arith::SelectOp::create(rewriter, op.getLoc(), lhsIsNaN,
rhs,
result);
79 return arith::SelectOp::create(rewriter, op.getLoc(), rhsIsNaN,
lhs,
85 ConversionPatternRewriter &rewriter) {
91 if (isa<tosa::AbsOp>(op) && isa<FloatType>(elementTy))
92 return math::AbsFOp::create(rewriter, loc, resultTypes, args);
94 if (isa<tosa::AbsOp>(op) && isa<IntegerType>(elementTy)) {
95 auto zero = arith::ConstantOp::create(rewriter, loc,
96 rewriter.getZeroAttr(elementTy));
97 auto neg = arith::SubIOp::create(rewriter, loc, zero, args[0]);
98 return arith::MaxSIOp::create(rewriter, loc, args[0], neg);
102 if (isa<tosa::AddOp>(op) && isa<FloatType>(elementTy))
103 return arith::AddFOp::create(rewriter, loc, resultTypes, args);
105 if (isa<tosa::AddOp>(op) && isa<IntegerType>(elementTy))
106 return arith::AddIOp::create(rewriter, loc, resultTypes, args);
109 if (isa<tosa::SubOp>(op) && isa<FloatType>(elementTy))
110 return arith::SubFOp::create(rewriter, loc, resultTypes, args);
112 if (isa<tosa::SubOp>(op) && isa<IntegerType>(elementTy))
113 return arith::SubIOp::create(rewriter, loc, resultTypes, args);
116 if (isa<tosa::IntDivOp>(op) && isa<IntegerType>(elementTy))
117 return arith::DivSIOp::create(rewriter, loc, resultTypes, args);
120 if (isa<tosa::ReciprocalOp>(op) && isa<FloatType>(elementTy)) {
122 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
123 return arith::DivFOp::create(rewriter, loc, one, args[0]);
127 if (isa<tosa::MulOp>(op)) {
128 auto shiftVal = cast<tosa::MulOp>(op).getShift();
130 bool shiftIsConstant =
true;
133 shift = shiftElem.
getValues<IntegerAttr>()[0].getInt();
135 shiftIsConstant =
false;
137 if (isa<FloatType>(elementTy)) {
139 (
void)rewriter.notifyMatchFailure(op,
140 "Cannot have shift value for float");
143 return arith::MulFOp::create(rewriter, loc, args[0], args[1]);
146 if (isa<IntegerType>(elementTy)) {
150 if (shift > 0 || !shiftIsConstant) {
157 a = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), a);
159 if (!
b.getType().isInteger(32))
160 b = arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(),
b);
162 auto shiftAmount = shiftIsConstant ? shiftConst : args[2];
163 auto roundingAttr = RoundingModeAttr::get(rewriter.getContext(),
164 RoundingMode::SINGLE_ROUND);
166 tosa::ApplyScaleOp::create(rewriter, loc, rewriter.getI32Type(), a,
167 b, shiftAmount, roundingAttr);
173 int bWidth =
b.getType().getIntOrFloatBitWidth();
174 int cWidth = resultTypes[0].getIntOrFloatBitWidth();
177 a = arith::ExtSIOp::create(rewriter, loc, resultTypes[0], a);
179 b = arith::ExtSIOp::create(rewriter, loc, resultTypes[0],
b);
181 return arith::MulIOp::create(rewriter, loc, resultTypes, a,
b);
186 if (isa<tosa::NegateOp>(op)) {
187 auto negate = cast<tosa::NegateOp>(op);
190 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
191 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
192 bool hasInZp = !failed(maybeInZp);
193 bool hasOutZp = !failed(maybeOutZp);
199 if (isa<FloatType>(elementTy))
200 return arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
202 if (isa<IntegerType>(elementTy)) {
204 Type intermediateType;
207 int intermediateBitWidth = 64;
209 if (hasInZp && hasOutZp) {
211 const int64_t zpAdd = inZp + outZp;
213 APInt::getSignedMaxValue(inputBitWidth).getSExtValue() +
218 if (maxValue <= APInt::getSignedMaxValue(16).getSExtValue()) {
219 intermediateBitWidth = 16;
220 }
else if (maxValue <= APInt::getSignedMaxValue(32).getSExtValue()) {
221 intermediateBitWidth = 32;
224 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
225 zpAddValue = arith::ConstantOp::create(
226 rewriter, loc, rewriter.getIntegerAttr(intermediateType, zpAdd));
228 intermediateType = rewriter.getIntegerType(intermediateBitWidth);
229 Value arg1 = args[1];
230 Value arg2 = args[2];
232 if (arg1.
getType() != intermediateType)
233 arg1 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg1);
234 if (arg2.
getType() != intermediateType)
235 arg2 = arith::ExtSIOp::create(rewriter, loc, intermediateType, arg2);
237 arith::AddIOp::create(rewriter, loc, intermediateType, arg1, arg2);
243 if (ext.
getType() != intermediateType)
244 ext = arith::ExtSIOp::create(rewriter, loc, intermediateType, ext);
245 auto sub = arith::SubIOp::create(rewriter, loc, zpAddValue, ext);
249 rewriter, loc, intermediateType,
250 APInt::getSignedMinValue(inputBitWidth).getSExtValue());
252 rewriter, loc, intermediateType,
253 APInt::getSignedMaxValue(inputBitWidth).getSExtValue());
257 if (
clamp.getType() == elementTy)
259 return arith::TruncIOp::create(rewriter, loc, elementTy,
clamp);
264 if (isa<tosa::BitwiseAndOp>(op) && isa<IntegerType>(elementTy))
265 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
268 if (isa<tosa::BitwiseOrOp>(op) && isa<IntegerType>(elementTy))
269 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
272 if (isa<tosa::BitwiseNotOp>(op) && isa<IntegerType>(elementTy)) {
273 auto allOnesAttr = rewriter.getIntegerAttr(
274 elementTy, APInt::getAllOnes(elementTy.getIntOrFloatBitWidth()));
275 auto allOnes = arith::ConstantOp::create(rewriter, loc, allOnesAttr);
276 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], allOnes);
280 if (isa<tosa::BitwiseXorOp>(op) && isa<IntegerType>(elementTy))
281 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
284 if (isa<tosa::LogicalLeftShiftOp>(op) && isa<IntegerType>(elementTy))
285 return arith::ShLIOp::create(rewriter, loc, resultTypes, args);
288 if (isa<tosa::LogicalRightShiftOp>(op) && isa<IntegerType>(elementTy))
289 return arith::ShRUIOp::create(rewriter, loc, resultTypes, args);
292 if (isa<tosa::ArithmeticRightShiftOp>(op) && isa<IntegerType>(elementTy)) {
293 auto result = arith::ShRSIOp::create(rewriter, loc, resultTypes, args);
294 bool round = cast<tosa::ArithmeticRightShiftOp>(op).getRound();
299 Type i1Ty = IntegerType::get(rewriter.getContext(), 1);
300 auto one = arith::ConstantOp::create(rewriter, loc,
301 IntegerAttr::get(elementTy, 1));
302 auto zero = arith::ConstantOp::create(rewriter, loc,
303 IntegerAttr::get(elementTy, 0));
305 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 0));
307 arith::ConstantOp::create(rewriter, loc, IntegerAttr::get(i1Ty, 1));
310 auto shiftValueGreaterThanZero = arith::CmpIOp::create(
311 rewriter, loc, arith::CmpIPredicate::sgt, args[1], zero);
315 arith::SubIOp::create(rewriter, loc, resultTypes, args[1], one);
317 arith::ShRSIOp::create(rewriter, loc, resultTypes, args[0], subtract)
319 auto truncated = arith::TruncIOp::create(rewriter, loc, i1Ty, shifted,
322 arith::AndIOp::create(rewriter, loc, i1Ty, truncated, i1one);
324 auto shouldRound = arith::SelectOp::create(
325 rewriter, loc, i1Ty, shiftValueGreaterThanZero, isInputOdd, i1zero);
327 arith::ExtUIOp::create(rewriter, loc, resultTypes, shouldRound);
328 return arith::AddIOp::create(rewriter, loc, resultTypes,
result, extended);
332 if (isa<tosa::ClzOp>(op) && isa<IntegerType>(elementTy)) {
333 return math::CountLeadingZerosOp::create(rewriter, loc, elementTy, args[0]);
337 if (isa<tosa::LogicalAndOp>(op) && elementTy.isInteger(1))
338 return arith::AndIOp::create(rewriter, loc, resultTypes, args);
341 if (isa<tosa::LogicalNotOp>(op) && elementTy.isInteger(1)) {
342 auto one = arith::ConstantOp::create(rewriter, loc,
343 rewriter.getIntegerAttr(elementTy, 1));
344 return arith::XOrIOp::create(rewriter, loc, resultTypes, args[0], one);
348 if (isa<tosa::LogicalOrOp>(op) && elementTy.isInteger(1))
349 return arith::OrIOp::create(rewriter, loc, resultTypes, args);
352 if (isa<tosa::LogicalXorOp>(op) && elementTy.isInteger(1))
353 return arith::XOrIOp::create(rewriter, loc, resultTypes, args);
356 if (isa<tosa::PowOp>(op) && isa<FloatType>(elementTy))
357 return mlir::math::PowFOp::create(rewriter, loc, resultTypes, args);
360 if (isa<tosa::RsqrtOp>(op) && isa<FloatType>(elementTy))
361 return mlir::math::RsqrtOp::create(rewriter, loc, resultTypes, args);
364 if (isa<tosa::LogOp>(op) && isa<FloatType>(elementTy))
365 return mlir::math::LogOp::create(rewriter, loc, resultTypes, args);
368 if (isa<tosa::ExpOp>(op) && isa<FloatType>(elementTy))
369 return mlir::math::ExpOp::create(rewriter, loc, resultTypes, args);
372 if (isa<tosa::SinOp>(op) && isa<FloatType>(elementTy))
373 return mlir::math::SinOp::create(rewriter, loc, resultTypes, args);
376 if (isa<tosa::CosOp>(op) && isa<FloatType>(elementTy))
377 return mlir::math::CosOp::create(rewriter, loc, resultTypes, args);
380 if (isa<tosa::TanhOp>(op) && isa<FloatType>(elementTy))
381 return mlir::math::TanhOp::create(rewriter, loc, resultTypes, args);
384 if (isa<tosa::ErfOp>(op) && llvm::isa<FloatType>(elementTy))
385 return mlir::math::ErfOp::create(rewriter, loc, resultTypes, args);
388 if (isa<tosa::GreaterOp>(op) && isa<FloatType>(elementTy))
389 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGT,
392 if (isa<tosa::GreaterOp>(op) && elementTy.isSignlessInteger())
393 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sgt,
397 if (isa<tosa::GreaterEqualOp>(op) && isa<FloatType>(elementTy))
398 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OGE,
401 if (isa<tosa::GreaterEqualOp>(op) && elementTy.isSignlessInteger())
402 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::sge,
406 if (isa<tosa::EqualOp>(op) && isa<FloatType>(elementTy))
407 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::OEQ,
410 if (isa<tosa::EqualOp>(op) && elementTy.isSignlessInteger())
411 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::eq,
415 if (isa<tosa::SelectOp>(op)) {
417 if (isa<FloatType>(elementTy) || isa<IntegerType>(elementTy))
418 return arith::SelectOp::create(rewriter, loc, args[0], args[1], args[2]);
422 if (isa<tosa::MaximumOp>(op) && isa<FloatType>(elementTy)) {
423 auto max = arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
425 rewriter, args[0], args[1],
max);
428 if (isa<tosa::MaximumOp>(op) && elementTy.isSignlessInteger()) {
429 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
433 if (isa<tosa::MinimumOp>(op) && isa<FloatType>(elementTy)) {
434 auto min = arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
436 rewriter, args[0], args[1],
min);
439 if (isa<tosa::MinimumOp>(op) && elementTy.isSignlessInteger()) {
440 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
444 if (isa<tosa::CeilOp>(op) && isa<FloatType>(elementTy))
445 return math::CeilOp::create(rewriter, loc, resultTypes, args);
448 if (isa<tosa::FloorOp>(op) && isa<FloatType>(elementTy))
449 return math::FloorOp::create(rewriter, loc, resultTypes, args);
452 if (isa<tosa::ClampOp>(op) && isa<FloatType>(elementTy)) {
453 bool losesInfo =
false;
454 auto clampOp = cast<tosa::ClampOp>(op);
455 APFloat minApf = cast<FloatAttr>(clampOp.getMinValAttr()).getValue();
456 APFloat maxApf = cast<FloatAttr>(clampOp.getMaxValAttr()).getValue();
457 minApf.convert(cast<FloatType>(elementTy).getFloatSemantics(),
458 APFloat::rmNearestTiesToEven, &losesInfo);
459 maxApf.convert(cast<FloatType>(elementTy).getFloatSemantics(),
460 APFloat::rmNearestTiesToEven, &losesInfo);
461 auto min = arith::ConstantOp::create(
462 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, minApf));
463 auto max = arith::ConstantOp::create(
464 rewriter, loc, elementTy, rewriter.getFloatAttr(elementTy, maxApf));
467 const auto nanMode = clampOp.getNanMode();
470 if (!isa<FloatType>(elementTy))
475 if (nanMode == NanPropagationMode::PROPAGATE)
489 Value isNaN = arith::CmpFOp::create(
490 rewriter, op->
getLoc(), arith::CmpFPredicate::UNO, args[0], args[0]);
493 return arith::SelectOp::create(rewriter, op->
getLoc(), isNaN,
min,
result);
496 if (isa<tosa::ClampOp>(op) && isa<IntegerType>(elementTy)) {
497 auto intTy = cast<IntegerType>(elementTy);
498 auto clampOp = cast<tosa::ClampOp>(op);
500 cast<IntegerAttr>(clampOp.getMinValAttr()).getValue().getSExtValue();
502 cast<IntegerAttr>(clampOp.getMaxValAttr()).getValue().getSExtValue();
504 int64_t minRepresentable = std::numeric_limits<int64_t>::min();
505 int64_t maxRepresentable = std::numeric_limits<int64_t>::max();
506 if (intTy.isUnsignedInteger()) {
507 minRepresentable = 0;
508 if (intTy.getIntOrFloatBitWidth() <= 63) {
510 (
int64_t)APInt::getMaxValue(intTy.getIntOrFloatBitWidth())
513 }
else if (intTy.getIntOrFloatBitWidth() <= 64) {
515 minRepresentable = APInt::getSignedMinValue(intTy.getIntOrFloatBitWidth())
517 maxRepresentable = APInt::getSignedMaxValue(intTy.getIntOrFloatBitWidth())
522 min = std::max(
min, minRepresentable);
523 max = std::max(
max, minRepresentable);
524 min = std::min(
min, maxRepresentable);
525 max = std::min(
max, maxRepresentable);
528 intTy.getIntOrFloatBitWidth());
530 intTy.getIntOrFloatBitWidth());
532 intTy.isUnsignedInteger());
536 if (isa<tosa::SigmoidOp>(op) && isa<FloatType>(elementTy)) {
538 arith::ConstantOp::create(rewriter, loc, FloatAttr::get(elementTy, 1));
539 auto negate = arith::NegFOp::create(rewriter, loc, resultTypes, args[0]);
540 auto exp = mlir::math::ExpOp::create(rewriter, loc, resultTypes, negate);
541 auto added = arith::AddFOp::create(rewriter, loc, exp, one);
542 return arith::DivFOp::create(rewriter, loc, one, added);
546 if (isa<tosa::CastOp>(op)) {
547 Type srcTy = elementTy;
548 Type dstTy = resultTypes.front();
550 (
void)rewriter.notifyMatchFailure(op,
"unsupported type");
560 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && bitExtend)
561 return arith::ExtFOp::create(rewriter, loc, resultTypes, args,
564 if (isa<FloatType>(srcTy) && isa<FloatType>(dstTy) && !bitExtend)
565 return arith::TruncFOp::create(rewriter, loc, resultTypes, args,
569 if (srcTy.
isInteger(1) && arith::UIToFPOp::areCastCompatible(srcTy, dstTy))
570 return arith::UIToFPOp::create(rewriter, loc, resultTypes, args,
573 if (srcTy.
isInteger(1) && isa<IntegerType>(dstTy) && bitExtend)
574 return arith::ExtUIOp::create(rewriter, loc, resultTypes, args,
580 auto unrealizedCast =
581 UnrealizedConversionCastOp::create(
585 return arith::UIToFPOp::create(rewriter, loc, resultTypes[0],
590 if (arith::SIToFPOp::areCastCompatible(srcTy, dstTy))
591 return arith::SIToFPOp::create(rewriter, loc, resultTypes, args,
595 if (isa<FloatType>(srcTy) && dstTy.
isInteger(1)) {
596 Value zero = arith::ConstantOp::create(rewriter, loc,
597 rewriter.getFloatAttr(srcTy, 0.0));
598 return arith::CmpFOp::create(rewriter, loc, arith::CmpFPredicate::UNE,
602 if (arith::FPToSIOp::areCastCompatible(srcTy, dstTy)) {
603 auto rounded = math::RoundEvenOp::create(rewriter, loc, args[0]);
605 const auto &fltSemantics = cast<FloatType>(srcTy).getFloatSemantics();
609 APFloat::semanticsMaxExponent(fltSemantics)) {
612 auto conv = arith::FPToSIOp::create(rewriter, loc, dstTy, rounded);
613 auto posInf = arith::ConstantOp::create(
616 APFloat::getInf(fltSemantics)));
617 auto negInf = arith::ConstantOp::create(
619 rewriter.getFloatAttr(
621 APFloat::getInf(fltSemantics,
true)));
622 auto overflow = arith::CmpFOp::create(
623 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, posInf);
624 auto underflow = arith::CmpFOp::create(
625 rewriter, loc, arith::CmpFPredicate::UEQ, rounded, negInf);
626 auto intMin = arith::ConstantOp::create(
628 rewriter.getIntegerAttr(
631 auto intMax = arith::ConstantOp::create(
633 rewriter.getIntegerAttr(
637 arith::SelectOp::create(rewriter, loc, overflow, intMax, conv);
638 return arith::SelectOp::create(rewriter, loc, underflow, intMin,
642 auto intMinFP = arith::ConstantOp::create(
644 rewriter.getFloatAttr(
650 if (cast<FloatType>(srcTy).getFPMantissaWidth() >=
656 auto intMaxFP = arith::ConstantOp::create(
658 rewriter.getFloatAttr(
665 return arith::FPToSIOp::create(rewriter, loc, dstTy, clamped);
672 auto intMaxPlusOneFP = arith::ConstantOp::create(
674 rewriter.getFloatAttr(
681 auto intMax = arith::ConstantOp::create(
683 rewriter.getIntegerAttr(
687 arith::MaximumFOp::create(rewriter, loc, rounded, intMinFP);
689 arith::FPToSIOp::create(rewriter, loc, dstTy, minClampedFP);
690 auto overflow = arith::CmpFOp::create(
691 rewriter, loc, arith::CmpFPredicate::UGE, rounded, intMaxPlusOneFP);
692 return arith::SelectOp::create(rewriter, loc, overflow, intMax,
698 if (isa<IntegerType>(srcTy) && dstTy.
isInteger(1)) {
701 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ne,
705 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && bitExtend)
706 return arith::ExtSIOp::create(rewriter, loc, resultTypes, args,
709 if (isa<IntegerType>(srcTy) && isa<IntegerType>(dstTy) && !bitExtend) {
710 return arith::TruncIOp::create(rewriter, loc, dstTy, args[0]);
714 (
void)rewriter.notifyMatchFailure(
715 op,
"unhandled op for linalg body calculation for elementwise op");
736 return tensor::DimOp::create(rewriter, loc,
tensor, indexValue).getResult();
742 auto shapedType = dyn_cast<ShapedType>(
tensor.getType());
743 assert(shapedType && shapedType.hasRank() &&
"expected a ranked shaped type");
744 assert(
index >= 0 &&
index < shapedType.getRank() &&
"index out of bounds");
745 if (shapedType.isDynamicDim(
index))
751 auto isRanked = [](
Value value) {
752 return isa<RankedTensorType>(value.getType());
754 return llvm::all_of(operation->
getOperands(), isRanked) &&
755 llvm::all_of(operation->
getResults(), isRanked);
768static std::pair<OpFoldResult, Value>
774 for (
auto operand : operands) {
775 auto size = cast<RankedTensorType>(operand.getType()).getDimSize(dim);
776 if (ShapedType::isStatic(size) && size > 1)
781 auto operandsWithDynamicDim =
782 llvm::filter_to_vector(operands, [&](
Value operand) {
783 return cast<RankedTensorType>(operand.
getType()).isDynamicDim(dim);
787 if (operandsWithDynamicDim.empty())
794 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[0], dim);
795 if (operandsWithDynamicDim.size() == 1)
796 return {targetSize, operandsWithDynamicDim[0]};
799 for (
size_t i = 1; i < operandsWithDynamicDim.size(); i++) {
801 getTensorDim(rewriter, loc, indexPool, operandsWithDynamicDim[i], dim);
802 targetSize = arith::MaxUIOp::create(rewriter, loc, targetSize, nextSize);
804 return {targetSize,
nullptr};
812 assert(!operands.empty());
813 auto rank = cast<RankedTensorType>(operands.front().
getType()).getRank();
816 for (
auto dim : llvm::seq<int64_t>(0, rank)) {
817 auto [targetSize, masterOperand] =
819 targetShape.push_back(targetSize);
820 masterOperands.push_back(masterOperand);
822 return {targetShape, masterOperands};
828 Value masterOperand) {
830 auto rankedTensorType = cast<RankedTensorType>(operand.
getType());
831 if (!rankedTensorType.isDynamicDim(dim))
838 if (operand == masterOperand)
842 auto rank = rankedTensorType.getRank();
844 for (
auto index : llvm::seq<int64_t>(0, rank)) {
847 affineExprs.push_back(affineExpr);
849 auto broadcastAffineMap =
855 auto one =
createIndex(rewriter, loc, indexPool, 1);
856 auto runtimeSize =
getTensorDim(rewriter, loc, indexPool, operand, dim);
857 auto broadcastNecessary = arith::CmpIOp::create(
858 rewriter, loc, arith::CmpIPredicate::eq, runtimeSize, one);
868 for (
auto index : llvm::seq<int64_t>(0, rank)) {
869 auto size =
index == dim ? targetSize
872 outputTensorShape.push_back(size);
874 Value outputTensor = tensor::EmptyOp::create(
875 opBuilder, loc, outputTensorShape, rankedTensorType.getElementType());
879 linalg::GenericOp::create(
880 opBuilder, loc, outputTensor.
getType(), operand, outputTensor,
884 linalg::YieldOp::create(opBuilder, loc, blockArgs.front());
889 auto castResultTensor = rewriter.
createOrFold<tensor::CastOp>(
890 loc, operand.
getType(), resultTensor);
893 scf::YieldOp::create(opBuilder, loc, castResultTensor);
898 scf::YieldOp::create(opBuilder, loc, operand);
902 auto ifOp = scf::IfOp::create(rewriter, loc, broadcastNecessary,
903 emitThenRegion, emitElseRegion);
904 return ifOp.getResult(0);
911 int64_t rank = cast<RankedTensorType>(operand.
getType()).getRank();
912 assert((
int64_t)targetShape.size() == rank);
913 assert((
int64_t)masterOperands.size() == rank);
914 for (
auto index : llvm::seq<int64_t>(0, rank))
927 if (operands.size() == 1)
931 bool hasDynamic =
false;
932 for (
auto op : operands) {
933 const auto tType = dyn_cast<RankedTensorType>(op.getType());
934 if (tType && !tType.hasStaticShape()) {
943 return llvm::map_to_vector(operands, [&](
Value operand) {
945 targetShape, masterOperands);
955 auto resultType = cast_or_null<RankedTensorType>(
958 return rewriter.notifyMatchFailure(operation,
"failed to convert type");
960 Value outputTensor = tensor::EmptyOp::create(rewriter, loc, targetShape,
961 resultType.getElementType());
966 auto rank = resultType.getRank();
967 auto affineMaps = llvm::map_to_vector(operands, [&](
Value operand) {
968 auto shape = cast<ShapedType>(operand.
getType()).getShape();
970 for (
auto it : llvm::enumerate(
shape)) {
974 bool requiresBroadcast =
975 (it.value() == 1 && resultType.getDimSize(it.index()) != 1);
976 auto affineExpr = requiresBroadcast
977 ? rewriter.getAffineConstantExpr(0)
978 : rewriter.getAffineDimExpr(it.index());
979 affineExprs.push_back(affineExpr);
981 return AffineMap::get(rank, 0, affineExprs, rewriter.getContext());
983 affineMaps.push_back(rewriter.getMultiDimIdentityMap(rank));
986 bool encounteredError =
false;
987 auto linalgOp = linalg::GenericOp::create(
988 rewriter, loc, outputTensor.
getType(), operands, outputTensor, affineMaps,
993 {resultType.getElementType()}, rewriter);
995 encounteredError =
true;
998 linalg::YieldOp::create(opBuilder, loc, opResult);
1000 if (encounteredError)
1001 return rewriter.notifyMatchFailure(
1002 operation,
"unable to create linalg.generic body for elementwise op");
1005 auto castResult = rewriter.createOrFold<tensor::CastOp>(
1006 loc, resultType, linalgOp->getResult(0));
1007 rewriter.replaceOp(operation, castResult);
1014 if (isa<tosa::MulOp>(operation)) {
1018 return operands.take_front(2);
1020 return operands.take_front(3);
1022 if (
auto negate = dyn_cast<tosa::NegateOp>(operation)) {
1023 FailureOr<int64_t> maybeInZp = negate.getInput1ZeroPoint();
1024 FailureOr<int64_t> maybeOutZp = negate.getOutputZeroPoint();
1025 if (failed(maybeOutZp) && failed(maybeInZp))
1028 return operands.take_front(1);
1035 ConversionPatternRewriter &rewriter,
1039 assert(operation->
getNumResults() == 1 &&
"elementwise op expects 1 result");
1041 "elementwise op expects at least 1 operand");
1043 return rewriter.notifyMatchFailure(operation,
1044 "Unranked tensors not supported");
1048 auto loc = operation->
getLoc();
1050 auto [targetShape, masterOperands] =
1052 auto broadcastOperands =
1054 targetShape, masterOperands);
1056 targetShape, converter);
1063 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy))
1066 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy))
1069 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy))
1072 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy))
1075 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy))
1077 elementTy, APFloat::getLargest(
1078 cast<FloatType>(elementTy).getFloatSemantics(),
false));
1080 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy))
1084 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy))
1086 elementTy, APFloat::getLargest(
1087 cast<FloatType>(elementTy).getFloatSemantics(),
true));
1089 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy))
1093 if (isa<tosa::ReduceAllOp>(op) && elementTy.
isInteger(1))
1096 if (isa<tosa::ReduceAnyOp>(op) && elementTy.
isInteger(1))
1099 if (isa<tosa::ArgMaxOp>(op) && isa<FloatType>(elementTy))
1101 elementTy, APFloat::getLargest(
1102 cast<FloatType>(elementTy).getFloatSemantics(),
true));
1104 if (isa<tosa::ArgMaxOp>(op) && isa<IntegerType>(elementTy))
1118 if (isa<tosa::ReduceSumOp>(op) && isa<FloatType>(elementTy)) {
1119 return arith::AddFOp::create(rewriter, loc, args);
1122 if (isa<tosa::ReduceSumOp>(op) && isa<IntegerType>(elementTy)) {
1123 return arith::AddIOp::create(rewriter, loc, args);
1126 if (isa<tosa::ReduceProductOp>(op) && isa<FloatType>(elementTy)) {
1127 return arith::MulFOp::create(rewriter, loc, args);
1130 if (isa<tosa::ReduceProductOp>(op) && isa<IntegerType>(elementTy)) {
1131 return arith::MulIOp::create(rewriter, loc, args);
1134 if (isa<tosa::ReduceMinOp>(op) && isa<FloatType>(elementTy)) {
1135 return arith::MinimumFOp::create(rewriter, loc, args[0], args[1]);
1138 if (isa<tosa::ReduceMinOp>(op) && isa<IntegerType>(elementTy)) {
1139 return arith::MinSIOp::create(rewriter, loc, args[0], args[1]);
1142 if (isa<tosa::ReduceMaxOp>(op) && isa<FloatType>(elementTy)) {
1143 return arith::MaximumFOp::create(rewriter, loc, args[0], args[1]);
1146 if (isa<tosa::ReduceMaxOp>(op) && isa<IntegerType>(elementTy)) {
1147 return arith::MaxSIOp::create(rewriter, loc, args[0], args[1]);
1150 if (isa<tosa::ReduceAllOp>(op) && elementTy.
isInteger(1))
1151 return arith::AndIOp::create(rewriter, loc, args);
1153 if (isa<tosa::ReduceAnyOp>(op) && elementTy.
isInteger(1))
1154 return arith::OrIOp::create(rewriter, loc, args);
1162template <
typename OpTy>
1165 auto loc = op->getLoc();
1166 auto inputTy = dyn_cast<RankedTensorType>(op->getOperand(0).getType());
1167 auto resultTy = dyn_cast<RankedTensorType>(op->getResult(0).getType());
1168 if (!inputTy || !resultTy)
1171 auto elementTy = resultTy.getElementType();
1172 Value input = op->getOperand(0);
1175 bool widenAccTy = std::is_same_v<OpTy, tosa::ReduceSumOp> &&
1176 isa<FloatType>(elementTy) &&
1177 cast<FloatType>(elementTy).isBF16();
1182 for (
unsigned i = 0; i < inputTy.getRank(); i++) {
1184 reduceShape.push_back(inputTy.getDimSize(i));
1185 if (inputTy.isDynamicDim(i))
1186 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1191 inputs.push_back(input);
1195 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1201 op,
"No initial value found for reduction operation");
1203 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
1205 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValue},
1208 outputs.push_back(filledTensor);
1210 bool isNanIgnoreMode =
false;
1211 if constexpr (std::is_same_v<OpTy, tosa::ReduceMinOp> ||
1212 std::is_same_v<OpTy, tosa::ReduceMaxOp>) {
1214 if (isa<FloatType>(elementTy) &&
1215 op.getNanMode() == NanPropagationMode::IGNORE) {
1216 isNanIgnoreMode =
true;
1222 auto trueValue = arith::ConstantOp::create(rewriter, loc, trueAttr);
1223 auto emptyBoolTensor =
1224 tensor::EmptyOp::create(rewriter, loc, reduceShape,
1225 trueValue.getType(), dynDims)
1227 auto allResultsNaNTensor =
1228 linalg::FillOp::create(rewriter, loc,
ValueRange{trueValue},
1240 inputs.push_back(input);
1241 outputs.push_back(allResultsNaNTensor);
1245 bool didEncounterError =
false;
1246 linalg::LinalgOp linalgOp = linalg::ReduceOp::create(
1247 rewriter, loc, inputs, outputs, axis,
1249 std::array<Value, 2> binaryArgs{
1250 blockArgs[0], isNanIgnoreMode ? blockArgs[2] : blockArgs[1]};
1253 if (binaryArgs[0].
getType() != accTy)
1254 binaryArgs[0] = arith::ExtFOp::create(nestedBuilder, nestedLoc, accTy,
1260 didEncounterError =
true;
1263 if (isNanIgnoreMode) {
1264 auto inputValue = blockArgs[0];
1265 auto initialValue = blockArgs[2];
1266 auto oldAllResultsNanFlagValue = blockArgs[3];
1269 Value isNaN = arith::CmpFOp::create(nestedBuilder, op->getLoc(),
1270 arith::CmpFPredicate::UNO,
1271 inputValue, inputValue);
1273 auto selectOp = arith::SelectOp::create(nestedBuilder, op->getLoc(),
1274 isNaN, initialValue,
result);
1277 auto newAllResultsNanFlagValue = arith::AndIOp::create(
1278 nestedBuilder, op->getLoc(), oldAllResultsNanFlagValue, isNaN);
1279 resultsToYield.push_back(selectOp);
1280 resultsToYield.push_back(newAllResultsNanFlagValue);
1282 resultsToYield.push_back(
result);
1284 linalg::YieldOp::create(nestedBuilder, loc, resultsToYield);
1287 if (!didEncounterError)
1289 op,
"unable to create linalg.generic body for reduce op");
1291 if (isNanIgnoreMode) {
1300 APFloat::getNaN(cast<FloatType>(elementTy).getFloatSemantics(),
false));
1301 auto nanValue = arith::ConstantOp::create(rewriter, loc, nanValueAttr);
1302 auto emptyNanTensor =
1303 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1305 auto nanFilledTensor =
1306 linalg::FillOp::create(rewriter, loc,
ValueRange{nanValue},
1312 auto finalEmptyTensor =
1313 tensor::EmptyOp::create(rewriter, loc, reduceShape, accTy, dynDims)
1319 ins.push_back(linalgOp->getOpResult(1));
1320 ins.push_back(nanFilledTensor);
1321 ins.push_back(linalgOp->getResult(0));
1322 outs.push_back(finalEmptyTensor);
1324 linalg::ElementwiseOp::create(rewriter, op->getLoc(), ins, outs,
1325 mlir::linalg::ElementwiseKind::select);
1326 linalgOp = linalgSelect;
1330 Value reducedRes = linalgOp->getResult(0);
1333 tensor::EmptyOp::create(rewriter, loc, reduceShape, elementTy, dynDims)
1336 const unsigned reducedRank =
1337 cast<ShapedType>(reducedRes.
getType()).getRank();
1340 linalg::GenericOp::create(
1346 Value truncf = arith::TruncFOp::create(nestedBuilder, nestedLoc,
1347 elementTy, args[0]);
1348 linalg::YieldOp::create(nestedBuilder, nestedLoc, truncf);
1354 uint64_t expandInputRank = cast<ShapedType>(reducedRes.
getType()).getRank();
1355 reassociationMap.resize(expandInputRank);
1357 for (uint64_t i = 0; i < expandInputRank; i++) {
1358 int32_t dimToPush = i > axis ? i + 1 : i;
1362 if (expandInputRank != 0) {
1363 int32_t expandedDim = axis < expandInputRank ? axis : expandInputRank - 1;
1364 reassociationMap[expandedDim].push_back(
1379template <
typename SrcOp>
1380class PointwiseConverter :
public OpConversionPattern<SrcOp> {
1382 using OpConversionPattern<SrcOp>::OpConversionPattern;
1383 using typename OpConversionPattern<SrcOp>::OpAdaptor;
1386 matchAndRewrite(SrcOp op, OpAdaptor operands,
1387 ConversionPatternRewriter &rewriter)
const final {
1389 op, operands.getOperands(), rewriter, *this->getTypeConverter());
1399 auto inputType = cast<RankedTensorType>(input.
getType());
1400 auto elemType = inputType.getElementType();
1401 auto collapsedType = RankedTensorType::get({}, elemType);
1403 return tensor::CollapseShapeOp::create(rewriter, loc, collapsedType, input,
1410 output.reserve(input.size());
1412 for (
auto v : llvm::map_range(
1413 input, [](int32_t val) {
return static_cast<int8_t
>(val); })) {
1414 output.push_back(v);
1426static void setupLinalgGenericOpInputAndIndexingMap(
1429 bool isConstant, tosa::RescaleOp op,
Value &constant,
int64_t &arg,
1430 bool isShift =
false) {
1432 auto loc = op.getLoc();
1433 auto inputTy = cast<ShapedType>(op.getInput().getType());
1434 unsigned rank = inputTy.getRank();
1440 if (values.size() == 1) {
1441 IntegerAttr intAttr = isShift
1444 constant = arith::ConstantOp::create(rewriter, loc, intAttr);
1448 auto tensorType = RankedTensorType::get(
1449 {
static_cast<int64_t>(values.size())}, elementType);
1455 genericInputs.push_back(
1456 arith::ConstantOp::create(rewriter, loc, EltAttr));
1464 auto operand = isShift ? op.getShift() : op.getMultiplier();
1465 auto tensorType = dyn_cast<RankedTensorType>(operand.getType());
1466 if (tensorType && tensorType.hasStaticShape() &&
1467 tensorType.getShape()[0] == 1) {
1472 genericInputs.push_back(collapse1xNTensorToN(rewriter, operand, loc));
1473 indexingMaps.push_back(broadcastMap);
1475 genericInputs.push_back(operand);
1481 arg = indexingMaps.size() - 1;
1486 FailureOr<int64_t> maybeZp,
Location loc,
1488 bool isOutputZp =
false) {
1491 const uint32_t attrBitwidth =
1492 isOutputZp ? 32 : (bitwidth > 32 ? bitwidth : 32);
1499 result = blockArgs[zpArg];
1500 auto zpTy =
result.getType();
1501 if (zpTy.getIntOrFloatBitWidth() < attrBitwidth) {
1504 if (zpTy.isUnsignedInteger()) {
1506 UnrealizedConversionCastOp::create(
1511 if (zpTy.isUnsignedInteger()) {
1512 return arith::ExtUIOp::create(builder, loc, extendType,
result);
1514 return arith::ExtSIOp::create(builder, loc, extendType,
result);
1518 return arith::ConstantOp::create(builder, loc,
1519 IntegerAttr::get(extendType, *maybeZp));
1526 using OpRewritePattern<tosa::RescaleOp>::OpRewritePattern;
1528 LogicalResult matchAndRewrite(tosa::RescaleOp op,
1529 PatternRewriter &rewriter)
const final {
1530 auto loc = op.getLoc();
1531 auto input = op.getInput();
1532 auto inputTy = cast<ShapedType>(op.getInput().getType());
1533 auto outputTy = cast<ShapedType>(op.getOutput().getType());
1534 unsigned rank = inputTy.getRank();
1537 if (op.getRoundingMode() == RoundingMode::INEXACT_ROUND)
1539 op,
"tosa.rescale with rounding mode = 'INEXACT_ROUND' is not "
1540 "currently supported");
1541 if (op.getRoundingMode() == RoundingMode::DOUBLE_ROUND && !op.getScale32())
1543 op,
"tosa.rescale requires scale32 for double_round to be true");
1545 if (!isa<IntegerType>(inputTy.getElementType()))
1548 SmallVector<Value> dynDims;
1549 for (
int i = 0; i < outputTy.getRank(); i++) {
1550 if (outputTy.isDynamicDim(i)) {
1551 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
1555 DenseElementsAttr shiftElems;
1556 bool isShiftConstant =
false;
1558 isShiftConstant =
true;
1560 DenseElementsAttr multiplierElems;
1561 bool isMultiplierConstant =
false;
1563 isMultiplierConstant =
true;
1565 llvm::SmallVector<int32_t> shiftValues;
1566 llvm::SmallVector<int32_t> multiplierValues;
1569 if (isMultiplierConstant && isShiftConstant) {
1571 shiftValues = llvm::map_to_vector(
1572 shiftElems.
getValues<IntegerAttr>(), [](IntegerAttr attr) -> int32_t {
1573 return static_cast<int32_t>(attr.getInt());
1576 llvm::map_to_vector(multiplierElems.
getValues<IntegerAttr>(),
1577 [](IntegerAttr attr) -> int32_t {
1578 return static_cast<int32_t>(attr.getInt());
1582 for (
int i = 0, s = multiplierValues.size(); i < s; i++) {
1583 if (shiftValues[i] > 63) {
1585 multiplierValues[i] = 0;
1590 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND &&
1591 llvm::any_of(shiftValues, [](int32_t v) {
return v > 31; });
1593 doubleRound = op.getRoundingMode() == RoundingMode::DOUBLE_ROUND;
1595 RoundingMode roundingMode =
1596 doubleRound ? RoundingMode::DOUBLE_ROUND : RoundingMode::SINGLE_ROUND;
1598 SmallVector<AffineMap> indexingMaps = {
1600 SmallVector<Value, 4> genericInputs = {input};
1604 Value multiplierConstant;
1605 int64_t multiplierArg = 0;
1606 setupLinalgGenericOpInputAndIndexingMap(
1607 rewriter, multiplierValues, genericInputs, indexingMaps,
1608 isMultiplierConstant, op, multiplierConstant, multiplierArg);
1612 Value shiftConstant;
1613 int64_t shiftArg = 0;
1614 setupLinalgGenericOpInputAndIndexingMap(
1615 rewriter, shiftValues, genericInputs, indexingMaps, isShiftConstant, op,
1616 shiftConstant, shiftArg,
true);
1621 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1622 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1632 genericInputs.push_back(
1633 collapse1xNTensorToN(rewriter, op->getOperand(3), loc));
1634 indexingMaps.push_back(broadcastMap);
1635 iZpArg = indexingMaps.size() - 1;
1639 genericInputs.push_back(
1640 collapse1xNTensorToN(rewriter, op->getOperand(4), loc));
1641 indexingMaps.push_back(broadcastMap);
1642 oZpArg = indexingMaps.size() - 1;
1649 Value emptyTensor = tensor::EmptyOp::create(
1650 rewriter, loc, outputTy.getShape(), outputTy.getElementType(),
1651 ArrayRef<Value>({dynDims}));
1653 auto linalgOp = linalg::GenericOp::create(
1654 rewriter, loc, outputTy, genericInputs,
ValueRange{emptyTensor},
1656 [&](OpBuilder &nestedBuilder, Location nestedLoc,
1658 Value value = blockArgs[0];
1659 Type valueTy = value.
getType();
1661 FailureOr<int64_t> maybeIZp = op.getInputZeroPoint();
1662 auto inputZp = getExtendZp(nestedBuilder, valueTy, maybeIZp,
1663 nestedLoc, blockArgs, iZpArg);
1665 FailureOr<int64_t> maybeOZp = op.getOutputZeroPoint();
1666 auto outputZp = getExtendZp(nestedBuilder, valueTy, maybeOZp,
1667 nestedLoc, blockArgs, oZpArg,
true);
1669 IntegerType outIntType =
1670 cast<IntegerType>(blockArgs.back().
getType());
1671 unsigned outBitWidth = outIntType.getWidth();
1672 assert(outBitWidth <= 32 &&
"Unexpected output zeropoint bitwidth");
1674 Value multiplier = multiplierConstant ? multiplierConstant
1675 : blockArgs[multiplierArg];
1676 Value shift = shiftConstant ? shiftConstant : blockArgs[shiftArg];
1679 value = UnrealizedConversionCastOp::create(
1680 nestedBuilder, nestedLoc,
1681 nestedBuilder.getIntegerType(
1687 if (op.getInputUnsigned()) {
1688 value = arith::ExtUIOp::create(nestedBuilder, nestedLoc,
1689 nestedBuilder.getI32Type(), value);
1691 value = arith::ExtSIOp::create(nestedBuilder, nestedLoc,
1692 nestedBuilder.getI32Type(), value);
1697 arith::SubIOp::create(nestedBuilder, nestedLoc, value, inputZp);
1699 value = tosa::ApplyScaleOp::create(nestedBuilder, loc,
1700 nestedBuilder.getI32Type(), value,
1701 multiplier, shift, roundingMode);
1705 arith::AddIOp::create(nestedBuilder, nestedLoc, value, outputZp);
1708 int32_t intMin = APInt::getSignedMinValue(outBitWidth).getSExtValue();
1709 int32_t intMax = APInt::getSignedMaxValue(outBitWidth).getSExtValue();
1712 if (op.getOutputUnsigned()) {
1714 intMax = APInt::getMaxValue(outBitWidth).getZExtValue();
1717 auto intMinVal = arith::ConstantOp::create(
1718 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMin));
1719 auto intMaxVal = arith::ConstantOp::create(
1720 nestedBuilder, loc, nestedBuilder.getI32IntegerAttr(intMax));
1723 nestedBuilder,
false);
1725 if (outIntType.getWidth() < 32) {
1726 value = arith::TruncIOp::create(
1727 nestedBuilder, nestedLoc,
1731 if (outIntType.isUnsignedInteger()) {
1732 value = UnrealizedConversionCastOp::create(nestedBuilder, nestedLoc,
1736 linalg::YieldOp::create(nestedBuilder, loc, value);
1739 rewriter.
replaceOp(op, linalgOp->getResults());
1749 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
1751 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1752 PatternRewriter &rewriter)
const final {
1753 Location loc = op.getLoc();
1754 ImplicitLocOpBuilder builder(loc, rewriter);
1755 auto input = op.getInput();
1756 auto inputTy = cast<RankedTensorType>(input.getType());
1757 auto resultTy = cast<RankedTensorType>(op.getType());
1758 const bool isBilinear = op.getMode() == ResizeMode::BILINEAR;
1760 auto inputH = inputTy.getDimSize(1);
1761 auto inputW = inputTy.getDimSize(2);
1762 auto outputH = resultTy.getDimSize(1);
1763 auto outputW = resultTy.getDimSize(2);
1765 if (inputH != 1 || inputW != 1 || outputH != 1 || outputW != 1)
1767 op,
"tosa.resize is not a pure 1x1->1x1 image operation");
1769 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
1770 op.getMode() != ResizeMode::BILINEAR)
1772 op,
"tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
1774 if (inputTy == resultTy) {
1779 SmallVector<int64_t> scale;
1785 SmallVector<ReassociationExprs, 4> reassociationMap(2);
1792 RankedTensorType::get({inputTy.getDimSize(0), inputTy.getDimSize(3)},
1793 inputTy.getElementType());
1794 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy, input,
1798 llvm::SmallVector<Value> outputDynSize;
1799 if (inputTy.isDynamicDim(0))
1800 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1801 if (inputTy.isDynamicDim(3))
1802 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1805 auto genericTy = collapseTy.clone(resultTy.getElementType());
1807 tensor::EmptyOp::create(builder, genericTy.getShape(),
1808 resultTy.getElementType(), outputDynSize);
1810 SmallVector<utils::IteratorType> iterators(genericTy.getRank(),
1811 utils::IteratorType::parallel);
1813 auto generic = linalg::GenericOp::create(
1815 ArrayRef<AffineMap>{genericMap, genericMap}, iterators,
1816 [=](OpBuilder &
b, Location loc,
ValueRange args) {
1817 Value value = args[0];
1819 if (inputTy.getElementType() != resultTy.getElementType()) {
1820 value = arith::ExtSIOp::create(
b, loc, resultTy.getElementType(),
1823 if (isBilinear && scale[0] != 0) {
1824 Value scaleY = arith::ConstantOp::create(
1825 b, loc,
b.getI32IntegerAttr(scale[0]));
1826 value = arith::MulIOp::create(
b, loc, value, scaleY);
1829 if (isBilinear && scale[2] != 0) {
1830 Value scaleX = arith::ConstantOp::create(
1831 b, loc,
b.getI32IntegerAttr(scale[2]));
1832 value = arith::MulIOp::create(
b, loc, value, scaleX);
1836 linalg::YieldOp::create(
b, loc, value);
1840 op, resultTy,
generic.getResults()[0], reassociationMap);
1852 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1856 auto input = op.getInput();
1857 auto inputTy = dyn_cast<RankedTensorType>(input.getType());
1858 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
1860 if (!inputTy || !resultTy)
1862 "requires ranked input/output types");
1864 auto batch = inputTy.getDimSize(0);
1865 auto channels = inputTy.getDimSize(3);
1866 auto inputH = inputTy.getDimSize(1);
1867 auto inputW = inputTy.getDimSize(2);
1868 auto outputH = resultTy.getDimSize(1);
1869 auto outputW = resultTy.getDimSize(2);
1871 if ((inputH != 1 || outputH == 1) && (inputW != 1 || outputW == 1))
1873 op,
"tosa.resize has no broadcasting behavior");
1878 resizeShape.push_back(batch);
1879 resizeShape.push_back(inputH == 1 ? 1 : outputH);
1880 resizeShape.push_back(inputW == 1 ? 1 : outputW);
1881 resizeShape.push_back(channels);
1883 auto resizeTy = resultTy.clone(resizeShape);
1885 tosa::ResizeOp::create(builder, resizeTy, input, op.getScale(),
1886 op.getOffset(), op.getBorder(), op.getMode());
1893 reassociationMap.push_back({});
1896 reassociationMap.push_back({});
1901 collapseShape.push_back(outputH);
1903 collapseShape.push_back(outputW);
1904 collapseShape.push_back(channels);
1906 auto collapseTy = resultTy.clone(collapseShape);
1907 Value collapse = tensor::CollapseShapeOp::create(builder, collapseTy,
1908 resize, reassociationMap);
1912 if (inputTy.isDynamicDim(0))
1913 outputDynSize.push_back(tensor::DimOp::create(builder, input, 0));
1914 if (inputTy.isDynamicDim(3))
1915 outputDynSize.push_back(tensor::DimOp::create(builder, input, 3));
1918 utils::IteratorType::parallel);
1919 Value empty = tensor::EmptyOp::create(
1920 builder, resultTy.getShape(), resultTy.getElementType(), outputDynSize);
1937 Value value = args[0];
1938 linalg::YieldOp::create(
b, loc, value);
1947 using OpRewritePattern<tosa::ResizeOp>::OpRewritePattern;
1949 LogicalResult matchAndRewrite(tosa::ResizeOp op,
1950 PatternRewriter &rewriter)
const final {
1951 Location loc = op.getLoc();
1952 ImplicitLocOpBuilder
b(loc, rewriter);
1953 auto input = op.getInput();
1954 auto inputTy = cast<ShapedType>(input.getType());
1955 auto resultTy = cast<ShapedType>(op.getType());
1956 auto resultETy = resultTy.getElementType();
1958 bool floatingPointMode = isa<FloatType>(resultETy);
1959 auto floatTy = resultETy;
1961 auto imageH = inputTy.getShape()[1];
1962 auto imageW = inputTy.getShape()[2];
1964 auto dynamicDimsOr =
1966 if (!dynamicDimsOr.has_value())
1968 op,
"unable to get dynamic dimensions of tosa.resize");
1970 if (op.getMode() != ResizeMode::NEAREST_NEIGHBOR &&
1971 op.getMode() != ResizeMode::BILINEAR)
1973 op,
"tosa.resize mode should be NEAREST_NEIGHBOR or BILINEAR");
1975 SmallVector<AffineMap, 2> affineMaps = {
1977 auto emptyTensor = tensor::EmptyOp::create(
b, resultTy.getShape(),
1978 resultETy, *dynamicDimsOr);
1979 auto genericOp = linalg::GenericOp::create(
1982 Value resize = genericOp.getResult(0);
1985 OpBuilder::InsertionGuard regionGuard(
b);
1986 b.createBlock(&genericOp.getRegion(), genericOp.getRegion().end(),
1988 Value batch = linalg::IndexOp::create(
b, 0);
1989 Value y = linalg::IndexOp::create(
b, 1);
1990 Value x = linalg::IndexOp::create(
b, 2);
1991 Value channel = linalg::IndexOp::create(
b, 3);
1994 arith::ConstantOp::create(
b,
b.getZeroAttr(
b.getI32Type()));
1995 Value zeroFp = arith::ConstantOp::create(
b,
b.getZeroAttr(floatTy));
1997 arith::ConstantOp::create(
b,
b.getI32IntegerAttr(imageH - 1));
1999 arith::ConstantOp::create(
b,
b.getI32IntegerAttr(imageW - 1));
2001 Value inY = arith::IndexCastOp::create(
b,
b.getI32Type(), y);
2002 Value inX = arith::IndexCastOp::create(
b,
b.getI32Type(), x);
2004 SmallVector<int64_t> scale, offset, border;
2009 op,
"tosa.resize scale/offset/border should have compile time "
2010 "constant values.");
2013 Value yScaleN, yScaleD, xScaleN, xScaleD;
2014 yScaleN = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[0]));
2015 yScaleD = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[1]));
2016 xScaleN = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[2]));
2017 xScaleD = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(scale[3]));
2019 Value yOffset, xOffset, yBorder, xBorder;
2020 yOffset = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(offset[0]));
2021 xOffset = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(offset[1]));
2022 yBorder = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(border[0]));
2023 xBorder = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(border[1]));
2026 auto getIndexAndDeltaFp = [&](Value &index, Value &delta, Value in,
2027 Value scaleN, Value scaleD, Value offset,
2028 int size, ImplicitLocOpBuilder &
b) {
2036 Value val = arith::MulIOp::create(
b, in, scaleD);
2037 val = arith::AddIOp::create(
b, val, offset);
2038 index = arith::FloorDivSIOp::create(
b, val, scaleN);
2041 Value scaledIndex = arith::MulIOp::create(
b, index, scaleN);
2042 Value r = arith::SubIOp::create(
b, val, scaledIndex);
2043 Value rFp = arith::SIToFPOp::create(
b, floatTy, r);
2046 Value scaleNfp = arith::UIToFPOp::create(
b, floatTy, scaleN);
2047 delta = arith::DivFOp::create(
b, rFp, scaleNfp);
2051 auto getIndexAndDeltaInt = [&](Value &index, Value &delta, Value in,
2052 Value scaleN, Value scaleD, Value offset,
2053 int size, ImplicitLocOpBuilder &
b) {
2062 Value val = arith::MulIOp::create(
b, in, scaleD);
2063 val = arith::AddIOp::create(
b, val, offset);
2064 index = arith::FloorDivSIOp::create(
b, val, scaleN);
2065 delta = arith::MulIOp::create(
b, index, scaleN);
2066 delta = arith::SubIOp::create(
b, val, delta);
2069 Value ix, iy, dx, dy;
2070 if (floatingPointMode) {
2071 getIndexAndDeltaFp(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH,
b);
2072 getIndexAndDeltaFp(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW,
b);
2074 getIndexAndDeltaInt(iy, dy, inY, yScaleN, yScaleD, yOffset, imageH,
b);
2075 getIndexAndDeltaInt(ix, dx, inX, xScaleN, xScaleD, xOffset, imageW,
b);
2078 if (op.getMode() == ResizeMode::NEAREST_NEIGHBOR) {
2079 auto one = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(1));
2081 auto getNearestIndexAndClamp = [&](Value val, Value dval, Value scale,
2082 Value
max,
int size,
2083 ImplicitLocOpBuilder &
b) -> Value {
2089 if (floatingPointMode) {
2091 arith::ConstantOp::create(
b,
b.getFloatAttr(floatTy, 0.5f));
2092 pred = arith::CmpFOp::create(
b, arith::CmpFPredicate::OGE, dval, h);
2094 Value dvalDouble = arith::ShLIOp::create(
b, dval, one);
2095 pred = arith::CmpIOp::create(
b, arith::CmpIPredicate::sge,
2099 auto offset = arith::SelectOp::create(
b, pred, one, zeroI32);
2100 val = arith::AddIOp::create(
b, val, offset);
2102 return arith::IndexCastOp::create(
b,
b.getIndexType(), val);
2105 iy = getNearestIndexAndClamp(iy, dy, yScaleN, hMax, imageH,
b);
2106 ix = getNearestIndexAndClamp(ix, dx, xScaleN, wMax, imageW,
b);
2108 Value
result = tensor::ExtractOp::create(
2111 linalg::YieldOp::create(
b,
result);
2114 assert(op.getMode() == ResizeMode::BILINEAR);
2116 auto oneVal = arith::ConstantOp::create(
b,
b.getI32IntegerAttr(1));
2118 auto getClampedIdxs = [&](Value &val0, Value &val1,
int size, Value in,
2119 Value
max, ImplicitLocOpBuilder &
b) {
2121 val1 = arith::AddIOp::create(
b, val0, oneVal);
2126 val0 = arith::IndexCastOp::create(
b,
b.getIndexType(), val0);
2127 val1 = arith::IndexCastOp::create(
b,
b.getIndexType(), val1);
2135 Value x0, x1, y0, y1;
2136 getClampedIdxs(y0, y1, imageH, iy, hMax,
b);
2137 getClampedIdxs(x0, x1, imageW, ix, wMax,
b);
2139 Value y0x0 = tensor::ExtractOp::create(
2141 Value y0x1 = tensor::ExtractOp::create(
2143 Value y1x0 = tensor::ExtractOp::create(
2145 Value y1x1 = tensor::ExtractOp::create(
2148 if (floatingPointMode) {
2150 arith::ConstantOp::create(
b,
b.getFloatAttr(floatTy, 1.0f));
2151 auto interpolate = [&](Value val0, Value val1, Value delta,
2153 ImplicitLocOpBuilder &
b) -> Value {
2156 Value oneMinusDelta = arith::SubFOp::create(
b, oneVal, delta);
2157 Value mul0 = arith::MulFOp::create(
b, val0, oneMinusDelta);
2158 Value mul1 = arith::MulFOp::create(
b, val1, delta);
2159 return arith::AddFOp::create(
b, mul0, mul1);
2165 Value topAcc = interpolate(y0x0, y0x1, dx, imageW,
b);
2170 Value bottomAcc = interpolate(y1x0, y1x1, dx, imageW,
b);
2174 Value
result = interpolate(topAcc, bottomAcc, dy, imageH,
b);
2175 linalg::YieldOp::create(
b,
result);
2178 y0x0 = arith::ExtSIOp::create(
b, resultETy, y0x0);
2179 y0x1 = arith::ExtSIOp::create(
b, resultETy, y0x1);
2180 y1x0 = arith::ExtSIOp::create(
b, resultETy, y1x0);
2181 y1x1 = arith::ExtSIOp::create(
b, resultETy, y1x1);
2184 if (resultETy.getIntOrFloatBitWidth() > deltaBitwidth) {
2185 dx = arith::ExtSIOp::create(
b, resultETy, dx);
2186 dy = arith::ExtSIOp::create(
b, resultETy, dy);
2189 Value yScaleNExt = yScaleN;
2190 Value xScaleNExt = xScaleN;
2192 const int64_t scaleBitwidth =
2194 if (resultETy.getIntOrFloatBitWidth() > scaleBitwidth) {
2195 yScaleNExt = arith::ExtSIOp::create(
b, resultETy, yScaleN);
2196 xScaleNExt = arith::ExtSIOp::create(
b, resultETy, xScaleN);
2199 auto interpolate = [](Value val0, Value val1, Value weight1,
2200 Value scale,
int inputSize,
2201 ImplicitLocOpBuilder &
b) -> Value {
2203 return arith::MulIOp::create(
b, val0, scale);
2204 Value weight0 = arith::SubIOp::create(
b, scale, weight1);
2205 Value mul0 = arith::MulIOp::create(
b, val0, weight0);
2206 Value mul1 = arith::MulIOp::create(
b, val1, weight1);
2207 return arith::AddIOp::create(
b, mul0, mul1);
2210 Value topAcc = interpolate(y0x0, y0x1, dx, xScaleNExt, imageW,
b);
2211 Value bottomAcc = interpolate(y1x0, y1x1, dx, xScaleNExt, imageW,
b);
2213 interpolate(topAcc, bottomAcc, dy, yScaleNExt, imageH,
b);
2214 linalg::YieldOp::create(
b,
result);
2227template <
typename SrcOp>
2230 using OpRewritePattern<SrcOp>::OpRewritePattern;
2232 LogicalResult matchAndRewrite(SrcOp op,
2233 PatternRewriter &rewriter)
const final {
2234 rewriter.
replaceOp(op, op.getOperation()->getOperands());
2239template <
typename SrcOp>
2242 using OpRewritePattern<SrcOp>::OpRewritePattern;
2244 LogicalResult matchAndRewrite(SrcOp reduceOp,
2245 PatternRewriter &rewriter)
const final {
2252 using OpRewritePattern<tosa::ReverseOp>::OpRewritePattern;
2254 LogicalResult matchAndRewrite(tosa::ReverseOp op,
2255 PatternRewriter &rewriter)
const final {
2256 auto loc = op.getLoc();
2257 Value input = op.getInput1();
2258 auto inputTy = cast<ShapedType>(input.
getType());
2259 auto resultTy = cast<ShapedType>(op.getType());
2260 auto axis = op.getAxis();
2262 SmallVector<Value> dynDims;
2263 for (
int i = 0; i < inputTy.getRank(); i++) {
2264 if (inputTy.isDynamicDim(i)) {
2265 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2269 Value axisDimSize = tensor::DimOp::create(rewriter, loc, input, axis);
2272 auto emptyTensor = tensor::EmptyOp::create(
2273 rewriter, loc, inputTy.getShape(),
2274 inputTy.getElementType(), ArrayRef<Value>({dynDims}))
2276 SmallVector<AffineMap, 2> affineMaps = {
2280 op, resultTy, ArrayRef<Value>({}),
ValueRange{emptyTensor}, affineMaps,
2282 [&](OpBuilder &nestedBuilder, Location nestedLoc,
ValueRange args) {
2283 llvm::SmallVector<Value>
indices;
2284 for (
unsigned int i = 0; i < inputTy.getRank(); i++) {
2286 linalg::IndexOp::create(rewriter, nestedLoc, i).getResult();
2290 arith::SubIOp::create(rewriter, nestedLoc, axisDimSize, one);
2291 index = arith::SubIOp::create(rewriter, nestedLoc, sizeMinusOne,
2298 auto extract = tensor::ExtractOp::create(nestedBuilder, nestedLoc,
2300 linalg::YieldOp::create(nestedBuilder, op.getLoc(),
2301 extract.getResult());
2311struct TileConverter :
public OpConversionPattern<tosa::TileOp> {
2312 using OpConversionPattern<tosa::TileOp>::OpConversionPattern;
2315 matchAndRewrite(tosa::TileOp op, OpAdaptor adaptor,
2316 ConversionPatternRewriter &rewriter)
const override {
2317 auto loc = op.getLoc();
2318 auto input = op.getInput1();
2319 auto inputTy = cast<ShapedType>(input.
getType());
2320 auto inputShape = inputTy.getShape();
2321 auto resultTy = cast<ShapedType>(op.getType());
2322 auto elementTy = inputTy.getElementType();
2323 int64_t rank = inputTy.getRank();
2325 SmallVector<int64_t> multiples;
2326 if (
failed(op.getConstantMultiples(multiples)))
2330 SmallVector<int64_t, 2> genericShape;
2331 for (
int i = 0; i < rank; i++) {
2332 int64_t dim = multiples[i];
2333 genericShape.push_back(dim == -1 ? ShapedType::kDynamic : dim);
2334 genericShape.push_back(inputShape[i]);
2337 SmallVector<Value> dynDims;
2338 for (
int i = 0; i < inputTy.getRank(); i++) {
2339 if (inputTy.isDynamicDim(i) || multiples[i] == -1) {
2340 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2344 auto emptyTensor = tensor::EmptyOp::create(
2345 rewriter, op.getLoc(), genericShape, elementTy, dynDims);
2348 SmallVector<AffineExpr, 4> dimExprs;
2349 dimExprs.reserve(rank);
2350 for (
unsigned i = 0; i < rank; ++i)
2351 dimExprs.push_back(rewriter.getAffineDimExpr(i * 2 + 1));
2353 auto readAffineMap =
2355 rewriter.getContext());
2357 SmallVector<AffineMap, 2> affineMaps = {
2358 readAffineMap, rewriter.getMultiDimIdentityMap(genericShape.size())};
2360 auto genericOp = linalg::GenericOp::create(
2361 rewriter, loc, RankedTensorType::get(genericShape, elementTy), input,
2364 [&](OpBuilder &nestedBuilder, Location nestedLoc,
ValueRange args) {
2365 linalg::YieldOp::create(nestedBuilder, op.getLoc(), *args.begin());
2370 rewriter.replaceOpWithNewOp<tosa::ReshapeOp>(
2371 op, resultTy, genericOp.getResult(0), shapeValue);
2391 using OpRewritePattern<tosa::ArgMaxOp>::OpRewritePattern;
2393 LogicalResult matchAndRewrite(tosa::ArgMaxOp argmaxOp,
2394 PatternRewriter &rewriter)
const final {
2395 auto loc = argmaxOp.getLoc();
2396 Value input = argmaxOp.getInput();
2397 auto inputTy = cast<ShapedType>(input.
getType());
2398 auto resultTy = cast<ShapedType>(argmaxOp.getOutput().getType());
2399 auto inElementTy = inputTy.getElementType();
2400 auto outElementTy = resultTy.getElementType();
2401 int axis = argmaxOp.getAxis();
2402 auto resultMaxTy = RankedTensorType::get(resultTy.getShape(), inElementTy);
2404 if (!isa<IntegerType>(outElementTy))
2405 return rewriter.notifyMatchFailure(
2407 "tosa.arg_max to linalg.* requires integer-like result type");
2409 SmallVector<Value> dynDims;
2410 for (
int i = 0; i < inputTy.getRank(); i++) {
2411 if (inputTy.isDynamicDim(i) && i != axis) {
2412 dynDims.push_back(tensor::DimOp::create(rewriter, loc, input, i));
2417 auto emptyTensorIdx =
2418 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2419 outElementTy, dynDims)
2421 auto fillValueIdx = arith::ConstantOp::create(
2422 rewriter, loc, rewriter.getIntegerAttr(outElementTy, 0));
2423 auto filledTensorIdx =
2424 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValueIdx},
2429 auto emptyTensorMax =
2430 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(), inElementTy,
2433 auto fillValueMaxAttr =
2436 if (!fillValueMaxAttr)
2437 return rewriter.notifyMatchFailure(
2438 argmaxOp,
"unsupported tosa.argmax element type");
2441 arith::ConstantOp::create(rewriter, loc, fillValueMaxAttr);
2442 auto filledTensorMax =
2443 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValueMax},
2449 SmallVector<utils::IteratorType, 4> iteratorTypes;
2450 iteratorTypes.resize(inputTy.getRank(), utils::IteratorType::parallel);
2451 iteratorTypes[axis] = utils::IteratorType::reduction;
2453 SmallVector<AffineExpr, 2> srcExprs;
2454 SmallVector<AffineExpr, 2> dstExprs;
2455 for (
int i = 0, rank = inputTy.getRank(); i != rank; ++i) {
2461 bool didEncounterError =
false;
2463 rewriter.getContext());
2464 auto linalgOp = linalg::GenericOp::create(
2465 rewriter, loc, ArrayRef<Type>({resultTy, resultMaxTy}), input,
2466 ValueRange({filledTensorIdx, filledTensorMax}), maps, iteratorTypes,
2467 [&](OpBuilder &nestedBuilder, Location nestedLoc,
2469 auto newValue = blockArgs[0];
2470 auto oldIndex = blockArgs[1];
2471 auto oldValue = blockArgs[2];
2473 Value newIndex = arith::IndexCastOp::create(
2474 rewriter, nestedLoc, oldIndex.getType(),
2475 linalg::IndexOp::create(rewriter, loc, axis));
2478 if (isa<FloatType>(inElementTy)) {
2479 if (argmaxOp.getNanMode() == NanPropagationMode::IGNORE) {
2482 predicate = arith::CmpFOp::create(rewriter, nestedLoc,
2483 arith::CmpFPredicate::OGT,
2484 newValue, oldValue);
2489 Value gt = arith::CmpFOp::create(rewriter, nestedLoc,
2490 arith::CmpFPredicate::UGT,
2491 newValue, oldValue);
2492 Value oldNonNaN = arith::CmpFOp::create(rewriter, nestedLoc,
2493 arith::CmpFPredicate::ORD,
2494 oldValue, oldValue);
2495 predicate = arith::AndIOp::create(
2496 rewriter, nestedLoc, rewriter.getI1Type(), gt, oldNonNaN);
2498 }
else if (isa<IntegerType>(inElementTy)) {
2499 predicate = arith::CmpIOp::create(rewriter, nestedLoc,
2500 arith::CmpIPredicate::sgt,
2501 newValue, oldValue);
2503 didEncounterError =
true;
2507 auto resultMax = arith::SelectOp::create(
2508 rewriter, nestedLoc, predicate, newValue, oldValue);
2509 auto resultIndex = arith::SelectOp::create(
2510 rewriter, nestedLoc, predicate, newIndex, oldIndex);
2511 linalg::YieldOp::create(nestedBuilder, nestedLoc,
2515 if (didEncounterError)
2516 return rewriter.notifyMatchFailure(
2517 argmaxOp,
"unsupported tosa.argmax element type");
2519 rewriter.replaceOp(argmaxOp, linalgOp.getResult(0));
2524class GatherConverter :
public OpConversionPattern<tosa::GatherOp> {
2526 using OpConversionPattern<tosa::GatherOp>::OpConversionPattern;
2528 matchAndRewrite(tosa::GatherOp op, OpAdaptor adaptor,
2529 ConversionPatternRewriter &rewriter)
const final {
2530 auto input = adaptor.getOperands()[0];
2531 auto indices = adaptor.getOperands()[1];
2533 auto valuesTy = dyn_cast<RankedTensorType>(op.getValues().getType());
2534 auto resultTy = dyn_cast<RankedTensorType>(op.getType());
2535 if (!valuesTy || !resultTy)
2536 return rewriter.notifyMatchFailure(op,
"unranked tensors not supported");
2538 auto dynamicDims = inferDynamicDimsForGather(
2539 rewriter, op.getLoc(), adaptor.getValues(), adaptor.getIndices());
2541 auto resultElementTy = resultTy.getElementType();
2543 auto loc = op.getLoc();
2545 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2546 resultElementTy, dynamicDims)
2549 SmallVector<AffineMap, 2> affineMaps = {
2551 resultTy.getRank(), 0,
2552 {rewriter.getAffineDimExpr(0), rewriter.getAffineDimExpr(1)},
2553 rewriter.getContext()),
2554 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2556 auto genericOp = linalg::GenericOp::create(
2560 [&](OpBuilder &
b, Location loc,
ValueRange args) {
2561 auto indexValue = args[0];
2562 auto index0 = linalg::IndexOp::create(rewriter, loc, 0);
2563 Value index1 = arith::IndexCastOp::create(
2564 rewriter, loc, rewriter.getIndexType(), indexValue);
2565 auto index2 = linalg::IndexOp::create(rewriter, loc, 2);
2566 Value extract = tensor::ExtractOp::create(
2567 rewriter, loc, input,
ValueRange{index0, index1, index2});
2568 linalg::YieldOp::create(rewriter, loc, extract);
2570 rewriter.replaceOp(op, genericOp.getResult(0));
2574 static llvm::SmallVector<Value> inferDynamicDimsForGather(OpBuilder &builder,
2578 llvm::SmallVector<Value> results;
2580 auto addDynamicDimension = [&](Value source, int64_t dim) {
2582 if (
auto dimValue = llvm::dyn_cast_if_present<Value>(sz))
2583 results.push_back(dimValue);
2586 addDynamicDimension(values, 0);
2587 addDynamicDimension(
indices, 1);
2588 addDynamicDimension(values, 2);
2598 using OpRewritePattern<tosa::TableOp>::OpRewritePattern;
2600 LogicalResult matchAndRewrite(tosa::TableOp op,
2601 PatternRewriter &rewriter)
const final {
2602 auto loc = op.getLoc();
2603 Value input = op.getInput1();
2604 Value table = op.getTable();
2605 auto inputTy = cast<ShapedType>(input.
getType());
2606 auto tableTy = cast<ShapedType>(table.
getType());
2607 auto resultTy = cast<ShapedType>(op.getType());
2609 auto inputElementTy = inputTy.getElementType();
2610 auto tableElementTy = tableTy.getElementType();
2611 auto resultElementTy = resultTy.getElementType();
2613 SmallVector<Value> dynDims;
2614 for (
int i = 0; i < resultTy.getRank(); ++i) {
2615 if (inputTy.isDynamicDim(i)) {
2617 tensor::DimOp::create(rewriter, loc, op.getOperand(0), i));
2622 tensor::EmptyOp::create(rewriter, loc, resultTy.getShape(),
2623 resultElementTy, dynDims)
2626 SmallVector<AffineMap, 2> affineMaps = {
2627 rewriter.getMultiDimIdentityMap(resultTy.getRank()),
2628 rewriter.getMultiDimIdentityMap(resultTy.getRank())};
2630 auto genericOp = linalg::GenericOp::create(
2633 rewriter.replaceOp(op, genericOp.getResult(0));
2636 OpBuilder::InsertionGuard regionGuard(rewriter);
2637 Block *block = rewriter.createBlock(
2638 &genericOp.getRegion(), genericOp.getRegion().end(),
2639 TypeRange({inputElementTy, resultElementTy}), {loc, loc});
2642 rewriter.setInsertionPointToStart(block);
2643 if (inputElementTy.isInteger(8) && tableElementTy.isInteger(8) &&
2644 resultElementTy.isInteger(8)) {
2645 Value index = arith::IndexCastOp::create(
2646 rewriter, loc, rewriter.getIndexType(), inputValue);
2648 index = arith::AddIOp::create(rewriter, loc, rewriter.getIndexType(),
2651 tensor::ExtractOp::create(rewriter, loc, table,
ValueRange{index});
2652 linalg::YieldOp::create(rewriter, loc, extract);
2656 if (inputElementTy.isInteger(16) && tableElementTy.isInteger(16) &&
2657 resultElementTy.isInteger(32)) {
2658 Value extend = arith::ExtSIOp::create(
2659 rewriter, loc, rewriter.getI32Type(), inputValue);
2661 auto offset = arith::ConstantOp::create(
2662 rewriter, loc, rewriter.getI32IntegerAttr(32768));
2663 auto seven = arith::ConstantOp::create(rewriter, loc,
2664 rewriter.getI32IntegerAttr(7));
2665 auto one = arith::ConstantOp::create(rewriter, loc,
2666 rewriter.getI32IntegerAttr(1));
2667 auto b1111111 = arith::ConstantOp::create(
2668 rewriter, loc, rewriter.getI32IntegerAttr(127));
2674 auto extendAdd = arith::AddIOp::create(rewriter, loc, extend, offset);
2675 Value index = arith::ShRUIOp::create(rewriter, loc, extendAdd, seven);
2677 arith::AndIOp::create(rewriter, loc, extendAdd, b1111111);
2682 Value indexPlusOne = arith::AddIOp::create(rewriter, loc, index, one);
2684 index = arith::IndexCastOp::create(rewriter, loc,
2685 rewriter.getIndexType(), index);
2686 indexPlusOne = arith::IndexCastOp::create(
2687 rewriter, loc, rewriter.getIndexType(), indexPlusOne);
2690 tensor::ExtractOp::create(rewriter, loc, table,
ValueRange{index});
2691 Value next = tensor::ExtractOp::create(rewriter, loc, table,
2695 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), base);
2697 arith::ExtSIOp::create(rewriter, loc, rewriter.getI32Type(), next);
2701 Value baseScaled = arith::ShLIOp::create(rewriter, loc, base, seven);
2702 Value diff = arith::SubIOp::create(rewriter, loc, next, base);
2703 Value diffScaled = arith::MulIOp::create(rewriter, loc, diff, fraction);
2705 arith::AddIOp::create(rewriter, loc, baseScaled, diffScaled);
2707 linalg::YieldOp::create(rewriter, loc,
result);
2713 return rewriter.notifyMatchFailure(
2714 op,
"unable to create body for tosa.table op");
2719 using OpRewritePattern<RFFT2dOp>::OpRewritePattern;
2721 static bool isRankedTensor(Type type) {
return isa<RankedTensorType>(type); }
2723 static OpFoldResult halfPlusOne(OpBuilder &builder, Location loc,
2729 auto divBy2 = builder.
createOrFold<arith::DivUIOp>(loc, value, two);
2730 auto plusOne = builder.
createOrFold<arith::AddIOp>(loc, divBy2, one);
2734 static RankedTensorType
2735 computeOutputShape(OpBuilder &builder, Location loc, Value input,
2736 llvm::SmallVectorImpl<Value> &dynamicSizes) {
2742 dims[2] = halfPlusOne(builder, loc, dims[2]);
2744 llvm::SmallVector<int64_t, 3> staticSizes;
2747 auto elementType = cast<RankedTensorType>(input.
getType()).getElementType();
2748 return RankedTensorType::get(staticSizes, elementType);
2751 static Value createZeroTensor(PatternRewriter &rewriter, Location loc,
2752 RankedTensorType type,
2753 llvm::ArrayRef<Value> dynamicSizes) {
2755 tensor::EmptyOp::create(rewriter, loc, type, dynamicSizes);
2756 auto fillValueAttr = rewriter.
getZeroAttr(type.getElementType());
2757 auto fillValue = arith::ConstantOp::create(rewriter, loc, fillValueAttr);
2759 linalg::FillOp::create(rewriter, loc,
ValueRange{fillValue},
2762 return filledTensor;
2765 static Value castIndexToFloat(OpBuilder &builder, Location loc,
2766 FloatType type, Value value) {
2767 auto integerVal = arith::IndexCastUIOp::create(
2769 type.getIntOrFloatBitWidth() > 32 ? builder.
getI64Type()
2773 return arith::UIToFPOp::create(builder, loc, type, integerVal);
2776 static Value createLinalgIndex(OpBuilder &builder, Location loc,
2777 FloatType type, int64_t index) {
2778 auto indexVal = linalg::IndexOp::create(builder, loc, index);
2779 return castIndexToFloat(builder, loc, type, indexVal);
2782 template <
typename... Args>
2783 static llvm::SmallVector<AffineExpr, 4> affineDimsExpr(OpBuilder &builder,
2788 LogicalResult matchAndRewrite(RFFT2dOp rfft2d,
2789 PatternRewriter &rewriter)
const override {
2790 if (!llvm::all_of(rfft2d->getOperandTypes(), isRankedTensor) ||
2791 !llvm::all_of(rfft2d->getResultTypes(), isRankedTensor)) {
2793 "only supports ranked tensors");
2796 auto loc = rfft2d.getLoc();
2797 auto input = rfft2d.getInputReal();
2799 dyn_cast<FloatType>(cast<ShapedType>(input.
getType()).getElementType());
2802 "only supports float element types");
2805 llvm::SmallVector<Value> dynamicSizes;
2806 auto outputType = computeOutputShape(rewriter, loc, input, dynamicSizes);
2809 llvm::SmallVector<utils::IteratorType, 5> iteratorTypes = {
2810 utils::IteratorType::parallel, utils::IteratorType::parallel,
2811 utils::IteratorType::parallel, utils::IteratorType::reduction,
2812 utils::IteratorType::reduction};
2815 llvm::SmallVector<Value> genericOpInputs = {input};
2816 llvm::SmallVector<Value> genericOpOutputs = {
2817 createZeroTensor(rewriter, loc, outputType, dynamicSizes),
2818 createZeroTensor(rewriter, loc, outputType, dynamicSizes)};
2822 llvm::ArrayRef{affineDimsExpr(rewriter, 0, 3, 4),
2823 affineDimsExpr(rewriter, 0, 1, 2),
2824 affineDimsExpr(rewriter, 0, 1, 2)},
2828 auto dimH = rewriter.
createOrFold<tensor::DimOp>(loc, input, 1);
2829 auto dimW = rewriter.
createOrFold<tensor::DimOp>(loc, input, 2);
2832 auto zeroFloat = arith::ConstantOp::create(
2833 rewriter, loc, rewriter.
getZeroAttr(elementType));
2834 auto twoPiAttr = rewriter.
getFloatAttr(elementType, 6.283185307179586);
2835 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
2840 auto constH = castIndexToFloat(rewriter, loc, elementType, dimH);
2841 auto constW = castIndexToFloat(rewriter, loc, elementType, dimW);
2842 auto halfH = index::DivUOp::create(rewriter, loc, dimH, twoIndex);
2843 auto halfW = index::DivUOp::create(rewriter, loc, dimW, twoIndex);
2845 auto buildBody = [&](OpBuilder &builder, Location loc,
ValueRange args) {
2846 Value valReal = args[0];
2847 Value sumReal = args[1];
2848 Value sumImag = args[2];
2851 Value oy = linalg::IndexOp::create(builder, loc, 1);
2852 Value ox = linalg::IndexOp::create(builder, loc, 2);
2853 Value iy = linalg::IndexOp::create(builder, loc, 3);
2854 Value ix = linalg::IndexOp::create(builder, loc, 4);
2859 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
2860 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
2862 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
2863 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
2865 auto iyRemFloat = castIndexToFloat(builder, loc, elementType, iyRem);
2866 auto ixRemFloat = castIndexToFloat(builder, loc, elementType, ixRem);
2868 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
2869 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
2870 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
2871 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
2878 auto iyIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2880 auto iyIsHalfH = arith::CmpIOp::create(
2881 builder, loc, arith::CmpIPredicate::eq, iyRem, halfH);
2882 auto ixIs0 = arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq,
2884 auto ixIsHalfW = arith::CmpIOp::create(
2885 builder, loc, arith::CmpIPredicate::eq, ixRem, halfW);
2887 auto iyIsSinSkippable =
2888 arith::OrIOp::create(builder, loc, iyIs0, iyIsHalfH);
2889 auto ixIsSinSkippable =
2890 arith::OrIOp::create(builder, loc, ixIs0, ixIsHalfW);
2891 auto shouldSkipSin = arith::AndIOp::create(builder, loc, iyIsSinSkippable,
2896 auto cosAngle = math::CosOp::create(builder, loc, angle);
2897 auto sinAngle = math::SinOp::create(builder, loc, angle);
2898 auto imagWeight = arith::SelectOp::create(builder, loc, shouldSkipSin,
2899 zeroFloat, sinAngle);
2900 auto realComponent =
2901 arith::MulFOp::create(builder, loc, valReal, cosAngle);
2902 auto imagComponent =
2903 arith::MulFOp::create(builder, loc, valReal, imagWeight);
2908 arith::AddFOp::create(builder, loc, sumReal, realComponent);
2910 arith::SubFOp::create(builder, loc, sumImag, imagComponent);
2912 linalg::YieldOp::create(builder, loc,
ValueRange{outReal, outImag});
2916 rfft2d, rfft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
2917 indexingMaps, iteratorTypes, buildBody);
2926 LogicalResult matchAndRewrite(FFT2dOp fft2d,
2927 PatternRewriter &rewriter)
const override {
2928 if (!llvm::all_of(fft2d->getOperandTypes(),
2929 RFFT2dConverter::isRankedTensor) ||
2930 !llvm::all_of(fft2d->getResultTypes(),
2931 RFFT2dConverter::isRankedTensor)) {
2935 Location loc = fft2d.getLoc();
2936 Value input_real = fft2d.getInputReal();
2937 Value input_imag = fft2d.getInputImag();
2938 BoolAttr inverse = fft2d.getInverseAttr();
2940 auto real_el_ty = cast<FloatType>(
2941 cast<ShapedType>(input_real.
getType()).getElementType());
2942 [[maybe_unused]]
auto imag_el_ty = cast<FloatType>(
2943 cast<ShapedType>(input_imag.
getType()).getElementType());
2945 assert(real_el_ty == imag_el_ty);
2948 SmallVector<Value> dynamicSizes;
2953 SmallVector<int64_t, 3> staticSizes;
2956 auto outputType = RankedTensorType::get(staticSizes, real_el_ty);
2959 SmallVector<utils::IteratorType, 5> iteratorTypes = {
2960 utils::IteratorType::parallel, utils::IteratorType::parallel,
2961 utils::IteratorType::parallel, utils::IteratorType::reduction,
2962 utils::IteratorType::reduction};
2965 SmallVector<Value> genericOpInputs = {input_real, input_imag};
2966 SmallVector<Value> genericOpOutputs = {
2967 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
2969 RFFT2dConverter::createZeroTensor(rewriter, loc, outputType,
2974 ArrayRef{RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
2975 RFFT2dConverter::affineDimsExpr(rewriter, 0, 3, 4),
2976 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2),
2977 RFFT2dConverter::affineDimsExpr(rewriter, 0, 1, 2)},
2981 auto dimH = rewriter.
createOrFold<tensor::DimOp>(loc, input_real, 1);
2982 auto dimW = rewriter.
createOrFold<tensor::DimOp>(loc, input_real, 2);
2985 auto twoPiAttr = rewriter.
getFloatAttr(real_el_ty, 6.283185307179586);
2986 auto twoPi = arith::ConstantOp::create(rewriter, loc, twoPiAttr);
2988 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimH);
2990 RFFT2dConverter::castIndexToFloat(rewriter, loc, real_el_ty, dimW);
2992 auto buildBody = [&](OpBuilder &builder, Location loc,
ValueRange args) {
2993 Value valReal = args[0];
2994 Value valImag = args[1];
2995 Value sumReal = args[2];
2996 Value sumImag = args[3];
2999 Value oy = linalg::IndexOp::create(builder, loc, 1);
3000 Value ox = linalg::IndexOp::create(builder, loc, 2);
3001 Value iy = linalg::IndexOp::create(builder, loc, 3);
3002 Value ix = linalg::IndexOp::create(builder, loc, 4);
3006 auto iyXoy = index::MulOp::create(builder, loc, iy, oy);
3007 auto ixXox = index::MulOp::create(builder, loc, ix, ox);
3009 auto iyRem = index::RemUOp::create(builder, loc, iyXoy, dimH);
3010 auto ixRem = index::RemUOp::create(builder, loc, ixXox, dimW);
3013 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, iyRem);
3015 RFFT2dConverter::castIndexToFloat(builder, loc, real_el_ty, ixRem);
3017 auto yComponent = arith::DivFOp::create(builder, loc, iyRemFloat, constH);
3018 auto xComponent = arith::DivFOp::create(builder, loc, ixRemFloat, constW);
3020 auto sumXY = arith::AddFOp::create(builder, loc, yComponent, xComponent);
3021 auto angle = arith::MulFOp::create(builder, loc, twoPi, sumXY);
3024 angle = arith::MulFOp::create(
3025 builder, loc, angle,
3026 arith::ConstantOp::create(rewriter, loc,
3032 auto cosAngle = math::CosOp::create(builder, loc, angle);
3033 auto sinAngle = math::SinOp::create(builder, loc, angle);
3035 auto rcos = arith::MulFOp::create(builder, loc, valReal, cosAngle);
3036 auto rsin = arith::MulFOp::create(builder, loc, valImag, sinAngle);
3037 auto realComponent = arith::AddFOp::create(builder, loc, rcos, rsin);
3039 auto icos = arith::MulFOp::create(builder, loc, valImag, cosAngle);
3040 auto isin = arith::MulFOp::create(builder, loc, valReal, sinAngle);
3042 auto imagComponent = arith::SubFOp::create(builder, loc, icos, isin);
3047 arith::AddFOp::create(builder, loc, sumReal, realComponent);
3049 arith::AddFOp::create(builder, loc, sumImag, imagComponent);
3051 linalg::YieldOp::create(builder, loc,
ValueRange{outReal, outImag});
3055 fft2d, fft2d.getResultTypes(), genericOpInputs, genericOpOutputs,
3056 indexingMaps, iteratorTypes, buildBody);
3068 patterns->
add<GenericResizeConverter>(patterns->
getContext(),
3072 patterns->
add<MaterializeResizeBroadcast>(patterns->
getContext(),
3077 PointwiseConverter<tosa::AddOp>,
3078 PointwiseConverter<tosa::SubOp>,
3079 PointwiseConverter<tosa::MulOp>,
3080 PointwiseConverter<tosa::IntDivOp>,
3081 PointwiseConverter<tosa::NegateOp>,
3082 PointwiseConverter<tosa::PowOp>,
3083 PointwiseConverter<tosa::ReciprocalOp>,
3084 PointwiseConverter<tosa::RsqrtOp>,
3085 PointwiseConverter<tosa::LogOp>,
3086 PointwiseConverter<tosa::ExpOp>,
3087 PointwiseConverter<tosa::AbsOp>,
3088 PointwiseConverter<tosa::SinOp>,
3089 PointwiseConverter<tosa::CosOp>,
3090 PointwiseConverter<tosa::TanhOp>,
3091 PointwiseConverter<tosa::ErfOp>,
3092 PointwiseConverter<tosa::BitwiseAndOp>,
3093 PointwiseConverter<tosa::BitwiseOrOp>,
3094 PointwiseConverter<tosa::BitwiseNotOp>,
3095 PointwiseConverter<tosa::BitwiseXorOp>,
3096 PointwiseConverter<tosa::LogicalAndOp>,
3097 PointwiseConverter<tosa::LogicalNotOp>,
3098 PointwiseConverter<tosa::LogicalOrOp>,
3099 PointwiseConverter<tosa::LogicalXorOp>,
3100 PointwiseConverter<tosa::CastOp>,
3101 PointwiseConverter<tosa::LogicalLeftShiftOp>,
3102 PointwiseConverter<tosa::LogicalRightShiftOp>,
3103 PointwiseConverter<tosa::ArithmeticRightShiftOp>,
3104 PointwiseConverter<tosa::ClzOp>,
3105 PointwiseConverter<tosa::SelectOp>,
3106 PointwiseConverter<tosa::GreaterOp>,
3107 PointwiseConverter<tosa::GreaterEqualOp>,
3108 PointwiseConverter<tosa::EqualOp>,
3109 PointwiseConverter<tosa::MaximumOp>,
3110 PointwiseConverter<tosa::MinimumOp>,
3111 PointwiseConverter<tosa::CeilOp>,
3112 PointwiseConverter<tosa::FloorOp>,
3113 PointwiseConverter<tosa::ClampOp>,
3114 PointwiseConverter<tosa::SigmoidOp>
3118 IdentityNConverter<tosa::IdentityOp>,
3119 ReduceConverter<tosa::ReduceAllOp>,
3120 ReduceConverter<tosa::ReduceAnyOp>,
3121 ReduceConverter<tosa::ReduceMinOp>,
3122 ReduceConverter<tosa::ReduceMaxOp>,
3123 ReduceConverter<tosa::ReduceSumOp>,
3124 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 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 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 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 TypedAttr createInitialValueForReduceOp(Operation *op, Type elementTy, PatternRewriter &rewriter)
static LogicalResult reduceMatchAndRewriteHelper(OpTy op, uint64_t axis, PatternRewriter &rewriter)
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 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)
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...
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)
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)
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter, RewritePatternSet *patterns)
Populates conversion passes from TOSA dialect to Linalg dialect.
std::optional< SmallVector< Value > > checkHasDynamicBatchDims(PatternRewriter &rewriter, Op op, ArrayRef< Value > params)
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...