31template <
typename Op,
typename... Args>
32static Op getOrDefineGlobal(ModuleOp &moduleOp,
const Location loc,
33 ConversionPatternRewriter &rewriter, StringRef name,
36 if (!(ret = moduleOp.lookupSymbol<
Op>(name))) {
37 ConversionPatternRewriter::InsertionGuard guard(rewriter);
38 rewriter.setInsertionPointToStart(moduleOp.getBody());
39 ret = Op::create(rewriter, loc, std::forward<Args>(args)...);
46 ConversionPatternRewriter &rewriter,
48 LLVM::LLVMFunctionType type) {
49 return getOrDefineGlobal<LLVM::LLVMFuncOp>(
50 moduleOp, loc, rewriter, name, name, type, LLVM::Linkage::External);
53std::pair<Value, Value> getRawPtrAndSize(
const Location loc,
54 ConversionPatternRewriter &rewriter,
57 Type ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
58 Type i32Type = rewriter.getI32Type();
59 auto descriptorType = cast<LLVM::LLVMStructType>(memRef.
getType());
62 auto indexType = cast<IntegerType>(descriptorType.getBody()[2]);
65 LLVM::ExtractValueOp::create(rewriter, loc, ptrType, memRef, 1);
67 LLVM::ExtractValueOp::create(rewriter, loc, indexType, memRef, 2);
69 LLVM::GEPOp::create(rewriter, loc, ptrType, elType, dataPtr, offset);
70 Value size = LLVM::ConstantOp::create(rewriter, loc, i32Type,
71 rewriter.getI32IntegerAttr(1));
72 if (descriptorType.getBody().size() > 3) {
73 for (
int64_t i = 0; i < rank; ++i) {
74 Value dim = LLVM::ExtractValueOp::create(rewriter, loc, memRef,
79 if (indexType.getWidth() > 32)
80 dim = LLVM::TruncOp::create(rewriter, loc, i32Type, dim);
81 else if (indexType.getWidth() < 32)
82 dim = LLVM::ZExtOp::create(rewriter, loc, i32Type, dim);
83 size = LLVM::MulOp::create(rewriter, loc, i32Type, dim, size);
86 return {resPtr, size};
100 static std::unique_ptr<MPIImplTraits>
get(ModuleOp &moduleOp);
102 explicit MPIImplTraits(ModuleOp &moduleOp) : moduleOp(moduleOp) {}
104 virtual ~MPIImplTraits() =
default;
106 ModuleOp &getModuleOp() {
return moduleOp; }
112 virtual Value getCommWorld(Location loc,
113 ConversionPatternRewriter &rewriter) = 0;
117 virtual Value castComm(Location loc, ConversionPatternRewriter &rewriter,
121 virtual intptr_t getStatusIgnore() = 0;
124 virtual void *getInPlace() = 0;
128 virtual Value getDataType(Location loc, ConversionPatternRewriter &rewriter,
133 virtual Value getMPIOp(Location loc, ConversionPatternRewriter &rewriter,
134 mpi::MPI_ReductionOpEnum opAttr) = 0;
141class MPICHImplTraits :
public MPIImplTraits {
142 static constexpr int MPI_FLOAT = 0x4c00040a;
143 static constexpr int MPI_DOUBLE = 0x4c00080b;
144 static constexpr int MPI_INT8_T = 0x4c000137;
145 static constexpr int MPI_INT16_T = 0x4c000238;
146 static constexpr int MPI_INT32_T = 0x4c000439;
147 static constexpr int MPI_INT64_T = 0x4c00083a;
148 static constexpr int MPI_UINT8_T = 0x4c00013b;
149 static constexpr int MPI_UINT16_T = 0x4c00023c;
150 static constexpr int MPI_UINT32_T = 0x4c00043d;
151 static constexpr int MPI_UINT64_T = 0x4c00083e;
152 static constexpr int MPI_MAX = 0x58000001;
153 static constexpr int MPI_MIN = 0x58000002;
154 static constexpr int MPI_SUM = 0x58000003;
155 static constexpr int MPI_PROD = 0x58000004;
156 static constexpr int MPI_LAND = 0x58000005;
157 static constexpr int MPI_BAND = 0x58000006;
158 static constexpr int MPI_LOR = 0x58000007;
159 static constexpr int MPI_BOR = 0x58000008;
160 static constexpr int MPI_LXOR = 0x58000009;
161 static constexpr int MPI_BXOR = 0x5800000a;
162 static constexpr int MPI_MINLOC = 0x5800000b;
163 static constexpr int MPI_MAXLOC = 0x5800000c;
164 static constexpr int MPI_REPLACE = 0x5800000d;
165 static constexpr int MPI_NO_OP = 0x5800000e;
168 using MPIImplTraits::MPIImplTraits;
170 ~MPICHImplTraits()
override =
default;
172 Value getCommWorld(
const Location loc,
173 ConversionPatternRewriter &rewriter)
override {
174 static constexpr int MPI_COMM_WORLD = 0x44000000;
175 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(),
179 Value castComm(
const Location loc, ConversionPatternRewriter &rewriter,
180 Value comm)
override {
181 return LLVM::TruncOp::create(rewriter, loc, rewriter.getI32Type(), comm);
184 intptr_t getStatusIgnore()
override {
return 1; }
186 void *getInPlace()
override {
return reinterpret_cast<void *
>(-1); }
188 Value getDataType(
const Location loc, ConversionPatternRewriter &rewriter,
189 Type type)
override {
193 else if (type.
isF64())
198 mtype = MPI_UINT64_T;
202 mtype = MPI_UINT32_T;
206 mtype = MPI_UINT16_T;
212 assert(
false &&
"unsupported type");
213 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(),
217 Value getMPIOp(
const Location loc, ConversionPatternRewriter &rewriter,
218 mpi::MPI_ReductionOpEnum opAttr)
override {
219 int32_t op = MPI_NO_OP;
221 case mpi::MPI_ReductionOpEnum::MPI_OP_NULL:
224 case mpi::MPI_ReductionOpEnum::MPI_MAX:
227 case mpi::MPI_ReductionOpEnum::MPI_MIN:
230 case mpi::MPI_ReductionOpEnum::MPI_SUM:
233 case mpi::MPI_ReductionOpEnum::MPI_PROD:
236 case mpi::MPI_ReductionOpEnum::MPI_LAND:
239 case mpi::MPI_ReductionOpEnum::MPI_BAND:
242 case mpi::MPI_ReductionOpEnum::MPI_LOR:
245 case mpi::MPI_ReductionOpEnum::MPI_BOR:
248 case mpi::MPI_ReductionOpEnum::MPI_LXOR:
251 case mpi::MPI_ReductionOpEnum::MPI_BXOR:
254 case mpi::MPI_ReductionOpEnum::MPI_MINLOC:
257 case mpi::MPI_ReductionOpEnum::MPI_MAXLOC:
260 case mpi::MPI_ReductionOpEnum::MPI_REPLACE:
264 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), op);
271class OMPIImplTraits :
public MPIImplTraits {
272 LLVM::GlobalOp getOrDefineExternalStruct(
const Location loc,
273 ConversionPatternRewriter &rewriter,
275 LLVM::LLVMStructType type) {
277 return getOrDefineGlobal<LLVM::GlobalOp>(
278 getModuleOp(), loc, rewriter, name, type,
false,
279 LLVM::Linkage::External, name,
284 using MPIImplTraits::MPIImplTraits;
286 ~OMPIImplTraits()
override =
default;
288 Value getCommWorld(
const Location loc,
289 ConversionPatternRewriter &rewriter)
override {
290 auto *context = rewriter.getContext();
293 LLVM::LLVMStructType::getOpaque(
"ompi_communicator_t", context);
294 StringRef name =
"ompi_mpi_comm_world";
297 getOrDefineExternalStruct(loc, rewriter, name, commStructT);
300 auto comm = LLVM::AddressOfOp::create(rewriter, loc,
301 LLVM::LLVMPointerType::get(context),
302 SymbolRefAttr::get(context, name));
303 return LLVM::PtrToIntOp::create(rewriter, loc, rewriter.getI64Type(), comm);
306 Value castComm(
const Location loc, ConversionPatternRewriter &rewriter,
307 Value comm)
override {
308 return LLVM::IntToPtrOp::create(
309 rewriter, loc, LLVM::LLVMPointerType::get(rewriter.getContext()), comm);
312 intptr_t getStatusIgnore()
override {
return 0; }
314 void *getInPlace()
override {
return reinterpret_cast<void *
>(1); }
316 Value getDataType(
const Location loc, ConversionPatternRewriter &rewriter,
317 Type type)
override {
320 mtype =
"ompi_mpi_float";
321 else if (type.
isF64())
322 mtype =
"ompi_mpi_double";
324 mtype =
"ompi_mpi_int64_t";
326 mtype =
"ompi_mpi_uint64_t";
328 mtype =
"ompi_mpi_int32_t";
330 mtype =
"ompi_mpi_uint32_t";
332 mtype =
"ompi_mpi_int16_t";
334 mtype =
"ompi_mpi_uint16_t";
336 mtype =
"ompi_mpi_int8_t";
338 mtype =
"ompi_mpi_uint8_t";
340 assert(
false &&
"unsupported type");
342 auto *context = rewriter.getContext();
345 LLVM::LLVMStructType::getOpaque(
"ompi_predefined_datatype_t", context);
347 getOrDefineExternalStruct(loc, rewriter, mtype, typeStructT);
349 return LLVM::AddressOfOp::create(rewriter, loc,
350 LLVM::LLVMPointerType::get(context),
351 SymbolRefAttr::get(context, mtype));
354 Value getMPIOp(
const Location loc, ConversionPatternRewriter &rewriter,
355 mpi::MPI_ReductionOpEnum opAttr)
override {
358 case mpi::MPI_ReductionOpEnum::MPI_OP_NULL:
359 op =
"ompi_mpi_no_op";
361 case mpi::MPI_ReductionOpEnum::MPI_MAX:
364 case mpi::MPI_ReductionOpEnum::MPI_MIN:
367 case mpi::MPI_ReductionOpEnum::MPI_SUM:
370 case mpi::MPI_ReductionOpEnum::MPI_PROD:
371 op =
"ompi_mpi_prod";
373 case mpi::MPI_ReductionOpEnum::MPI_LAND:
374 op =
"ompi_mpi_land";
376 case mpi::MPI_ReductionOpEnum::MPI_BAND:
377 op =
"ompi_mpi_band";
379 case mpi::MPI_ReductionOpEnum::MPI_LOR:
382 case mpi::MPI_ReductionOpEnum::MPI_BOR:
385 case mpi::MPI_ReductionOpEnum::MPI_LXOR:
386 op =
"ompi_mpi_lxor";
388 case mpi::MPI_ReductionOpEnum::MPI_BXOR:
389 op =
"ompi_mpi_bxor";
391 case mpi::MPI_ReductionOpEnum::MPI_MINLOC:
392 op =
"ompi_mpi_minloc";
394 case mpi::MPI_ReductionOpEnum::MPI_MAXLOC:
395 op =
"ompi_mpi_maxloc";
397 case mpi::MPI_ReductionOpEnum::MPI_REPLACE:
398 op =
"ompi_mpi_replace";
401 auto *context = rewriter.getContext();
404 LLVM::LLVMStructType::getOpaque(
"ompi_predefined_op_t", context);
406 getOrDefineExternalStruct(loc, rewriter, op, opStructT);
408 return LLVM::AddressOfOp::create(rewriter, loc,
409 LLVM::LLVMPointerType::get(context),
410 SymbolRefAttr::get(context, op));
414std::unique_ptr<MPIImplTraits> MPIImplTraits::get(ModuleOp &moduleOp) {
415 auto attr =
dlti::query(moduleOp, {
"MPI:Implementation"},
false);
417 return std::make_unique<MPICHImplTraits>(moduleOp);
418 auto strAttr = dyn_cast<StringAttr>(attr.value());
419 if (strAttr && strAttr.getValue() ==
"OpenMPI")
420 return std::make_unique<OMPIImplTraits>(moduleOp);
421 if (!strAttr || strAttr.getValue() !=
"MPICH")
422 moduleOp.emitWarning() <<
"Unknown \"MPI:Implementation\" value in DLTI ("
423 << (strAttr ? strAttr.getValue() :
"<NULL>")
424 <<
"), defaulting to MPICH";
425 return std::make_unique<MPICHImplTraits>(moduleOp);
432struct InitOpLowering :
public ConvertOpToLLVMPattern<mpi::InitOp> {
436 matchAndRewrite(mpi::InitOp op, OpAdaptor adaptor,
437 ConversionPatternRewriter &rewriter)
const override {
438 Location loc = op.getLoc();
441 Type ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
444 auto nullPtrOp = LLVM::ZeroOp::create(rewriter, loc, ptrType);
445 Value llvmnull = nullPtrOp.getRes();
448 auto moduleOp = op->getParentOfType<ModuleOp>();
452 LLVM::LLVMFunctionType::get(rewriter.getI32Type(), {ptrType, ptrType});
454 LLVM::LLVMFuncOp initDecl =
458 rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, initDecl,
469struct FinalizeOpLowering :
public ConvertOpToLLVMPattern<mpi::FinalizeOp> {
473 matchAndRewrite(mpi::FinalizeOp op, OpAdaptor adaptor,
474 ConversionPatternRewriter &rewriter)
const override {
476 Location loc = op.getLoc();
479 auto moduleOp = op->getParentOfType<ModuleOp>();
482 auto initFuncType = LLVM::LLVMFunctionType::get(rewriter.getI32Type(), {});
485 moduleOp, loc, rewriter,
"MPI_Finalize", initFuncType);
488 rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, initDecl,
ValueRange{});
498struct CommWorldOpLowering :
public ConvertOpToLLVMPattern<mpi::CommWorldOp> {
502 matchAndRewrite(mpi::CommWorldOp op, OpAdaptor adaptor,
503 ConversionPatternRewriter &rewriter)
const override {
505 auto moduleOp = op->getParentOfType<ModuleOp>();
506 auto mpiTraits = MPIImplTraits::get(moduleOp);
508 rewriter.replaceOp(op, mpiTraits->getCommWorld(op.getLoc(), rewriter));
518struct CommSplitOpLowering :
public ConvertOpToLLVMPattern<mpi::CommSplitOp> {
522 matchAndRewrite(mpi::CommSplitOp op, OpAdaptor adaptor,
523 ConversionPatternRewriter &rewriter)
const override {
525 auto moduleOp = op->getParentOfType<ModuleOp>();
526 auto mpiTraits = MPIImplTraits::get(moduleOp);
527 Type i32 = rewriter.getI32Type();
528 Type ptrType = LLVM::LLVMPointerType::get(op->getContext());
529 Location loc = op.getLoc();
532 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
533 auto one = LLVM::ConstantOp::create(rewriter, loc, i32, 1);
535 LLVM::AllocaOp::create(rewriter, loc, ptrType, comm.
getType(), one);
539 LLVM::LLVMFunctionType::get(i32, {comm.
getType(), i32, i32, ptrType});
542 "MPI_Comm_split", funcType);
545 LLVM::CallOp::create(rewriter, loc, funcDecl,
547 adaptor.getKey(), outPtr.getRes()});
550 Value res = LLVM::LoadOp::create(rewriter, loc, i32, outPtr.getResult());
551 res = LLVM::SExtOp::create(rewriter, loc, rewriter.getI64Type(), res);
555 SmallVector<Value> replacements;
557 replacements.push_back(callOp.getResult());
560 replacements.push_back(res);
561 rewriter.replaceOp(op, replacements);
571struct CommRankOpLowering :
public ConvertOpToLLVMPattern<mpi::CommRankOp> {
575 matchAndRewrite(mpi::CommRankOp op, OpAdaptor adaptor,
576 ConversionPatternRewriter &rewriter)
const override {
578 Location loc = op.getLoc();
579 MLIRContext *context = rewriter.getContext();
580 Type i32 = rewriter.getI32Type();
583 Type ptrType = LLVM::LLVMPointerType::get(context);
586 auto moduleOp = op->getParentOfType<ModuleOp>();
588 auto mpiTraits = MPIImplTraits::get(moduleOp);
590 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
594 LLVM::LLVMFunctionType::get(i32, {comm.
getType(), ptrType});
597 moduleOp, loc, rewriter,
"MPI_Comm_rank", rankFuncType);
600 auto one = LLVM::ConstantOp::create(rewriter, loc, i32, 1);
601 auto rankptr = LLVM::AllocaOp::create(rewriter, loc, ptrType, i32, one);
602 auto callOp = LLVM::CallOp::create(rewriter, loc, initDecl,
607 LLVM::LoadOp::create(rewriter, loc, i32, rankptr.getResult());
611 SmallVector<Value> replacements;
613 replacements.push_back(callOp.getResult());
616 replacements.push_back(loadedRank.getRes());
617 rewriter.replaceOp(op, replacements);
627static Value createOrFoldCommSize(ConversionPatternRewriter &rewriter,
628 Location loc, Value commOrg,
630 auto i32 = rewriter.getI32Type();
631 auto nRanksOp = mpi::CommSizeOp::create(rewriter, loc, i32, commOrg);
632 if (succeeded(
FoldToDLTIConst(nRanksOp,
"MPI:comm_world_size", rewriter)))
633 return nRanksOp.getSize();
634 rewriter.eraseOp(nRanksOp);
635 return mpi::CommSizeOp::create(rewriter, loc, i32, commAdapt).getSize();
638struct CommSizeOpLowering :
public ConvertOpToLLVMPattern<mpi::CommSizeOp> {
642 matchAndRewrite(mpi::CommSizeOp op, OpAdaptor adaptor,
643 ConversionPatternRewriter &rewriter)
const override {
645 Location loc = op.getLoc();
646 MLIRContext *context = rewriter.getContext();
647 Type i32 = rewriter.getI32Type();
650 Type ptrType = LLVM::LLVMPointerType::get(context);
653 auto moduleOp = op->getParentOfType<ModuleOp>();
655 auto mpiTraits = MPIImplTraits::get(moduleOp);
657 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
661 LLVM::LLVMFunctionType::get(i32, {comm.
getType(), ptrType});
664 moduleOp, loc, rewriter,
"MPI_Comm_size", SizeFuncType);
667 auto one = LLVM::ConstantOp::create(rewriter, loc, i32, 1);
668 auto sizeptr = LLVM::AllocaOp::create(rewriter, loc, ptrType, i32, one);
669 auto callOp = LLVM::CallOp::create(rewriter, loc, initDecl,
674 LLVM::LoadOp::create(rewriter, loc, i32, sizeptr.getResult());
678 SmallVector<Value> replacements;
680 replacements.push_back(callOp.getResult());
683 replacements.push_back(loadedSize.getRes());
684 rewriter.replaceOp(op, replacements);
694struct SendOpLowering :
public ConvertOpToLLVMPattern<mpi::SendOp> {
698 matchAndRewrite(mpi::SendOp op, OpAdaptor adaptor,
699 ConversionPatternRewriter &rewriter)
const override {
701 Location loc = op.getLoc();
702 MLIRContext *context = rewriter.getContext();
703 Type i32 = rewriter.getI32Type();
704 Type elemType = op.getRef().getType().getElementType();
705 int64_t rank = op.getRef().getType().getRank();
708 Type ptrType = LLVM::LLVMPointerType::get(context);
711 auto moduleOp = op->getParentOfType<ModuleOp>();
714 auto [dataPtr, size] =
715 getRawPtrAndSize(loc, rewriter, adaptor.getRef(), rank, elemType);
716 auto mpiTraits = MPIImplTraits::get(moduleOp);
717 Value dataType = mpiTraits->getDataType(loc, rewriter, elemType);
718 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
722 auto funcType = LLVM::LLVMFunctionType::get(
723 i32, {ptrType, i32, dataType.
getType(), i32, i32, comm.
getType()});
725 LLVM::LLVMFuncOp funcDecl =
729 auto funcCall = LLVM::CallOp::create(rewriter, loc, funcDecl,
732 adaptor.getTag(), comm});
734 rewriter.replaceOp(op, funcCall.getResult());
736 rewriter.eraseOp(op);
746struct RecvOpLowering :
public ConvertOpToLLVMPattern<mpi::RecvOp> {
750 matchAndRewrite(mpi::RecvOp op, OpAdaptor adaptor,
751 ConversionPatternRewriter &rewriter)
const override {
753 Location loc = op.getLoc();
754 MLIRContext *context = rewriter.getContext();
755 Type i32 = rewriter.getI32Type();
756 Type i64 = rewriter.getI64Type();
757 Type elemType = op.getRef().getType().getElementType();
758 int64_t rank = op.getRef().getType().getRank();
761 Type ptrType = LLVM::LLVMPointerType::get(context);
764 auto moduleOp = op->getParentOfType<ModuleOp>();
767 auto [dataPtr, size] =
768 getRawPtrAndSize(loc, rewriter, adaptor.getRef(), rank, elemType);
769 auto mpiTraits = MPIImplTraits::get(moduleOp);
770 Value dataType = mpiTraits->getDataType(loc, rewriter, elemType);
771 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
772 Value statusIgnore = LLVM::ConstantOp::create(rewriter, loc, i64,
773 mpiTraits->getStatusIgnore());
775 LLVM::IntToPtrOp::create(rewriter, loc, ptrType, statusIgnore);
780 LLVM::LLVMFunctionType::get(i32, {ptrType, i32, dataType.
getType(), i32,
781 i32, comm.
getType(), ptrType});
783 LLVM::LLVMFuncOp funcDecl =
787 auto funcCall = LLVM::CallOp::create(
788 rewriter, loc, funcDecl,
789 ValueRange{dataPtr, size, dataType, adaptor.getSource(),
790 adaptor.getTag(), comm, statusIgnore});
792 rewriter.replaceOp(op, funcCall.getResult());
794 rewriter.eraseOp(op);
804struct AllGatherOpLowering :
public ConvertOpToLLVMPattern<mpi::AllGatherOp> {
808 matchAndRewrite(mpi::AllGatherOp op, OpAdaptor adaptor,
809 ConversionPatternRewriter &rewriter)
const override {
810 Location loc = op.getLoc();
811 MLIRContext *context = rewriter.getContext();
812 Type sElemType = op.getSendbuf().getType().getElementType();
813 Type rElemType = op.getRecvbuf().getType().getElementType();
814 int64_t sRank = op.getSendbuf().getType().getRank();
815 int64_t rRank = op.getRecvbuf().getType().getRank();
816 auto [sendPtr, sendSize] =
817 getRawPtrAndSize(loc, rewriter, adaptor.getSendbuf(), sRank, sElemType);
818 auto [recvPtr, recvSize] =
819 getRawPtrAndSize(loc, rewriter, adaptor.getRecvbuf(), rRank, rElemType);
821 auto moduleOp = op->getParentOfType<ModuleOp>();
822 auto mpiTraits = MPIImplTraits::get(moduleOp);
823 Value sDataType = mpiTraits->getDataType(loc, rewriter, sElemType);
824 Value rDataType = mpiTraits->getDataType(loc, rewriter, rElemType);
825 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
827 Type ptrType = LLVM::LLVMPointerType::get(context);
828 Type i32 = rewriter.getI32Type();
833 auto funcType = LLVM::LLVMFunctionType::get(
834 i32, {ptrType, i32, sDataType.
getType(), ptrType, i32,
837 LLVM::LLVMFuncOp funcDecl =
842 createOrFoldCommSize(rewriter, loc, op.getComm(), adaptor.getComm());
843 Value recvCountPerRank =
844 LLVM::UDivOp::create(rewriter, loc, i32, recvSize, nRanks);
848 LLVM::CallOp::create(rewriter, loc, funcDecl,
849 ValueRange{sendPtr, sendSize, sDataType, recvPtr,
850 recvCountPerRank, rDataType, comm});
853 rewriter.replaceOp(op, funcCall.getResult());
855 rewriter.eraseOp(op);
865struct AllReduceOpLowering :
public ConvertOpToLLVMPattern<mpi::AllReduceOp> {
869 matchAndRewrite(mpi::AllReduceOp op, OpAdaptor adaptor,
870 ConversionPatternRewriter &rewriter)
const override {
871 Location loc = op.getLoc();
872 MLIRContext *context = rewriter.getContext();
873 Type i32 = rewriter.getI32Type();
874 Type i64 = rewriter.getI64Type();
875 Type elemType = op.getSendbuf().getType().getElementType();
876 int64_t sRank = op.getSendbuf().getType().getRank();
877 int64_t rRank = op.getRecvbuf().getType().getRank();
880 Type ptrType = LLVM::LLVMPointerType::get(context);
881 auto moduleOp = op->getParentOfType<ModuleOp>();
882 auto mpiTraits = MPIImplTraits::get(moduleOp);
883 auto [sendPtr, sendSize] =
884 getRawPtrAndSize(loc, rewriter, adaptor.getSendbuf(), sRank, elemType);
885 auto [recvPtr, recvSize] =
886 getRawPtrAndSize(loc, rewriter, adaptor.getRecvbuf(), rRank, elemType);
889 if (adaptor.getSendbuf() == adaptor.getRecvbuf()) {
890 sendPtr = LLVM::ConstantOp::create(
892 reinterpret_cast<int64_t
>(mpiTraits->getInPlace()));
893 sendPtr = LLVM::IntToPtrOp::create(rewriter, loc, ptrType, sendPtr);
896 Value dataType = mpiTraits->getDataType(loc, rewriter, elemType);
897 Value mpiOp = mpiTraits->getMPIOp(loc, rewriter, op.getOp());
898 Value commWorld = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
902 auto funcType = LLVM::LLVMFunctionType::get(
906 LLVM::LLVMFuncOp funcDecl =
910 auto funcCall = LLVM::CallOp::create(
911 rewriter, loc, funcDecl,
912 ValueRange{sendPtr, recvPtr, sendSize, dataType, mpiOp, commWorld});
915 rewriter.replaceOp(op, funcCall.getResult());
917 rewriter.eraseOp(op);
927struct ReduceScatterBlockOpLowering
928 :
public ConvertOpToLLVMPattern<mpi::ReduceScatterBlockOp> {
932 matchAndRewrite(mpi::ReduceScatterBlockOp op, OpAdaptor adaptor,
933 ConversionPatternRewriter &rewriter)
const override {
934 Location loc = op.getLoc();
935 MLIRContext *context = rewriter.getContext();
936 Type i32 = rewriter.getI32Type();
937 Type i64 = rewriter.getI64Type();
938 Type elemType = op.getSendbuf().getType().getElementType();
939 int64_t sRank = op.getSendbuf().getType().getRank();
940 int64_t rRank = op.getRecvbuf().getType().getRank();
943 Type ptrType = LLVM::LLVMPointerType::get(context);
944 auto moduleOp = op->getParentOfType<ModuleOp>();
945 auto mpiTraits = MPIImplTraits::get(moduleOp);
946 auto [sendPtr, sendSize] =
947 getRawPtrAndSize(loc, rewriter, adaptor.getSendbuf(), sRank, elemType);
948 auto [recvPtr, recvSize] =
949 getRawPtrAndSize(loc, rewriter, adaptor.getRecvbuf(), rRank, elemType);
952 if (adaptor.getSendbuf() == adaptor.getRecvbuf()) {
953 sendPtr = LLVM::ConstantOp::create(
955 reinterpret_cast<int64_t
>(mpiTraits->getInPlace()));
956 sendPtr = LLVM::IntToPtrOp::create(rewriter, loc, ptrType, sendPtr);
959 Value dataType = mpiTraits->getDataType(loc, rewriter, elemType);
960 Value mpiOp = mpiTraits->getMPIOp(loc, rewriter, op.getOp());
961 Value comm = mpiTraits->castComm(loc, rewriter, adaptor.getComm());
964 createOrFoldCommSize(rewriter, loc, op.getComm(), adaptor.getComm());
965 Value totalExpected =
966 LLVM::MulOp::create(rewriter, loc, i32, recvSize, nRanks);
967 Value sizeIsValid = LLVM::ICmpOp::create(
968 rewriter, loc, LLVM::ICmpPredicate::eq, sendSize, totalExpected);
969 cf::AssertOp::create(rewriter, loc, sizeIsValid,
970 "Send buffer's size must be the receive buffer's size "
971 "times the number of ranks");
975 auto funcType = LLVM::LLVMFunctionType::get(
980 moduleOp, loc, rewriter,
"MPI_Reduce_scatter_block", funcType);
983 auto funcCall = LLVM::CallOp::create(
984 rewriter, loc, funcDecl,
985 ValueRange{sendPtr, recvPtr, recvSize, dataType, mpiOp, comm});
988 rewriter.replaceOp(op, funcCall.getResult());
990 rewriter.eraseOp(op);
1001struct FuncToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
1002 FuncToLLVMDialectInterface(Dialect *dialect)
1003 : ConvertToLLVMPatternInterface(dialect) {}
1007 void populateConvertToLLVMConversionPatterns(
1008 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
1009 RewritePatternSet &patterns)
const final {
1024 converter.addConversion([](mpi::CommType type) {
1025 return IntegerType::get(type.getContext(), 64);
1027 patterns.
add<CommRankOpLowering, CommSizeOpLowering, CommSplitOpLowering,
1028 CommWorldOpLowering, FinalizeOpLowering, InitOpLowering,
1029 SendOpLowering, RecvOpLowering, AllGatherOpLowering,
1030 AllReduceOpLowering, ReduceScatterBlockOpLowering>(converter);
1035 dialect->addInterfaces<FuncToLLVMDialectInterface>();
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This provides public APIs that all operations should have.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
FailureOr< Attribute > query(Operation *op, ArrayRef< DataLayoutEntryKey > keys, bool emitError=false)
Perform a DLTI-query at op, recursively querying each key of keys on query interface-implementing att...
LogicalResult FoldToDLTIConst(OpT op, const char *key, mlir::PatternRewriter &b)
void populateMPIToLLVMConversionPatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns)
void registerConvertMPIToLLVMInterface(DialectRegistry ®istry)
Include the generated interface declarations.
LLVM::LLVMFuncOp getOrDefineFunction(Operation *moduleOp, Location loc, OpBuilder &b, StringRef name, LLVM::LLVMFunctionType type)
Note that these functions don't take a SymbolTable because GPU module lowerings can have name collisi...
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...