MLIR 24.0.0git
SCFToOpenMP.cpp
Go to the documentation of this file.
1//===- SCFToOpenMP.cpp - Structured Control Flow to OpenMP conversion -----===//
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 a pass to convert scf.parallel operations into OpenMP
10// parallel loops.
11//
12//===----------------------------------------------------------------------===//
13
15
23#include "mlir/IR/SymbolTable.h"
24#include "mlir/Pass/Pass.h"
26
27namespace mlir {
28#define GEN_PASS_DEF_CONVERTSCFTOOPENMPPASS
29#include "mlir/Conversion/Passes.h.inc"
30} // namespace mlir
31
32using namespace mlir;
33
34/// Matches a block containing a "simple" reduction. The expected shape of the
35/// block is as follows.
36///
37/// ^bb(%arg0, %arg1):
38/// %0 = OpTy(%arg0, %arg1)
39/// scf.reduce.return %0
40template <typename... OpTy>
41static bool matchSimpleReduction(Block &block) {
42 if (block.empty() || llvm::hasSingleElement(block) ||
43 std::next(block.begin(), 2) != block.end())
44 return false;
45
46 if (block.getNumArguments() != 2)
47 return false;
48
50 Value reducedVal = matchReduction({block.getArguments()[1]},
51 /*redPos=*/0, combinerOps);
52
53 if (!reducedVal || !isa<BlockArgument>(reducedVal) || combinerOps.size() != 1)
54 return false;
55
56 return isa<OpTy...>(combinerOps[0]) &&
57 isa<scf::ReduceReturnOp>(block.back()) &&
58 block.front().getOperands() == block.getArguments();
59}
60
61/// Matches a block containing a select-based min/max reduction. The types of
62/// select and compare operations are provided as template arguments. The
63/// comparison predicates suitable for min and max are provided as function
64/// arguments. If a reduction is matched, `ifMin` will be set if the reduction
65/// compute the minimum and unset if it computes the maximum, otherwise it
66/// remains unmodified. The expected shape of the block is as follows.
67///
68/// ^bb(%arg0, %arg1):
69/// %0 = CompareOpTy(<one-of-predicates>, %arg0, %arg1)
70/// %1 = SelectOpTy(%0, %arg0, %arg1) // %arg0, %arg1 may be swapped here.
71/// scf.reduce.return %1
72template <
73 typename CompareOpTy, typename SelectOpTy,
74 typename Predicate = decltype(std::declval<CompareOpTy>().getPredicate())>
75static bool
77 ArrayRef<Predicate> greaterThanPredicates, bool &isMin) {
78 static_assert(
79 llvm::is_one_of<SelectOpTy, arith::SelectOp, LLVM::SelectOp>::value,
80 "only arithmetic and llvm select ops are supported");
81
82 // Expect exactly three operations in the block.
83 if (block.empty() || llvm::hasSingleElement(block) ||
84 std::next(block.begin(), 2) == block.end() ||
85 std::next(block.begin(), 3) != block.end())
86 return false;
87
88 // Check op kinds.
89 auto compare = dyn_cast<CompareOpTy>(block.front());
90 auto select = dyn_cast<SelectOpTy>(block.front().getNextNode());
91 auto terminator = dyn_cast<scf::ReduceReturnOp>(block.back());
92 if (!compare || !select || !terminator)
93 return false;
94
95 // Block arguments must be compared.
96 if (compare->getOperands() != block.getArguments())
97 return false;
98
99 // Detect whether the comparison is less-than or greater-than, otherwise bail.
100 bool isLess;
101 if (llvm::is_contained(lessThanPredicates, compare.getPredicate())) {
102 isLess = true;
103 } else if (llvm::is_contained(greaterThanPredicates,
104 compare.getPredicate())) {
105 isLess = false;
106 } else {
107 return false;
108 }
109
110 if (select.getCondition() != compare.getResult())
111 return false;
112
113 // Detect if the operands are swapped between cmpf and select. Match the
114 // comparison type with the requested type or with the opposite of the
115 // requested type if the operands are swapped. Use generic accessors because
116 // std and LLVM versions of select have different operand names but identical
117 // positions.
118 constexpr unsigned kTrueValue = 1;
119 constexpr unsigned kFalseValue = 2;
120 bool sameOperands = select.getOperand(kTrueValue) == compare.getLhs() &&
121 select.getOperand(kFalseValue) == compare.getRhs();
122 bool swappedOperands = select.getOperand(kTrueValue) == compare.getRhs() &&
123 select.getOperand(kFalseValue) == compare.getLhs();
124 if (!sameOperands && !swappedOperands)
125 return false;
126
127 if (select.getResult() != terminator.getResult())
128 return false;
129
130 // The reduction is a min if it uses less-than predicates with same operands
131 // or greather-than predicates with swapped operands. Similarly for max.
132 isMin = (isLess && sameOperands) || (!isLess && swappedOperands);
133 return isMin || (isLess & swappedOperands) || (!isLess && sameOperands);
134}
135
136/// Returns the float semantics for the given float type.
137static const llvm::fltSemantics &fltSemanticsForType(FloatType type) {
138 if (type.isF16())
139 return llvm::APFloat::IEEEhalf();
140 if (type.isF32())
141 return llvm::APFloat::IEEEsingle();
142 if (type.isF64())
143 return llvm::APFloat::IEEEdouble();
144 if (type.isF128())
145 return llvm::APFloat::IEEEquad();
146 if (type.isBF16())
147 return llvm::APFloat::BFloat();
148 if (type.isF80())
149 return llvm::APFloat::x87DoubleExtended();
150 llvm_unreachable("unknown float type");
151}
152
153/// Helper to create a splat attribute for vector types, or return the scalar
154/// attribute for scalar types.
156 if (auto vecType = dyn_cast<VectorType>(type))
157 return DenseElementsAttr::get(vecType, val);
158 return val;
159}
160
161/// Returns an attribute with the minimum (if `min` is set) or the maximum value
162/// (otherwise) for the given float type.
164 Type elType = getElementTypeOrSelf(type);
165 auto fltType = cast<FloatType>(elType);
166 auto val = llvm::APFloat::getLargest(fltSemanticsForType(fltType), min);
167
168 return getSplatOrScalarAttr(type, FloatAttr::get(elType, val));
169}
170
171/// Returns an attribute with the signed integer minimum (if `min` is set) or
172/// the maximum value (otherwise) for the given integer type, regardless of its
173/// signedness semantics (only the width is considered).
175 Type elType = getElementTypeOrSelf(type);
176 auto intType = cast<IntegerType>(elType);
177 unsigned bitwidth = intType.getWidth();
178 auto val = min ? llvm::APInt::getSignedMinValue(bitwidth)
179 : llvm::APInt::getSignedMaxValue(bitwidth);
180
181 return getSplatOrScalarAttr(type, IntegerAttr::get(elType, val));
182}
183
184/// Returns an attribute with the unsigned integer minimum (if `min` is set) or
185/// the maximum value (otherwise) for the given integer type, regardless of its
186/// signedness semantics (only the width is considered).
188 Type elType = getElementTypeOrSelf(type);
189 auto intType = cast<IntegerType>(elType);
190 unsigned bitwidth = intType.getWidth();
191 auto val =
192 min ? llvm::APInt::getZero(bitwidth) : llvm::APInt::getAllOnes(bitwidth);
193
194 return getSplatOrScalarAttr(type, IntegerAttr::get(elType, val));
195}
196
197/// Creates an OpenMP reduction declaration and inserts it into the provided
198/// symbol table. The declaration has a constant initializer with the neutral
199/// value `initValue`, and the `reductionIndex`-th reduction combiner carried
200/// over from `reduce`.
201static omp::DeclareReductionOp
203 scf::ReduceOp reduce, int64_t reductionIndex, Attribute initValue) {
204 OpBuilder::InsertionGuard guard(builder);
205 Type type = reduce.getOperands()[reductionIndex].getType();
206 auto decl = omp::DeclareReductionOp::create(builder, reduce.getLoc(),
207 "__scf_reduction",
208 /*sym_visibility=*/nullptr, type,
209 /*byref_element_type=*/{});
210 symbolTable.insert(decl);
211
212 builder.createBlock(&decl.getInitializerRegion(),
213 decl.getInitializerRegion().end(), {type},
214 {reduce.getOperands()[reductionIndex].getLoc()});
215 builder.setInsertionPointToEnd(&decl.getInitializerRegion().back());
216 Value init =
217 LLVM::ConstantOp::create(builder, reduce.getLoc(), type, initValue);
218 omp::YieldOp::create(builder, reduce.getLoc(), init);
219
220 Operation *terminator =
221 &reduce.getReductions()[reductionIndex].front().back();
222 assert(isa<scf::ReduceReturnOp>(terminator) &&
223 "expected reduce op to be terminated by reduce return");
224 builder.setInsertionPoint(terminator);
225 builder.replaceOpWithNewOp<omp::YieldOp>(terminator,
226 terminator->getOperands());
227 builder.inlineRegionBefore(reduce.getReductions()[reductionIndex],
228 decl.getReductionRegion(),
229 decl.getReductionRegion().end());
230 return decl;
231}
232
233/// Adds an atomic reduction combiner to the given OpenMP reduction declaration
234/// using llvm.atomicrmw of the given kind.
235static omp::DeclareReductionOp addAtomicRMW(OpBuilder &builder,
236 LLVM::AtomicBinOp atomicKind,
237 omp::DeclareReductionOp decl,
238 scf::ReduceOp reduce,
239 int64_t reductionIndex) {
240 OpBuilder::InsertionGuard guard(builder);
241 auto ptrType = LLVM::LLVMPointerType::get(builder.getContext());
242 Location reduceOperandLoc = reduce.getOperands()[reductionIndex].getLoc();
243 builder.createBlock(&decl.getAtomicReductionRegion(),
244 decl.getAtomicReductionRegion().end(), {ptrType, ptrType},
245 {reduceOperandLoc, reduceOperandLoc});
246 Block *atomicBlock = &decl.getAtomicReductionRegion().back();
247 builder.setInsertionPointToEnd(atomicBlock);
248 Value loaded = LLVM::LoadOp::create(builder, reduce.getLoc(), decl.getType(),
249 atomicBlock->getArgument(1));
250 LLVM::AtomicRMWOp::create(builder, reduce.getLoc(), atomicKind,
251 atomicBlock->getArgument(0), loaded,
252 LLVM::AtomicOrdering::monotonic);
253 omp::YieldOp::create(builder, reduce.getLoc(), ArrayRef<Value>());
254 return decl;
255}
256
257/// Returns true if the type is supported by llvm.atomicrmw.
258/// LLVM IR currently does not support atomic operations on vector types.
259/// See LLVM Language Reference Manual on 'atomicrmw'.
260static bool supportsAtomic(Type type) { return !isa<VectorType>(type); }
261
262/// Creates an OpenMP reduction declaration that corresponds to the given SCF
263/// reduction and returns it. Recognizes common reductions in order to identify
264/// the neutral value, necessary for the OpenMP declaration. If the reduction
265/// cannot be recognized, returns null.
266static omp::DeclareReductionOp declareReduction(PatternRewriter &builder,
267 scf::ReduceOp reduce,
268 int64_t reductionIndex) {
270 SymbolTable symbolTable(container);
271
272 // Insert reduction declarations in the symbol-table ancestor before the
273 // ancestor of the current insertion point.
274 Operation *insertionPoint = reduce;
275 while (insertionPoint->getParentOp() != container)
276 insertionPoint = insertionPoint->getParentOp();
277 OpBuilder::InsertionGuard guard(builder);
278 builder.setInsertionPoint(insertionPoint);
279
280 assert(llvm::hasSingleElement(reduce.getReductions()[reductionIndex]) &&
281 "expected reduction region to have a single element");
282
283 // Match simple binary reductions that can be expressed with atomicrmw.
284 Type type = reduce.getOperands()[reductionIndex].getType();
285 Block &reduction = reduce.getReductions()[reductionIndex].front();
286
287 // Handle scalar element type extraction for vector bitwidth safety.
288 Type elType = getElementTypeOrSelf(type);
289
290 // Arithmetic Reductions
292 omp::DeclareReductionOp decl = createDecl(
293 builder, symbolTable, reduce, reductionIndex,
294 getSplatOrScalarAttr(type, builder.getFloatAttr(elType, 0.0)));
295 return supportsAtomic(type) ? addAtomicRMW(builder, LLVM::AtomicBinOp::fadd,
296 decl, reduce, reductionIndex)
297 : decl;
298 }
300 omp::DeclareReductionOp decl = createDecl(
301 builder, symbolTable, reduce, reductionIndex,
302 getSplatOrScalarAttr(type, builder.getIntegerAttr(elType, 0)));
303 return supportsAtomic(type) ? addAtomicRMW(builder, LLVM::AtomicBinOp::add,
304 decl, reduce, reductionIndex)
305 : decl;
306 }
308 omp::DeclareReductionOp decl = createDecl(
309 builder, symbolTable, reduce, reductionIndex,
310 getSplatOrScalarAttr(type, builder.getIntegerAttr(elType, 0)));
311 return supportsAtomic(type) ? addAtomicRMW(builder, LLVM::AtomicBinOp::_or,
312 decl, reduce, reductionIndex)
313 : decl;
314 }
316 omp::DeclareReductionOp decl = createDecl(
317 builder, symbolTable, reduce, reductionIndex,
318 getSplatOrScalarAttr(type, builder.getIntegerAttr(elType, 0)));
319 return supportsAtomic(type) ? addAtomicRMW(builder, LLVM::AtomicBinOp::_xor,
320 decl, reduce, reductionIndex)
321 : decl;
322 }
324 APInt allOnes = llvm::APInt::getAllOnes(elType.getIntOrFloatBitWidth());
325 omp::DeclareReductionOp decl = createDecl(
326 builder, symbolTable, reduce, reductionIndex,
327 getSplatOrScalarAttr(type, builder.getIntegerAttr(elType, allOnes)));
328 return supportsAtomic(type) ? addAtomicRMW(builder, LLVM::AtomicBinOp::_and,
329 decl, reduce, reductionIndex)
330 : decl;
331 }
332
333 // Match simple binary reductions that cannot be expressed with atomicrmw.
334 // TODO: add atomic region using cmpxchg (which needs atomic load to be
335 // available as an op).
337 return createDecl(
338 builder, symbolTable, reduce, reductionIndex,
339 getSplatOrScalarAttr(type, builder.getFloatAttr(elType, 1.0)));
340 }
341
343 return createDecl(
344 builder, symbolTable, reduce, reductionIndex,
345 getSplatOrScalarAttr(type, builder.getIntegerAttr(elType, 1)));
346 }
347
348 // Match select-based min/max reductions.
349 bool isMin;
350 // Floating Point Min/Max
351 if (matchSelectReduction<arith::CmpFOp, arith::SelectOp,
352 arith::CmpFPredicate>(
353 reduction, {arith::CmpFPredicate::OLT, arith::CmpFPredicate::OLE},
354 {arith::CmpFPredicate::OGT, arith::CmpFPredicate::OGE}, isMin) ||
355 matchSelectReduction<arith::CmpFOp, arith::SelectOp,
356 arith::CmpFPredicate>(
357 reduction, {arith::CmpFPredicate::OGT, arith::CmpFPredicate::OGE},
358 {arith::CmpFPredicate::OLT, arith::CmpFPredicate::OLE}, isMin)) {
359 return createDecl(builder, symbolTable, reduce, reductionIndex,
360 minMaxValueForFloat(type, !isMin));
361 }
362
363 // Integer Min/Max
364 if (matchSelectReduction<arith::CmpIOp, arith::SelectOp,
365 arith::CmpIPredicate>(
366 reduction, {arith::CmpIPredicate::slt, arith::CmpIPredicate::sle},
367 {arith::CmpIPredicate::sgt, arith::CmpIPredicate::sge}, isMin) ||
368 matchSelectReduction<arith::CmpIOp, arith::SelectOp,
369 arith::CmpIPredicate>(
370 reduction, {arith::CmpIPredicate::sgt, arith::CmpIPredicate::sge},
371 {arith::CmpIPredicate::slt, arith::CmpIPredicate::sle}, isMin)) {
372 omp::DeclareReductionOp decl =
373 createDecl(builder, symbolTable, reduce, reductionIndex,
374 minMaxValueForSignedInt(type, !isMin));
375 return supportsAtomic(type) ? addAtomicRMW(builder,
376 isMin ? LLVM::AtomicBinOp::min
377 : LLVM::AtomicBinOp::max,
378 decl, reduce, reductionIndex)
379 : decl;
380 }
381
382 // Unsigned Integer Min/Max
383 if (matchSelectReduction<arith::CmpIOp, arith::SelectOp,
384 arith::CmpIPredicate>(
385 reduction, {arith::CmpIPredicate::ult, arith::CmpIPredicate::ule},
386 {arith::CmpIPredicate::ugt, arith::CmpIPredicate::uge}, isMin) ||
387 matchSelectReduction<arith::CmpIOp, arith::SelectOp,
388 arith::CmpIPredicate>(
389 reduction, {arith::CmpIPredicate::ugt, arith::CmpIPredicate::uge},
390 {arith::CmpIPredicate::ult, arith::CmpIPredicate::ule}, isMin)) {
391 omp::DeclareReductionOp decl =
392 createDecl(builder, symbolTable, reduce, reductionIndex,
393 minMaxValueForUnsignedInt(type, !isMin));
394 return supportsAtomic(type) ? addAtomicRMW(builder,
395 isMin ? LLVM::AtomicBinOp::umin
396 : LLVM::AtomicBinOp::umax,
397 decl, reduce, reductionIndex)
398 : decl;
399 }
400
401 return nullptr;
402}
403
404namespace {
405
406struct ParallelOpLowering : public OpRewritePattern<scf::ParallelOp> {
407 static constexpr unsigned kUseOpenMPDefaultNumThreads = 0;
408 unsigned numThreads;
409
410 ParallelOpLowering(MLIRContext *context,
411 unsigned numThreads = kUseOpenMPDefaultNumThreads)
412 : OpRewritePattern<scf::ParallelOp>(context), numThreads(numThreads) {}
413
414 LogicalResult matchAndRewrite(scf::ParallelOp parallelOp,
415 PatternRewriter &rewriter) const override {
416 if (parallelOp.getUnsignedCmp())
417 return rewriter.notifyMatchFailure(
418 parallelOp, "unsigned loop bounds are not supported");
419 // Bail out early if any reduction init value has a type that is not
420 // compatible with LLVM (e.g. index), since we cannot allocate a reduction
421 // variable for such types.
422 for (Value init : parallelOp.getInitVals()) {
423 if (!LLVM::isCompatibleType(init.getType()) &&
424 !isa<LLVM::PointerElementTypeInterface>(init.getType()))
425 return rewriter.notifyMatchFailure(
426 parallelOp, "reduction init type is not an LLVM-compatible type");
427 }
428
429 // Declare reductions.
430 // TODO: consider checking it here is already a compatible reduction
431 // declaration and use it instead of redeclaring.
432 SmallVector<Attribute> reductionSyms;
433 SmallVector<omp::DeclareReductionOp> ompReductionDecls;
434 auto reduce = cast<scf::ReduceOp>(parallelOp.getBody()->getTerminator());
435 for (int64_t i = 0, e = parallelOp.getNumReductions(); i < e; ++i) {
436 omp::DeclareReductionOp decl = declareReduction(rewriter, reduce, i);
437 ompReductionDecls.push_back(decl);
438 if (!decl)
439 return failure();
440 reductionSyms.push_back(
441 SymbolRefAttr::get(rewriter.getContext(), decl.getSymName()));
442 }
443
444 // Allocate reduction variables. Make sure the we don't overflow the stack
445 // with local `alloca`s by saving and restoring the stack pointer.
446 Location loc = parallelOp.getLoc();
447 Value one =
448 LLVM::ConstantOp::create(rewriter, loc, rewriter.getIntegerType(64),
449 rewriter.getI64IntegerAttr(1));
450 SmallVector<Value> reductionVariables;
451 reductionVariables.reserve(parallelOp.getNumReductions());
452 auto ptrType = LLVM::LLVMPointerType::get(parallelOp.getContext());
453 for (Value init : parallelOp.getInitVals()) {
454 Value storage = LLVM::AllocaOp::create(rewriter, loc, ptrType,
455 init.getType(), one, 0);
456 LLVM::StoreOp::create(rewriter, loc, init, storage);
457 reductionVariables.push_back(storage);
458 }
459
460 // Replace the reduction operations contained in this loop. Must be done
461 // here rather than in a separate pattern to have access to the list of
462 // reduction variables.
463 for (auto [x, y, rD] : llvm::zip_equal(
464 reductionVariables, reduce.getOperands(), ompReductionDecls)) {
465 OpBuilder::InsertionGuard guard(rewriter);
466 rewriter.setInsertionPoint(reduce);
467 Region &redRegion = rD.getReductionRegion();
468 // The SCF dialect by definition contains only structured operations
469 // and hence the SCF reduction region will contain a single block.
470 // The ompReductionDecls region is a copy of the SCF reduction region
471 // and hence has the same property.
472 assert(redRegion.hasOneBlock() &&
473 "expect reduction region to have one block");
474 Value pvtRedVar = parallelOp.getRegion().addArgument(x.getType(), loc);
475 Value pvtRedVal = LLVM::LoadOp::create(rewriter, reduce.getLoc(),
476 rD.getType(), pvtRedVar);
477 // Make a copy of the reduction combiner region in the body
478 mlir::OpBuilder builder(rewriter.getContext());
479 builder.setInsertionPoint(reduce);
480 mlir::IRMapping mapper;
481 assert(redRegion.getNumArguments() == 2 &&
482 "expect reduction region to have two arguments");
483 mapper.map(redRegion.getArgument(0), pvtRedVal);
484 mapper.map(redRegion.getArgument(1), y);
485 for (auto &op : redRegion.getOps()) {
486 Operation *cloneOp = builder.clone(op, mapper);
487 if (auto yieldOp = dyn_cast<omp::YieldOp>(*cloneOp)) {
488 assert(yieldOp && yieldOp.getResults().size() == 1 &&
489 "expect YieldOp in reduction region to return one result");
490 Value redVal = yieldOp.getResults()[0];
491 LLVM::StoreOp::create(rewriter, loc, redVal, pvtRedVar);
492 rewriter.eraseOp(yieldOp);
493 break;
494 }
495 }
496 }
497 rewriter.eraseOp(reduce);
498
499 SmallVector<Value> numThreadsVars;
500 if (numThreads > 0) {
501 Value numThreadsVar = LLVM::ConstantOp::create(
502 rewriter, loc, rewriter.getI32IntegerAttr(numThreads));
503 numThreadsVars.push_back(numThreadsVar);
504 }
505 // Create the parallel wrapper.
506 auto ompParallel = omp::ParallelOp::create(
507 rewriter, loc,
508 /* allocate_vars = */ llvm::SmallVector<Value>{},
509 /* allocator_vars = */ llvm::SmallVector<Value>{},
510 /* allocate_alignments = */ nullptr,
511 /* allocate_private_indices = */ nullptr,
512 /* if_expr = */ Value{},
513 /* num_threads_vars = */ numThreadsVars,
514 /* private_vars = */ ValueRange(),
515 /* private_syms = */ nullptr,
516 /* private_needs_barrier = */ false,
517 /* proc_bind_kind = */ omp::ClauseProcBindKindAttr{},
518 /* reduction_mod = */ nullptr,
519 /* reduction_vars = */ llvm::SmallVector<Value>{},
520 /* reduction_byref = */ DenseBoolArrayAttr{},
521 /* reduction_syms = */ ArrayAttr{});
522 {
523
524 OpBuilder::InsertionGuard guard(rewriter);
525 rewriter.createBlock(&ompParallel.getRegion());
526
527 // Replace the loop.
528 {
529 OpBuilder::InsertionGuard allocaGuard(rewriter);
530 // Create worksharing loop wrapper.
531 auto wsloopOp = omp::WsloopOp::create(rewriter, parallelOp.getLoc());
532 if (!reductionVariables.empty()) {
533 wsloopOp.setReductionSymsAttr(
534 ArrayAttr::get(rewriter.getContext(), reductionSyms));
535 wsloopOp.getReductionVarsMutable().append(reductionVariables);
536 llvm::SmallVector<bool> reductionByRef;
537 // false because these reductions always reduce scalars and so do
538 // not need to pass by reference
539 reductionByRef.resize(reductionVariables.size(), false);
540 wsloopOp.setReductionByref(
541 DenseBoolArrayAttr::get(rewriter.getContext(), reductionByRef));
542 }
543 omp::TerminatorOp::create(rewriter, loc); // omp.parallel terminator.
544
545 // The wrapper's entry block arguments will define the reduction
546 // variables.
547 llvm::SmallVector<mlir::Type> reductionTypes;
548 reductionTypes.reserve(reductionVariables.size());
549 llvm::transform(reductionVariables, std::back_inserter(reductionTypes),
550 [](mlir::Value v) { return v.getType(); });
551 rewriter.createBlock(
552 &wsloopOp.getRegion(), {}, reductionTypes,
553 llvm::SmallVector<mlir::Location>(reductionVariables.size(),
554 parallelOp.getLoc()));
555
556 // Create loop nest and populate region with contents of scf.parallel.
557 auto loopOp = omp::LoopNestOp::create(
558 rewriter, parallelOp.getLoc(), parallelOp.getLowerBound().size(),
559 parallelOp.getLowerBound(), parallelOp.getUpperBound(),
560 parallelOp.getStep(), /*loop_inclusive=*/false,
561 /*tile_sizes=*/nullptr);
562
563 rewriter.inlineRegionBefore(parallelOp.getRegion(), loopOp.getRegion(),
564 loopOp.getRegion().begin());
565
566 // Remove reduction-related block arguments from omp.loop_nest and
567 // redirect uses to the corresponding omp.wsloop block argument.
568 mlir::Block &loopOpEntryBlock = loopOp.getRegion().front();
569 unsigned numLoops = parallelOp.getNumLoops();
570 rewriter.replaceAllUsesWith(
571 loopOpEntryBlock.getArguments().drop_front(numLoops),
572 wsloopOp.getRegion().getArguments());
573 loopOpEntryBlock.eraseArguments(
574 numLoops, loopOpEntryBlock.getNumArguments() - numLoops);
575
576 Block *ops =
577 rewriter.splitBlock(&loopOpEntryBlock, loopOpEntryBlock.begin());
578 rewriter.setInsertionPointToStart(&loopOpEntryBlock);
579
580 auto scope = memref::AllocaScopeOp::create(
581 rewriter, parallelOp.getLoc(), TypeRange());
582 omp::YieldOp::create(rewriter, loc, ValueRange());
583 Block *scopeBlock = rewriter.createBlock(&scope.getBodyRegion());
584 rewriter.mergeBlocks(ops, scopeBlock);
585 rewriter.setInsertionPointToEnd(&*scope.getBodyRegion().begin());
586 memref::AllocaScopeReturnOp::create(rewriter, loc, ValueRange());
587 }
588 }
589
590 // Load loop results.
591 SmallVector<Value> results;
592 results.reserve(reductionVariables.size());
593 for (auto [variable, type] :
594 llvm::zip(reductionVariables, parallelOp.getResultTypes())) {
595 Value res = LLVM::LoadOp::create(rewriter, loc, type, variable);
596 results.push_back(res);
597 }
598 rewriter.replaceOp(parallelOp, results);
599
600 return success();
601 }
602};
603
604/// Applies the conversion patterns in the given function.
605static LogicalResult applyPatterns(ModuleOp module, unsigned numThreads) {
606 RewritePatternSet patterns(module.getContext());
607 patterns.add<ParallelOpLowering>(module.getContext(), numThreads);
608 FrozenRewritePatternSet frozen(std::move(patterns));
609 walkAndApplyPatterns(module, frozen);
610 auto status = module.walk([](Operation *op) {
611 if (isa<scf::ReduceOp, scf::ReduceReturnOp, scf::ParallelOp>(op)) {
612 op->emitError("unconverted operation found");
613 return WalkResult::interrupt();
614 }
615 return WalkResult::advance();
616 });
617 return failure(status.wasInterrupted());
618}
619
620/// A pass converting SCF operations to OpenMP operations.
621struct SCFToOpenMPPass
623
624 using Base::Base;
625
626 /// Pass entry point.
627 void runOnOperation() override {
628 if (failed(applyPatterns(getOperation(), numThreads)))
629 signalPassFailure();
630 }
631};
632
633} // namespace
return success()
static Value reduce(OpBuilder &builder, Location loc, Value input, Value output, int64_t dim)
ArrayAttr()
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
static void applyPatterns(Region &region, const FrozenRewritePatternSet &patterns, ArrayRef< ReductionNode::Range > rangeToKeep, bool eraseOpNotInRange)
We implicitly number each operation in the region and if an operation's number falls into rangeToKeep...
static Attribute minMaxValueForFloat(Type type, bool min)
Returns an attribute with the minimum (if min is set) or the maximum value (otherwise) for the given ...
static omp::DeclareReductionOp addAtomicRMW(OpBuilder &builder, LLVM::AtomicBinOp atomicKind, omp::DeclareReductionOp decl, scf::ReduceOp reduce, int64_t reductionIndex)
Adds an atomic reduction combiner to the given OpenMP reduction declaration using llvm....
static bool matchSimpleReduction(Block &block)
Matches a block containing a "simple" reduction.
static const llvm::fltSemantics & fltSemanticsForType(FloatType type)
Returns the float semantics for the given float type.
static Attribute getSplatOrScalarAttr(Type type, Attribute val)
Helper to create a splat attribute for vector types, or return the scalar attribute for scalar types.
static bool supportsAtomic(Type type)
Returns true if the type is supported by llvm.atomicrmw.
static omp::DeclareReductionOp declareReduction(PatternRewriter &builder, scf::ReduceOp reduce, int64_t reductionIndex)
Creates an OpenMP reduction declaration that corresponds to the given SCF reduction and returns it.
static bool matchSelectReduction(Block &block, ArrayRef< Predicate > lessThanPredicates, ArrayRef< Predicate > greaterThanPredicates, bool &isMin)
Matches a block containing a select-based min/max reduction.
static omp::DeclareReductionOp createDecl(PatternRewriter &builder, SymbolTable &symbolTable, scf::ReduceOp reduce, int64_t reductionIndex, Attribute initValue)
Creates an OpenMP reduction declaration and inserts it into the provided symbol table.
static Attribute minMaxValueForSignedInt(Type type, bool min)
Returns an attribute with the signed integer minimum (if min is set) or the maximum value (otherwise)...
static Attribute minMaxValueForUnsignedInt(Type type, bool min)
Returns an attribute with the unsigned integer minimum (if min is set) or the maximum value (otherwis...
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
bool empty()
Definition Block.h:173
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
Operation & front()
Definition Block.h:178
Operation & back()
Definition Block.h:177
void eraseArguments(unsigned start, unsigned num)
Erases 'num' arguments from the index 'start'.
Definition Block.cpp:206
BlockArgListType getArguments()
Definition Block.h:112
iterator end()
Definition Block.h:169
iterator begin()
Definition Block.h:168
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
MLIRContext * getContext() const
Definition Builders.h:56
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class represents a frozen set of patterns that can be processed by a pattern applicator.
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
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
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.
Definition Builders.cpp:439
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
Definition Builders.h:439
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
Operation * clone(IRMapping &mapper, const CloneOptions &options=CloneOptions::all())
Create a deep copy of this operation, remapping any operands that use values outside of the operation...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
iterator_range< OpIterator > getOps()
Definition Region.h:180
unsigned getNumArguments()
Definition Region.h:136
BlockArgument getArgument(unsigned i)
Definition Region.h:137
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
Block * splitBlock(Block *block, Block::iterator before)
Split the operations starting at "before" (inclusive) out of the given block into a new 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 eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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 inlineRegionBefore(Region &region, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
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 allows for representing and managing the symbol table used by operations with the 'SymbolT...
Definition SymbolTable.h:25
StringAttr insert(Operation *symbol, Block::iterator insertPt={})
Insert a new symbol into the table, and rename it as necessary to avoid collisions.
static Operation * getNearestSymbolTable(Operation *from)
Returns the nearest symbol table from a given operation from.
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< bool > content)
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
Value matchReduction(ArrayRef< BlockArgument > iterCarriedArgs, unsigned redPos, SmallVectorImpl< Operation * > &combinerOps)
Utility to match a generic reduction given a list of iteration-carried arguments, iterCarriedArgs and...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void walkAndApplyPatterns(Operation *op, const FrozenRewritePatternSet &patterns, RewriterBase::Listener *listener=nullptr)
A fast walk-based pattern rewrite driver.
detail::DenseArrayAttrImpl< bool > DenseBoolArrayAttr
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...