25#include "llvm/ADT/STLExtras.h"
26#include "llvm/ADT/SmallBitVector.h"
27#include "llvm/ADT/SmallVectorExtras.h"
28#include "llvm/ADT/TypeSwitch.h"
29#include "llvm/Support/DebugLog.h"
30#include "llvm/Support/LogicalResult.h"
31#include "llvm/Support/MathExtras.h"
38using llvm::divideCeilSigned;
39using llvm::divideFloorSigned;
42#define DEBUG_TYPE "affine-ops"
44#include "mlir/Dialect/Affine/IR/AffineOpsDialect.cpp.inc"
51 if (
auto arg = dyn_cast<BlockArgument>(value))
52 return arg.getParentRegion() == region;
75 if (llvm::isa<BlockArgument>(value))
76 return legalityCheck(mapping.
lookup(value), dest);
83 bool isDimLikeOp = isa<ShapedDimOpInterface>(value.
getDefiningOp());
94 return llvm::all_of(values, [&](
Value v) {
101template <
typename OpTy>
104 static_assert(llvm::is_one_of<OpTy, AffineReadOpInterface,
105 AffineWriteOpInterface>::value,
106 "only ops with affine read/write interface are supported");
113 dimOperands, src, dest, mapping,
117 symbolOperands, src, dest, mapping,
134 op.getMapOperands(), src, dest, mapping,
139 op.getMapOperands(), src, dest, mapping,
150struct AffineInlinerInterface :
public DialectInlinerInterface {
151 using DialectInlinerInterface::DialectInlinerInterface;
162 IRMapping &valueMapping)
const final {
166 if (!isa<AffineParallelOp, AffineForOp, AffineIfOp>(destOp))
177 for (Operation &op : srcBlock) {
179 if (
auto iface = dyn_cast<MemoryEffectOpInterface>(op)) {
180 if (iface.hasNoEffect())
187 llvm::TypeSwitch<Operation *, bool>(&op)
188 .Case<AffineApplyOp, AffineReadOpInterface,
189 AffineWriteOpInterface>([&](
auto op) {
192 .Default([](Operation *) {
206 bool isLegalToInline(Operation *op, Region *region,
bool wouldBeCloned,
207 IRMapping &valueMapping)
const final {
212 Operation *parentOp = region->getParentOp();
213 return parentOp->
hasTrait<OpTrait::AffineScope>() ||
214 isa<AffineForOp, AffineParallelOp, AffineIfOp>(parentOp);
218 bool shouldAnalyzeRecursively(Operation *op)
const final {
return true; }
226void AffineDialect::initialize() {
229#include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
231 addInterfaces<AffineInlinerInterface>();
232 declarePromisedInterfaces<ValueBoundsOpInterface, AffineApplyOp, AffineMaxOp,
241 if (
auto poison = dyn_cast<ub::PoisonAttr>(value))
242 return ub::PoisonOp::create(builder, loc, type, poison);
243 return arith::ConstantOp::materialize(builder, value, type, loc);
251 if (
auto arg = dyn_cast<BlockArgument>(value)) {
267 while (
auto *parentOp = curOp->getParentOp()) {
269 return curOp->getParentRegion();
278 if (!isa<AffineForOp, AffineIfOp, AffineParallelOp>(parentOp))
303 auto *parentOp = llvm::cast<BlockArgument>(value).getOwner()->getParentOp();
331 if (
auto applyOp = dyn_cast<AffineApplyOp>(op))
332 return applyOp.isValidDim(region);
335 if (isa<AffineDelinearizeIndexOp, AffineLinearizeIndexOp>(op))
336 return llvm::all_of(op->getOperands(),
337 [&](
Value arg) { return ::isValidDim(arg, region); });
340 if (
auto dimOp = dyn_cast<ShapedDimOpInterface>(op))
348template <
typename AnyMemRefDefOp>
351 MemRefType memRefType = memrefDefOp.getType();
354 if (
index >= memRefType.getRank()) {
359 if (!memRefType.isDynamicDim(
index))
362 unsigned dynamicDimPos = memRefType.getDynamicDimIndex(
index);
363 return isValidSymbol(*(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos),
375 if (llvm::isa<BlockArgument>(dimOp.getShapedValue()))
383 if (!
index.has_value())
387 Operation *op = dimOp.getShapedValue().getDefiningOp();
388 while (
auto castOp = dyn_cast<memref::CastOp>(op)) {
390 if (isa<UnrankedMemRefType>(castOp.getSource().getType()))
392 op = castOp.getSource().getDefiningOp();
399 .Case<memref::ViewOp, memref::SubViewOp, memref::AllocOp>(
401 .Default([](
Operation *) {
return false; });
435 if (parentRegion == region)
476 if (
isPure(defOp) && llvm::all_of(defOp->getOperands(), [&](
Value operand) {
477 return affine::isValidSymbol(operand, region);
483 if (
auto dimOp = dyn_cast<ShapedDimOpInterface>(defOp))
501 printer <<
'(' << operands.take_front(numDims) <<
')';
502 if (operands.size() > numDims)
503 printer <<
'[' << operands.drop_front(numDims) <<
']';
513 numDims = opInfos.size();
527template <
typename OpTy>
532 for (
auto operand : operands) {
533 if (opIt++ < numDims) {
535 return op.emitOpError(
"operand cannot be used as a dimension id");
537 return op.emitOpError(
"operand cannot be used as a symbol");
548 return AffineValueMap(getAffineMap(), getOperands(), getResult());
555 AffineMapAttr mapAttr;
561 auto map = mapAttr.getValue();
563 if (map.getNumDims() != numDims ||
564 numDims + map.getNumSymbols() !=
result.operands.size()) {
566 "dimension or symbol index mismatch");
569 result.types.append(map.getNumResults(), indexTy);
574 p <<
" " << getMapAttr();
576 getAffineMap().getNumDims(), p);
580LogicalResult AffineApplyOp::verify() {
587 "operand count and affine map dimension and symbol count must match");
591 return emitOpError(
"mapping must produce one value");
597 for (
Value operand : getMapOperands().drop_front(affineMap.
getNumDims())) {
599 return emitError(
"dimensional operand cannot be used as a symbol");
607bool AffineApplyOp::isValidDim() {
608 return llvm::all_of(getOperands(),
615bool AffineApplyOp::isValidDim(
Region *region) {
616 return llvm::all_of(getOperands(),
617 [&](
Value op) { return ::isValidDim(op, region); });
622bool AffineApplyOp::isValidSymbol() {
623 return llvm::all_of(getOperands(),
629bool AffineApplyOp::isValidSymbol(
Region *region) {
630 return llvm::all_of(getOperands(), [&](
Value operand) {
636 auto map = getAffineMap();
639 auto expr = map.getResult(0);
640 if (
auto dim = dyn_cast<AffineDimExpr>(expr))
641 return getOperand(dim.getPosition());
642 if (
auto sym = dyn_cast<AffineSymbolExpr>(expr))
643 return getOperand(map.getNumDims() + sym.getPosition());
647 bool hasPoison =
false;
649 map.constantFold(adaptor.getMapOperands(),
result, &hasPoison);
669 auto dimExpr = dyn_cast<AffineDimExpr>(e);
679 Value operand = operands[dimExpr.getPosition()];
684 if (forOp.hasConstantLowerBound() && forOp.getConstantLowerBound() == 0) {
685 operandDivisor = forOp.getStepAsInt();
687 uint64_t lbLargestKnownDivisor =
688 forOp.getLowerBoundMap().getLargestKnownDivisorOfMapExprs();
689 operandDivisor = std::gcd(lbLargestKnownDivisor, forOp.getStepAsInt());
692 return operandDivisor;
699 if (
auto constExpr = dyn_cast<AffineConstantExpr>(e)) {
700 int64_t constVal = constExpr.getValue();
701 return constVal >= 0 && constVal < k;
703 auto dimExpr = dyn_cast<AffineDimExpr>(e);
706 Value operand = operands[dimExpr.getPosition()];
710 if (forOp.hasConstantLowerBound() && forOp.getConstantLowerBound() >= 0 &&
711 forOp.hasConstantUpperBound() && forOp.getConstantUpperBound() <= k) {
727 auto bin = dyn_cast<AffineBinaryOpExpr>(e);
735 quotientTimesDiv = llhs;
741 quotientTimesDiv = rlhs;
751 if (forOp && forOp.hasConstantLowerBound())
752 return forOp.getConstantLowerBound();
759 if (!forOp || !forOp.hasConstantUpperBound())
764 if (forOp.hasConstantLowerBound()) {
765 return forOp.getConstantUpperBound() - 1 -
766 (forOp.getConstantUpperBound() - forOp.getConstantLowerBound() - 1) %
767 forOp.getStepAsInt();
769 return forOp.getConstantUpperBound() - 1;
780 constLowerBounds.reserve(operands.size());
781 constUpperBounds.reserve(operands.size());
782 for (
Value operand : operands) {
787 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr))
788 return constExpr.getValue();
803 constLowerBounds.reserve(operands.size());
804 constUpperBounds.reserve(operands.size());
805 for (
Value operand : operands) {
810 std::optional<int64_t> lowerBound;
811 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr)) {
812 lowerBound = constExpr.getValue();
815 constLowerBounds, constUpperBounds,
826 auto binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
837 binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
845 lhs = binExpr.getLHS();
846 rhs = binExpr.getRHS();
847 auto rhsConst = dyn_cast<AffineConstantExpr>(
rhs);
851 int64_t rhsConstVal = rhsConst.getValue();
853 if (rhsConstVal <= 0)
858 std::optional<int64_t> lhsLbConst =
860 std::optional<int64_t> lhsUbConst =
862 if (lhsLbConst && lhsUbConst) {
863 int64_t lhsLbConstVal = *lhsLbConst;
864 int64_t lhsUbConstVal = *lhsUbConst;
868 divideFloorSigned(lhsLbConstVal, rhsConstVal) ==
869 divideFloorSigned(lhsUbConstVal, rhsConstVal)) {
871 divideFloorSigned(lhsLbConstVal, rhsConstVal), context);
877 divideCeilSigned(lhsLbConstVal, rhsConstVal) ==
878 divideCeilSigned(lhsUbConstVal, rhsConstVal)) {
885 lhsLbConstVal < rhsConstVal && lhsUbConstVal < rhsConstVal) {
898 if (rhsConstVal % divisor == 0 &&
900 expr = quotientTimesDiv.
floorDiv(rhsConst);
901 }
else if (divisor % rhsConstVal == 0 &&
903 expr =
rem % rhsConst;
929 if (operands.empty())
935 constLowerBounds.reserve(operands.size());
936 constUpperBounds.reserve(operands.size());
937 for (
Value operand : operands) {
951 if (
auto constExpr = dyn_cast<AffineConstantExpr>(e)) {
952 lowerBounds.push_back(constExpr.getValue());
953 upperBounds.push_back(constExpr.getValue());
955 lowerBounds.push_back(
957 constLowerBounds, constUpperBounds,
959 upperBounds.push_back(
961 constLowerBounds, constUpperBounds,
968 for (
auto exprEn : llvm::enumerate(map.
getResults())) {
970 unsigned i = exprEn.index();
972 if (lowerBounds[i] && upperBounds[i] && *lowerBounds[i] == *upperBounds[i])
977 if (!upperBounds[i]) {
978 irredundantExprs.push_back(e);
983 if (!llvm::any_of(llvm::enumerate(lowerBounds), [&](
const auto &en) {
984 auto otherLowerBound = en.value();
985 unsigned pos = en.index();
986 if (pos == i || !otherLowerBound)
988 if (*otherLowerBound > *upperBounds[i])
990 if (*otherLowerBound < *upperBounds[i])
995 if (upperBounds[pos] && lowerBounds[i] &&
996 lowerBounds[i] == upperBounds[i] &&
997 otherLowerBound == *upperBounds[pos] && i < pos)
1001 irredundantExprs.push_back(e);
1003 if (!lowerBounds[i]) {
1004 irredundantExprs.push_back(e);
1008 if (!llvm::any_of(llvm::enumerate(upperBounds), [&](
const auto &en) {
1009 auto otherUpperBound = en.value();
1010 unsigned pos = en.index();
1011 if (pos == i || !otherUpperBound)
1013 if (*otherUpperBound < *lowerBounds[i])
1015 if (*otherUpperBound > *lowerBounds[i])
1017 if (lowerBounds[pos] && upperBounds[i] &&
1018 lowerBounds[i] == upperBounds[i] &&
1019 otherUpperBound == lowerBounds[pos] && i < pos)
1023 irredundantExprs.push_back(e);
1037 assert(map.
getNumInputs() == operands.size() &&
"invalid operands for map");
1043 newResults.push_back(expr);
1066 LDBG() <<
"replaceAffineMinBoundingBoxExpression: `" << minOp <<
"`";
1067 AffineMap affineMinMap = minOp.getAffineMap();
1070 for (
unsigned i = 0, e = affineMinMap.
getNumResults(); i < e; ++i) {
1076 minOp.getOperands())))
1084 for (
auto [i, dim] : llvm::enumerate(minOp.getDimOperands())) {
1085 auto it = llvm::find(dims, dim);
1086 if (it == dims.end()) {
1087 unmappedDims.push_back(i);
1093 for (
auto [i, sym] : llvm::enumerate(minOp.getSymbolOperands())) {
1094 auto it = llvm::find(syms, sym);
1095 if (it == syms.end()) {
1096 unmappedSyms.push_back(i);
1109 if (llvm::any_of(unmappedDims,
1110 [&](
unsigned i) {
return expr.isFunctionOfDim(i); }) ||
1111 llvm::any_of(unmappedSyms,
1112 [&](
unsigned i) {
return expr.isFunctionOfSymbol(i); }))
1118 repl[dimOrSym.
ceilDiv(convertedExpr)] = c1;
1120 repl[(dimOrSym + convertedExpr - 1).floorDiv(convertedExpr)] = c1;
1125 return success(*map != initialMap);
1134 AffineExpr e,
const llvm::SmallDenseSet<AffineExpr, 4> &exprsToRemove,
1136 auto binOp = dyn_cast<AffineBinaryOpExpr>(e);
1147 llvm::SmallDenseSet<AffineExpr, 4> ourTracker(exprsToRemove);
1152 if (!ourTracker.erase(thisTerm)) {
1153 toPreserve.push_back(thisTerm);
1157 auto nextBinOp = dyn_cast_if_present<AffineBinaryOpExpr>(nextTerm);
1159 thisTerm = nextTerm;
1162 thisTerm = nextBinOp.getRHS();
1163 nextTerm = nextBinOp.getLHS();
1166 if (!ourTracker.empty())
1171 for (
AffineExpr preserved : llvm::reverse(toPreserve))
1172 newExpr = newExpr + preserved;
1173 replacementsMap.insert({e, newExpr});
1191 AffineDelinearizeIndexOp delinOp,
Value resultToReplace,
AffineMap *map,
1193 if (!delinOp.getDynamicBasis().empty())
1195 if (resultToReplace != delinOp.getMultiIndex().back())
1200 for (
auto [pos, dim] : llvm::enumerate(dims)) {
1201 auto asResult = dyn_cast_if_present<OpResult>(dim);
1204 if (asResult.getOwner() == delinOp.getOperation())
1207 for (
auto [pos, sym] : llvm::enumerate(syms)) {
1208 auto asResult = dyn_cast_if_present<OpResult>(sym);
1211 if (asResult.getOwner() == delinOp.getOperation())
1214 if (llvm::is_contained(resToExpr,
AffineExpr()))
1217 bool isDimReplacement = llvm::all_of(resToExpr, llvm::IsaPred<AffineDimExpr>);
1219 llvm::SmallDenseSet<AffineExpr, 4> expectedExprs;
1222 for (
auto [binding, size] : llvm::zip(
1223 llvm::reverse(resToExpr), llvm::reverse(delinOp.getStaticBasis()))) {
1227 if (resToExpr.size() != delinOp.getStaticBasis().size())
1228 expectedExprs.insert(resToExpr[0] * stride);
1237 if (replacements.empty())
1241 if (isDimReplacement)
1242 dims.push_back(delinOp.getLinearIndex());
1244 syms.push_back(delinOp.getLinearIndex());
1245 *map = origMap.
replace(replacements, dims.size(), syms.size());
1249 if (
auto d = dyn_cast<AffineDimExpr>(e)) {
1250 unsigned pos = d.getPosition();
1252 dims[pos] =
nullptr;
1254 if (
auto s = dyn_cast<AffineSymbolExpr>(e)) {
1255 unsigned pos = s.getPosition();
1257 syms[pos] =
nullptr;
1276 unsigned dimOrSymbolPosition,
1279 bool replaceAffineMin) {
1281 bool isDimReplacement = (dimOrSymbolPosition < dims.size());
1282 unsigned pos = isDimReplacement ? dimOrSymbolPosition
1283 : dimOrSymbolPosition - dims.size();
1284 Value &v = isDimReplacement ? dims[pos] : syms[pos];
1288 if (
auto minOp = v.
getDefiningOp<AffineMinOp>(); minOp && replaceAffineMin) {
1295 if (
auto delinOp = v.
getDefiningOp<affine::AffineDelinearizeIndexOp>()) {
1309 AffineMap composeMap = affineApply.getAffineMap();
1310 assert(composeMap.
getNumResults() == 1 &&
"affine.apply with >1 results");
1312 affineApply.getMapOperands().end());
1326 dims.append(composeDims.begin(), composeDims.end());
1327 syms.append(composeSyms.begin(), composeSyms.end());
1328 *map = map->
replace(toReplace, replacementExpr, dims.size(), syms.size());
1338 bool composeAffineMin =
false) {
1357 bool changed =
false;
1358 for (
unsigned pos = 0; pos != dims.size() + syms.size(); ++pos)
1371 unsigned nDims = 0, nSyms = 0;
1373 dimReplacements.reserve(dims.size());
1374 symReplacements.reserve(syms.size());
1375 for (
auto *container : {&dims, &syms}) {
1376 bool isDim = (container == &dims);
1377 auto &repls = isDim ? dimReplacements : symReplacements;
1378 for (
const auto &en : llvm::enumerate(*container)) {
1379 Value v = en.value();
1383 "map is function of unexpected expr@pos");
1389 operands->push_back(v);
1402 while (llvm::any_of(*operands, [](
Value v) {
1408 if (composeAffineMin && llvm::any_of(*operands, [](
Value v) {
1418 bool composeAffineMin) {
1423 return AffineApplyOp::create(
b, loc, map, valueOperands);
1429 bool composeAffineMin) {
1434 operands, composeAffineMin);
1441 bool composeAffineMin =
false) {
1447 for (
unsigned i : llvm::seq<unsigned>(0, map.
getNumResults())) {
1455 llvm::append_range(dims,
1457 llvm::append_range(symbols,
1464 operands = llvm::to_vector(llvm::concat<Value>(dims, symbols));
1471 bool composeAffineMin) {
1472 assert(map.
getNumResults() == 1 &&
"building affine.apply with !=1 result");
1482 AffineApplyOp applyOp =
1487 for (
unsigned i = 0, e = constOperands.size(); i != e; ++i)
1492 if (failed(applyOp->fold(constOperands, foldResults)) ||
1493 foldResults.empty()) {
1495 listener->notifyOperationInserted(applyOp, {});
1496 return applyOp.getResult();
1500 return llvm::getSingleElement(foldResults);
1510 operands, composeAffineMin);
1516 bool composeAffineMin) {
1517 return llvm::map_to_vector(
1518 llvm::seq<unsigned>(0, map.
getNumResults()), [&](
unsigned i) {
1519 return makeComposedFoldedAffineApply(b, loc, map.getSubMap({i}),
1520 operands, composeAffineMin);
1524template <
typename OpTy>
1530 return OpTy::create(
b, loc,
b.getIndexType(), map, valueOperands);
1539template <
typename OpTy>
1555 for (
unsigned i = 0, e = constOperands.size(); i != e; ++i)
1560 if (failed(minMaxOp->fold(constOperands, foldResults)) ||
1561 foldResults.empty()) {
1563 listener->notifyOperationInserted(minMaxOp, {});
1564 return minMaxOp.getResult();
1568 return llvm::getSingleElement(foldResults);
1587template <
class MapOrSet>
1590 if (!mapOrSet || operands->empty())
1593 assert(mapOrSet->getNumInputs() == operands->size() &&
1594 "map/set inputs must match number of operands");
1596 auto *context = mapOrSet->getContext();
1598 resultOperands.reserve(operands->size());
1600 remappedSymbols.reserve(operands->size());
1601 unsigned nextDim = 0;
1602 unsigned nextSym = 0;
1603 unsigned oldNumSyms = mapOrSet->getNumSymbols();
1605 for (
unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) {
1606 if (i < mapOrSet->getNumDims()) {
1610 remappedSymbols.push_back((*operands)[i]);
1613 resultOperands.push_back((*operands)[i]);
1616 resultOperands.push_back((*operands)[i]);
1620 resultOperands.append(remappedSymbols.begin(), remappedSymbols.end());
1621 *operands = resultOperands;
1622 *mapOrSet = mapOrSet->replaceDimsAndSymbols(
1623 dimRemapping, {}, nextDim, oldNumSyms + nextSym);
1625 assert(mapOrSet->getNumInputs() == operands->size() &&
1626 "map/set inputs must match number of operands");
1635template <
class MapOrSet>
1638 if (!mapOrSet || operands.empty())
1641 unsigned numOperands = operands.size();
1643 assert(mapOrSet.getNumInputs() == numOperands &&
1644 "map/set inputs must match number of operands");
1646 auto *context = mapOrSet.getContext();
1648 resultOperands.reserve(numOperands);
1650 remappedDims.reserve(numOperands);
1652 symOperands.reserve(mapOrSet.getNumSymbols());
1653 unsigned nextSym = 0;
1654 unsigned nextDim = 0;
1655 unsigned oldNumDims = mapOrSet.getNumDims();
1657 resultOperands.assign(operands.begin(), operands.begin() + oldNumDims);
1658 for (
unsigned i = oldNumDims, e = mapOrSet.getNumInputs(); i != e; ++i) {
1661 symRemapping[i - oldNumDims] =
1663 remappedDims.push_back(operands[i]);
1666 symOperands.push_back(operands[i]);
1670 append_range(resultOperands, remappedDims);
1671 append_range(resultOperands, symOperands);
1672 operands = resultOperands;
1673 mapOrSet = mapOrSet.replaceDimsAndSymbols(
1674 {}, symRemapping, oldNumDims + nextDim, nextSym);
1676 assert(mapOrSet.getNumInputs() == operands.size() &&
1677 "map/set inputs must match number of operands");
1681template <
class MapOrSet>
1684 static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value,
1685 "Argument must be either of AffineMap or IntegerSet type");
1687 if (!mapOrSet || operands->empty())
1690 assert(mapOrSet->getNumInputs() == operands->size() &&
1691 "map/set inputs must match number of operands");
1697 llvm::SmallBitVector usedDims(mapOrSet->getNumDims());
1698 llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols());
1700 if (
auto dimExpr = dyn_cast<AffineDimExpr>(expr))
1701 usedDims[dimExpr.getPosition()] =
true;
1702 else if (
auto symExpr = dyn_cast<AffineSymbolExpr>(expr))
1703 usedSyms[symExpr.getPosition()] =
true;
1706 auto *context = mapOrSet->getContext();
1709 resultOperands.reserve(operands->size());
1711 llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims;
1713 unsigned nextDim = 0;
1714 for (
unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) {
1717 auto it = seenDims.find((*operands)[i]);
1718 if (it == seenDims.end()) {
1720 resultOperands.push_back((*operands)[i]);
1721 seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i]));
1723 dimRemapping[i] = it->second;
1727 llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols;
1729 unsigned nextSym = 0;
1730 for (
unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) {
1736 IntegerAttr operandCst;
1737 if (
matchPattern((*operands)[i + mapOrSet->getNumDims()],
1744 auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]);
1745 if (it == seenSymbols.end()) {
1747 resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]);
1748 seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()],
1751 symRemapping[i] = it->second;
1754 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping,
1756 *operands = resultOperands;
1773template <
typename AffineOpTy>
1782 LogicalResult matchAndRewrite(AffineOpTy affineOp,
1785 llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp,
1786 AffineStoreOp, AffineApplyOp, AffineMinOp, AffineMaxOp,
1787 AffineVectorStoreOp, AffineVectorLoadOp>::value,
1788 "affine load/store/vectorstore/vectorload/apply/prefetch/min/max op "
1790 auto map = affineOp.getAffineMap();
1792 auto oldOperands = affineOp.getMapOperands();
1797 if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(),
1798 resultOperands.begin()))
1801 replaceAffineOp(rewriter, affineOp, map, resultOperands);
1809void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp(
1813 mapOperands,
load.getMaybeAlign());
1816void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp(
1820 prefetch, prefetch.getMemref(), map, mapOperands, prefetch.getIsWrite(),
1821 prefetch.getLocalityHint(), prefetch.getIsDataCache());
1824void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp(
1828 store, store.getValueToStore(), store.getMemRef(), map, mapOperands,
1829 store.getMaybeAlign());
1832void SimplifyAffineOp<AffineVectorLoadOp>::replaceAffineOp(
1836 vectorload, vectorload.getVectorType(), vectorload.getMemRef(), map,
1837 mapOperands, vectorload.getMaybeAlign());
1840void SimplifyAffineOp<AffineVectorStoreOp>::replaceAffineOp(
1844 vectorstore, vectorstore.getValueToStore(), vectorstore.getMemRef(), map,
1845 mapOperands, vectorstore.getMaybeAlign());
1849template <
typename AffineOpTy>
1850void SimplifyAffineOp<AffineOpTy>::replaceAffineOp(
1859 results.
add<SimplifyAffineOp<AffineApplyOp>>(context);
1874 result.addOperands(srcMemRef);
1875 result.addAttribute(getSrcMapAttrStrName(), AffineMapAttr::get(srcMap));
1876 result.addOperands(srcIndices);
1877 result.addOperands(destMemRef);
1878 result.addAttribute(getDstMapAttrStrName(), AffineMapAttr::get(dstMap));
1879 result.addOperands(destIndices);
1880 result.addOperands(tagMemRef);
1881 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
1882 result.addOperands(tagIndices);
1883 result.addOperands(numElements);
1885 result.addOperands({stride, elementsPerStride});
1890 p <<
" " << getSrcMemRef() <<
'[';
1892 p <<
"], " << getDstMemRef() <<
'[';
1894 p <<
"], " << getTagMemRef() <<
'[';
1898 p <<
", " << getStride();
1899 p <<
", " << getNumElementsPerStride();
1901 p <<
" : " << getSrcMemRefType() <<
", " << getDstMemRefType() <<
", "
1902 << getTagMemRefType();
1911ParseResult AffineDmaStartOp::parse(
OpAsmParser &parser,
1914 AffineMapAttr srcMapAttr;
1917 AffineMapAttr dstMapAttr;
1920 AffineMapAttr tagMapAttr;
1935 getSrcMapAttrStrName(),
1939 getDstMapAttrStrName(),
1943 getTagMapAttrStrName(),
1952 if (!strideInfo.empty() && strideInfo.size() != 2) {
1954 "expected two stride related operands");
1956 bool isStrided = strideInfo.size() == 2;
1961 if (types.size() != 3)
1979 if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() ||
1980 dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() ||
1981 tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1983 "memref operand count not equal to map.numInputs");
1987LogicalResult AffineDmaStartOp::verify() {
1988 if (!llvm::isa<MemRefType>(getOperand(getSrcMemRefOperandIndex()).
getType()))
1989 return emitOpError(
"expected DMA source to be of memref type");
1990 if (!llvm::isa<MemRefType>(getOperand(getDstMemRefOperandIndex()).
getType()))
1991 return emitOpError(
"expected DMA destination to be of memref type");
1992 if (!llvm::isa<MemRefType>(getOperand(getTagMemRefOperandIndex()).
getType()))
1993 return emitOpError(
"expected DMA tag to be of memref type");
1995 unsigned numInputsAllMaps = getSrcMap().getNumInputs() +
1996 getDstMap().getNumInputs() +
1997 getTagMap().getNumInputs();
1998 if (getNumOperands() != numInputsAllMaps + 3 + 1 &&
1999 getNumOperands() != numInputsAllMaps + 3 + 1 + 2) {
2000 return emitOpError(
"incorrect number of operands");
2004 for (
auto idx : getSrcIndices()) {
2005 if (!idx.getType().isIndex())
2006 return emitOpError(
"src index to dma_start must have 'index' type");
2009 "src index must be a valid dimension or symbol identifier");
2011 for (
auto idx : getDstIndices()) {
2012 if (!idx.getType().isIndex())
2013 return emitOpError(
"dst index to dma_start must have 'index' type");
2016 "dst index must be a valid dimension or symbol identifier");
2018 for (
auto idx : getTagIndices()) {
2019 if (!idx.getType().isIndex())
2020 return emitOpError(
"tag index to dma_start must have 'index' type");
2023 "tag index must be a valid dimension or symbol identifier");
2028LogicalResult AffineDmaStartOp::fold(FoldAdaptor adaptor,
2034void AffineDmaStartOp::getEffects(
2053 result.addOperands(tagMemRef);
2054 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
2055 result.addOperands(tagIndices);
2056 result.addOperands(numElements);
2060 p <<
" " << getTagMemRef() <<
'[';
2065 p <<
" : " << getTagMemRef().getType();
2073ParseResult AffineDmaWaitOp::parse(
OpAsmParser &parser,
2076 AffineMapAttr tagMapAttr;
2085 getTagMapAttrStrName(),
2094 if (!llvm::isa<MemRefType>(type))
2096 "expected tag to be of memref type");
2098 if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
2100 "tag memref operand count != to map.numInputs");
2104LogicalResult AffineDmaWaitOp::verify() {
2105 if (!llvm::isa<MemRefType>(getOperand(0).
getType()))
2106 return emitOpError(
"expected DMA tag to be of memref type");
2108 for (
auto idx : getTagIndices()) {
2109 if (!idx.getType().isIndex())
2110 return emitOpError(
"index to dma_wait must have 'index' type");
2113 "index must be a valid dimension or symbol identifier");
2118LogicalResult AffineDmaWaitOp::fold(FoldAdaptor adaptor,
2124void AffineDmaWaitOp::getEffects(
2140 ValueRange iterArgs, BodyBuilderFn bodyBuilder) {
2141 assert(((!lbMap && lbOperands.empty()) ||
2143 "lower bound operand count does not match the affine map");
2144 assert(((!ubMap && ubOperands.empty()) ||
2146 "upper bound operand count does not match the affine map");
2147 assert(step > 0 &&
"step has to be a positive integer constant");
2149 OpBuilder::InsertionGuard guard(builder);
2153 getOperandSegmentSizeAttr(),
2155 static_cast<int32_t>(ubOperands.size()),
2156 static_cast<int32_t>(iterArgs.size())}));
2158 for (Value val : iterArgs)
2159 result.addTypes(val.getType());
2166 result.addAttribute(getLowerBoundMapAttrName(
result.name),
2167 AffineMapAttr::get(lbMap));
2168 result.addOperands(lbOperands);
2171 result.addAttribute(getUpperBoundMapAttrName(
result.name),
2172 AffineMapAttr::get(ubMap));
2173 result.addOperands(ubOperands);
2175 result.addOperands(iterArgs);
2178 Region *bodyRegion =
result.addRegion();
2180 Value inductionVar =
2182 for (Value val : iterArgs)
2183 bodyBlock->
addArgument(val.getType(), val.getLoc());
2188 if (iterArgs.empty() && !bodyBuilder) {
2189 ensureTerminator(*bodyRegion, builder,
result.location);
2190 }
else if (bodyBuilder) {
2191 OpBuilder::InsertionGuard guard(builder);
2193 bodyBuilder(builder,
result.location, inductionVar,
2200 BodyBuilderFn bodyBuilder) {
2203 return build(builder,
result, {}, lbMap, {}, ubMap, step, iterArgs,
2207LogicalResult AffineForOp::verify() {
2208 auto *body = getBody();
2209 if (body->getNumArguments() == 0 || !getInductionVar().
getType().isIndex())
2210 return emitOpError(
"expected body to have an index argument for the "
2211 "induction variable");
2216LogicalResult AffineForOp::verifyRegions() {
2218 if (getStepAsInt() <= 0)
2219 return emitOpError(
"expected step to be a positive integer, got ")
2224 if (getLowerBoundMap().getNumInputs() > 0)
2226 getLowerBoundMap().getNumDims())))
2229 if (getUpperBoundMap().getNumInputs() > 0)
2231 getUpperBoundMap().getNumDims())))
2233 if (getLowerBoundMap().getNumResults() < 1)
2234 return emitOpError(
"expected lower bound map to have at least one result");
2235 if (getUpperBoundMap().getNumResults() < 1)
2236 return emitOpError(
"expected upper bound map to have at least one result");
2238 unsigned opNumResults = getNumResults();
2239 if (opNumResults == 0)
2245 if (getNumIterOperands() != opNumResults)
2247 "mismatch between the number of loop-carried values and results");
2248 if (getNumRegionIterArgs() != opNumResults)
2250 "mismatch between the number of basic block args and results");
2260 bool failedToParsedMinMax =
2264 auto boundAttrStrName =
2265 isLower ? AffineForOp::getLowerBoundMapAttrName(
result.name)
2266 : AffineForOp::getUpperBoundMapAttrName(
result.name);
2273 if (!boundOpInfos.empty()) {
2275 if (boundOpInfos.size() > 1)
2277 "expected only one loop bound operand");
2289 result.addAttribute(boundAttrStrName, AffineMapAttr::get(map));
2302 if (
auto affineMapAttr = dyn_cast<AffineMapAttr>(boundAttr)) {
2303 unsigned currentNumOperands =
result.operands.size();
2308 auto map = affineMapAttr.getValue();
2309 if (map.getNumDims() != numDims)
2312 "dim operand count and affine map dim count must match");
2314 unsigned numDimAndSymbolOperands =
2315 result.operands.size() - currentNumOperands;
2316 if (numDims + map.getNumSymbols() != numDimAndSymbolOperands)
2319 "symbol operand count and affine map symbol count must match");
2323 if (map.getNumResults() > 1 && failedToParsedMinMax) {
2325 return p.
emitError(attrLoc,
"lower loop bound affine map with "
2326 "multiple results requires 'max' prefix");
2328 return p.
emitError(attrLoc,
"upper loop bound affine map with multiple "
2329 "results requires 'min' prefix");
2335 if (
auto integerAttr = dyn_cast<IntegerAttr>(boundAttr)) {
2336 result.attributes.pop_back();
2345 "expected valid affine map representation for loop bounds");
2350 OpAsmParser::Argument inductionVariable;
2357 int64_t numOperands =
result.operands.size();
2360 int64_t numLbOperands =
result.operands.size() - numOperands;
2363 numOperands =
result.operands.size();
2366 int64_t numUbOperands =
result.operands.size() - numOperands;
2371 getStepAttrName(
result.name),
2375 IntegerAttr stepAttr;
2377 getStepAttrName(
result.name).data(),
2381 if (!stepAttr.getValue().isStrictlyPositive())
2384 "expected step to be representable as a positive signed integer");
2388 SmallVector<OpAsmParser::Argument, 4> regionArgs;
2389 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
2392 regionArgs.push_back(inductionVariable);
2400 for (
auto argOperandType :
2401 llvm::zip(llvm::drop_begin(regionArgs), operands,
result.types)) {
2402 Type type = std::get<2>(argOperandType);
2403 std::get<0>(argOperandType).type = type;
2411 getOperandSegmentSizeAttr(),
2413 static_cast<int32_t>(numUbOperands),
2414 static_cast<int32_t>(operands.size())}));
2417 Region *body =
result.addRegion();
2418 if (regionArgs.size() !=
result.types.size() + 1)
2421 "mismatch between the number of loop-carried values and results");
2425 AffineForOp::ensureTerminator(*body, builder,
result.location);
2447 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr)) {
2448 p << constExpr.getValue();
2456 if (isa<AffineSymbolExpr>(expr)) {
2472unsigned AffineForOp::getNumIterOperands() {
2473 AffineMap lbMap = getLowerBoundMapAttr().getValue();
2474 AffineMap ubMap = getUpperBoundMapAttr().getValue();
2479std::optional<MutableArrayRef<OpOperand>>
2480AffineForOp::getYieldedValuesMutable() {
2481 return cast<AffineYieldOp>(getBody()->getTerminator()).getOperandsMutable();
2493 if (getStepAsInt() != 1)
2494 p <<
" step " << getStepAsInt();
2496 bool printBlockTerminators =
false;
2497 if (getNumIterOperands() > 0) {
2499 auto regionArgs = getRegionIterArgs();
2500 auto operands = getInits();
2502 llvm::interleaveComma(llvm::zip(regionArgs, operands), p, [&](
auto it) {
2503 p << std::get<0>(it) <<
" = " << std::get<1>(it);
2505 p <<
") -> (" << getResultTypes() <<
")";
2506 printBlockTerminators =
true;
2511 printBlockTerminators);
2513 (*this)->getAttrs(),
2514 {getLowerBoundMapAttrName(getOperation()->getName()),
2515 getUpperBoundMapAttrName(getOperation()->getName()),
2516 getStepAttrName(getOperation()->getName()),
2517 getOperandSegmentSizeAttr()});
2522 auto foldLowerOrUpperBound = [&forOp](
bool lower) {
2526 auto boundOperands =
2527 lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands();
2528 for (
auto operand : boundOperands) {
2531 operandConstants.push_back(operandCst);
2535 lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap();
2537 "bound maps should have at least one result");
2539 if (failed(boundMap.
constantFold(operandConstants, foldedResults)))
2543 assert(!foldedResults.empty() &&
"bounds should have at least one result");
2544 auto maxOrMin = llvm::cast<IntegerAttr>(foldedResults[0]).getValue();
2545 for (
unsigned i = 1, e = foldedResults.size(); i < e; i++) {
2546 auto foldedResult = llvm::cast<IntegerAttr>(foldedResults[i]).getValue();
2547 maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult)
2548 : llvm::APIntOps::smin(maxOrMin, foldedResult);
2550 lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue())
2551 : forOp.setConstantUpperBound(maxOrMin.getSExtValue());
2556 bool folded =
false;
2557 if (!forOp.hasConstantLowerBound())
2558 folded |= succeeded(foldLowerOrUpperBound(
true));
2561 if (!forOp.hasConstantUpperBound())
2562 folded |= succeeded(foldLowerOrUpperBound(
false));
2568 int64_t step = forOp.getStepAsInt();
2569 if (!forOp.hasConstantBounds() || step <= 0)
2570 return std::nullopt;
2571 int64_t lb = forOp.getConstantLowerBound();
2572 int64_t ub = forOp.getConstantUpperBound();
2573 return ub - lb <= 0 ? 0 : (
ub - lb + step - 1) / step;
2578 if (!llvm::hasSingleElement(*forOp.getBody()))
2580 if (forOp.getNumResults() == 0)
2583 if (tripCount == 0) {
2586 return forOp.getInits();
2589 auto yieldOp = cast<AffineYieldOp>(forOp.getBody()->getTerminator());
2590 auto iterArgs = forOp.getRegionIterArgs();
2591 bool hasValDefinedOutsideLoop =
false;
2592 bool iterArgsNotInOrder =
false;
2593 for (
unsigned i = 0, e = yieldOp->getNumOperands(); i < e; ++i) {
2594 Value val = yieldOp.getOperand(i);
2598 if (val == forOp.getInductionVar())
2600 if (iterArgIt == iterArgs.end()) {
2602 assert(forOp.isDefinedOutsideOfLoop(val) &&
2603 "must be defined outside of the loop");
2604 hasValDefinedOutsideLoop =
true;
2605 replacements.push_back(val);
2607 unsigned pos = std::distance(iterArgs.begin(), iterArgIt);
2609 iterArgsNotInOrder =
true;
2610 replacements.push_back(forOp.getInits()[pos]);
2615 if (!tripCount.has_value() &&
2616 (hasValDefinedOutsideLoop || iterArgsNotInOrder))
2620 if (tripCount.has_value() && tripCount.value() >= 2 && iterArgsNotInOrder)
2622 return llvm::to_vector_of<OpFoldResult>(replacements);
2630 auto lbMap = forOp.getLowerBoundMap();
2631 auto ubMap = forOp.getUpperBoundMap();
2632 auto prevLbMap = lbMap;
2633 auto prevUbMap = ubMap;
2646 if (lbMap == prevLbMap && ubMap == prevUbMap)
2649 if (lbMap != prevLbMap)
2650 forOp.setLowerBound(lbOperands, lbMap);
2651 if (ubMap != prevUbMap)
2652 forOp.setUpperBound(ubOperands, ubMap);
2661LogicalResult AffineForOp::fold(FoldAdaptor adaptor,
2671 results.assign(getInits().begin(), getInits().end());
2675 if (!foldResults.empty()) {
2676 results.assign(foldResults);
2685 "invalid region point");
2692void AffineForOp::getSuccessorRegions(
2697 "expected loop region");
2703 if (tripCount.has_value()) {
2707 if (tripCount == 1) {
2708 regions.push_back(RegionSuccessor(getOperation()));
2714 if (tripCount.value() > 0) {
2715 regions.push_back(RegionSuccessor(&getRegion()));
2718 if (tripCount.value() == 0) {
2719 regions.push_back(RegionSuccessor(getOperation()));
2727 regions.push_back(RegionSuccessor(&getRegion()));
2728 regions.push_back(RegionSuccessor(getOperation()));
2733 return getResults();
2734 return getRegionIterArgs();
2747 assert(map.
getNumResults() >= 1 &&
"bound map has at least one result");
2748 getLowerBoundOperandsMutable().assign(lbOperands);
2749 setLowerBoundMap(map);
2754 assert(map.
getNumResults() >= 1 &&
"bound map has at least one result");
2755 getUpperBoundOperandsMutable().assign(ubOperands);
2756 setUpperBoundMap(map);
2759bool AffineForOp::hasConstantLowerBound() {
2760 return getLowerBoundMap().isSingleConstant();
2763bool AffineForOp::hasConstantUpperBound() {
2764 return getUpperBoundMap().isSingleConstant();
2767int64_t AffineForOp::getConstantLowerBound() {
2768 return getLowerBoundMap().getSingleConstantResult();
2771int64_t AffineForOp::getConstantUpperBound() {
2772 return getUpperBoundMap().getSingleConstantResult();
2775void AffineForOp::setConstantLowerBound(
int64_t value) {
2779void AffineForOp::setConstantUpperBound(
int64_t value) {
2783AffineForOp::operand_range AffineForOp::getControlOperands() {
2788bool AffineForOp::matchingBoundOperandList() {
2789 auto lbMap = getLowerBoundMap();
2790 auto ubMap = getUpperBoundMap();
2796 for (
unsigned i = 0, e = lbMap.
getNumInputs(); i < e; i++) {
2798 if (getOperand(i) != getOperand(numOperands + i))
2806std::optional<SmallVector<Value>> AffineForOp::getLoopInductionVars() {
2807 return SmallVector<Value>{getInductionVar()};
2810std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopLowerBounds() {
2811 if (!hasConstantLowerBound())
2812 return std::nullopt;
2814 return SmallVector<OpFoldResult>{
2815 OpFoldResult(
b.getI64IntegerAttr(getConstantLowerBound()))};
2818std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopSteps() {
2820 return SmallVector<OpFoldResult>{
2821 OpFoldResult(
b.getI64IntegerAttr(getStepAsInt()))};
2824std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopUpperBounds() {
2825 if (!hasConstantUpperBound())
2828 return SmallVector<OpFoldResult>{
2829 OpFoldResult(
b.getI64IntegerAttr(getConstantUpperBound()))};
2832std::optional<APInt> AffineForOp::getStaticTripCount() {
2834 int64_t step = getStepAsInt();
2836 return std::nullopt;
2838 if (hasConstantBounds()) {
2839 int64_t lb = getConstantLowerBound();
2840 int64_t ub = getConstantUpperBound();
2841 int64_t loopSpan = ub - lb;
2844 return APInt(64, llvm::divideCeilSigned(loopSpan, step));
2847 auto lbMap = getLowerBoundMap();
2848 auto ubMap = getUpperBoundMap();
2850 return std::nullopt;
2857 SmallVector<AffineExpr, 4> lbSplatExpr(ubValueMap.getNumResults(),
2860 lbSplatExpr, context);
2863 AffineValueMap tripCountValueMap;
2867 std::optional<uint64_t> tripCount;
2868 for (
unsigned i = 0, e = tripCountValueMap.
getNumResults(); i < e; ++i) {
2870 if (
auto constExpr = llvm::dyn_cast<AffineConstantExpr>(expr)) {
2871 uint64_t value = constExpr.getValue();
2872 if (tripCount.has_value())
2873 tripCount = std::min(*tripCount, value);
2877 return std::nullopt;
2881 if (tripCount.has_value())
2882 return APInt(64, *tripCount);
2884 return std::nullopt;
2887FailureOr<LoopLikeOpInterface> AffineForOp::replaceWithAdditionalYields(
2889 bool replaceInitOperandUsesInLoop,
2892 OpBuilder::InsertionGuard g(rewriter);
2894 auto inits = llvm::to_vector(getInits());
2895 inits.append(newInitOperands.begin(), newInitOperands.end());
2896 AffineForOp newLoop = AffineForOp::create(
2901 newLoop->setDiscardableAttrs(getOperation()->getDiscardableAttrDictionary());
2904 auto yieldOp = cast<AffineYieldOp>(getBody()->getTerminator());
2905 ArrayRef<BlockArgument> newIterArgs =
2906 newLoop.getBody()->getArguments().take_back(newInitOperands.size());
2908 OpBuilder::InsertionGuard g(rewriter);
2910 SmallVector<Value> newYieldedValues =
2911 newYieldValuesFn(rewriter, getLoc(), newIterArgs);
2912 assert(newInitOperands.size() == newYieldedValues.size() &&
2913 "expected as many new yield values as new iter operands");
2915 yieldOp.getOperandsMutable().append(newYieldedValues);
2920 rewriter.
mergeBlocks(getBody(), newLoop.getBody(),
2921 newLoop.getBody()->getArguments().take_front(
2922 getBody()->getNumArguments()));
2924 if (replaceInitOperandUsesInLoop) {
2927 for (
auto it : llvm::zip(newInitOperands, newIterArgs)) {
2929 [&](OpOperand &use) {
2931 return newLoop->isProperAncestor(user);
2938 newLoop->getResults().take_front(getNumResults()));
2939 return cast<LoopLikeOpInterface>(newLoop.getOperation());
2967 auto ivArg = dyn_cast<BlockArgument>(val);
2968 if (!ivArg || !ivArg.getOwner() || !ivArg.getOwner()->getParent())
2969 return AffineForOp();
2971 ivArg.getOwner()->getParent()->getParentOfType<AffineForOp>())
2973 return forOp.getInductionVar() == val ? forOp : AffineForOp();
2974 return AffineForOp();
2978 auto ivArg = dyn_cast<BlockArgument>(val);
2979 if (!ivArg || !ivArg.getOwner())
2982 auto parallelOp = dyn_cast_if_present<AffineParallelOp>(containingOp);
2983 if (parallelOp && llvm::is_contained(parallelOp.getIVs(), val))
2992 ivs->reserve(forInsts.size());
2993 for (
auto forInst : forInsts)
2994 ivs->push_back(forInst.getInductionVar());
2999 ivs.reserve(affineOps.size());
3002 if (
auto forOp = dyn_cast<AffineForOp>(op))
3003 ivs.push_back(forOp.getInductionVar());
3004 else if (
auto parallelOp = dyn_cast<AffineParallelOp>(op))
3005 for (
size_t i = 0; i < parallelOp.getBody()->getNumArguments(); i++)
3006 ivs.push_back(parallelOp.getBody()->getArgument(i));
3012template <
typename BoundListTy,
typename LoopCreatorTy>
3017 LoopCreatorTy &&loopCreatorFn) {
3018 assert(lbs.size() == ubs.size() &&
"Mismatch in number of arguments");
3019 assert(lbs.size() == steps.size() &&
"Mismatch in number of arguments");
3031 ivs.reserve(lbs.size());
3032 for (
unsigned i = 0, e = lbs.size(); i < e; ++i) {
3038 if (i == e - 1 && bodyBuilderFn) {
3040 bodyBuilderFn(nestedBuilder, nestedLoc, ivs);
3042 AffineYieldOp::create(nestedBuilder, nestedLoc);
3047 auto loop = loopCreatorFn(builder, loc, lbs[i], ubs[i], steps[i], loopBody);
3056 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3057 return AffineForOp::create(builder, loc, lb,
ub, step,
3065 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3068 if (lbConst && ubConst)
3070 ubConst.value(), step, bodyBuilderFn);
3101 LogicalResult matchAndRewrite(AffineIfOp ifOp,
3103 if (ifOp.getElseRegion().empty() ||
3104 !llvm::hasSingleElement(*ifOp.getElseBlock()) || ifOp.getNumResults())
3117 using OpRewritePattern<AffineIfOp>::OpRewritePattern;
3119 LogicalResult matchAndRewrite(AffineIfOp op,
3120 PatternRewriter &rewriter)
const override {
3122 auto isTriviallyFalse = [](IntegerSet iSet) {
3123 return iSet.isEmptyIntegerSet();
3126 auto isTriviallyTrue = [](IntegerSet iSet) {
3127 return (iSet.getNumEqualities() == 1 && iSet.getNumInequalities() == 0 &&
3128 iSet.getConstraint(0) == 0);
3131 IntegerSet affineIfConditions = op.getIntegerSet();
3133 if (isTriviallyFalse(affineIfConditions)) {
3137 if (op.getNumResults() == 0 && !op.hasElse()) {
3143 blockToMove = op.getElseBlock();
3144 }
else if (isTriviallyTrue(affineIfConditions)) {
3145 blockToMove = op.getThenBlock();
3149 Operation *blockToMoveTerminator = blockToMove->
getTerminator();
3163 rewriter.
eraseOp(blockToMoveTerminator);
3171void AffineIfOp::getSuccessorRegions(
3179 if (getElseRegion().empty()) {
3194 return getResults();
3195 if (successor == &getThenRegion())
3196 return getThenRegion().getArguments();
3197 if (successor == &getElseRegion())
3198 return getElseRegion().getArguments();
3199 llvm_unreachable(
"invalid region successor");
3202LogicalResult AffineIfOp::verify() {
3205 auto conditionAttr =
3206 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3208 return emitOpError(
"requires an integer set attribute named 'condition'");
3211 IntegerSet condition = conditionAttr.getValue();
3213 return emitOpError(
"operand count and condition integer set dimension and "
3214 "symbol count must match");
3226 IntegerSetAttr conditionAttr;
3229 AffineIfOp::getConditionAttrStrName(),
3235 auto set = conditionAttr.getValue();
3236 if (set.getNumDims() != numDims)
3239 "dim operand count and integer set dim count must match");
3240 if (numDims + set.getNumSymbols() !=
result.operands.size())
3243 "symbol operand count and integer set symbol count must match");
3250 result.regions.reserve(2);
3257 AffineIfOp::ensureTerminator(*thenRegion, parser.
getBuilder(),
3264 AffineIfOp::ensureTerminator(*elseRegion, parser.
getBuilder(),
3276 auto conditionAttr =
3277 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3278 p <<
" " << conditionAttr;
3280 conditionAttr.getValue().getNumDims(), p);
3287 auto &elseRegion = this->getElseRegion();
3288 if (!elseRegion.
empty()) {
3297 getConditionAttrStrName());
3302 ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName())
3306void AffineIfOp::setIntegerSet(
IntegerSet newSet) {
3307 (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet));
3312 (*this)->setOperands(operands);
3317 bool withElseRegion) {
3318 assert(resultTypes.empty() || withElseRegion);
3321 result.addTypes(resultTypes);
3322 result.addOperands(args);
3323 result.addAttribute(getConditionAttrStrName(), IntegerSetAttr::get(set));
3327 if (resultTypes.empty())
3328 AffineIfOp::ensureTerminator(*thenRegion, builder,
result.location);
3331 if (withElseRegion) {
3333 if (resultTypes.empty())
3334 AffineIfOp::ensureTerminator(*elseRegion, builder,
result.location);
3340 AffineIfOp::build(builder,
result, {}, set, args,
3349 bool composeAffineMin =
false) {
3356 if (llvm::none_of(operands,
3367 auto set = getIntegerSet();
3373 if (getIntegerSet() == set && llvm::equal(operands, getOperands()))
3376 setConditional(set, operands);
3382 results.
add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context);
3387 StringAttr attrName, llvm::MaybeAlign alignment) {
3389 result.addAttribute(attrName,
3399 llvm::MaybeAlign alignment) {
3400 assert(operands.size() == 1 + map.
getNumInputs() &&
"inconsistent operands");
3401 result.addOperands(operands);
3403 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3406 auto memrefType = llvm::cast<MemRefType>(operands[0].
getType());
3407 result.types.push_back(memrefType.getElementType());
3412 llvm::MaybeAlign alignment) {
3413 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
3415 result.addOperands(mapOperands);
3416 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3417 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3420 result.types.push_back(memrefType.getElementType());
3425 llvm::MaybeAlign alignment) {
3426 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3427 int64_t rank = memrefType.getRank();
3441 AffineMapAttr mapAttr;
3446 AffineLoadOp::getMapAttrStrName(),
3457 if (AffineMapAttr mapAttr =
3458 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3462 {getMapAttrStrName()});
3468template <
typename AffineMemOpTy>
3472 MemRefType memrefType,
unsigned numIndexOperands) {
3475 return op->emitOpError(
"affine map num results must equal memref rank");
3477 return op->emitOpError(
"expects as many subscripts as affine map inputs");
3479 for (
auto idx : mapOperands) {
3480 if (!idx.getType().isIndex())
3481 return op->emitOpError(
"index to load must have 'index' type");
3489LogicalResult AffineLoadOp::verify() {
3491 if (
getType() != memrefType.getElementType())
3492 return emitOpError(
"result type must match element type of memref");
3495 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3496 getMapOperands(), memrefType,
3497 getNumOperands() - 1)))
3505 results.
add<SimplifyAffineOp<AffineLoadOp>>(context);
3514 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
3519 getGlobalOp, getGlobalOp.getNameAttr());
3525 dyn_cast_or_null<DenseElementsAttr>(global.getConstantInitValue());
3529 if (
auto splatAttr = dyn_cast<SplatElementsAttr>(cstAttr))
3530 return splatAttr.getSplatValue<
Attribute>();
3532 if (!getAffineMap().isConstant())
3535 llvm::map_to_vector<4>(getAffineMap().getConstantResults(),
3536 [](
int64_t v) -> uint64_t {
return v; });
3546 ValueRange mapOperands, llvm::MaybeAlign alignment) {
3547 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
3548 result.addOperands(valueToStore);
3550 result.addOperands(mapOperands);
3551 result.getOrAddProperties<Properties>().map = AffineMapAttr::get(map);
3559 llvm::MaybeAlign alignment) {
3560 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3561 int64_t rank = memrefType.getRank();
3575 AffineMapAttr mapAttr;
3580 mapOperands, mapAttr, AffineStoreOp::getMapAttrStrName(),
3591 p <<
" " << getValueToStore();
3593 if (AffineMapAttr mapAttr =
3594 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3598 {getMapAttrStrName()});
3602LogicalResult AffineStoreOp::verify() {
3605 if (getValueToStore().
getType() != memrefType.getElementType())
3607 "value to store must have the same type as memref element type");
3610 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3611 getMapOperands(), memrefType,
3612 getNumOperands() - 2)))
3620 results.
add<SimplifyAffineOp<AffineStoreOp>>(context);
3623LogicalResult AffineStoreOp::fold(FoldAdaptor adaptor,
3633template <
typename T>
3636 if (op.getNumOperands() !=
3637 op.getMap().getNumDims() + op.getMap().getNumSymbols())
3638 return op.emitOpError(
3639 "operand count and affine map dimension and symbol count must match");
3641 if (op.getMap().getNumResults() == 0)
3642 return op.emitOpError(
"affine map expect at least one result");
3646template <
typename T>
3648 p <<
' ' << op->getAttr(T::getMapAttrStrName());
3649 auto operands = op.getOperands();
3650 unsigned numDims = op.getMap().getNumDims();
3651 p <<
'(' << operands.take_front(numDims) <<
')';
3653 if (operands.size() != numDims)
3654 p <<
'[' << operands.drop_front(numDims) <<
']';
3656 {T::getMapAttrStrName()});
3659template <
typename T>
3666 AffineMapAttr mapAttr;
3682template <
typename T>
3684 static_assert(llvm::is_one_of<T, AffineMinOp, AffineMaxOp>::value,
3685 "expected affine min or max op");
3691 auto foldedMap = op.getMap().partialConstantFold(operands, &results);
3693 if (foldedMap.getNumSymbols() == 1 && foldedMap.isSymbolIdentity())
3694 return op.getOperand(0);
3697 if (results.empty()) {
3699 if (foldedMap == op.getMap())
3701 op->setAttr(
"map", AffineMapAttr::get(foldedMap));
3702 return op.getResult();
3706 auto resultIt = std::is_same<T, AffineMinOp>::value
3707 ? llvm::min_element(results)
3708 : llvm::max_element(results);
3709 if (resultIt == results.end())
3711 return IntegerAttr::get(IndexType::get(op.getContext()), *resultIt);
3715template <
typename T>
3721 AffineMap oldMap = affineOp.getAffineMap();
3727 if (!llvm::is_contained(newExprs, expr))
3728 newExprs.push_back(expr);
3758template <
typename T>
3764 AffineMap oldMap = affineOp.getAffineMap();
3766 affineOp.getMapOperands().take_front(oldMap.
getNumDims());
3768 affineOp.getMapOperands().take_back(oldMap.
getNumSymbols());
3770 auto newDimOperands = llvm::to_vector<8>(dimOperands);
3771 auto newSymOperands = llvm::to_vector<8>(symOperands);
3779 if (
auto symExpr = dyn_cast<AffineSymbolExpr>(expr)) {
3780 Value symValue = symOperands[symExpr.getPosition()];
3782 producerOps.push_back(producerOp);
3785 }
else if (
auto dimExpr = dyn_cast<AffineDimExpr>(expr)) {
3786 Value dimValue = dimOperands[dimExpr.getPosition()];
3788 producerOps.push_back(producerOp);
3795 newExprs.push_back(expr);
3798 if (producerOps.empty())
3805 for (T producerOp : producerOps) {
3806 AffineMap producerMap = producerOp.getAffineMap();
3807 unsigned numProducerDims = producerMap.
getNumDims();
3812 producerOp.getMapOperands().take_front(numProducerDims);
3814 producerOp.getMapOperands().take_back(numProducerSyms);
3815 newDimOperands.append(dimValues.begin(), dimValues.end());
3816 newSymOperands.append(symValues.begin(), symValues.end());
3820 newExprs.push_back(expr.
shiftDims(numProducerDims, numUsedDims)
3824 numUsedDims += numProducerDims;
3825 numUsedSyms += numProducerSyms;
3831 llvm::to_vector<8>(llvm::concat<Value>(newDimOperands, newSymOperands));
3850 if (!resultExpr.isPureAffine())
3855 if (failed(flattenResult))
3868 if (llvm::is_sorted(flattenedExprs))
3873 llvm::to_vector(llvm::seq<unsigned>(0, map.
getNumResults()));
3874 llvm::sort(resultPermutation, [&](
unsigned lhs,
unsigned rhs) {
3875 return flattenedExprs[
lhs] < flattenedExprs[
rhs];
3878 for (
unsigned idx : resultPermutation)
3899template <
typename T>
3905 AffineMap map = affineOp.getAffineMap();
3913template <
typename T>
3919 if (affineOp.getMap().getNumResults() != 1)
3922 affineOp.getOperands());
3990ParseResult AffinePrefetchOp::parse(
OpAsmParser &parser,
3997 IntegerAttr hintInfo;
3999 StringRef readOrWrite, cacheType;
4001 AffineMapAttr mapAttr;
4005 AffinePrefetchOp::getMapAttrStrName(),
4011 AffinePrefetchOp::getLocalityHintAttrStrName(),
4021 if (readOrWrite !=
"read" && readOrWrite !=
"write")
4023 "rw specifier has to be 'read' or 'write'");
4024 result.addAttribute(AffinePrefetchOp::getIsWriteAttrStrName(),
4027 if (cacheType !=
"data" && cacheType !=
"instr")
4029 "cache type has to be 'data' or 'instr'");
4031 result.addAttribute(AffinePrefetchOp::getIsDataCacheAttrStrName(),
4038 p <<
" " << getMemref() <<
'[';
4039 AffineMapAttr mapAttr =
4040 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4043 p <<
']' <<
", " << (getIsWrite() ?
"write" :
"read") <<
", " <<
"locality<"
4044 << getLocalityHint() <<
">, " << (getIsDataCache() ?
"data" :
"instr");
4046 (*this)->getAttrs(),
4047 {getMapAttrStrName(), getLocalityHintAttrStrName(),
4048 getIsDataCacheAttrStrName(), getIsWriteAttrStrName()});
4052LogicalResult AffinePrefetchOp::verify() {
4053 auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4057 return emitOpError(
"affine.prefetch affine map num results must equal"
4062 if (getNumOperands() != 1)
4067 for (
auto idx : getMapOperands()) {
4070 "index must be a valid dimension or symbol identifier");
4078 results.
add<SimplifyAffineOp<AffinePrefetchOp>>(context);
4081LogicalResult AffinePrefetchOp::fold(FoldAdaptor adaptor,
4096 auto ubs = llvm::map_to_vector<4>(ranges, [&](
int64_t value) {
4100 build(builder,
result, resultTypes, reductions, lbs, {}, ubs,
4110 assert(llvm::all_of(lbMaps,
4112 return m.
getNumDims() == lbMaps[0].getNumDims() &&
4115 "expected all lower bounds maps to have the same number of dimensions "
4117 assert(llvm::all_of(ubMaps,
4119 return m.
getNumDims() == ubMaps[0].getNumDims() &&
4122 "expected all upper bounds maps to have the same number of dimensions "
4124 assert((lbMaps.empty() || lbMaps[0].getNumInputs() == lbArgs.size()) &&
4125 "expected lower bound maps to have as many inputs as lower bound "
4127 assert((ubMaps.empty() || ubMaps[0].getNumInputs() == ubArgs.size()) &&
4128 "expected upper bound maps to have as many inputs as upper bound "
4132 result.addTypes(resultTypes);
4136 for (arith::AtomicRMWKind reduction : reductions)
4137 reductionAttrs.push_back(
4139 result.addAttribute(getReductionsAttrStrName(),
4149 groups.reserve(groups.size() + maps.size());
4150 exprs.reserve(maps.size());
4155 return AffineMap::get(maps[0].getNumDims(), maps[0].getNumSymbols(), exprs,
4161 AffineMap lbMap = concatMapsSameInput(lbMaps, lbGroups);
4162 AffineMap ubMap = concatMapsSameInput(ubMaps, ubGroups);
4163 result.addAttribute(getLowerBoundsMapAttrStrName(),
4164 AffineMapAttr::get(lbMap));
4165 result.addAttribute(getLowerBoundsGroupsAttrStrName(),
4167 result.addAttribute(getUpperBoundsMapAttrStrName(),
4168 AffineMapAttr::get(ubMap));
4169 result.addAttribute(getUpperBoundsGroupsAttrStrName(),
4172 result.addOperands(lbArgs);
4173 result.addOperands(ubArgs);
4176 auto *bodyRegion =
result.addRegion();
4180 for (
unsigned i = 0, e = steps.size(); i < e; ++i)
4182 if (resultTypes.empty())
4183 ensureTerminator(*bodyRegion, builder,
result.location);
4187 return {&getRegion()};
4190unsigned AffineParallelOp::getNumDims() {
return getSteps().size(); }
4192AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() {
4193 return getOperands().take_front(getLowerBoundsMap().getNumInputs());
4196AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() {
4197 return getOperands().drop_front(getLowerBoundsMap().getNumInputs());
4200AffineMap AffineParallelOp::getLowerBoundMap(
unsigned pos) {
4201 auto values = getLowerBoundsGroups().getValues<int32_t>();
4203 for (
unsigned i = 0; i < pos; ++i)
4205 return getLowerBoundsMap().getSliceMap(start, values[pos]);
4208AffineMap AffineParallelOp::getUpperBoundMap(
unsigned pos) {
4209 auto values = getUpperBoundsGroups().getValues<int32_t>();
4211 for (
unsigned i = 0; i < pos; ++i)
4213 return getUpperBoundsMap().getSliceMap(start, values[pos]);
4217 return AffineValueMap(getLowerBoundsMap(), getLowerBoundsOperands());
4221 return AffineValueMap(getUpperBoundsMap(), getUpperBoundsOperands());
4224std::optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() {
4225 if (hasMinMaxBounds())
4226 return std::nullopt;
4234 for (
unsigned i = 0, e = rangesValueMap.
getNumResults(); i < e; ++i) {
4235 auto expr = rangesValueMap.
getResult(i);
4236 auto cst = dyn_cast<AffineConstantExpr>(expr);
4238 return std::nullopt;
4239 out.push_back(cst.getValue());
4244Block *AffineParallelOp::getBody() {
return &getRegion().
front(); }
4246OpBuilder AffineParallelOp::getBodyBuilder() {
4247 return OpBuilder(getBody(), std::prev(getBody()->end()));
4252 "operands to map must match number of inputs");
4254 auto ubOperands = getUpperBoundsOperands();
4257 newOperands.append(ubOperands.begin(), ubOperands.end());
4258 (*this)->setOperands(newOperands);
4260 setLowerBoundsMapAttr(AffineMapAttr::get(map));
4265 "operands to map must match number of inputs");
4268 newOperands.append(ubOperands.begin(), ubOperands.end());
4269 (*this)->setOperands(newOperands);
4271 setUpperBoundsMapAttr(AffineMapAttr::get(map));
4280 arith::AtomicRMWKind op) {
4282 case arith::AtomicRMWKind::addf:
4283 return isa<FloatType>(resultType);
4284 case arith::AtomicRMWKind::addi:
4285 return isa<IntegerType>(resultType);
4286 case arith::AtomicRMWKind::assign:
4288 case arith::AtomicRMWKind::mulf:
4289 return isa<FloatType>(resultType);
4290 case arith::AtomicRMWKind::muli:
4291 return isa<IntegerType>(resultType);
4292 case arith::AtomicRMWKind::maximumf:
4293 case arith::AtomicRMWKind::maxnumf:
4294 case arith::AtomicRMWKind::minimumf:
4295 case arith::AtomicRMWKind::minnumf:
4296 return isa<FloatType>(resultType);
4297 case arith::AtomicRMWKind::maxs: {
4298 auto intType = dyn_cast<IntegerType>(resultType);
4299 return intType && intType.isSigned();
4301 case arith::AtomicRMWKind::mins: {
4302 auto intType = dyn_cast<IntegerType>(resultType);
4303 return intType && intType.isSigned();
4305 case arith::AtomicRMWKind::maxu: {
4306 auto intType = dyn_cast<IntegerType>(resultType);
4307 return intType && intType.isUnsigned();
4309 case arith::AtomicRMWKind::minu: {
4310 auto intType = dyn_cast<IntegerType>(resultType);
4311 return intType && intType.isUnsigned();
4313 case arith::AtomicRMWKind::ori:
4314 case arith::AtomicRMWKind::andi:
4315 case arith::AtomicRMWKind::xori:
4316 return isa<IntegerType>(resultType);
4318 llvm_unreachable(
"Unhandled atomic rmw kind");
4321LogicalResult AffineParallelOp::verify() {
4322 auto numDims = getNumDims();
4325 getSteps().size() != numDims || getBody()->getNumArguments() != numDims) {
4326 return emitOpError() <<
"the number of region arguments ("
4327 << getBody()->getNumArguments()
4328 <<
") and the number of map groups for lower ("
4329 << getLowerBoundsGroups().getNumElements()
4330 <<
") and upper bound ("
4331 << getUpperBoundsGroups().getNumElements()
4332 <<
"), and the number of steps (" << getSteps().size()
4333 <<
") must all match";
4336 unsigned expectedNumLBResults = 0;
4337 for (APInt v : getLowerBoundsGroups()) {
4338 unsigned results = v.getZExtValue();
4341 <<
"expected lower bound map to have at least one result";
4342 expectedNumLBResults += results;
4344 if (expectedNumLBResults != getLowerBoundsMap().getNumResults())
4345 return emitOpError() <<
"expected lower bounds map to have "
4346 << expectedNumLBResults <<
" results";
4347 unsigned expectedNumUBResults = 0;
4348 for (APInt v : getUpperBoundsGroups()) {
4349 unsigned results = v.getZExtValue();
4352 <<
"expected upper bound map to have at least one result";
4353 expectedNumUBResults += results;
4355 if (expectedNumUBResults != getUpperBoundsMap().getNumResults())
4356 return emitOpError() <<
"expected upper bounds map to have "
4357 << expectedNumUBResults <<
" results";
4359 if (getReductions().size() != getNumResults())
4360 return emitOpError(
"a reduction must be specified for each output");
4364 for (
auto it : llvm::enumerate((getReductions()))) {
4366 auto intAttr = dyn_cast<IntegerAttr>(attr);
4367 if (!intAttr || !arith::symbolizeAtomicRMWKind(intAttr.getInt()))
4368 return emitOpError(
"invalid reduction attribute");
4369 auto kind = arith::symbolizeAtomicRMWKind(intAttr.getInt()).value();
4371 return emitOpError(
"result type cannot match reduction attribute");
4377 getLowerBoundsMap().getNumDims())))
4381 getUpperBoundsMap().getNumDims())))
4390 if (newMap ==
getAffineMap() && newOperands == operands)
4392 reset(newMap, newOperands);
4402 bool ubCanonicalized = succeeded(
ub.canonicalize());
4405 if (!lbCanonicalized && !ubCanonicalized)
4408 if (lbCanonicalized)
4410 if (ubCanonicalized)
4411 op.setUpperBounds(
ub.getOperands(),
ub.getAffineMap());
4416LogicalResult AffineParallelOp::fold(FoldAdaptor adaptor,
4417 SmallVectorImpl<OpFoldResult> &results) {
4428 StringRef keyword) {
4431 ValueRange dimOperands = operands.take_front(numDims);
4432 ValueRange symOperands = operands.drop_front(numDims);
4434 for (llvm::APInt groupSize : group) {
4438 unsigned size = groupSize.getZExtValue();
4443 p << keyword <<
'(';
4452void AffineParallelOp::print(OpAsmPrinter &p) {
4453 p <<
" (" << getBody()->getArguments() <<
") = (";
4455 getLowerBoundsOperands(),
"max");
4458 getUpperBoundsOperands(),
"min");
4460 SmallVector<int64_t, 8> steps = getSteps();
4461 bool elideSteps = llvm::all_of(steps, [](int64_t step) {
return step == 1; });
4464 llvm::interleaveComma(steps, p);
4467 if (getNumResults()) {
4469 llvm::interleaveComma(getReductions(), p, [&](
auto &attr) {
4470 arith::AtomicRMWKind sym = *arith::symbolizeAtomicRMWKind(
4471 llvm::cast<IntegerAttr>(attr).getInt());
4472 p <<
"\"" << arith::stringifyAtomicRMWKind(sym) <<
"\"";
4474 p <<
") -> (" << getResultTypes() <<
")";
4481 (*this)->getAttrs(),
4482 {AffineParallelOp::getReductionsAttrStrName(),
4483 AffineParallelOp::getLowerBoundsMapAttrStrName(),
4484 AffineParallelOp::getLowerBoundsGroupsAttrStrName(),
4485 AffineParallelOp::getUpperBoundsMapAttrStrName(),
4486 AffineParallelOp::getUpperBoundsGroupsAttrStrName(),
4487 AffineParallelOp::getStepsAttrStrName()});
4494static ParseResult deduplicateAndResolveOperands(
4495 OpAsmParser &parser,
4496 ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> operands,
4497 SmallVectorImpl<Value> &uniqueOperands,
4498 SmallVectorImpl<AffineExpr> &replacements,
AffineExprKind kind) {
4500 "expected operands to be dim or symbol expression");
4503 for (
const auto &list : operands) {
4504 SmallVector<Value> valueOperands;
4507 for (Value operand : valueOperands) {
4508 unsigned pos = std::distance(uniqueOperands.begin(),
4509 llvm::find(uniqueOperands, operand));
4510 if (pos == uniqueOperands.size())
4511 uniqueOperands.push_back(operand);
4512 replacements.push_back(
4522enum class MinMaxKind { Min, Max };
4541static ParseResult parseAffineMapWithMinMax(OpAsmParser &parser,
4546 const llvm::StringLiteral tmpAttrStrName =
"__pseudo_bound_map";
4548 StringRef mapName = kind == MinMaxKind::Min
4549 ? AffineParallelOp::getUpperBoundsMapAttrStrName()
4550 : AffineParallelOp::getLowerBoundsMapAttrStrName();
4551 StringRef groupsName =
4552 kind == MinMaxKind::Min
4553 ? AffineParallelOp::getUpperBoundsGroupsAttrStrName()
4554 : AffineParallelOp::getLowerBoundsGroupsAttrStrName();
4560 result.addAttribute(
4561 mapName, AffineMapAttr::get(parser.getBuilder().getEmptyAffineMap()));
4562 result.addAttribute(groupsName, parser.getBuilder().getI32TensorAttr({}));
4566 SmallVector<AffineExpr> flatExprs;
4567 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatDimOperands;
4568 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatSymOperands;
4569 SmallVector<int32_t> numMapsPerGroup;
4570 SmallVector<OpAsmParser::UnresolvedOperand> mapOperands;
4571 auto parseOperands = [&]() {
4573 kind == MinMaxKind::Min ?
"min" :
"max"))) {
4574 mapOperands.clear();
4580 result.attributes.erase(tmpAttrStrName);
4581 llvm::append_range(flatExprs, map.getValue().getResults());
4582 auto operandsRef = llvm::ArrayRef(mapOperands);
4583 auto dimsRef = operandsRef.take_front(map.getValue().getNumDims());
4584 SmallVector<OpAsmParser::UnresolvedOperand> dims(dimsRef);
4585 auto symsRef = operandsRef.drop_front(map.getValue().getNumDims());
4586 SmallVector<OpAsmParser::UnresolvedOperand> syms(symsRef);
4587 flatDimOperands.append(map.getValue().getNumResults(), dims);
4588 flatSymOperands.append(map.getValue().getNumResults(), syms);
4589 numMapsPerGroup.push_back(map.getValue().getNumResults());
4592 flatSymOperands.emplace_back(),
4593 flatExprs.emplace_back())))
4595 numMapsPerGroup.push_back(1);
4602 unsigned totalNumDims = 0;
4603 unsigned totalNumSyms = 0;
4604 for (
unsigned i = 0, e = flatExprs.size(); i < e; ++i) {
4605 unsigned numDims = flatDimOperands[i].size();
4606 unsigned numSyms = flatSymOperands[i].size();
4607 flatExprs[i] = flatExprs[i]
4608 .shiftDims(numDims, totalNumDims)
4609 .shiftSymbols(numSyms, totalNumSyms);
4610 totalNumDims += numDims;
4611 totalNumSyms += numSyms;
4615 SmallVector<Value> dimOperands, symOperands;
4616 SmallVector<AffineExpr> dimRplacements, symRepacements;
4617 if (deduplicateAndResolveOperands(parser, flatDimOperands, dimOperands,
4619 deduplicateAndResolveOperands(parser, flatSymOperands, symOperands,
4623 result.operands.append(dimOperands.begin(), dimOperands.end());
4624 result.operands.append(symOperands.begin(), symOperands.end());
4627 auto flatMap =
AffineMap::get(totalNumDims, totalNumSyms, flatExprs,
4629 flatMap = flatMap.replaceDimsAndSymbols(
4630 dimRplacements, symRepacements, dimOperands.size(), symOperands.size());
4632 result.addAttribute(mapName, AffineMapAttr::get(flatMap));
4642ParseResult AffineParallelOp::parse(OpAsmParser &parser,
4643 OperationState &
result) {
4646 SmallVector<OpAsmParser::Argument, 4> ivs;
4649 parseAffineMapWithMinMax(parser,
result, MinMaxKind::Max) ||
4651 parseAffineMapWithMinMax(parser,
result, MinMaxKind::Min))
4654 AffineMapAttr stepsMapAttr;
4655 NamedAttrList stepsAttrs;
4656 SmallVector<OpAsmParser::UnresolvedOperand, 4> stepsMapOperands;
4658 SmallVector<int64_t, 4> steps(ivs.size(), 1);
4659 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4663 AffineParallelOp::getStepsAttrStrName(),
4669 SmallVector<int64_t, 4> steps;
4670 auto stepsMap = stepsMapAttr.getValue();
4671 for (
const auto &
result : stepsMap.getResults()) {
4672 auto constExpr = dyn_cast<AffineConstantExpr>(
result);
4675 "steps must be constant integers");
4676 steps.push_back(constExpr.getValue());
4678 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4684 SmallVector<Attribute, 4> reductions;
4688 auto parseAttributes = [&]() -> ParseResult {
4693 NamedAttrList attrStorage;
4698 std::optional<arith::AtomicRMWKind> reduction =
4699 arith::symbolizeAtomicRMWKind(attrVal.getValue());
4701 return parser.
emitError(loc,
"invalid reduction value: ") << attrVal;
4702 reductions.push_back(
4710 result.addAttribute(AffineParallelOp::getReductionsAttrStrName(),
4718 Region *body =
result.addRegion();
4719 for (
auto &iv : ivs)
4720 iv.type = indexType;
4726 AffineParallelOp::ensureTerminator(*body, builder,
result.location);
4734LogicalResult AffineYieldOp::verify() {
4735 auto *parentOp = (*this)->getParentOp();
4736 auto results = parentOp->getResults();
4737 auto operands = getOperands();
4739 if (!isa<AffineParallelOp, AffineIfOp, AffineForOp>(parentOp))
4740 return emitOpError() <<
"only terminates affine.if/for/parallel regions";
4741 if (parentOp->getNumResults() != getNumOperands())
4742 return emitOpError() <<
"parent of yield must have same number of "
4743 "results as the yield operands";
4744 for (
auto it : llvm::zip(results, operands)) {
4746 return emitOpError() <<
"types mismatch between yield op and its parent";
4756void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4757 VectorType resultType, AffineMap map,
4759 llvm::MaybeAlign alignment) {
4760 assert(operands.size() == 1 + map.
getNumInputs() &&
"inconsistent operands");
4761 result.addOperands(operands);
4763 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4766 result.types.push_back(resultType);
4769void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4770 VectorType resultType, Value memref,
4772 llvm::MaybeAlign alignment) {
4773 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
4774 result.addOperands(memref);
4775 result.addOperands(mapOperands);
4776 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4779 result.types.push_back(resultType);
4782void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4783 VectorType resultType, Value memref,
4785 auto memrefType = llvm::cast<MemRefType>(memref.
getType());
4786 int64_t rank = memrefType.getRank();
4791 build(builder,
result, resultType, memref, map,
indices, alignment);
4794void AffineVectorLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
4795 MLIRContext *context) {
4796 results.
add<SimplifyAffineOp<AffineVectorLoadOp>>(context);
4799ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser,
4800 OperationState &
result) {
4804 MemRefType memrefType;
4805 VectorType resultType;
4806 OpAsmParser::UnresolvedOperand memrefInfo;
4807 AffineMapAttr mapAttr;
4808 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4812 AffineVectorLoadOp::getMapAttrStrName(),
4822void AffineVectorLoadOp::print(OpAsmPrinter &p) {
4824 if (AffineMapAttr mapAttr =
4825 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4829 {getMapAttrStrName()});
4834static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType,
4835 VectorType vectorType) {
4837 if (memrefType.getElementType() != vectorType.getElementType())
4839 "requires memref and vector types of the same elemental type");
4843LogicalResult AffineVectorLoadOp::verify() {
4846 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4847 getMapOperands(), memrefType,
4848 getNumOperands() - 1)))
4861void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &
result,
4862 Value valueToStore, Value memref, AffineMap map,
4864 llvm::MaybeAlign alignment) {
4865 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
4866 result.addOperands(valueToStore);
4867 result.addOperands(memref);
4868 result.addOperands(mapOperands);
4869 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4875void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &
result,
4876 Value valueToStore, Value memref,
4878 llvm::MaybeAlign alignment) {
4879 auto memrefType = llvm::cast<MemRefType>(memref.
getType());
4880 int64_t rank = memrefType.getRank();
4885 build(builder,
result, valueToStore, memref, map,
indices, alignment);
4887void AffineVectorStoreOp::getCanonicalizationPatterns(
4888 RewritePatternSet &results, MLIRContext *context) {
4889 results.
add<SimplifyAffineOp<AffineVectorStoreOp>>(context);
4892ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser,
4893 OperationState &
result) {
4896 MemRefType memrefType;
4897 VectorType resultType;
4898 OpAsmParser::UnresolvedOperand storeValueInfo;
4899 OpAsmParser::UnresolvedOperand memrefInfo;
4900 AffineMapAttr mapAttr;
4901 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4906 AffineVectorStoreOp::getMapAttrStrName(),
4916void AffineVectorStoreOp::print(OpAsmPrinter &p) {
4917 p <<
" " << getValueToStore();
4919 if (AffineMapAttr mapAttr =
4920 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4924 {getMapAttrStrName()});
4925 p <<
" : " <<
getMemRefType() <<
", " << getValueToStore().getType();
4928LogicalResult AffineVectorStoreOp::verify() {
4931 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4932 getMapOperands(), memrefType,
4933 getNumOperands() - 2)))
4946void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4947 OperationState &odsState,
4949 ArrayRef<int64_t> staticBasis,
4950 bool hasOuterBound) {
4951 SmallVector<Type> returnTypes(hasOuterBound ? staticBasis.size()
4952 : staticBasis.size() + 1,
4954 build(odsBuilder, odsState, returnTypes, linearIndex, dynamicBasis,
4958void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4959 OperationState &odsState,
4961 bool hasOuterBound) {
4962 if (hasOuterBound && !basis.empty() && basis.front() ==
nullptr) {
4963 hasOuterBound =
false;
4964 basis = basis.drop_front();
4966 SmallVector<Value> dynamicBasis;
4967 SmallVector<int64_t> staticBasis;
4970 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4974void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4975 OperationState &odsState,
4977 ArrayRef<OpFoldResult> basis,
4978 bool hasOuterBound) {
4979 if (hasOuterBound && !basis.empty() && basis.front() == OpFoldResult()) {
4980 hasOuterBound =
false;
4981 basis = basis.drop_front();
4983 SmallVector<Value> dynamicBasis;
4984 SmallVector<int64_t> staticBasis;
4986 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4990void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4991 OperationState &odsState,
4992 Value linearIndex, ArrayRef<int64_t> basis,
4993 bool hasOuterBound) {
4994 build(odsBuilder, odsState, linearIndex,
ValueRange{}, basis, hasOuterBound);
4997LogicalResult AffineDelinearizeIndexOp::verify() {
4998 ArrayRef<int64_t> staticBasis = getStaticBasis();
4999 if (getNumResults() != staticBasis.size() &&
5000 getNumResults() != staticBasis.size() + 1)
5001 return emitOpError(
"should return an index for each basis element and up "
5002 "to one extra index");
5004 auto dynamicMarkersCount = llvm::count_if(staticBasis, ShapedType::isDynamic);
5005 if (
static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5007 "mismatch between dynamic and static basis (kDynamic marker but no "
5008 "corresponding dynamic basis entry) -- this can only happen due to an "
5009 "incorrect fold/rewrite");
5011 if (!llvm::all_of(staticBasis, [](int64_t v) {
5012 return v > 0 || ShapedType::isDynamic(v);
5014 return emitOpError(
"no basis element may be statically non-positive");
5023static std::optional<SmallVector<int64_t>>
5027 uint64_t dynamicBasisIndex = 0;
5033 if (basis && isa<IntegerAttr>(basis)) {
5034 mutableDynamicBasis.
erase(dynamicBasisIndex);
5036 ++dynamicBasisIndex;
5041 if (dynamicBasisIndex == dynamicBasis.size())
5042 return std::nullopt;
5048 staticBasis.push_back(ShapedType::kDynamic);
5050 staticBasis.push_back(*basisVal);
5057AffineDelinearizeIndexOp::fold(FoldAdaptor adaptor,
5058 SmallVectorImpl<OpFoldResult> &
result) {
5059 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5061 adaptor.getDynamicBasis());
5062 if (maybeStaticBasis) {
5063 setStaticBasis(*maybeStaticBasis);
5068 if (getNumResults() == 1) {
5069 result.push_back(getLinearIndex());
5073 if (adaptor.getLinearIndex() ==
nullptr)
5076 if (!adaptor.getDynamicBasis().empty())
5079 int64_t highPart = cast<IntegerAttr>(adaptor.getLinearIndex()).getInt();
5080 Type attrType = getLinearIndex().getType();
5082 ArrayRef<int64_t> staticBasis = getStaticBasis();
5083 if (hasOuterBound())
5084 staticBasis = staticBasis.drop_front();
5085 for (int64_t modulus : llvm::reverse(staticBasis)) {
5086 result.push_back(IntegerAttr::get(attrType, llvm::mod(highPart, modulus)));
5087 highPart = llvm::divideFloorSigned(highPart, modulus);
5089 result.push_back(IntegerAttr::get(attrType, highPart));
5094SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getEffectiveBasis() {
5096 if (hasOuterBound()) {
5097 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5099 getDynamicBasis().drop_front(), builder);
5101 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5105 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5108SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getPaddedBasis() {
5109 SmallVector<OpFoldResult> ret = getMixedBasis();
5110 if (!hasOuterBound())
5111 ret.insert(ret.begin(), OpFoldResult());
5118struct DropUnitExtentBasis
5119 :
public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5122 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5123 PatternRewriter &rewriter)
const override {
5124 SmallVector<Value> replacements(delinearizeOp->getNumResults(),
nullptr);
5125 std::optional<Value> zero = std::nullopt;
5126 Location loc = delinearizeOp->getLoc();
5127 Type indexType = delinearizeOp.getLinearIndex().getType();
5128 auto getZero = [&]() -> Value {
5130 zero = arith::ConstantOp::create(rewriter, loc,
5132 return zero.value();
5137 SmallVector<OpFoldResult> newBasis;
5138 for (
auto [index, basis] :
5139 llvm::enumerate(delinearizeOp.getPaddedBasis())) {
5140 std::optional<int64_t> basisVal =
5143 replacements[index] =
getZero();
5145 newBasis.push_back(basis);
5148 if (newBasis.size() == delinearizeOp.getNumResults())
5150 "no unit basis elements");
5152 if (!newBasis.empty()) {
5154 auto newDelinearizeOp = affine::AffineDelinearizeIndexOp::create(
5155 rewriter, loc, delinearizeOp.getLinearIndex(), newBasis);
5161 replacement = newDelinearizeOp->getResult(newIndex++);
5165 rewriter.
replaceOp(delinearizeOp, replacements);
5180struct CancelDelinearizeOfLinearizeDisjointExactTail
5181 :
public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5184 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5185 PatternRewriter &rewriter)
const override {
5186 auto linearizeOp = delinearizeOp.getLinearIndex()
5187 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5190 "index doesn't come from linearize");
5192 if (!linearizeOp.getDisjoint())
5195 ValueRange linearizeIns = linearizeOp.getMultiIndex();
5197 SmallVector<OpFoldResult> linearizeBasis = linearizeOp.getMixedBasis();
5198 SmallVector<OpFoldResult> delinearizeBasis = delinearizeOp.getMixedBasis();
5199 size_t numMatches = 0;
5200 for (
auto [linSize, delinSize] : llvm::zip(
5201 llvm::reverse(linearizeBasis), llvm::reverse(delinearizeBasis))) {
5202 if (linSize != delinSize)
5207 if (numMatches == 0)
5209 delinearizeOp,
"final basis element doesn't match linearize");
5212 if (numMatches == linearizeBasis.size() &&
5213 numMatches == delinearizeBasis.size() &&
5214 linearizeIns.size() == delinearizeOp.getNumResults()) {
5215 rewriter.
replaceOp(delinearizeOp, linearizeOp.getMultiIndex());
5219 Value newLinearize = affine::AffineLinearizeIndexOp::create(
5220 rewriter, linearizeOp.getLoc(), linearizeIns.drop_back(numMatches),
5221 ArrayRef<OpFoldResult>{linearizeBasis}.drop_back(numMatches),
5222 linearizeOp.getDisjoint());
5223 auto newDelinearize = affine::AffineDelinearizeIndexOp::create(
5224 rewriter, delinearizeOp.getLoc(), newLinearize,
5225 ArrayRef<OpFoldResult>{delinearizeBasis}.drop_back(numMatches),
5226 delinearizeOp.hasOuterBound());
5227 SmallVector<Value> mergedResults(newDelinearize.getResults());
5228 mergedResults.append(linearizeIns.take_back(numMatches).begin(),
5229 linearizeIns.take_back(numMatches).end());
5230 rewriter.
replaceOp(delinearizeOp, mergedResults);
5248struct SplitDelinearizeSpanningLastLinearizeArg final
5249 : OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5252 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5253 PatternRewriter &rewriter)
const override {
5254 auto linearizeOp = delinearizeOp.getLinearIndex()
5255 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5258 "index doesn't come from linearize");
5260 if (!linearizeOp.getDisjoint())
5262 "linearize isn't disjoint");
5267 if (linearizeOp.getStaticBasis().empty())
5269 linearizeOp,
"linearize has no basis elements (no inputs)");
5271 int64_t
target = linearizeOp.getStaticBasis().back();
5272 if (ShapedType::isDynamic(
target))
5274 linearizeOp,
"linearize ends with dynamic basis value");
5276 int64_t sizeToSplit = 1;
5277 size_t elemsToSplit = 0;
5278 ArrayRef<int64_t> basis = delinearizeOp.getStaticBasis();
5279 for (int64_t basisElem : llvm::reverse(basis)) {
5280 if (ShapedType::isDynamic(basisElem))
5282 delinearizeOp,
"dynamic basis element while scanning for split");
5283 sizeToSplit *= basisElem;
5286 if (sizeToSplit >
target)
5288 "overshot last argument size");
5289 if (sizeToSplit ==
target)
5293 if (sizeToSplit <
target)
5295 delinearizeOp,
"product of known basis elements doesn't exceed last "
5296 "linearize argument");
5298 if (elemsToSplit < 2)
5301 "need at least two elements to form the basis product");
5303 Value linearizeWithoutBack = affine::AffineLinearizeIndexOp::create(
5304 rewriter, linearizeOp.getLoc(), linearizeOp.getLinearIndex().getType(),
5305 linearizeOp.getMultiIndex().drop_back(), linearizeOp.getDynamicBasis(),
5306 linearizeOp.getStaticBasis().drop_back(), linearizeOp.getDisjoint());
5307 auto delinearizeWithoutSplitPart = affine::AffineDelinearizeIndexOp::create(
5308 rewriter, delinearizeOp.getLoc(), linearizeWithoutBack,
5309 delinearizeOp.getDynamicBasis(), basis.drop_back(elemsToSplit),
5310 delinearizeOp.hasOuterBound());
5311 auto delinearizeBack = affine::AffineDelinearizeIndexOp::create(
5312 rewriter, delinearizeOp.getLoc(), linearizeOp.getMultiIndex().back(),
5313 basis.take_back(elemsToSplit),
true);
5314 SmallVector<Value> results = llvm::to_vector(
5315 llvm::concat<Value>(delinearizeWithoutSplitPart.getResults(),
5316 delinearizeBack.getResults()));
5317 rewriter.
replaceOp(delinearizeOp, results);
5324void affine::AffineDelinearizeIndexOp::getCanonicalizationPatterns(
5325 RewritePatternSet &patterns, MLIRContext *context) {
5327 .
insert<CancelDelinearizeOfLinearizeDisjointExactTail,
5328 DropUnitExtentBasis, SplitDelinearizeSpanningLastLinearizeArg>(
5339 if (multiIndex.empty())
5340 return IndexType::get(ctx);
5341 return multiIndex.front().
getType();
5344void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5345 OperationState &odsState,
5348 if (!basis.empty() && basis.front() == Value())
5349 basis = basis.drop_front();
5350 SmallVector<Value> dynamicBasis;
5351 SmallVector<int64_t> staticBasis;
5355 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5359void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5360 OperationState &odsState,
5362 ArrayRef<OpFoldResult> basis,
5364 if (!basis.empty() && basis.front() == OpFoldResult())
5365 basis = basis.drop_front();
5366 SmallVector<Value> dynamicBasis;
5367 SmallVector<int64_t> staticBasis;
5370 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5374void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5375 OperationState &odsState,
5377 ArrayRef<int64_t> basis,
bool disjoint) {
5379 build(odsBuilder, odsState, resultType, multiIndex,
ValueRange{}, basis,
5383LogicalResult AffineLinearizeIndexOp::verify() {
5384 size_t numIndexes = getMultiIndex().size();
5385 size_t numBasisElems = getStaticBasis().size();
5386 if (numIndexes != numBasisElems && numIndexes != numBasisElems + 1)
5387 return emitOpError(
"should be passed a basis element for each index except "
5388 "possibly the first");
5390 auto dynamicMarkersCount =
5391 llvm::count_if(getStaticBasis(), ShapedType::isDynamic);
5392 if (
static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5394 "mismatch between dynamic and static basis (kDynamic marker but no "
5395 "corresponding dynamic basis entry) -- this can only happen due to an "
5396 "incorrect fold/rewrite");
5401OpFoldResult AffineLinearizeIndexOp::fold(FoldAdaptor adaptor) {
5402 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5404 adaptor.getDynamicBasis());
5405 if (maybeStaticBasis) {
5406 setStaticBasis(*maybeStaticBasis);
5410 if (getMultiIndex().empty())
5411 return IntegerAttr::get(getResult().
getType(), 0);
5414 if (getMultiIndex().size() == 1)
5415 return getMultiIndex().front();
5420 if (llvm::any_of(adaptor.getMultiIndex(), [](Attribute a) {
5421 return !isa_and_nonnull<IntegerAttr>(a);
5425 if (!adaptor.getDynamicBasis().empty())
5430 for (
auto [length, indexAttr] :
5431 llvm::zip_first(llvm::reverse(getStaticBasis()),
5432 llvm::reverse(adaptor.getMultiIndex()))) {
5433 result =
result + cast<IntegerAttr>(indexAttr).getInt() * stride;
5434 stride = stride * length;
5437 if (!hasOuterBound())
5440 cast<IntegerAttr>(adaptor.getMultiIndex().front()).getInt() * stride;
5445SmallVector<OpFoldResult> AffineLinearizeIndexOp::getEffectiveBasis() {
5447 if (hasOuterBound()) {
5448 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5450 getDynamicBasis().drop_front(), builder);
5452 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5456 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5459SmallVector<OpFoldResult> AffineLinearizeIndexOp::getPaddedBasis() {
5460 SmallVector<OpFoldResult> ret = getMixedBasis();
5461 if (!hasOuterBound())
5462 ret.insert(ret.begin(), OpFoldResult());
5477struct DropLinearizeUnitComponentsIfDisjointOrZero final
5478 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5481 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5482 PatternRewriter &rewriter)
const override {
5484 size_t numIndices = multiIndex.size();
5485 SmallVector<Value> newIndices;
5486 newIndices.reserve(numIndices);
5487 SmallVector<OpFoldResult> newBasis;
5488 newBasis.reserve(numIndices);
5490 if (!op.hasOuterBound()) {
5491 newIndices.push_back(multiIndex.front());
5492 multiIndex = multiIndex.drop_front();
5495 SmallVector<OpFoldResult> basis = op.getMixedBasis();
5496 for (
auto [index, basisElem] : llvm::zip_equal(multiIndex, basis)) {
5498 if (!basisEntry || *basisEntry != 1) {
5499 newIndices.push_back(index);
5500 newBasis.push_back(basisElem);
5505 if (!op.getDisjoint() && (!indexValue || *indexValue != 0)) {
5506 newIndices.push_back(index);
5507 newBasis.push_back(basisElem);
5511 if (newIndices.size() == numIndices)
5513 "no unit basis entries to replace");
5515 if (newIndices.empty()) {
5517 op, rewriter.
getZeroAttr(op.getLinearIndex().getType()));
5521 op, newIndices, newBasis, op.getDisjoint());
5527 ArrayRef<OpFoldResult> terms) {
5528 int64_t nDynamic = 0;
5529 SmallVector<Value> dynamicPart;
5531 for (OpFoldResult term : terms) {
5538 dynamicPart.push_back(cast<Value>(term));
5542 if (
auto constant = dyn_cast<AffineConstantExpr>(
result))
5544 return AffineApplyOp::create(builder, loc,
result, dynamicPart).getResult();
5574struct CancelLinearizeOfDelinearizePortion final
5575 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5585 unsigned linStart = 0;
5586 unsigned delinStart = 0;
5587 unsigned length = 0;
5591 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp linearizeOp,
5592 PatternRewriter &rewriter)
const override {
5593 SmallVector<Match> matches;
5595 const SmallVector<OpFoldResult> linBasis = linearizeOp.getPaddedBasis();
5596 ArrayRef<OpFoldResult> linBasisRef = linBasis;
5598 ValueRange multiIndex = linearizeOp.getMultiIndex();
5599 unsigned numLinArgs = multiIndex.size();
5600 unsigned linArgIdx = 0;
5603 llvm::SmallPtrSet<Operation *, 2> alreadyMatchedDelinearize;
5604 while (linArgIdx < numLinArgs) {
5605 auto asResult = dyn_cast<OpResult>(multiIndex[linArgIdx]);
5611 auto delinearizeOp =
5612 dyn_cast<AffineDelinearizeIndexOp>(asResult.getOwner());
5613 if (!delinearizeOp) {
5630 unsigned delinArgIdx = asResult.getResultNumber();
5631 SmallVector<OpFoldResult> delinBasis = delinearizeOp.getPaddedBasis();
5632 OpFoldResult firstDelinBound = delinBasis[delinArgIdx];
5633 OpFoldResult firstLinBound = linBasis[linArgIdx];
5634 bool boundsMatch = firstDelinBound == firstLinBound;
5635 bool bothAtFront = linArgIdx == 0 && delinArgIdx == 0;
5636 bool knownByDisjoint =
5637 linearizeOp.getDisjoint() && delinArgIdx == 0 && !firstDelinBound;
5638 if (!boundsMatch && !bothAtFront && !knownByDisjoint) {
5644 unsigned numDelinOuts = delinearizeOp.getNumResults();
5645 for (; j + linArgIdx < numLinArgs && j + delinArgIdx < numDelinOuts;
5647 if (multiIndex[linArgIdx + j] !=
5648 delinearizeOp.getResult(delinArgIdx + j))
5650 if (linBasis[linArgIdx + j] != delinBasis[delinArgIdx + j])
5656 if (j <= 1 || !alreadyMatchedDelinearize.insert(delinearizeOp).second) {
5660 matches.push_back(Match{delinearizeOp, linArgIdx, delinArgIdx, j});
5664 if (matches.empty())
5666 linearizeOp,
"no run of delinearize outputs to deal with");
5671 SmallVector<SmallVector<Value>> delinearizeReplacements;
5673 SmallVector<Value> newIndex;
5674 newIndex.reserve(numLinArgs);
5675 SmallVector<OpFoldResult> newBasis;
5676 newBasis.reserve(numLinArgs);
5677 unsigned prevMatchEnd = 0;
5678 for (Match m : matches) {
5679 unsigned gap = m.linStart - prevMatchEnd;
5680 llvm::append_range(newIndex, multiIndex.slice(prevMatchEnd, gap));
5681 llvm::append_range(newBasis, linBasisRef.slice(prevMatchEnd, gap));
5683 prevMatchEnd = m.linStart + m.length;
5685 PatternRewriter::InsertionGuard g(rewriter);
5688 ArrayRef<OpFoldResult> basisToMerge =
5689 linBasisRef.slice(m.linStart, m.length);
5692 OpFoldResult newSize =
5697 newIndex.push_back(m.delinearize.getLinearIndex());
5698 newBasis.push_back(newSize);
5700 delinearizeReplacements.push_back(SmallVector<Value>());
5704 SmallVector<Value> newDelinResults;
5705 SmallVector<OpFoldResult> newDelinBasis = m.delinearize.getPaddedBasis();
5706 newDelinBasis.erase(newDelinBasis.begin() + m.delinStart,
5707 newDelinBasis.begin() + m.delinStart + m.length);
5708 newDelinBasis.insert(newDelinBasis.begin() + m.delinStart, newSize);
5709 auto newDelinearize = AffineDelinearizeIndexOp::create(
5710 rewriter, m.delinearize.getLoc(), m.delinearize.getLinearIndex(),
5716 Value combinedElem = newDelinearize.getResult(m.delinStart);
5717 auto residualDelinearize = AffineDelinearizeIndexOp::create(
5718 rewriter, m.delinearize.getLoc(), combinedElem, basisToMerge);
5723 llvm::append_range(newDelinResults,
5724 newDelinearize.getResults().take_front(m.delinStart));
5725 llvm::append_range(newDelinResults, residualDelinearize.getResults());
5728 newDelinearize.getResults().drop_front(m.delinStart + 1));
5730 delinearizeReplacements.push_back(newDelinResults);
5731 newIndex.push_back(combinedElem);
5732 newBasis.push_back(newSize);
5734 llvm::append_range(newIndex, multiIndex.drop_front(prevMatchEnd));
5735 llvm::append_range(newBasis, linBasisRef.drop_front(prevMatchEnd));
5737 linearizeOp, newIndex, newBasis, linearizeOp.getDisjoint());
5739 for (
auto [m, newResults] :
5740 llvm::zip_equal(matches, delinearizeReplacements)) {
5741 if (newResults.empty())
5743 rewriter.
replaceOp(m.delinearize, newResults);
5754struct DropLinearizeLeadingZero final
5755 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5758 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5759 PatternRewriter &rewriter)
const override {
5760 Value leadingIdx = op.getMultiIndex().front();
5764 if (op.getMultiIndex().size() == 1) {
5769 SmallVector<OpFoldResult> mixedBasis = op.getMixedBasis();
5770 ArrayRef<OpFoldResult> newMixedBasis = mixedBasis;
5771 if (op.hasOuterBound())
5772 newMixedBasis = newMixedBasis.drop_front();
5775 op, op.getMultiIndex().drop_front(), newMixedBasis, op.getDisjoint());
5781void affine::AffineLinearizeIndexOp::getCanonicalizationPatterns(
5782 RewritePatternSet &patterns, MLIRContext *context) {
5783 patterns.
add<CancelLinearizeOfDelinearizePortion, DropLinearizeLeadingZero,
5784 DropLinearizeUnitComponentsIfDisjointOrZero>(context);
5791#define GET_OP_CLASSES
5792#include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
static AffineForOp buildAffineLoopFromConstants(OpBuilder &builder, Location loc, int64_t lb, int64_t ub, int64_t step, AffineForOp::BodyBuilderFn bodyBuilderFn)
Creates an affine loop from the bounds known to be constants.
static bool hasTrivialZeroTripCount(AffineForOp op)
Returns true if the affine.for has zero iterations in trivial cases.
static Type inferIndexType(MLIRContext *ctx, ValueRange multiIndex)
Infer the index type from a set of multi-index values. Returns the common type (index or vector<....
static LogicalResult verifyMemoryOpIndexing(AffineMemOpTy op, AffineMapAttr mapAttr, Operation::operand_range mapOperands, MemRefType memrefType, unsigned numIndexOperands)
Verify common indexing invariants of affine.load, affine.store, affine.vector_load and affine....
static void printAffineMinMaxOp(OpAsmPrinter &p, T op)
static bool isResultTypeMatchAtomicRMWKind(Type resultType, arith::AtomicRMWKind op)
static bool remainsLegalAfterInline(Value value, Region *src, Region *dest, const IRMapping &mapping, function_ref< bool(Value, Region *)> legalityCheck)
Checks if value known to be a legal affine dimension or symbol in src region remains legal if the ope...
static void printMinMaxBound(OpAsmPrinter &p, AffineMapAttr mapAttr, DenseIntElementsAttr group, ValueRange operands, StringRef keyword)
Prints a lower(upper) bound of an affine parallel loop with max(min) conditions in it.
static OpFoldResult foldMinMaxOp(T op, ArrayRef< Attribute > operands)
Fold an affine min or max operation with the given operands.
static bool isTopLevelValueOrAbove(Value value, Region *region)
A utility function to check if a value is defined at the top level of region or is an argument of reg...
static LogicalResult canonicalizeLoopBounds(AffineForOp forOp)
Canonicalize the bounds of the given loop.
static void simplifyExprAndOperands(AffineExpr &expr, unsigned numDims, unsigned numSymbols, ArrayRef< Value > operands)
Simplify expr while exploiting information from the values in operands.
static bool isValidAffineIndexOperand(Value value, Region *region)
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet, SmallVectorImpl< Value > *operands)
static ParseResult parseBound(bool isLower, OperationState &result, OpAsmParser &p)
Parse a for operation loop bounds.
static std::optional< SmallVector< int64_t > > foldCstValueToCstAttrBasis(ArrayRef< OpFoldResult > mixedBasis, MutableOperandRange mutableDynamicBasis, ArrayRef< Attribute > dynamicBasis)
Given mixed basis of affine.delinearize_index/linearize_index replace constant SSA values with the co...
static void canonicalizePromotedSymbols(MapOrSet *mapOrSet, SmallVectorImpl< Value > *operands)
static void simplifyMinOrMaxExprWithOperands(AffineMap &map, ArrayRef< Value > operands, bool isMax)
Simplify the expressions in map while making use of lower or upper bounds of its operands.
static ParseResult parseAffineMinMaxOp(OpAsmParser &parser, OperationState &result)
static LogicalResult replaceAffineDelinearizeIndexInverseExpression(AffineDelinearizeIndexOp delinOp, Value resultToReplace, AffineMap *map, SmallVectorImpl< Value > &dims, SmallVectorImpl< Value > &syms)
If this map contains of the expression x_1 + x_1 * C_1 + ... x_n * C_N + / ... (not necessarily in or...
static void composeSetAndOperands(IntegerSet &set, SmallVectorImpl< Value > &operands, bool composeAffineMin=false)
Compose any affine.apply ops feeding into operands of the integer set set by composing the maps of su...
static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, unsigned index, Region *region)
Returns true if the 'index' dimension of the memref defined by memrefDefOp is a statically shaped one...
static bool isNonNegativeBoundedBy(AffineExpr e, ArrayRef< Value > operands, int64_t k)
Check if e is known to be: 0 <= e < k.
static AffineForOp buildAffineLoopFromValues(OpBuilder &builder, Location loc, Value lb, Value ub, int64_t step, AffineForOp::BodyBuilderFn bodyBuilderFn)
Creates an affine loop from the bounds that may or may not be constants.
static void simplifyMapWithOperands(AffineMap &map, ArrayRef< Value > operands)
Simplify the map while exploiting information on the values in operands.
static void printDimAndSymbolList(Operation::operand_iterator begin, Operation::operand_iterator end, unsigned numDims, OpAsmPrinter &printer)
Prints dimension and symbol list.
static int64_t getLargestKnownDivisor(AffineExpr e, ArrayRef< Value > operands)
Returns the largest known divisor of e.
static void composeAffineMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands, bool composeAffineMin=false)
Iterate over operands and fold away all those produced by an AffineApplyOp iteratively.
static void legalizeDemotedDims(MapOrSet &mapOrSet, SmallVectorImpl< Value > &operands)
A valid affine dimension may appear as a symbol in affine.apply operations.
static OpTy makeComposedMinMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
static std::optional< int64_t > getUpperBound(Value iv)
Gets the constant upper bound on an affine.for iv.
static void buildAffineLoopNestImpl(OpBuilder &builder, Location loc, BoundListTy lbs, BoundListTy ubs, ArrayRef< int64_t > steps, function_ref< void(OpBuilder &, Location, ValueRange)> bodyBuilderFn, LoopCreatorTy &&loopCreatorFn)
Builds an affine loop nest, using "loopCreatorFn" to create individual loop operations.
static LogicalResult foldLoopBounds(AffineForOp forOp)
Fold the constant bounds of a loop.
static LogicalResult replaceAffineMinBoundingBoxExpression(AffineMinOp minOp, AffineExpr dimOrSym, AffineMap *map, ValueRange dims, ValueRange syms)
Assuming dimOrSym is a quantity in the apply op map map and defined by minOp = affine_min(x_1,...
static void addAlignmentAttr(OpBuilder &builder, OperationState &result, StringAttr attrName, llvm::MaybeAlign alignment)
Adds the optional alignment attribute to result, if one is given.
static SmallVector< OpFoldResult > AffineForEmptyLoopFolder(AffineForOp forOp)
Fold the empty loop.
static LogicalResult verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands, unsigned numDims)
Utility function to verify that a set of operands are valid dimension and symbol identifiers.
static OpFoldResult makeComposedFoldedMinMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
static bool isDimOpValidSymbol(ShapedDimOpInterface dimOp, Region *region)
Returns true if the result of the dim op is a valid symbol for region.
static bool isQTimesDPlusR(AffineExpr e, ArrayRef< Value > operands, int64_t &div, AffineExpr "ientTimesDiv, AffineExpr &rem)
Check if expression e is of the form d*e_1 + e_2 where 0 <= e_2 < d.
static LogicalResult replaceDimOrSym(AffineMap *map, unsigned dimOrSymbolPosition, SmallVectorImpl< Value > &dims, SmallVectorImpl< Value > &syms, bool replaceAffineMin)
Replace all occurrences of AffineExpr at position pos in map by the defining AffineApplyOp expression...
static std::optional< int64_t > getLowerBound(Value iv)
Gets the constant lower bound on an iv.
static std::optional< uint64_t > getTrivialConstantTripCount(AffineForOp forOp)
Returns constant trip count in trivial cases.
static LogicalResult verifyAffineMinMaxOp(T op)
static void printBound(AffineMapAttr boundMap, Operation::operand_range boundOperands, const char *prefix, OpAsmPrinter &p)
static void shortenAddChainsContainingAll(AffineExpr e, const llvm::SmallDenseSet< AffineExpr, 4 > &exprsToRemove, AffineExpr newVal, DenseMap< AffineExpr, AffineExpr > &replacementsMap)
Recursively traverse e.
static void composeMultiResultAffineMap(AffineMap &map, SmallVectorImpl< Value > &operands, bool composeAffineMin=false)
Composes the given affine map with the given list of operands, pulling in the maps from any affine....
static LogicalResult canonicalizeMapExprAndTermOrder(AffineMap &map)
Canonicalize the result expression order of an affine map and return success if the order changed.
static Value getZero(OpBuilder &b, Location loc, Type elementType)
Get zero value for an element type.
static Value getMemRef(Operation *memOp)
Returns the memref being read/written by a memref/affine load/store op.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static Operation::operand_range getLowerBoundOperands(AffineForOp forOp)
static Operation::operand_range getUpperBoundOperands(AffineForOp forOp)
static VectorType getVectorType(Type scalarTy, const VectorizationStrategy *strategy)
Returns the vector type resulting from applying the provided vectorization strategy on the scalar typ...
RetTy walkPostOrder(AffineExpr expr)
Base type for affine expression.
AffineExpr shiftDims(unsigned numDims, unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
AffineExpr shiftSymbols(unsigned numSymbols, unsigned shift, unsigned offset=0) const
Replace symbols[offset ... numSymbols) by symbols[offset + shift ... shift + numSymbols).
AffineExpr floorDiv(uint64_t v) const
AffineExprKind getKind() const
Return the classification for this type.
int64_t getLargestKnownDivisor() const
Returns the greatest known integral divisor of this affine expression.
MLIRContext * getContext() const
AffineExpr replace(AffineExpr expr, AffineExpr replacement) const
Sparse replace method.
AffineExpr ceilDiv(uint64_t v) 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.
MLIRContext * getContext() const
bool isFunctionOfDim(unsigned position) const
Return true if any affine expression involves AffineDimExpr position.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
AffineMap shiftDims(unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
bool isFunctionOfSymbol(unsigned position) const
Return true if any affine expression involves AffineSymbolExpr position.
unsigned getNumResults() const
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...
AffineMap replaceDimsAndSymbols(ArrayRef< AffineExpr > dimReplacements, ArrayRef< AffineExpr > symReplacements, unsigned numResultDims, unsigned numResultSyms) const
This method substitutes any uses of dimensions and symbols (e.g.
unsigned getNumInputs() const
AffineMap shiftSymbols(unsigned shift, unsigned offset=0) const
Replace symbols[offset ... numSymbols) by symbols[offset + shift ... shift + numSymbols).
AffineExpr getResult(unsigned idx) const
AffineMap replace(AffineExpr expr, AffineExpr replacement, unsigned numResultDims, unsigned numResultSyms) const
Sparse replace method.
static AffineMap getConstantMap(int64_t val, MLIRContext *context)
Returns a single constant result affine map.
AffineMap getSubMap(ArrayRef< unsigned > resultPos) const
Returns the map consisting of the resultPos subset.
LogicalResult constantFold(ArrayRef< Attribute > operandConstants, SmallVectorImpl< Attribute > &results, bool *hasPoison=nullptr) const
Folds the results of the application of an affine map on the provided operands to a constant if possi...
@ Paren
Parens surrounding zero or more operands.
@ OptionalSquare
Square brackets supporting zero or more ops, or nothing.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
ParseResult addTypeToList(Type type, SmallVectorImpl< Type > &result)
Add the specified type to the end of the specified type list and return success.
virtual ParseResult parseOptionalRParen()=0
Parse a ) token if present.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
virtual ParseResult parseArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an arrow followed by a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
void printOptionalArrowTypeList(TypeRange &&types)
Print an optional arrow followed by a type list.
Attributes are known-constant values of operations.
This class represents an argument of a Block.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
BlockArgListType getArguments()
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
IntegerAttr getIntegerAttr(Type type, int64_t value)
AffineMap getDimIdentityMap()
AffineMap getMultiDimIdentityMap(unsigned rank)
AffineExpr getAffineSymbolExpr(unsigned position)
AffineExpr getAffineConstantExpr(int64_t constant)
DenseIntElementsAttr getI32TensorAttr(ArrayRef< int32_t > values)
Tensor-typed DenseIntElementsAttr getters.
IntegerAttr getI64IntegerAttr(int64_t value)
IntegerType getIntegerType(unsigned width)
BoolAttr getBoolAttr(bool value)
AffineMap getEmptyAffineMap()
Returns a zero result affine map with no dimensions or symbols: () -> ().
TypedAttr getZeroAttr(Type type)
AffineMap getConstantAffineMap(int64_t val)
Returns a single constant result affine map with 0 dimensions and 0 symbols.
AffineMap getSymbolIdentityMap()
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
An attribute that represents a reference to a dense integer vector or tensor object.
This is a utility class for mapping one set of IR entities to another.
auto lookup(T from) const
Lookup a mapped value within the map.
An integer set representing a conjunction of one or more affine equalities and inequalities.
unsigned getNumDims() const
static IntegerSet get(unsigned dimCount, unsigned symbolCount, ArrayRef< AffineExpr > constraints, ArrayRef< bool > eqFlags)
MLIRContext * getContext() const
unsigned getNumInputs() const
ArrayRef< AffineExpr > getConstraints() const
ArrayRef< bool > getEqFlags() const
Returns the equality bits, which specify whether each of the constraints is an equality or inequality...
unsigned getNumSymbols() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class provides a mutable adaptor for a range of operands.
void erase(unsigned subStart, unsigned subLen=1)
Erase the operands within the given sub-range.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgument(Argument &result, bool allowType=false, bool allowAttrs=false)=0
Parse a single argument with the following syntax:
ParseResult parseTrailingOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None)
Parse zero or more trailing SSA comma-separated trailing operand references with a specified surround...
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
virtual ParseResult parseAffineMapOfSSAIds(SmallVectorImpl< UnresolvedOperand > &operands, Attribute &map, StringRef attrName, NamedAttrList &attrs, Delimiter delimiter=Delimiter::Square)=0
Parses an affine map attribute where dims and symbols are SSA operands.
ParseResult parseAssignmentList(SmallVectorImpl< Argument > &lhs, SmallVectorImpl< UnresolvedOperand > &rhs)
Parse a list of assignments of the form (x1 = y1, x2 = y2, ...)
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseAffineExprOfSSAIds(SmallVectorImpl< UnresolvedOperand > &dimOperands, SmallVectorImpl< UnresolvedOperand > &symbOperands, AffineExpr &expr)=0
Parses an affine expression where dims and symbols are SSA operands.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printAffineExprOfSSAIds(AffineExpr expr, ValueRange dimOperands, ValueRange symOperands)=0
Prints an affine expression of SSA ids with SSA id names used instead of dims and symbols.
virtual void printAffineMapOfSSAIds(AffineMapAttr mapAttr, ValueRange operands)=0
Prints an affine map of SSA ids, where SSA id names are used in place of dims/symbols.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
virtual void printRegionArgument(BlockArgument arg, ArrayRef< NamedAttribute > argAttrs={}, bool omitType=false)=0
Print a block argument in the usual format of: ssaName : type {attr1=42} loc("here") where location p...
virtual void printOperand(Value value)=0
Print implementations for various things an operation contains.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents a single result from folding an operation.
A trait of region holding operations that defines a new scope for polyhedral optimization purposes.
This class provides the API for ops that are known to be isolated from above.
This class implements the operand iterators for the Operation class.
Operation is the basic unit of execution within MLIR.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
OperandRange operand_range
operand_range getOperands()
Returns an iterator on the underlying Value's.
Region * getParentRegion()
Returns the region to which the instruction belongs.
operand_range::iterator operand_iterator
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
RegionBranchTerminatorOpInterface getTerminatorPredecessorOrNull() const
Returns the terminator if branching from a region.
This class represents a successor of a region.
Region * getSuccessor() const
Return the given region successor.
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Operation * getParentOp()
Return the parent operation this region is attached to.
bool hasOneBlock()
Return true if this region has exactly one block.
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void eraseBlock(Block *block)
This method erases all operations in a block.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
virtual void replaceUsesWithIf(Value from, Value to, function_ref< bool(OpOperand &)> functor, bool *allUsesReplaced=nullptr)
Find uses of from and replace them with to if the functor returns true.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
void mergeBlocks(Block *source, Block *dest, ValueRange argValues={})
Inline the operations of block 'source' into the end of block 'dest'.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void startOpModification(Operation *op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class represents a specific instance of an effect.
static DerivedEffect * get()
static DefaultResource * get()
std::vector< SmallVector< int64_t, 8 > > operandExprStack
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
A variable that can be added to the constraint set as a "column".
static bool compare(const Variable &lhs, ComparisonOperator cmp, const Variable &rhs)
Return "true" if "lhs cmp rhs" was proven to hold.
This class provides an abstraction over the different types of ranges over Values.
type_range getType() const
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.
Region * getParentRegion()
Return the Region in which this Value is defined.
AffineBound represents a lower or upper bound in the for operation.
An AffineValueMap is an affine map plus its ML value operands and results for analysis purposes.
LogicalResult canonicalize()
Attempts to canonicalize the map and operands.
ArrayRef< Value > getOperands() const
AffineExpr getResult(unsigned i)
AffineMap getAffineMap() const
void reset(AffineMap map, ValueRange operands, ValueRange results={})
unsigned getNumResults() const
static void difference(const AffineValueMap &a, const AffineValueMap &b, AffineValueMap *res)
Return the value map that is the difference of value maps 'a' and 'b', represented as an affine map a...
Operation * getOwner() const
Return the owner of this operand.
constexpr auto RecursivelySpeculatable
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto NotSpeculatable
void buildAffineLoopNest(OpBuilder &builder, Location loc, ArrayRef< int64_t > lbs, ArrayRef< int64_t > ubs, ArrayRef< int64_t > steps, function_ref< void(OpBuilder &, Location, ValueRange)> bodyBuilderFn=nullptr)
Builds a perfect nest of affine.for loops, i.e., each loop except the innermost one contains only ano...
AffineApplyOp makeComposedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Returns a composed AffineApplyOp by composing map and operands with other AffineApplyOps supplying th...
void extractForInductionVars(ArrayRef< AffineForOp > forInsts, SmallVectorImpl< Value > *ivs)
Extracts the induction variables from a list of AffineForOps and places them in the output argument i...
bool isValidDim(Value value)
Returns true if the given Value can be used as a dimension id in the region of the closest surroundin...
bool isAffineInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp or AffineParallelOp.
SmallVector< OpFoldResult > makeComposedFoldedMultiResultAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Variant of makeComposedFoldedAffineApply suitable for multi-result maps.
OpFoldResult computeProduct(Location loc, OpBuilder &builder, ArrayRef< OpFoldResult > terms)
Return the product of terms, creating an affine.apply if any of them are non-constant values.
AffineForOp getForInductionVarOwner(Value val)
Returns the loop parent of an induction variable.
void canonicalizeMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands)
Modifies both map and operands in-place so as to:
OpFoldResult makeComposedFoldedAffineMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Constructs an AffineMinOp that computes a maximum across the results of applying map to operands,...
bool isAffineForInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp.
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
OpFoldResult makeComposedFoldedAffineMin(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Constructs an AffineMinOp that computes a minimum across the results of applying map to operands,...
bool isTopLevelValue(Value value)
A utility function to check if a value is defined at the top level of an op with trait AffineScope or...
Region * getAffineAnalysisScope(Operation *op)
Returns the closest region enclosing op that is held by a non-affine operation; nullptr if there is n...
void fullyComposeAffineMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands, bool composeAffineMin=false)
Given an affine map map and its input operands, this method composes into map, maps of AffineApplyOps...
void canonicalizeSetAndOperands(IntegerSet *set, SmallVectorImpl< Value > *operands)
Canonicalizes an integer set the same way canonicalizeMapAndOperands does for affine maps.
void extractInductionVars(ArrayRef< Operation * > affineOps, SmallVectorImpl< Value > &ivs)
Extracts the induction variables from a list of either AffineForOp or AffineParallelOp and places the...
bool isValidSymbol(Value value)
Returns true if the given value can be used as a symbol in the region of the closest surrounding op t...
AffineParallelOp getAffineParallelInductionVarOwner(Value val)
Returns true if the provided value is among the induction variables of an AffineParallelOp.
Region * getAffineScope(Operation *op)
Returns the closest region enclosing op that is held by an operation with trait AffineScope; nullptr ...
ParseResult parseDimAndSymbolList(OpAsmParser &parser, SmallVectorImpl< Value > &operands, unsigned &numDims)
Parses dimension and symbol list.
bool isAffineParallelInductionVar(Value val)
Returns true if val is the induction variable of an AffineParallelOp.
AffineMinOp makeComposedAffineMin(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Returns an AffineMinOp obtained by composing map and operands with AffineApplyOps supplying those ope...
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
MemRefType getMemRefType(T &&t)
Convenience method to abbreviate casting getType().
Include the generated interface declarations.
AffineMap simplifyAffineMap(AffineMap map)
Simplifies an affine map by simplifying its underlying AffineExpr results.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
OpFoldResult getAsIndexOpFoldResult(MLIRContext *ctx, int64_t val)
Convert int64_t to integer attributes of index type and return them as OpFoldResult.
AffineMap removeDuplicateExprs(AffineMap map)
Returns a map with the same dimension and symbol count as map, but whose results are the unique affin...
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
std::function< SmallVector< Value >( OpBuilder &b, Location loc, ArrayRef< BlockArgument > newBbArgs)> NewYieldValuesFn
A function that returns the additional yielded values during replaceWithAdditionalYields.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
std::optional< int64_t > getBoundForAffineExpr(AffineExpr expr, unsigned numDims, unsigned numSymbols, ArrayRef< std::optional< int64_t > > constLowerBounds, ArrayRef< std::optional< int64_t > > constUpperBounds, bool isUpper)
Get a lower or upper (depending on isUpper) bound for expr while using the constant lower and upper b...
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isPure(Operation *op)
Returns true if the given operation is pure, i.e., is speculatable that does not touch memory.
@ CeilDiv
RHS of ceildiv is always a constant or a symbolic expression.
@ Mod
RHS of mod is always a constant or a symbolic expression with a positive value.
@ DimId
Dimensional identifier.
@ FloorDiv
RHS of floordiv is always a constant or a symbolic expression.
@ SymbolId
Symbolic identifier.
AffineExpr getAffineBinaryOpExpr(AffineExprKind kind, AffineExpr lhs, AffineExpr rhs)
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
llvm::TypeSwitch< T, ResultT > TypeSwitch
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
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.
AffineMap foldAttributesIntoMap(Builder &b, AffineMap map, ArrayRef< OpFoldResult > operands, SmallVector< Value > &remainingValues)
Fold all attributes among the given operands into the affine map.
llvm::function_ref< Fn > function_ref
AffineExpr getAffineSymbolExpr(unsigned position, MLIRContext *context)
Canonicalize the affine map result expression order of an affine min/max operation.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
Remove duplicated expressions in affine min/max ops.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
Merge an affine min/max op to its consumers if its consumer is also an affine min/max op.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
This is the representation of an operand reference.
This class represents a listener that may be used to hook into various actions within an OpBuilder.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.