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);
554 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects))
555 iface->populateConvertToLLVMConversionPatterns(
target, converter, patterns);
559 target.addLegalOp<gpu::GPUModuleOp, gpu::BinaryOp>();
561 target.addDynamicallyLegalOp<gpu::LaunchFuncOp>(
562 [&](gpu::LaunchFuncOp op) ->
bool {
return converter.isLegal(op); });
570 kernelBarePtrCallConv,
571 kernelIntersperseSizeCallConv);
574 applyPartialConversion(getOperation(),
target, std::move(patterns))))
580 auto module = builder.getBlock()->getParent()->getParentOfType<ModuleOp>();
581 auto function = [&] {
582 if (
auto function = module.lookupSymbol<LLVM::LLVMFuncOp>(
functionName))
587 return LLVM::CallOp::create(builder, loc, function, arguments);
604 llvm_unreachable(
"unsupported type");
610 if (llvm::isa<ComplexType>(type)) {
612 auto elementType = cast<ComplexType>(type).getElementType();
613 if (elementType.isBF16())
615 if (elementType.isF16())
617 if (elementType.isF32())
619 if (elementType.isF64())
621 if (elementType.isInteger(8))
623 if (elementType.isInteger(16))
625 if (elementType.isInteger(32))
643 llvm_unreachable(
"unsupported element type");
647 return spMat.
getDefiningOp<gpu::Create2To4SpMatOp>().getPruneFlag();
672 llvm_unreachable(
"cannot find spmat def");
677 auto spmmOp = dyn_cast<gpu::SpMMOp>(user);
689 ConversionPatternRewriter &rewriter) {
690 if (!llvm::all_of(operands, [](
Value value) {
693 return rewriter.notifyMatchFailure(
694 op,
"Cannot convert if operands aren't of LLVM type.");
700 gpu::AsyncOpInterface op) {
701 if (op.getAsyncDependencies().size() != 1)
702 return rewriter.notifyMatchFailure(
703 op,
"Can only convert with exactly one async dependency.");
705 if (!op.getAsyncToken())
706 return rewriter.notifyMatchFailure(op,
"Can convert only async version.");
711LogicalResult ConvertHostRegisterOpToGpuRuntimeCallPattern::matchAndRewrite(
712 gpu::HostRegisterOp hostRegisterOp, OpAdaptor adaptor,
713 ConversionPatternRewriter &rewriter)
const {
714 auto *op = hostRegisterOp.getOperation();
718 Location loc = op->getLoc();
720 auto memRefType = hostRegisterOp.getValue().getType();
721 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
724 auto arguments = getTypeConverter()->promoteOperands(
725 loc, op->getOperands(), adaptor.getOperands(), rewriter);
726 arguments.push_back(elementSize);
727 hostRegisterCallBuilder.create(loc, rewriter, arguments);
729 rewriter.eraseOp(op);
733LogicalResult ConvertHostUnregisterOpToGpuRuntimeCallPattern::matchAndRewrite(
734 gpu::HostUnregisterOp hostUnregisterOp, OpAdaptor adaptor,
735 ConversionPatternRewriter &rewriter)
const {
736 Operation *op = hostUnregisterOp.getOperation();
740 Location loc = op->
getLoc();
742 auto memRefType = hostUnregisterOp.getValue().getType();
743 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
746 auto arguments = getTypeConverter()->promoteOperands(
747 loc, op->
getOperands(), adaptor.getOperands(), rewriter);
748 arguments.push_back(elementSize);
749 hostUnregisterCallBuilder.create(loc, rewriter, arguments);
751 rewriter.eraseOp(op);
755LogicalResult ConvertAllocOpToGpuRuntimeCallPattern::matchAndRewrite(
756 gpu::AllocOp allocOp, OpAdaptor adaptor,
757 ConversionPatternRewriter &rewriter)
const {
759 MemRefType memRefType = allocOp.getType();
762 !isConvertibleAndHasIdentityMaps(memRefType))
765 auto loc = allocOp.getLoc();
767 bool isShared = allocOp.getHostShared();
769 if (isShared && allocOp.getAsyncToken())
770 return rewriter.notifyMatchFailure(
771 allocOp,
"Host Shared allocation cannot be done async");
772 if (adaptor.getAsyncDependencies().size() > 1)
773 return rewriter.notifyMatchFailure(
774 allocOp,
"Can convert with at most one async dependency.");
778 SmallVector<Value, 4> shape;
779 SmallVector<Value, 4> strides;
781 getMemRefDescriptorSizes(loc, memRefType, adaptor.getDynamicSizes(), rewriter,
782 shape, strides, sizeBytes);
786 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
787 Value stream = adaptor.getAsyncDependencies().empty()
789 : adaptor.getAsyncDependencies().front();
791 auto isHostShared = mlir::LLVM::ConstantOp::create(
792 rewriter, loc, llvmInt8Type, rewriter.getI8IntegerAttr(isShared));
795 allocCallBuilder.create(loc, rewriter, {sizeBytes, stream, isHostShared})
800 unsigned dstAddrSpace = memRefType.getMemorySpaceAsInt();
801 unsigned srcAddrSpace =
802 cast<LLVM::LLVMPointerType>(allocatedPtr.
getType()).getAddressSpace();
803 if (dstAddrSpace != srcAddrSpace) {
805 LLVM::LLVMPointerType::get(rewriter.getContext(), dstAddrSpace);
807 LLVM::AddrSpaceCastOp::create(rewriter, loc, targetPtrTy, allocatedPtr);
811 Value alignedPtr = allocatedPtr;
814 auto memRefDescriptor = this->createMemRefDescriptor(
815 loc, memRefType, allocatedPtr, alignedPtr, shape, strides, rewriter);
817 if (allocOp.getAsyncToken()) {
819 rewriter.replaceOp(allocOp, {memRefDescriptor, stream});
821 rewriter.replaceOp(allocOp, {memRefDescriptor});
827LogicalResult ConvertDeallocOpToGpuRuntimeCallPattern::matchAndRewrite(
828 gpu::DeallocOp deallocOp, OpAdaptor adaptor,
829 ConversionPatternRewriter &rewriter)
const {
832 if (adaptor.getAsyncDependencies().size() > 1)
833 return rewriter.notifyMatchFailure(
834 deallocOp,
"Can convert with at most one async dependency.");
836 Location loc = deallocOp.getLoc();
839 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
840 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
841 Value stream = adaptor.getAsyncDependencies().empty()
843 : adaptor.getAsyncDependencies().front();
844 deallocCallBuilder.create(loc, rewriter, {pointer, stream});
846 if (deallocOp.getAsyncToken()) {
848 rewriter.replaceOp(deallocOp, {stream});
851 rewriter.eraseOp(deallocOp);
857 return isa<gpu::AsyncTokenType>(value.
getType());
872LogicalResult ConvertAsyncYieldToGpuRuntimeCallPattern::matchAndRewrite(
873 async::YieldOp yieldOp, OpAdaptor adaptor,
874 ConversionPatternRewriter &rewriter)
const {
876 return rewriter.notifyMatchFailure(yieldOp,
"no gpu async token operand");
878 Location loc = yieldOp.getLoc();
879 SmallVector<Value, 4> newOperands(adaptor.getOperands());
880 llvm::SmallDenseSet<Value> streams;
881 for (
auto &operand : yieldOp->getOpOperands()) {
884 auto idx = operand.getOperandNumber();
885 auto stream = adaptor.getOperands()[idx];
886 auto event = eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
887 eventRecordCallBuilder.create(loc, rewriter, {event, stream});
888 newOperands[idx] = event;
889 streams.insert(stream);
891 for (
auto stream : streams)
892 streamDestroyCallBuilder.create(loc, rewriter, {stream});
894 rewriter.modifyOpInPlace(yieldOp, [&] { yieldOp->setOperands(newOperands); });
900 assert(isa<LLVM::LLVMPointerType>(value.
getType()));
902 return *defOp.getCallee() == functionName;
910LogicalResult ConvertWaitOpToGpuRuntimeCallPattern::matchAndRewrite(
911 gpu::WaitOp waitOp, OpAdaptor adaptor,
912 ConversionPatternRewriter &rewriter)
const {
913 if (waitOp.getAsyncToken())
914 return rewriter.notifyMatchFailure(waitOp,
"Cannot convert async op.");
916 Location loc = waitOp.getLoc();
918 for (
auto operand : adaptor.getOperands()) {
921 streamSynchronizeCallBuilder.create(loc, rewriter, {operand});
922 streamDestroyCallBuilder.create(loc, rewriter, {operand});
926 eventSynchronizeCallBuilder.create(loc, rewriter, {operand});
927 eventDestroyCallBuilder.create(loc, rewriter, {operand});
931 rewriter.eraseOp(waitOp);
940LogicalResult ConvertWaitAsyncOpToGpuRuntimeCallPattern::matchAndRewrite(
941 gpu::WaitOp waitOp, OpAdaptor adaptor,
942 ConversionPatternRewriter &rewriter)
const {
943 if (!waitOp.getAsyncToken())
944 return rewriter.notifyMatchFailure(waitOp,
"Can only convert async op.");
946 Location loc = waitOp.getLoc();
948 auto insertionPoint = rewriter.saveInsertionPoint();
949 SmallVector<Value, 1> events;
951 llvm::zip(waitOp.getAsyncDependencies(), adaptor.getOperands())) {
952 auto operand = std::get<1>(pair);
956 auto *defOp = std::get<0>(pair).getDefiningOp();
957 rewriter.setInsertionPointAfter(defOp);
958 auto event = eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
959 eventRecordCallBuilder.create(loc, rewriter, {event, operand});
960 events.push_back(event);
964 events.push_back(operand);
967 rewriter.restoreInsertionPoint(insertionPoint);
968 auto stream = streamCreateCallBuilder.create(loc, rewriter, {}).getResult();
969 for (
auto event : events)
970 streamWaitEventCallBuilder.create(loc, rewriter, {stream,
event});
971 for (
auto event : events)
972 eventDestroyCallBuilder.create(loc, rewriter, {
event});
973 rewriter.replaceOp(waitOp, {stream});
979LogicalResult LegalizeLaunchFuncOpPattern::matchAndRewrite(
980 gpu::LaunchFuncOp launchOp, OpAdaptor adaptor,
981 ConversionPatternRewriter &rewriter)
const {
988 if (!launchOp.getAsyncToken() && !launchOp.getAsyncDependencies().empty())
989 return rewriter.notifyMatchFailure(
990 launchOp,
"Cannot convert non-async op with async dependencies.");
992 Location loc = launchOp.getLoc();
994 Value stream = Value();
995 if (!adaptor.getAsyncDependencies().empty()) {
996 stream = adaptor.getAsyncDependencies().front();
999 if (adaptor.getAsyncDependencies().size() > 1) {
1000 auto insertionPoint = rewriter.saveInsertionPoint();
1001 SmallVector<Value, 4> events;
1002 for (
auto [origDep, convertedDep] :
1003 llvm::zip(launchOp.getAsyncDependencies().drop_front(),
1004 adaptor.getAsyncDependencies().drop_front())) {
1006 streamCreateCallBuilder.functionName)) {
1007 events.push_back(convertedDep);
1010 Operation *defOp = origDep.getDefiningOp();
1011 rewriter.setInsertionPointAfter(defOp);
1013 eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
1014 eventRecordCallBuilder.create(loc, rewriter, {event, convertedDep});
1015 events.push_back(event);
1017 rewriter.restoreInsertionPoint(insertionPoint);
1018 for (Value event : events)
1019 streamWaitEventCallBuilder.create(loc, rewriter, {stream,
event});
1020 for (Value event : events)
1021 eventDestroyCallBuilder.create(loc, rewriter, {
event});
1026 else if (launchOp.getAsyncToken())
1027 stream = streamCreateCallBuilder.create(loc, rewriter, {}).getResult();
1032 OperandRange origArguments = launchOp.getKernelOperands();
1033 bool effectiveBarePtr = kernelBarePtrCallConv ||
1034 getTypeConverter()->getOptions().useBarePtrCallConv;
1035 if (effectiveBarePtr) {
1036 for (Value arg : origArguments) {
1037 if (isa<UnrankedMemRefType>(arg.getType()))
1038 return rewriter.notifyMatchFailure(
1039 loc,
"unranked memref kernel argument is not supported with "
1040 "the bare-pointer calling convention");
1043 SmallVector<Value, 8> llvmArguments = getTypeConverter()->promoteOperands(
1044 loc, origArguments, adaptor.getKernelOperands(), rewriter,
1045 kernelBarePtrCallConv);
1046 SmallVector<Value, 8> llvmArgumentsWithSizes;
1049 if (kernelIntersperseSizeCallConv) {
1050 if (origArguments.size() != llvmArguments.size()) {
1052 return rewriter.notifyMatchFailure(
1054 "Cannot add sizes to arguments with one-to-many LLVM IR expansion.");
1057 llvmArgumentsWithSizes.reserve(llvmArguments.size() * 2);
1058 for (
auto [llvmArg, origArg] : zip_equal(llvmArguments, origArguments)) {
1059 auto memrefTy = dyn_cast<MemRefType>(origArg.getType());
1061 return rewriter.notifyMatchFailure(
1062 launchOp,
"Operand to launch op is not a memref.");
1065 if (!memrefTy.hasStaticShape() ||
1066 !memrefTy.getElementType().isIntOrFloat()) {
1067 return rewriter.notifyMatchFailure(
1068 launchOp,
"Operand to launch op is not a memref with a static "
1069 "shape and an integer or float element type.");
1072 unsigned bitwidth = memrefTy.getElementTypeBitWidth();
1073 if (bitwidth % 8 != 0) {
1074 return rewriter.notifyMatchFailure(
1075 launchOp,
"Operand to launch op is not a memref with a "
1076 "byte-aligned element type.");
1079 uint64_t staticSize =
static_cast<uint64_t
>(bitwidth / 8) *
1080 static_cast<uint64_t
>(memrefTy.getNumElements());
1084 llvmArgumentsWithSizes.push_back(llvmArg);
1085 llvmArgumentsWithSizes.push_back(sizeArg);
1089 std::optional<gpu::KernelDim3> clusterSize = std::nullopt;
1090 if (launchOp.hasClusterSize()) {
1092 gpu::KernelDim3{adaptor.getClusterSizeX(), adaptor.getClusterSizeY(),
1093 adaptor.getClusterSizeZ()};
1095 auto newLaunchOp = gpu::LaunchFuncOp::create(
1096 rewriter, launchOp.getLoc(), launchOp.getKernelAttr(),
1097 gpu::KernelDim3{adaptor.getGridSizeX(), adaptor.getGridSizeY(),
1098 adaptor.getGridSizeZ()},
1099 gpu::KernelDim3{adaptor.getBlockSizeX(), adaptor.getBlockSizeY(),
1100 adaptor.getBlockSizeZ()},
1101 adaptor.getDynamicSharedMemorySize(),
1102 llvmArgumentsWithSizes.empty() ? llvmArguments : llvmArgumentsWithSizes,
1103 nullptr, {}, stream, clusterSize);
1104 if (launchOp.getCooperative())
1105 newLaunchOp.setCooperative(
true);
1106 if (launchOp.getAsyncToken())
1107 rewriter.replaceOp(launchOp, {stream});
1109 rewriter.eraseOp(launchOp);
1114 ConversionPatternRewriter &rewriter,
1115 LLVM::LLVMPointerType destinationType,
1118 auto sourceTy = cast<LLVM::LLVMPointerType>(sourcePtr.
getType());
1119 if (destinationType.getAddressSpace() != sourceTy.getAddressSpace())
1120 sourcePtr = LLVM::AddrSpaceCastOp::create(
1122 LLVM::LLVMPointerType::get(rewriter.getContext(),
1123 destinationType.getAddressSpace()),
1128LogicalResult ConvertMemcpyOpToGpuRuntimeCallPattern::matchAndRewrite(
1129 gpu::MemcpyOp memcpyOp, OpAdaptor adaptor,
1130 ConversionPatternRewriter &rewriter)
const {
1131 auto memRefType = cast<MemRefType>(memcpyOp.getSrc().getType());
1134 !isConvertibleAndHasIdentityMaps(memRefType) ||
1138 auto loc = memcpyOp.getLoc();
1140 MemRefDescriptor srcDesc(adaptor.getSrc());
1141 Value numElements =
getNumElements(rewriter, loc, memRefType, srcDesc);
1143 Type elementPtrType = getElementPtrType(memRefType);
1144 Value nullPtr = LLVM::ZeroOp::create(rewriter, loc, elementPtrType);
1145 Value gepPtr = LLVM::GEPOp::create(
1146 rewriter, loc, elementPtrType,
1147 typeConverter->convertType(memRefType.getElementType()), nullPtr,
1150 LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), gepPtr);
1153 srcDesc.alignedPtr(rewriter, loc),
1154 *getTypeConverter());
1156 loc, rewriter, llvmPointerType,
1157 MemRefDescriptor(adaptor.getDst()).alignedPtr(rewriter, loc),
1158 *getTypeConverter());
1160 auto stream = adaptor.getAsyncDependencies().front();
1161 memcpyCallBuilder.create(loc, rewriter, {dst, src, sizeBytes, stream});
1163 rewriter.replaceOp(memcpyOp, {stream});
1168LogicalResult ConvertMemsetOpToGpuRuntimeCallPattern::matchAndRewrite(
1169 gpu::MemsetOp memsetOp, OpAdaptor adaptor,
1170 ConversionPatternRewriter &rewriter)
const {
1171 auto memRefType = cast<MemRefType>(memsetOp.getDst().getType());
1174 !isConvertibleAndHasIdentityMaps(memRefType) ||
1178 auto loc = memsetOp.getLoc();
1180 Type valueType = adaptor.getValue().getType();
1183 if (!valueType.
isIntOrFloat() || (bitWidth != 16 && bitWidth != 32)) {
1184 return rewriter.notifyMatchFailure(
1185 memsetOp,
"value must be a 16 or 32 bit int or float");
1189 Type bitCastType = valueTypeWidth == 32 ? llvmInt32Type : llvmInt16Type;
1191 MemRefDescriptor dstDesc(adaptor.getDst());
1192 Value numElements =
getNumElements(rewriter, loc, memRefType, dstDesc);
1195 LLVM::BitcastOp::create(rewriter, loc, bitCastType, adaptor.getValue());
1197 dstDesc.alignedPtr(rewriter, loc),
1198 *getTypeConverter());
1200 auto stream = adaptor.getAsyncDependencies().front();
1201 FunctionCallBuilder builder =
1202 valueTypeWidth == 32 ? memset32CallBuilder : memset16CallBuilder;
1203 builder.
create(loc, rewriter, {dst, value, numElements, stream});
1205 rewriter.replaceOp(memsetOp, {stream});
1209LogicalResult ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern::matchAndRewrite(
1210 gpu::SetDefaultDeviceOp op, OpAdaptor adaptor,
1211 ConversionPatternRewriter &rewriter)
const {
1212 Location loc = op.getLoc();
1213 auto call = setDefaultDeviceCallBuilder.create(loc, rewriter,
1214 {adaptor.getDevIndex()});
1215 rewriter.replaceOp(op, call);
1219template <
typename T>
1222 return LLVM::ConstantOp::create(builder, loc, llvmInt32Type,
1223 static_cast<int32_t
>(tValue));
1226LogicalResult ConvertCreateDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1227 gpu::CreateDnTensorOp op, OpAdaptor adaptor,
1228 ConversionPatternRewriter &rewriter)
const {
1232 Location loc = op.getLoc();
1233 auto stream = adaptor.getAsyncDependencies().front();
1235 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1236 Type dType = op.getMemref().
getType().getElementType();
1239 SmallVector<Value, 4> dims;
1240 for (Value dim : adaptor.getDims()) {
1241 dims.push_back(dim);
1251 if (dims.size() == 2) {
1255 handle = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1256 llvmInt8Type, handleSz, 16);
1257 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1259 createLtDnMatCallBuilder
1260 .create(loc, rewriter,
1261 {handle, dims[0], dims[1], pTensor, dtp, stream})
1265 createDnMatCallBuilder
1266 .create(loc, rewriter, {dims[0], dims[1], pTensor, dtp, stream})
1270 assert(dims.size() == 1 &&
"Only 1D and 2D tensors are supported");
1271 handle = createDnVecCallBuilder
1272 .create(loc, rewriter, {dims[0], pTensor, dtp, stream})
1275 rewriter.replaceOp(op, {handle, stream});
1279LogicalResult ConvertDestroyDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1280 gpu::DestroyDnTensorOp op, OpAdaptor adaptor,
1281 ConversionPatternRewriter &rewriter)
const {
1285 Location loc = op.getLoc();
1286 auto stream = adaptor.getAsyncDependencies().front();
1287 auto definingOp = op.getDnTensor().
getDefiningOp<gpu::CreateDnTensorOp>();
1288 SmallVector<Value, 4> dims;
1289 for (Value dim : definingOp.getDims()) {
1290 dims.push_back(dim);
1292 if (dims.size() == 2) {
1296 destroyCuSparseLtDnMatBuilder.create(loc, rewriter,
1297 {adaptor.getDnTensor(), stream});
1299 destroyDnMatCallBuilder.create(loc, rewriter,
1300 {adaptor.getDnTensor(), stream});
1303 assert(dims.size() == 1 &&
"Only 1D and 2D tensors are supported");
1304 destroyDnVecCallBuilder.create(loc, rewriter,
1305 {adaptor.getDnTensor(), stream});
1307 rewriter.replaceOp(op, {stream});
1311LogicalResult ConvertCreateCooOpToGpuRuntimeCallPattern::matchAndRewrite(
1312 gpu::CreateCooOp op, OpAdaptor adaptor,
1313 ConversionPatternRewriter &rewriter)
const {
1317 Location loc = op.getLoc();
1318 auto stream = adaptor.getAsyncDependencies().front();
1320 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1322 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1324 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1326 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1328 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1332 createCooCallBuilder
1333 .create(loc, rewriter,
1334 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1335 pRowIdxs, pColIdxs, pValues, itp, dtp, stream})
1337 rewriter.replaceOp(op, {handle, stream});
1341LogicalResult ConvertCreateCooAoSOpToGpuRuntimeCallPattern::matchAndRewrite(
1342 gpu::CreateCooAoSOp op, OpAdaptor adaptor,
1343 ConversionPatternRewriter &rewriter)
const {
1347 Location loc = op.getLoc();
1348 auto stream = adaptor.getAsyncDependencies().front();
1349 Value pIdxs = MemRefDescriptor(adaptor.getIdxs()).allocatedPtr(rewriter, loc);
1351 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1352 Type iType = llvm::cast<MemRefType>(op.getIdxs().getType()).getElementType();
1354 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1358 createCooAoSCallBuilder
1359 .create(loc, rewriter,
1360 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1361 pIdxs, pValues, itp, dtp, stream})
1363 rewriter.replaceOp(op, {handle, stream});
1367LogicalResult ConvertCreateCsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1368 gpu::CreateCsrOp op, OpAdaptor adaptor,
1369 ConversionPatternRewriter &rewriter)
const {
1373 Location loc = op.getLoc();
1374 auto stream = adaptor.getAsyncDependencies().front();
1376 MemRefDescriptor(adaptor.getRowPos()).allocatedPtr(rewriter, loc);
1378 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1380 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1382 llvm::cast<MemRefType>(op.getRowPos().getType()).getElementType();
1384 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1386 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1391 createCsrCallBuilder
1392 .create(loc, rewriter,
1393 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1394 pRowPos, pColIdxs, pValues, ptp, itp, dtp, stream})
1396 rewriter.replaceOp(op, {handle, stream});
1400LogicalResult ConvertCreate2To4SpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1401 gpu::Create2To4SpMatOp op, OpAdaptor adaptor,
1402 ConversionPatternRewriter &rewriter)
const {
1406 Location loc = op.getLoc();
1407 auto stream = adaptor.getAsyncDependencies().front();
1409 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1411 llvm::cast<MemRefType>(op.getMemref().getType()).getElementType();
1416 Value handle = LLVM::AllocaOp::create(
1417 rewriter, loc, llvmPointerType, llvmInt8Type, handleSz, 16);
1418 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1420 create2To4SpMatCallBuilder
1421 .create(loc, rewriter,
1422 {handle, adaptor.getRows(), adaptor.getCols(), pMat, dtp, stream})
1424 rewriter.replaceOp(op, {handle, stream});
1428LogicalResult ConvertDestroySpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1429 gpu::DestroySpMatOp op, OpAdaptor adaptor,
1430 ConversionPatternRewriter &rewriter)
const {
1434 Location loc = op.getLoc();
1435 auto stream = adaptor.getAsyncDependencies().front();
1438 destroyCuSparseLtSpMatBuilder.create(loc, rewriter,
1439 {adaptor.getSpmat(), stream});
1442 destroySpMatCallBuilder.create(loc, rewriter, {adaptor.getSpmat(), stream});
1444 rewriter.replaceOp(op, {stream});
1448LogicalResult ConvertSpMVBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1449 gpu::SpMVBufferSizeOp op, OpAdaptor adaptor,
1450 ConversionPatternRewriter &rewriter)
const {
1454 Location loc = op.getLoc();
1458 auto stream = adaptor.getAsyncDependencies().front();
1459 auto bufferSize = spMVBufferSizeCallBuilder
1460 .create(loc, rewriter,
1461 {modeA, adaptor.getSpmatA(), adaptor.getDnX(),
1462 adaptor.getDnY(), computeType, stream})
1464 rewriter.replaceOp(op, {bufferSize, stream});
1468LogicalResult ConvertSpMVOpToGpuRuntimeCallPattern::matchAndRewrite(
1469 gpu::SpMVOp op, OpAdaptor adaptor,
1470 ConversionPatternRewriter &rewriter)
const {
1474 Location loc = op.getLoc();
1478 auto stream = adaptor.getAsyncDependencies().front();
1480 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1481 spMVCallBuilder.create(loc, rewriter,
1482 {modeA, adaptor.getSpmatA(), adaptor.getDnX(),
1483 adaptor.getDnY(), computeType, pBuf, stream});
1484 rewriter.replaceOp(op, {stream});
1488LogicalResult ConvertSpMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1489 gpu::SpMMBufferSizeOp op, OpAdaptor adaptor,
1490 ConversionPatternRewriter &rewriter)
const {
1494 Location loc = op.getLoc();
1497 auto stream = adaptor.getAsyncDependencies().front();
1506 LLVM::AllocaOp::create(rewriter, loc, llvmPointerType, llvmPointerType,
1508 createCuSparseLtSpMMBufferSizeBuilder
1509 .create(loc, rewriter,
1510 {bufferSize, modeA, modeB, adaptor.getSpmatA(),
1511 adaptor.getDnmatB(), adaptor.getDnmatC(), computeType,
1515 auto bufferSizePtr1 = LLVM::GEPOp::create(
1516 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1518 auto bufferSizePtr2 = LLVM::GEPOp::create(
1519 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1522 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSize);
1524 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr1);
1526 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr2);
1528 rewriter.replaceOp(op, {bufferSize0, bufferSize1, bufferSize2, stream});
1533 createSpMMBufferSizeCallBuilder
1534 .create(loc, rewriter,
1535 {modeA, modeB, adaptor.getSpmatA(), adaptor.getDnmatB(),
1536 adaptor.getDnmatC(), computeType, stream})
1538 rewriter.replaceOp(op, {bufferSize, stream});
1543LogicalResult ConvertSDDMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1544 gpu::SDDMMBufferSizeOp op, OpAdaptor adaptor,
1545 ConversionPatternRewriter &rewriter)
const {
1549 Location loc = op.getLoc();
1554 auto stream = adaptor.getAsyncDependencies().front();
1556 createSDDMMBufferSizeCallBuilder
1557 .create(loc, rewriter,
1558 {modeA, modeB, adaptor.getDnmatA(), adaptor.getDnmatB(),
1559 adaptor.getSpmatC(), computeType, stream})
1561 rewriter.replaceOp(op, {bufferSize, stream});
1565LogicalResult ConvertSpMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1566 gpu::SpMMOp op, OpAdaptor adaptor,
1567 ConversionPatternRewriter &rewriter)
const {
1571 Location loc = op.getLoc();
1577 auto stream = adaptor.getAsyncDependencies().front();
1581 SmallVector<Value> pBufs;
1582 for (Value buffer : adaptor.getBuffers()) {
1583 Value pBuf = MemRefDescriptor(buffer).allocatedPtr(rewriter, loc);
1584 pBufs.push_back(pBuf);
1586 createCuSparseLtSpMMBuilder.create(
1588 {adaptor.getSpmatA(), adaptor.getDnmatB(), adaptor.getDnmatC(),
1589 pBufs[0], pBufs[1], pBufs[2], stream});
1591 Value pBuf = MemRefDescriptor(adaptor.getBuffers().front())
1592 .allocatedPtr(rewriter, loc);
1593 createSpMMCallBuilder.create(loc, rewriter,
1594 {modeA, modeB, adaptor.getSpmatA(),
1595 adaptor.getDnmatB(), adaptor.getDnmatC(),
1596 computeType, pBuf, stream});
1598 rewriter.replaceOp(op, {stream});
1602template <
typename T>
1604 converter.addConversion([&converter](T) ->
Type {
1605 return LLVM::LLVMPointerType::get(&converter.
getContext());
1609LogicalResult ConvertSDDMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1610 gpu::SDDMMOp op, OpAdaptor adaptor,
1611 ConversionPatternRewriter &rewriter)
const {
1615 Location loc = op.getLoc();
1620 auto stream = adaptor.getAsyncDependencies().front();
1622 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1623 createSDDMMCallBuilder.create(loc, rewriter,
1624 {modeA, modeB, adaptor.getDnmatA(),
1625 adaptor.getDnmatB(), adaptor.getSpmatC(),
1626 computeType, pBuf, stream});
1627 rewriter.replaceOp(op, {stream});
1632ConvertSpGEMMCreateDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1633 gpu::SpGEMMCreateDescrOp op, OpAdaptor adaptor,
1634 ConversionPatternRewriter &rewriter)
const {
1638 Location loc = op.getLoc();
1639 auto stream = adaptor.getAsyncDependencies().front();
1640 Value descr = createSpGEMMCreateDescrBuilder.create(loc, rewriter, {stream})
1642 rewriter.replaceOp(op, {descr, stream});
1647ConvertSpGEMMDestroyDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1648 gpu::SpGEMMDestroyDescrOp op, OpAdaptor adaptor,
1649 ConversionPatternRewriter &rewriter)
const {
1653 Location loc = op.getLoc();
1654 auto stream = adaptor.getAsyncDependencies().front();
1655 createSpGEMMDestroyDescrBuilder.create(loc, rewriter,
1656 {adaptor.getDesc(), stream});
1657 rewriter.replaceOp(op, {stream});
1662ConvertSpGEMMWorkEstimationOrComputeOpToGpuRuntimeCallPattern::matchAndRewrite(
1663 gpu::SpGEMMWorkEstimationOrComputeOp op, OpAdaptor adaptor,
1664 ConversionPatternRewriter &rewriter)
const {
1668 Location loc = op.getLoc();
1673 auto stream = adaptor.getAsyncDependencies().front();
1676 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1677 Value bufferSizeNew;
1679 if (adaptor.getKind() ==
1680 gpu::SpGEMMWorkEstimationOrComputeKind::WORK_ESTIMATION) {
1682 createSpGEMMWorkEstimationBuilder
1683 .create(loc, rewriter,
1684 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1685 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1686 adaptor.getBufferSz(), pBuf, stream})
1690 createSpGEMMComputeBuilder
1691 .create(loc, rewriter,
1692 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1693 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1694 adaptor.getBufferSz(), pBuf, stream})
1697 rewriter.replaceOp(op, {bufferSizeNew, stream});
1701LogicalResult ConvertSpGEMMCopyOpToGpuRuntimeCallPattern::matchAndRewrite(
1702 gpu::SpGEMMCopyOp op, OpAdaptor adaptor,
1703 ConversionPatternRewriter &rewriter)
const {
1707 Location loc = op.getLoc();
1712 auto stream = adaptor.getAsyncDependencies().front();
1713 createSpGEMMCopyBuilder.create(loc, rewriter,
1714 {adaptor.getDesc(), modeA, modeB,
1715 adaptor.getSpmatA(), adaptor.getSpmatB(),
1716 adaptor.getSpmatC(), computeType, stream});
1717 rewriter.replaceOp(op, {stream});
1721LogicalResult ConvertSpMatGetSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1722 gpu::SpMatGetSizeOp op, OpAdaptor adaptor,
1723 ConversionPatternRewriter &rewriter)
const {
1727 Location loc = op.getLoc();
1728 auto stream = adaptor.getAsyncDependencies().front();
1731 auto buffer = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1732 llvmInt64Type, three, 16);
1734 auto rowsPtr = LLVM::GEPOp::create(
1735 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1737 auto colsPtr = LLVM::GEPOp::create(
1738 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1740 auto nnzsPtr = LLVM::GEPOp::create(
1741 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1743 createSpMatGetSizeBuilder.create(
1744 loc, rewriter, {adaptor.getSpmat(), rowsPtr, colsPtr, nnzsPtr, stream});
1745 auto rows = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, rowsPtr);
1746 auto cols = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, colsPtr);
1747 auto nnzs = LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, nnzsPtr);
1749 rewriter.replaceOp(op, {rows, cols, nnzs, stream});
1753LogicalResult ConvertSetCsrPointersOpToGpuRuntimeCallPattern::matchAndRewrite(
1754 gpu::SetCsrPointersOp op, OpAdaptor adaptor,
1755 ConversionPatternRewriter &rewriter)
const {
1759 Location loc = op.getLoc();
1760 auto stream = adaptor.getAsyncDependencies().front();
1762 MemRefDescriptor(adaptor.getPositions()).allocatedPtr(rewriter, loc);
1764 MemRefDescriptor(adaptor.getCoordinates()).allocatedPtr(rewriter, loc);
1766 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1767 createSetCsrPointersBuilder.create(
1768 loc, rewriter, {adaptor.getSpmat(), pPos, pCrd, pVal, stream});
1769 rewriter.replaceOp(op, {stream});
1773LogicalResult ConvertCreateCscOpToGpuRuntimeCallPattern::matchAndRewrite(
1774 gpu::CreateCscOp op, OpAdaptor adaptor,
1775 ConversionPatternRewriter &rewriter)
const {
1779 Location loc = op.getLoc();
1780 auto stream = adaptor.getAsyncDependencies().front();
1782 MemRefDescriptor(adaptor.getColPos()).allocatedPtr(rewriter, loc);
1784 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1786 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1788 llvm::cast<MemRefType>(op.getColPos().getType()).getElementType();
1790 llvm::cast<MemRefType>(op.getRowIdxs().getType()).getElementType();
1792 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1797 createCscCallBuilder
1798 .create(loc, rewriter,
1799 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1800 pColPos, pRowIdxs, pValues, ptp, itp, dtp, stream})
1802 rewriter.replaceOp(op, {handle, stream});
1806LogicalResult ConvertCreateBsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1807 gpu::CreateBsrOp op, OpAdaptor adaptor,
1808 ConversionPatternRewriter &rewriter)
const {
1812 Location loc = op.getLoc();
1813 auto stream = adaptor.getAsyncDependencies().front();
1815 MemRefDescriptor(adaptor.getBRowPos()).allocatedPtr(rewriter, loc);
1817 MemRefDescriptor(adaptor.getBColIdxs()).allocatedPtr(rewriter, loc);
1819 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1821 llvm::cast<MemRefType>(op.getBRowPos().getType()).getElementType();
1823 llvm::cast<MemRefType>(op.getBColIdxs().getType()).getElementType();
1825 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1830 createBsrCallBuilder
1831 .create(loc, rewriter,
1832 {adaptor.getBrows(), adaptor.getBcols(), adaptor.getBnnz(),
1833 adaptor.getRBlockSize(), adaptor.getCBlockSize(), pRowPos,
1834 pColIdxs, pValues, ptp, itp, dtp, stream})
1836 rewriter.replaceOp(op, {handle, stream});
1842 bool kernelBarePtrCallConv,
bool kernelIntersperseSizeCallConv) {
1852 patterns.
add<ConvertAsyncYieldToGpuRuntimeCallPattern>(converter,
1855 patterns.
add<ConvertAllocOpToGpuRuntimeCallPattern,
1856 ConvertDeallocOpToGpuRuntimeCallPattern,
1857 ConvertHostRegisterOpToGpuRuntimeCallPattern,
1858 ConvertHostUnregisterOpToGpuRuntimeCallPattern,
1859 ConvertMemcpyOpToGpuRuntimeCallPattern,
1860 ConvertMemsetOpToGpuRuntimeCallPattern,
1861 ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern,
1862 ConvertWaitAsyncOpToGpuRuntimeCallPattern,
1863 ConvertWaitOpToGpuRuntimeCallPattern,
1864 ConvertCreateDnTensorOpToGpuRuntimeCallPattern,
1865 ConvertDestroyDnTensorOpToGpuRuntimeCallPattern,
1866 ConvertCreateCooOpToGpuRuntimeCallPattern,
1867 ConvertCreateCooAoSOpToGpuRuntimeCallPattern,
1868 ConvertCreateCsrOpToGpuRuntimeCallPattern,
1869 ConvertCreateCscOpToGpuRuntimeCallPattern,
1870 ConvertCreateBsrOpToGpuRuntimeCallPattern,
1871 ConvertCreate2To4SpMatOpToGpuRuntimeCallPattern,
1872 ConvertDestroySpMatOpToGpuRuntimeCallPattern,
1873 ConvertSpMVBufferSizeOpToGpuRuntimeCallPattern,
1874 ConvertSpMVOpToGpuRuntimeCallPattern,
1875 ConvertSpMMBufferSizeOpToGpuRuntimeCallPattern,
1876 ConvertSDDMMBufferSizeOpToGpuRuntimeCallPattern,
1877 ConvertSpMMOpToGpuRuntimeCallPattern,
1878 ConvertSDDMMOpToGpuRuntimeCallPattern,
1879 ConvertSpGEMMCreateDescrOpToGpuRuntimeCallPattern,
1880 ConvertSpGEMMDestroyDescrOpToGpuRuntimeCallPattern,
1881 ConvertSpGEMMWorkEstimationOrComputeOpToGpuRuntimeCallPattern,
1882 ConvertSpGEMMCopyOpToGpuRuntimeCallPattern,
1883 ConvertSpMatGetSizeOpToGpuRuntimeCallPattern,
1884 ConvertSetCsrPointersOpToGpuRuntimeCallPattern>(converter);
1885 patterns.
add<LegalizeLaunchFuncOpPattern>(converter, kernelBarePtrCallConv,
1886 kernelIntersperseSizeCallConv);
1894struct GPUModuleOpConvertToLLVMInterface
1895 :
public ConvertToLLVMOpInterface::ExternalModel<
1896 GPUModuleOpConvertToLLVMInterface, gpu::GPUModuleOp> {
1898 void getConvertToLLVMConversionAttrs(
1903void GPUModuleOpConvertToLLVMInterface::getConvertToLLVMConversionAttrs(
1904 Operation *op, SmallVectorImpl<ConvertToLLVMAttrInterface> &attrs)
const {
1905 auto module = cast<gpu::GPUModuleOp>(op);
1906 ArrayAttr targetsAttr =
module.getTargetsAttr();
1908 if (!targetsAttr || targetsAttr.size() != 1)
1910 if (
auto patternAttr = dyn_cast<ConvertToLLVMAttrInterface>(targetsAttr[0]))
1911 attrs.push_back(patternAttr);
1916 gpu::GPUModuleOp::attachInterface<GPUModuleOpConvertToLLVMInterface>(*ctx);
static void addOpaquePointerConversion(LLVMTypeConverter &converter)
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)
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