36#include "llvm/ADT/STLExtras.h"
37#include "llvm/ADT/SetVector.h"
38#include "llvm/ADT/SmallBitVector.h"
39#include "llvm/ADT/TypeSwitch.h"
40#include "llvm/Support/InterleavedRange.h"
42#include "llvm/Support/DebugLog.h"
46#define DEBUG_TYPE "parallel-loop-fusion"
49#define GEN_PASS_DEF_SCFPARALLELLOOPFUSION
50#include "mlir/Dialect/SCF/Transforms/Passes.h.inc"
60 return walkResult.wasInterrupted();
65 ParallelOp secondPloop) {
66 if (firstPloop.getNumLoops() != secondPloop.getNumLoops() ||
67 firstPloop.getUnsignedCmp() != secondPloop.getUnsignedCmp())
77 return std::equal(lhs.begin(), lhs.end(), rhs.begin(),
79 if (lhsValue == rhsValue)
81 std::optional<int64_t> lhsConst =
82 getConstantIntValue(lhsValue);
83 std::optional<int64_t> rhsConst =
84 getConstantIntValue(rhsValue);
85 return lhsConst && rhsConst && *lhsConst == *rhsConst;
88 return matchOperands(firstPloop.getLowerBound(),
89 secondPloop.getLowerBound()) &&
90 matchOperands(firstPloop.getUpperBound(),
91 secondPloop.getUpperBound()) &&
92 matchOperands(firstPloop.getStep(), secondPloop.getStep());
103 if (!isa<memref::StoreOp, vector::TransferWriteOp, vector::StoreOp>(op1))
105 bool opsAreIdentical =
107 .Case([&](memref::StoreOp storeOp1) {
108 auto storeOp2 = cast<memref::StoreOp>(op2);
109 return (storeOp1.getMemRef() == storeOp2.getMemRef()) &&
110 (storeOp1.getIndices() == storeOp2.getIndices());
112 .Case([&](vector::TransferWriteOp writeOp1) {
113 auto writeOp2 = cast<vector::TransferWriteOp>(op2);
114 return (writeOp1.getBase() == writeOp2.getBase()) &&
115 (writeOp1.getIndices() == writeOp2.getIndices()) &&
116 (writeOp1.getMask() == writeOp2.getMask()) &&
117 (writeOp1.getValueToStore().
getType() ==
118 writeOp2.getValueToStore().getType()) &&
119 (writeOp1.getInBounds() == writeOp2.getInBounds());
121 .Case([&](vector::StoreOp vecStoreOp1) {
122 auto vecStoreOp2 = cast<vector::StoreOp>(op2);
123 return (vecStoreOp1.getBase() == vecStoreOp2.getBase()) &&
124 (vecStoreOp1.getIndices() == vecStoreOp2.getIndices()) &&
125 (vecStoreOp1.getValueToStore().
getType() ==
126 vecStoreOp2.getValueToStore().getType()) &&
127 (vecStoreOp1.getAlignment() == vecStoreOp2.getAlignment()) &&
128 (vecStoreOp1.getNontemporal() ==
129 vecStoreOp2.getNontemporal());
131 .Default([](
Operation *) {
return false; });
132 return opsAreIdentical;
145 if (!val1DefOp || !val2DefOp)
150 val1DefOp, val2DefOp,
165 return constOp.value();
168 return constOp.value();
175 return constOp.value();
178 return constOp.value();
182 if (
auto applyOp = expr.
getDefiningOp<affine::AffineApplyOp>()) {
190 auto bin = dyn_cast<AffineBinaryOpExpr>(
result);
193 auto lhsDim = dyn_cast<AffineDimExpr>(bin.getLHS());
194 auto rhsDim = dyn_cast<AffineDimExpr>(bin.getRHS());
195 auto lhsConst = dyn_cast<AffineConstantExpr>(bin.getLHS());
196 auto rhsConst = dyn_cast<AffineConstantExpr>(bin.getRHS());
197 if (lhsConst && rhsDim)
198 return lhsConst.getValue();
199 if (rhsConst && lhsDim)
200 return rhsConst.getValue();
236 auto getConstLoopBoundsForIV =
237 [](
Value index) -> std::optional<std::tuple<int64_t, int64_t, int64_t>> {
238 auto blockArg = dyn_cast<BlockArgument>(
index);
241 auto *parentOp = blockArg.getOwner()->getParentOp();
242 auto loopLike = dyn_cast<LoopLikeOpInterface>(parentOp);
249 auto ivs = loopLike.getLoopInductionVars();
252 auto it = llvm::find(*ivs, blockArg);
253 if (it == ivs->end())
255 unsigned pos = std::distance(ivs->begin(), it);
256 if (pos >= ranges.size())
258 auto [lb,
ub, step] = ranges[pos];
259 return std::make_tuple(lb,
ub, step);
263 std::optional<int64_t> writeConst =
265 if (!writeConst && writeIndex) {
267 if (
auto bounds = getConstLoopBoundsForIV(writeIndex)) {
268 auto [lb,
ub, step] = *bounds;
269 if (step > 0 &&
ub == lb + step)
277 if (rangeExtent <= 0 || step <= 0)
281 int64_t rangeEnd = rangeStart + rangeExtent;
282 return lb >= rangeStart &&
ub <= rangeEnd;
285 if (offsetConst && writeConst) {
287 int64_t start = *offsetConst + *writeConst;
289 return (*loadConst >= start && *loadConst < start + extent);
290 if (
auto bounds = getConstLoopBoundsForIV(loadIndex)) {
291 auto [lb,
ub, step] = *bounds;
292 return loopIVWithinRange(lb,
ub, step, start, extent);
298 if (offsetConst && *offsetConst == 0 &&
301 if (
auto addConst =
getAddConstant(loadIndex, writeIndex, loopsIVsMap)) {
305 return (*addConst >= start && *addConst < start + extent);
311 if (
auto offsetVal = dyn_cast<Value>(offset)) {
324 .Case([&](memref::LoadOp
load) {
return load.getMemRef(); })
325 .Case([&](memref::StoreOp store) {
return store.getMemRef(); })
326 .Case([&](vector::TransferReadOp read) {
return read.getBase(); })
327 .Case([&](vector::TransferWriteOp write) {
return write.getBase(); })
328 .Case([&](vector::LoadOp
load) {
return load.getBase(); })
329 .Case([&](vector::StoreOp store) {
return store.getBase(); })
352 Value base = writeBase;
356 llvm::SmallBitVector droppedDims;
357 bool hasSubview =
false;
358 auto *ctx = loadOp.getContext();
359 if (
auto subView = base.
getDefiningOp<memref::SubViewOp>()) {
360 if (!subView.hasUnitStride())
362 baseMemref = cast<MemrefValue>(subView.getSource());
363 offsets = llvm::to_vector(subView.getMixedOffsets());
364 droppedDims = subView.getDroppedDims();
367 baseMemref = dyn_cast<MemrefValue>(base);
372 auto loadIndices = loadOp.getIndices();
373 unsigned baseRank = baseMemref.getType().getRank();
374 if ((loadOp.getMemref() != baseMemref) || (loadIndices.size() != baseRank))
377 unsigned writeRank = writeIndices.size();
378 if ((!hasSubview && writeRank != baseRank) ||
379 (hasSubview && offsets.size() != baseRank) ||
380 (vectorDimForWriteDim.size() != writeRank))
383 auto zeroAttr = IntegerAttr::get(IndexType::get(ctx), 0);
384 unsigned writeMemrefDim = 0;
385 for (
unsigned baseDim : llvm::seq(baseRank)) {
386 bool wasDropped = (hasSubview && droppedDims.test(baseDim));
387 int64_t vectorDim = !wasDropped ? vectorDimForWriteDim[writeMemrefDim] : -1;
389 if (vectorDim >= 0) {
390 int64_t dimSize = vecTy.getDimSize(vectorDim);
391 if (dimSize == ShapedType::kDynamic)
395 Value writeIndex = !wasDropped ? writeIndices[writeMemrefDim] :
Value();
411 vector::TransferWriteOp writeOp,
413 auto vecTy = dyn_cast<VectorType>(writeOp.getVector().getType());
417 unsigned writeRank = writeOp.getIndices().size();
425 for (
unsigned vecDim = 0; vecDim < permutationMap.
getNumResults(); ++vecDim) {
426 auto dimExpr = dyn_cast<AffineDimExpr>(permutationMap.
getResult(vecDim));
429 unsigned writeDim = dimExpr.getPosition();
430 if (writeDim >= writeRank || vectorDimForWriteDim[writeDim] != -1)
432 vectorDimForWriteDim[writeDim] = vecDim;
436 vecTy, vectorDimForWriteDim, ivsMap);
441 vector::StoreOp storeOp,
443 auto vecTy = dyn_cast<VectorType>(storeOp.getValueToStore().getType());
447 unsigned writeRank = storeOp.getIndices().size();
448 if (vecTy.getRank() > writeRank)
452 unsigned vecRank = vecTy.getRank();
453 for (
unsigned i = 0; i < vecRank; ++i) {
454 unsigned writeDim = writeRank - vecRank + i;
455 vectorDimForWriteDim[writeDim] = i;
459 vecTy, vectorDimForWriteDim, ivsMap);
469template <
typename OpTy1,
typename OpTy2>
471 OpTy1 op1, OpTy2 op2,
const IRMapping &firstToSecondPloopIVsMap,
475 if (!base1 || !base2)
478 auto accessThroughTrivialSubviewIsSame =
479 [&
b](memref::SubViewOp subView,
ValueRange subViewAccess,
482 LogicalResult resolved = resolveSourceIndicesRankReducingSubview(
483 subView.getLoc(),
b, subView, subViewAccess, resolvedSubviewAccess);
484 if (failed(resolved) ||
485 (resolvedSubviewAccess.size() != sourceAccess.size()))
487 for (
auto [dimIdx, resolvedIndex] :
488 llvm::enumerate(resolvedSubviewAccess)) {
497 if (
auto subView = base1.template getDefiningOp<memref::SubViewOp>();
500 base2, cast<MemrefValue>(subView.getSource())) &&
501 accessThroughTrivialSubviewIsSame(subView, op1.getIndices(),
503 firstToSecondPloopIVsMap))
507 if (
auto subView = base2.template getDefiningOp<memref::SubViewOp>();
510 base1, cast<MemrefValue>(subView.getSource())) &&
511 accessThroughTrivialSubviewIsSame(subView, op2.getIndices(),
513 firstToSecondPloopIVsMap))
522template <
typename OpTy1,
typename OpTy2>
525 auto indices1 = op1.getIndices();
526 auto indices2 = op2.getIndices();
527 if (indices1.size() != indices2.size())
529 for (
auto [idx1, idx2] : llvm::zip(indices1, indices2)) {
540 const IRMapping &firstToSecondPloopIVsMap,
542 if (!loadOp || !storeOp)
545 if (!isa<memref::LoadOp, vector::TransferReadOp, vector::LoadOp>(loadOp))
547 bool accessSameMemory =
549 .Case([&](memref::LoadOp memLoadOp) {
550 if (
auto memStoreOp = dyn_cast<memref::StoreOp>(storeOp))
552 firstToSecondPloopIVsMap,
b);
553 if (
auto vecWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp))
555 firstToSecondPloopIVsMap);
556 if (
auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp))
558 firstToSecondPloopIVsMap);
561 .Case([&](vector::TransferReadOp vecReadOp) {
562 auto vecWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp);
566 firstToSecondPloopIVsMap,
b) &&
567 (vecReadOp.getMask() == vecWriteOp.getMask()) &&
568 (vecReadOp.getInBounds() == vecWriteOp.getInBounds());
570 .Case([&](vector::LoadOp vecLoadOp) {
571 auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp);
575 firstToSecondPloopIVsMap,
b) &&
576 (vecLoadOp.getAlignment() == vecStoreOp.getAlignment());
578 .Default([](
Operation *) {
return false; });
579 return accessSameMemory;
584 .Case([&](memref::StoreOp storeOp) {
return storeOp.getMemRef(); })
585 .Case([&](vector::TransferWriteOp writeOp) {
return writeOp.getBase(); })
586 .Case([&](vector::StoreOp vecStoreOp) {
return vecStoreOp.getBase(); })
595 if (
auto transfWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp);
596 transfWriteOp && isa<memref::LoadOp>(loadOp))
599 if (
auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp);
600 vecStoreOp && isa<memref::LoadOp>(loadOp))
610 ParallelOp firstPloop, ParallelOp secondPloop,
611 const IRMapping &firstToSecondPloopIndices,
616 llvm::SmallSetVector<Value, 4> buffersWrittenInFirstPloop;
618 auto collectStoreOpsInWalk = [&](
Operation *op) {
619 auto memOpInterf = dyn_cast_if_present<MemoryEffectOpInterface>(op);
632 MemrefValue storeOpBaseMemref = dyn_cast<MemrefValue>(storeOpBase);
633 if (!storeOpBaseMemref)
637 bufferStoresInFirstPloop[buffer].push_back(op);
638 buffersWrittenInFirstPloop.insert(buffer);
644 if (firstPloop.getBody()->walk(collectStoreOpsInWalk).wasInterrupted())
653 auto checkLoadInWalkHasNoIncompatibleDataDeps = [&](
Operation *loadOp) {
654 auto memOpInterf = dyn_cast_if_present<MemoryEffectOpInterface>(loadOp);
670 if (!isa<memref::LoadOp, vector::TransferReadOp, vector::LoadOp>(loadOp) ||
677 for (
Value storedMem : buffersWrittenInFirstPloop)
678 if ((storedMem != loadedOrigBuf) &&
mayAlias(storedMem, loadedOrigBuf) &&
679 !llvm::all_of(bufferStoresInFirstPloop[storedMem],
682 firstToSecondPloopIndices);
687 auto writeOpsIt = bufferStoresInFirstPloop.find(loadedOrigBuf);
688 if (writeOpsIt == bufferStoresInFirstPloop.end())
694 if (writeOps.empty())
701 if (!llvm::all_of(writeOps, [&](
Operation *otherWriteOp) {
710 firstToSecondPloopIndices,
b)) {
719 return !secondPloop.getBody()
720 ->walk(checkLoadInWalkHasNoIncompatibleDataDeps)
729 const IRMapping &firstToSecondPloopIndices,
733 firstPloop, secondPloop, firstToSecondPloopIndices,
mayAlias,
b))
737 secondToFirstPloopIndices.
map(secondPloop.getBody()->getArguments(),
738 firstPloop.getBody()->getArguments());
740 secondPloop, firstPloop, secondToFirstPloopIndices,
mayAlias,
b);
747 const IRMapping &firstToSecondPloopIndices,
768static std::optional<ParallelOp>
771 assert(loop.getNumLoops() ==
indices.size());
772 if (loop.getNumLoops() < 2)
784 ParallelOp::create(builder, loop.getLoc(), newLB, newUB, newStep,
785 loop.getInitVals(),
nullptr, loop.getUnsignedCmp());
786 auto ivs = loop.getInductionVars();
790 for (
auto [iv, riv] : llvm::zip(ivs, newIvs)) {
791 mapping.
map(iv, riv);
796 for (
auto &o : loop.getNumReductions()
797 ? loop.getBodyRegion().front()
798 : loop.getBodyRegion().front().without_terminator()) {
820 return llvm::hash_combine(
833 ParallelOp &secondPloop,
834 int permBudget = 120) {
836 if (firstPloop.getNumLoops() < 2 ||
837 firstPloop.getNumLoops() != secondPloop.getNumLoops())
842 llvm::SmallSetVector<LoopIV, 6> unique;
843 for (
unsigned index : llvm::seq(firstPloop.getNumLoops())) {
844 firstIVs[
index].lBound = firstPloop.getLowerBound()[
index];
845 firstIVs[
index].uBound = firstPloop.getUpperBound()[
index];
846 firstIVs[
index].step = firstPloop.getStep()[
index];
847 secondIVs[
index].lBound = secondPloop.getLowerBound()[
index];
848 secondIVs[
index].uBound = secondPloop.getUpperBound()[
index];
849 secondIVs[
index].step = secondPloop.getStep()[
index];
850 unique.insert(firstIVs[
index]);
855 llvm::zip(firstIVs, secondIVs), diffIVs.begin(),
856 [](
auto const &pair) { return std::get<0>(pair) != std::get<1>(pair); });
859 for (
auto [idx, val] : enumerate(diffIVs))
869 std::iota(basic.begin(), basic.end(), 0);
871 if (
indices.empty() && unique.size() == firstIVs.size())
881 if (fIdx != sIdx && firstIVs[fIdx] == secondIVs[sIdx] &&
882 remaps.end() == std::find(remaps.begin(), remaps.end(), sIdx)) {
883 remaps.push_back(sIdx);
890 if (
indices.size() != remaps.size())
894 for (
auto [from, to] : zip(
indices, remaps)) {
898 LDBG() <<
"Collected basic permutations: "
899 << llvm::interleaved_array(basic);
902 if (unique.size() == firstIVs.size()) {
909 assert(unique.size() != firstIVs.size() &&
910 "Expected at least two equal axes");
915 for (
auto iv : unique) {
917 for (
unsigned index : llvm::seq(firstIVs.size())) {
918 if (firstIVs[
index] == iv)
919 group.push_back(
index);
921 if (group.size() > 1)
922 groups.push_back(std::move(group));
928 while (repeat && permBudget) {
930 for (
auto const &[group, groupRemaps] : zip(groups, rmpdGroups)) {
931 repeat |= std::next_permutation(groupRemaps.begin(), groupRemaps.end());
938 for (
auto const &[group, groupRemaps] : zip(groups, rmpdGroups)) {
939 for (
auto [from, to] : zip(group, groupRemaps))
940 extra[from] = basic[to];
942 if (basic != extra) {
943 LDBG() <<
"Collected extra permutations: "
944 << llvm::interleaved_array(extra);
946 extraResults.push_back(std::move(extra));
959 Block *block1 = firstPloop.getBody();
960 Block *block2 = secondPloop.getBody();
962 ValueRange inits2 = secondPloop.getInitVals();
965 newInitVars.append(inits2.begin(), inits2.end());
968 b.setInsertionPoint(secondPloop);
969 auto newSecondPloop =
970 ParallelOp::create(
b, secondPloop.getLoc(), secondPloop.getLowerBound(),
971 secondPloop.getUpperBound(), secondPloop.getStep(),
972 newInitVars,
nullptr, secondPloop.getUnsignedCmp());
974 Block *newBlock = newSecondPloop.getBody();
978 b.inlineBlockBefore(block2, newBlock, newBlock->
begin(),
980 b.inlineBlockBefore(block1, newBlock, newBlock->
begin(),
983 ValueRange results = newSecondPloop.getResults();
984 if (!results.empty()) {
985 b.setInsertionPointToEnd(newBlock);
990 newReduceArgs.append(reduceArgs2.begin(), reduceArgs2.end());
992 auto newReduceOp = scf::ReduceOp::create(
b, term2.getLoc(), newReduceArgs);
994 for (
auto &&[i, reg] : llvm::enumerate(llvm::concat<Region>(
995 term1.getReductions(), term2.getReductions()))) {
997 Block &newRedBlock = newReduceOp.getReductions()[i].
front();
998 b.inlineBlockBefore(&oldRedBlock, &newRedBlock, newRedBlock.
begin(),
1002 firstPloop.replaceAllUsesWith(results.take_front(inits1.size()));
1003 secondPloop.replaceAllUsesWith(results.take_back(inits2.size()));
1008 secondPloop.erase();
1009 secondPloop = newSecondPloop;
1013static void fuseIfLegal(ParallelOp firstPloop, ParallelOp &secondPloop,
1016 Block *block1 = firstPloop.getBody();
1017 Block *block2 = secondPloop.getBody();
1021 if (
isFusionLegal(firstPloop, secondPloop, firstToSecondPloopIndices,
1033 LDBG() <<
"Applied permutation: " << llvm::interleaved_array(perms);
1036 firstToSecondPloopIndices.
clear();
1038 newLoop->getBody()->getArguments());
1039 if (!
isFusionLegal(firstPloop, *newLoop, firstToSecondPloopIndices,
1041 LDBG() <<
"Rejected: " << newLoop;
1047 secondPloop.replaceAllUsesWith(newLoop->getResults());
1048 secondPloop->erase();
1049 secondPloop = *newLoop;
1060 for (
auto &block : region) {
1061 ploopChains.clear();
1062 ploopChains.push_back({});
1067 bool noSideEffects =
true;
1068 for (
auto &op : block) {
1069 if (
auto ploop = dyn_cast<ParallelOp>(op)) {
1070 if (noSideEffects) {
1071 ploopChains.back().push_back(ploop);
1073 ploopChains.push_back({ploop});
1074 noSideEffects =
true;
1082 for (
int i = 0, e = ploops.size(); i + 1 < e; ++i)
1089struct ParallelLoopFusion
1091 void runOnOperation()
override {
1092 auto &aa = getAnalysis<AliasAnalysis>();
1099 auto val2Def = val2.getDefiningOp();
1103 val2Def ? val2Def->getParentOfType<ParallelOp>() :
nullptr;
1104 if (val1Loop != val2Loop)
1107 return !aa.alias(val1, val2).isNo();
1110 getOperation()->walk([&](
Operation *child) {
1119 return std::make_unique<ParallelLoopFusion>();
static bool mayAlias(Value first, Value second)
Returns true if two values may be referencing aliasing memory.
static bool canResolveAlias(Operation *loadOp, Operation *storeOp, const IRMapping &loopsIVsMap)
To be called when mayAlias(val1, val2) is true.
static std::optional< ParallelOp > interchangeLoops(OpBuilder &builder, ParallelOp &loop, const ArrayRef< int64_t > &indices)
static bool equalIterationSpaces(ParallelOp firstPloop, ParallelOp secondPloop)
Verify equal iteration spaces.
static bool isLoadOnWrittenVector(memref::LoadOp loadOp, Value writeBase, ValueRange writeIndices, VectorType vecTy, ArrayRef< int64_t > vectorDimForWriteDim, const IRMapping &ivsMap)
Recognize scalar memref.load of an element produced by a vector write (vector.transfer_write or vecto...
static bool loadMatchesVectorWrite(memref::LoadOp loadOp, vector::TransferWriteOp writeOp, const IRMapping &ivsMap)
Recognize scalar memref.load of an element produced by a vector.transfer_write.
static std::optional< int64_t > getAddConstant(Value expr, Value base, const IRMapping &loopsIVsMap)
If the expr value is the result of an integer addition of base and a constant, return the constant.
static bool opsAccessSameIndices(OpTy1 op1, OpTy2 op2, const IRMapping &loopsIVsMap, OpBuilder &b)
Check if both memory read/write operations access the same indices (considering also the mapping of i...
static Value getStoreOpTargetBuffer(Operation *op)
static void applyLoopFusion(ParallelOp &firstPloop, ParallelOp &secondPloop, OpBuilder &builder)
Prepend operations of firstPloop's body into secondPloop's body.
static bool haveNoDataDependenciesExceptSameIndex(ParallelOp firstPloop, ParallelOp secondPloop, const IRMapping &firstToSecondPloopIndices, llvm::function_ref< bool(Value, Value)> mayAlias, OpBuilder &b)
Check that the parallel loops have no mixed access to the same buffers.
static Value getBaseMemref(Operation *op)
Return the base memref value used by the given memory op.
static bool loadsFromSameMemoryLocationWrittenBy(Operation *loadOp, Operation *storeOp, const IRMapping &firstToSecondPloopIVsMap, OpBuilder &b)
Check if the loadOp reads from the same memory location (same buffer, same indices and same propertie...
static SmallVector< SmallVector< int64_t > > computeCandidateInterchangePermutations(ParallelOp &firstPloop, ParallelOp &secondPloop, int permBudget=120)
static bool loadIndexWithinWriteRange(Value loadIndex, OpFoldResult offset, Value writeIndex, int64_t extent, const IRMapping &loopsIVsMap)
static bool opsWriteSameMemLocation(Operation *op1, Operation *op2)
Check if both operations are the same type of memory write op and write to the same memory location (...
static bool noIncompatibleDataDependencies(ParallelOp firstPloop, ParallelOp secondPloop, const IRMapping &firstToSecondPloopIndices, llvm::function_ref< bool(Value, Value)> mayAlias, OpBuilder &b)
Check that in each loop there are no read ops on the buffers written by the other loop,...
static bool valsAreEquivalent(Value val1, Value val2, const IRMapping &loopsIVsMap)
Check if val1 (from the first parallel loop) and val2 (from the second) are equivalent,...
static bool isFusionLegal(ParallelOp firstPloop, ParallelOp secondPloop, const IRMapping &firstToSecondPloopIndices, llvm::function_ref< bool(Value, Value)> mayAlias, OpBuilder &b)
Check if fusion of the two parallel loops is legal: i.e.
static bool opsAccessSameIndicesViaRankReducingSubview(OpTy1 op1, OpTy2 op2, const IRMapping &firstToSecondPloopIVsMap, OpBuilder &b)
Check if both operations access the same positions of the same buffer, but one of the two does it thr...
static bool loadMatchesVectorStore(memref::LoadOp loadOp, vector::StoreOp storeOp, const IRMapping &ivsMap)
Recognize scalar memref.load of an element produced by a vector.store.
static bool hasNestedParallelOp(ParallelOp ploop)
Verify there are no nested ParallelOps.
static void fuseIfLegal(ParallelOp firstPloop, ParallelOp &secondPloop, OpBuilder builder, llvm::function_ref< bool(Value, Value)> mayAlias)
Check fusion pre-conditions and call fusion if it is possible.
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumSymbols() const
unsigned getNumDims() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgListType getArguments()
A class for computing basic dominance information.
bool properlyDominates(Operation *a, Operation *b, bool enclosingOpOk=true) const
Return true if operation A properly dominates operation B, i.e.
This is a utility class for mapping one set of IR entities to another.
auto lookupOrDefault(T from) const
Lookup a mapped value within the map.
void clear()
Clears all mappings held by the mapper.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
static OpBuilder atBlockBegin(Block *block, Listener *listener=nullptr)
Create a builder and set the insertion point to before the first operation in the block but still ins...
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.
This trait indicates that the memory effects of an operation includes the effects of operations neste...
This class implements the operand iterators for the Operation class.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
unsigned getNumRegions()
Returns the number of regions held by this operation.
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
OperationName getName()
The name of an operation is the key identifier for it.
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
user_range getUsers()
Returns a range of all users.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static WalkResult advance()
static WalkResult interrupt()
MemrefValue skipFullyAliasingOperations(MemrefValue source)
Walk up the source chain until an operation that changes/defines the view of memory is found (i....
bool isSameViewOrTrivialAlias(MemrefValue a, MemrefValue b)
Checks if two (memref) values are the same or statically known to alias the same region of memory.
void naivelyFuseParallelOps(Region ®ion, llvm::function_ref< bool(Value, Value)> mayAlias)
Fuses all adjacent scf.parallel operations with identical bounds and step into one scf....
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
SmallVector< T > applyPermutation(ArrayRef< T > input, ArrayRef< int64_t > permutation)
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
llvm::SmallVector< std::tuple< int64_t, int64_t, int64_t > > getConstLoopBounds(mlir::LoopLikeOpInterface loopOp)
Get constant loop bounds and steps for each of the induction variables of the given loop operation,...
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
TypedValue< BaseMemRefType > MemrefValue
A value with a memref type.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
std::unique_ptr< Pass > createParallelLoopFusionPass()
Creates a loop fusion pass which fuses parallel loops.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
bool operator==(LoopIV const &other) const
bool operator!=(LoopIV const &other) const
static bool isEqual(const LoopIV &lhs, const LoopIV &rhs)
static unsigned getHashValue(const LoopIV &val)
The following effect indicates that the operation frees some resource that has been allocated.
The following effect indicates that the operation reads from some resource.
The following effect indicates that the operation writes to some resource.
static bool isEquivalentTo(Operation *lhs, Operation *rhs, function_ref< LogicalResult(Value, Value)> checkEquivalent, function_ref< void(Value, Value)> markEquivalent=nullptr, Flags flags=Flags::None, function_ref< LogicalResult(ValueRange, ValueRange)> checkCommutativeEquivalent=nullptr)
Compare two operations (including their regions) and return if they are equivalent.