22#include "llvm/ADT/STLExtras.h"
23#include "llvm/ADT/SetOperations.h"
24#include "llvm/ADT/SmallBitVector.h"
25#include "llvm/ADT/SmallVector.h"
26#include "llvm/Support/Casting.h"
27#include "llvm/Support/raw_ostream.h"
34#include "mlir/Dialect/Linalg/IR/LinalgInterfaces.cpp.inc"
43 for (
auto &opOperand : linalgOp->getOpOperands()) {
44 if (llvm::is_contained(droppedOperands, &opOperand))
46 indexingMaps.push_back(linalgOp.getMatchingIndexingMap(&opOperand));
48 if (indexingMaps.empty()) {
51 return linalgOp.getNumLoops() == 0;
54 indexingMaps, linalgOp.getContext())) !=
AffineMap();
63 if (!op.isAllParallelLoops() || !op.isSingleInputOutput())
66 auto mapRange = op.getIndexingMapsArray();
67 if (mapRange.size() != 2 || !mapRange.front().isIdentity() ||
68 !mapRange.back().isIdentity()) {
72 Block *body = op.getBlock();
75 auto yieldOp = dyn_cast<linalg::YieldOp>(body->
back());
76 if (!yieldOp || yieldOp.getNumOperands() != 1)
78 return yieldOp->getOperand(0) == body->
getArgument(0);
88 if (!op.isAllParallelLoops() || op.getNumDpsInits() != 1 ||
89 op.getNumDpsInputs() != 0)
93 if (op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
96 Block *body = op.getBody();
100 auto yieldOp = dyn_cast<linalg::YieldOp>(body->
back());
101 if (!yieldOp || yieldOp.getNumOperands() != 1)
104 Value yieldOperand = yieldOp->getOperand(0);
116 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
117 !op.isSingleYieldOp())
121 if (!op.payloadUsesValueFromOperand(op.getDpsInputOperand(0)) ||
122 op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
125 OpOperand *value = op.getDpsInputOperand(0);
126 if (!op.isScalar(value))
140std::optional<SmallVector<int64_t>>
142 if (
auto broadcastOp = dyn_cast<BroadcastOp>(linalgOp.getOperation()))
144 broadcastOp.getDimensions().end());
146 auto op = dyn_cast<GenericOp>(linalgOp.getOperation());
151 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
152 !op.isSingleYieldOp())
155 auto srcTy = op.getDpsInputOperand(0)->get().getType();
156 auto dstTy = op.getDpsInitOperand(0)->get().getType();
157 if (!isa<MemRefType, RankedTensorType>(srcTy) ||
158 !isa<MemRefType, RankedTensorType>(dstTy))
164 auto dstMap = op.getIndexingMapsArray()[1];
165 if (!dstMap.isIdentity())
169 auto srcMap = op.getIndexingMapsArray()[0];
171 if (srcMap.getResults().size() >= dstMap.getResults().size())
175 for (
unsigned i = 0; i < srcMap.getNumResults(); ++i) {
176 auto expr = llvm::dyn_cast<AffineDimExpr>(srcMap.getResults()[i]);
179 int64_t pos = expr.getPosition();
180 if (i > 0 && pos <= position[i - 1])
182 position.push_back(expr.getPosition());
186 auto numDims = srcMap.getNumDims();
188 for (
auto dim : llvm::seq<int64_t>(0, numDims)) {
189 if (!llvm::is_contained(position, dim))
190 broadcastedDims.push_back(dim);
192 return broadcastedDims;
198std::optional<SmallVector<int64_t>>
203 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
204 !op.isSingleYieldOp())
207 auto mapRange = op.getIndexingMapsArray();
208 if (mapRange.size() != 2)
211 auto mapOfInput = mapRange.front();
212 auto mapOfResult = mapRange.back();
216 if (!mapOfResult.isIdentity() || !mapOfInput.isPermutation())
220 for (
unsigned i = 0; i < mapOfInput.getNumDims(); ++i) {
221 auto expr = llvm::cast<AffineDimExpr>(mapOfInput.getResults()[i]);
222 permutation[expr.getPosition()] = i;
233 if (!op.isAllParallelLoops() || op.getNumLoops() < 1)
238 if (op.getNumDpsInputs() != arity || op.getNumDpsInits() != 1)
242 if (op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
249 Block *body = op.getBody();
261 auto yieldOp = dyn_cast<linalg::YieldOp>(body->
back());
262 return !(!yieldOp || yieldOp.getNumOperands() != 1 ||
263 yieldOp->getOperand(0).getDefiningOp() != oper);
272 if (!op.payloadUsesValueFromOperand(op.getDpsInputOperand(0)))
283 OpOperand *inputOpOperand0 = op.getDpsInputOperand(0);
284 OpOperand *inputOpOperand1 = op.getDpsInputOperand(1);
285 return !(!op.payloadUsesValueFromOperand(inputOpOperand0) ||
286 !op.payloadUsesValueFromOperand(inputOpOperand1));
303 OpOperand *inputOpOperand0 = op.getDpsInputOperand(0);
304 OpOperand *inputOpOperand1 = op.getDpsInputOperand(1);
305 OpOperand *inputOpOperand2 = op.getDpsInputOperand(2);
306 return !(!op.payloadUsesValueFromOperand(inputOpOperand0) ||
307 !op.payloadUsesValueFromOperand(inputOpOperand1) ||
308 !op.payloadUsesValueFromOperand(inputOpOperand2));
322 auto iface = dyn_cast<MemoryEffectOpInterface>(op);
323 if (!iface || !iface.hasNoEffect())
333 llvm::raw_ostream &errs) {
335 errs <<
"no terminator in the block";
340 errs <<
"expected block with 3 arguments";
346 errs <<
"expected terminator with 1 operand";
354 errs <<
"expected reduction op to be binary";
363 errs <<
"expected reduction to take block argument #2 as one of the "
364 "operands (modulo unary casts)";
369 isa<BlockArgument>(reductionLHS) ? reductionRHS : reductionLHS);
373 errs <<
"expected elementwise op to be binary";
377 if (!isaPair(elementwiseOp, reductionOp)) {
378 errs <<
"expected reduction/elementwise op kind not satisfied";
391 errs <<
"expected elementwise op to apply to block arguments (modulo unary "
398template <
typename AddOpTy,
typename MulOpTy,
typename... Args>
400 static_assert(
sizeof...(Args) % 2 == 0,
401 "expected an even number of template arguments");
402 if (isa<AddOpTy>(
add) && isa<MulOpTy>(
mul))
405 if constexpr (
sizeof...(Args) > 0)
413template <
typename... Args>
425static llvm::SmallDenseSet<int64_t>
428 utils::IteratorType iter) {
429 assert(iterators.size() == indexingMap.
getNumDims());
430 llvm::SmallDenseSet<int64_t> res;
432 if (
auto d = dyn_cast<AffineDimExpr>(e)) {
433 if (iterators[d.getPosition()] == iter &&
435 return e.isFunctionOfDim(d.getPosition());
437 res.insert(d.getPosition());
444auto par = utils::IteratorType::parallel;
445auto red = utils::IteratorType::reduction;
452static FailureOr<SmallVector<utils::IteratorType>>
458 if (
auto dim = dyn_cast<AffineDimExpr>(expr))
459 iterators[dim.getPosition()] = par;
474static FailureOr<ContractionDimensions>
477 llvm::SmallDenseSet<int64_t> a =
479 llvm::SmallDenseSet<int64_t>
b =
481 llvm::SmallDenseSet<int64_t> c =
485 llvm::SmallDenseSet<int64_t> ac = a;
486 llvm::set_intersect(ac, c);
487 llvm::set_subtract(ac,
b);
489 llvm::SmallDenseSet<int64_t> bc =
b;
490 llvm::set_intersect(bc, c);
491 llvm::set_subtract(bc, a);
493 llvm::SmallDenseSet<int64_t> batches = a;
494 llvm::set_intersect(batches,
b);
495 llvm::set_intersect(batches, c);
498 llvm::SmallDenseSet<int64_t> ra =
500 llvm::SmallDenseSet<int64_t> rb =
502 llvm::set_intersect(ra, rb);
510 llvm::sort(dimensions.
batch);
511 llvm::sort(dimensions.
m);
512 llvm::sort(dimensions.
n);
513 llvm::sort(dimensions.
k);
517FailureOr<ContractionDimensions>
519 if (linalgOp.getNumDpsInits() != 1 || linalgOp.getNumDpsInputs() != 2)
522 linalgOp.getIteratorTypesArray());
525FailureOr<ContractionDimensions>
527 if (indexingMaps.size() != 3)
530 if (failed(iterators))
549 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
552 if (linalgOp.getNumDpsInputs() != 2 || linalgOp.getNumDpsInits() != 1)
554 auto mapRange = linalgOp.getIndexingMapsArray();
555 if (linalgOp.getNumReductionLoops() == 0)
557 if (llvm::any_of(mapRange,
563 arith::MulFOp, arith::AddFOp,
564 arith::MulIOp, arith::AddIOp,
565 complex::MulOp, complex::AddOp,
566 arith::AndIOp, arith::OrIOp>(
567 *linalgOp.getBlock())) {
574 assert(succeeded(res) &&
"unexpected failure to infer contraction dims");
584 return "expected a LinalgOp";
586 return "expected op with 2 inputs and 1 output";
588 return "expected at least 1 reduction";
590 return "expected indexing maps to be projected permutations";
592 return "expected add/mul op in the body";
596 llvm_unreachable(
"unhandled MatchContractionResult case");
603 return isa<ContractionOpInterface>(op) ||
636 return isa<T>(lhs) ? cast<T>(lhs) : (isa<T>(rhs) ? cast<T>(rhs) :
nullptr);
648struct ConvAccessExprWalker
651 llvm::SmallDenseSet<int64_t> convolvedDims;
653 llvm::SmallDenseMap<int64_t, int64_t> convolvedDimMapping;
655 llvm::SmallDenseSet<int64_t> unConvolvedDims;
657 llvm::SmallDenseMap<int64_t, AffineExpr> strideAndDilationMapping;
661 void clearMultiUseDims(AffineMap map) {
662 for (
int dimPos = 0, e = map.
getNumDims(); dimPos < e; ++dimPos) {
663 if (llvm::count_if(map.
getResults(), [dimPos](AffineExpr e) {
664 return e.isFunctionOfDim(dimPos);
666 convolvedDims.erase(dimPos);
667 unConvolvedDims.erase(dimPos);
670 auto it = convolvedDimMapping.find(dimPos);
671 if (it != convolvedDimMapping.end()) {
672 int64_t pairedDim = it->second;
673 convolvedDims.erase(pairedDim);
674 unConvolvedDims.erase(pairedDim);
675 strideAndDilationMapping.erase(pairedDim);
676 convolvedDimMapping.erase(dimPos);
677 convolvedDimMapping.erase(pairedDim);
683 LogicalResult visitDimExpr(AffineDimExpr dimExpr) {
685 if (unConvolvedDims.count(position) || convolvedDims.count(position)) {
688 unConvolvedDims.insert(position);
692 LogicalResult visitSymbolExpr(AffineSymbolExpr expr) {
return failure(); }
694 LogicalResult visitConstantExpr(AffineConstantExpr expr) {
return failure(); }
696 LogicalResult visitAffineBinaryOpExpr(AffineBinaryOpExpr binaryExpr) {
698 if (binaryExpr.
getKind() != AffineExprKind::Add)
700 auto lhsDimPos = getDimExprOrMulExprDimPos(binaryExpr.
getLHS());
701 auto rhsDimPos = getDimExprOrMulExprDimPos(binaryExpr.
getRHS());
704 convolvedDimMapping[*lhsDimPos] = *rhsDimPos;
705 convolvedDimMapping[*rhsDimPos] = *lhsDimPos;
709 FailureOr<int64_t> getDimExprOrMulExprDimPos(AffineExpr expr) {
710 if (
auto dimExpr = dyn_cast<AffineDimExpr>(expr)) {
712 if (convolvedDims.count(dim) || unConvolvedDims.count(dim))
715 strideAndDilationMapping[dim] =
717 convolvedDims.insert(dim);
720 if (
auto symbolMulExpr = dyn_cast<AffineBinaryOpExpr>(expr)) {
721 if (symbolMulExpr.getKind() != AffineExprKind::Mul)
723 auto lhsExpr = symbolMulExpr.getLHS();
724 auto rhsExpr = symbolMulExpr.getRHS();
733 if (!mulExpr || !dimExpr)
736 if (convolvedDims.count(dim) || unConvolvedDims.count(dim))
738 strideAndDilationMapping[dim] = mulExpr;
739 convolvedDims.insert(dim);
749 "expected map to have projected permutations");
750 llvm::SmallDenseSet<int64_t> preservedDims;
752 preservedDims.insert(cast<AffineDimExpr>(expr).getPosition());
753 return preservedDims;
759 for (
auto e : exprs) {
760 auto constantExpr = dyn_cast<AffineConstantExpr>(e);
761 assert(constantExpr &&
"Found non-constant stride/dilation");
762 vals.push_back(constantExpr.getValue());
787 ConvAccessExprWalker &inputExprWalker,
bool allowEmptyConvolvedDims,
790 AffineMap outputMap = indexingMaps.back();
791 llvm::SmallDenseSet<int64_t> filterDims =
793 llvm::SmallDenseSet<int64_t> outputDims =
797 llvm::SmallDenseSet<int64_t> batch = inputExprWalker.unConvolvedDims;
798 llvm::set_intersect(batch, outputDims);
799 llvm::set_subtract(batch, filterDims);
802 llvm::SmallDenseSet<int64_t> oi = inputExprWalker.convolvedDims;
803 llvm::set_intersect(oi, outputDims);
806 llvm::SmallDenseSet<int64_t> oc = filterDims;
807 llvm::set_intersect(oc, outputDims);
808 llvm::set_subtract(oc, inputExprWalker.unConvolvedDims);
811 llvm::SmallDenseSet<int64_t> depth = filterDims;
812 llvm::set_intersect(depth, outputDims);
813 llvm::set_intersect(depth, inputExprWalker.unConvolvedDims);
815 llvm::SmallDenseSet<int64_t> filterReducedDims =
819 llvm::SmallDenseSet<int64_t> fl = inputExprWalker.convolvedDims;
820 llvm::set_intersect(fl, filterReducedDims);
823 llvm::SmallDenseSet<int64_t> ic = inputExprWalker.unConvolvedDims;
824 llvm::set_intersect(ic, filterReducedDims);
826 if (oi.empty() && !allowEmptyConvolvedDims)
840 llvm::sort(dimensions.
batch);
844 llvm::sort(dimensions.
depth);
850 dimensions.
filterLoop.push_back(inputExprWalker.convolvedDimMapping[oiDim]);
853 if (!nativeStrides) {
856 strideExprs.push_back(inputExprWalker.strideAndDilationMapping[oiDim]);
859 dimensions.
strides = llvm::to_vector<2>(nativeStrides.getValues<
int64_t>());
861 if (!nativeDilations) {
864 dilationExprs.push_back(inputExprWalker.strideAndDilationMapping[flDim]);
868 llvm::to_vector<2>(nativeDilations.getValues<
int64_t>());
875 return dyn_cast_or_null<DenseIntElementsAttr>(
876 linalgOp->getInherentAttr(name).value_or(
Attribute{}));
907FailureOr<ConvolutionDimensions>
909 if (linalgOp.getNumDpsInits() != 1 || linalgOp.getNumDpsInputs() != 2)
912 auto indexingMaps = linalgOp.getIndexingMapsArray();
915 ConvAccessExprWalker inputExprWalker;
916 for (
AffineExpr expr : indexingMaps[0].getResults())
917 (
void)inputExprWalker.visit(expr);
918 inputExprWalker.clearMultiUseDims(indexingMaps[0]);
921 indexingMaps, linalgOp.getIteratorTypesArray(), inputExprWalker,
927FailureOr<ConvolutionDimensions>
929 if (indexingMaps.size() != 3)
933 FailureOr<SmallVector<utils::IteratorType>> iterators =
935 if (failed(iterators))
939 ConvAccessExprWalker inputExprWalker;
940 for (
AffineExpr expr : indexingMaps[0].getResults())
941 (
void)inputExprWalker.visit(expr);
942 inputExprWalker.clearMultiUseDims(indexingMaps[0]);
968 bool allowEmptyConvolvedDims) {
969 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
972 if (linalgOp.getNumDpsInputs() < 2 || linalgOp.getNumDpsInits() != 1)
975 auto indexingMaps = linalgOp.getIndexingMapsArray();
978 ConvAccessExprWalker inputExprWalker;
979 if (llvm::any_of(indexingMaps[0].getResults(),
981 return failed(inputExprWalker.visit(expr));
987 if (!indexingMaps[1].isProjectedPermutation() ||
988 !indexingMaps.back().isProjectedPermutation())
991 auto iteratorTypes = linalgOp.getIteratorTypesArray();
993 llvm::SmallDenseSet<int64_t> outputDims =
995 llvm::SmallDenseSet<int64_t> filterDims =
getPreservedDims(indexingMaps[1]);
1009 llvm::SmallDenseSet<int64_t> allLoopDims;
1010 for (
auto outputExpr : indexingMaps.back().getResults()) {
1011 int64_t outputDim = cast<AffineDimExpr>(outputExpr).getPosition();
1012 if (inputExprWalker.unConvolvedDims.count(outputDim) &&
1013 !filterDims.count(outputDim)) {
1015 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1017 allLoopDims.insert(outputDim);
1020 if (inputExprWalker.convolvedDims.count(outputDim) &&
1021 !filterDims.count(outputDim)) {
1023 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1025 allLoopDims.insert(outputDim);
1028 if (!inputExprWalker.convolvedDims.count(outputDim) &&
1029 !inputExprWalker.unConvolvedDims.count(outputDim) &&
1030 filterDims.count(outputDim)) {
1032 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1034 allLoopDims.insert(outputDim);
1037 if (inputExprWalker.unConvolvedDims.count(outputDim) &&
1038 filterDims.count(outputDim)) {
1040 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1042 allLoopDims.insert(outputDim);
1047 for (
auto filterExpr : indexingMaps[1].getResults()) {
1048 int64_t filterDim = cast<AffineDimExpr>(filterExpr).getPosition();
1049 if (outputDims.count(filterDim) &&
1050 !inputExprWalker.unConvolvedDims.count(filterDim) &&
1051 !inputExprWalker.convolvedDims.count(filterDim)) {
1055 if (inputExprWalker.convolvedDims.count(filterDim) &&
1056 !outputDims.count(filterDim)) {
1058 if (iteratorTypes[filterDim] != utils::IteratorType::reduction)
1060 if (allLoopDims.count(filterDim))
1062 allLoopDims.insert(filterDim);
1065 if (inputExprWalker.unConvolvedDims.count(filterDim) &&
1066 !outputDims.count(filterDim)) {
1068 if (iteratorTypes[filterDim] != utils::IteratorType::reduction)
1070 if (allLoopDims.count(filterDim))
1072 allLoopDims.insert(filterDim);
1075 if (inputExprWalker.unConvolvedDims.count(filterDim) &&
1076 outputDims.count(filterDim)) {
1083 if (allLoopDims.size() != linalgOp.getNumLoops())
1086 if (!allowEmptyConvolvedDims && inputExprWalker.convolvedDims.empty())
1091 indexingMaps, iteratorTypes, inputExprWalker, allowEmptyConvolvedDims,
1094 assert(succeeded(res) &&
"unexpected failure to infer convolution dims");
1105 return "expected a LinalgOp";
1107 return "expected op with 2 inputs and 1 output";
1109 return "unexpected input index map for convolutions";
1111 return "expected output/filter indexing maps to be projected permutations";
1113 return "unexpected loop dimension for convolution op";
1115 return "expected all iterators used to access outputs to be parallel";
1117 return "expected all iterators not used to access outputs to be reduction";
1119 return "expected convolved dim to be non-empty";
1123 llvm_unreachable(
"unhandled MatchConvolutionResult case");
1127 bool allowEmptyConvolvedDims) {
1129 linalgOp.getOperation(),
nullptr, allowEmptyConvolvedDims) ==
1145enum class MatchFillResult {
1155 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
1157 return MatchFillResult::NotLinalgOp;
1158 if (linalgOp.getNumDpsInputs() != 1 || linalgOp.getNumDpsInits() != 1)
1159 return MatchFillResult::WrongNumOperands;
1161 OpOperand *value = linalgOp.getDpsInputOperand(0);
1162 if (!linalgOp.isScalar(value))
1163 return MatchFillResult::NotScalarInput;
1166 OpOperand *output = linalgOp.getDpsInitOperand(0);
1169 if (scalarType != outputElementType)
1170 return MatchFillResult::TypeMismatch;
1172 return MatchFillResult::Success;
1177 if (res == MatchFillResult::NotLinalgOp)
1178 return op->
emitError(
"expected a LinalgOp");
1179 if (res == MatchFillResult::WrongNumOperands)
1180 return op->
emitError(
"expected op with 1 input and 1 output");
1181 if (res == MatchFillResult::NotScalarInput)
1182 return op->
emitError(
"expected op with scalar input");
1183 if (res == MatchFillResult::TypeMismatch) {
1184 auto linalgOp = cast<linalg::LinalgOp>(op);
1185 Type scalarType = linalgOp.getDpsInputOperand(0)->get().getType();
1186 Type outputElementType =
1188 return op->
emitOpError(
"expected fill value type (")
1189 << scalarType <<
") to match output element type ("
1190 << outputElementType <<
")";
1203 for (
OpOperand &opOperand : getOperation()->getOpOperands()) {
1204 for (
int64_t i = 0, e = getRank(&opOperand); i < e; ++i)
1212 assert(!hasDynamicShape() &&
"expected operands to have static shapes");
1213 for (
OpOperand &opOperand : getOperation()->getOpOperands())
1214 llvm::append_range(res,
getShape(&opOperand));
1219 AffineMap map = getLoopsToShapesMap();
1221 auto viewSizes = createFlatListOfOperandDims(
b, loc);
1222 SmallVector<Range, 4> res(numDims);
1223 for (
unsigned idx = 0; idx < numRes; ++idx) {
1225 if (
auto d = dyn_cast<AffineDimExpr>(
result)) {
1226 if (res[d.getPosition()].offset)
1228 res[d.getPosition()] =
1229 Range{
b.getIndexAttr(0), viewSizes[idx],
b.getIndexAttr(1)};
1240 : positions(std::move(positions)) {}
1255 llvm::SmallBitVector positions;
1258static std::pair<int64_t, int64_t>
1262 for (
OpOperand *input : op.getDpsInputOperands())
1263 inputRankSum += op.getRank(input);
1264 for (
OpOperand &output : op.getDpsInitsMutable())
1265 outputRankSum += op.getRank(&output);
1266 return {inputRankSum, inputRankSum + outputRankSum};
1281 AffineMap loopsToShapesMap = getLoopsToShapesMap();
1289 AffineMap loopToResultsShapeMap = loopsToShapesMap.
getSliceMap(
1290 resultShapesSubMapPos.first,
1291 resultShapesSubMapPos.second - resultShapesSubMapPos.first);
1292 AffineMap resultShapesFromInputShapesMap =
1293 loopToResultsShapeMap.
compose(getShapesToLoopsMap());
1297 llvm::SmallBitVector outputDims(resultShapesFromInputShapesMap.
getNumDims());
1298 outputDims.set(resultShapesSubMapPos.first, resultShapesSubMapPos.second);
1299 HasAffineDimExprVisitor checkDimExpr(std::move(outputDims));
1300 Location loc = getOperation()->getLoc();
1301 IRRewriter rewriter(
b);
1302 SmallVector<OpFoldResult> allResultDimValues =
1303 affine::makeComposedFoldedMultiResultAffineApply(
1304 rewriter, loc, resultShapesFromInputShapesMap,
1305 createFlatListOfOperandDims(
b, loc));
1307 ArrayRef<AffineExpr> shapeExprs = resultShapesFromInputShapesMap.
getResults();
1308 for (OpOperand &opOperand : getDpsInitsMutable()) {
1309 SmallVector<OpFoldResult> shapes;
1310 for (int64_t dim : llvm::seq<int64_t>(0, getRank(&opOperand))) {
1311 auto shapedType = llvm::cast<ShapedType>(opOperand.get().getType());
1312 if (!shapedType.isDynamicDim(dim)) {
1314 shapes.push_back(
b.getIndexAttr(shapedType.getDimSize(dim)));
1317 OpFoldResult ofr = checkDimExpr.visit(shapeExprs[pos])
1319 : allResultDimValues[pos];
1324 reifiedReturnShapes.emplace_back(std::move(shapes));
1333 auto dpsIface = cast<DestinationStyleOpInterface>(*this->getOperation());
1334 if (!dpsIface.isDpsInput(opOperand))
1335 return operandNumber;
1336 unsigned start = dpsIface.getDpsInits().getBeginOperandIndex();
1337 assert(!dpsIface.isDpsInit(opOperand));
1340 return cast<DestinationStyleOpInterface>(*this->getOperation())
1341 .getNumDpsInputs() +
1342 operandNumber - start;
1346 LinalgOp linalgOp = cast<LinalgOp>(op);
1348 if (!linalgOp.hasPureTensorSemantics() &&
1350 return op->
emitOpError(
"expected to have pure tensor or buffer semantics");
1354 if (linalgOp.hasDynamicIndexingMaps())
1355 if (failed(linalgOp.verifyIndexingMapRequiredAttributes()))
1359 if (failed(cast<IndexingMapOpInterface>(op).verifyImpl()))
1364 for (
OpOperand &opOperand : linalgOp->getOpOperands()) {
1365 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(&opOperand);
1367 unsigned numLoops = linalgOp.getNumLoops();
1371 <<
" dim(s) to match the number of loops";
1374 linalgOp.getReductionDims(redDims);
1376 if (!linalgOp.getShapesToLoopsMap())
1377 return op->
emitOpError(
"expected the shape-to-loops map to be non-null");
1380 if (linalgOp->getNumRegions() != 1 || !linalgOp->getRegion(0).hasOneBlock())
1381 return op->
emitOpError(
"expects to have 1 region with 1 block");
1389 Block &block = linalgOp->getRegion(0).front();
1391 if (linalgOp.getOpOperandsMatchingBBargs().size() != block.
getNumArguments())
1392 return op->
emitOpError(
"expected as many non-induction variable region "
1393 "arguments as the number of input/output operands");
1395 for (
OpOperand *opOperand : linalgOp.getOpOperandsMatchingBBargs()) {
1397 if (isa<MemRefType, RankedTensorType>(elementType))
1400 if (elementType != argType)
1401 return op->
emitOpError(
"expected type of bb argument #")
1403 <<
" to match element or self type of the corresponding operand ("
1404 << elementType <<
")";
static FailureOr< ContractionDimensions > inferContractionDimsImpl(ArrayRef< AffineMap > indexingMaps, ArrayRef< utils::IteratorType > iterators)
Find 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcomputation ...
static Value getSourceSkipUnary(Value value)
If the value is defined by a chain of unary side effect-free, go up the use-def chain until the first...
static llvm::SmallDenseSet< int64_t > getPreservedDims(AffineMap map)
static T getAffineExprOfType(AffineExpr lhs, AffineExpr rhs)
Of the given two expressions returns one that is of type T (lhs gets preference over rhs)
static FailureOr< ConvolutionDimensions > inferConvolutionDimsImpl(ArrayRef< AffineMap > indexingMaps, ArrayRef< utils::IteratorType > iterators, ConvAccessExprWalker &inputExprWalker, bool allowEmptyConvolvedDims, DenseIntElementsAttr nativeStrides, DenseIntElementsAttr nativeDilations)
Classifies dimensions in the indexingMaps used by a convolution subcomputation, as captured by inputE...
static std::pair< int64_t, int64_t > getResultsPositionInLoopsToShapeMap(LinalgOp &op)
static bool isPairTemplateImpl(Operation *add, Operation *mul)
Returns true if the two operations are of the kinds specified by a pair of consecutive template argum...
static bool isaElemwiseSingleOpInterface(linalg::GenericOp op, unsigned arity)
static MatchFillResult isFillInterfaceImpl(Operation *op)
static bool isContractionBody(Block &block)
Returns true if the block is a body of a contraction with the kinds of operations given pairwise by t...
static std::optional< Value > isaExternalFillOp(GenericOp op)
Detects if a linalg.generic operation represents an external scalar input.
static FailureOr< SmallVector< utils::IteratorType > > inferIteratorsFromOutMap(AffineMap map)
Infer the iterator types from the init affine map.
static DenseIntElementsAttr getInherentConvolutionAttr(LinalgOp linalgOp, StringRef name)
static llvm::SmallDenseSet< int64_t > findPermutationsIndexingOperand(AffineMap indexingMap, ArrayRef< utils::IteratorType > iterators, utils::IteratorType iter)
Given an indexingMap and its corresponding iterators, returns the positions of the iterators of type ...
static SmallVector< int64_t, 2 > getConstantsFromExprList(const SmallVector< AffineExpr, 2 > &exprs)
static std::optional< Value > isaInlinedFillOp(GenericOp op)
Detects if a linalg.generic operation represents a fill with an inlined constant.
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Affine binary operation expression.
AffineExpr getLHS() const
AffineExpr getRHS() const
An integer constant appearing in affine expression.
A dimensional identifier appearing in an affine expression.
unsigned getPosition() const
bool visit(AffineExpr expr)
See documentation for AffineExprVisitorBase.
Base type for affine expression.
AffineExprKind getKind() const
Return the classification for this type.
MLIRContext * getContext() const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
AffineMap getSliceMap(unsigned start, unsigned length) const
Returns the map consisting of length expressions starting from start.
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
A symbolic identifier appearing in an affine expression.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
OpListType & getOperations()
Operation * getTerminator()
Get the terminator operation of this block.
An attribute that represents a reference to a dense integer vector or tensor object.
IRValueT get() const
Return the current value being used by this operand.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
This class provides the API for ops that are known to be terminators.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool mightHaveTrait()
Returns true if the operation might have the provided trait.
unsigned getNumOperands()
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
MatchConvolutionResult isConvolutionInterfaceImpl(Operation *op, ConvolutionDimensions *dimensions=nullptr, bool allowEmptyConvolvedDims=false)
Checks whether op conforms to ConvolutionOpInterface and populates dimensions with indexes of the dif...
@ NotProjectedPermutations
bool isContractionBody(Block &block, function_ref< bool(Operation *, Operation *)> isaPair, llvm::raw_ostream &errs=mlir::thread_safe_nulls())
Returns true if the block contains a contraction of the following form:
StringRef getMatchConvolutionMessage(MatchConvolutionResult res)
Returns the error message corresponding to the convolution checking return code.
bool canOpOperandsBeDroppedImpl(linalg::LinalgOp linalgOp, ArrayRef< OpOperand * > droppedOperands)
Implementation of the method that check if given operands can be dropped, i.e.
MatchContractionResult isContractionInterfaceImpl(Operation *op, ContractionDimensions *dimensions=nullptr)
Checks whether op conforms to ContractionOpInterface and populates dimensions with indexes of the dif...
LogicalResult verifyContractionInterface(Operation *op)
Verify that op conforms to ContractionOpInterface.
@ NotProjectedPermutations
@ NonOutputDimNotReduction
LogicalResult verifyFillInterface(Operation *op)
Verify that op conforms to the FillOpInterface.
StringRef getMatchContractionMessage(MatchContractionResult res)
Returns the error message corresponding to the contraction checking return code.
LogicalResult verifyStructuredOpInterface(Operation *op)
Verify that op conforms to the invariants of StructuredOpInterface.
LogicalResult verifyConvolutionInterface(Operation *op)
Verify that op conforms to the ConvolutionOpInterface.
std::optional< SmallVector< int64_t > > isaTransposeOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a linalg.transpose.
bool isaElemwiseSingleUnaryOpInterface(GenericOp genericOp)
Checks whether a given genericOp is semantically equivalent to a single linalg elementwise unary op,...
bool isaCopyOpInterface(LinalgOp linalgOp)
Checks whether linalgOp is semantically equivalent to a linalg.copyOp.
FailureOr< ConvolutionDimensions > inferConvolutionDims(LinalgOp linalgOp)
Find at least 1 parallel (output_image) and reduction (filter_loop) dimension candidates that form a ...
OpFoldResult createFoldedDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
bool isaConvolutionOpInterface(LinalgOp linalgOp, bool allowEmptyConvolvedDims=false)
Checks whether linalgOp conforms to ConvolutionOpInterface.
bool isaElemwiseSingleTernaryOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a single linalg elementwise ternary op e....
std::optional< SmallVector< int64_t > > isaBroadcastOpInterface(LinalgOp linalgOp)
Checks whether linalgOp is semantically equivalent to a broadcast operation.
FailureOr< ContractionDimensions > inferContractionDims(LinalgOp linalgOp)
Find at least 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcom...
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
bool isaContractionOpInterface(LinalgOp linalgOp)
Checks whether linalgOp conforms to ContractionOpInterface.
std::optional< Value > isaFillOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a linalg.fill.
bool isaElemwiseSingleBinaryOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a single linalg elementwise binary op e....
Include the generated interface declarations.
AffineMap concatAffineMaps(ArrayRef< AffineMap > maps, MLIRContext *context)
Concatenates a list of maps into a single AffineMap, stepping over potentially empty maps.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::function_ref< Fn > function_ref
HasAffineDimExprVisitor(llvm::SmallBitVector positions)
bool visitDimExpr(AffineDimExpr dimExpr)
bool visitAffineBinaryOpExpr(AffineBinaryOpExpr binaryOpExpr)
bool visitSymbolExpr(AffineSymbolExpr symbolExpr)
bool visitConstantExpr(AffineConstantExpr constExpr)
Positions of a Linalg op loops that correspond to different kinds of a contraction dimension.
SmallVector< unsigned, 2 > batch
SmallVector< unsigned, 2 > m
SmallVector< unsigned, 2 > n
SmallVector< unsigned, 2 > k
Positions of a Linalg op loops that correspond to different kinds of a convolution dimension.
SmallVector< unsigned, 2 > depth
SmallVector< unsigned, 2 > outputImage
SmallVector< unsigned, 2 > outputChannel
SmallVector< int64_t, 2 > dilations
SmallVector< int64_t, 2 > strides
SmallVector< unsigned, 2 > inputChannel
SmallVector< unsigned, 2 > batch
SmallVector< unsigned, 2 > filterLoop