40#include "llvm/ADT/STLExtras.h"
42#define DEBUG_TYPE "gpu-to-llvm"
45#define GEN_PASS_DEF_GPUTOLLVMCONVERSIONPASS
46#include "mlir/Conversion/Passes.h.inc"
52class GpuToLLVMConversionPass
56 void getDependentDialects(DialectRegistry ®istry)
const final {
57 Base::getDependentDialects(registry);
61 void runOnOperation()
override;
64template <
typename OpTy>
67 explicit ConvertOpToGpuRuntimeCallPattern(
68 const LLVMTypeConverter &typeConverter, PatternBenefit benefit = 1)
69 : ConvertOpToLLVMPattern<OpTy>(typeConverter, benefit) {}
72 Value
getNumElements(ConversionPatternRewriter &rewriter, Location loc,
73 MemRefType type, MemRefDescriptor desc)
const {
75 if (type.hasStaticShape())
77 rewriter, loc, indexType, type.getNumElements());
79 uint64_t rank = type.getRank();
80 Value numElements = desc.
size(rewriter, loc, 0);
81 for (
unsigned i = 1; i < rank; i++)
82 numElements = LLVM::MulOp::create(rewriter, loc, numElements,
83 desc.
size(rewriter, loc, i));
87 MLIRContext *context = &this->getTypeConverter()->
getContext();
89 Type llvmVoidType = LLVM::LLVMVoidType::get(context);
90 LLVM::LLVMPointerType llvmPointerType = LLVM::LLVMPointerType::get(context);
91 Type llvmInt8Type = IntegerType::get(context, 8);
92 Type llvmInt16Type = IntegerType::get(context, 16);
93 Type llvmInt32Type = IntegerType::get(context, 32);
94 Type llvmInt64Type = IntegerType::get(context, 64);
95 Type llvmFloat32Type = Float32Type::get(context);
96 Type llvmIntPtrType = IntegerType::get(
97 context, this->getTypeConverter()->getPointerBitwidth(0));
99 FunctionCallBuilder streamCreateCallBuilder = {
100 "mgpuStreamCreate", llvmPointerType , {}};
101 FunctionCallBuilder streamDestroyCallBuilder = {
102 "mgpuStreamDestroy", llvmVoidType, {llvmPointerType }};
103 FunctionCallBuilder streamSynchronizeCallBuilder = {
104 "mgpuStreamSynchronize",
107 FunctionCallBuilder streamWaitEventCallBuilder = {
108 "mgpuStreamWaitEvent",
110 {llvmPointerType , llvmPointerType }};
111 FunctionCallBuilder eventCreateCallBuilder = {
112 "mgpuEventCreate", llvmPointerType , {}};
113 FunctionCallBuilder eventDestroyCallBuilder = {
114 "mgpuEventDestroy", llvmVoidType, {llvmPointerType }};
115 FunctionCallBuilder eventSynchronizeCallBuilder = {
116 "mgpuEventSynchronize",
119 FunctionCallBuilder eventRecordCallBuilder = {
122 {llvmPointerType , llvmPointerType }};
123 FunctionCallBuilder hostRegisterCallBuilder = {
124 "mgpuMemHostRegisterMemRef",
129 FunctionCallBuilder hostUnregisterCallBuilder = {
130 "mgpuMemHostUnregisterMemRef",
135 FunctionCallBuilder allocCallBuilder = {
141 FunctionCallBuilder deallocCallBuilder = {
144 {llvmPointerType , llvmPointerType }};
145 FunctionCallBuilder memcpyCallBuilder = {
148 {llvmPointerType , llvmPointerType ,
151 FunctionCallBuilder memset16CallBuilder = {
158 FunctionCallBuilder memset32CallBuilder = {
161 {llvmPointerType , llvmInt32Type ,
164 FunctionCallBuilder setDefaultDeviceCallBuilder = {
165 "mgpuSetDefaultDevice",
168 FunctionCallBuilder createDnVecCallBuilder = {
171 {llvmIntPtrType, llvmPointerType, llvmInt32Type,
173 FunctionCallBuilder destroyDnVecCallBuilder = {
176 {llvmPointerType, llvmPointerType }};
177 FunctionCallBuilder createDnMatCallBuilder = {
180 {llvmIntPtrType, llvmIntPtrType, llvmPointerType, llvmInt32Type,
182 FunctionCallBuilder destroyDnMatCallBuilder = {
185 {llvmPointerType, llvmPointerType }};
186 FunctionCallBuilder createCooCallBuilder = {
189 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
190 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
192 FunctionCallBuilder createCooAoSCallBuilder = {
195 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
196 llvmPointerType, llvmInt32Type, llvmInt32Type,
198 FunctionCallBuilder createCsrCallBuilder = {
201 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
202 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
203 llvmInt32Type, llvmPointerType }};
204 FunctionCallBuilder createCscCallBuilder = {
207 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
208 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
209 llvmInt32Type, llvmPointerType }};
210 FunctionCallBuilder createBsrCallBuilder = {
213 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmIntPtrType,
214 llvmIntPtrType, llvmPointerType, llvmPointerType, llvmPointerType,
215 llvmInt32Type, llvmInt32Type, llvmInt32Type,
217 FunctionCallBuilder destroySpMatCallBuilder = {
220 {llvmPointerType, llvmPointerType }};
221 FunctionCallBuilder spMVBufferSizeCallBuilder = {
222 "mgpuSpMVBufferSize",
224 {llvmInt32Type, llvmPointerType, llvmPointerType, llvmPointerType,
225 llvmInt32Type, llvmPointerType }};
226 FunctionCallBuilder spMVCallBuilder = {
229 {llvmInt32Type, llvmPointerType, llvmPointerType, llvmPointerType,
230 llvmInt32Type, llvmPointerType, llvmPointerType }};
231 FunctionCallBuilder createSpMMBufferSizeCallBuilder = {
232 "mgpuSpMMBufferSize",
234 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
235 llvmPointerType, llvmInt32Type, llvmPointerType }};
236 FunctionCallBuilder createSpMMCallBuilder = {
239 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
240 llvmPointerType, llvmInt32Type, llvmPointerType,
242 FunctionCallBuilder createSDDMMBufferSizeCallBuilder = {
243 "mgpuSDDMMBufferSize",
245 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
246 llvmPointerType, llvmInt32Type, llvmPointerType }};
247 FunctionCallBuilder createSDDMMCallBuilder = {
250 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
251 llvmPointerType, llvmInt32Type, llvmPointerType,
253 FunctionCallBuilder createLtDnMatCallBuilder = {
254 "mgpuCreateCuSparseLtDnMat",
256 {llvmPointerType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
257 llvmInt32Type, llvmPointerType }};
258 FunctionCallBuilder destroyCuSparseLtSpMatBuilder = {
259 "mgpuDestroyCuSparseLtSpMat",
261 {llvmPointerType, llvmPointerType }};
262 FunctionCallBuilder destroyCuSparseLtDnMatBuilder = {
263 "mgpuDestroyCuSparseLtDnMat",
265 {llvmPointerType, llvmPointerType }};
266 FunctionCallBuilder create2To4SpMatCallBuilder = {
267 "mgpuCusparseLtCreate2To4SpMat",
269 {llvmPointerType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
270 llvmInt32Type, llvmPointerType }};
271 FunctionCallBuilder createCuSparseLtSpMMBufferSizeBuilder = {
272 "mgpuCuSparseLtSpMMBufferSize",
274 {llvmPointerType, llvmInt32Type, llvmInt32Type, llvmPointerType,
275 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
277 FunctionCallBuilder createCuSparseLtSpMMBuilder = {
278 "mgpuCuSparseLtSpMM",
280 {llvmPointerType, llvmPointerType, llvmPointerType, llvmPointerType,
281 llvmPointerType, llvmPointerType, llvmPointerType }};
282 FunctionCallBuilder createSpGEMMCreateDescrBuilder = {
283 "mgpuSpGEMMCreateDescr",
286 FunctionCallBuilder createSpGEMMDestroyDescrBuilder = {
287 "mgpuSpGEMMDestroyDescr",
289 {llvmPointerType , llvmPointerType }};
290 FunctionCallBuilder createSpGEMMWorkEstimationBuilder = {
291 "mgpuSpGEMMWorkEstimation",
293 {llvmPointerType , llvmInt32Type , llvmInt32Type ,
294 llvmPointerType , llvmPointerType , llvmPointerType ,
295 llvmInt32Type , llvmIntPtrType , llvmPointerType ,
297 FunctionCallBuilder createSpGEMMComputeBuilder = {
300 {llvmPointerType , llvmInt32Type , llvmInt32Type ,
301 llvmPointerType , llvmPointerType , llvmPointerType ,
302 llvmInt32Type , llvmIntPtrType , llvmPointerType ,
304 FunctionCallBuilder createSpGEMMCopyBuilder = {
307 {llvmPointerType , llvmInt32Type , llvmInt32Type ,
308 llvmPointerType , llvmPointerType , llvmPointerType ,
309 llvmInt32Type , llvmPointerType }};
310 FunctionCallBuilder createSpMatGetSizeBuilder = {
313 {llvmPointerType , llvmPointerType , llvmPointerType ,
314 llvmPointerType , llvmPointerType }};
315 FunctionCallBuilder createSetCsrPointersBuilder = {
316 "mgpuSetCsrPointers",
318 {llvmPointerType , llvmPointerType ,
319 llvmPointerType , llvmPointerType ,
325class ConvertHostRegisterOpToGpuRuntimeCallPattern
326 :
public ConvertOpToGpuRuntimeCallPattern<gpu::HostRegisterOp> {
328 ConvertHostRegisterOpToGpuRuntimeCallPattern(
329 const LLVMTypeConverter &typeConverter)
330 : ConvertOpToGpuRuntimeCallPattern<gpu::HostRegisterOp>(typeConverter) {}
334 matchAndRewrite(gpu::HostRegisterOp hostRegisterOp, OpAdaptor adaptor,
335 ConversionPatternRewriter &rewriter)
const override;
338class ConvertHostUnregisterOpToGpuRuntimeCallPattern
339 :
public ConvertOpToGpuRuntimeCallPattern<gpu::HostUnregisterOp> {
341 ConvertHostUnregisterOpToGpuRuntimeCallPattern(
342 const LLVMTypeConverter &typeConverter)
343 : ConvertOpToGpuRuntimeCallPattern<gpu::HostUnregisterOp>(typeConverter) {
348 matchAndRewrite(gpu::HostUnregisterOp hostUnregisterOp, OpAdaptor adaptor,
349 ConversionPatternRewriter &rewriter)
const override;
354class ConvertAllocOpToGpuRuntimeCallPattern
355 :
public ConvertOpToGpuRuntimeCallPattern<gpu::AllocOp> {
357 ConvertAllocOpToGpuRuntimeCallPattern(
const LLVMTypeConverter &typeConverter)
358 : ConvertOpToGpuRuntimeCallPattern<gpu::AllocOp>(typeConverter) {}
362 matchAndRewrite(gpu::AllocOp allocOp, OpAdaptor adaptor,
363 ConversionPatternRewriter &rewriter)
const override;
368class ConvertDeallocOpToGpuRuntimeCallPattern
369 :
public ConvertOpToGpuRuntimeCallPattern<gpu::DeallocOp> {
371 ConvertDeallocOpToGpuRuntimeCallPattern(
372 const LLVMTypeConverter &typeConverter)
373 : ConvertOpToGpuRuntimeCallPattern<gpu::DeallocOp>(typeConverter) {}
377 matchAndRewrite(gpu::DeallocOp deallocOp, OpAdaptor adaptor,
378 ConversionPatternRewriter &rewriter)
const override;
381class ConvertAsyncYieldToGpuRuntimeCallPattern
382 :
public ConvertOpToGpuRuntimeCallPattern<async::YieldOp> {
384 ConvertAsyncYieldToGpuRuntimeCallPattern(
385 const LLVMTypeConverter &typeConverter, PatternBenefit benefit = 1)
386 : ConvertOpToGpuRuntimeCallPattern<async::YieldOp>(typeConverter,
391 matchAndRewrite(async::YieldOp yieldOp, OpAdaptor adaptor,
392 ConversionPatternRewriter &rewriter)
const override;
397class ConvertWaitOpToGpuRuntimeCallPattern
398 :
public ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp> {
400 ConvertWaitOpToGpuRuntimeCallPattern(
const LLVMTypeConverter &typeConverter)
401 : ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp>(typeConverter) {}
405 matchAndRewrite(gpu::WaitOp waitOp, OpAdaptor adaptor,
406 ConversionPatternRewriter &rewriter)
const override;
411class ConvertWaitAsyncOpToGpuRuntimeCallPattern
412 :
public ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp> {
414 ConvertWaitAsyncOpToGpuRuntimeCallPattern(
415 const LLVMTypeConverter &typeConverter)
416 : ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp>(typeConverter) {}
420 matchAndRewrite(gpu::WaitOp waitOp, OpAdaptor adaptor,
421 ConversionPatternRewriter &rewriter)
const override;
425class LegalizeLaunchFuncOpPattern
426 :
public ConvertOpToGpuRuntimeCallPattern<gpu::LaunchFuncOp> {
428 LegalizeLaunchFuncOpPattern(
const LLVMTypeConverter &typeConverter,
429 bool kernelBarePtrCallConv,
430 bool kernelIntersperseSizeCallConv)
431 : ConvertOpToGpuRuntimeCallPattern<gpu::LaunchFuncOp>(typeConverter),
432 kernelBarePtrCallConv(kernelBarePtrCallConv),
433 kernelIntersperseSizeCallConv(kernelIntersperseSizeCallConv) {}
437 matchAndRewrite(gpu::LaunchFuncOp launchOp, OpAdaptor adaptor,
438 ConversionPatternRewriter &rewriter)
const override;
440 bool kernelBarePtrCallConv;
441 bool kernelIntersperseSizeCallConv;
446class ConvertMemcpyOpToGpuRuntimeCallPattern
447 :
public ConvertOpToGpuRuntimeCallPattern<gpu::MemcpyOp> {
449 ConvertMemcpyOpToGpuRuntimeCallPattern(
const LLVMTypeConverter &typeConverter)
450 : ConvertOpToGpuRuntimeCallPattern<gpu::MemcpyOp>(typeConverter) {}
454 matchAndRewrite(gpu::MemcpyOp memcpyOp, OpAdaptor adaptor,
455 ConversionPatternRewriter &rewriter)
const override;
460class ConvertMemsetOpToGpuRuntimeCallPattern
461 :
public ConvertOpToGpuRuntimeCallPattern<gpu::MemsetOp> {
463 ConvertMemsetOpToGpuRuntimeCallPattern(
const LLVMTypeConverter &typeConverter)
464 : ConvertOpToGpuRuntimeCallPattern<gpu::MemsetOp>(typeConverter) {}
468 matchAndRewrite(gpu::MemsetOp memsetOp, OpAdaptor adaptor,
469 ConversionPatternRewriter &rewriter)
const override;
474class ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern
475 :
public ConvertOpToGpuRuntimeCallPattern<gpu::SetDefaultDeviceOp> {
477 ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern(
478 const LLVMTypeConverter &typeConverter)
479 : ConvertOpToGpuRuntimeCallPattern<gpu::SetDefaultDeviceOp>(
483 matchAndRewrite(gpu::SetDefaultDeviceOp op, OpAdaptor adaptor,
484 ConversionPatternRewriter &rewriter)
const override;
489#define DECLARE_CONVERT_OP_TO_GPU_RUNTIME_CALL_PATTERN(op_name) \
490 class Convert##op_name##ToGpuRuntimeCallPattern \
491 : public ConvertOpToGpuRuntimeCallPattern<gpu::op_name> { \
493 Convert##op_name##ToGpuRuntimeCallPattern( \
494 const LLVMTypeConverter &typeConverter) \
495 : ConvertOpToGpuRuntimeCallPattern<gpu::op_name>(typeConverter) {} \
499 matchAndRewrite(gpu::op_name op, OpAdaptor adaptor, \
500 ConversionPatternRewriter &rewriter) const override; \
527void GpuToLLVMConversionPass::runOnOperation() {
538 vector::populateVectorFromElementsUnrollPatterns(patterns);
540 return signalPassFailure();
543 LowerToLLVMOptions
options(context);
544 options.useBarePtrCallConv = hostBarePtrCallConv;
545 RewritePatternSet patterns(context);
546 ConversionTarget
target(*context);
547 target.addLegalDialect<LLVM::LLVMDialect>();
548 LLVMTypeConverter converter(context,
options);
553 auto *iface = dyn_cast<ConvertToLLVMPatternInterface>(dialect);
556 iface->populateConvertToLLVMConversionPatterns(
target, converter, patterns);
561 target.addLegalOp<gpu::GPUModuleOp, gpu::BinaryOp>();
563 target.addDynamicallyLegalOp<gpu::LaunchFuncOp>(
564 [&](gpu::LaunchFuncOp op) ->
bool {
return converter.isLegal(op); });
572 kernelBarePtrCallConv,
573 kernelIntersperseSizeCallConv);
576 applyPartialConversion(getOperation(),
target, std::move(patterns))))
582 auto module = builder.getBlock()->getParent()->getParentOfType<ModuleOp>();
583 auto function = [&] {
584 if (
auto function = module.lookupSymbol<LLVM::LLVMFuncOp>(
functionName))
589 return LLVM::CallOp::create(builder, loc, function, arguments);
606 llvm_unreachable(
"unsupported type");
612 if (llvm::isa<ComplexType>(type)) {
614 auto elementType = cast<ComplexType>(type).getElementType();
615 if (elementType.isBF16())
617 if (elementType.isF16())
619 if (elementType.isF32())
621 if (elementType.isF64())
623 if (elementType.isInteger(8))
625 if (elementType.isInteger(16))
627 if (elementType.isInteger(32))
645 llvm_unreachable(
"unsupported element type");
649 return spMat.
getDefiningOp<gpu::Create2To4SpMatOp>().getPruneFlag();
674 llvm_unreachable(
"cannot find spmat def");
679 auto spmmOp = dyn_cast<gpu::SpMMOp>(user);
691 ConversionPatternRewriter &rewriter) {
692 if (!llvm::all_of(operands, [](
Value value) {
695 return rewriter.notifyMatchFailure(
696 op,
"Cannot convert if operands aren't of LLVM type.");
702 gpu::AsyncOpInterface op) {
703 if (op.getAsyncDependencies().size() != 1)
704 return rewriter.notifyMatchFailure(
705 op,
"Can only convert with exactly one async dependency.");
707 if (!op.getAsyncToken())
708 return rewriter.notifyMatchFailure(op,
"Can convert only async version.");
713LogicalResult ConvertHostRegisterOpToGpuRuntimeCallPattern::matchAndRewrite(
714 gpu::HostRegisterOp hostRegisterOp, OpAdaptor adaptor,
715 ConversionPatternRewriter &rewriter)
const {
716 auto *op = hostRegisterOp.getOperation();
720 Location loc = op->getLoc();
722 auto memRefType = hostRegisterOp.getValue().getType();
723 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
726 auto arguments = getTypeConverter()->promoteOperands(
727 loc, op->getOperands(), adaptor.getOperands(), rewriter);
728 arguments.push_back(elementSize);
729 hostRegisterCallBuilder.create(loc, rewriter, arguments);
731 rewriter.eraseOp(op);
735LogicalResult ConvertHostUnregisterOpToGpuRuntimeCallPattern::matchAndRewrite(
736 gpu::HostUnregisterOp hostUnregisterOp, OpAdaptor adaptor,
737 ConversionPatternRewriter &rewriter)
const {
738 Operation *op = hostUnregisterOp.getOperation();
742 Location loc = op->
getLoc();
744 auto memRefType = hostUnregisterOp.getValue().getType();
745 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
748 auto arguments = getTypeConverter()->promoteOperands(
749 loc, op->
getOperands(), adaptor.getOperands(), rewriter);
750 arguments.push_back(elementSize);
751 hostUnregisterCallBuilder.create(loc, rewriter, arguments);
753 rewriter.eraseOp(op);
757LogicalResult ConvertAllocOpToGpuRuntimeCallPattern::matchAndRewrite(
758 gpu::AllocOp allocOp, OpAdaptor adaptor,
759 ConversionPatternRewriter &rewriter)
const {
761 MemRefType memRefType = allocOp.getType();
764 !isConvertibleAndHasIdentityMaps(memRefType))
767 auto loc = allocOp.getLoc();
769 bool isShared = allocOp.getHostShared();
771 if (isShared && allocOp.getAsyncToken())
772 return rewriter.notifyMatchFailure(
773 allocOp,
"Host Shared allocation cannot be done async");
774 if (adaptor.getAsyncDependencies().size() > 1)
775 return rewriter.notifyMatchFailure(
776 allocOp,
"Can convert with at most one async dependency.");
780 SmallVector<Value, 4> shape;
781 SmallVector<Value, 4> strides;
783 getMemRefDescriptorSizes(loc, memRefType, adaptor.getDynamicSizes(), rewriter,
784 shape, strides, sizeBytes);
788 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
789 Value stream = adaptor.getAsyncDependencies().empty()
791 : adaptor.getAsyncDependencies().front();
793 auto isHostShared = mlir::LLVM::ConstantOp::create(
794 rewriter, loc, llvmInt8Type, rewriter.getI8IntegerAttr(isShared));
797 allocCallBuilder.create(loc, rewriter, {sizeBytes, stream, isHostShared})
802 unsigned dstAddrSpace = memRefType.getMemorySpaceAsInt();
803 unsigned srcAddrSpace =
804 cast<LLVM::LLVMPointerType>(allocatedPtr.
getType()).getAddressSpace();
805 if (dstAddrSpace != srcAddrSpace) {
807 LLVM::LLVMPointerType::get(rewriter.getContext(), dstAddrSpace);
809 LLVM::AddrSpaceCastOp::create(rewriter, loc, targetPtrTy, allocatedPtr);
813 Value alignedPtr = allocatedPtr;
816 auto memRefDescriptor = this->createMemRefDescriptor(
817 loc, memRefType, allocatedPtr, alignedPtr, shape, strides, rewriter);
819 if (allocOp.getAsyncToken()) {
821 rewriter.replaceOp(allocOp, {memRefDescriptor, stream});
823 rewriter.replaceOp(allocOp, {memRefDescriptor});
829LogicalResult ConvertDeallocOpToGpuRuntimeCallPattern::matchAndRewrite(
830 gpu::DeallocOp deallocOp, OpAdaptor adaptor,
831 ConversionPatternRewriter &rewriter)
const {
834 if (adaptor.getAsyncDependencies().size() > 1)
835 return rewriter.notifyMatchFailure(
836 deallocOp,
"Can convert with at most one async dependency.");
838 Location loc = deallocOp.getLoc();
841 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
842 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
843 Value stream = adaptor.getAsyncDependencies().empty()
845 : adaptor.getAsyncDependencies().front();
846 deallocCallBuilder.create(loc, rewriter, {pointer, stream});
848 if (deallocOp.getAsyncToken()) {
850 rewriter.replaceOp(deallocOp, {stream});
853 rewriter.eraseOp(deallocOp);
859 return isa<gpu::AsyncTokenType>(value.
getType());
874LogicalResult ConvertAsyncYieldToGpuRuntimeCallPattern::matchAndRewrite(
875 async::YieldOp yieldOp, OpAdaptor adaptor,
876 ConversionPatternRewriter &rewriter)
const {
878 return rewriter.notifyMatchFailure(yieldOp,
"no gpu async token operand");
880 Location loc = yieldOp.getLoc();
881 SmallVector<Value, 4> newOperands(adaptor.getOperands());
882 llvm::SmallDenseSet<Value> streams;
883 for (
auto &operand : yieldOp->getOpOperands()) {
886 auto idx = operand.getOperandNumber();
887 auto stream = adaptor.getOperands()[idx];
888 auto event = eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
889 eventRecordCallBuilder.create(loc, rewriter, {event, stream});
890 newOperands[idx] = event;
891 streams.insert(stream);
893 for (
auto stream : streams)
894 streamDestroyCallBuilder.create(loc, rewriter, {stream});
896 rewriter.modifyOpInPlace(yieldOp, [&] { yieldOp->setOperands(newOperands); });
902 assert(isa<LLVM::LLVMPointerType>(value.
getType()));
904 return *defOp.getCallee() == functionName;
912LogicalResult ConvertWaitOpToGpuRuntimeCallPattern::matchAndRewrite(
913 gpu::WaitOp waitOp, OpAdaptor adaptor,
914 ConversionPatternRewriter &rewriter)
const {
915 if (waitOp.getAsyncToken())
916 return rewriter.notifyMatchFailure(waitOp,
"Cannot convert async op.");
918 Location loc = waitOp.getLoc();
920 for (
auto operand : adaptor.getOperands()) {
923 streamSynchronizeCallBuilder.create(loc, rewriter, {operand});
924 streamDestroyCallBuilder.create(loc, rewriter, {operand});
928 eventSynchronizeCallBuilder.create(loc, rewriter, {operand});
929 eventDestroyCallBuilder.create(loc, rewriter, {operand});
933 rewriter.eraseOp(waitOp);
942LogicalResult ConvertWaitAsyncOpToGpuRuntimeCallPattern::matchAndRewrite(
943 gpu::WaitOp waitOp, OpAdaptor adaptor,
944 ConversionPatternRewriter &rewriter)
const {
945 if (!waitOp.getAsyncToken())
946 return rewriter.notifyMatchFailure(waitOp,
"Can only convert async op.");
948 Location loc = waitOp.getLoc();
950 auto insertionPoint = rewriter.saveInsertionPoint();
951 SmallVector<Value, 1> events;
953 llvm::zip(waitOp.getAsyncDependencies(), adaptor.getOperands())) {
954 auto operand = std::get<1>(pair);
958 auto *defOp = std::get<0>(pair).getDefiningOp();
959 rewriter.setInsertionPointAfter(defOp);
960 auto event = eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
961 eventRecordCallBuilder.create(loc, rewriter, {event, operand});
962 events.push_back(event);
966 events.push_back(operand);
969 rewriter.restoreInsertionPoint(insertionPoint);
970 auto stream = streamCreateCallBuilder.create(loc, rewriter, {}).getResult();
971 for (
auto event : events)
972 streamWaitEventCallBuilder.create(loc, rewriter, {stream,
event});
973 for (
auto event : events)
974 eventDestroyCallBuilder.create(loc, rewriter, {
event});
975 rewriter.replaceOp(waitOp, {stream});
981LogicalResult LegalizeLaunchFuncOpPattern::matchAndRewrite(
982 gpu::LaunchFuncOp launchOp, OpAdaptor adaptor,
983 ConversionPatternRewriter &rewriter)
const {
990 if (!launchOp.getAsyncToken() && !launchOp.getAsyncDependencies().empty())
991 return rewriter.notifyMatchFailure(
992 launchOp,
"Cannot convert non-async op with async dependencies.");
994 Location loc = launchOp.getLoc();
996 Value stream = Value();
997 if (!adaptor.getAsyncDependencies().empty()) {
998 stream = adaptor.getAsyncDependencies().front();
1001 if (adaptor.getAsyncDependencies().size() > 1) {
1002 auto insertionPoint = rewriter.saveInsertionPoint();
1003 SmallVector<Value, 4> events;
1004 for (
auto [origDep, convertedDep] :
1005 llvm::zip(launchOp.getAsyncDependencies().drop_front(),
1006 adaptor.getAsyncDependencies().drop_front())) {
1008 streamCreateCallBuilder.functionName)) {
1009 events.push_back(convertedDep);
1012 Operation *defOp = origDep.getDefiningOp();
1013 rewriter.setInsertionPointAfter(defOp);
1015 eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
1016 eventRecordCallBuilder.create(loc, rewriter, {event, convertedDep});
1017 events.push_back(event);
1019 rewriter.restoreInsertionPoint(insertionPoint);
1020 for (Value event : events)
1021 streamWaitEventCallBuilder.create(loc, rewriter, {stream,
event});
1022 for (Value event : events)
1023 eventDestroyCallBuilder.create(loc, rewriter, {
event});
1028 else if (launchOp.getAsyncToken())
1029 stream = streamCreateCallBuilder.create(loc, rewriter, {}).getResult();
1034 OperandRange origArguments = launchOp.getKernelOperands();
1035 bool effectiveBarePtr = kernelBarePtrCallConv ||
1036 getTypeConverter()->getOptions().useBarePtrCallConv;
1037 if (effectiveBarePtr) {
1038 for (Value arg : origArguments) {
1039 if (isa<UnrankedMemRefType>(arg.getType()))
1040 return rewriter.notifyMatchFailure(
1041 loc,
"unranked memref kernel argument is not supported with "
1042 "the bare-pointer calling convention");
1045 SmallVector<Value, 8> llvmArguments = getTypeConverter()->promoteOperands(
1046 loc, origArguments, adaptor.getKernelOperands(), rewriter,
1047 kernelBarePtrCallConv);
1048 SmallVector<Value, 8> llvmArgumentsWithSizes;
1051 if (kernelIntersperseSizeCallConv) {
1052 if (origArguments.size() != llvmArguments.size()) {
1054 return rewriter.notifyMatchFailure(
1056 "Cannot add sizes to arguments with one-to-many LLVM IR expansion.");
1059 llvmArgumentsWithSizes.reserve(llvmArguments.size() * 2);
1060 for (
auto [llvmArg, origArg] : zip_equal(llvmArguments, origArguments)) {
1061 auto memrefTy = dyn_cast<MemRefType>(origArg.getType());
1063 return rewriter.notifyMatchFailure(
1064 launchOp,
"Operand to launch op is not a memref.");
1067 if (!memrefTy.hasStaticShape() ||
1068 !memrefTy.getElementType().isIntOrFloat()) {
1069 return rewriter.notifyMatchFailure(
1070 launchOp,
"Operand to launch op is not a memref with a static "
1071 "shape and an integer or float element type.");
1074 unsigned bitwidth = memrefTy.getElementTypeBitWidth();
1075 if (bitwidth % 8 != 0) {
1076 return rewriter.notifyMatchFailure(
1077 launchOp,
"Operand to launch op is not a memref with a "
1078 "byte-aligned element type.");
1081 uint64_t staticSize =
static_cast<uint64_t
>(bitwidth / 8) *
1082 static_cast<uint64_t
>(memrefTy.getNumElements());
1086 llvmArgumentsWithSizes.push_back(llvmArg);
1087 llvmArgumentsWithSizes.push_back(sizeArg);
1091 std::optional<gpu::KernelDim3> clusterSize = std::nullopt;
1092 if (launchOp.hasClusterSize()) {
1094 gpu::KernelDim3{adaptor.getClusterSizeX(), adaptor.getClusterSizeY(),
1095 adaptor.getClusterSizeZ()};
1097 auto newLaunchOp = gpu::LaunchFuncOp::create(
1098 rewriter, launchOp.getLoc(), launchOp.getKernelAttr(),
1099 gpu::KernelDim3{adaptor.getGridSizeX(), adaptor.getGridSizeY(),
1100 adaptor.getGridSizeZ()},
1101 gpu::KernelDim3{adaptor.getBlockSizeX(), adaptor.getBlockSizeY(),
1102 adaptor.getBlockSizeZ()},
1103 adaptor.getDynamicSharedMemorySize(),
1104 llvmArgumentsWithSizes.empty() ? llvmArguments : llvmArgumentsWithSizes,
1105 nullptr, {}, stream, clusterSize);
1106 if (launchOp.getCooperative())
1107 newLaunchOp.setCooperative(
true);
1108 if (launchOp.getAsyncToken())
1109 rewriter.replaceOp(launchOp, {stream});
1111 rewriter.eraseOp(launchOp);
1116 ConversionPatternRewriter &rewriter,
1117 LLVM::LLVMPointerType destinationType,
1120 auto sourceTy = cast<LLVM::LLVMPointerType>(sourcePtr.
getType());
1121 if (destinationType.getAddressSpace() != sourceTy.getAddressSpace())
1122 sourcePtr = LLVM::AddrSpaceCastOp::create(
1124 LLVM::LLVMPointerType::get(rewriter.getContext(),
1125 destinationType.getAddressSpace()),
1130LogicalResult ConvertMemcpyOpToGpuRuntimeCallPattern::matchAndRewrite(
1131 gpu::MemcpyOp memcpyOp, OpAdaptor adaptor,
1132 ConversionPatternRewriter &rewriter)
const {
1133 auto memRefType = cast<MemRefType>(memcpyOp.getSrc().getType());
1136 !isConvertibleAndHasIdentityMaps(memRefType) ||
1140 auto loc = memcpyOp.getLoc();
1142 MemRefDescriptor srcDesc(adaptor.getSrc());
1143 Value numElements =
getNumElements(rewriter, loc, memRefType, srcDesc);
1145 Type elementPtrType = getElementPtrType(memRefType);
1146 Value nullPtr = LLVM::ZeroOp::create(rewriter, loc, elementPtrType);
1147 Value gepPtr = LLVM::GEPOp::create(
1148 rewriter, loc, elementPtrType,
1149 typeConverter->convertType(memRefType.getElementType()), nullPtr,
1152 LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), gepPtr);
1155 srcDesc.alignedPtr(rewriter, loc),
1156 *getTypeConverter());
1158 loc, rewriter, llvmPointerType,
1159 MemRefDescriptor(adaptor.getDst()).alignedPtr(rewriter, loc),
1160 *getTypeConverter());
1162 auto stream = adaptor.getAsyncDependencies().front();
1163 memcpyCallBuilder.create(loc, rewriter, {dst, src, sizeBytes, stream});
1165 rewriter.replaceOp(memcpyOp, {stream});
1170LogicalResult ConvertMemsetOpToGpuRuntimeCallPattern::matchAndRewrite(
1171 gpu::MemsetOp memsetOp, OpAdaptor adaptor,
1172 ConversionPatternRewriter &rewriter)
const {
1173 auto memRefType = cast<MemRefType>(memsetOp.getDst().getType());
1176 !isConvertibleAndHasIdentityMaps(memRefType) ||
1180 auto loc = memsetOp.getLoc();
1182 Type valueType = adaptor.getValue().getType();
1185 if (!valueType.
isIntOrFloat() || (bitWidth != 16 && bitWidth != 32)) {
1186 return rewriter.notifyMatchFailure(
1187 memsetOp,
"value must be a 16 or 32 bit int or float");
1191 Type bitCastType = valueTypeWidth == 32 ? llvmInt32Type : llvmInt16Type;
1193 MemRefDescriptor dstDesc(adaptor.getDst());
1194 Value numElements =
getNumElements(rewriter, loc, memRefType, dstDesc);
1197 LLVM::BitcastOp::create(rewriter, loc, bitCastType, adaptor.getValue());
1199 dstDesc.alignedPtr(rewriter, loc),
1200 *getTypeConverter());
1202 auto stream = adaptor.getAsyncDependencies().front();
1203 FunctionCallBuilder builder =
1204 valueTypeWidth == 32 ? memset32CallBuilder : memset16CallBuilder;
1205 builder.
create(loc, rewriter, {dst, value, numElements, stream});
1207 rewriter.replaceOp(memsetOp, {stream});
1211LogicalResult ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern::matchAndRewrite(
1212 gpu::SetDefaultDeviceOp op, OpAdaptor adaptor,
1213 ConversionPatternRewriter &rewriter)
const {
1214 Location loc = op.getLoc();
1215 auto call = setDefaultDeviceCallBuilder.create(loc, rewriter,
1216 {adaptor.getDevIndex()});
1217 rewriter.replaceOp(op, call);
1221template <
typename T>
1224 return LLVM::ConstantOp::create(builder, loc, llvmInt32Type,
1225 static_cast<int32_t
>(tValue));
1228template <
typename T>
1231 return LLVM::ConstantOp::create(
1232 builder, loc, llvmFloat32Type,
1236LogicalResult ConvertCreateDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1237 gpu::CreateDnTensorOp op, OpAdaptor adaptor,
1238 ConversionPatternRewriter &rewriter)
const {
1242 Location loc = op.getLoc();
1243 auto stream = adaptor.getAsyncDependencies().front();
1245 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1246 Type dType = op.getMemref().
getType().getElementType();
1249 SmallVector<Value, 4> dims;
1250 for (Value dim : adaptor.getDims()) {
1251 dims.push_back(dim);
1261 if (dims.size() == 2) {
1265 handle = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1266 llvmInt8Type, handleSz, 16);
1267 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1269 createLtDnMatCallBuilder
1270 .create(loc, rewriter,
1271 {handle, dims[0], dims[1], pTensor, dtp, stream})
1275 createDnMatCallBuilder
1276 .create(loc, rewriter, {dims[0], dims[1], pTensor, dtp, stream})
1280 assert(dims.size() == 1 &&
"Only 1D and 2D tensors are supported");
1281 handle = createDnVecCallBuilder
1282 .create(loc, rewriter, {dims[0], pTensor, dtp, stream})
1285 rewriter.replaceOp(op, {handle, stream});
1289LogicalResult ConvertDestroyDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1290 gpu::DestroyDnTensorOp op, OpAdaptor adaptor,
1291 ConversionPatternRewriter &rewriter)
const {
1295 Location loc = op.getLoc();
1296 auto stream = adaptor.getAsyncDependencies().front();
1297 auto definingOp = op.getDnTensor().
getDefiningOp<gpu::CreateDnTensorOp>();
1298 SmallVector<Value, 4> dims;
1299 for (Value dim : definingOp.getDims()) {
1300 dims.push_back(dim);
1302 if (dims.size() == 2) {
1306 destroyCuSparseLtDnMatBuilder.create(loc, rewriter,
1307 {adaptor.getDnTensor(), stream});
1309 destroyDnMatCallBuilder.create(loc, rewriter,
1310 {adaptor.getDnTensor(), stream});
1313 assert(dims.size() == 1 &&
"Only 1D and 2D tensors are supported");
1314 destroyDnVecCallBuilder.create(loc, rewriter,
1315 {adaptor.getDnTensor(), stream});
1317 rewriter.replaceOp(op, {stream});
1321LogicalResult ConvertCreateCooOpToGpuRuntimeCallPattern::matchAndRewrite(
1322 gpu::CreateCooOp op, OpAdaptor adaptor,
1323 ConversionPatternRewriter &rewriter)
const {
1327 Location loc = op.getLoc();
1328 auto stream = adaptor.getAsyncDependencies().front();
1330 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1332 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1334 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1336 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1338 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1342 createCooCallBuilder
1343 .create(loc, rewriter,
1344 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1345 pRowIdxs, pColIdxs, pValues, itp, dtp, stream})
1347 rewriter.replaceOp(op, {handle, stream});
1351LogicalResult ConvertCreateCooAoSOpToGpuRuntimeCallPattern::matchAndRewrite(
1352 gpu::CreateCooAoSOp op, OpAdaptor adaptor,
1353 ConversionPatternRewriter &rewriter)
const {
1357 Location loc = op.getLoc();
1358 auto stream = adaptor.getAsyncDependencies().front();
1359 Value pIdxs = MemRefDescriptor(adaptor.getIdxs()).allocatedPtr(rewriter, loc);
1361 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1362 Type iType = llvm::cast<MemRefType>(op.getIdxs().getType()).getElementType();
1364 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1368 createCooAoSCallBuilder
1369 .create(loc, rewriter,
1370 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1371 pIdxs, pValues, itp, dtp, stream})
1373 rewriter.replaceOp(op, {handle, stream});
1377LogicalResult ConvertCreateCsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1378 gpu::CreateCsrOp op, OpAdaptor adaptor,
1379 ConversionPatternRewriter &rewriter)
const {
1383 Location loc = op.getLoc();
1384 auto stream = adaptor.getAsyncDependencies().front();
1386 MemRefDescriptor(adaptor.getRowPos()).allocatedPtr(rewriter, loc);
1388 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1390 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1392 llvm::cast<MemRefType>(op.getRowPos().getType()).getElementType();
1394 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1396 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1401 createCsrCallBuilder
1402 .create(loc, rewriter,
1403 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1404 pRowPos, pColIdxs, pValues, ptp, itp, dtp, stream})
1406 rewriter.replaceOp(op, {handle, stream});
1410LogicalResult ConvertCreate2To4SpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1411 gpu::Create2To4SpMatOp op, OpAdaptor adaptor,
1412 ConversionPatternRewriter &rewriter)
const {
1416 Location loc = op.getLoc();
1417 auto stream = adaptor.getAsyncDependencies().front();
1419 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1421 llvm::cast<MemRefType>(op.getMemref().getType()).getElementType();
1426 Value handle = LLVM::AllocaOp::create(
1427 rewriter, loc, llvmPointerType, llvmInt8Type, handleSz, 16);
1428 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1430 create2To4SpMatCallBuilder
1431 .create(loc, rewriter,
1432 {handle, adaptor.getRows(), adaptor.getCols(), pMat, dtp, stream})
1434 rewriter.replaceOp(op, {handle, stream});
1438LogicalResult ConvertDestroySpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1439 gpu::DestroySpMatOp op, OpAdaptor adaptor,
1440 ConversionPatternRewriter &rewriter)
const {
1444 Location loc = op.getLoc();
1445 auto stream = adaptor.getAsyncDependencies().front();
1448 destroyCuSparseLtSpMatBuilder.create(loc, rewriter,
1449 {adaptor.getSpmat(), stream});
1452 destroySpMatCallBuilder.create(loc, rewriter, {adaptor.getSpmat(), stream});
1454 rewriter.replaceOp(op, {stream});
1458LogicalResult ConvertSpMVBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1459 gpu::SpMVBufferSizeOp op, OpAdaptor adaptor,
1460 ConversionPatternRewriter &rewriter)
const {
1464 Location loc = op.getLoc();
1468 auto stream = adaptor.getAsyncDependencies().front();
1469 auto bufferSize = spMVBufferSizeCallBuilder
1470 .create(loc, rewriter,
1471 {modeA, adaptor.getSpmatA(), adaptor.getDnX(),
1472 adaptor.getDnY(), computeType, stream})
1474 rewriter.replaceOp(op, {bufferSize, stream});
1478LogicalResult ConvertSpMVOpToGpuRuntimeCallPattern::matchAndRewrite(
1479 gpu::SpMVOp op, OpAdaptor adaptor,
1480 ConversionPatternRewriter &rewriter)
const {
1484 Location loc = op.getLoc();
1488 auto stream = adaptor.getAsyncDependencies().front();
1490 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1491 spMVCallBuilder.create(loc, rewriter,
1492 {modeA, adaptor.getSpmatA(), adaptor.getDnX(),
1493 adaptor.getDnY(), computeType, pBuf, stream});
1494 rewriter.replaceOp(op, {stream});
1498LogicalResult ConvertSpMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1499 gpu::SpMMBufferSizeOp op, OpAdaptor adaptor,
1500 ConversionPatternRewriter &rewriter)
const {
1504 Location loc = op.getLoc();
1507 auto stream = adaptor.getAsyncDependencies().front();
1516 LLVM::AllocaOp::create(rewriter, loc, llvmPointerType, llvmPointerType,
1518 createCuSparseLtSpMMBufferSizeBuilder
1519 .create(loc, rewriter,
1520 {bufferSize, modeA, modeB, adaptor.getSpmatA(),
1521 adaptor.getDnmatB(), adaptor.getDnmatC(), computeType,
1525 auto bufferSizePtr1 = LLVM::GEPOp::create(
1526 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1528 auto bufferSizePtr2 = LLVM::GEPOp::create(
1529 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1532 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSize);
1534 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr1);
1536 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr2);
1538 rewriter.replaceOp(op, {bufferSize0, bufferSize1, bufferSize2, stream});
1543 createSpMMBufferSizeCallBuilder
1544 .create(loc, rewriter,
1545 {modeA, modeB, adaptor.getSpmatA(), adaptor.getDnmatB(),
1546 adaptor.getDnmatC(), computeType, stream})
1548 rewriter.replaceOp(op, {bufferSize, stream});
1553LogicalResult ConvertSDDMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1554 gpu::SDDMMBufferSizeOp op, OpAdaptor adaptor,
1555 ConversionPatternRewriter &rewriter)
const {
1559 Location loc = op.getLoc();
1564 auto stream = adaptor.getAsyncDependencies().front();
1566 createSDDMMBufferSizeCallBuilder
1567 .create(loc, rewriter,
1568 {modeA, modeB, adaptor.getDnmatA(), adaptor.getDnmatB(),
1569 adaptor.getSpmatC(), computeType, stream})
1571 rewriter.replaceOp(op, {bufferSize, stream});
1575LogicalResult ConvertSpMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1576 gpu::SpMMOp op, OpAdaptor adaptor,
1577 ConversionPatternRewriter &rewriter)
const {
1581 Location loc = op.getLoc();
1587 auto stream = adaptor.getAsyncDependencies().front();
1591 SmallVector<Value> pBufs;
1592 for (Value buffer : adaptor.getBuffers()) {
1593 Value pBuf = MemRefDescriptor(buffer).allocatedPtr(rewriter, loc);
1594 pBufs.push_back(pBuf);
1596 createCuSparseLtSpMMBuilder.create(
1598 {adaptor.getSpmatA(), adaptor.getDnmatB(), adaptor.getDnmatC(),
1599 pBufs[0], pBufs[1], pBufs[2], stream});
1601 Value pBuf = MemRefDescriptor(adaptor.getBuffers().front())
1602 .allocatedPtr(rewriter, loc);
1603 createSpMMCallBuilder.create(loc, rewriter,
1604 {modeA, modeB, adaptor.getSpmatA(),
1605 adaptor.getDnmatB(), adaptor.getDnmatC(),
1606 computeType, pBuf, stream});
1608 rewriter.replaceOp(op, {stream});
1612template <
typename T>
1614 converter.addConversion([&converter](T) ->
Type {
1615 return LLVM::LLVMPointerType::get(&converter.
getContext());
1619LogicalResult ConvertSDDMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1620 gpu::SDDMMOp op, OpAdaptor adaptor,
1621 ConversionPatternRewriter &rewriter)
const {
1625 Location loc = op.getLoc();
1630 auto stream = adaptor.getAsyncDependencies().front();
1632 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1633 createSDDMMCallBuilder.create(loc, rewriter,
1634 {modeA, modeB, adaptor.getDnmatA(),
1635 adaptor.getDnmatB(), adaptor.getSpmatC(),
1636 computeType, pBuf, stream});
1637 rewriter.replaceOp(op, {stream});
1642ConvertSpGEMMCreateDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1643 gpu::SpGEMMCreateDescrOp op, OpAdaptor adaptor,
1644 ConversionPatternRewriter &rewriter)
const {
1648 Location loc = op.getLoc();
1649 auto stream = adaptor.getAsyncDependencies().front();
1650 Value descr = createSpGEMMCreateDescrBuilder.create(loc, rewriter, {stream})
1652 rewriter.replaceOp(op, {descr, stream});
1657ConvertSpGEMMDestroyDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1658 gpu::SpGEMMDestroyDescrOp op, OpAdaptor adaptor,
1659 ConversionPatternRewriter &rewriter)
const {
1663 Location loc = op.getLoc();
1664 auto stream = adaptor.getAsyncDependencies().front();
1665 createSpGEMMDestroyDescrBuilder.create(loc, rewriter,
1666 {adaptor.getDesc(), stream});
1667 rewriter.replaceOp(op, {stream});
1672ConvertSpGEMMWorkEstimationOrComputeOpToGpuRuntimeCallPattern::matchAndRewrite(
1673 gpu::SpGEMMWorkEstimationOrComputeOp op, OpAdaptor adaptor,
1674 ConversionPatternRewriter &rewriter)
const {
1678 Location loc = op.getLoc();
1683 auto stream = adaptor.getAsyncDependencies().front();
1686 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1687 Value bufferSizeNew;
1689 if (adaptor.getKind() ==
1690 gpu::SpGEMMWorkEstimationOrComputeKind::WORK_ESTIMATION) {
1692 createSpGEMMWorkEstimationBuilder
1693 .create(loc, rewriter,
1694 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1695 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1696 adaptor.getBufferSz(), pBuf, stream})
1700 createSpGEMMComputeBuilder
1701 .create(loc, rewriter,
1702 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1703 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1704 adaptor.getBufferSz(), pBuf, stream})
1707 rewriter.replaceOp(op, {bufferSizeNew, stream});
1711LogicalResult ConvertSpGEMMCopyOpToGpuRuntimeCallPattern::matchAndRewrite(
1712 gpu::SpGEMMCopyOp op, OpAdaptor adaptor,
1713 ConversionPatternRewriter &rewriter)
const {
1717 Location loc = op.getLoc();
1722 auto stream = adaptor.getAsyncDependencies().front();
1723 createSpGEMMCopyBuilder.create(loc, rewriter,
1724 {adaptor.getDesc(), modeA, modeB,
1725 adaptor.getSpmatA(), adaptor.getSpmatB(),
1726 adaptor.getSpmatC(), computeType, stream});
1727 rewriter.replaceOp(op, {stream});
1731LogicalResult ConvertSpMatGetSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1732 gpu::SpMatGetSizeOp op, OpAdaptor adaptor,
1733 ConversionPatternRewriter &rewriter)
const {
1737 Location loc = op.getLoc();
1738 auto stream = adaptor.getAsyncDependencies().front();
1741 auto buffer = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1742 llvmInt64Type, three, 16);
1744 auto rowsPtr = LLVM::GEPOp::create(
1745 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1747 auto colsPtr = LLVM::GEPOp::create(
1748 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1750 auto nnzsPtr = LLVM::GEPOp::create(
1751 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1753 createSpMatGetSizeBuilder.create(
1754 loc, rewriter, {adaptor.getSpmat(), rowsPtr, colsPtr, nnzsPtr, stream});
1755 auto rows = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, rowsPtr);
1756 auto cols = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, colsPtr);
1757 auto nnzs = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, nnzsPtr);
1759 rewriter.replaceOp(op, {rows, cols, nnzs, stream});
1763LogicalResult ConvertSetCsrPointersOpToGpuRuntimeCallPattern::matchAndRewrite(
1764 gpu::SetCsrPointersOp op, OpAdaptor adaptor,
1765 ConversionPatternRewriter &rewriter)
const {
1769 Location loc = op.getLoc();
1770 auto stream = adaptor.getAsyncDependencies().front();
1772 MemRefDescriptor(adaptor.getPositions()).allocatedPtr(rewriter, loc);
1774 MemRefDescriptor(adaptor.getCoordinates()).allocatedPtr(rewriter, loc);
1776 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1777 createSetCsrPointersBuilder.create(
1778 loc, rewriter, {adaptor.getSpmat(), pPos, pCrd, pVal, stream});
1779 rewriter.replaceOp(op, {stream});
1783LogicalResult ConvertCreateCscOpToGpuRuntimeCallPattern::matchAndRewrite(
1784 gpu::CreateCscOp op, OpAdaptor adaptor,
1785 ConversionPatternRewriter &rewriter)
const {
1789 Location loc = op.getLoc();
1790 auto stream = adaptor.getAsyncDependencies().front();
1792 MemRefDescriptor(adaptor.getColPos()).allocatedPtr(rewriter, loc);
1794 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1796 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1798 llvm::cast<MemRefType>(op.getColPos().getType()).getElementType();
1800 llvm::cast<MemRefType>(op.getRowIdxs().getType()).getElementType();
1802 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1807 createCscCallBuilder
1808 .create(loc, rewriter,
1809 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1810 pColPos, pRowIdxs, pValues, ptp, itp, dtp, stream})
1812 rewriter.replaceOp(op, {handle, stream});
1816LogicalResult ConvertCreateBsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1817 gpu::CreateBsrOp op, OpAdaptor adaptor,
1818 ConversionPatternRewriter &rewriter)
const {
1822 Location loc = op.getLoc();
1823 auto stream = adaptor.getAsyncDependencies().front();
1825 MemRefDescriptor(adaptor.getBRowPos()).allocatedPtr(rewriter, loc);
1827 MemRefDescriptor(adaptor.getBColIdxs()).allocatedPtr(rewriter, loc);
1829 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1831 llvm::cast<MemRefType>(op.getBRowPos().getType()).getElementType();
1833 llvm::cast<MemRefType>(op.getBColIdxs().getType()).getElementType();
1835 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1840 createBsrCallBuilder
1841 .create(loc, rewriter,
1842 {adaptor.getBrows(), adaptor.getBcols(), adaptor.getBnnz(),
1843 adaptor.getRBlockSize(), adaptor.getCBlockSize(), pRowPos,
1844 pColIdxs, pValues, ptp, itp, dtp, stream})
1846 rewriter.replaceOp(op, {handle, stream});
1852 bool kernelBarePtrCallConv,
bool kernelIntersperseSizeCallConv) {
1862 patterns.
add<ConvertAsyncYieldToGpuRuntimeCallPattern>(converter,
1865 patterns.
add<ConvertAllocOpToGpuRuntimeCallPattern,
1866 ConvertDeallocOpToGpuRuntimeCallPattern,
1867 ConvertHostRegisterOpToGpuRuntimeCallPattern,
1868 ConvertHostUnregisterOpToGpuRuntimeCallPattern,
1869 ConvertMemcpyOpToGpuRuntimeCallPattern,
1870 ConvertMemsetOpToGpuRuntimeCallPattern,
1871 ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern,
1872 ConvertWaitAsyncOpToGpuRuntimeCallPattern,
1873 ConvertWaitOpToGpuRuntimeCallPattern,
1874 ConvertCreateDnTensorOpToGpuRuntimeCallPattern,
1875 ConvertDestroyDnTensorOpToGpuRuntimeCallPattern,
1876 ConvertCreateCooOpToGpuRuntimeCallPattern,
1877 ConvertCreateCooAoSOpToGpuRuntimeCallPattern,
1878 ConvertCreateCsrOpToGpuRuntimeCallPattern,
1879 ConvertCreateCscOpToGpuRuntimeCallPattern,
1880 ConvertCreateBsrOpToGpuRuntimeCallPattern,
1881 ConvertCreate2To4SpMatOpToGpuRuntimeCallPattern,
1882 ConvertDestroySpMatOpToGpuRuntimeCallPattern,
1883 ConvertSpMVBufferSizeOpToGpuRuntimeCallPattern,
1884 ConvertSpMVOpToGpuRuntimeCallPattern,
1885 ConvertSpMMBufferSizeOpToGpuRuntimeCallPattern,
1886 ConvertSDDMMBufferSizeOpToGpuRuntimeCallPattern,
1887 ConvertSpMMOpToGpuRuntimeCallPattern,
1888 ConvertSDDMMOpToGpuRuntimeCallPattern,
1889 ConvertSpGEMMCreateDescrOpToGpuRuntimeCallPattern,
1890 ConvertSpGEMMDestroyDescrOpToGpuRuntimeCallPattern,
1891 ConvertSpGEMMWorkEstimationOrComputeOpToGpuRuntimeCallPattern,
1892 ConvertSpGEMMCopyOpToGpuRuntimeCallPattern,
1893 ConvertSpMatGetSizeOpToGpuRuntimeCallPattern,
1894 ConvertSetCsrPointersOpToGpuRuntimeCallPattern>(converter);
1895 patterns.
add<LegalizeLaunchFuncOpPattern>(converter, kernelBarePtrCallConv,
1896 kernelIntersperseSizeCallConv);
1904struct GPUModuleOpConvertToLLVMInterface
1905 :
public ConvertToLLVMOpInterface::ExternalModel<
1906 GPUModuleOpConvertToLLVMInterface, gpu::GPUModuleOp> {
1908 void getConvertToLLVMConversionAttrs(
1913void GPUModuleOpConvertToLLVMInterface::getConvertToLLVMConversionAttrs(
1914 Operation *op, SmallVectorImpl<ConvertToLLVMAttrInterface> &attrs)
const {
1915 auto module = cast<gpu::GPUModuleOp>(op);
1916 ArrayAttr targetsAttr =
module.getTargetsAttr();
1918 if (!targetsAttr || targetsAttr.size() != 1)
1920 if (
auto patternAttr = dyn_cast<ConvertToLLVMAttrInterface>(targetsAttr[0]))
1921 attrs.push_back(patternAttr);
1926 gpu::GPUModuleOp::attachInterface<GPUModuleOpConvertToLLVMInterface>(*ctx);
static void addOpaquePointerConversion(LLVMTypeConverter &converter)
static Value genConstFloat32From(OpBuilder &builder, Location loc, T tValue)
static int32_t getCuSparseDataTypeFrom(Type type)
static LogicalResult areAllLLVMTypes(Operation *op, ValueRange operands, ConversionPatternRewriter &rewriter)
static Value genConstInt32From(OpBuilder &builder, Location loc, T tValue)
static gpu::Prune2To4SpMatFlag get2To4PruneFlag(Value spMat)
static bool isGpuAsyncTokenType(Value value)
#define DECLARE_CONVERT_OP_TO_GPU_RUNTIME_CALL_PATTERN(op_name)
Generic rewriting rule for operation on sparse matrices.
static int32_t getCuSparseLtDataTypeFrom(Type type)
static bool isDefinedByCallTo(Value value, StringRef functionName)
static Value bitAndAddrspaceCast(Location loc, ConversionPatternRewriter &rewriter, LLVM::LLVMPointerType destinationType, Value sourcePtr, const LLVMTypeConverter &typeConverter)
static bool isSpMMCusparseLtOp(Value op)
static int32_t getCuSparseIndexTypeFrom(Type type)
static bool is2To4Sparsity(Value spMat)
static LogicalResult isAsyncWithOneDependency(ConversionPatternRewriter &rewriter, gpu::AsyncOpInterface op)
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
static llvm::Value * getSizeInBytes(DataLayout &dl, const mlir::Type &type, Operation *clauseOp, llvm::Value *basePointer, llvm::Type *baseType, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::ManagedStatic< PassManagerOptions > options
IntegerType getIntegerType(unsigned width)
FloatAttr getF32FloatAttr(float value)
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Type getIndexType() const
Gets the MLIR type wrapping the LLVM integer type whose bit width is defined by the used type convert...
static Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Create a constant Op producing a value of resultType from an index-typed integer attribute.
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.
MLIRContext & getContext() const
Returns the MLIR context.
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.
std::vector< Dialect * > getLoadedDialects()
Return information about all IR dialects loaded in the context.
Value size(OpBuilder &builder, Location loc, unsigned pos)
Builds IR extracting the pos-th size from the descriptor.
This class helps build Operations.
static OpBuilder atBlockEnd(Block *block, Listener *listener=nullptr)
Create a builder and set the insertion point to after the last operation in the block but still insid...
Operation is the basic unit of execution within MLIR.
Location getLoc()
The source location the operation was defined or derived from.
void print(raw_ostream &os, const OpPrintingFlags &flags={})
operand_range getOperands()
Returns an iterator on the underlying Value's.
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 isInteger() const
Return true if this is an integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Type getType() const
Return the type of this value.
user_range getUsers() const
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Creates an llvm.mlir.constant producing value as resultType, which is expected to be the converted in...
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
void registerConvertGpuToLLVMInterface(DialectRegistry ®istry)
Registers the ConvertToLLVMOpInterface interface on the gpu::GPUModuleOP operation.
void populateVectorTransferLoweringPatterns(RewritePatternSet &patterns, std::optional< unsigned > maxTransferRank=std::nullopt, PatternBenefit benefit=1)
Populate the pattern set with the following patterns:
Include the generated interface declarations.
void populateVectorToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, bool reassociateFPReductions=false, bool force32BitVectorIndices=false, bool useVectorAlignment=false, bool enableGEPInboundsNuw=false)
Collect a set of patterns to convert from the Vector dialect to LLVM.
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
void populateFinalizeMemRefToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, SymbolTableCollection *symbolTables=nullptr)
Collect a set of patterns to convert memory-related operations from the MemRef dialect to the LLVM di...
void populateGpuToLLVMConversionPatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns, bool kernelBarePtrCallConv=false, bool kernelIntersperseSizeCallConv=false)
Collect a set of patterns to convert from the GPU dialect to LLVM and populate converter for gpu type...
void registerConvertToLLVMDependentDialectLoading(DialectRegistry ®istry)
Register the extension that will load dependent dialects for LLVM conversion.
void populateAsyncStructuralTypeConversionsAndLegality(TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target)
Populates patterns for async structural type conversions.
LLVM::LLVMFunctionType functionType
LLVM::CallOp create(Location loc, OpBuilder &builder, ArrayRef< Value > arguments) const