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;
778 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr))
779 return constExpr.getValue();
783 constLowerBounds.reserve(operands.size());
784 constUpperBounds.reserve(operands.size());
785 for (
Value operand : operands) {
801 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr))
802 return constExpr.getValue();
806 constLowerBounds.reserve(operands.size());
807 constUpperBounds.reserve(operands.size());
808 for (
Value operand : operands) {
823 auto binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
834 binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
842 lhs = binExpr.getLHS();
843 rhs = binExpr.getRHS();
844 auto rhsConst = dyn_cast<AffineConstantExpr>(
rhs);
848 int64_t rhsConstVal = rhsConst.getValue();
850 if (rhsConstVal <= 0)
855 std::optional<int64_t> lhsLbConst =
857 std::optional<int64_t> lhsUbConst =
859 if (lhsLbConst && lhsUbConst) {
860 int64_t lhsLbConstVal = *lhsLbConst;
861 int64_t lhsUbConstVal = *lhsUbConst;
865 divideFloorSigned(lhsLbConstVal, rhsConstVal) ==
866 divideFloorSigned(lhsUbConstVal, rhsConstVal)) {
868 divideFloorSigned(lhsLbConstVal, rhsConstVal), context);
874 divideCeilSigned(lhsLbConstVal, rhsConstVal) ==
875 divideCeilSigned(lhsUbConstVal, rhsConstVal)) {
882 lhsLbConstVal < rhsConstVal && lhsUbConstVal < rhsConstVal) {
895 if (rhsConstVal % divisor == 0 &&
897 expr = quotientTimesDiv.
floorDiv(rhsConst);
898 }
else if (divisor % rhsConstVal == 0 &&
900 expr =
rem % rhsConst;
926 if (operands.empty())
932 constLowerBounds.reserve(operands.size());
933 constUpperBounds.reserve(operands.size());
934 for (
Value operand : operands) {
948 if (
auto constExpr = dyn_cast<AffineConstantExpr>(e)) {
949 lowerBounds.push_back(constExpr.getValue());
950 upperBounds.push_back(constExpr.getValue());
952 lowerBounds.push_back(
954 constLowerBounds, constUpperBounds,
956 upperBounds.push_back(
958 constLowerBounds, constUpperBounds,
965 for (
auto exprEn : llvm::enumerate(map.
getResults())) {
967 unsigned i = exprEn.index();
969 if (lowerBounds[i] && upperBounds[i] && *lowerBounds[i] == *upperBounds[i])
974 if (!upperBounds[i]) {
975 irredundantExprs.push_back(e);
980 if (!llvm::any_of(llvm::enumerate(lowerBounds), [&](
const auto &en) {
981 auto otherLowerBound = en.value();
982 unsigned pos = en.index();
983 if (pos == i || !otherLowerBound)
985 if (*otherLowerBound > *upperBounds[i])
987 if (*otherLowerBound < *upperBounds[i])
992 if (upperBounds[pos] && lowerBounds[i] &&
993 lowerBounds[i] == upperBounds[i] &&
994 otherLowerBound == *upperBounds[pos] && i < pos)
998 irredundantExprs.push_back(e);
1000 if (!lowerBounds[i]) {
1001 irredundantExprs.push_back(e);
1005 if (!llvm::any_of(llvm::enumerate(upperBounds), [&](
const auto &en) {
1006 auto otherUpperBound = en.value();
1007 unsigned pos = en.index();
1008 if (pos == i || !otherUpperBound)
1010 if (*otherUpperBound < *lowerBounds[i])
1012 if (*otherUpperBound > *lowerBounds[i])
1014 if (lowerBounds[pos] && upperBounds[i] &&
1015 lowerBounds[i] == upperBounds[i] &&
1016 otherUpperBound == lowerBounds[pos] && i < pos)
1020 irredundantExprs.push_back(e);
1034 assert(map.
getNumInputs() == operands.size() &&
"invalid operands for map");
1040 newResults.push_back(expr);
1063 LDBG() <<
"replaceAffineMinBoundingBoxExpression: `" << minOp <<
"`";
1064 AffineMap affineMinMap = minOp.getAffineMap();
1067 for (
unsigned i = 0, e = affineMinMap.
getNumResults(); i < e; ++i) {
1073 minOp.getOperands())))
1081 for (
auto [i, dim] : llvm::enumerate(minOp.getDimOperands())) {
1082 auto it = llvm::find(dims, dim);
1083 if (it == dims.end()) {
1084 unmappedDims.push_back(i);
1090 for (
auto [i, sym] : llvm::enumerate(minOp.getSymbolOperands())) {
1091 auto it = llvm::find(syms, sym);
1092 if (it == syms.end()) {
1093 unmappedSyms.push_back(i);
1106 if (llvm::any_of(unmappedDims,
1107 [&](
unsigned i) {
return expr.isFunctionOfDim(i); }) ||
1108 llvm::any_of(unmappedSyms,
1109 [&](
unsigned i) {
return expr.isFunctionOfSymbol(i); }))
1115 repl[dimOrSym.
ceilDiv(convertedExpr)] = c1;
1117 repl[(dimOrSym + convertedExpr - 1).floorDiv(convertedExpr)] = c1;
1122 return success(*map != initialMap);
1131 AffineExpr e,
const llvm::SmallDenseSet<AffineExpr, 4> &exprsToRemove,
1133 auto binOp = dyn_cast<AffineBinaryOpExpr>(e);
1144 llvm::SmallDenseSet<AffineExpr, 4> ourTracker(exprsToRemove);
1149 if (!ourTracker.erase(thisTerm)) {
1150 toPreserve.push_back(thisTerm);
1154 auto nextBinOp = dyn_cast_if_present<AffineBinaryOpExpr>(nextTerm);
1156 thisTerm = nextTerm;
1159 thisTerm = nextBinOp.getRHS();
1160 nextTerm = nextBinOp.getLHS();
1163 if (!ourTracker.empty())
1168 for (
AffineExpr preserved : llvm::reverse(toPreserve))
1169 newExpr = newExpr + preserved;
1170 replacementsMap.insert({e, newExpr});
1188 AffineDelinearizeIndexOp delinOp,
Value resultToReplace,
AffineMap *map,
1190 if (!delinOp.getDynamicBasis().empty())
1192 if (resultToReplace != delinOp.getMultiIndex().back())
1197 for (
auto [pos, dim] : llvm::enumerate(dims)) {
1198 auto asResult = dyn_cast_if_present<OpResult>(dim);
1201 if (asResult.getOwner() == delinOp.getOperation())
1204 for (
auto [pos, sym] : llvm::enumerate(syms)) {
1205 auto asResult = dyn_cast_if_present<OpResult>(sym);
1208 if (asResult.getOwner() == delinOp.getOperation())
1211 if (llvm::is_contained(resToExpr,
AffineExpr()))
1214 bool isDimReplacement = llvm::all_of(resToExpr, llvm::IsaPred<AffineDimExpr>);
1216 llvm::SmallDenseSet<AffineExpr, 4> expectedExprs;
1219 for (
auto [binding, size] : llvm::zip(
1220 llvm::reverse(resToExpr), llvm::reverse(delinOp.getStaticBasis()))) {
1224 if (resToExpr.size() != delinOp.getStaticBasis().size())
1225 expectedExprs.insert(resToExpr[0] * stride);
1234 if (replacements.empty())
1238 if (isDimReplacement)
1239 dims.push_back(delinOp.getLinearIndex());
1241 syms.push_back(delinOp.getLinearIndex());
1242 *map = origMap.
replace(replacements, dims.size(), syms.size());
1246 if (
auto d = dyn_cast<AffineDimExpr>(e)) {
1247 unsigned pos = d.getPosition();
1249 dims[pos] =
nullptr;
1251 if (
auto s = dyn_cast<AffineSymbolExpr>(e)) {
1252 unsigned pos = s.getPosition();
1254 syms[pos] =
nullptr;
1273 unsigned dimOrSymbolPosition,
1276 bool replaceAffineMin) {
1278 bool isDimReplacement = (dimOrSymbolPosition < dims.size());
1279 unsigned pos = isDimReplacement ? dimOrSymbolPosition
1280 : dimOrSymbolPosition - dims.size();
1281 Value &v = isDimReplacement ? dims[pos] : syms[pos];
1285 if (
auto minOp = v.
getDefiningOp<AffineMinOp>(); minOp && replaceAffineMin) {
1292 if (
auto delinOp = v.
getDefiningOp<affine::AffineDelinearizeIndexOp>()) {
1306 AffineMap composeMap = affineApply.getAffineMap();
1307 assert(composeMap.
getNumResults() == 1 &&
"affine.apply with >1 results");
1309 affineApply.getMapOperands().end());
1323 dims.append(composeDims.begin(), composeDims.end());
1324 syms.append(composeSyms.begin(), composeSyms.end());
1325 *map = map->
replace(toReplace, replacementExpr, dims.size(), syms.size());
1335 bool composeAffineMin =
false) {
1354 bool changed =
false;
1355 for (
unsigned pos = 0; pos != dims.size() + syms.size(); ++pos)
1368 unsigned nDims = 0, nSyms = 0;
1370 dimReplacements.reserve(dims.size());
1371 symReplacements.reserve(syms.size());
1372 for (
auto *container : {&dims, &syms}) {
1373 bool isDim = (container == &dims);
1374 auto &repls = isDim ? dimReplacements : symReplacements;
1375 for (
const auto &en : llvm::enumerate(*container)) {
1376 Value v = en.value();
1380 "map is function of unexpected expr@pos");
1386 operands->push_back(v);
1399 while (llvm::any_of(*operands, [](
Value v) {
1405 if (composeAffineMin && llvm::any_of(*operands, [](
Value v) {
1415 bool composeAffineMin) {
1420 return AffineApplyOp::create(
b, loc, map, valueOperands);
1426 bool composeAffineMin) {
1431 operands, composeAffineMin);
1438 bool composeAffineMin =
false) {
1444 for (
unsigned i : llvm::seq<unsigned>(0, map.
getNumResults())) {
1452 llvm::append_range(dims,
1454 llvm::append_range(symbols,
1461 operands = llvm::to_vector(llvm::concat<Value>(dims, symbols));
1468 bool composeAffineMin) {
1469 assert(map.
getNumResults() == 1 &&
"building affine.apply with !=1 result");
1479 AffineApplyOp applyOp =
1484 for (
unsigned i = 0, e = constOperands.size(); i != e; ++i)
1489 if (failed(applyOp->fold(constOperands, foldResults)) ||
1490 foldResults.empty()) {
1492 listener->notifyOperationInserted(applyOp, {});
1493 return applyOp.getResult();
1497 return llvm::getSingleElement(foldResults);
1507 operands, composeAffineMin);
1513 bool composeAffineMin) {
1514 return llvm::map_to_vector(
1515 llvm::seq<unsigned>(0, map.
getNumResults()), [&](
unsigned i) {
1516 return makeComposedFoldedAffineApply(b, loc, map.getSubMap({i}),
1517 operands, composeAffineMin);
1521template <
typename OpTy>
1527 return OpTy::create(
b, loc,
b.getIndexType(), map, valueOperands);
1536template <
typename OpTy>
1552 for (
unsigned i = 0, e = constOperands.size(); i != e; ++i)
1557 if (failed(minMaxOp->fold(constOperands, foldResults)) ||
1558 foldResults.empty()) {
1560 listener->notifyOperationInserted(minMaxOp, {});
1561 return minMaxOp.getResult();
1565 return llvm::getSingleElement(foldResults);
1584template <
class MapOrSet>
1587 if (!mapOrSet || operands->empty())
1590 assert(mapOrSet->getNumInputs() == operands->size() &&
1591 "map/set inputs must match number of operands");
1593 auto *context = mapOrSet->getContext();
1595 resultOperands.reserve(operands->size());
1597 remappedSymbols.reserve(operands->size());
1598 unsigned nextDim = 0;
1599 unsigned nextSym = 0;
1600 unsigned oldNumSyms = mapOrSet->getNumSymbols();
1602 for (
unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) {
1603 if (i < mapOrSet->getNumDims()) {
1607 remappedSymbols.push_back((*operands)[i]);
1610 resultOperands.push_back((*operands)[i]);
1613 resultOperands.push_back((*operands)[i]);
1617 resultOperands.append(remappedSymbols.begin(), remappedSymbols.end());
1618 *operands = resultOperands;
1619 *mapOrSet = mapOrSet->replaceDimsAndSymbols(
1620 dimRemapping, {}, nextDim, oldNumSyms + nextSym);
1622 assert(mapOrSet->getNumInputs() == operands->size() &&
1623 "map/set inputs must match number of operands");
1632template <
class MapOrSet>
1635 if (!mapOrSet || operands.empty())
1638 unsigned numOperands = operands.size();
1640 assert(mapOrSet.getNumInputs() == numOperands &&
1641 "map/set inputs must match number of operands");
1643 auto *context = mapOrSet.getContext();
1645 resultOperands.reserve(numOperands);
1647 remappedDims.reserve(numOperands);
1649 symOperands.reserve(mapOrSet.getNumSymbols());
1650 unsigned nextSym = 0;
1651 unsigned nextDim = 0;
1652 unsigned oldNumDims = mapOrSet.getNumDims();
1654 resultOperands.assign(operands.begin(), operands.begin() + oldNumDims);
1655 for (
unsigned i = oldNumDims, e = mapOrSet.getNumInputs(); i != e; ++i) {
1658 symRemapping[i - oldNumDims] =
1660 remappedDims.push_back(operands[i]);
1663 symOperands.push_back(operands[i]);
1667 append_range(resultOperands, remappedDims);
1668 append_range(resultOperands, symOperands);
1669 operands = resultOperands;
1670 mapOrSet = mapOrSet.replaceDimsAndSymbols(
1671 {}, symRemapping, oldNumDims + nextDim, nextSym);
1673 assert(mapOrSet.getNumInputs() == operands.size() &&
1674 "map/set inputs must match number of operands");
1678template <
class MapOrSet>
1681 static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value,
1682 "Argument must be either of AffineMap or IntegerSet type");
1684 if (!mapOrSet || operands->empty())
1687 assert(mapOrSet->getNumInputs() == operands->size() &&
1688 "map/set inputs must match number of operands");
1694 llvm::SmallBitVector usedDims(mapOrSet->getNumDims());
1695 llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols());
1697 if (
auto dimExpr = dyn_cast<AffineDimExpr>(expr))
1698 usedDims[dimExpr.getPosition()] =
true;
1699 else if (
auto symExpr = dyn_cast<AffineSymbolExpr>(expr))
1700 usedSyms[symExpr.getPosition()] =
true;
1703 auto *context = mapOrSet->getContext();
1706 resultOperands.reserve(operands->size());
1708 llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims;
1710 unsigned nextDim = 0;
1711 for (
unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) {
1714 auto it = seenDims.find((*operands)[i]);
1715 if (it == seenDims.end()) {
1717 resultOperands.push_back((*operands)[i]);
1718 seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i]));
1720 dimRemapping[i] = it->second;
1724 llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols;
1726 unsigned nextSym = 0;
1727 for (
unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) {
1733 IntegerAttr operandCst;
1734 if (
matchPattern((*operands)[i + mapOrSet->getNumDims()],
1741 auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]);
1742 if (it == seenSymbols.end()) {
1744 resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]);
1745 seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()],
1748 symRemapping[i] = it->second;
1751 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping,
1753 *operands = resultOperands;
1770template <
typename AffineOpTy>
1779 LogicalResult matchAndRewrite(AffineOpTy affineOp,
1782 llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp,
1783 AffineStoreOp, AffineApplyOp, AffineMinOp, AffineMaxOp,
1784 AffineVectorStoreOp, AffineVectorLoadOp>::value,
1785 "affine load/store/vectorstore/vectorload/apply/prefetch/min/max op "
1787 auto map = affineOp.getAffineMap();
1789 auto oldOperands = affineOp.getMapOperands();
1794 if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(),
1795 resultOperands.begin()))
1798 replaceAffineOp(rewriter, affineOp, map, resultOperands);
1806void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp(
1810 mapOperands,
load.getMaybeAlign());
1813void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp(
1817 prefetch, prefetch.getMemref(), map, mapOperands, prefetch.getIsWrite(),
1818 prefetch.getLocalityHint(), prefetch.getIsDataCache());
1821void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp(
1825 store, store.getValueToStore(), store.getMemRef(), map, mapOperands,
1826 store.getMaybeAlign());
1829void SimplifyAffineOp<AffineVectorLoadOp>::replaceAffineOp(
1833 vectorload, vectorload.getVectorType(), vectorload.getMemRef(), map,
1834 mapOperands, vectorload.getMaybeAlign());
1837void SimplifyAffineOp<AffineVectorStoreOp>::replaceAffineOp(
1841 vectorstore, vectorstore.getValueToStore(), vectorstore.getMemRef(), map,
1842 mapOperands, vectorstore.getMaybeAlign());
1846template <
typename AffineOpTy>
1847void SimplifyAffineOp<AffineOpTy>::replaceAffineOp(
1856 results.
add<SimplifyAffineOp<AffineApplyOp>>(context);
1871 result.addOperands(srcMemRef);
1872 result.addAttribute(getSrcMapAttrStrName(), AffineMapAttr::get(srcMap));
1873 result.addOperands(srcIndices);
1874 result.addOperands(destMemRef);
1875 result.addAttribute(getDstMapAttrStrName(), AffineMapAttr::get(dstMap));
1876 result.addOperands(destIndices);
1877 result.addOperands(tagMemRef);
1878 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
1879 result.addOperands(tagIndices);
1880 result.addOperands(numElements);
1882 result.addOperands({stride, elementsPerStride});
1887 p <<
" " << getSrcMemRef() <<
'[';
1889 p <<
"], " << getDstMemRef() <<
'[';
1891 p <<
"], " << getTagMemRef() <<
'[';
1895 p <<
", " << getStride();
1896 p <<
", " << getNumElementsPerStride();
1898 p <<
" : " << getSrcMemRefType() <<
", " << getDstMemRefType() <<
", "
1899 << getTagMemRefType();
1908ParseResult AffineDmaStartOp::parse(
OpAsmParser &parser,
1911 AffineMapAttr srcMapAttr;
1914 AffineMapAttr dstMapAttr;
1917 AffineMapAttr tagMapAttr;
1932 getSrcMapAttrStrName(),
1936 getDstMapAttrStrName(),
1940 getTagMapAttrStrName(),
1949 if (!strideInfo.empty() && strideInfo.size() != 2) {
1951 "expected two stride related operands");
1953 bool isStrided = strideInfo.size() == 2;
1958 if (types.size() != 3)
1976 if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() ||
1977 dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() ||
1978 tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1980 "memref operand count not equal to map.numInputs");
1984LogicalResult AffineDmaStartOp::verify() {
1985 if (!llvm::isa<MemRefType>(getOperand(getSrcMemRefOperandIndex()).
getType()))
1986 return emitOpError(
"expected DMA source to be of memref type");
1987 if (!llvm::isa<MemRefType>(getOperand(getDstMemRefOperandIndex()).
getType()))
1988 return emitOpError(
"expected DMA destination to be of memref type");
1989 if (!llvm::isa<MemRefType>(getOperand(getTagMemRefOperandIndex()).
getType()))
1990 return emitOpError(
"expected DMA tag to be of memref type");
1992 unsigned numInputsAllMaps = getSrcMap().getNumInputs() +
1993 getDstMap().getNumInputs() +
1994 getTagMap().getNumInputs();
1995 if (getNumOperands() != numInputsAllMaps + 3 + 1 &&
1996 getNumOperands() != numInputsAllMaps + 3 + 1 + 2) {
1997 return emitOpError(
"incorrect number of operands");
2001 for (
auto idx : getSrcIndices()) {
2002 if (!idx.getType().isIndex())
2003 return emitOpError(
"src index to dma_start must have 'index' type");
2006 "src index must be a valid dimension or symbol identifier");
2008 for (
auto idx : getDstIndices()) {
2009 if (!idx.getType().isIndex())
2010 return emitOpError(
"dst index to dma_start must have 'index' type");
2013 "dst index must be a valid dimension or symbol identifier");
2015 for (
auto idx : getTagIndices()) {
2016 if (!idx.getType().isIndex())
2017 return emitOpError(
"tag index to dma_start must have 'index' type");
2020 "tag index must be a valid dimension or symbol identifier");
2025LogicalResult AffineDmaStartOp::fold(FoldAdaptor adaptor,
2031void AffineDmaStartOp::getEffects(
2050 result.addOperands(tagMemRef);
2051 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
2052 result.addOperands(tagIndices);
2053 result.addOperands(numElements);
2057 p <<
" " << getTagMemRef() <<
'[';
2062 p <<
" : " << getTagMemRef().getType();
2070ParseResult AffineDmaWaitOp::parse(
OpAsmParser &parser,
2073 AffineMapAttr tagMapAttr;
2082 getTagMapAttrStrName(),
2091 if (!llvm::isa<MemRefType>(type))
2093 "expected tag to be of memref type");
2095 if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
2097 "tag memref operand count != to map.numInputs");
2101LogicalResult AffineDmaWaitOp::verify() {
2102 if (!llvm::isa<MemRefType>(getOperand(0).
getType()))
2103 return emitOpError(
"expected DMA tag to be of memref type");
2105 for (
auto idx : getTagIndices()) {
2106 if (!idx.getType().isIndex())
2107 return emitOpError(
"index to dma_wait must have 'index' type");
2110 "index must be a valid dimension or symbol identifier");
2115LogicalResult AffineDmaWaitOp::fold(FoldAdaptor adaptor,
2121void AffineDmaWaitOp::getEffects(
2137 ValueRange iterArgs, BodyBuilderFn bodyBuilder) {
2138 assert(((!lbMap && lbOperands.empty()) ||
2140 "lower bound operand count does not match the affine map");
2141 assert(((!ubMap && ubOperands.empty()) ||
2143 "upper bound operand count does not match the affine map");
2144 assert(step > 0 &&
"step has to be a positive integer constant");
2146 OpBuilder::InsertionGuard guard(builder);
2150 getOperandSegmentSizeAttr(),
2152 static_cast<int32_t>(ubOperands.size()),
2153 static_cast<int32_t>(iterArgs.size())}));
2155 for (Value val : iterArgs)
2156 result.addTypes(val.getType());
2163 result.addAttribute(getLowerBoundMapAttrName(
result.name),
2164 AffineMapAttr::get(lbMap));
2165 result.addOperands(lbOperands);
2168 result.addAttribute(getUpperBoundMapAttrName(
result.name),
2169 AffineMapAttr::get(ubMap));
2170 result.addOperands(ubOperands);
2172 result.addOperands(iterArgs);
2175 Region *bodyRegion =
result.addRegion();
2177 Value inductionVar =
2179 for (Value val : iterArgs)
2180 bodyBlock->
addArgument(val.getType(), val.getLoc());
2185 if (iterArgs.empty() && !bodyBuilder) {
2186 ensureTerminator(*bodyRegion, builder,
result.location);
2187 }
else if (bodyBuilder) {
2188 OpBuilder::InsertionGuard guard(builder);
2190 bodyBuilder(builder,
result.location, inductionVar,
2197 BodyBuilderFn bodyBuilder) {
2200 return build(builder,
result, {}, lbMap, {}, ubMap, step, iterArgs,
2204LogicalResult AffineForOp::verify() {
2205 auto *body = getBody();
2206 if (body->getNumArguments() == 0 || !getInductionVar().
getType().isIndex())
2207 return emitOpError(
"expected body to have an index argument for the "
2208 "induction variable");
2213LogicalResult AffineForOp::verifyRegions() {
2215 if (getStepAsInt() <= 0)
2216 return emitOpError(
"expected step to be a positive integer, got ")
2221 if (getLowerBoundMap().getNumInputs() > 0)
2223 getLowerBoundMap().getNumDims())))
2226 if (getUpperBoundMap().getNumInputs() > 0)
2228 getUpperBoundMap().getNumDims())))
2230 if (getLowerBoundMap().getNumResults() < 1)
2231 return emitOpError(
"expected lower bound map to have at least one result");
2232 if (getUpperBoundMap().getNumResults() < 1)
2233 return emitOpError(
"expected upper bound map to have at least one result");
2235 unsigned opNumResults = getNumResults();
2236 if (opNumResults == 0)
2242 if (getNumIterOperands() != opNumResults)
2244 "mismatch between the number of loop-carried values and results");
2245 if (getNumRegionIterArgs() != opNumResults)
2247 "mismatch between the number of basic block args and results");
2257 bool failedToParsedMinMax =
2261 auto boundAttrStrName =
2262 isLower ? AffineForOp::getLowerBoundMapAttrName(
result.name)
2263 : AffineForOp::getUpperBoundMapAttrName(
result.name);
2270 if (!boundOpInfos.empty()) {
2272 if (boundOpInfos.size() > 1)
2274 "expected only one loop bound operand");
2286 result.addAttribute(boundAttrStrName, AffineMapAttr::get(map));
2299 if (
auto affineMapAttr = dyn_cast<AffineMapAttr>(boundAttr)) {
2300 unsigned currentNumOperands =
result.operands.size();
2305 auto map = affineMapAttr.getValue();
2306 if (map.getNumDims() != numDims)
2309 "dim operand count and affine map dim count must match");
2311 unsigned numDimAndSymbolOperands =
2312 result.operands.size() - currentNumOperands;
2313 if (numDims + map.getNumSymbols() != numDimAndSymbolOperands)
2316 "symbol operand count and affine map symbol count must match");
2320 if (map.getNumResults() > 1 && failedToParsedMinMax) {
2322 return p.
emitError(attrLoc,
"lower loop bound affine map with "
2323 "multiple results requires 'max' prefix");
2325 return p.
emitError(attrLoc,
"upper loop bound affine map with multiple "
2326 "results requires 'min' prefix");
2332 if (
auto integerAttr = dyn_cast<IntegerAttr>(boundAttr)) {
2333 result.attributes.pop_back();
2342 "expected valid affine map representation for loop bounds");
2347 OpAsmParser::Argument inductionVariable;
2354 int64_t numOperands =
result.operands.size();
2357 int64_t numLbOperands =
result.operands.size() - numOperands;
2360 numOperands =
result.operands.size();
2363 int64_t numUbOperands =
result.operands.size() - numOperands;
2368 getStepAttrName(
result.name),
2372 IntegerAttr stepAttr;
2374 getStepAttrName(
result.name).data(),
2378 if (!stepAttr.getValue().isStrictlyPositive())
2381 "expected step to be representable as a positive signed integer");
2385 SmallVector<OpAsmParser::Argument, 4> regionArgs;
2386 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
2389 regionArgs.push_back(inductionVariable);
2397 for (
auto argOperandType :
2398 llvm::zip(llvm::drop_begin(regionArgs), operands,
result.types)) {
2399 Type type = std::get<2>(argOperandType);
2400 std::get<0>(argOperandType).type = type;
2408 getOperandSegmentSizeAttr(),
2410 static_cast<int32_t>(numUbOperands),
2411 static_cast<int32_t>(operands.size())}));
2414 Region *body =
result.addRegion();
2415 if (regionArgs.size() !=
result.types.size() + 1)
2418 "mismatch between the number of loop-carried values and results");
2422 AffineForOp::ensureTerminator(*body, builder,
result.location);
2444 if (
auto constExpr = dyn_cast<AffineConstantExpr>(expr)) {
2445 p << constExpr.getValue();
2453 if (isa<AffineSymbolExpr>(expr)) {
2469unsigned AffineForOp::getNumIterOperands() {
2470 AffineMap lbMap = getLowerBoundMapAttr().getValue();
2471 AffineMap ubMap = getUpperBoundMapAttr().getValue();
2476std::optional<MutableArrayRef<OpOperand>>
2477AffineForOp::getYieldedValuesMutable() {
2478 return cast<AffineYieldOp>(getBody()->getTerminator()).getOperandsMutable();
2490 if (getStepAsInt() != 1)
2491 p <<
" step " << getStepAsInt();
2493 bool printBlockTerminators =
false;
2494 if (getNumIterOperands() > 0) {
2496 auto regionArgs = getRegionIterArgs();
2497 auto operands = getInits();
2499 llvm::interleaveComma(llvm::zip(regionArgs, operands), p, [&](
auto it) {
2500 p << std::get<0>(it) <<
" = " << std::get<1>(it);
2502 p <<
") -> (" << getResultTypes() <<
")";
2503 printBlockTerminators =
true;
2508 printBlockTerminators);
2510 (*this)->getAttrs(),
2511 {getLowerBoundMapAttrName(getOperation()->getName()),
2512 getUpperBoundMapAttrName(getOperation()->getName()),
2513 getStepAttrName(getOperation()->getName()),
2514 getOperandSegmentSizeAttr()});
2519 auto foldLowerOrUpperBound = [&forOp](
bool lower) {
2523 auto boundOperands =
2524 lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands();
2525 for (
auto operand : boundOperands) {
2528 operandConstants.push_back(operandCst);
2532 lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap();
2534 "bound maps should have at least one result");
2536 if (failed(boundMap.
constantFold(operandConstants, foldedResults)))
2540 assert(!foldedResults.empty() &&
"bounds should have at least one result");
2541 auto maxOrMin = llvm::cast<IntegerAttr>(foldedResults[0]).getValue();
2542 for (
unsigned i = 1, e = foldedResults.size(); i < e; i++) {
2543 auto foldedResult = llvm::cast<IntegerAttr>(foldedResults[i]).getValue();
2544 maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult)
2545 : llvm::APIntOps::smin(maxOrMin, foldedResult);
2547 lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue())
2548 : forOp.setConstantUpperBound(maxOrMin.getSExtValue());
2553 bool folded =
false;
2554 if (!forOp.hasConstantLowerBound())
2555 folded |= succeeded(foldLowerOrUpperBound(
true));
2558 if (!forOp.hasConstantUpperBound())
2559 folded |= succeeded(foldLowerOrUpperBound(
false));
2565 int64_t step = forOp.getStepAsInt();
2566 if (!forOp.hasConstantBounds() || step <= 0)
2567 return std::nullopt;
2568 int64_t lb = forOp.getConstantLowerBound();
2569 int64_t ub = forOp.getConstantUpperBound();
2570 return ub - lb <= 0 ? 0 : (
ub - lb + step - 1) / step;
2575 if (!llvm::hasSingleElement(*forOp.getBody()))
2577 if (forOp.getNumResults() == 0)
2580 if (tripCount == 0) {
2583 return forOp.getInits();
2586 auto yieldOp = cast<AffineYieldOp>(forOp.getBody()->getTerminator());
2587 auto iterArgs = forOp.getRegionIterArgs();
2588 bool hasValDefinedOutsideLoop =
false;
2589 bool iterArgsNotInOrder =
false;
2590 for (
unsigned i = 0, e = yieldOp->getNumOperands(); i < e; ++i) {
2591 Value val = yieldOp.getOperand(i);
2595 if (val == forOp.getInductionVar())
2597 if (iterArgIt == iterArgs.end()) {
2599 assert(forOp.isDefinedOutsideOfLoop(val) &&
2600 "must be defined outside of the loop");
2601 hasValDefinedOutsideLoop =
true;
2602 replacements.push_back(val);
2604 unsigned pos = std::distance(iterArgs.begin(), iterArgIt);
2606 iterArgsNotInOrder =
true;
2607 replacements.push_back(forOp.getInits()[pos]);
2612 if (!tripCount.has_value() &&
2613 (hasValDefinedOutsideLoop || iterArgsNotInOrder))
2617 if (tripCount.has_value() && tripCount.value() >= 2 && iterArgsNotInOrder)
2619 return llvm::to_vector_of<OpFoldResult>(replacements);
2627 auto lbMap = forOp.getLowerBoundMap();
2628 auto ubMap = forOp.getUpperBoundMap();
2629 auto prevLbMap = lbMap;
2630 auto prevUbMap = ubMap;
2643 if (lbMap == prevLbMap && ubMap == prevUbMap)
2646 if (lbMap != prevLbMap)
2647 forOp.setLowerBound(lbOperands, lbMap);
2648 if (ubMap != prevUbMap)
2649 forOp.setUpperBound(ubOperands, ubMap);
2658LogicalResult AffineForOp::fold(FoldAdaptor adaptor,
2668 results.assign(getInits().begin(), getInits().end());
2672 if (!foldResults.empty()) {
2673 results.assign(foldResults);
2682 "invalid region point");
2689void AffineForOp::getSuccessorRegions(
2694 "expected loop region");
2700 if (tripCount.has_value()) {
2704 if (tripCount == 1) {
2705 regions.push_back(RegionSuccessor(getOperation()));
2711 if (tripCount.value() > 0) {
2712 regions.push_back(RegionSuccessor(&getRegion()));
2715 if (tripCount.value() == 0) {
2716 regions.push_back(RegionSuccessor(getOperation()));
2724 regions.push_back(RegionSuccessor(&getRegion()));
2725 regions.push_back(RegionSuccessor(getOperation()));
2730 return getResults();
2731 return getRegionIterArgs();
2744 assert(map.
getNumResults() >= 1 &&
"bound map has at least one result");
2745 getLowerBoundOperandsMutable().assign(lbOperands);
2746 setLowerBoundMap(map);
2751 assert(map.
getNumResults() >= 1 &&
"bound map has at least one result");
2752 getUpperBoundOperandsMutable().assign(ubOperands);
2753 setUpperBoundMap(map);
2756bool AffineForOp::hasConstantLowerBound() {
2757 return getLowerBoundMap().isSingleConstant();
2760bool AffineForOp::hasConstantUpperBound() {
2761 return getUpperBoundMap().isSingleConstant();
2764int64_t AffineForOp::getConstantLowerBound() {
2765 return getLowerBoundMap().getSingleConstantResult();
2768int64_t AffineForOp::getConstantUpperBound() {
2769 return getUpperBoundMap().getSingleConstantResult();
2772void AffineForOp::setConstantLowerBound(
int64_t value) {
2776void AffineForOp::setConstantUpperBound(
int64_t value) {
2780AffineForOp::operand_range AffineForOp::getControlOperands() {
2785bool AffineForOp::matchingBoundOperandList() {
2786 auto lbMap = getLowerBoundMap();
2787 auto ubMap = getUpperBoundMap();
2793 for (
unsigned i = 0, e = lbMap.
getNumInputs(); i < e; i++) {
2795 if (getOperand(i) != getOperand(numOperands + i))
2803std::optional<SmallVector<Value>> AffineForOp::getLoopInductionVars() {
2804 return SmallVector<Value>{getInductionVar()};
2807std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopLowerBounds() {
2808 if (!hasConstantLowerBound())
2809 return std::nullopt;
2811 return SmallVector<OpFoldResult>{
2812 OpFoldResult(
b.getI64IntegerAttr(getConstantLowerBound()))};
2815std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopSteps() {
2817 return SmallVector<OpFoldResult>{
2818 OpFoldResult(
b.getI64IntegerAttr(getStepAsInt()))};
2821std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopUpperBounds() {
2822 if (!hasConstantUpperBound())
2825 return SmallVector<OpFoldResult>{
2826 OpFoldResult(
b.getI64IntegerAttr(getConstantUpperBound()))};
2829std::optional<APInt> AffineForOp::getStaticTripCount() {
2831 int64_t step = getStepAsInt();
2833 return std::nullopt;
2835 if (hasConstantBounds()) {
2836 int64_t lb = getConstantLowerBound();
2837 int64_t ub = getConstantUpperBound();
2838 int64_t loopSpan = ub - lb;
2841 return APInt(64, llvm::divideCeilSigned(loopSpan, step));
2844 auto lbMap = getLowerBoundMap();
2845 auto ubMap = getUpperBoundMap();
2847 return std::nullopt;
2854 SmallVector<AffineExpr, 4> lbSplatExpr(ubValueMap.getNumResults(),
2857 lbSplatExpr, context);
2860 AffineValueMap tripCountValueMap;
2864 std::optional<uint64_t> tripCount;
2865 for (
unsigned i = 0, e = tripCountValueMap.
getNumResults(); i < e; ++i) {
2867 if (
auto constExpr = llvm::dyn_cast<AffineConstantExpr>(expr)) {
2868 uint64_t value = constExpr.getValue();
2869 if (tripCount.has_value())
2870 tripCount = std::min(*tripCount, value);
2874 return std::nullopt;
2878 if (tripCount.has_value())
2879 return APInt(64, *tripCount);
2881 return std::nullopt;
2884FailureOr<LoopLikeOpInterface> AffineForOp::replaceWithAdditionalYields(
2886 bool replaceInitOperandUsesInLoop,
2889 OpBuilder::InsertionGuard g(rewriter);
2891 auto inits = llvm::to_vector(getInits());
2892 inits.append(newInitOperands.begin(), newInitOperands.end());
2893 AffineForOp newLoop = AffineForOp::create(
2898 newLoop->setDiscardableAttrs(getOperation()->getDiscardableAttrDictionary());
2901 auto yieldOp = cast<AffineYieldOp>(getBody()->getTerminator());
2902 ArrayRef<BlockArgument> newIterArgs =
2903 newLoop.getBody()->getArguments().take_back(newInitOperands.size());
2905 OpBuilder::InsertionGuard g(rewriter);
2907 SmallVector<Value> newYieldedValues =
2908 newYieldValuesFn(rewriter, getLoc(), newIterArgs);
2909 assert(newInitOperands.size() == newYieldedValues.size() &&
2910 "expected as many new yield values as new iter operands");
2912 yieldOp.getOperandsMutable().append(newYieldedValues);
2917 rewriter.
mergeBlocks(getBody(), newLoop.getBody(),
2918 newLoop.getBody()->getArguments().take_front(
2919 getBody()->getNumArguments()));
2921 if (replaceInitOperandUsesInLoop) {
2924 for (
auto it : llvm::zip(newInitOperands, newIterArgs)) {
2926 [&](OpOperand &use) {
2928 return newLoop->isProperAncestor(user);
2935 newLoop->getResults().take_front(getNumResults()));
2936 return cast<LoopLikeOpInterface>(newLoop.getOperation());
2964 auto ivArg = dyn_cast<BlockArgument>(val);
2965 if (!ivArg || !ivArg.getOwner() || !ivArg.getOwner()->getParent())
2966 return AffineForOp();
2968 ivArg.getOwner()->getParent()->getParentOfType<AffineForOp>())
2970 return forOp.getInductionVar() == val ? forOp : AffineForOp();
2971 return AffineForOp();
2975 auto ivArg = dyn_cast<BlockArgument>(val);
2976 if (!ivArg || !ivArg.getOwner())
2979 auto parallelOp = dyn_cast_if_present<AffineParallelOp>(containingOp);
2980 if (parallelOp && llvm::is_contained(parallelOp.getIVs(), val))
2989 ivs->reserve(forInsts.size());
2990 for (
auto forInst : forInsts)
2991 ivs->push_back(forInst.getInductionVar());
2996 ivs.reserve(affineOps.size());
2999 if (
auto forOp = dyn_cast<AffineForOp>(op))
3000 ivs.push_back(forOp.getInductionVar());
3001 else if (
auto parallelOp = dyn_cast<AffineParallelOp>(op))
3002 for (
size_t i = 0; i < parallelOp.getBody()->getNumArguments(); i++)
3003 ivs.push_back(parallelOp.getBody()->getArgument(i));
3009template <
typename BoundListTy,
typename LoopCreatorTy>
3014 LoopCreatorTy &&loopCreatorFn) {
3015 assert(lbs.size() == ubs.size() &&
"Mismatch in number of arguments");
3016 assert(lbs.size() == steps.size() &&
"Mismatch in number of arguments");
3028 ivs.reserve(lbs.size());
3029 for (
unsigned i = 0, e = lbs.size(); i < e; ++i) {
3035 if (i == e - 1 && bodyBuilderFn) {
3037 bodyBuilderFn(nestedBuilder, nestedLoc, ivs);
3039 AffineYieldOp::create(nestedBuilder, nestedLoc);
3044 auto loop = loopCreatorFn(builder, loc, lbs[i], ubs[i], steps[i], loopBody);
3053 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3054 return AffineForOp::create(builder, loc, lb,
ub, step,
3062 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3065 if (lbConst && ubConst)
3067 ubConst.value(), step, bodyBuilderFn);
3098 LogicalResult matchAndRewrite(AffineIfOp ifOp,
3100 if (ifOp.getElseRegion().empty() ||
3101 !llvm::hasSingleElement(*ifOp.getElseBlock()) || ifOp.getNumResults())
3114 using OpRewritePattern<AffineIfOp>::OpRewritePattern;
3116 LogicalResult matchAndRewrite(AffineIfOp op,
3117 PatternRewriter &rewriter)
const override {
3119 auto isTriviallyFalse = [](IntegerSet iSet) {
3120 return iSet.isEmptyIntegerSet();
3123 auto isTriviallyTrue = [](IntegerSet iSet) {
3124 return (iSet.getNumEqualities() == 1 && iSet.getNumInequalities() == 0 &&
3125 iSet.getConstraint(0) == 0);
3128 IntegerSet affineIfConditions = op.getIntegerSet();
3130 if (isTriviallyFalse(affineIfConditions)) {
3134 if (op.getNumResults() == 0 && !op.hasElse()) {
3140 blockToMove = op.getElseBlock();
3141 }
else if (isTriviallyTrue(affineIfConditions)) {
3142 blockToMove = op.getThenBlock();
3146 Operation *blockToMoveTerminator = blockToMove->
getTerminator();
3160 rewriter.
eraseOp(blockToMoveTerminator);
3168void AffineIfOp::getSuccessorRegions(
3176 if (getElseRegion().empty()) {
3191 return getResults();
3192 if (successor == &getThenRegion())
3193 return getThenRegion().getArguments();
3194 if (successor == &getElseRegion())
3195 return getElseRegion().getArguments();
3196 llvm_unreachable(
"invalid region successor");
3199LogicalResult AffineIfOp::verify() {
3202 auto conditionAttr =
3203 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3205 return emitOpError(
"requires an integer set attribute named 'condition'");
3208 IntegerSet condition = conditionAttr.getValue();
3210 return emitOpError(
"operand count and condition integer set dimension and "
3211 "symbol count must match");
3223 IntegerSetAttr conditionAttr;
3226 AffineIfOp::getConditionAttrStrName(),
3232 auto set = conditionAttr.getValue();
3233 if (set.getNumDims() != numDims)
3236 "dim operand count and integer set dim count must match");
3237 if (numDims + set.getNumSymbols() !=
result.operands.size())
3240 "symbol operand count and integer set symbol count must match");
3247 result.regions.reserve(2);
3254 AffineIfOp::ensureTerminator(*thenRegion, parser.
getBuilder(),
3261 AffineIfOp::ensureTerminator(*elseRegion, parser.
getBuilder(),
3273 auto conditionAttr =
3274 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3275 p <<
" " << conditionAttr;
3277 conditionAttr.getValue().getNumDims(), p);
3284 auto &elseRegion = this->getElseRegion();
3285 if (!elseRegion.
empty()) {
3294 getConditionAttrStrName());
3299 ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName())
3303void AffineIfOp::setIntegerSet(
IntegerSet newSet) {
3304 (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet));
3309 (*this)->setOperands(operands);
3314 bool withElseRegion) {
3315 assert(resultTypes.empty() || withElseRegion);
3318 result.addTypes(resultTypes);
3319 result.addOperands(args);
3320 result.addAttribute(getConditionAttrStrName(), IntegerSetAttr::get(set));
3324 if (resultTypes.empty())
3325 AffineIfOp::ensureTerminator(*thenRegion, builder,
result.location);
3328 if (withElseRegion) {
3330 if (resultTypes.empty())
3331 AffineIfOp::ensureTerminator(*elseRegion, builder,
result.location);
3337 AffineIfOp::build(builder,
result, {}, set, args,
3346 bool composeAffineMin =
false) {
3353 if (llvm::none_of(operands,
3364 auto set = getIntegerSet();
3370 if (getIntegerSet() == set && llvm::equal(operands, getOperands()))
3373 setConditional(set, operands);
3379 results.
add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context);
3384 StringAttr attrName, llvm::MaybeAlign alignment) {
3386 result.addAttribute(attrName,
3396 llvm::MaybeAlign alignment) {
3397 assert(operands.size() == 1 + map.
getNumInputs() &&
"inconsistent operands");
3398 result.addOperands(operands);
3400 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3403 auto memrefType = llvm::cast<MemRefType>(operands[0].
getType());
3404 result.types.push_back(memrefType.getElementType());
3409 llvm::MaybeAlign alignment) {
3410 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
3412 result.addOperands(mapOperands);
3413 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3414 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3417 result.types.push_back(memrefType.getElementType());
3422 llvm::MaybeAlign alignment) {
3423 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3424 int64_t rank = memrefType.getRank();
3438 AffineMapAttr mapAttr;
3443 AffineLoadOp::getMapAttrStrName(),
3454 if (AffineMapAttr mapAttr =
3455 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3459 {getMapAttrStrName()});
3465template <
typename AffineMemOpTy>
3469 MemRefType memrefType,
unsigned numIndexOperands) {
3472 return op->emitOpError(
"affine map num results must equal memref rank");
3474 return op->emitOpError(
"expects as many subscripts as affine map inputs");
3476 for (
auto idx : mapOperands) {
3477 if (!idx.getType().isIndex())
3478 return op->emitOpError(
"index to load must have 'index' type");
3486LogicalResult AffineLoadOp::verify() {
3488 if (
getType() != memrefType.getElementType())
3489 return emitOpError(
"result type must match element type of memref");
3492 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3493 getMapOperands(), memrefType,
3494 getNumOperands() - 1)))
3502 results.
add<SimplifyAffineOp<AffineLoadOp>>(context);
3511 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
3516 getGlobalOp, getGlobalOp.getNameAttr());
3522 dyn_cast_or_null<DenseElementsAttr>(global.getConstantInitValue());
3526 if (
auto splatAttr = dyn_cast<SplatElementsAttr>(cstAttr))
3527 return splatAttr.getSplatValue<
Attribute>();
3529 if (!getAffineMap().isConstant())
3532 llvm::map_to_vector<4>(getAffineMap().getConstantResults(),
3533 [](
int64_t v) -> uint64_t {
return v; });
3543 ValueRange mapOperands, llvm::MaybeAlign alignment) {
3544 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
3545 result.addOperands(valueToStore);
3547 result.addOperands(mapOperands);
3548 result.getOrAddProperties<Properties>().map = AffineMapAttr::get(map);
3556 llvm::MaybeAlign alignment) {
3557 auto memrefType = llvm::cast<MemRefType>(
memref.getType());
3558 int64_t rank = memrefType.getRank();
3572 AffineMapAttr mapAttr;
3577 mapOperands, mapAttr, AffineStoreOp::getMapAttrStrName(),
3588 p <<
" " << getValueToStore();
3590 if (AffineMapAttr mapAttr =
3591 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3595 {getMapAttrStrName()});
3599LogicalResult AffineStoreOp::verify() {
3602 if (getValueToStore().
getType() != memrefType.getElementType())
3604 "value to store must have the same type as memref element type");
3607 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3608 getMapOperands(), memrefType,
3609 getNumOperands() - 2)))
3617 results.
add<SimplifyAffineOp<AffineStoreOp>>(context);
3620LogicalResult AffineStoreOp::fold(FoldAdaptor adaptor,
3630template <
typename T>
3633 if (op.getNumOperands() !=
3634 op.getMap().getNumDims() + op.getMap().getNumSymbols())
3635 return op.emitOpError(
3636 "operand count and affine map dimension and symbol count must match");
3638 if (op.getMap().getNumResults() == 0)
3639 return op.emitOpError(
"affine map expect at least one result");
3643template <
typename T>
3645 p <<
' ' << op->getAttr(T::getMapAttrStrName());
3646 auto operands = op.getOperands();
3647 unsigned numDims = op.getMap().getNumDims();
3648 p <<
'(' << operands.take_front(numDims) <<
')';
3650 if (operands.size() != numDims)
3651 p <<
'[' << operands.drop_front(numDims) <<
']';
3653 {T::getMapAttrStrName()});
3656template <
typename T>
3663 AffineMapAttr mapAttr;
3679template <
typename T>
3681 static_assert(llvm::is_one_of<T, AffineMinOp, AffineMaxOp>::value,
3682 "expected affine min or max op");
3688 auto foldedMap = op.getMap().partialConstantFold(operands, &results);
3690 if (foldedMap.getNumSymbols() == 1 && foldedMap.isSymbolIdentity())
3691 return op.getOperand(0);
3694 if (results.empty()) {
3696 if (foldedMap == op.getMap())
3698 op->setAttr(
"map", AffineMapAttr::get(foldedMap));
3699 return op.getResult();
3703 auto resultIt = std::is_same<T, AffineMinOp>::value
3704 ? llvm::min_element(results)
3705 : llvm::max_element(results);
3706 if (resultIt == results.end())
3708 return IntegerAttr::get(IndexType::get(op.getContext()), *resultIt);
3712template <
typename T>
3718 AffineMap oldMap = affineOp.getAffineMap();
3724 if (!llvm::is_contained(newExprs, expr))
3725 newExprs.push_back(expr);
3755template <
typename T>
3761 AffineMap oldMap = affineOp.getAffineMap();
3763 affineOp.getMapOperands().take_front(oldMap.
getNumDims());
3765 affineOp.getMapOperands().take_back(oldMap.
getNumSymbols());
3767 auto newDimOperands = llvm::to_vector<8>(dimOperands);
3768 auto newSymOperands = llvm::to_vector<8>(symOperands);
3776 if (
auto symExpr = dyn_cast<AffineSymbolExpr>(expr)) {
3777 Value symValue = symOperands[symExpr.getPosition()];
3779 producerOps.push_back(producerOp);
3782 }
else if (
auto dimExpr = dyn_cast<AffineDimExpr>(expr)) {
3783 Value dimValue = dimOperands[dimExpr.getPosition()];
3785 producerOps.push_back(producerOp);
3792 newExprs.push_back(expr);
3795 if (producerOps.empty())
3802 for (T producerOp : producerOps) {
3803 AffineMap producerMap = producerOp.getAffineMap();
3804 unsigned numProducerDims = producerMap.
getNumDims();
3809 producerOp.getMapOperands().take_front(numProducerDims);
3811 producerOp.getMapOperands().take_back(numProducerSyms);
3812 newDimOperands.append(dimValues.begin(), dimValues.end());
3813 newSymOperands.append(symValues.begin(), symValues.end());
3817 newExprs.push_back(expr.
shiftDims(numProducerDims, numUsedDims)
3821 numUsedDims += numProducerDims;
3822 numUsedSyms += numProducerSyms;
3828 llvm::to_vector<8>(llvm::concat<Value>(newDimOperands, newSymOperands));
3847 if (!resultExpr.isPureAffine())
3852 if (failed(flattenResult))
3865 if (llvm::is_sorted(flattenedExprs))
3870 llvm::to_vector(llvm::seq<unsigned>(0, map.
getNumResults()));
3871 llvm::sort(resultPermutation, [&](
unsigned lhs,
unsigned rhs) {
3872 return flattenedExprs[
lhs] < flattenedExprs[
rhs];
3875 for (
unsigned idx : resultPermutation)
3896template <
typename T>
3902 AffineMap map = affineOp.getAffineMap();
3910template <
typename T>
3916 if (affineOp.getMap().getNumResults() != 1)
3919 affineOp.getOperands());
3987ParseResult AffinePrefetchOp::parse(
OpAsmParser &parser,
3994 IntegerAttr hintInfo;
3996 StringRef readOrWrite, cacheType;
3998 AffineMapAttr mapAttr;
4002 AffinePrefetchOp::getMapAttrStrName(),
4008 AffinePrefetchOp::getLocalityHintAttrStrName(),
4018 if (readOrWrite !=
"read" && readOrWrite !=
"write")
4020 "rw specifier has to be 'read' or 'write'");
4021 result.addAttribute(AffinePrefetchOp::getIsWriteAttrStrName(),
4024 if (cacheType !=
"data" && cacheType !=
"instr")
4026 "cache type has to be 'data' or 'instr'");
4028 result.addAttribute(AffinePrefetchOp::getIsDataCacheAttrStrName(),
4035 p <<
" " << getMemref() <<
'[';
4036 AffineMapAttr mapAttr =
4037 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4040 p <<
']' <<
", " << (getIsWrite() ?
"write" :
"read") <<
", " <<
"locality<"
4041 << getLocalityHint() <<
">, " << (getIsDataCache() ?
"data" :
"instr");
4043 (*this)->getAttrs(),
4044 {getMapAttrStrName(), getLocalityHintAttrStrName(),
4045 getIsDataCacheAttrStrName(), getIsWriteAttrStrName()});
4049LogicalResult AffinePrefetchOp::verify() {
4050 auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4054 return emitOpError(
"affine.prefetch affine map num results must equal"
4059 if (getNumOperands() != 1)
4064 for (
auto idx : getMapOperands()) {
4067 "index must be a valid dimension or symbol identifier");
4075 results.
add<SimplifyAffineOp<AffinePrefetchOp>>(context);
4078LogicalResult AffinePrefetchOp::fold(FoldAdaptor adaptor,
4093 auto ubs = llvm::map_to_vector<4>(ranges, [&](
int64_t value) {
4097 build(builder,
result, resultTypes, reductions, lbs, {}, ubs,
4107 assert(llvm::all_of(lbMaps,
4109 return m.
getNumDims() == lbMaps[0].getNumDims() &&
4112 "expected all lower bounds maps to have the same number of dimensions "
4114 assert(llvm::all_of(ubMaps,
4116 return m.
getNumDims() == ubMaps[0].getNumDims() &&
4119 "expected all upper bounds maps to have the same number of dimensions "
4121 assert((lbMaps.empty() || lbMaps[0].getNumInputs() == lbArgs.size()) &&
4122 "expected lower bound maps to have as many inputs as lower bound "
4124 assert((ubMaps.empty() || ubMaps[0].getNumInputs() == ubArgs.size()) &&
4125 "expected upper bound maps to have as many inputs as upper bound "
4129 result.addTypes(resultTypes);
4133 for (arith::AtomicRMWKind reduction : reductions)
4134 reductionAttrs.push_back(
4136 result.addAttribute(getReductionsAttrStrName(),
4146 groups.reserve(groups.size() + maps.size());
4147 exprs.reserve(maps.size());
4152 return AffineMap::get(maps[0].getNumDims(), maps[0].getNumSymbols(), exprs,
4158 AffineMap lbMap = concatMapsSameInput(lbMaps, lbGroups);
4159 AffineMap ubMap = concatMapsSameInput(ubMaps, ubGroups);
4160 result.addAttribute(getLowerBoundsMapAttrStrName(),
4161 AffineMapAttr::get(lbMap));
4162 result.addAttribute(getLowerBoundsGroupsAttrStrName(),
4164 result.addAttribute(getUpperBoundsMapAttrStrName(),
4165 AffineMapAttr::get(ubMap));
4166 result.addAttribute(getUpperBoundsGroupsAttrStrName(),
4169 result.addOperands(lbArgs);
4170 result.addOperands(ubArgs);
4173 auto *bodyRegion =
result.addRegion();
4177 for (
unsigned i = 0, e = steps.size(); i < e; ++i)
4179 if (resultTypes.empty())
4180 ensureTerminator(*bodyRegion, builder,
result.location);
4184 return {&getRegion()};
4187unsigned AffineParallelOp::getNumDims() {
return getSteps().size(); }
4189AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() {
4190 return getOperands().take_front(getLowerBoundsMap().getNumInputs());
4193AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() {
4194 return getOperands().drop_front(getLowerBoundsMap().getNumInputs());
4197AffineMap AffineParallelOp::getLowerBoundMap(
unsigned pos) {
4198 auto values = getLowerBoundsGroups().getValues<int32_t>();
4200 for (
unsigned i = 0; i < pos; ++i)
4202 return getLowerBoundsMap().getSliceMap(start, values[pos]);
4205AffineMap AffineParallelOp::getUpperBoundMap(
unsigned pos) {
4206 auto values = getUpperBoundsGroups().getValues<int32_t>();
4208 for (
unsigned i = 0; i < pos; ++i)
4210 return getUpperBoundsMap().getSliceMap(start, values[pos]);
4214 return AffineValueMap(getLowerBoundsMap(), getLowerBoundsOperands());
4218 return AffineValueMap(getUpperBoundsMap(), getUpperBoundsOperands());
4221std::optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() {
4222 if (hasMinMaxBounds())
4223 return std::nullopt;
4231 for (
unsigned i = 0, e = rangesValueMap.
getNumResults(); i < e; ++i) {
4232 auto expr = rangesValueMap.
getResult(i);
4233 auto cst = dyn_cast<AffineConstantExpr>(expr);
4235 return std::nullopt;
4236 out.push_back(cst.getValue());
4241Block *AffineParallelOp::getBody() {
return &getRegion().
front(); }
4243OpBuilder AffineParallelOp::getBodyBuilder() {
4244 return OpBuilder(getBody(), std::prev(getBody()->end()));
4249 "operands to map must match number of inputs");
4251 auto ubOperands = getUpperBoundsOperands();
4254 newOperands.append(ubOperands.begin(), ubOperands.end());
4255 (*this)->setOperands(newOperands);
4257 setLowerBoundsMapAttr(AffineMapAttr::get(map));
4262 "operands to map must match number of inputs");
4265 newOperands.append(ubOperands.begin(), ubOperands.end());
4266 (*this)->setOperands(newOperands);
4268 setUpperBoundsMapAttr(AffineMapAttr::get(map));
4277 arith::AtomicRMWKind op) {
4279 case arith::AtomicRMWKind::addf:
4280 return isa<FloatType>(resultType);
4281 case arith::AtomicRMWKind::addi:
4282 return isa<IntegerType>(resultType);
4283 case arith::AtomicRMWKind::assign:
4285 case arith::AtomicRMWKind::mulf:
4286 return isa<FloatType>(resultType);
4287 case arith::AtomicRMWKind::muli:
4288 return isa<IntegerType>(resultType);
4289 case arith::AtomicRMWKind::maximumf:
4290 case arith::AtomicRMWKind::maxnumf:
4291 case arith::AtomicRMWKind::minimumf:
4292 case arith::AtomicRMWKind::minnumf:
4293 return isa<FloatType>(resultType);
4294 case arith::AtomicRMWKind::maxs: {
4295 auto intType = dyn_cast<IntegerType>(resultType);
4296 return intType && intType.isSigned();
4298 case arith::AtomicRMWKind::mins: {
4299 auto intType = dyn_cast<IntegerType>(resultType);
4300 return intType && intType.isSigned();
4302 case arith::AtomicRMWKind::maxu: {
4303 auto intType = dyn_cast<IntegerType>(resultType);
4304 return intType && intType.isUnsigned();
4306 case arith::AtomicRMWKind::minu: {
4307 auto intType = dyn_cast<IntegerType>(resultType);
4308 return intType && intType.isUnsigned();
4310 case arith::AtomicRMWKind::ori:
4311 case arith::AtomicRMWKind::andi:
4312 case arith::AtomicRMWKind::xori:
4313 return isa<IntegerType>(resultType);
4315 llvm_unreachable(
"Unhandled atomic rmw kind");
4318LogicalResult AffineParallelOp::verify() {
4319 auto numDims = getNumDims();
4322 getSteps().size() != numDims || getBody()->getNumArguments() != numDims) {
4323 return emitOpError() <<
"the number of region arguments ("
4324 << getBody()->getNumArguments()
4325 <<
") and the number of map groups for lower ("
4326 << getLowerBoundsGroups().getNumElements()
4327 <<
") and upper bound ("
4328 << getUpperBoundsGroups().getNumElements()
4329 <<
"), and the number of steps (" << getSteps().size()
4330 <<
") must all match";
4333 unsigned expectedNumLBResults = 0;
4334 for (APInt v : getLowerBoundsGroups()) {
4335 unsigned results = v.getZExtValue();
4338 <<
"expected lower bound map to have at least one result";
4339 expectedNumLBResults += results;
4341 if (expectedNumLBResults != getLowerBoundsMap().getNumResults())
4342 return emitOpError() <<
"expected lower bounds map to have "
4343 << expectedNumLBResults <<
" results";
4344 unsigned expectedNumUBResults = 0;
4345 for (APInt v : getUpperBoundsGroups()) {
4346 unsigned results = v.getZExtValue();
4349 <<
"expected upper bound map to have at least one result";
4350 expectedNumUBResults += results;
4352 if (expectedNumUBResults != getUpperBoundsMap().getNumResults())
4353 return emitOpError() <<
"expected upper bounds map to have "
4354 << expectedNumUBResults <<
" results";
4356 if (getReductions().size() != getNumResults())
4357 return emitOpError(
"a reduction must be specified for each output");
4361 for (
auto it : llvm::enumerate((getReductions()))) {
4363 auto intAttr = dyn_cast<IntegerAttr>(attr);
4364 if (!intAttr || !arith::symbolizeAtomicRMWKind(intAttr.getInt()))
4365 return emitOpError(
"invalid reduction attribute");
4366 auto kind = arith::symbolizeAtomicRMWKind(intAttr.getInt()).value();
4368 return emitOpError(
"result type cannot match reduction attribute");
4374 getLowerBoundsMap().getNumDims())))
4378 getUpperBoundsMap().getNumDims())))
4387 if (newMap ==
getAffineMap() && newOperands == operands)
4389 reset(newMap, newOperands);
4399 bool ubCanonicalized = succeeded(
ub.canonicalize());
4402 if (!lbCanonicalized && !ubCanonicalized)
4405 if (lbCanonicalized)
4407 if (ubCanonicalized)
4408 op.setUpperBounds(
ub.getOperands(),
ub.getAffineMap());
4413LogicalResult AffineParallelOp::fold(FoldAdaptor adaptor,
4414 SmallVectorImpl<OpFoldResult> &results) {
4425 StringRef keyword) {
4428 ValueRange dimOperands = operands.take_front(numDims);
4429 ValueRange symOperands = operands.drop_front(numDims);
4431 for (llvm::APInt groupSize : group) {
4435 unsigned size = groupSize.getZExtValue();
4440 p << keyword <<
'(';
4449void AffineParallelOp::print(OpAsmPrinter &p) {
4450 p <<
" (" << getBody()->getArguments() <<
") = (";
4452 getLowerBoundsOperands(),
"max");
4455 getUpperBoundsOperands(),
"min");
4457 SmallVector<int64_t, 8> steps = getSteps();
4458 bool elideSteps = llvm::all_of(steps, [](int64_t step) {
return step == 1; });
4461 llvm::interleaveComma(steps, p);
4464 if (getNumResults()) {
4466 llvm::interleaveComma(getReductions(), p, [&](
auto &attr) {
4467 arith::AtomicRMWKind sym = *arith::symbolizeAtomicRMWKind(
4468 llvm::cast<IntegerAttr>(attr).getInt());
4469 p <<
"\"" << arith::stringifyAtomicRMWKind(sym) <<
"\"";
4471 p <<
") -> (" << getResultTypes() <<
")";
4478 (*this)->getAttrs(),
4479 {AffineParallelOp::getReductionsAttrStrName(),
4480 AffineParallelOp::getLowerBoundsMapAttrStrName(),
4481 AffineParallelOp::getLowerBoundsGroupsAttrStrName(),
4482 AffineParallelOp::getUpperBoundsMapAttrStrName(),
4483 AffineParallelOp::getUpperBoundsGroupsAttrStrName(),
4484 AffineParallelOp::getStepsAttrStrName()});
4491static ParseResult deduplicateAndResolveOperands(
4492 OpAsmParser &parser,
4493 ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> operands,
4494 SmallVectorImpl<Value> &uniqueOperands,
4495 SmallVectorImpl<AffineExpr> &replacements,
AffineExprKind kind) {
4497 "expected operands to be dim or symbol expression");
4500 for (
const auto &list : operands) {
4501 SmallVector<Value> valueOperands;
4504 for (Value operand : valueOperands) {
4505 unsigned pos = std::distance(uniqueOperands.begin(),
4506 llvm::find(uniqueOperands, operand));
4507 if (pos == uniqueOperands.size())
4508 uniqueOperands.push_back(operand);
4509 replacements.push_back(
4519enum class MinMaxKind { Min, Max };
4538static ParseResult parseAffineMapWithMinMax(OpAsmParser &parser,
4543 const llvm::StringLiteral tmpAttrStrName =
"__pseudo_bound_map";
4545 StringRef mapName = kind == MinMaxKind::Min
4546 ? AffineParallelOp::getUpperBoundsMapAttrStrName()
4547 : AffineParallelOp::getLowerBoundsMapAttrStrName();
4548 StringRef groupsName =
4549 kind == MinMaxKind::Min
4550 ? AffineParallelOp::getUpperBoundsGroupsAttrStrName()
4551 : AffineParallelOp::getLowerBoundsGroupsAttrStrName();
4557 result.addAttribute(
4558 mapName, AffineMapAttr::get(parser.getBuilder().getEmptyAffineMap()));
4559 result.addAttribute(groupsName, parser.getBuilder().getI32TensorAttr({}));
4563 SmallVector<AffineExpr> flatExprs;
4564 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatDimOperands;
4565 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatSymOperands;
4566 SmallVector<int32_t> numMapsPerGroup;
4567 SmallVector<OpAsmParser::UnresolvedOperand> mapOperands;
4568 auto parseOperands = [&]() {
4570 kind == MinMaxKind::Min ?
"min" :
"max"))) {
4571 mapOperands.clear();
4577 result.attributes.erase(tmpAttrStrName);
4578 llvm::append_range(flatExprs, map.getValue().getResults());
4579 auto operandsRef = llvm::ArrayRef(mapOperands);
4580 auto dimsRef = operandsRef.take_front(map.getValue().getNumDims());
4581 SmallVector<OpAsmParser::UnresolvedOperand> dims(dimsRef);
4582 auto symsRef = operandsRef.drop_front(map.getValue().getNumDims());
4583 SmallVector<OpAsmParser::UnresolvedOperand> syms(symsRef);
4584 flatDimOperands.append(map.getValue().getNumResults(), dims);
4585 flatSymOperands.append(map.getValue().getNumResults(), syms);
4586 numMapsPerGroup.push_back(map.getValue().getNumResults());
4589 flatSymOperands.emplace_back(),
4590 flatExprs.emplace_back())))
4592 numMapsPerGroup.push_back(1);
4599 unsigned totalNumDims = 0;
4600 unsigned totalNumSyms = 0;
4601 for (
unsigned i = 0, e = flatExprs.size(); i < e; ++i) {
4602 unsigned numDims = flatDimOperands[i].size();
4603 unsigned numSyms = flatSymOperands[i].size();
4604 flatExprs[i] = flatExprs[i]
4605 .shiftDims(numDims, totalNumDims)
4606 .shiftSymbols(numSyms, totalNumSyms);
4607 totalNumDims += numDims;
4608 totalNumSyms += numSyms;
4612 SmallVector<Value> dimOperands, symOperands;
4613 SmallVector<AffineExpr> dimRplacements, symRepacements;
4614 if (deduplicateAndResolveOperands(parser, flatDimOperands, dimOperands,
4616 deduplicateAndResolveOperands(parser, flatSymOperands, symOperands,
4620 result.operands.append(dimOperands.begin(), dimOperands.end());
4621 result.operands.append(symOperands.begin(), symOperands.end());
4624 auto flatMap =
AffineMap::get(totalNumDims, totalNumSyms, flatExprs,
4626 flatMap = flatMap.replaceDimsAndSymbols(
4627 dimRplacements, symRepacements, dimOperands.size(), symOperands.size());
4629 result.addAttribute(mapName, AffineMapAttr::get(flatMap));
4639ParseResult AffineParallelOp::parse(OpAsmParser &parser,
4640 OperationState &
result) {
4643 SmallVector<OpAsmParser::Argument, 4> ivs;
4646 parseAffineMapWithMinMax(parser,
result, MinMaxKind::Max) ||
4648 parseAffineMapWithMinMax(parser,
result, MinMaxKind::Min))
4651 AffineMapAttr stepsMapAttr;
4652 NamedAttrList stepsAttrs;
4653 SmallVector<OpAsmParser::UnresolvedOperand, 4> stepsMapOperands;
4655 SmallVector<int64_t, 4> steps(ivs.size(), 1);
4656 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4660 AffineParallelOp::getStepsAttrStrName(),
4666 SmallVector<int64_t, 4> steps;
4667 auto stepsMap = stepsMapAttr.getValue();
4668 for (
const auto &
result : stepsMap.getResults()) {
4669 auto constExpr = dyn_cast<AffineConstantExpr>(
result);
4672 "steps must be constant integers");
4673 steps.push_back(constExpr.getValue());
4675 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4681 SmallVector<Attribute, 4> reductions;
4685 auto parseAttributes = [&]() -> ParseResult {
4690 NamedAttrList attrStorage;
4695 std::optional<arith::AtomicRMWKind> reduction =
4696 arith::symbolizeAtomicRMWKind(attrVal.getValue());
4698 return parser.
emitError(loc,
"invalid reduction value: ") << attrVal;
4699 reductions.push_back(
4707 result.addAttribute(AffineParallelOp::getReductionsAttrStrName(),
4715 Region *body =
result.addRegion();
4716 for (
auto &iv : ivs)
4717 iv.type = indexType;
4723 AffineParallelOp::ensureTerminator(*body, builder,
result.location);
4731LogicalResult AffineYieldOp::verify() {
4732 auto *parentOp = (*this)->getParentOp();
4733 auto results = parentOp->getResults();
4734 auto operands = getOperands();
4736 if (!isa<AffineParallelOp, AffineIfOp, AffineForOp>(parentOp))
4737 return emitOpError() <<
"only terminates affine.if/for/parallel regions";
4738 if (parentOp->getNumResults() != getNumOperands())
4739 return emitOpError() <<
"parent of yield must have same number of "
4740 "results as the yield operands";
4741 for (
auto it : llvm::zip(results, operands)) {
4743 return emitOpError() <<
"types mismatch between yield op and its parent";
4753void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4754 VectorType resultType, AffineMap map,
4756 llvm::MaybeAlign alignment) {
4757 assert(operands.size() == 1 + map.
getNumInputs() &&
"inconsistent operands");
4758 result.addOperands(operands);
4760 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4763 result.types.push_back(resultType);
4766void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4767 VectorType resultType, Value memref,
4769 llvm::MaybeAlign alignment) {
4770 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
4771 result.addOperands(memref);
4772 result.addOperands(mapOperands);
4773 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4776 result.types.push_back(resultType);
4779void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &
result,
4780 VectorType resultType, Value memref,
4782 auto memrefType = llvm::cast<MemRefType>(memref.
getType());
4783 int64_t rank = memrefType.getRank();
4788 build(builder,
result, resultType, memref, map,
indices, alignment);
4791void AffineVectorLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
4792 MLIRContext *context) {
4793 results.
add<SimplifyAffineOp<AffineVectorLoadOp>>(context);
4796ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser,
4797 OperationState &
result) {
4801 MemRefType memrefType;
4802 VectorType resultType;
4803 OpAsmParser::UnresolvedOperand memrefInfo;
4804 AffineMapAttr mapAttr;
4805 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4809 AffineVectorLoadOp::getMapAttrStrName(),
4819void AffineVectorLoadOp::print(OpAsmPrinter &p) {
4821 if (AffineMapAttr mapAttr =
4822 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4826 {getMapAttrStrName()});
4831static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType,
4832 VectorType vectorType) {
4834 if (memrefType.getElementType() != vectorType.getElementType())
4836 "requires memref and vector types of the same elemental type");
4840LogicalResult AffineVectorLoadOp::verify() {
4843 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4844 getMapOperands(), memrefType,
4845 getNumOperands() - 1)))
4858void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &
result,
4859 Value valueToStore, Value memref, AffineMap map,
4861 llvm::MaybeAlign alignment) {
4862 assert(map.
getNumInputs() == mapOperands.size() &&
"inconsistent index info");
4863 result.addOperands(valueToStore);
4864 result.addOperands(memref);
4865 result.addOperands(mapOperands);
4866 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4872void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &
result,
4873 Value valueToStore, Value memref,
4875 llvm::MaybeAlign alignment) {
4876 auto memrefType = llvm::cast<MemRefType>(memref.
getType());
4877 int64_t rank = memrefType.getRank();
4882 build(builder,
result, valueToStore, memref, map,
indices, alignment);
4884void AffineVectorStoreOp::getCanonicalizationPatterns(
4885 RewritePatternSet &results, MLIRContext *context) {
4886 results.
add<SimplifyAffineOp<AffineVectorStoreOp>>(context);
4889ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser,
4890 OperationState &
result) {
4893 MemRefType memrefType;
4894 VectorType resultType;
4895 OpAsmParser::UnresolvedOperand storeValueInfo;
4896 OpAsmParser::UnresolvedOperand memrefInfo;
4897 AffineMapAttr mapAttr;
4898 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4903 AffineVectorStoreOp::getMapAttrStrName(),
4913void AffineVectorStoreOp::print(OpAsmPrinter &p) {
4914 p <<
" " << getValueToStore();
4916 if (AffineMapAttr mapAttr =
4917 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4921 {getMapAttrStrName()});
4922 p <<
" : " <<
getMemRefType() <<
", " << getValueToStore().getType();
4925LogicalResult AffineVectorStoreOp::verify() {
4928 *
this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4929 getMapOperands(), memrefType,
4930 getNumOperands() - 2)))
4943void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4944 OperationState &odsState,
4946 ArrayRef<int64_t> staticBasis,
4947 bool hasOuterBound) {
4948 SmallVector<Type> returnTypes(hasOuterBound ? staticBasis.size()
4949 : staticBasis.size() + 1,
4951 build(odsBuilder, odsState, returnTypes, linearIndex, dynamicBasis,
4955void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4956 OperationState &odsState,
4958 bool hasOuterBound) {
4959 if (hasOuterBound && !basis.empty() && basis.front() ==
nullptr) {
4960 hasOuterBound =
false;
4961 basis = basis.drop_front();
4963 SmallVector<Value> dynamicBasis;
4964 SmallVector<int64_t> staticBasis;
4967 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4971void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4972 OperationState &odsState,
4974 ArrayRef<OpFoldResult> basis,
4975 bool hasOuterBound) {
4976 if (hasOuterBound && !basis.empty() && basis.front() == OpFoldResult()) {
4977 hasOuterBound =
false;
4978 basis = basis.drop_front();
4980 SmallVector<Value> dynamicBasis;
4981 SmallVector<int64_t> staticBasis;
4983 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4987void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4988 OperationState &odsState,
4989 Value linearIndex, ArrayRef<int64_t> basis,
4990 bool hasOuterBound) {
4991 build(odsBuilder, odsState, linearIndex,
ValueRange{}, basis, hasOuterBound);
4994LogicalResult AffineDelinearizeIndexOp::verify() {
4995 ArrayRef<int64_t> staticBasis = getStaticBasis();
4996 if (getNumResults() != staticBasis.size() &&
4997 getNumResults() != staticBasis.size() + 1)
4998 return emitOpError(
"should return an index for each basis element and up "
4999 "to one extra index");
5001 auto dynamicMarkersCount = llvm::count_if(staticBasis, ShapedType::isDynamic);
5002 if (
static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5004 "mismatch between dynamic and static basis (kDynamic marker but no "
5005 "corresponding dynamic basis entry) -- this can only happen due to an "
5006 "incorrect fold/rewrite");
5008 if (!llvm::all_of(staticBasis, [](int64_t v) {
5009 return v > 0 || ShapedType::isDynamic(v);
5011 return emitOpError(
"no basis element may be statically non-positive");
5020static std::optional<SmallVector<int64_t>>
5024 uint64_t dynamicBasisIndex = 0;
5030 if (basis && isa<IntegerAttr>(basis)) {
5031 mutableDynamicBasis.
erase(dynamicBasisIndex);
5033 ++dynamicBasisIndex;
5038 if (dynamicBasisIndex == dynamicBasis.size())
5039 return std::nullopt;
5045 staticBasis.push_back(ShapedType::kDynamic);
5047 staticBasis.push_back(*basisVal);
5054AffineDelinearizeIndexOp::fold(FoldAdaptor adaptor,
5055 SmallVectorImpl<OpFoldResult> &
result) {
5056 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5058 adaptor.getDynamicBasis());
5059 if (maybeStaticBasis) {
5060 setStaticBasis(*maybeStaticBasis);
5065 if (getNumResults() == 1) {
5066 result.push_back(getLinearIndex());
5070 if (adaptor.getLinearIndex() ==
nullptr)
5073 if (!adaptor.getDynamicBasis().empty())
5076 int64_t highPart = cast<IntegerAttr>(adaptor.getLinearIndex()).getInt();
5077 Type attrType = getLinearIndex().getType();
5079 ArrayRef<int64_t> staticBasis = getStaticBasis();
5080 if (hasOuterBound())
5081 staticBasis = staticBasis.drop_front();
5082 for (int64_t modulus : llvm::reverse(staticBasis)) {
5083 result.push_back(IntegerAttr::get(attrType, llvm::mod(highPart, modulus)));
5084 highPart = llvm::divideFloorSigned(highPart, modulus);
5086 result.push_back(IntegerAttr::get(attrType, highPart));
5091SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getEffectiveBasis() {
5093 if (hasOuterBound()) {
5094 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5096 getDynamicBasis().drop_front(), builder);
5098 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5102 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5105SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getPaddedBasis() {
5106 SmallVector<OpFoldResult> ret = getMixedBasis();
5107 if (!hasOuterBound())
5108 ret.insert(ret.begin(), OpFoldResult());
5115struct DropUnitExtentBasis
5116 :
public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5119 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5120 PatternRewriter &rewriter)
const override {
5121 SmallVector<Value> replacements(delinearizeOp->getNumResults(),
nullptr);
5122 std::optional<Value> zero = std::nullopt;
5123 Location loc = delinearizeOp->getLoc();
5124 Type indexType = delinearizeOp.getLinearIndex().getType();
5125 auto getZero = [&]() -> Value {
5127 zero = arith::ConstantOp::create(rewriter, loc,
5129 return zero.value();
5134 SmallVector<OpFoldResult> newBasis;
5135 for (
auto [index, basis] :
5136 llvm::enumerate(delinearizeOp.getPaddedBasis())) {
5137 std::optional<int64_t> basisVal =
5140 replacements[index] =
getZero();
5142 newBasis.push_back(basis);
5145 if (newBasis.size() == delinearizeOp.getNumResults())
5147 "no unit basis elements");
5149 if (!newBasis.empty()) {
5151 auto newDelinearizeOp = affine::AffineDelinearizeIndexOp::create(
5152 rewriter, loc, delinearizeOp.getLinearIndex(), newBasis);
5158 replacement = newDelinearizeOp->getResult(newIndex++);
5162 rewriter.
replaceOp(delinearizeOp, replacements);
5177struct CancelDelinearizeOfLinearizeDisjointExactTail
5178 :
public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5181 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5182 PatternRewriter &rewriter)
const override {
5183 auto linearizeOp = delinearizeOp.getLinearIndex()
5184 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5187 "index doesn't come from linearize");
5189 if (!linearizeOp.getDisjoint())
5192 ValueRange linearizeIns = linearizeOp.getMultiIndex();
5194 SmallVector<OpFoldResult> linearizeBasis = linearizeOp.getMixedBasis();
5195 SmallVector<OpFoldResult> delinearizeBasis = delinearizeOp.getMixedBasis();
5196 size_t numMatches = 0;
5197 for (
auto [linSize, delinSize] : llvm::zip(
5198 llvm::reverse(linearizeBasis), llvm::reverse(delinearizeBasis))) {
5199 if (linSize != delinSize)
5204 if (numMatches == 0)
5206 delinearizeOp,
"final basis element doesn't match linearize");
5209 if (numMatches == linearizeBasis.size() &&
5210 numMatches == delinearizeBasis.size() &&
5211 linearizeIns.size() == delinearizeOp.getNumResults()) {
5212 rewriter.
replaceOp(delinearizeOp, linearizeOp.getMultiIndex());
5216 Value newLinearize = affine::AffineLinearizeIndexOp::create(
5217 rewriter, linearizeOp.getLoc(), linearizeIns.drop_back(numMatches),
5218 ArrayRef<OpFoldResult>{linearizeBasis}.drop_back(numMatches),
5219 linearizeOp.getDisjoint());
5220 auto newDelinearize = affine::AffineDelinearizeIndexOp::create(
5221 rewriter, delinearizeOp.getLoc(), newLinearize,
5222 ArrayRef<OpFoldResult>{delinearizeBasis}.drop_back(numMatches),
5223 delinearizeOp.hasOuterBound());
5224 SmallVector<Value> mergedResults(newDelinearize.getResults());
5225 mergedResults.append(linearizeIns.take_back(numMatches).begin(),
5226 linearizeIns.take_back(numMatches).end());
5227 rewriter.
replaceOp(delinearizeOp, mergedResults);
5245struct SplitDelinearizeSpanningLastLinearizeArg final
5246 : OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5249 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5250 PatternRewriter &rewriter)
const override {
5251 auto linearizeOp = delinearizeOp.getLinearIndex()
5252 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5255 "index doesn't come from linearize");
5257 if (!linearizeOp.getDisjoint())
5259 "linearize isn't disjoint");
5264 if (linearizeOp.getStaticBasis().empty())
5266 linearizeOp,
"linearize has no basis elements (no inputs)");
5268 int64_t
target = linearizeOp.getStaticBasis().back();
5269 if (ShapedType::isDynamic(
target))
5271 linearizeOp,
"linearize ends with dynamic basis value");
5273 int64_t sizeToSplit = 1;
5274 size_t elemsToSplit = 0;
5275 ArrayRef<int64_t> basis = delinearizeOp.getStaticBasis();
5276 for (int64_t basisElem : llvm::reverse(basis)) {
5277 if (ShapedType::isDynamic(basisElem))
5279 delinearizeOp,
"dynamic basis element while scanning for split");
5280 sizeToSplit *= basisElem;
5283 if (sizeToSplit >
target)
5285 "overshot last argument size");
5286 if (sizeToSplit ==
target)
5290 if (sizeToSplit <
target)
5292 delinearizeOp,
"product of known basis elements doesn't exceed last "
5293 "linearize argument");
5295 if (elemsToSplit < 2)
5298 "need at least two elements to form the basis product");
5300 Value linearizeWithoutBack = affine::AffineLinearizeIndexOp::create(
5301 rewriter, linearizeOp.getLoc(), linearizeOp.getLinearIndex().getType(),
5302 linearizeOp.getMultiIndex().drop_back(), linearizeOp.getDynamicBasis(),
5303 linearizeOp.getStaticBasis().drop_back(), linearizeOp.getDisjoint());
5304 auto delinearizeWithoutSplitPart = affine::AffineDelinearizeIndexOp::create(
5305 rewriter, delinearizeOp.getLoc(), linearizeWithoutBack,
5306 delinearizeOp.getDynamicBasis(), basis.drop_back(elemsToSplit),
5307 delinearizeOp.hasOuterBound());
5308 auto delinearizeBack = affine::AffineDelinearizeIndexOp::create(
5309 rewriter, delinearizeOp.getLoc(), linearizeOp.getMultiIndex().back(),
5310 basis.take_back(elemsToSplit),
true);
5311 SmallVector<Value> results = llvm::to_vector(
5312 llvm::concat<Value>(delinearizeWithoutSplitPart.getResults(),
5313 delinearizeBack.getResults()));
5314 rewriter.
replaceOp(delinearizeOp, results);
5321void affine::AffineDelinearizeIndexOp::getCanonicalizationPatterns(
5322 RewritePatternSet &patterns, MLIRContext *context) {
5324 .
insert<CancelDelinearizeOfLinearizeDisjointExactTail,
5325 DropUnitExtentBasis, SplitDelinearizeSpanningLastLinearizeArg>(
5336 if (multiIndex.empty())
5337 return IndexType::get(ctx);
5338 return multiIndex.front().
getType();
5341void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5342 OperationState &odsState,
5345 if (!basis.empty() && basis.front() == Value())
5346 basis = basis.drop_front();
5347 SmallVector<Value> dynamicBasis;
5348 SmallVector<int64_t> staticBasis;
5352 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5356void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5357 OperationState &odsState,
5359 ArrayRef<OpFoldResult> basis,
5361 if (!basis.empty() && basis.front() == OpFoldResult())
5362 basis = basis.drop_front();
5363 SmallVector<Value> dynamicBasis;
5364 SmallVector<int64_t> staticBasis;
5367 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5371void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5372 OperationState &odsState,
5374 ArrayRef<int64_t> basis,
bool disjoint) {
5376 build(odsBuilder, odsState, resultType, multiIndex,
ValueRange{}, basis,
5380LogicalResult AffineLinearizeIndexOp::verify() {
5381 size_t numIndexes = getMultiIndex().size();
5382 size_t numBasisElems = getStaticBasis().size();
5383 if (numIndexes != numBasisElems && numIndexes != numBasisElems + 1)
5384 return emitOpError(
"should be passed a basis element for each index except "
5385 "possibly the first");
5387 auto dynamicMarkersCount =
5388 llvm::count_if(getStaticBasis(), ShapedType::isDynamic);
5389 if (
static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5391 "mismatch between dynamic and static basis (kDynamic marker but no "
5392 "corresponding dynamic basis entry) -- this can only happen due to an "
5393 "incorrect fold/rewrite");
5398OpFoldResult AffineLinearizeIndexOp::fold(FoldAdaptor adaptor) {
5399 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5401 adaptor.getDynamicBasis());
5402 if (maybeStaticBasis) {
5403 setStaticBasis(*maybeStaticBasis);
5407 if (getMultiIndex().empty())
5408 return IntegerAttr::get(getResult().
getType(), 0);
5411 if (getMultiIndex().size() == 1)
5412 return getMultiIndex().front();
5417 if (llvm::any_of(adaptor.getMultiIndex(), [](Attribute a) {
5418 return !isa_and_nonnull<IntegerAttr>(a);
5422 if (!adaptor.getDynamicBasis().empty())
5427 for (
auto [length, indexAttr] :
5428 llvm::zip_first(llvm::reverse(getStaticBasis()),
5429 llvm::reverse(adaptor.getMultiIndex()))) {
5430 result =
result + cast<IntegerAttr>(indexAttr).getInt() * stride;
5431 stride = stride * length;
5434 if (!hasOuterBound())
5437 cast<IntegerAttr>(adaptor.getMultiIndex().front()).getInt() * stride;
5442SmallVector<OpFoldResult> AffineLinearizeIndexOp::getEffectiveBasis() {
5444 if (hasOuterBound()) {
5445 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5447 getDynamicBasis().drop_front(), builder);
5449 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5453 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5456SmallVector<OpFoldResult> AffineLinearizeIndexOp::getPaddedBasis() {
5457 SmallVector<OpFoldResult> ret = getMixedBasis();
5458 if (!hasOuterBound())
5459 ret.insert(ret.begin(), OpFoldResult());
5474struct DropLinearizeUnitComponentsIfDisjointOrZero final
5475 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5478 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5479 PatternRewriter &rewriter)
const override {
5481 size_t numIndices = multiIndex.size();
5482 SmallVector<Value> newIndices;
5483 newIndices.reserve(numIndices);
5484 SmallVector<OpFoldResult> newBasis;
5485 newBasis.reserve(numIndices);
5487 if (!op.hasOuterBound()) {
5488 newIndices.push_back(multiIndex.front());
5489 multiIndex = multiIndex.drop_front();
5492 SmallVector<OpFoldResult> basis = op.getMixedBasis();
5493 for (
auto [index, basisElem] : llvm::zip_equal(multiIndex, basis)) {
5495 if (!basisEntry || *basisEntry != 1) {
5496 newIndices.push_back(index);
5497 newBasis.push_back(basisElem);
5502 if (!op.getDisjoint() && (!indexValue || *indexValue != 0)) {
5503 newIndices.push_back(index);
5504 newBasis.push_back(basisElem);
5508 if (newIndices.size() == numIndices)
5510 "no unit basis entries to replace");
5512 if (newIndices.empty()) {
5514 op, rewriter.
getZeroAttr(op.getLinearIndex().getType()));
5518 op, newIndices, newBasis, op.getDisjoint());
5524 ArrayRef<OpFoldResult> terms) {
5525 int64_t nDynamic = 0;
5526 SmallVector<Value> dynamicPart;
5528 for (OpFoldResult term : terms) {
5535 dynamicPart.push_back(cast<Value>(term));
5539 if (
auto constant = dyn_cast<AffineConstantExpr>(
result))
5541 return AffineApplyOp::create(builder, loc,
result, dynamicPart).getResult();
5571struct CancelLinearizeOfDelinearizePortion final
5572 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5582 unsigned linStart = 0;
5583 unsigned delinStart = 0;
5584 unsigned length = 0;
5588 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp linearizeOp,
5589 PatternRewriter &rewriter)
const override {
5590 SmallVector<Match> matches;
5592 const SmallVector<OpFoldResult> linBasis = linearizeOp.getPaddedBasis();
5593 ArrayRef<OpFoldResult> linBasisRef = linBasis;
5595 ValueRange multiIndex = linearizeOp.getMultiIndex();
5596 unsigned numLinArgs = multiIndex.size();
5597 unsigned linArgIdx = 0;
5600 llvm::SmallPtrSet<Operation *, 2> alreadyMatchedDelinearize;
5601 while (linArgIdx < numLinArgs) {
5602 auto asResult = dyn_cast<OpResult>(multiIndex[linArgIdx]);
5608 auto delinearizeOp =
5609 dyn_cast<AffineDelinearizeIndexOp>(asResult.getOwner());
5610 if (!delinearizeOp) {
5627 unsigned delinArgIdx = asResult.getResultNumber();
5628 SmallVector<OpFoldResult> delinBasis = delinearizeOp.getPaddedBasis();
5629 OpFoldResult firstDelinBound = delinBasis[delinArgIdx];
5630 OpFoldResult firstLinBound = linBasis[linArgIdx];
5631 bool boundsMatch = firstDelinBound == firstLinBound;
5632 bool bothAtFront = linArgIdx == 0 && delinArgIdx == 0;
5633 bool knownByDisjoint =
5634 linearizeOp.getDisjoint() && delinArgIdx == 0 && !firstDelinBound;
5635 if (!boundsMatch && !bothAtFront && !knownByDisjoint) {
5641 unsigned numDelinOuts = delinearizeOp.getNumResults();
5642 for (; j + linArgIdx < numLinArgs && j + delinArgIdx < numDelinOuts;
5644 if (multiIndex[linArgIdx + j] !=
5645 delinearizeOp.getResult(delinArgIdx + j))
5647 if (linBasis[linArgIdx + j] != delinBasis[delinArgIdx + j])
5653 if (j <= 1 || !alreadyMatchedDelinearize.insert(delinearizeOp).second) {
5657 matches.push_back(Match{delinearizeOp, linArgIdx, delinArgIdx, j});
5661 if (matches.empty())
5663 linearizeOp,
"no run of delinearize outputs to deal with");
5668 SmallVector<SmallVector<Value>> delinearizeReplacements;
5670 SmallVector<Value> newIndex;
5671 newIndex.reserve(numLinArgs);
5672 SmallVector<OpFoldResult> newBasis;
5673 newBasis.reserve(numLinArgs);
5674 unsigned prevMatchEnd = 0;
5675 for (Match m : matches) {
5676 unsigned gap = m.linStart - prevMatchEnd;
5677 llvm::append_range(newIndex, multiIndex.slice(prevMatchEnd, gap));
5678 llvm::append_range(newBasis, linBasisRef.slice(prevMatchEnd, gap));
5680 prevMatchEnd = m.linStart + m.length;
5682 PatternRewriter::InsertionGuard g(rewriter);
5685 ArrayRef<OpFoldResult> basisToMerge =
5686 linBasisRef.slice(m.linStart, m.length);
5689 OpFoldResult newSize =
5694 newIndex.push_back(m.delinearize.getLinearIndex());
5695 newBasis.push_back(newSize);
5697 delinearizeReplacements.push_back(SmallVector<Value>());
5701 SmallVector<Value> newDelinResults;
5702 SmallVector<OpFoldResult> newDelinBasis = m.delinearize.getPaddedBasis();
5703 newDelinBasis.erase(newDelinBasis.begin() + m.delinStart,
5704 newDelinBasis.begin() + m.delinStart + m.length);
5705 newDelinBasis.insert(newDelinBasis.begin() + m.delinStart, newSize);
5706 auto newDelinearize = AffineDelinearizeIndexOp::create(
5707 rewriter, m.delinearize.getLoc(), m.delinearize.getLinearIndex(),
5713 Value combinedElem = newDelinearize.getResult(m.delinStart);
5714 auto residualDelinearize = AffineDelinearizeIndexOp::create(
5715 rewriter, m.delinearize.getLoc(), combinedElem, basisToMerge);
5720 llvm::append_range(newDelinResults,
5721 newDelinearize.getResults().take_front(m.delinStart));
5722 llvm::append_range(newDelinResults, residualDelinearize.getResults());
5725 newDelinearize.getResults().drop_front(m.delinStart + 1));
5727 delinearizeReplacements.push_back(newDelinResults);
5728 newIndex.push_back(combinedElem);
5729 newBasis.push_back(newSize);
5731 llvm::append_range(newIndex, multiIndex.drop_front(prevMatchEnd));
5732 llvm::append_range(newBasis, linBasisRef.drop_front(prevMatchEnd));
5734 linearizeOp, newIndex, newBasis, linearizeOp.getDisjoint());
5736 for (
auto [m, newResults] :
5737 llvm::zip_equal(matches, delinearizeReplacements)) {
5738 if (newResults.empty())
5740 rewriter.
replaceOp(m.delinearize, newResults);
5751struct DropLinearizeLeadingZero final
5752 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5755 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5756 PatternRewriter &rewriter)
const override {
5757 Value leadingIdx = op.getMultiIndex().front();
5761 if (op.getMultiIndex().size() == 1) {
5766 SmallVector<OpFoldResult> mixedBasis = op.getMixedBasis();
5767 ArrayRef<OpFoldResult> newMixedBasis = mixedBasis;
5768 if (op.hasOuterBound())
5769 newMixedBasis = newMixedBasis.drop_front();
5772 op, op.getMultiIndex().drop_front(), newMixedBasis, op.getDisjoint());
5778void affine::AffineLinearizeIndexOp::getCanonicalizationPatterns(
5779 RewritePatternSet &patterns, MLIRContext *context) {
5780 patterns.
add<CancelLinearizeOfDelinearizePortion, DropLinearizeLeadingZero,
5781 DropLinearizeUnitComponentsIfDisjointOrZero>(context);
5788#define GET_OP_CLASSES
5789#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.