MLIR 24.0.0git
ParallelLoopFusion.cpp
Go to the documentation of this file.
1//===- ParallelLoopFusion.cpp - Code to perform loop fusion ---------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements loop fusion on parallel loops.
10//
11//===----------------------------------------------------------------------===//
12
14
26#include "mlir/IR/Builders.h"
28#include "mlir/IR/IRMapping.h"
29#include "mlir/IR/Matchers.h"
33#include "mlir/IR/Value.h"
35
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"
41
42#include "llvm/Support/DebugLog.h"
43#include <numeric>
44#include <optional>
45#include <tuple>
46#define DEBUG_TYPE "parallel-loop-fusion"
47
48namespace mlir {
49#define GEN_PASS_DEF_SCFPARALLELLOOPFUSION
50#include "mlir/Dialect/SCF/Transforms/Passes.h.inc"
51} // namespace mlir
52
53using namespace mlir;
54using namespace mlir::scf;
55
56/// Verify there are no nested ParallelOps.
57static bool hasNestedParallelOp(ParallelOp ploop) {
58 auto walkResult =
59 ploop.getBody()->walk([](ParallelOp) { return WalkResult::interrupt(); });
60 return walkResult.wasInterrupted();
61}
62
63/// Verify equal iteration spaces.
64static bool equalIterationSpaces(ParallelOp firstPloop,
65 ParallelOp secondPloop) {
66 if (firstPloop.getNumLoops() != secondPloop.getNumLoops() ||
67 firstPloop.getUnsignedCmp() != secondPloop.getUnsignedCmp())
68 return false;
69
70 // Two bounds match if they are the same value, or if both are constants
71 // holding the same value. The latter matters because equivalent bounds are
72 // often materialized by distinct `arith.constant` ops, which leaves the
73 // iteration spaces equal even though the SSA values differ.
74 auto matchOperands = [&](const OperandRange &lhs,
75 const OperandRange &rhs) -> bool {
76 // TODO: Extend this to support aliases.
77 return std::equal(lhs.begin(), lhs.end(), rhs.begin(),
78 [](Value lhsValue, Value rhsValue) {
79 if (lhsValue == rhsValue)
80 return true;
81 std::optional<int64_t> lhsConst =
82 getConstantIntValue(lhsValue);
83 std::optional<int64_t> rhsConst =
84 getConstantIntValue(rhsValue);
85 return lhsConst && rhsConst && *lhsConst == *rhsConst;
86 });
87 };
88 return matchOperands(firstPloop.getLowerBound(),
89 secondPloop.getLowerBound()) &&
90 matchOperands(firstPloop.getUpperBound(),
91 secondPloop.getUpperBound()) &&
92 matchOperands(firstPloop.getStep(), secondPloop.getStep());
93}
94
95/// Check if both operations are the same type of memory write op and
96/// write to the same memory location (same buffer and same indices).
98 if (!op1 || !op2 || op1->getName() != op2->getName())
99 return false;
100 if (op1 == op2)
101 return true;
102 // support only these memory-writing ops for now
103 if (!isa<memref::StoreOp, vector::TransferWriteOp, vector::StoreOp>(op1))
104 return false;
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());
111 })
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());
120 })
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());
130 })
131 .Default([](Operation *) { return false; });
132 return opsAreIdentical;
133}
134
135/// Check if val1 (from the first parallel loop) and val2 (from the
136/// second) are equivalent, considering the mapping of induction variables from
137/// the first to the second parallel loop.
138static bool valsAreEquivalent(Value val1, Value val2,
139 const IRMapping &loopsIVsMap) {
140 if (val1 == val2 || loopsIVsMap.lookupOrDefault(val1) == val2 ||
141 loopsIVsMap.lookupOrDefault(val2) == val1)
142 return true;
143 Operation *val1DefOp = val1.getDefiningOp();
144 Operation *val2DefOp = val2.getDefiningOp();
145 if (!val1DefOp || !val2DefOp)
146 return false;
147 if (!isMemoryEffectFree(val1DefOp) || !isMemoryEffectFree(val2DefOp))
148 return false;
150 val1DefOp, val2DefOp,
151 [&](Value v1, Value v2) {
152 return success(loopsIVsMap.lookupOrDefault(v1) == v2 ||
153 loopsIVsMap.lookupOrDefault(v2) == v1);
154 },
155 /*markEquivalent=*/nullptr, OperationEquivalence::Flags::IgnoreLocations);
156}
157
158/// If the `expr` value is the result of an integer addition of `base` and a
159/// constant, return the constant.
160static std::optional<int64_t> getAddConstant(Value expr, Value base,
161 const IRMapping &loopsIVsMap) {
162 if (auto addOp = expr.getDefiningOp<arith::AddIOp>()) {
163 if (auto constOp = getConstantIntValue(addOp.getLhs());
164 constOp && valsAreEquivalent(addOp.getRhs(), base, loopsIVsMap))
165 return constOp.value();
166 if (auto constOp = getConstantIntValue(addOp.getRhs());
167 constOp && valsAreEquivalent(addOp.getLhs(), base, loopsIVsMap))
168 return constOp.value();
169 return std::nullopt;
170 }
171
172 if (auto addOp = expr.getDefiningOp<index::AddOp>()) {
173 if (auto constOp = getConstantIntValue(addOp.getLhs());
174 constOp && valsAreEquivalent(addOp.getRhs(), base, loopsIVsMap))
175 return constOp.value();
176 if (auto constOp = getConstantIntValue(addOp.getRhs());
177 constOp && valsAreEquivalent(addOp.getLhs(), base, loopsIVsMap))
178 return constOp.value();
179 return std::nullopt;
180 }
181
182 if (auto applyOp = expr.getDefiningOp<affine::AffineApplyOp>()) {
183 AffineMap map = applyOp.getAffineMap();
184 if (map.getNumResults() != 1 || map.getNumDims() != 1 ||
185 map.getNumSymbols() != 0)
186 return std::nullopt;
187 if (!valsAreEquivalent(applyOp.getOperand(0), base, loopsIVsMap))
188 return std::nullopt;
189 AffineExpr result = map.getResult(0);
190 auto bin = dyn_cast<AffineBinaryOpExpr>(result);
191 if (!bin || bin.getKind() != AffineExprKind::Add)
192 return std::nullopt;
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();
201 }
202 return std::nullopt;
203}
204
205// Return true if the scalar load index may hit any element covered by a
206// vector.store/transfer_write along a single memref dimension. Supported cases:
207//
208// 1) Direct index match (with optional offset):
209// vector.transfer_write %v, %A[%i] : vector<4xf32>, memref<...>
210// %x = memref.load %A[%i] : memref<...>
211//
212// 2) Loop IV range intersects the write range:
213// vector.transfer_write %v, %A[%c0] : vector<4xf32>, memref<...>
214// scf.for %k = %c0 to %c4 step %c1 { %x = memref.load %A[%k] }
215//
216// 3) Constant index (or IV + constant) within the write range:
217// vector.transfer_write %v, %A[%c0] : vector<4xf32>, memref<...>
218// %x = memref.load %A[%c2] : memref<...>
219// %y = memref.load %A[%i + %c1] : memref<...>
220//
221// Args:
222// - loadIndex: index used by the scalar load for this dimension.
223// - offset: subview offset for the base memref dimension (if any).
224// - writeIndex: index used by the transfer_write for this dimension. Can be
225// null if the dim was dropped by a rank reducing subview, whose result is
226// written by the vector.write.
227// - extent: vector size along this dimension (number of elements written).
228// - loopsIVsMap: IV equivalence map between fused loops.
229static bool loadIndexWithinWriteRange(Value loadIndex, OpFoldResult offset,
230 Value writeIndex, int64_t extent,
231 const IRMapping &loopsIVsMap) {
232 if (extent <= 0)
233 return false;
234
235 // Extract constant loop bounds for loop IVs (e.g. from scf.for).
236 auto getConstLoopBoundsForIV =
237 [](Value index) -> std::optional<std::tuple<int64_t, int64_t, int64_t>> {
238 auto blockArg = dyn_cast<BlockArgument>(index);
239 if (!blockArg)
240 return std::nullopt;
241 auto *parentOp = blockArg.getOwner()->getParentOp();
242 auto loopLike = dyn_cast<LoopLikeOpInterface>(parentOp);
243 if (!loopLike)
244 return std::nullopt;
245 auto ranges = getConstLoopBounds(loopLike);
246 if (ranges.empty())
247 return std::nullopt;
248
249 auto ivs = loopLike.getLoopInductionVars();
250 if (!ivs)
251 return std::nullopt;
252 auto it = llvm::find(*ivs, blockArg);
253 if (it == ivs->end())
254 return std::nullopt;
255 unsigned pos = std::distance(ivs->begin(), it);
256 if (pos >= ranges.size())
257 return std::nullopt;
258 auto [lb, ub, step] = ranges[pos];
259 return std::make_tuple(lb, ub, step);
260 };
261
262 std::optional<int64_t> offsetConst = getConstantIntValue(offset);
263 std::optional<int64_t> writeConst =
264 writeIndex ? getConstantIntValue(writeIndex) : std::optional<int64_t>(0);
265 if (!writeConst && writeIndex) {
266 // Treat single-iteration IVs as constants for matching.
267 if (auto bounds = getConstLoopBoundsForIV(writeIndex)) {
268 auto [lb, ub, step] = *bounds;
269 if (step > 0 && ub == lb + step)
270 writeConst = lb;
271 }
272 }
273
274 // Check whether a loop IV is fully contained in a constant write range.
275 auto loopIVWithinRange = [](int64_t lb, int64_t ub, int64_t step,
276 int64_t rangeStart, int64_t rangeExtent) -> bool {
277 if (rangeExtent <= 0 || step <= 0)
278 return false;
279 if (ub <= lb)
280 return false;
281 int64_t rangeEnd = rangeStart + rangeExtent;
282 return lb >= rangeStart && ub <= rangeEnd;
283 };
284
285 if (offsetConst && writeConst) {
286 // Constant start of the write range; check constant load or loop IV range.
287 int64_t start = *offsetConst + *writeConst;
288 if (auto loadConst = getConstantIntValue(loadIndex))
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);
293 }
294 }
295
296 if (writeIndex) {
297 // Direct IV match (or IV + constant) against the write index.
298 if (offsetConst && *offsetConst == 0 &&
299 valsAreEquivalent(loadIndex, writeIndex, loopsIVsMap))
300 return true;
301 if (auto addConst = getAddConstant(loadIndex, writeIndex, loopsIVsMap)) {
302 // Match load index of the form writeIndex + C within the write extent.
303 if (offsetConst) {
304 int64_t start = *offsetConst;
305 return (*addConst >= start && *addConst < start + extent);
306 }
307 }
308 return false;
309 }
310
311 if (auto offsetVal = dyn_cast<Value>(offset)) {
312 // Exact match when extent is 1 and the load hits the offset value.
313 if (extent == 1 && valsAreEquivalent(loadIndex, offsetVal, loopsIVsMap))
314 return true;
315 }
316
317 return false;
318}
319
320/// Return the base memref value used by the given memory op.
322 // TODO: use the common interface for memory ops once available.
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(); })
330 .Default([](Operation *) { return Value(); });
331}
332
333/// Recognize scalar memref.load of an element produced by a vector write
334/// (vector.transfer_write or vector.store, optionally through a rank-reducing
335/// unit-stride subview) of the same buffer. This covers the pattern where a
336/// vector write stores a full lane pack and a subsequent scalar load reads an
337/// element from that lane pack. EXAMPLE:
338/// vector.transfer_write %V, %arg[%x, %y, ..., 0] {in_bounds = [true]} :
339/// vector<4xf32>, memref<4xf32, strided<[1], offset: ?>>
340/// scf.for %iter = %c0 to %c4 step %c1 iter_args(...) -> (f32) {
341/// %0 = memref.load %arg[%x, %y, ..., %iter] : memref<1x128x16x4xf32>
342/// ...
343/// }
344///
345static bool isLoadOnWrittenVector(memref::LoadOp loadOp, Value writeBase,
346 ValueRange writeIndices, VectorType vecTy,
347 ArrayRef<int64_t> vectorDimForWriteDim,
348 const IRMapping &ivsMap) {
349 if (!vecTy)
350 return false;
351
352 Value base = writeBase;
353 // The write base if there is no subview, or the subview source otherwise.
354 MemrefValue baseMemref = nullptr;
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())
361 return false;
362 baseMemref = cast<MemrefValue>(subView.getSource());
363 offsets = llvm::to_vector(subView.getMixedOffsets());
364 droppedDims = subView.getDroppedDims();
365 hasSubview = true;
366 } else {
367 baseMemref = dyn_cast<MemrefValue>(base);
368 if (!baseMemref)
369 return false;
370 }
371
372 auto loadIndices = loadOp.getIndices();
373 unsigned baseRank = baseMemref.getType().getRank();
374 if ((loadOp.getMemref() != baseMemref) || (loadIndices.size() != baseRank))
375 return false;
376
377 unsigned writeRank = writeIndices.size();
378 if ((!hasSubview && writeRank != baseRank) ||
379 (hasSubview && offsets.size() != baseRank) ||
380 (vectorDimForWriteDim.size() != writeRank))
381 return false;
382
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;
388 int64_t extent = 1;
389 if (vectorDim >= 0) {
390 int64_t dimSize = vecTy.getDimSize(vectorDim);
391 if (dimSize == ShapedType::kDynamic)
392 return false;
393 extent = dimSize;
394 }
395 Value writeIndex = !wasDropped ? writeIndices[writeMemrefDim] : Value();
396 OpFoldResult offset =
397 hasSubview ? offsets[baseDim] : OpFoldResult(zeroAttr);
398 if (!loadIndexWithinWriteRange(loadIndices[baseDim], offset, writeIndex,
399 extent, ivsMap))
400 return false;
401 if (!wasDropped)
402 ++writeMemrefDim;
403 }
404
405 return true;
406}
407
408/// Recognize scalar memref.load of an element produced by a
409/// vector.transfer_write
410static bool loadMatchesVectorWrite(memref::LoadOp loadOp,
411 vector::TransferWriteOp writeOp,
412 const IRMapping &ivsMap) {
413 auto vecTy = dyn_cast<VectorType>(writeOp.getVector().getType());
414 if (!vecTy)
415 return false;
416
417 unsigned writeRank = writeOp.getIndices().size();
418 AffineMap permutationMap = writeOp.getPermutationMap();
419 if (!permutationMap.isProjectedPermutation() ||
420 permutationMap.getNumResults() != vecTy.getRank() ||
421 permutationMap.getNumDims() != writeRank)
422 return false;
423
424 SmallVector<int64_t> vectorDimForWriteDim(writeRank, -1);
425 for (unsigned vecDim = 0; vecDim < permutationMap.getNumResults(); ++vecDim) {
426 auto dimExpr = dyn_cast<AffineDimExpr>(permutationMap.getResult(vecDim));
427 if (!dimExpr)
428 return false;
429 unsigned writeDim = dimExpr.getPosition();
430 if (writeDim >= writeRank || vectorDimForWriteDim[writeDim] != -1)
431 return false;
432 vectorDimForWriteDim[writeDim] = vecDim;
433 }
434
435 return isLoadOnWrittenVector(loadOp, writeOp.getBase(), writeOp.getIndices(),
436 vecTy, vectorDimForWriteDim, ivsMap);
437}
438
439/// Recognize scalar memref.load of an element produced by a vector.store
440static bool loadMatchesVectorStore(memref::LoadOp loadOp,
441 vector::StoreOp storeOp,
442 const IRMapping &ivsMap) {
443 auto vecTy = dyn_cast<VectorType>(storeOp.getValueToStore().getType());
444 if (!vecTy)
445 return false;
446
447 unsigned writeRank = storeOp.getIndices().size();
448 if (vecTy.getRank() > writeRank)
449 return false;
450
451 SmallVector<int64_t> vectorDimForWriteDim(writeRank, -1);
452 unsigned vecRank = vecTy.getRank();
453 for (unsigned i = 0; i < vecRank; ++i) {
454 unsigned writeDim = writeRank - vecRank + i;
455 vectorDimForWriteDim[writeDim] = i;
456 }
457
458 return isLoadOnWrittenVector(loadOp, storeOp.getBase(), storeOp.getIndices(),
459 vecTy, vectorDimForWriteDim, ivsMap);
460}
461
462/// Check if both operations access the same positions of the same
463/// buffer, but one of the two does it through a rank-reducing full subview of
464/// the buffer (the other's base). EXAMPLE:
465/// memref.store %a, %buf[%c0, %i, %j] : memref<1x2x2xf32>
466/// %alias = memref.subview %buf[0, 0, 0][1, 2, 2][1, 1, 1]: memref<1x2x2xf32>
467/// to memref<2x2xf32>
468/// %val = memref.load %alias[%i, %j] : memref<2x2xf32>
469template <typename OpTy1, typename OpTy2>
471 OpTy1 op1, OpTy2 op2, const IRMapping &firstToSecondPloopIVsMap,
472 OpBuilder &b) {
473 auto base1 = cast<MemrefValue>(getBaseMemref(op1));
474 auto base2 = cast<MemrefValue>(getBaseMemref(op2));
475 if (!base1 || !base2)
476 return false;
477
478 auto accessThroughTrivialSubviewIsSame =
479 [&b](memref::SubViewOp subView, ValueRange subViewAccess,
480 ValueRange sourceAccess, const IRMapping &ivsMap) -> bool {
481 SmallVector<Value> resolvedSubviewAccess;
482 LogicalResult resolved = resolveSourceIndicesRankReducingSubview(
483 subView.getLoc(), b, subView, subViewAccess, resolvedSubviewAccess);
484 if (failed(resolved) ||
485 (resolvedSubviewAccess.size() != sourceAccess.size()))
486 return false;
487 for (auto [dimIdx, resolvedIndex] :
488 llvm::enumerate(resolvedSubviewAccess)) {
489 if (!matchPattern(resolvedIndex, m_Zero()) &&
490 !valsAreEquivalent(resolvedIndex, sourceAccess[dimIdx], ivsMap))
491 return false;
492 }
493 return true;
494 };
495
496 // Case 1: op1 uses a subview of op2's base.
497 if (auto subView = base1.template getDefiningOp<memref::SubViewOp>();
498 subView &&
500 base2, cast<MemrefValue>(subView.getSource())) &&
501 accessThroughTrivialSubviewIsSame(subView, op1.getIndices(),
502 op2.getIndices(),
503 firstToSecondPloopIVsMap))
504 return true;
505
506 // Case 2: op2 uses a subview of op1's base.
507 if (auto subView = base2.template getDefiningOp<memref::SubViewOp>();
508 subView &&
510 base1, cast<MemrefValue>(subView.getSource())) &&
511 accessThroughTrivialSubviewIsSame(subView, op2.getIndices(),
512 op1.getIndices(),
513 firstToSecondPloopIVsMap))
514 return true;
515
516 return false;
517}
518
519/// Check if both memory read/write operations access the same indices
520/// (considering also the mapping of induction variables from the first to the
521/// second parallel loop).
522template <typename OpTy1, typename OpTy2>
523static bool opsAccessSameIndices(OpTy1 op1, OpTy2 op2,
524 const IRMapping &loopsIVsMap, OpBuilder &b) {
525 auto indices1 = op1.getIndices();
526 auto indices2 = op2.getIndices();
527 if (indices1.size() != indices2.size())
528 return opsAccessSameIndicesViaRankReducingSubview(op1, op2, loopsIVsMap, b);
529 for (auto [idx1, idx2] : llvm::zip(indices1, indices2)) {
530 if (!valsAreEquivalent(idx1, idx2, loopsIVsMap))
531 return false;
532 }
533 return true;
534}
535
536/// Check if the loadOp reads from the same memory location (same buffer,
537/// same indices and same properties) as written by the storeOp.
538static bool
540 const IRMapping &firstToSecondPloopIVsMap,
541 OpBuilder &b) {
542 if (!loadOp || !storeOp)
543 return false;
544 // Support only these memory-reading ops for now
545 if (!isa<memref::LoadOp, vector::TransferReadOp, vector::LoadOp>(loadOp))
546 return false;
547 bool accessSameMemory =
549 .Case([&](memref::LoadOp memLoadOp) {
550 if (auto memStoreOp = dyn_cast<memref::StoreOp>(storeOp))
551 return opsAccessSameIndices(memLoadOp, memStoreOp,
552 firstToSecondPloopIVsMap, b);
553 if (auto vecWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp))
554 return loadMatchesVectorWrite(memLoadOp, vecWriteOp,
555 firstToSecondPloopIVsMap);
556 if (auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp))
557 return loadMatchesVectorStore(memLoadOp, vecStoreOp,
558 firstToSecondPloopIVsMap);
559 return false;
560 })
561 .Case([&](vector::TransferReadOp vecReadOp) {
562 auto vecWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp);
563 if (!vecWriteOp)
564 return false;
565 return opsAccessSameIndices(vecReadOp, vecWriteOp,
566 firstToSecondPloopIVsMap, b) &&
567 (vecReadOp.getMask() == vecWriteOp.getMask()) &&
568 (vecReadOp.getInBounds() == vecWriteOp.getInBounds());
569 })
570 .Case([&](vector::LoadOp vecLoadOp) {
571 auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp);
572 if (!vecStoreOp)
573 return false;
574 return opsAccessSameIndices(vecLoadOp, vecStoreOp,
575 firstToSecondPloopIVsMap, b) &&
576 (vecLoadOp.getAlignment() == vecStoreOp.getAlignment());
577 })
578 .Default([](Operation *) { return false; });
579 return accessSameMemory;
581
584 .Case([&](memref::StoreOp storeOp) { return storeOp.getMemRef(); })
585 .Case([&](vector::TransferWriteOp writeOp) { return writeOp.getBase(); })
586 .Case([&](vector::StoreOp vecStoreOp) { return vecStoreOp.getBase(); })
587 .Default([](Operation *) { return Value(); });
589
590/// To be called when `mayAlias(val1, val2)` is true. Check if the potential
591/// aliasing between the loadOp and storeOp can be resolved by analyzing their
592/// access patterns.
593static bool canResolveAlias(Operation *loadOp, Operation *storeOp,
594 const IRMapping &loopsIVsMap) {
595 if (auto transfWriteOp = dyn_cast<vector::TransferWriteOp>(storeOp);
596 transfWriteOp && isa<memref::LoadOp>(loadOp))
597 return loadMatchesVectorWrite(cast<memref::LoadOp>(loadOp), transfWriteOp,
598 loopsIVsMap);
599 if (auto vecStoreOp = dyn_cast<vector::StoreOp>(storeOp);
600 vecStoreOp && isa<memref::LoadOp>(loadOp))
601 return loadMatchesVectorStore(cast<memref::LoadOp>(loadOp), vecStoreOp,
602 loopsIVsMap);
603 return false;
604}
605
606/// Check that the parallel loops have no mixed access to the same buffers.
607/// Return `true` if the second parallel loop does not read or write the buffers
608/// written by the first loop using different indices.
610 ParallelOp firstPloop, ParallelOp secondPloop,
611 const IRMapping &firstToSecondPloopIndices,
613 // Map buffers to their store/write ops in the firstPloop
614 DenseMap<Value, SmallVector<Operation *>> bufferStoresInFirstPloop;
615 // Record all the memory buffers used in store/write ops found in firstPloop
616 llvm::SmallSetVector<Value, 4> buffersWrittenInFirstPloop;
617
618 auto collectStoreOpsInWalk = [&](Operation *op) {
619 auto memOpInterf = dyn_cast_if_present<MemoryEffectOpInterface>(op);
620 // Ignore ops that don't write to memory
621 if (!memOpInterf || (!memOpInterf.hasEffect<MemoryEffects::Write>() &&
622 !memOpInterf.hasEffect<MemoryEffects::Free>()))
623 return WalkResult::advance();
624
625 // Only these memory-writing ops are supported for now:
626 // memref.store, vector.transfer_write, vector.store
627 Value storeOpBase = getStoreOpTargetBuffer(op);
628 if (!storeOpBase)
629 return WalkResult::interrupt();
630
631 // Expect the base operand to be a Memref
632 MemrefValue storeOpBaseMemref = dyn_cast<MemrefValue>(storeOpBase);
633 if (!storeOpBaseMemref)
634 return WalkResult::interrupt();
635 // Get the original memref buffer, skipping full view-like ops
636 Value buffer = memref::skipFullyAliasingOperations(storeOpBaseMemref);
637 bufferStoresInFirstPloop[buffer].push_back(op);
638 buffersWrittenInFirstPloop.insert(buffer);
639 return WalkResult::advance();
640 };
641
642 // Walk the first parallel loop to collect all store/write ops and their
643 // target buffers
644 if (firstPloop.getBody()->walk(collectStoreOpsInWalk).wasInterrupted())
645 return false;
646
647 // Check that this load/read op encountered while walking the second parallel
648 // loop does not have incompatible data dependencies with the store/write ops
649 // collected from the first parallel loop: the loops can be fused only if in
650 // the 2nd loop there are no loads/stores from/to the buffers written in the
651 // 1st loop, except when on the same exact memory location (same indices) as
652 // written in the 1st loop.
653 auto checkLoadInWalkHasNoIncompatibleDataDeps = [&](Operation *loadOp) {
654 auto memOpInterf = dyn_cast_if_present<MemoryEffectOpInterface>(loadOp);
655 // To be conservative, we should stop on ops that don't advertise their
656 // memory effects. However, many ops don't implement MemoryEffectOpInterface
657 // yet, so for now we just skip them.
658 // TODO: once more ops add MemoryEffectOpInterface, interrupt the walk here.
659 if (!memOpInterf &&
661 return WalkResult::advance();
662 // Ignore ops that don't read from memory, and wrapping ops that have nested
663 // memory effects (e.g. loops, conditionals) as they will be analyzed when
664 // visiting their nested ops.
665 if ((!memOpInterf &&
667 (memOpInterf && !memOpInterf.hasEffect<MemoryEffects::Read>()))
668 return WalkResult::advance();
669 // Support only these memory-reading ops for now
670 if (!isa<memref::LoadOp, vector::TransferReadOp, vector::LoadOp>(loadOp) ||
671 !isa<MemrefValue>(loadOp->getOperand(0)))
672 return WalkResult::interrupt();
673
674 MemrefValue loadOpBase = cast<MemrefValue>(loadOp->getOperand(0));
675 MemrefValue loadedOrigBuf = memref::skipFullyAliasingOperations(loadOpBase);
676
677 for (Value storedMem : buffersWrittenInFirstPloop)
678 if ((storedMem != loadedOrigBuf) && mayAlias(storedMem, loadedOrigBuf) &&
679 !llvm::all_of(bufferStoresInFirstPloop[storedMem],
680 [&](Operation *storeOp) {
681 return canResolveAlias(loadOp, storeOp,
682 firstToSecondPloopIndices);
683 })) {
684 return WalkResult::interrupt();
685 }
686
687 auto writeOpsIt = bufferStoresInFirstPloop.find(loadedOrigBuf);
688 if (writeOpsIt == bufferStoresInFirstPloop.end())
689 return WalkResult::advance();
690 // Store/write ops to this buffer in the firstPloop
691 SmallVector<mlir::Operation *> &writeOps = writeOpsIt->second;
692
693 // If the first loop has no writes to this buffer, continue
694 if (writeOps.empty())
695 return WalkResult::advance();
696
697 Operation *writeOp = writeOps.front();
698
699 // In the first parallel loop, multiple writes to the same memref are
700 // allowed only on the same memory location
701 if (!llvm::all_of(writeOps, [&](Operation *otherWriteOp) {
702 return opsWriteSameMemLocation(writeOp, otherWriteOp);
703 })) {
704 return WalkResult::interrupt();
705 }
706
707 // Check that the load in secondPloop reads from the same memory location as
708 // written by the corresponding store in firstPloop
709 if (!loadsFromSameMemoryLocationWrittenBy(loadOp, writeOp,
710 firstToSecondPloopIndices, b)) {
711 return WalkResult::interrupt();
712 }
713
714 return WalkResult::advance();
715 };
716
717 // Walk the second parallel loop to check load/read ops against the stores
718 // collected from the first parallel loop.
719 return !secondPloop.getBody()
720 ->walk(checkLoadInWalkHasNoIncompatibleDataDeps)
721 .wasInterrupted();
722}
723
724/// Check that in each loop there are no read ops on the buffers written
725/// by the other loop, except when reading from the same exact memory location
726/// (same indices) as written in the other loop.
727static bool
728noIncompatibleDataDependencies(ParallelOp firstPloop, ParallelOp secondPloop,
729 const IRMapping &firstToSecondPloopIndices,
731 OpBuilder &b) {
733 firstPloop, secondPloop, firstToSecondPloopIndices, mayAlias, b))
734 return false;
735
736 IRMapping secondToFirstPloopIndices;
737 secondToFirstPloopIndices.map(secondPloop.getBody()->getArguments(),
738 firstPloop.getBody()->getArguments());
740 secondPloop, firstPloop, secondToFirstPloopIndices, mayAlias, b);
741}
742
743/// Check if fusion of the two parallel loops is legal:
744/// i.e. no nested parallel loops, equal iteration spaces,
745/// and no incompatible data dependencies between the loops.
746static bool isFusionLegal(ParallelOp firstPloop, ParallelOp secondPloop,
747 const IRMapping &firstToSecondPloopIndices,
749 OpBuilder &b) {
750 if (hasNestedParallelOp(firstPloop) || hasNestedParallelOp(secondPloop) ||
751 !equalIterationSpaces(firstPloop, secondPloop) ||
752 !noIncompatibleDataDependencies(firstPloop, secondPloop,
753 firstToSecondPloopIndices, mayAlias, b))
754 return false;
755
756 // We are fusing first loop into second, make sure there are no users of the
757 // first loop results between loops.
758 DominanceInfo dom;
759 for (Operation *user : firstPloop->getUsers()) {
760 if (!dom.properlyDominates(secondPloop, user, /*enclosingOpOk*/ false))
761 return false;
762 }
763 return true;
764}
765
766// Returns new parallel loop where two loops matching indices param are
767// interchanged
768static std::optional<ParallelOp>
769interchangeLoops(OpBuilder &builder, ParallelOp &loop,
770 const ArrayRef<int64_t> &indices) {
771 assert(loop.getNumLoops() == indices.size());
772 if (loop.getNumLoops() < 2)
773 return std::nullopt;
774
775 // Replace the parallel loop with the same parallel loop.
776 builder.setInsertionPoint(loop);
777 SmallVector<Value> newLB =
778 applyPermutation(SmallVector<Value>(loop.getLowerBound()), indices);
779 SmallVector<Value> newUB =
780 applyPermutation(SmallVector<Value>(loop.getUpperBound()), indices);
781 SmallVector<Value> newStep =
783 auto newOp =
784 ParallelOp::create(builder, loop.getLoc(), newLB, newUB, newStep,
785 loop.getInitVals(), nullptr, loop.getUnsignedCmp());
786 auto ivs = loop.getInductionVars();
788 newOp.getInductionVars(), invertPermutationVector(indices));
789 IRMapping mapping;
790 for (auto [iv, riv] : llvm::zip(ivs, newIvs)) {
791 mapping.map(iv, riv);
792 }
793
794 // Copy parallel loop body
795 auto b = OpBuilder::atBlockBegin(newOp.getBody());
796 for (auto &o : loop.getNumReductions()
797 ? loop.getBodyRegion().front()
798 : loop.getBodyRegion().front().without_terminator()) {
799 b.clone(o, mapping);
800 }
801 return newOp;
802}
803
804struct LoopIV {
806 bool operator!=(LoopIV const &other) const { return !(*this == other); }
807 bool operator==(LoopIV const &other) const {
808 return lBound == other.lBound && uBound == other.uBound &&
809 step == other.step;
810 }
811};
812
813template <>
815 static inline bool isEqual(const LoopIV &lhs, const LoopIV &rhs) {
816 return (lhs == rhs);
817 }
818
819 static inline unsigned getHashValue(const LoopIV &val) {
820 return llvm::hash_combine(
824 }
825};
826
827// Returns vector of candidate permutation indices vectors,
828// can be empty. Caps the number of extra candidate permutations
829// explored to avoid combinatorial explosion. This makes the search
830// intentionally incomplete.
833 ParallelOp &secondPloop,
834 int permBudget = 120) {
835 // Check preconditions
836 if (firstPloop.getNumLoops() < 2 ||
837 firstPloop.getNumLoops() != secondPloop.getNumLoops())
838 return {};
839
840 SmallVector<LoopIV> firstIVs(firstPloop.getNumLoops());
841 SmallVector<LoopIV> secondIVs(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]);
851 }
852
853 SmallVector<bool> diffIVs(firstPloop.getNumLoops());
854 llvm::transform(
855 llvm::zip(firstIVs, secondIVs), diffIVs.begin(),
856 [](auto const &pair) { return std::get<0>(pair) != std::get<1>(pair); });
857
859 for (auto [idx, val] : enumerate(diffIVs))
860 if (val)
861 indices.push_back(idx);
862
863 // Not a permutation shortcut
864 if (indices.size() == 1)
865 return {};
866
867 // Initialize with identity permutations
868 SmallVector<int64_t> basic(firstIVs.size());
869 std::iota(basic.begin(), basic.end(), 0);
870
871 if (indices.empty() && unique.size() == firstIVs.size())
872 return {};
873
874 if (indices.size() > 1) {
875 // Determine whether the iteration space of the first loop is a permutation
876 // of the second and collect remaps.
878 for (auto fIdx : indices) {
879 for (auto sIdx : indices) {
880 // can be remapped
881 if (fIdx != sIdx && firstIVs[fIdx] == secondIVs[sIdx] &&
882 remaps.end() == std::find(remaps.begin(), remaps.end(), sIdx)) {
883 remaps.push_back(sIdx);
884 break;
885 }
886 }
887 }
888
889 // Not a permutation
890 if (indices.size() != remaps.size())
891 return {};
892
893 // compose permutation indices
894 for (auto [from, to] : zip(indices, remaps)) {
895 basic[from] = to;
896 }
897
898 LDBG() << "Collected basic permutations: "
899 << llvm::interleaved_array(basic);
900
901 // All axes are unique, no further permutatons needed
902 if (unique.size() == firstIVs.size()) {
903 return {basic};
904 }
905 }
906
907 //
908 // Permute equal axes
909 assert(unique.size() != firstIVs.size() &&
910 "Expected at least two equal axes");
911
912 // Collect equal axes to groups
913 SmallVector<SmallVector<int64_t>> extraResults{basic};
915 for (auto iv : unique) {
917 for (unsigned index : llvm::seq(firstIVs.size())) {
918 if (firstIVs[index] == iv)
919 group.push_back(index);
920 }
921 if (group.size() > 1)
922 groups.push_back(std::move(group));
923 }
924
925 // Permute axes groups
926 SmallVector<SmallVector<int64_t>> rmpdGroups(groups);
927 bool repeat = true;
928 while (repeat && permBudget) {
929 repeat = false;
930 for (auto const &[group, groupRemaps] : zip(groups, rmpdGroups)) {
931 repeat |= std::next_permutation(groupRemaps.begin(), groupRemaps.end());
932 if (repeat)
933 break;
934 }
935
936 if (repeat) {
937 SmallVector<int64_t> extra(basic);
938 for (auto const &[group, groupRemaps] : zip(groups, rmpdGroups)) {
939 for (auto [from, to] : zip(group, groupRemaps))
940 extra[from] = basic[to];
941 }
942 if (basic != extra) {
943 LDBG() << "Collected extra permutations: "
944 << llvm::interleaved_array(extra);
945
946 extraResults.push_back(std::move(extra));
947 permBudget--;
948 }
949 }
950 }
951
952 return extraResults;
953}
954
955/// Prepend operations of firstPloop's body into secondPloop's body.
956/// Update secondPloop with new loop.
957static void applyLoopFusion(ParallelOp &firstPloop, ParallelOp &secondPloop,
958 OpBuilder &builder) {
959 Block *block1 = firstPloop.getBody();
960 Block *block2 = secondPloop.getBody();
961 ValueRange inits1 = firstPloop.getInitVals();
962 ValueRange inits2 = secondPloop.getInitVals();
963
964 SmallVector<Value> newInitVars(inits1.begin(), inits1.end());
965 newInitVars.append(inits2.begin(), inits2.end());
966
967 IRRewriter b(builder);
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());
973
974 Block *newBlock = newSecondPloop.getBody();
975 auto term1 = cast<ReduceOp>(block1->getTerminator());
976 auto term2 = cast<ReduceOp>(block2->getTerminator());
977
978 b.inlineBlockBefore(block2, newBlock, newBlock->begin(),
979 newBlock->getArguments());
980 b.inlineBlockBefore(block1, newBlock, newBlock->begin(),
981 newBlock->getArguments());
982
983 ValueRange results = newSecondPloop.getResults();
984 if (!results.empty()) {
985 b.setInsertionPointToEnd(newBlock);
986
987 ValueRange reduceArgs1 = term1.getOperands();
988 ValueRange reduceArgs2 = term2.getOperands();
989 SmallVector<Value> newReduceArgs(reduceArgs1.begin(), reduceArgs1.end());
990 newReduceArgs.append(reduceArgs2.begin(), reduceArgs2.end());
991
992 auto newReduceOp = scf::ReduceOp::create(b, term2.getLoc(), newReduceArgs);
993
994 for (auto &&[i, reg] : llvm::enumerate(llvm::concat<Region>(
995 term1.getReductions(), term2.getReductions()))) {
996 Block &oldRedBlock = reg.front();
997 Block &newRedBlock = newReduceOp.getReductions()[i].front();
998 b.inlineBlockBefore(&oldRedBlock, &newRedBlock, newRedBlock.begin(),
999 newRedBlock.getArguments());
1000 }
1001
1002 firstPloop.replaceAllUsesWith(results.take_front(inits1.size()));
1003 secondPloop.replaceAllUsesWith(results.take_back(inits2.size()));
1004 }
1005 term1->erase();
1006 term2->erase();
1007 firstPloop.erase();
1008 secondPloop.erase();
1009 secondPloop = newSecondPloop;
1010}
1011
1012/// Check fusion pre-conditions and call fusion if it is possible
1013static void fuseIfLegal(ParallelOp firstPloop, ParallelOp &secondPloop,
1014 OpBuilder builder,
1016 Block *block1 = firstPloop.getBody();
1017 Block *block2 = secondPloop.getBody();
1018 IRMapping firstToSecondPloopIndices;
1019 firstToSecondPloopIndices.map(block1->getArguments(), block2->getArguments());
1020
1021 if (isFusionLegal(firstPloop, secondPloop, firstToSecondPloopIndices,
1022 mayAlias, builder)) {
1023 applyLoopFusion(firstPloop, secondPloop, builder);
1024 return;
1025 }
1026
1027 // If iteration space of the second parallel loop is a permutation of the
1028 // first one then interchange iteration space of the second parallel loop
1029 // and re-asses possibility of fusion.
1030 for (auto &perms :
1031 computeCandidateInterchangePermutations(firstPloop, secondPloop)) {
1032 OpBuilder::InsertionGuard guard(builder);
1033 LDBG() << "Applied permutation: " << llvm::interleaved_array(perms);
1034
1035 auto newLoop = interchangeLoops(builder, secondPloop, perms);
1036 firstToSecondPloopIndices.clear();
1037 firstToSecondPloopIndices.map(block1->getArguments(),
1038 newLoop->getBody()->getArguments());
1039 if (!isFusionLegal(firstPloop, *newLoop, firstToSecondPloopIndices,
1040 mayAlias, builder)) {
1041 LDBG() << "Rejected: " << newLoop;
1042
1043 newLoop->erase();
1044 continue;
1045 }
1046
1047 secondPloop.replaceAllUsesWith(newLoop->getResults());
1048 secondPloop->erase();
1049 secondPloop = *newLoop;
1050 applyLoopFusion(firstPloop, secondPloop, builder);
1051 break;
1052 }
1053}
1054
1056 Region &region, llvm::function_ref<bool(Value, Value)> mayAlias) {
1057 OpBuilder b(region);
1058 // Consider every single block and attempt to fuse adjacent loops.
1060 for (auto &block : region) {
1061 ploopChains.clear();
1062 ploopChains.push_back({});
1063
1064 // Not using `walk()` to traverse only top-level parallel loops and also
1065 // make sure that there are no side-effecting ops between the parallel
1066 // loops.
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);
1072 } else {
1073 ploopChains.push_back({ploop});
1074 noSideEffects = true;
1075 }
1076 continue;
1077 }
1078 // TODO: Handle region side effects properly.
1079 noSideEffects &= isMemoryEffectFree(&op) && op.getNumRegions() == 0;
1080 }
1081 for (MutableArrayRef<ParallelOp> ploops : ploopChains) {
1082 for (int i = 0, e = ploops.size(); i + 1 < e; ++i)
1083 fuseIfLegal(ploops[i], ploops[i + 1], b, mayAlias);
1084 }
1085 }
1086}
1087
1088namespace {
1089struct ParallelLoopFusion
1090 : public impl::SCFParallelLoopFusionBase<ParallelLoopFusion> {
1091 void runOnOperation() override {
1092 auto &aa = getAnalysis<AliasAnalysis>();
1093
1094 auto mayAlias = [&](Value val1, Value val2) -> bool {
1095 // If the memref is defined in one of the parallel loops body, careful
1096 // alias analysis is needed.
1097 // TODO: check if this is still needed as a separate check.
1098 auto val1Def = val1.getDefiningOp();
1099 auto val2Def = val2.getDefiningOp();
1100 auto val1Loop =
1101 val1Def ? val1Def->getParentOfType<ParallelOp>() : nullptr;
1102 auto val2Loop =
1103 val2Def ? val2Def->getParentOfType<ParallelOp>() : nullptr;
1104 if (val1Loop != val2Loop)
1105 return true;
1106
1107 return !aa.alias(val1, val2).isNo();
1108 };
1109
1110 getOperation()->walk([&](Operation *child) {
1111 for (Region &region : child->getRegions())
1113 });
1114 }
1115};
1116} // namespace
1117
1118std::unique_ptr<Pass> mlir::createParallelLoopFusionPass() {
1119 return std::make_unique<ParallelLoopFusion>();
1120}
return success()
static bool mayAlias(Value first, Value second)
Returns true if two values may be referencing aliasing memory.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
auto load
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.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
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.
Definition Block.h:34
Operation & front()
Definition Block.h:178
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgListType getArguments()
Definition Block.h:112
iterator begin()
Definition Block.h:168
A class for computing basic dominance information.
Definition Dominance.h:143
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.
Definition IRMapping.h:26
auto lookupOrDefault(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:65
void clear()
Clears all mappings held by the mapper.
Definition IRMapping.h:79
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
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.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
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...
Definition Builders.h:243
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
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.
Definition ValueRange.h:44
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
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 &region, 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.
Definition Matchers.h:490
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.
Definition Utils.cpp:311
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,...
Definition Utils.cpp:1659
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
TypedValue< BaseMemRefType > MemrefValue
A value with a memref type.
Definition MemRefUtils.h:26
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
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.