MLIR 24.0.0git
GPUToLLVMConversion.cpp
Go to the documentation of this file.
1//===- ConvertLaunchFuncToGpuRuntimeCalls.cpp - MLIR GPU lowering passes --===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements a pass to convert gpu.launch_func op into a sequence of
10// GPU runtime calls. As most of GPU runtimes does not have a stable published
11// ABI, this pass uses a slim runtime layer that builds on top of the public
12// API from GPU runtime headers.
13//
14//===----------------------------------------------------------------------===//
15
17
34#include "mlir/IR/Attributes.h"
35#include "mlir/IR/Builders.h"
36#include "mlir/IR/BuiltinOps.h"
39
40#include "llvm/ADT/STLExtras.h"
41
42#define DEBUG_TYPE "gpu-to-llvm"
43
44namespace mlir {
45#define GEN_PASS_DEF_GPUTOLLVMCONVERSIONPASS
46#include "mlir/Conversion/Passes.h.inc"
47} // namespace mlir
48
49using namespace mlir;
50
51namespace {
52class GpuToLLVMConversionPass
53 : public impl::GpuToLLVMConversionPassBase<GpuToLLVMConversionPass> {
54public:
55 using Base::Base;
56 void getDependentDialects(DialectRegistry &registry) const final {
57 Base::getDependentDialects(registry);
59 }
60 // Run the dialect converter on the module.
61 void runOnOperation() override;
62};
63
64template <typename OpTy>
65class ConvertOpToGpuRuntimeCallPattern : public ConvertOpToLLVMPattern<OpTy> {
66public:
67 explicit ConvertOpToGpuRuntimeCallPattern(
68 const LLVMTypeConverter &typeConverter, PatternBenefit benefit = 1)
69 : ConvertOpToLLVMPattern<OpTy>(typeConverter, benefit) {}
70
71protected:
72 Value getNumElements(ConversionPatternRewriter &rewriter, Location loc,
73 MemRefType type, MemRefDescriptor desc) const {
74 Type indexType = ConvertToLLVMPattern::getIndexType();
75 if (type.hasStaticShape())
77 rewriter, loc, indexType, type.getNumElements());
78 // Compute the number of elements by multiplying all the dim sizes.
79 uint64_t rank = type.getRank();
80 Value numElements = desc.size(rewriter, loc, /*pos=*/0);
81 for (unsigned i = 1; i < rank; i++)
82 numElements = LLVM::MulOp::create(rewriter, loc, numElements,
83 desc.size(rewriter, loc, /*pos=*/i));
84 return numElements;
85 }
86
87 MLIRContext *context = &this->getTypeConverter()->getContext();
88
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));
98
99 FunctionCallBuilder streamCreateCallBuilder = {
100 "mgpuStreamCreate", llvmPointerType /* void *stream */, {}};
101 FunctionCallBuilder streamDestroyCallBuilder = {
102 "mgpuStreamDestroy", llvmVoidType, {llvmPointerType /* void *stream */}};
103 FunctionCallBuilder streamSynchronizeCallBuilder = {
104 "mgpuStreamSynchronize",
105 llvmVoidType,
106 {llvmPointerType /* void *stream */}};
107 FunctionCallBuilder streamWaitEventCallBuilder = {
108 "mgpuStreamWaitEvent",
109 llvmVoidType,
110 {llvmPointerType /* void *stream */, llvmPointerType /* void *event */}};
111 FunctionCallBuilder eventCreateCallBuilder = {
112 "mgpuEventCreate", llvmPointerType /* void *event */, {}};
113 FunctionCallBuilder eventDestroyCallBuilder = {
114 "mgpuEventDestroy", llvmVoidType, {llvmPointerType /* void *event */}};
115 FunctionCallBuilder eventSynchronizeCallBuilder = {
116 "mgpuEventSynchronize",
117 llvmVoidType,
118 {llvmPointerType /* void *event */}};
119 FunctionCallBuilder eventRecordCallBuilder = {
120 "mgpuEventRecord",
121 llvmVoidType,
122 {llvmPointerType /* void *event */, llvmPointerType /* void *stream */}};
123 FunctionCallBuilder hostRegisterCallBuilder = {
124 "mgpuMemHostRegisterMemRef",
125 llvmVoidType,
126 {llvmIntPtrType /* intptr_t rank */,
127 llvmPointerType /* void *memrefDesc */,
128 llvmIntPtrType /* intptr_t elementSizeBytes */}};
129 FunctionCallBuilder hostUnregisterCallBuilder = {
130 "mgpuMemHostUnregisterMemRef",
131 llvmVoidType,
132 {llvmIntPtrType /* intptr_t rank */,
133 llvmPointerType /* void *memrefDesc */,
134 llvmIntPtrType /* intptr_t elementSizeBytes */}};
135 FunctionCallBuilder allocCallBuilder = {
136 "mgpuMemAlloc",
137 llvmPointerType /* void * */,
138 {llvmIntPtrType /* intptr_t sizeBytes */,
139 llvmPointerType /* void *stream */,
140 llvmInt8Type /* bool isHostShared */}};
141 FunctionCallBuilder deallocCallBuilder = {
142 "mgpuMemFree",
143 llvmVoidType,
144 {llvmPointerType /* void *ptr */, llvmPointerType /* void *stream */}};
145 FunctionCallBuilder memcpyCallBuilder = {
146 "mgpuMemcpy",
147 llvmVoidType,
148 {llvmPointerType /* void *dst */, llvmPointerType /* void *src */,
149 llvmIntPtrType /* intptr_t sizeBytes */,
150 llvmPointerType /* void *stream */}};
151 FunctionCallBuilder memset16CallBuilder = {
152 "mgpuMemset16",
153 llvmVoidType,
154 {llvmPointerType /* void *dst */,
155 llvmInt16Type /* unsigned short value */,
156 llvmIntPtrType /* intptr_t sizeBytes */,
157 llvmPointerType /* void *stream */}};
158 FunctionCallBuilder memset32CallBuilder = {
159 "mgpuMemset32",
160 llvmVoidType,
161 {llvmPointerType /* void *dst */, llvmInt32Type /* unsigned int value */,
162 llvmIntPtrType /* intptr_t sizeBytes */,
163 llvmPointerType /* void *stream */}};
164 FunctionCallBuilder setDefaultDeviceCallBuilder = {
165 "mgpuSetDefaultDevice",
166 llvmVoidType,
167 {llvmInt32Type /* uint32_t devIndex */}};
168 FunctionCallBuilder createDnVecCallBuilder = {
169 "mgpuCreateDnVec",
170 llvmPointerType,
171 {llvmIntPtrType, llvmPointerType, llvmInt32Type,
172 llvmPointerType /* void *stream */}};
173 FunctionCallBuilder destroyDnVecCallBuilder = {
174 "mgpuDestroyDnVec",
175 llvmVoidType,
176 {llvmPointerType, llvmPointerType /* void *stream */}};
177 FunctionCallBuilder createDnMatCallBuilder = {
178 "mgpuCreateDnMat",
179 llvmPointerType,
180 {llvmIntPtrType, llvmIntPtrType, llvmPointerType, llvmInt32Type,
181 llvmPointerType /* void *stream */}};
182 FunctionCallBuilder destroyDnMatCallBuilder = {
183 "mgpuDestroyDnMat",
184 llvmVoidType,
185 {llvmPointerType, llvmPointerType /* void *stream */}};
186 FunctionCallBuilder createCooCallBuilder = {
187 "mgpuCreateCoo",
188 llvmPointerType,
189 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
190 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
191 llvmPointerType /* void *stream */}};
192 FunctionCallBuilder createCooAoSCallBuilder = {
193 "mgpuCreateCooAoS", // deprecated in cuSPARSE 11.2
194 llvmPointerType,
195 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
196 llvmPointerType, llvmInt32Type, llvmInt32Type,
197 llvmPointerType /* void *stream */}};
198 FunctionCallBuilder createCsrCallBuilder = {
199 "mgpuCreateCsr",
200 llvmPointerType,
201 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
202 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
203 llvmInt32Type, llvmPointerType /* void *stream */}};
204 FunctionCallBuilder createCscCallBuilder = {
205 "mgpuCreateCsc",
206 llvmPointerType,
207 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
208 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
209 llvmInt32Type, llvmPointerType /* void *stream */}};
210 FunctionCallBuilder createBsrCallBuilder = {
211 "mgpuCreateBsr",
212 llvmPointerType,
213 {llvmIntPtrType, llvmIntPtrType, llvmIntPtrType, llvmIntPtrType,
214 llvmIntPtrType, llvmPointerType, llvmPointerType, llvmPointerType,
215 llvmInt32Type, llvmInt32Type, llvmInt32Type,
216 llvmPointerType /* void *stream */}};
217 FunctionCallBuilder destroySpMatCallBuilder = {
218 "mgpuDestroySpMat",
219 llvmVoidType,
220 {llvmPointerType, llvmPointerType /* void *stream */}};
221 FunctionCallBuilder spMVBufferSizeCallBuilder = {
222 "mgpuSpMVBufferSize",
223 llvmIntPtrType,
224 {llvmInt32Type, llvmPointerType, llvmPointerType, llvmPointerType,
225 llvmInt32Type, llvmPointerType /* void *stream */}};
226 FunctionCallBuilder spMVCallBuilder = {
227 "mgpuSpMV",
228 llvmVoidType,
229 {llvmInt32Type, llvmPointerType, llvmPointerType, llvmPointerType,
230 llvmInt32Type, llvmPointerType, llvmPointerType /* void *stream */}};
231 FunctionCallBuilder createSpMMBufferSizeCallBuilder = {
232 "mgpuSpMMBufferSize",
233 llvmIntPtrType,
234 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
235 llvmPointerType, llvmInt32Type, llvmPointerType /* void *stream */}};
236 FunctionCallBuilder createSpMMCallBuilder = {
237 "mgpuSpMM",
238 llvmVoidType,
239 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
240 llvmPointerType, llvmInt32Type, llvmPointerType,
241 llvmPointerType /* void *stream */}};
242 FunctionCallBuilder createSDDMMBufferSizeCallBuilder = {
243 "mgpuSDDMMBufferSize",
244 llvmIntPtrType,
245 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
246 llvmPointerType, llvmInt32Type, llvmPointerType /* void *stream */}};
247 FunctionCallBuilder createSDDMMCallBuilder = {
248 "mgpuSDDMM",
249 llvmVoidType,
250 {llvmInt32Type, llvmInt32Type, llvmPointerType, llvmPointerType,
251 llvmPointerType, llvmInt32Type, llvmPointerType,
252 llvmPointerType /* void *stream */}};
253 FunctionCallBuilder createLtDnMatCallBuilder = {
254 "mgpuCreateCuSparseLtDnMat",
255 llvmVoidType,
256 {llvmPointerType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
257 llvmInt32Type, llvmPointerType /* void *stream */}};
258 FunctionCallBuilder destroyCuSparseLtSpMatBuilder = {
259 "mgpuDestroyCuSparseLtSpMat",
260 llvmVoidType,
261 {llvmPointerType, llvmPointerType /* void *stream */}};
262 FunctionCallBuilder destroyCuSparseLtDnMatBuilder = {
263 "mgpuDestroyCuSparseLtDnMat",
264 llvmVoidType,
265 {llvmPointerType, llvmPointerType /* void *stream */}};
266 FunctionCallBuilder create2To4SpMatCallBuilder = {
267 "mgpuCusparseLtCreate2To4SpMat",
268 llvmVoidType,
269 {llvmPointerType, llvmIntPtrType, llvmIntPtrType, llvmPointerType,
270 llvmInt32Type, llvmPointerType /* void *stream */}};
271 FunctionCallBuilder createCuSparseLtSpMMBufferSizeBuilder = {
272 "mgpuCuSparseLtSpMMBufferSize",
273 llvmVoidType,
274 {llvmPointerType, llvmInt32Type, llvmInt32Type, llvmPointerType,
275 llvmPointerType, llvmPointerType, llvmInt32Type, llvmInt32Type,
276 llvmPointerType /*void *stream*/}};
277 FunctionCallBuilder createCuSparseLtSpMMBuilder = {
278 "mgpuCuSparseLtSpMM",
279 llvmVoidType,
280 {llvmPointerType, llvmPointerType, llvmPointerType, llvmPointerType,
281 llvmPointerType, llvmPointerType, llvmPointerType /*void *stream*/}};
282 FunctionCallBuilder createSpGEMMCreateDescrBuilder = {
283 "mgpuSpGEMMCreateDescr",
284 llvmPointerType,
285 {llvmPointerType /*void *stream*/}};
286 FunctionCallBuilder createSpGEMMDestroyDescrBuilder = {
287 "mgpuSpGEMMDestroyDescr",
288 llvmVoidType,
289 {llvmPointerType /*s*/, llvmPointerType /*void *stream*/}};
290 FunctionCallBuilder createSpGEMMWorkEstimationBuilder = {
291 "mgpuSpGEMMWorkEstimation",
292 llvmIntPtrType,
293 {llvmPointerType /*s*/, llvmInt32Type /*ma*/, llvmInt32Type /*mb*/,
294 llvmPointerType /*a*/, llvmPointerType /*b*/, llvmPointerType /*c*/,
295 llvmInt32Type /*ctp*/, llvmIntPtrType /*bs*/, llvmPointerType /*buf*/,
296 llvmPointerType /*void *stream*/}};
297 FunctionCallBuilder createSpGEMMComputeBuilder = {
298 "mgpuSpGEMMCompute",
299 llvmIntPtrType,
300 {llvmPointerType /*s*/, llvmInt32Type /*ma*/, llvmInt32Type /*mb*/,
301 llvmPointerType /*a*/, llvmPointerType /*b*/, llvmPointerType /*c*/,
302 llvmInt32Type /*ctp*/, llvmIntPtrType /*bs*/, llvmPointerType /*buf*/,
303 llvmPointerType /*void *stream*/}};
304 FunctionCallBuilder createSpGEMMCopyBuilder = {
305 "mgpuSpGEMMCopy",
306 llvmVoidType,
307 {llvmPointerType /*s*/, llvmInt32Type /*ma*/, llvmInt32Type /*mb*/,
308 llvmPointerType /*a*/, llvmPointerType /*b*/, llvmPointerType /*c*/,
309 llvmInt32Type /*ctp*/, llvmPointerType /*void *stream*/}};
310 FunctionCallBuilder createSpMatGetSizeBuilder = {
311 "mgpuSpMatGetSize",
312 llvmVoidType,
313 {llvmPointerType /*mc*/, llvmPointerType /*rc*/, llvmPointerType /*cc*/,
314 llvmPointerType /*nc*/, llvmPointerType /*void *stream*/}};
315 FunctionCallBuilder createSetCsrPointersBuilder = {
316 "mgpuSetCsrPointers",
317 llvmVoidType,
318 {llvmPointerType /*spmat*/, llvmPointerType /*pos*/,
319 llvmPointerType /*crd*/, llvmPointerType /*val*/,
320 llvmPointerType /*void *stream*/}};
321};
322
323/// A rewrite pattern to convert gpu.host_register operations into a GPU runtime
324/// call. Currently it supports CUDA and ROCm (HIP).
325class ConvertHostRegisterOpToGpuRuntimeCallPattern
326 : public ConvertOpToGpuRuntimeCallPattern<gpu::HostRegisterOp> {
327public:
328 ConvertHostRegisterOpToGpuRuntimeCallPattern(
329 const LLVMTypeConverter &typeConverter)
330 : ConvertOpToGpuRuntimeCallPattern<gpu::HostRegisterOp>(typeConverter) {}
331
332private:
333 LogicalResult
334 matchAndRewrite(gpu::HostRegisterOp hostRegisterOp, OpAdaptor adaptor,
335 ConversionPatternRewriter &rewriter) const override;
336};
337
338class ConvertHostUnregisterOpToGpuRuntimeCallPattern
339 : public ConvertOpToGpuRuntimeCallPattern<gpu::HostUnregisterOp> {
340public:
341 ConvertHostUnregisterOpToGpuRuntimeCallPattern(
342 const LLVMTypeConverter &typeConverter)
343 : ConvertOpToGpuRuntimeCallPattern<gpu::HostUnregisterOp>(typeConverter) {
344 }
345
346private:
347 LogicalResult
348 matchAndRewrite(gpu::HostUnregisterOp hostUnregisterOp, OpAdaptor adaptor,
349 ConversionPatternRewriter &rewriter) const override;
350};
351
352/// A rewrite pattern to convert gpu.alloc operations into a GPU runtime
353/// call. Currently it supports CUDA and ROCm (HIP).
354class ConvertAllocOpToGpuRuntimeCallPattern
355 : public ConvertOpToGpuRuntimeCallPattern<gpu::AllocOp> {
356public:
357 ConvertAllocOpToGpuRuntimeCallPattern(const LLVMTypeConverter &typeConverter)
358 : ConvertOpToGpuRuntimeCallPattern<gpu::AllocOp>(typeConverter) {}
359
360private:
361 LogicalResult
362 matchAndRewrite(gpu::AllocOp allocOp, OpAdaptor adaptor,
363 ConversionPatternRewriter &rewriter) const override;
364};
365
366/// A rewrite pattern to convert gpu.dealloc operations into a GPU runtime
367/// call. Currently it supports CUDA and ROCm (HIP).
368class ConvertDeallocOpToGpuRuntimeCallPattern
369 : public ConvertOpToGpuRuntimeCallPattern<gpu::DeallocOp> {
370public:
371 ConvertDeallocOpToGpuRuntimeCallPattern(
372 const LLVMTypeConverter &typeConverter)
373 : ConvertOpToGpuRuntimeCallPattern<gpu::DeallocOp>(typeConverter) {}
374
375private:
376 LogicalResult
377 matchAndRewrite(gpu::DeallocOp deallocOp, OpAdaptor adaptor,
378 ConversionPatternRewriter &rewriter) const override;
379};
380
381class ConvertAsyncYieldToGpuRuntimeCallPattern
382 : public ConvertOpToGpuRuntimeCallPattern<async::YieldOp> {
383public:
384 ConvertAsyncYieldToGpuRuntimeCallPattern(
385 const LLVMTypeConverter &typeConverter, PatternBenefit benefit = 1)
386 : ConvertOpToGpuRuntimeCallPattern<async::YieldOp>(typeConverter,
387 benefit) {}
388
389private:
390 LogicalResult
391 matchAndRewrite(async::YieldOp yieldOp, OpAdaptor adaptor,
392 ConversionPatternRewriter &rewriter) const override;
393};
394
395/// A rewrite pattern to convert gpu.wait operations into a GPU runtime
396/// call. Currently it supports CUDA and ROCm (HIP).
397class ConvertWaitOpToGpuRuntimeCallPattern
398 : public ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp> {
399public:
400 ConvertWaitOpToGpuRuntimeCallPattern(const LLVMTypeConverter &typeConverter)
401 : ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp>(typeConverter) {}
402
403private:
404 LogicalResult
405 matchAndRewrite(gpu::WaitOp waitOp, OpAdaptor adaptor,
406 ConversionPatternRewriter &rewriter) const override;
407};
408
409/// A rewrite pattern to convert gpu.wait async operations into a GPU runtime
410/// call. Currently it supports CUDA and ROCm (HIP).
411class ConvertWaitAsyncOpToGpuRuntimeCallPattern
412 : public ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp> {
413public:
414 ConvertWaitAsyncOpToGpuRuntimeCallPattern(
415 const LLVMTypeConverter &typeConverter)
416 : ConvertOpToGpuRuntimeCallPattern<gpu::WaitOp>(typeConverter) {}
417
418private:
419 LogicalResult
420 matchAndRewrite(gpu::WaitOp waitOp, OpAdaptor adaptor,
421 ConversionPatternRewriter &rewriter) const override;
422};
423
424/// A rewrite patter to legalize gpu.launch_func with LLVM types.
425class LegalizeLaunchFuncOpPattern
426 : public ConvertOpToGpuRuntimeCallPattern<gpu::LaunchFuncOp> {
427public:
428 LegalizeLaunchFuncOpPattern(const LLVMTypeConverter &typeConverter,
429 bool kernelBarePtrCallConv,
430 bool kernelIntersperseSizeCallConv)
431 : ConvertOpToGpuRuntimeCallPattern<gpu::LaunchFuncOp>(typeConverter),
432 kernelBarePtrCallConv(kernelBarePtrCallConv),
433 kernelIntersperseSizeCallConv(kernelIntersperseSizeCallConv) {}
434
435private:
436 LogicalResult
437 matchAndRewrite(gpu::LaunchFuncOp launchOp, OpAdaptor adaptor,
438 ConversionPatternRewriter &rewriter) const override;
439
440 bool kernelBarePtrCallConv;
441 bool kernelIntersperseSizeCallConv;
442};
443
444/// A rewrite pattern to convert gpu.memcpy operations into a GPU runtime
445/// call. Currently it supports CUDA and ROCm (HIP).
446class ConvertMemcpyOpToGpuRuntimeCallPattern
447 : public ConvertOpToGpuRuntimeCallPattern<gpu::MemcpyOp> {
448public:
449 ConvertMemcpyOpToGpuRuntimeCallPattern(const LLVMTypeConverter &typeConverter)
450 : ConvertOpToGpuRuntimeCallPattern<gpu::MemcpyOp>(typeConverter) {}
451
452private:
453 LogicalResult
454 matchAndRewrite(gpu::MemcpyOp memcpyOp, OpAdaptor adaptor,
455 ConversionPatternRewriter &rewriter) const override;
456};
457
458/// A rewrite pattern to convert gpu.memset operations into a GPU runtime
459/// call. Currently it supports CUDA and ROCm (HIP).
460class ConvertMemsetOpToGpuRuntimeCallPattern
461 : public ConvertOpToGpuRuntimeCallPattern<gpu::MemsetOp> {
462public:
463 ConvertMemsetOpToGpuRuntimeCallPattern(const LLVMTypeConverter &typeConverter)
464 : ConvertOpToGpuRuntimeCallPattern<gpu::MemsetOp>(typeConverter) {}
465
466private:
467 LogicalResult
468 matchAndRewrite(gpu::MemsetOp memsetOp, OpAdaptor adaptor,
469 ConversionPatternRewriter &rewriter) const override;
470};
471
472/// A rewrite pattern to convert gpu.set_default_device to a GPU runtime call.
473/// Currently supports CUDA and ROCm (HIP)
474class ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern
475 : public ConvertOpToGpuRuntimeCallPattern<gpu::SetDefaultDeviceOp> {
476public:
477 ConvertSetDefaultDeviceOpToGpuRuntimeCallPattern(
478 const LLVMTypeConverter &typeConverter)
479 : ConvertOpToGpuRuntimeCallPattern<gpu::SetDefaultDeviceOp>(
480 typeConverter) {}
481
482 LogicalResult
483 matchAndRewrite(gpu::SetDefaultDeviceOp op, OpAdaptor adaptor,
484 ConversionPatternRewriter &rewriter) const override;
485};
486
487/// Generic rewriting rule for operation on sparse matrices.
488/// Currently supports CUDA (by means of cuSparse and cuSparseLt).
489#define DECLARE_CONVERT_OP_TO_GPU_RUNTIME_CALL_PATTERN(op_name) \
490 class Convert##op_name##ToGpuRuntimeCallPattern \
491 : public ConvertOpToGpuRuntimeCallPattern<gpu::op_name> { \
492 public: \
493 Convert##op_name##ToGpuRuntimeCallPattern( \
494 const LLVMTypeConverter &typeConverter) \
495 : ConvertOpToGpuRuntimeCallPattern<gpu::op_name>(typeConverter) {} \
496 \
497 private: \
498 LogicalResult \
499 matchAndRewrite(gpu::op_name op, OpAdaptor adaptor, \
500 ConversionPatternRewriter &rewriter) const override; \
501 };
502
520DECLARE_CONVERT_OP_TO_GPU_RUNTIME_CALL_PATTERN(SpGEMMWorkEstimationOrComputeOp)
524
525} // namespace
526
527void GpuToLLVMConversionPass::runOnOperation() {
528 MLIRContext *context = &getContext();
529
530 // Perform progressive lowering of vector transfer operations.
531 {
532 RewritePatternSet patterns(&getContext());
533 // Vector transfer ops with rank > 1 should be lowered with VectorToSCF.
535 /*maxTransferRank=*/1);
536 // Transform N-D vector.from_elements to 1-D vector.from_elements before
537 // conversion.
538 vector::populateVectorFromElementsUnrollPatterns(patterns);
539 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))
540 return signalPassFailure();
541 }
542
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);
549
550 // Populate all patterns from all dialects that implement the
551 // `ConvertToLLVMPatternInterface` interface.
552 std::vector<Dialect *> dialects = context->getLoadedDialects();
553 for (auto *iface :
554 llvm::make_isa_range<ConvertToLLVMPatternInterface>(dialects))
555 iface->populateConvertToLLVMConversionPatterns(target, converter, patterns);
556
557 // Preserve GPU modules and binaries. Modules are preserved as they can be
558 // converted later by `gpu-module-to-binary`.
559 target.addLegalOp<gpu::GPUModuleOp, gpu::BinaryOp>();
560 // Accept as legal LaunchFuncOps if the operands have been lowered.
561 target.addDynamicallyLegalOp<gpu::LaunchFuncOp>(
562 [&](gpu::LaunchFuncOp op) -> bool { return converter.isLegal(op); });
563
564 // These aren't covered by the ConvertToLLVMPatternInterface right now.
565 populateVectorToLLVMConversionPatterns(converter, patterns);
568 target);
569 populateGpuToLLVMConversionPatterns(converter, patterns,
570 kernelBarePtrCallConv,
571 kernelIntersperseSizeCallConv);
572
573 if (failed(
574 applyPartialConversion(getOperation(), target, std::move(patterns))))
575 signalPassFailure();
576}
577
579 ArrayRef<Value> arguments) const {
580 auto module = builder.getBlock()->getParent()->getParentOfType<ModuleOp>();
581 auto function = [&] {
582 if (auto function = module.lookupSymbol<LLVM::LLVMFuncOp>(functionName))
583 return function;
584 auto builder = OpBuilder::atBlockEnd(module.getBody());
585 return LLVM::LLVMFuncOp::create(builder, loc, functionName, functionType);
586 }();
587 return LLVM::CallOp::create(builder, loc, function, arguments);
588}
589
590// Corresponding to cusparseIndexType_t defined in cusparse.h.
591static int32_t getCuSparseIndexTypeFrom(Type type) {
592 if (type.isInteger(16))
593 return 1; // CUSPARSE_INDEX_16U
594 if (type.isInteger(32))
595 return 2; // CUSPARSE_INDEX_32I
596 return 3; // CUSPARSE_INDEX_64I
597}
598
599static int32_t getCuSparseLtDataTypeFrom(Type type) {
600 if (type.isF16())
601 return 0; // CUSPARSE_COMPUTE_16F,
602 if (type.isInteger(32))
603 return 1; // CUSPARSE_COMPUTE_32I
604 llvm_unreachable("unsupported type");
605 // TODO: add support to TF32
606}
607
608// Corresponding to cudaDataType_t defined in CUDA library_types.h.
609static int32_t getCuSparseDataTypeFrom(Type type) {
610 if (llvm::isa<ComplexType>(type)) {
611 // get the element type
612 auto elementType = cast<ComplexType>(type).getElementType();
613 if (elementType.isBF16())
614 return 15; // CUDA_C_16BF
615 if (elementType.isF16())
616 return 6; // CUDA_C_16F
617 if (elementType.isF32())
618 return 4; // CUDA_C_32F
619 if (elementType.isF64())
620 return 5; // CUDA_C_64F
621 if (elementType.isInteger(8))
622 return 7; // CUDA_C_8I
623 if (elementType.isInteger(16))
624 return 21; // CUDA_C_16I
625 if (elementType.isInteger(32))
626 return 11; // CUDA_C_32I
627 }
628 if (type.isBF16())
629 return 14; // CUDA_R_16BF
630 if (type.isF16())
631 return 2; // CUDA_R_16F
632 if (type.isF32())
633 return 0; // CUDA_R_32F
634 if (type.isF64())
635 return 1; // CUDA_R_64F
636 if (type.isInteger(8))
637 return 3; // CUDA_R_8I
638 if (type.isInteger(16))
639 return 20; // CUDA_R_16I
640 if (type.isInteger(32))
641 return 10; // CUDA_R_32I
642
643 llvm_unreachable("unsupported element type");
644}
645
646static gpu::Prune2To4SpMatFlag get2To4PruneFlag(Value spMat) {
647 return spMat.getDefiningOp<gpu::Create2To4SpMatOp>().getPruneFlag();
648}
649
650// TODO: We may want a run-time (of the mlir compiler) disablement/warning:
651// cusparseLt currently won't work for cuda architecture <8.0 and will trigger a
652// runtime (of the CUDA program) error , but it might be great if we could at
653// least output a warning when we found the target architecture is <8.0 and the
654// user still wants to use cusparseLt. to make sure when lowering gpu sparse
655// dialect to llvm calls, the cusparselt calls are disabled for cuda
656// architecture <8.0
657static bool is2To4Sparsity(Value spMat) {
658 if (auto op = spMat.getDefiningOp<gpu::Create2To4SpMatOp>())
659 return true;
660 if (auto op = spMat.getDefiningOp<gpu::CreateCooOp>())
661 return false;
662 if (auto op = spMat.getDefiningOp<gpu::CreateCooAoSOp>())
663 return false;
664 if (auto op = spMat.getDefiningOp<gpu::CreateCsrOp>())
665 return false;
666 if (auto op = spMat.getDefiningOp<gpu::CreateCscOp>())
667 return false;
668 if (auto op = spMat.getDefiningOp<gpu::CreateBsrOp>())
669 return false;
670 // Print the spMat defining op
671 spMat.getDefiningOp()->print(llvm::errs());
672 llvm_unreachable("cannot find spmat def");
673}
674
675static bool isSpMMCusparseLtOp(Value op) {
676 for (Operation *user : op.getUsers()) {
677 auto spmmOp = dyn_cast<gpu::SpMMOp>(user);
678 // If the other operator is 50% sparsity then we should use cusparseLt
679 if (!spmmOp)
680 continue;
681 if (is2To4Sparsity(spmmOp.getSpmatA()))
682 return true;
683 }
684 return false;
685}
686
687// Returns whether all operands are of LLVM type.
688static LogicalResult areAllLLVMTypes(Operation *op, ValueRange operands,
689 ConversionPatternRewriter &rewriter) {
690 if (!llvm::all_of(operands, [](Value value) {
691 return LLVM::isCompatibleType(value.getType());
692 }))
693 return rewriter.notifyMatchFailure(
694 op, "Cannot convert if operands aren't of LLVM type.");
695 return success();
696}
697
698static LogicalResult
699isAsyncWithOneDependency(ConversionPatternRewriter &rewriter,
700 gpu::AsyncOpInterface op) {
701 if (op.getAsyncDependencies().size() != 1)
702 return rewriter.notifyMatchFailure(
703 op, "Can only convert with exactly one async dependency.");
704
705 if (!op.getAsyncToken())
706 return rewriter.notifyMatchFailure(op, "Can convert only async version.");
707
708 return success();
709}
710
711LogicalResult ConvertHostRegisterOpToGpuRuntimeCallPattern::matchAndRewrite(
712 gpu::HostRegisterOp hostRegisterOp, OpAdaptor adaptor,
713 ConversionPatternRewriter &rewriter) const {
714 auto *op = hostRegisterOp.getOperation();
715 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)))
716 return failure();
717
718 Location loc = op->getLoc();
719
720 auto memRefType = hostRegisterOp.getValue().getType();
721 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
722 auto elementSize = getSizeInBytes(loc, elementType, rewriter);
723
724 auto arguments = getTypeConverter()->promoteOperands(
725 loc, op->getOperands(), adaptor.getOperands(), rewriter);
726 arguments.push_back(elementSize);
727 hostRegisterCallBuilder.create(loc, rewriter, arguments);
728
729 rewriter.eraseOp(op);
730 return success();
731}
732
733LogicalResult ConvertHostUnregisterOpToGpuRuntimeCallPattern::matchAndRewrite(
734 gpu::HostUnregisterOp hostUnregisterOp, OpAdaptor adaptor,
735 ConversionPatternRewriter &rewriter) const {
736 Operation *op = hostUnregisterOp.getOperation();
737 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)))
738 return failure();
739
740 Location loc = op->getLoc();
741
742 auto memRefType = hostUnregisterOp.getValue().getType();
743 auto elementType = cast<UnrankedMemRefType>(memRefType).getElementType();
744 auto elementSize = getSizeInBytes(loc, elementType, rewriter);
745
746 auto arguments = getTypeConverter()->promoteOperands(
747 loc, op->getOperands(), adaptor.getOperands(), rewriter);
748 arguments.push_back(elementSize);
749 hostUnregisterCallBuilder.create(loc, rewriter, arguments);
750
751 rewriter.eraseOp(op);
752 return success();
753}
754
755LogicalResult ConvertAllocOpToGpuRuntimeCallPattern::matchAndRewrite(
756 gpu::AllocOp allocOp, OpAdaptor adaptor,
757 ConversionPatternRewriter &rewriter) const {
758
759 MemRefType memRefType = allocOp.getType();
760
761 if (failed(areAllLLVMTypes(allocOp, adaptor.getOperands(), rewriter)) ||
762 !isConvertibleAndHasIdentityMaps(memRefType))
763 return failure();
764
765 auto loc = allocOp.getLoc();
766
767 bool isShared = allocOp.getHostShared();
768
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.");
775
776 // Get shape of the memref as values: static sizes are constant
777 // values and dynamic sizes are passed to 'alloc' as operands.
778 SmallVector<Value, 4> shape;
779 SmallVector<Value, 4> strides;
780 Value sizeBytes;
781 getMemRefDescriptorSizes(loc, memRefType, adaptor.getDynamicSizes(), rewriter,
782 shape, strides, sizeBytes);
783
784 // Allocate the underlying buffer and store a pointer to it in the MemRef
785 // descriptor.
786 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
787 Value stream = adaptor.getAsyncDependencies().empty()
788 ? nullPtr
789 : adaptor.getAsyncDependencies().front();
790
791 auto isHostShared = mlir::LLVM::ConstantOp::create(
792 rewriter, loc, llvmInt8Type, rewriter.getI8IntegerAttr(isShared));
793
794 Value allocatedPtr =
795 allocCallBuilder.create(loc, rewriter, {sizeBytes, stream, isHostShared})
796 .getResult();
797
798 // Cast the runtime-returned pointer to the memref's address space if it
799 // isn't there already, so it fits the descriptor's pointer slots.
800 unsigned dstAddrSpace = memRefType.getMemorySpaceAsInt();
801 unsigned srcAddrSpace =
802 cast<LLVM::LLVMPointerType>(allocatedPtr.getType()).getAddressSpace();
803 if (dstAddrSpace != srcAddrSpace) {
804 auto targetPtrTy =
805 LLVM::LLVMPointerType::get(rewriter.getContext(), dstAddrSpace);
806 allocatedPtr =
807 LLVM::AddrSpaceCastOp::create(rewriter, loc, targetPtrTy, allocatedPtr);
808 }
809
810 // No alignment.
811 Value alignedPtr = allocatedPtr;
812
813 // Create the MemRef descriptor.
814 auto memRefDescriptor = this->createMemRefDescriptor(
815 loc, memRefType, allocatedPtr, alignedPtr, shape, strides, rewriter);
816
817 if (allocOp.getAsyncToken()) {
818 // Async alloc: make dependent ops use the same stream.
819 rewriter.replaceOp(allocOp, {memRefDescriptor, stream});
820 } else {
821 rewriter.replaceOp(allocOp, {memRefDescriptor});
822 }
823
824 return success();
825}
826
827LogicalResult ConvertDeallocOpToGpuRuntimeCallPattern::matchAndRewrite(
828 gpu::DeallocOp deallocOp, OpAdaptor adaptor,
829 ConversionPatternRewriter &rewriter) const {
830 if (failed(areAllLLVMTypes(deallocOp, adaptor.getOperands(), rewriter)))
831 return failure();
832 if (adaptor.getAsyncDependencies().size() > 1)
833 return rewriter.notifyMatchFailure(
834 deallocOp, "Can convert with at most one async dependency.");
835
836 Location loc = deallocOp.getLoc();
837
838 Value pointer =
839 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
840 auto nullPtr = mlir::LLVM::ZeroOp::create(rewriter, loc, llvmPointerType);
841 Value stream = adaptor.getAsyncDependencies().empty()
842 ? nullPtr
843 : adaptor.getAsyncDependencies().front();
844 deallocCallBuilder.create(loc, rewriter, {pointer, stream});
845
846 if (deallocOp.getAsyncToken()) {
847 // Async dealloc: propagate the stream as the async token replacement.
848 rewriter.replaceOp(deallocOp, {stream});
849 } else {
850 // Sync dealloc: no results to replace, just remove the op.
851 rewriter.eraseOp(deallocOp);
852 }
853 return success();
854}
855
856static bool isGpuAsyncTokenType(Value value) {
857 return isa<gpu::AsyncTokenType>(value.getType());
858}
859
860// Converts !gpu.async.token operands of `async.yield` to runtime calls. The
861// !gpu.async.token are lowered to stream within the async.execute region, but
862// are passed as events between them. For each !gpu.async.token operand, we
863// create an event and record it on the stream.
864//
865// This pattern is registered with a higher benefit than the structural
866// async.yield rewriter from populateAsyncStructuralTypeConversionsAndLegality
867// so it wins when both match. Without that benefit override, the structural
868// pattern can win and silently retype gpu.async.token operands without
869// recording an event, leaving the host await to call cuEventSynchronize on
870// a stream pointer (a no-op that returns an error), racing the host against
871// the GPU.
872LogicalResult ConvertAsyncYieldToGpuRuntimeCallPattern::matchAndRewrite(
873 async::YieldOp yieldOp, OpAdaptor adaptor,
874 ConversionPatternRewriter &rewriter) const {
875 if (llvm::none_of(yieldOp.getOperands(), isGpuAsyncTokenType))
876 return rewriter.notifyMatchFailure(yieldOp, "no gpu async token operand");
877
878 Location loc = yieldOp.getLoc();
879 SmallVector<Value, 4> newOperands(adaptor.getOperands());
880 llvm::SmallDenseSet<Value> streams;
881 for (auto &operand : yieldOp->getOpOperands()) {
882 if (!isGpuAsyncTokenType(operand.get()))
883 continue;
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);
890 }
891 for (auto stream : streams)
892 streamDestroyCallBuilder.create(loc, rewriter, {stream});
893
894 rewriter.modifyOpInPlace(yieldOp, [&] { yieldOp->setOperands(newOperands); });
895 return success();
896}
897
898// Returns whether `value` is the result of an LLVM::CallOp to `functionName`.
899static bool isDefinedByCallTo(Value value, StringRef functionName) {
900 assert(isa<LLVM::LLVMPointerType>(value.getType()));
901 if (auto defOp = value.getDefiningOp<LLVM::CallOp>())
902 return *defOp.getCallee() == functionName;
903 return false;
904}
905
906// Converts `gpu.wait` to runtime calls. The converted op synchronizes the host
907// with the stream/event operands. The operands are destroyed. That is, it
908// assumes that it is not used afterwards or elsewhere. Otherwise we will get a
909// runtime error. Eventually, we should guarantee this property.
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.");
915
916 Location loc = waitOp.getLoc();
917
918 for (auto operand : adaptor.getOperands()) {
919 if (isDefinedByCallTo(operand, streamCreateCallBuilder.functionName)) {
920 // The converted operand's definition created a stream.
921 streamSynchronizeCallBuilder.create(loc, rewriter, {operand});
922 streamDestroyCallBuilder.create(loc, rewriter, {operand});
923 } else {
924 // Otherwise the converted operand is an event. This assumes that we use
925 // events in control flow code as well.
926 eventSynchronizeCallBuilder.create(loc, rewriter, {operand});
927 eventDestroyCallBuilder.create(loc, rewriter, {operand});
928 }
929 }
930
931 rewriter.eraseOp(waitOp);
932 return success();
933}
934
935// Converts `gpu.wait async` to runtime calls. The converted op creates a new
936// stream that is synchronized with stream/event operands. The operands are
937// destroyed. That is, it assumes that it is not used afterwards or elsewhere.
938// Otherwise we will get a runtime error. Eventually, we should guarantee this
939// property.
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.");
945
946 Location loc = waitOp.getLoc();
947
948 auto insertionPoint = rewriter.saveInsertionPoint();
949 SmallVector<Value, 1> events;
950 for (auto pair :
951 llvm::zip(waitOp.getAsyncDependencies(), adaptor.getOperands())) {
952 auto operand = std::get<1>(pair);
953 if (isDefinedByCallTo(operand, streamCreateCallBuilder.functionName)) {
954 // The converted operand's definition created a stream. Insert an event
955 // into the stream just after the last use of the original token operand.
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);
961 } else {
962 // Otherwise the converted operand is an event. This assumes that we use
963 // events in control flow code as well.
964 events.push_back(operand);
965 }
966 }
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});
974
975 return success();
976}
977
978// Legalize the op's operands.
979LogicalResult LegalizeLaunchFuncOpPattern::matchAndRewrite(
980 gpu::LaunchFuncOp launchOp, OpAdaptor adaptor,
981 ConversionPatternRewriter &rewriter) const {
982 if (failed(areAllLLVMTypes(launchOp, adaptor.getOperands(), rewriter)))
983 return failure();
984
985 // Fail when the synchronous version of the op has async dependencies. The
986 // lowering destroys the stream, and we do not want to check that there is no
987 // use of the stream after this op.
988 if (!launchOp.getAsyncToken() && !launchOp.getAsyncDependencies().empty())
989 return rewriter.notifyMatchFailure(
990 launchOp, "Cannot convert non-async op with async dependencies.");
991
992 Location loc = launchOp.getLoc();
993
994 Value stream = Value();
995 if (!adaptor.getAsyncDependencies().empty()) {
996 stream = adaptor.getAsyncDependencies().front();
997 // Synchronize additional async dependencies onto the primary stream using
998 // events, following the same approach as gpu.wait async lowering.
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())) {
1005 if (!isDefinedByCallTo(convertedDep,
1006 streamCreateCallBuilder.functionName)) {
1007 events.push_back(convertedDep);
1008 continue;
1009 }
1010 Operation *defOp = origDep.getDefiningOp();
1011 rewriter.setInsertionPointAfter(defOp);
1012 Value event =
1013 eventCreateCallBuilder.create(loc, rewriter, {}).getResult();
1014 eventRecordCallBuilder.create(loc, rewriter, {event, convertedDep});
1015 events.push_back(event);
1016 }
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});
1022 }
1023 }
1024 // If the async keyword is present and there are no dependencies, then a
1025 // stream must be created to pass to subsequent operations.
1026 else if (launchOp.getAsyncToken())
1027 stream = streamCreateCallBuilder.create(loc, rewriter, {}).getResult();
1028
1029 // Lower the kernel operands to match kernel parameters.
1030 // Note: If `useBarePtrCallConv` is set in the type converter's options,
1031 // the value of `kernelBarePtrCallConv` will be ignored.
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");
1041 }
1042 }
1043 SmallVector<Value, 8> llvmArguments = getTypeConverter()->promoteOperands(
1044 loc, origArguments, adaptor.getKernelOperands(), rewriter,
1045 /*useBarePtrCallConv=*/kernelBarePtrCallConv);
1046 SmallVector<Value, 8> llvmArgumentsWithSizes;
1047
1048 // Intersperse size information if requested.
1049 if (kernelIntersperseSizeCallConv) {
1050 if (origArguments.size() != llvmArguments.size()) {
1051 // This shouldn't happen if the bare-pointer calling convention is used.
1052 return rewriter.notifyMatchFailure(
1053 launchOp,
1054 "Cannot add sizes to arguments with one-to-many LLVM IR expansion.");
1055 }
1056
1057 llvmArgumentsWithSizes.reserve(llvmArguments.size() * 2);
1058 for (auto [llvmArg, origArg] : zip_equal(llvmArguments, origArguments)) {
1059 auto memrefTy = dyn_cast<MemRefType>(origArg.getType());
1060 if (!memrefTy) {
1061 return rewriter.notifyMatchFailure(
1062 launchOp, "Operand to launch op is not a memref.");
1063 }
1064
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.");
1070 }
1071
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.");
1077 }
1078
1079 uint64_t staticSize = static_cast<uint64_t>(bitwidth / 8) *
1080 static_cast<uint64_t>(memrefTy.getNumElements());
1081
1082 Value sizeArg =
1083 createIndexAttrConstant(rewriter, loc, getIndexType(), staticSize);
1084 llvmArgumentsWithSizes.push_back(llvmArg); // Presumably a bare pointer.
1085 llvmArgumentsWithSizes.push_back(sizeArg);
1086 }
1087 }
1088
1089 std::optional<gpu::KernelDim3> clusterSize = std::nullopt;
1090 if (launchOp.hasClusterSize()) {
1091 clusterSize =
1092 gpu::KernelDim3{adaptor.getClusterSizeX(), adaptor.getClusterSizeY(),
1093 adaptor.getClusterSizeZ()};
1094 }
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});
1108 else
1109 rewriter.eraseOp(launchOp);
1110 return success();
1111}
1112
1114 ConversionPatternRewriter &rewriter,
1115 LLVM::LLVMPointerType destinationType,
1116 Value sourcePtr,
1117 const LLVMTypeConverter &typeConverter) {
1118 auto sourceTy = cast<LLVM::LLVMPointerType>(sourcePtr.getType());
1119 if (destinationType.getAddressSpace() != sourceTy.getAddressSpace())
1120 sourcePtr = LLVM::AddrSpaceCastOp::create(
1121 rewriter, loc,
1122 LLVM::LLVMPointerType::get(rewriter.getContext(),
1123 destinationType.getAddressSpace()),
1124 sourcePtr);
1125 return sourcePtr;
1126}
1127
1128LogicalResult ConvertMemcpyOpToGpuRuntimeCallPattern::matchAndRewrite(
1129 gpu::MemcpyOp memcpyOp, OpAdaptor adaptor,
1130 ConversionPatternRewriter &rewriter) const {
1131 auto memRefType = cast<MemRefType>(memcpyOp.getSrc().getType());
1132
1133 if (failed(areAllLLVMTypes(memcpyOp, adaptor.getOperands(), rewriter)) ||
1134 !isConvertibleAndHasIdentityMaps(memRefType) ||
1135 failed(isAsyncWithOneDependency(rewriter, memcpyOp)))
1136 return failure();
1137
1138 auto loc = memcpyOp.getLoc();
1139
1140 MemRefDescriptor srcDesc(adaptor.getSrc());
1141 Value numElements = getNumElements(rewriter, loc, memRefType, srcDesc);
1142
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,
1148 numElements);
1149 auto sizeBytes =
1150 LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), gepPtr);
1151
1152 auto src = bitAndAddrspaceCast(loc, rewriter, llvmPointerType,
1153 srcDesc.alignedPtr(rewriter, loc),
1154 *getTypeConverter());
1155 auto dst = bitAndAddrspaceCast(
1156 loc, rewriter, llvmPointerType,
1157 MemRefDescriptor(adaptor.getDst()).alignedPtr(rewriter, loc),
1158 *getTypeConverter());
1159
1160 auto stream = adaptor.getAsyncDependencies().front();
1161 memcpyCallBuilder.create(loc, rewriter, {dst, src, sizeBytes, stream});
1162
1163 rewriter.replaceOp(memcpyOp, {stream});
1164
1165 return success();
1166}
1167
1168LogicalResult ConvertMemsetOpToGpuRuntimeCallPattern::matchAndRewrite(
1169 gpu::MemsetOp memsetOp, OpAdaptor adaptor,
1170 ConversionPatternRewriter &rewriter) const {
1171 auto memRefType = cast<MemRefType>(memsetOp.getDst().getType());
1172
1173 if (failed(areAllLLVMTypes(memsetOp, adaptor.getOperands(), rewriter)) ||
1174 !isConvertibleAndHasIdentityMaps(memRefType) ||
1175 failed(isAsyncWithOneDependency(rewriter, memsetOp)))
1176 return failure();
1177
1178 auto loc = memsetOp.getLoc();
1179
1180 Type valueType = adaptor.getValue().getType();
1181 unsigned bitWidth = valueType.getIntOrFloatBitWidth();
1182 // Ints and floats of 16 or 32 bit width are allowed.
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");
1186 }
1187
1188 unsigned valueTypeWidth = valueType.getIntOrFloatBitWidth();
1189 Type bitCastType = valueTypeWidth == 32 ? llvmInt32Type : llvmInt16Type;
1190
1191 MemRefDescriptor dstDesc(adaptor.getDst());
1192 Value numElements = getNumElements(rewriter, loc, memRefType, dstDesc);
1193
1194 auto value =
1195 LLVM::BitcastOp::create(rewriter, loc, bitCastType, adaptor.getValue());
1196 auto dst = bitAndAddrspaceCast(loc, rewriter, llvmPointerType,
1197 dstDesc.alignedPtr(rewriter, loc),
1198 *getTypeConverter());
1199
1200 auto stream = adaptor.getAsyncDependencies().front();
1201 FunctionCallBuilder builder =
1202 valueTypeWidth == 32 ? memset32CallBuilder : memset16CallBuilder;
1203 builder.create(loc, rewriter, {dst, value, numElements, stream});
1204
1205 rewriter.replaceOp(memsetOp, {stream});
1206 return success();
1207}
1208
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);
1216 return success();
1217}
1218
1219template <typename T>
1220static Value genConstInt32From(OpBuilder &builder, Location loc, T tValue) {
1221 Type llvmInt32Type = builder.getIntegerType(32);
1222 return LLVM::ConstantOp::create(builder, loc, llvmInt32Type,
1223 static_cast<int32_t>(tValue));
1224}
1225
1226LogicalResult ConvertCreateDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1227 gpu::CreateDnTensorOp op, OpAdaptor adaptor,
1228 ConversionPatternRewriter &rewriter) const {
1229 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1230 failed(isAsyncWithOneDependency(rewriter, op)))
1231 return failure();
1232 Location loc = op.getLoc();
1233 auto stream = adaptor.getAsyncDependencies().front();
1234 Value pTensor =
1235 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1236 Type dType = op.getMemref().getType().getElementType();
1237 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1238
1239 SmallVector<Value, 4> dims;
1240 for (Value dim : adaptor.getDims()) {
1241 dims.push_back(dim);
1242 }
1243
1244 Value handle;
1245 // TODO: For now, we track the use of the handle and lower it to cusparse /
1246 // cusparseLt accordingly. If in a block, both cusparse and cusparseLt are
1247 // used, we require two separate Creation ops to be the correct logic. In
1248 // future, we may add support to using one handle in sparse tensor / GPU
1249 // dialect in both cusparse and cusparseLt. use the cusparseLt create call if
1250 // the dnmat is used with spmat with 2:4 sparsity
1251 if (dims.size() == 2) {
1252 if (isSpMMCusparseLtOp(op.getDnTensor())) {
1253 auto handleSz =
1254 createIndexAttrConstant(rewriter, loc, getIndexType(), 11032);
1255 handle = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1256 llvmInt8Type, handleSz, /*alignment=*/16);
1257 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1258
1259 createLtDnMatCallBuilder
1260 .create(loc, rewriter,
1261 {handle, dims[0], dims[1], pTensor, dtp, stream})
1262 .getResult();
1263 } else {
1264 handle =
1265 createDnMatCallBuilder
1266 .create(loc, rewriter, {dims[0], dims[1], pTensor, dtp, stream})
1267 .getResult();
1268 }
1269 } else {
1270 assert(dims.size() == 1 && "Only 1D and 2D tensors are supported");
1271 handle = createDnVecCallBuilder
1272 .create(loc, rewriter, {dims[0], pTensor, dtp, stream})
1273 .getResult();
1274 }
1275 rewriter.replaceOp(op, {handle, stream});
1276 return success();
1277}
1278
1279LogicalResult ConvertDestroyDnTensorOpToGpuRuntimeCallPattern::matchAndRewrite(
1280 gpu::DestroyDnTensorOp op, OpAdaptor adaptor,
1281 ConversionPatternRewriter &rewriter) const {
1282 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1283 failed(isAsyncWithOneDependency(rewriter, op)))
1284 return failure();
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);
1291 }
1292 if (dims.size() == 2) {
1293 // Use the cusparseLt destroy call if the dnmat is used with spmat with
1294 // 2:4 sparsity
1295 if (isSpMMCusparseLtOp(op.getDnTensor())) {
1296 destroyCuSparseLtDnMatBuilder.create(loc, rewriter,
1297 {adaptor.getDnTensor(), stream});
1298 } else {
1299 destroyDnMatCallBuilder.create(loc, rewriter,
1300 {adaptor.getDnTensor(), stream});
1301 }
1302 } else {
1303 assert(dims.size() == 1 && "Only 1D and 2D tensors are supported");
1304 destroyDnVecCallBuilder.create(loc, rewriter,
1305 {adaptor.getDnTensor(), stream});
1306 }
1307 rewriter.replaceOp(op, {stream});
1308 return success();
1309}
1310
1311LogicalResult ConvertCreateCooOpToGpuRuntimeCallPattern::matchAndRewrite(
1312 gpu::CreateCooOp op, OpAdaptor adaptor,
1313 ConversionPatternRewriter &rewriter) const {
1314 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1315 failed(isAsyncWithOneDependency(rewriter, op)))
1316 return failure();
1317 Location loc = op.getLoc();
1318 auto stream = adaptor.getAsyncDependencies().front();
1319 Value pRowIdxs =
1320 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1321 Value pColIdxs =
1322 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1323 Value pValues =
1324 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1325 Type iType =
1326 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1327 Type dType =
1328 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1329 auto itp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(iType));
1330 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1331 auto handle =
1332 createCooCallBuilder
1333 .create(loc, rewriter,
1334 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1335 pRowIdxs, pColIdxs, pValues, itp, dtp, stream})
1336 .getResult();
1337 rewriter.replaceOp(op, {handle, stream});
1338 return success();
1339}
1340
1341LogicalResult ConvertCreateCooAoSOpToGpuRuntimeCallPattern::matchAndRewrite(
1342 gpu::CreateCooAoSOp op, OpAdaptor adaptor,
1343 ConversionPatternRewriter &rewriter) const {
1344 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1345 failed(isAsyncWithOneDependency(rewriter, op)))
1346 return failure();
1347 Location loc = op.getLoc();
1348 auto stream = adaptor.getAsyncDependencies().front();
1349 Value pIdxs = MemRefDescriptor(adaptor.getIdxs()).allocatedPtr(rewriter, loc);
1350 Value pValues =
1351 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1352 Type iType = llvm::cast<MemRefType>(op.getIdxs().getType()).getElementType();
1353 Type dType =
1354 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1355 auto itp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(iType));
1356 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1357 auto handle =
1358 createCooAoSCallBuilder
1359 .create(loc, rewriter,
1360 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1361 pIdxs, pValues, itp, dtp, stream})
1362 .getResult();
1363 rewriter.replaceOp(op, {handle, stream});
1364 return success();
1365}
1366
1367LogicalResult ConvertCreateCsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1368 gpu::CreateCsrOp op, OpAdaptor adaptor,
1369 ConversionPatternRewriter &rewriter) const {
1370 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1371 failed(isAsyncWithOneDependency(rewriter, op)))
1372 return failure();
1373 Location loc = op.getLoc();
1374 auto stream = adaptor.getAsyncDependencies().front();
1375 Value pRowPos =
1376 MemRefDescriptor(adaptor.getRowPos()).allocatedPtr(rewriter, loc);
1377 Value pColIdxs =
1378 MemRefDescriptor(adaptor.getColIdxs()).allocatedPtr(rewriter, loc);
1379 Value pValues =
1380 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1381 Type pType =
1382 llvm::cast<MemRefType>(op.getRowPos().getType()).getElementType();
1383 Type iType =
1384 llvm::cast<MemRefType>(op.getColIdxs().getType()).getElementType();
1385 Type dType =
1386 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1387 auto ptp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(pType));
1388 auto itp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(iType));
1389 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1390 auto handle =
1391 createCsrCallBuilder
1392 .create(loc, rewriter,
1393 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1394 pRowPos, pColIdxs, pValues, ptp, itp, dtp, stream})
1395 .getResult();
1396 rewriter.replaceOp(op, {handle, stream});
1397 return success();
1398}
1399
1400LogicalResult ConvertCreate2To4SpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1401 gpu::Create2To4SpMatOp op, OpAdaptor adaptor,
1402 ConversionPatternRewriter &rewriter) const {
1403 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1404 failed(isAsyncWithOneDependency(rewriter, op)))
1405 return failure();
1406 Location loc = op.getLoc();
1407 auto stream = adaptor.getAsyncDependencies().front();
1408 Value pMat =
1409 MemRefDescriptor(adaptor.getMemref()).allocatedPtr(rewriter, loc);
1410 Type dType =
1411 llvm::cast<MemRefType>(op.getMemref().getType()).getElementType();
1412 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1413
1414 // CUDA runner asserts the size is 44104 bytes.
1415 auto handleSz = createIndexAttrConstant(rewriter, loc, getIndexType(), 44104);
1416 Value handle = LLVM::AllocaOp::create(
1417 rewriter, loc, llvmPointerType, llvmInt8Type, handleSz, /*alignment=*/16);
1418 handle = LLVM::BitcastOp::create(rewriter, loc, llvmPointerType, handle);
1419
1420 create2To4SpMatCallBuilder
1421 .create(loc, rewriter,
1422 {handle, adaptor.getRows(), adaptor.getCols(), pMat, dtp, stream})
1423 .getResult();
1424 rewriter.replaceOp(op, {handle, stream});
1425 return success();
1426}
1427
1428LogicalResult ConvertDestroySpMatOpToGpuRuntimeCallPattern::matchAndRewrite(
1429 gpu::DestroySpMatOp op, OpAdaptor adaptor,
1430 ConversionPatternRewriter &rewriter) const {
1431 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1432 failed(isAsyncWithOneDependency(rewriter, op)))
1433 return failure();
1434 Location loc = op.getLoc();
1435 auto stream = adaptor.getAsyncDependencies().front();
1436 // Use the cusparseLt destroy call if the spmat is 2:4 sparsity
1437 if (is2To4Sparsity(op.getSpmat())) {
1438 destroyCuSparseLtSpMatBuilder.create(loc, rewriter,
1439 {adaptor.getSpmat(), stream});
1440
1441 } else {
1442 destroySpMatCallBuilder.create(loc, rewriter, {adaptor.getSpmat(), stream});
1443 }
1444 rewriter.replaceOp(op, {stream});
1445 return success();
1446}
1447
1448LogicalResult ConvertSpMVBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1449 gpu::SpMVBufferSizeOp op, OpAdaptor adaptor,
1450 ConversionPatternRewriter &rewriter) const {
1451 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1452 failed(isAsyncWithOneDependency(rewriter, op)))
1453 return failure();
1454 Location loc = op.getLoc();
1455 auto modeA = genConstInt32From(rewriter, loc, op.getModeA());
1456 auto computeType = genConstInt32From(
1457 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
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})
1463 .getResult();
1464 rewriter.replaceOp(op, {bufferSize, stream});
1465 return success();
1466}
1467
1468LogicalResult ConvertSpMVOpToGpuRuntimeCallPattern::matchAndRewrite(
1469 gpu::SpMVOp op, OpAdaptor adaptor,
1470 ConversionPatternRewriter &rewriter) const {
1471 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1472 failed(isAsyncWithOneDependency(rewriter, op)))
1473 return failure();
1474 Location loc = op.getLoc();
1475 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1476 auto computeType = genConstInt32From(
1477 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1478 auto stream = adaptor.getAsyncDependencies().front();
1479 Value pBuf =
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});
1485 return success();
1486}
1487
1488LogicalResult ConvertSpMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1489 gpu::SpMMBufferSizeOp op, OpAdaptor adaptor,
1490 ConversionPatternRewriter &rewriter) const {
1491 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1492 failed(isAsyncWithOneDependency(rewriter, op)))
1493 return failure();
1494 Location loc = op.getLoc();
1495 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1496 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
1497 auto stream = adaptor.getAsyncDependencies().front();
1498 Value bufferSize;
1499 if (is2To4Sparsity(op.getSpmatA())) {
1500 auto pruneFlag =
1501 genConstInt32From(rewriter, loc, get2To4PruneFlag(op.getSpmatA()));
1502 auto computeType = genConstInt32From(
1503 rewriter, loc, getCuSparseLtDataTypeFrom(adaptor.getComputeType()));
1504 auto three = createIndexAttrConstant(rewriter, loc, getIndexType(), 3);
1505 auto bufferSize =
1506 LLVM::AllocaOp::create(rewriter, loc, llvmPointerType, llvmPointerType,
1507 three, /*alignment=*/16);
1508 createCuSparseLtSpMMBufferSizeBuilder
1509 .create(loc, rewriter,
1510 {bufferSize, modeA, modeB, adaptor.getSpmatA(),
1511 adaptor.getDnmatB(), adaptor.getDnmatC(), computeType,
1512 pruneFlag, stream})
1513 .getResult();
1514
1515 auto bufferSizePtr1 = LLVM::GEPOp::create(
1516 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1517 ValueRange{createIndexAttrConstant(rewriter, loc, getIndexType(), 1)});
1518 auto bufferSizePtr2 = LLVM::GEPOp::create(
1519 rewriter, loc, llvmPointerType, llvmPointerType, bufferSize,
1520 ValueRange{createIndexAttrConstant(rewriter, loc, getIndexType(), 2)});
1521 auto bufferSize0 =
1522 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSize);
1523 auto bufferSize1 =
1524 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr1);
1525 auto bufferSize2 =
1526 LLVM::LoadOp::create(rewriter, loc, llvmInt64Type, bufferSizePtr2);
1527
1528 rewriter.replaceOp(op, {bufferSize0, bufferSize1, bufferSize2, stream});
1529 } else {
1530 auto computeType = genConstInt32From(
1531 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1532 bufferSize =
1533 createSpMMBufferSizeCallBuilder
1534 .create(loc, rewriter,
1535 {modeA, modeB, adaptor.getSpmatA(), adaptor.getDnmatB(),
1536 adaptor.getDnmatC(), computeType, stream})
1537 .getResult();
1538 rewriter.replaceOp(op, {bufferSize, stream});
1539 }
1540 return success();
1541}
1542
1543LogicalResult ConvertSDDMMBufferSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1544 gpu::SDDMMBufferSizeOp op, OpAdaptor adaptor,
1545 ConversionPatternRewriter &rewriter) const {
1546 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1547 failed(isAsyncWithOneDependency(rewriter, op)))
1548 return failure();
1549 Location loc = op.getLoc();
1550 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1551 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
1552 auto computeType = genConstInt32From(
1553 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1554 auto stream = adaptor.getAsyncDependencies().front();
1555 auto bufferSize =
1556 createSDDMMBufferSizeCallBuilder
1557 .create(loc, rewriter,
1558 {modeA, modeB, adaptor.getDnmatA(), adaptor.getDnmatB(),
1559 adaptor.getSpmatC(), computeType, stream})
1560 .getResult();
1561 rewriter.replaceOp(op, {bufferSize, stream});
1562 return success();
1563}
1564
1565LogicalResult ConvertSpMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1566 gpu::SpMMOp op, OpAdaptor adaptor,
1567 ConversionPatternRewriter &rewriter) const {
1568 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1569 failed(isAsyncWithOneDependency(rewriter, op)))
1570 return failure();
1571 Location loc = op.getLoc();
1572 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1573 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
1574 auto computeType = genConstInt32From(
1575 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1576
1577 auto stream = adaptor.getAsyncDependencies().front();
1578
1579 // Lower to cusparseLt if applicable
1580 if (is2To4Sparsity(op.getSpmatA())) {
1581 SmallVector<Value> pBufs;
1582 for (Value buffer : adaptor.getBuffers()) {
1583 Value pBuf = MemRefDescriptor(buffer).allocatedPtr(rewriter, loc);
1584 pBufs.push_back(pBuf);
1585 }
1586 createCuSparseLtSpMMBuilder.create(
1587 loc, rewriter,
1588 {adaptor.getSpmatA(), adaptor.getDnmatB(), adaptor.getDnmatC(),
1589 pBufs[0], pBufs[1], pBufs[2], stream});
1590 } else {
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});
1597 }
1598 rewriter.replaceOp(op, {stream});
1599 return success();
1600}
1601
1602template <typename T>
1604 converter.addConversion([&converter](T) -> Type {
1605 return LLVM::LLVMPointerType::get(&converter.getContext());
1606 });
1607}
1608
1609LogicalResult ConvertSDDMMOpToGpuRuntimeCallPattern::matchAndRewrite(
1610 gpu::SDDMMOp op, OpAdaptor adaptor,
1611 ConversionPatternRewriter &rewriter) const {
1612 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1613 failed(isAsyncWithOneDependency(rewriter, op)))
1614 return failure();
1615 Location loc = op.getLoc();
1616 auto computeType = genConstInt32From(
1617 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1618 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1619 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
1620 auto stream = adaptor.getAsyncDependencies().front();
1621 Value pBuf =
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});
1628 return success();
1629}
1630
1631LogicalResult
1632ConvertSpGEMMCreateDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1633 gpu::SpGEMMCreateDescrOp op, OpAdaptor adaptor,
1634 ConversionPatternRewriter &rewriter) const {
1635 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1636 failed(isAsyncWithOneDependency(rewriter, op)))
1637 return failure();
1638 Location loc = op.getLoc();
1639 auto stream = adaptor.getAsyncDependencies().front();
1640 Value descr = createSpGEMMCreateDescrBuilder.create(loc, rewriter, {stream})
1641 .getResult();
1642 rewriter.replaceOp(op, {descr, stream});
1643 return success();
1644}
1645
1646LogicalResult
1647ConvertSpGEMMDestroyDescrOpToGpuRuntimeCallPattern::matchAndRewrite(
1648 gpu::SpGEMMDestroyDescrOp op, OpAdaptor adaptor,
1649 ConversionPatternRewriter &rewriter) const {
1650 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1651 failed(isAsyncWithOneDependency(rewriter, op)))
1652 return failure();
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});
1658 return success();
1659}
1660
1661LogicalResult
1662ConvertSpGEMMWorkEstimationOrComputeOpToGpuRuntimeCallPattern::matchAndRewrite(
1663 gpu::SpGEMMWorkEstimationOrComputeOp op, OpAdaptor adaptor,
1664 ConversionPatternRewriter &rewriter) const {
1665 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1666 failed(isAsyncWithOneDependency(rewriter, op)))
1667 return failure();
1668 Location loc = op.getLoc();
1669 auto computeType = genConstInt32From(
1670 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1671 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1672 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
1673 auto stream = adaptor.getAsyncDependencies().front();
1674
1675 Value pBuf =
1676 MemRefDescriptor(adaptor.getBuffer()).allocatedPtr(rewriter, loc);
1677 Value bufferSizeNew;
1678
1679 if (adaptor.getKind() ==
1680 gpu::SpGEMMWorkEstimationOrComputeKind::WORK_ESTIMATION) {
1681 bufferSizeNew =
1682 createSpGEMMWorkEstimationBuilder
1683 .create(loc, rewriter,
1684 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1685 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1686 adaptor.getBufferSz(), pBuf, stream})
1687 .getResult();
1688 } else {
1689 bufferSizeNew =
1690 createSpGEMMComputeBuilder
1691 .create(loc, rewriter,
1692 {adaptor.getDesc(), modeA, modeB, adaptor.getSpmatA(),
1693 adaptor.getSpmatB(), adaptor.getSpmatC(), computeType,
1694 adaptor.getBufferSz(), pBuf, stream})
1695 .getResult();
1696 }
1697 rewriter.replaceOp(op, {bufferSizeNew, stream});
1698 return success();
1699}
1700
1701LogicalResult ConvertSpGEMMCopyOpToGpuRuntimeCallPattern::matchAndRewrite(
1702 gpu::SpGEMMCopyOp op, OpAdaptor adaptor,
1703 ConversionPatternRewriter &rewriter) const {
1704 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1705 failed(isAsyncWithOneDependency(rewriter, op)))
1706 return failure();
1707 Location loc = op.getLoc();
1708 auto computeType = genConstInt32From(
1709 rewriter, loc, getCuSparseDataTypeFrom(adaptor.getComputeType()));
1710 auto modeA = genConstInt32From(rewriter, loc, adaptor.getModeA());
1711 auto modeB = genConstInt32From(rewriter, loc, adaptor.getModeB());
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});
1718 return success();
1719}
1720
1721LogicalResult ConvertSpMatGetSizeOpToGpuRuntimeCallPattern::matchAndRewrite(
1722 gpu::SpMatGetSizeOp op, OpAdaptor adaptor,
1723 ConversionPatternRewriter &rewriter) const {
1724 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1725 failed(isAsyncWithOneDependency(rewriter, op)))
1726 return failure();
1727 Location loc = op.getLoc();
1728 auto stream = adaptor.getAsyncDependencies().front();
1729
1730 auto three = createIndexAttrConstant(rewriter, loc, getIndexType(), 3);
1731 auto buffer = LLVM::AllocaOp::create(rewriter, loc, llvmPointerType,
1732 llvmInt64Type, three, /*alignment=*/16);
1733
1734 auto rowsPtr = LLVM::GEPOp::create(
1735 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1736 ValueRange{createIndexAttrConstant(rewriter, loc, getIndexType(), 0)});
1737 auto colsPtr = LLVM::GEPOp::create(
1738 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1739 ValueRange{createIndexAttrConstant(rewriter, loc, getIndexType(), 1)});
1740 auto nnzsPtr = LLVM::GEPOp::create(
1741 rewriter, loc, llvmPointerType, llvmPointerType, buffer,
1742 ValueRange{createIndexAttrConstant(rewriter, loc, getIndexType(), 2)});
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);
1748
1749 rewriter.replaceOp(op, {rows, cols, nnzs, stream});
1750 return success();
1751}
1752
1753LogicalResult ConvertSetCsrPointersOpToGpuRuntimeCallPattern::matchAndRewrite(
1754 gpu::SetCsrPointersOp op, OpAdaptor adaptor,
1755 ConversionPatternRewriter &rewriter) const {
1756 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1757 failed(isAsyncWithOneDependency(rewriter, op)))
1758 return failure();
1759 Location loc = op.getLoc();
1760 auto stream = adaptor.getAsyncDependencies().front();
1761 Value pPos =
1762 MemRefDescriptor(adaptor.getPositions()).allocatedPtr(rewriter, loc);
1763 Value pCrd =
1764 MemRefDescriptor(adaptor.getCoordinates()).allocatedPtr(rewriter, loc);
1765 Value pVal =
1766 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1767 createSetCsrPointersBuilder.create(
1768 loc, rewriter, {adaptor.getSpmat(), pPos, pCrd, pVal, stream});
1769 rewriter.replaceOp(op, {stream});
1770 return success();
1771}
1772
1773LogicalResult ConvertCreateCscOpToGpuRuntimeCallPattern::matchAndRewrite(
1774 gpu::CreateCscOp op, OpAdaptor adaptor,
1775 ConversionPatternRewriter &rewriter) const {
1776 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1777 failed(isAsyncWithOneDependency(rewriter, op)))
1778 return failure();
1779 Location loc = op.getLoc();
1780 auto stream = adaptor.getAsyncDependencies().front();
1781 Value pColPos =
1782 MemRefDescriptor(adaptor.getColPos()).allocatedPtr(rewriter, loc);
1783 Value pRowIdxs =
1784 MemRefDescriptor(adaptor.getRowIdxs()).allocatedPtr(rewriter, loc);
1785 Value pValues =
1786 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1787 Type pType =
1788 llvm::cast<MemRefType>(op.getColPos().getType()).getElementType();
1789 Type iType =
1790 llvm::cast<MemRefType>(op.getRowIdxs().getType()).getElementType();
1791 Type dType =
1792 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1793 auto ptp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(pType));
1794 auto itp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(iType));
1795 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1796 auto handle =
1797 createCscCallBuilder
1798 .create(loc, rewriter,
1799 {adaptor.getRows(), adaptor.getCols(), adaptor.getNnz(),
1800 pColPos, pRowIdxs, pValues, ptp, itp, dtp, stream})
1801 .getResult();
1802 rewriter.replaceOp(op, {handle, stream});
1803 return success();
1804}
1805
1806LogicalResult ConvertCreateBsrOpToGpuRuntimeCallPattern::matchAndRewrite(
1807 gpu::CreateBsrOp op, OpAdaptor adaptor,
1808 ConversionPatternRewriter &rewriter) const {
1809 if (failed(areAllLLVMTypes(op, adaptor.getOperands(), rewriter)) ||
1810 failed(isAsyncWithOneDependency(rewriter, op)))
1811 return failure();
1812 Location loc = op.getLoc();
1813 auto stream = adaptor.getAsyncDependencies().front();
1814 Value pRowPos =
1815 MemRefDescriptor(adaptor.getBRowPos()).allocatedPtr(rewriter, loc);
1816 Value pColIdxs =
1817 MemRefDescriptor(adaptor.getBColIdxs()).allocatedPtr(rewriter, loc);
1818 Value pValues =
1819 MemRefDescriptor(adaptor.getValues()).allocatedPtr(rewriter, loc);
1820 Type pType =
1821 llvm::cast<MemRefType>(op.getBRowPos().getType()).getElementType();
1822 Type iType =
1823 llvm::cast<MemRefType>(op.getBColIdxs().getType()).getElementType();
1824 Type dType =
1825 llvm::cast<MemRefType>(op.getValues().getType()).getElementType();
1826 auto ptp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(pType));
1827 auto itp = genConstInt32From(rewriter, loc, getCuSparseIndexTypeFrom(iType));
1828 auto dtp = genConstInt32From(rewriter, loc, getCuSparseDataTypeFrom(dType));
1829 auto handle =
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})
1835 .getResult();
1836 rewriter.replaceOp(op, {handle, stream});
1837 return success();
1838}
1839
1841 LLVMTypeConverter &converter, RewritePatternSet &patterns,
1842 bool kernelBarePtrCallConv, bool kernelIntersperseSizeCallConv) {
1847
1848 // Higher benefit so this pattern wins over the structural async.yield
1849 // rewriter from populateAsyncStructuralTypeConversionsAndLegality on yields
1850 // with gpu.async.token operands. The structural rewriter would silently
1851 // retype operands without recording an event on the underlying stream.
1852 patterns.add<ConvertAsyncYieldToGpuRuntimeCallPattern>(converter,
1853 /*benefit=*/2);
1854
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);
1887}
1888
1889//===----------------------------------------------------------------------===//
1890// GPUModuleOp convert to LLVM op interface
1891//===----------------------------------------------------------------------===//
1892
1893namespace {
1894struct GPUModuleOpConvertToLLVMInterface
1895 : public ConvertToLLVMOpInterface::ExternalModel<
1896 GPUModuleOpConvertToLLVMInterface, gpu::GPUModuleOp> {
1897 /// Get the conversion patterns from the target attribute.
1898 void getConvertToLLVMConversionAttrs(
1900};
1901} // namespace
1902
1903void GPUModuleOpConvertToLLVMInterface::getConvertToLLVMConversionAttrs(
1904 Operation *op, SmallVectorImpl<ConvertToLLVMAttrInterface> &attrs) const {
1905 auto module = cast<gpu::GPUModuleOp>(op);
1906 ArrayAttr targetsAttr = module.getTargetsAttr();
1907 // Fail if there are no target attributes or there is more than one target.
1908 if (!targetsAttr || targetsAttr.size() != 1)
1909 return;
1910 if (auto patternAttr = dyn_cast<ConvertToLLVMAttrInterface>(targetsAttr[0]))
1911 attrs.push_back(patternAttr);
1912}
1913
1915 registry.addExtension(+[](MLIRContext *ctx, gpu::GPUDialect *dialect) {
1916 gpu::GPUModuleOp::attachInterface<GPUModuleOpConvertToLLVMInterface>(*ctx);
1917 });
1918}
return success()
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.
ArrayAttr()
b getContext())
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)
Definition Builders.cpp:75
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
Type getIndexType() const
Gets the MLIR type wrapping the LLVM integer type whose bit width is defined by the used type convert...
Definition Pattern.cpp:38
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.
Definition Pattern.cpp:64
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...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
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.
Definition Builders.h:210
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...
Definition Builders.h:249
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
void print(raw_ostream &os, const OpPrintingFlags &flags={})
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
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...
Definition Types.h:74
bool isF64() const
Definition Types.cpp:41
bool isF32() const
Definition Types.cpp:40
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
bool isF16() const
Definition Types.cpp:38
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isBF16() const
Definition Types.cpp:37
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Definition Value.h:108
Type getType() const
Return the type of this value.
Definition Value.h:105
user_range getUsers() const
Definition Value.h:218
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
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...
Definition Pattern.cpp:58
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
void registerConvertGpuToLLVMInterface(DialectRegistry &registry)
Registers the ConvertToLLVMOpInterface interface on the gpu::GPUModuleOP operation.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:733
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 &region, 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 &registry)
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