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 // Bail out early if any reduction init value has a type that is not
417 // compatible with LLVM (e.g. index), since we cannot allocate a reduction
418 // variable for such types.
419 for (Value init : parallelOp.getInitVals()) {
420 if (!LLVM::isCompatibleType(init.getType()) &&
421 !isa<LLVM::PointerElementTypeInterface>(init.getType()))
422 return rewriter.notifyMatchFailure(
423 parallelOp, "reduction init type is not an LLVM-compatible type");
424 }
425
426 // Declare reductions.
427 // TODO: consider checking it here is already a compatible reduction
428 // declaration and use it instead of redeclaring.
429 SmallVector<Attribute> reductionSyms;
430 SmallVector<omp::DeclareReductionOp> ompReductionDecls;
431 auto reduce = cast<scf::ReduceOp>(parallelOp.getBody()->getTerminator());
432 for (int64_t i = 0, e = parallelOp.getNumReductions(); i < e; ++i) {
433 omp::DeclareReductionOp decl = declareReduction(rewriter, reduce, i);
434 ompReductionDecls.push_back(decl);
435 if (!decl)
436 return failure();
437 reductionSyms.push_back(
438 SymbolRefAttr::get(rewriter.getContext(), decl.getSymName()));
439 }
440
441 // Allocate reduction variables. Make sure the we don't overflow the stack
442 // with local `alloca`s by saving and restoring the stack pointer.
443 Location loc = parallelOp.getLoc();
444 Value one =
445 LLVM::ConstantOp::create(rewriter, loc, rewriter.getIntegerType(64),
446 rewriter.getI64IntegerAttr(1));
447 SmallVector<Value> reductionVariables;
448 reductionVariables.reserve(parallelOp.getNumReductions());
449 auto ptrType = LLVM::LLVMPointerType::get(parallelOp.getContext());
450 for (Value init : parallelOp.getInitVals()) {
451 Value storage = LLVM::AllocaOp::create(rewriter, loc, ptrType,
452 init.getType(), one, 0);
453 LLVM::StoreOp::create(rewriter, loc, init, storage);
454 reductionVariables.push_back(storage);
455 }
456
457 // Replace the reduction operations contained in this loop. Must be done
458 // here rather than in a separate pattern to have access to the list of
459 // reduction variables.
460 for (auto [x, y, rD] : llvm::zip_equal(
461 reductionVariables, reduce.getOperands(), ompReductionDecls)) {
462 OpBuilder::InsertionGuard guard(rewriter);
463 rewriter.setInsertionPoint(reduce);
464 Region &redRegion = rD.getReductionRegion();
465 // The SCF dialect by definition contains only structured operations
466 // and hence the SCF reduction region will contain a single block.
467 // The ompReductionDecls region is a copy of the SCF reduction region
468 // and hence has the same property.
469 assert(redRegion.hasOneBlock() &&
470 "expect reduction region to have one block");
471 Value pvtRedVar = parallelOp.getRegion().addArgument(x.getType(), loc);
472 Value pvtRedVal = LLVM::LoadOp::create(rewriter, reduce.getLoc(),
473 rD.getType(), pvtRedVar);
474 // Make a copy of the reduction combiner region in the body
475 mlir::OpBuilder builder(rewriter.getContext());
476 builder.setInsertionPoint(reduce);
477 mlir::IRMapping mapper;
478 assert(redRegion.getNumArguments() == 2 &&
479 "expect reduction region to have two arguments");
480 mapper.map(redRegion.getArgument(0), pvtRedVal);
481 mapper.map(redRegion.getArgument(1), y);
482 for (auto &op : redRegion.getOps()) {
483 Operation *cloneOp = builder.clone(op, mapper);
484 if (auto yieldOp = dyn_cast<omp::YieldOp>(*cloneOp)) {
485 assert(yieldOp && yieldOp.getResults().size() == 1 &&
486 "expect YieldOp in reduction region to return one result");
487 Value redVal = yieldOp.getResults()[0];
488 LLVM::StoreOp::create(rewriter, loc, redVal, pvtRedVar);
489 rewriter.eraseOp(yieldOp);
490 break;
491 }
492 }
493 }
494 rewriter.eraseOp(reduce);
495
496 SmallVector<Value> numThreadsVars;
497 if (numThreads > 0) {
498 Value numThreadsVar = LLVM::ConstantOp::create(
499 rewriter, loc, rewriter.getI32IntegerAttr(numThreads));
500 numThreadsVars.push_back(numThreadsVar);
501 }
502 // Create the parallel wrapper.
503 auto ompParallel = omp::ParallelOp::create(
504 rewriter, loc,
505 /* allocate_vars = */ llvm::SmallVector<Value>{},
506 /* allocator_vars = */ llvm::SmallVector<Value>{},
507 /* allocate_alignments = */ nullptr,
508 /* allocate_private_indices = */ nullptr,
509 /* if_expr = */ Value{},
510 /* num_threads_vars = */ numThreadsVars,
511 /* private_vars = */ ValueRange(),
512 /* private_syms = */ nullptr,
513 /* private_needs_barrier = */ nullptr,
514 /* proc_bind_kind = */ omp::ClauseProcBindKindAttr{},
515 /* reduction_mod = */ nullptr,
516 /* reduction_vars = */ llvm::SmallVector<Value>{},
517 /* reduction_byref = */ DenseBoolArrayAttr{},
518 /* reduction_syms = */ ArrayAttr{});
519 {
520
521 OpBuilder::InsertionGuard guard(rewriter);
522 rewriter.createBlock(&ompParallel.getRegion());
523
524 // Replace the loop.
525 {
526 OpBuilder::InsertionGuard allocaGuard(rewriter);
527 // Create worksharing loop wrapper.
528 auto wsloopOp = omp::WsloopOp::create(rewriter, parallelOp.getLoc());
529 if (!reductionVariables.empty()) {
530 wsloopOp.setReductionSymsAttr(
531 ArrayAttr::get(rewriter.getContext(), reductionSyms));
532 wsloopOp.getReductionVarsMutable().append(reductionVariables);
533 llvm::SmallVector<bool> reductionByRef;
534 // false because these reductions always reduce scalars and so do
535 // not need to pass by reference
536 reductionByRef.resize(reductionVariables.size(), false);
537 wsloopOp.setReductionByref(
538 DenseBoolArrayAttr::get(rewriter.getContext(), reductionByRef));
539 }
540 omp::TerminatorOp::create(rewriter, loc); // omp.parallel terminator.
541
542 // The wrapper's entry block arguments will define the reduction
543 // variables.
544 llvm::SmallVector<mlir::Type> reductionTypes;
545 reductionTypes.reserve(reductionVariables.size());
546 llvm::transform(reductionVariables, std::back_inserter(reductionTypes),
547 [](mlir::Value v) { return v.getType(); });
548 rewriter.createBlock(
549 &wsloopOp.getRegion(), {}, reductionTypes,
550 llvm::SmallVector<mlir::Location>(reductionVariables.size(),
551 parallelOp.getLoc()));
552
553 // Create loop nest and populate region with contents of scf.parallel.
554 auto loopOp = omp::LoopNestOp::create(
555 rewriter, parallelOp.getLoc(), parallelOp.getLowerBound().size(),
556 parallelOp.getLowerBound(), parallelOp.getUpperBound(),
557 parallelOp.getStep(), /*loop_inclusive=*/false,
558 /*tile_sizes=*/nullptr);
559
560 rewriter.inlineRegionBefore(parallelOp.getRegion(), loopOp.getRegion(),
561 loopOp.getRegion().begin());
562
563 // Remove reduction-related block arguments from omp.loop_nest and
564 // redirect uses to the corresponding omp.wsloop block argument.
565 mlir::Block &loopOpEntryBlock = loopOp.getRegion().front();
566 unsigned numLoops = parallelOp.getNumLoops();
567 rewriter.replaceAllUsesWith(
568 loopOpEntryBlock.getArguments().drop_front(numLoops),
569 wsloopOp.getRegion().getArguments());
570 loopOpEntryBlock.eraseArguments(
571 numLoops, loopOpEntryBlock.getNumArguments() - numLoops);
572
573 Block *ops =
574 rewriter.splitBlock(&loopOpEntryBlock, loopOpEntryBlock.begin());
575 rewriter.setInsertionPointToStart(&loopOpEntryBlock);
576
577 auto scope = memref::AllocaScopeOp::create(
578 rewriter, parallelOp.getLoc(), TypeRange());
579 omp::YieldOp::create(rewriter, loc, ValueRange());
580 Block *scopeBlock = rewriter.createBlock(&scope.getBodyRegion());
581 rewriter.mergeBlocks(ops, scopeBlock);
582 rewriter.setInsertionPointToEnd(&*scope.getBodyRegion().begin());
583 memref::AllocaScopeReturnOp::create(rewriter, loc, ValueRange());
584 }
585 }
586
587 // Load loop results.
588 SmallVector<Value> results;
589 results.reserve(reductionVariables.size());
590 for (auto [variable, type] :
591 llvm::zip(reductionVariables, parallelOp.getResultTypes())) {
592 Value res = LLVM::LoadOp::create(rewriter, loc, type, variable);
593 results.push_back(res);
594 }
595 rewriter.replaceOp(parallelOp, results);
596
597 return success();
598 }
599};
600
601/// Applies the conversion patterns in the given function.
602static LogicalResult applyPatterns(ModuleOp module, unsigned numThreads) {
603 RewritePatternSet patterns(module.getContext());
604 patterns.add<ParallelOpLowering>(module.getContext(), numThreads);
605 FrozenRewritePatternSet frozen(std::move(patterns));
606 walkAndApplyPatterns(module, frozen);
607 auto status = module.walk([](Operation *op) {
608 if (isa<scf::ReduceOp, scf::ReduceReturnOp, scf::ParallelOp>(op)) {
609 op->emitError("unconverted operation found");
610 return WalkResult::interrupt();
611 }
612 return WalkResult::advance();
613 });
614 return failure(status.wasInterrupted());
615}
616
617/// A pass converting SCF operations to OpenMP operations.
618struct SCFToOpenMPPass
619 : public impl::ConvertSCFToOpenMPPassBase<SCFToOpenMPPass> {
620
621 using Base::Base;
622
623 /// Pass entry point.
624 void runOnOperation() override {
625 if (failed(applyPatterns(getOperation(), numThreads)))
626 signalPassFailure();
627 }
628};
629
630} // 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:33
bool empty()
Definition Block.h:172
BlockArgument getArgument(unsigned i)
Definition Block.h:153
unsigned getNumArguments()
Definition Block.h:152
Operation & front()
Definition Block.h:177
Operation & back()
Definition Block.h:176
void eraseArguments(unsigned start, unsigned num)
Erases 'num' arguments from the index 'start'.
Definition Block.cpp:206
BlockArgListType getArguments()
Definition Block.h:111
iterator end()
Definition Block.h:168
iterator begin()
Definition Block.h:167
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:185
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:24
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:732
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...