178template <
typename ValueRange>
181 for (
unsigned i = 0; i < values.size(); ++i)
191 assert(!tripCounts.empty() &&
"tripCounts must be not empty");
193 for (ssize_t i = tripCounts.size() - 1; i >= 0; --i) {
194 coords[i] = arith::RemSIOp::create(
b,
index, tripCounts[i]);
195 index = arith::DivSIOp::create(
b,
index, tripCounts[i]);
204static ParallelComputeFunctionType
211 inputs.reserve(2 + 4 * op.getNumLoops() + captures.size());
216 inputs.push_back(indexTy);
217 inputs.push_back(indexTy);
220 for (
unsigned i = 0; i < op.getNumLoops(); ++i)
221 inputs.push_back(indexTy);
226 for (
unsigned i = 0; i < op.getNumLoops(); ++i) {
227 inputs.push_back(indexTy);
228 inputs.push_back(indexTy);
229 inputs.push_back(indexTy);
233 for (
Value capture : captures)
234 inputs.push_back(capture.getType());
243 scf::ParallelOp op,
const ParallelComputeFunctionBounds &bounds,
248 ModuleOp module = op->getParentOfType<ModuleOp>();
250 ParallelComputeFunctionType computeFuncType =
253 FunctionType type = computeFuncType.type;
254 func::FuncOp
func = func::FuncOp::create(
256 numBlockAlignedInnerLoops > 0 ?
"parallel_compute_fn_with_aligned_loops"
257 :
"parallel_compute_fn",
268 b.createBlock(&
func.getBody(),
func.begin(), type.getInputs(),
270 b.setInsertionPointToEnd(block);
272 ParallelComputeFunctionArgs args = {op.getNumLoops(),
func.getArguments()};
284 return llvm::map_to_vector(llvm::zip(args, attrs),
285 [&](
auto tuple) ->
Value {
286 if (IntegerAttr attr = std::get<1>(tuple))
287 return arith::ConstantOp::create(
b, attr);
288 return std::get<0>(tuple);
293 auto tripCounts = values(args.tripCounts(), bounds.tripCounts);
296 auto lowerBounds = values(args.lowerBounds(), bounds.lowerBounds);
297 auto steps = values(args.steps(), bounds.steps);
304 Value tripCount = tripCounts[0];
305 for (
unsigned i = 1; i < tripCounts.size(); ++i)
306 tripCount = arith::MulIOp::create(
b, tripCount, tripCounts[i]);
310 Value blockFirstIndex = arith::MulIOp::create(
b, blockIndex, blockSize);
314 Value blockEnd0 = arith::AddIOp::create(
b, blockFirstIndex, blockSize);
315 Value blockEnd1 = arith::MinSIOp::create(
b, blockEnd0, tripCount);
316 Value blockLastIndex = arith::SubIOp::create(
b, blockEnd1, c1);
319 auto blockFirstCoord =
delinearize(
b, blockFirstIndex, tripCounts);
320 auto blockLastCoord =
delinearize(
b, blockLastIndex, tripCounts);
328 for (
size_t i = 0; i < blockLastCoord.size(); ++i)
329 blockEndCoord[i] = arith::AddIOp::create(
b, blockLastCoord[i], c1);
333 using LoopBodyBuilder =
335 using LoopNestBuilder = std::function<LoopBodyBuilder(
size_t loopIdx)>;
366 LoopNestBuilder workLoopBuilder = [&](
size_t loopIdx) -> LoopBodyBuilder {
372 computeBlockInductionVars[loopIdx] =
373 arith::AddIOp::create(
b, lowerBounds[loopIdx],
374 arith::MulIOp::create(
b, iv, steps[loopIdx]));
377 isBlockFirstCoord[loopIdx] = arith::CmpIOp::create(
378 b, arith::CmpIPredicate::eq, iv, blockFirstCoord[loopIdx]);
379 isBlockLastCoord[loopIdx] = arith::CmpIOp::create(
380 b, arith::CmpIPredicate::eq, iv, blockLastCoord[loopIdx]);
384 isBlockFirstCoord[loopIdx] = arith::AndIOp::create(
385 b, isBlockFirstCoord[loopIdx], isBlockFirstCoord[loopIdx - 1]);
386 isBlockLastCoord[loopIdx] = arith::AndIOp::create(
387 b, isBlockLastCoord[loopIdx], isBlockLastCoord[loopIdx - 1]);
391 if (loopIdx < op.getNumLoops() - 1) {
392 if (loopIdx + 1 >= op.getNumLoops() - numBlockAlignedInnerLoops) {
395 scf::ForOp::create(
b, c0, tripCounts[loopIdx + 1], c1,
ValueRange(),
396 workLoopBuilder(loopIdx + 1));
401 auto lb = arith::SelectOp::create(
b, isBlockFirstCoord[loopIdx],
402 blockFirstCoord[loopIdx + 1], c0);
404 auto ub = arith::SelectOp::create(
b, isBlockLastCoord[loopIdx],
405 blockEndCoord[loopIdx + 1],
406 tripCounts[loopIdx + 1]);
409 workLoopBuilder(loopIdx + 1));
412 scf::YieldOp::create(
b, loc);
418 mapping.
map(op.getInductionVars(), computeBlockInductionVars);
419 mapping.
map(computeFuncType.captures, captures);
421 for (
auto &bodyOp : op.getRegion().front().without_terminator())
422 b.clone(bodyOp, mapping);
423 scf::YieldOp::create(
b, loc);
427 scf::ForOp::create(
b, blockFirstCoord[0], blockEndCoord[0], c1,
ValueRange(),
431 return {op.getNumLoops(),
func, std::move(computeFuncType.captures)};
457 Location loc = computeFunc.func.getLoc();
460 ModuleOp module = computeFunc.func->getParentOfType<ModuleOp>();
463 computeFunc.func.getFunctionType().getInputs();
470 inputTypes.push_back(async::GroupType::get(rewriter.
getContext()));
472 inputTypes.append(computeFuncInputTypes.begin(), computeFuncInputTypes.end());
475 func::FuncOp
func = func::FuncOp::create(loc,
"async_dispatch_fn", type);
484 Block *block =
b.createBlock(&
func.getBody(),
func.begin(), type.getInputs(),
486 b.setInsertionPointToEnd(block);
488 Type indexTy =
b.getIndexType();
505 scf::WhileOp whileOp = scf::WhileOp::create(
b, types, operands);
506 Block *before =
b.createBlock(&whileOp.getBefore(), {}, types, locations);
507 Block *after =
b.createBlock(&whileOp.getAfter(), {}, types, locations);
512 b.setInsertionPointToEnd(before);
515 Value distance = arith::SubIOp::create(
b, end, start);
517 arith::CmpIOp::create(
b, arith::CmpIPredicate::sgt, distance, c1);
518 scf::ConditionOp::create(
b, dispatch, before->
getArguments());
524 b.setInsertionPointToEnd(after);
527 Value distance = arith::SubIOp::create(
b, end, start);
528 Value halfDistance = arith::DivSIOp::create(
b, distance, c2);
529 Value midIndex = arith::AddIOp::create(
b, start, halfDistance);
532 auto executeBodyBuilder = [&](
OpBuilder &executeBuilder,
537 operands[1] = midIndex;
540 func::CallOp::create(executeBuilder, executeLoc,
func.getSymName(),
541 func.getResultTypes(), operands);
542 async::YieldOp::create(executeBuilder, executeLoc,
ValueRange());
548 AddToGroupOp::create(
b, indexTy, execute.getToken(), group);
549 scf::YieldOp::create(
b,
ValueRange({start, midIndex}));
554 b.setInsertionPointAfter(whileOp);
557 auto forwardedInputs = block->
getArguments().drop_front(3);
559 computeFuncOperands.append(forwardedInputs.begin(), forwardedInputs.end());
561 func::CallOp::create(
b, computeFunc.func.getSymName(),
562 computeFunc.func.getResultTypes(), computeFuncOperands);
570 ParallelComputeFunction ¶llelComputeFunction,
571 scf::ParallelOp op,
Value blockSize,
578 func::FuncOp asyncDispatchFunction =
587 operands.append(tripCounts);
588 operands.append(op.getLowerBound().begin(), op.getLowerBound().end());
589 operands.append(op.getUpperBound().begin(), op.getUpperBound().end());
590 operands.append(op.getStep().begin(), op.getStep().end());
591 operands.append(parallelComputeFunction.captures);
597 Value isSingleBlock =
598 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, blockCount, c1);
605 appendBlockComputeOperands(operands);
607 func::CallOp::create(
b, parallelComputeFunction.func.getSymName(),
608 parallelComputeFunction.func.getResultTypes(),
610 scf::YieldOp::create(
b);
619 Value groupSize = arith::SubIOp::create(
b, blockCount, c1);
620 Value group = CreateGroupOp::create(
b, GroupType::get(ctx), groupSize);
624 appendBlockComputeOperands(operands);
626 func::CallOp::create(
b, asyncDispatchFunction.getSymName(),
627 asyncDispatchFunction.getResultTypes(), operands);
630 AwaitAllOp::create(
b, group);
632 scf::YieldOp::create(
b);
636 scf::IfOp::create(
b, isSingleBlock, syncDispatch, asyncDispatch);
643 ParallelComputeFunction ¶llelComputeFunction,
644 scf::ParallelOp op,
Value blockSize,
Value blockCount,
648 func::FuncOp compute = parallelComputeFunction.func;
656 Value groupSize = arith::SubIOp::create(
b, blockCount, c1);
657 Value group = CreateGroupOp::create(
b, GroupType::get(ctx), groupSize);
660 using LoopBodyBuilder =
666 computeFuncOperands.append(tripCounts);
667 computeFuncOperands.append(op.getLowerBound().begin(),
668 op.getLowerBound().end());
669 computeFuncOperands.append(op.getUpperBound().begin(),
670 op.getUpperBound().end());
671 computeFuncOperands.append(op.getStep().begin(), op.getStep().end());
672 computeFuncOperands.append(parallelComputeFunction.captures);
673 return computeFuncOperands;
682 auto executeBodyBuilder = [&](
OpBuilder &executeBuilder,
684 func::CallOp::create(executeBuilder, executeLoc, compute.getSymName(),
685 compute.getResultTypes(), computeFuncOperands(iv));
686 async::YieldOp::create(executeBuilder, executeLoc,
ValueRange());
692 AddToGroupOp::create(
b, rewriter.
getIndexType(), execute.getToken(), group);
693 scf::YieldOp::create(
b);
697 scf::ForOp::create(
b, c1, blockCount, c1,
ValueRange(), loopBuilder);
700 func::CallOp::create(
b, compute.getSymName(), compute.getResultTypes(),
701 computeFuncOperands(c0));
704 AwaitAllOp::create(
b, group);
708AsyncParallelForRewrite::matchAndRewrite(scf::ParallelOp op,
709 PatternRewriter &rewriter)
const {
711 if (op.getNumReductions() != 0)
714 if (op.getUnsignedCmp())
716 op,
"unsigned loop bounds are not supported");
718 ImplicitLocOpBuilder
b(op.getLoc(), rewriter);
723 Value minTaskSize = computeMinTaskSize(
b, op);
731 SmallVector<Value> tripCounts(op.getNumLoops());
732 for (
size_t i = 0; i < op.getNumLoops(); ++i) {
733 auto lb = op.getLowerBound()[i];
734 auto ub = op.getUpperBound()[i];
735 auto step = op.getStep()[i];
736 auto range =
b.createOrFold<arith::SubIOp>(ub, lb);
737 tripCounts[i] =
b.createOrFold<arith::CeilDivSIOp>(range, step);
742 Value tripCount = tripCounts[0];
743 for (
size_t i = 1; i < tripCounts.size(); ++i)
744 tripCount = arith::MulIOp::create(
b, tripCount, tripCounts[i]);
749 Value isZeroIterations =
750 arith::CmpIOp::create(
b, arith::CmpIPredicate::eq, tripCount, c0);
753 auto noOp = [&](OpBuilder &nestedBuilder, Location loc) {
754 scf::YieldOp::create(nestedBuilder, loc);
759 auto dispatch = [&](OpBuilder &nestedBuilder, Location loc) {
760 ImplicitLocOpBuilder
b(loc, nestedBuilder);
767 ParallelComputeFunctionBounds staticBounds = {
781 static constexpr int64_t maxUnrollableIterations = 512;
785 int numUnrollableLoops = 0;
787 auto getInt = [](IntegerAttr attr) {
return attr ? attr.getInt() : 0; };
789 SmallVector<int64_t> numIterations(op.getNumLoops());
790 numIterations.back() = getInt(staticBounds.tripCounts.back());
792 for (
int i = op.getNumLoops() - 2; i >= 0; --i) {
793 int64_t tripCount = getInt(staticBounds.tripCounts[i]);
794 int64_t innerIterations = numIterations[i + 1];
795 numIterations[i] = tripCount * innerIterations;
798 if (innerIterations > 0 && innerIterations <= maxUnrollableIterations)
799 numUnrollableLoops++;
802 Value numWorkerThreadsVal;
803 if (numWorkerThreads >= 0)
806 numWorkerThreadsVal = async::RuntimeNumWorkerThreadsOp::create(
b);
821 const SmallVector<std::pair<int, float>> overshardingBrackets = {
822 {4, 4.0f}, {8, 2.0f}, {16, 1.0f}, {32, 0.8f}, {64, 0.6f}};
823 const float initialOvershardingFactor = 8.0f;
826 b,
b.getF32Type(), llvm::APFloat(initialOvershardingFactor));
827 for (
const std::pair<int, float> &p : overshardingBrackets) {
829 Value inBracket = arith::CmpIOp::create(
830 b, arith::CmpIPredicate::sgt, numWorkerThreadsVal, bracketBegin);
832 b,
b.getF32Type(), llvm::APFloat(p.second));
833 scalingFactor = arith::SelectOp::create(
834 b, inBracket, bracketScalingFactor, scalingFactor);
836 Value numWorkersIndex =
837 arith::IndexCastOp::create(
b,
b.getI32Type(), numWorkerThreadsVal);
838 Value numWorkersFloat =
839 arith::SIToFPOp::create(
b,
b.getF32Type(), numWorkersIndex);
840 Value scaledNumWorkers =
841 arith::MulFOp::create(
b, scalingFactor, numWorkersFloat);
843 arith::FPToSIOp::create(
b,
b.getI32Type(), scaledNumWorkers);
844 Value scaledWorkers =
845 arith::IndexCastOp::create(
b,
b.getIndexType(), scaledNumInt);
847 Value maxComputeBlocks = arith::MaxSIOp::create(
854 Value bs0 = arith::CeilDivSIOp::create(
b, tripCount, maxComputeBlocks);
855 Value bs1 = arith::MaxSIOp::create(
b, bs0, minTaskSize);
856 Value blockSize = arith::MinSIOp::create(
b, tripCount, bs1);
866 Value blockCount = arith::CeilDivSIOp::create(
b, tripCount, blockSize);
869 auto dispatchDefault = [&](OpBuilder &nestedBuilder, Location loc) {
870 ParallelComputeFunction compute =
873 ImplicitLocOpBuilder
b(loc, nestedBuilder);
874 doDispatch(
b, rewriter, compute, op, blockSize, blockCount, tripCounts);
875 scf::YieldOp::create(
b);
879 auto dispatchBlockAligned = [&](OpBuilder &nestedBuilder, Location loc) {
881 op, staticBounds, numUnrollableLoops, rewriter);
883 ImplicitLocOpBuilder
b(loc, nestedBuilder);
887 b, numIterations[op.getNumLoops() - numUnrollableLoops]);
888 Value alignedBlockSize = arith::MulIOp::create(
889 b, arith::CeilDivSIOp::create(
b, blockSize, numIters), numIters);
890 doDispatch(
b, rewriter, compute, op, alignedBlockSize, blockCount,
892 scf::YieldOp::create(
b);
898 if (numUnrollableLoops > 0) {
900 b, numIterations[op.getNumLoops() - numUnrollableLoops]);
901 Value useBlockAlignedComputeFn = arith::CmpIOp::create(
902 b, arith::CmpIPredicate::sge, blockSize, numIters);
904 scf::IfOp::create(
b, useBlockAlignedComputeFn, dispatchBlockAligned,
906 scf::YieldOp::create(
b);
908 dispatchDefault(
b, loc);
913 scf::IfOp::create(
b, isZeroIterations, noOp, dispatch);
921void AsyncParallelForPass::runOnOperation() {
924 RewritePatternSet patterns(ctx);
926 patterns, asyncDispatch, numWorkerThreads,
927 [&](ImplicitLocOpBuilder builder, scf::ParallelOp op) {
938 patterns.
add<AsyncParallelForRewrite>(ctx, asyncDispatch, numWorkerThreads,