1 //===- LowerGpuOpsToNVVMOps.cpp - MLIR GPU to NVVM 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 generate NVVMIR operations for higher-level
10 // GPU operations.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "mlir/Conversion/GPUToNVVM/GPUToNVVMPass.h"
15 
16 #include "mlir/Conversion/ArithmeticToLLVM/ArithmeticToLLVM.h"
17 #include "mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h"
18 #include "mlir/Conversion/FuncToLLVM/ConvertFuncToLLVM.h"
19 #include "mlir/Conversion/LLVMCommon/ConversionTarget.h"
20 #include "mlir/Conversion/LLVMCommon/LoweringOptions.h"
21 #include "mlir/Conversion/LLVMCommon/TypeConverter.h"
22 #include "mlir/Conversion/MemRefToLLVM/MemRefToLLVM.h"
23 #include "mlir/Dialect/Arithmetic/IR/Arithmetic.h"
24 #include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
25 #include "mlir/Dialect/Func/IR/FuncOps.h"
26 #include "mlir/Dialect/GPU/GPUDialect.h"
27 #include "mlir/Dialect/GPU/Passes.h"
28 #include "mlir/Dialect/LLVMIR/NVVMDialect.h"
29 #include "mlir/Dialect/Math/IR/Math.h"
30 #include "mlir/Dialect/MemRef/IR/MemRef.h"
31 #include "mlir/IR/BlockAndValueMapping.h"
32 #include "mlir/Transforms/DialectConversion.h"
33 #include "mlir/Transforms/GreedyPatternRewriteDriver.h"
34 #include "llvm/Support/FormatVariadic.h"
35 
36 #include "../GPUCommon/GPUOpsLowering.h"
37 #include "../GPUCommon/IndexIntrinsicsOpLowering.h"
38 #include "../GPUCommon/OpToFuncCallLowering.h"
39 #include "../PassDetail.h"
40 
41 using namespace mlir;
42 
43 namespace {
44 
45 /// NVVM memory space identifiers.
46 enum NVVMMemorySpace {
47   /// Global memory space identifier.
48   kGlobalMemorySpace = 1,
49   /// Shared memory space identifier.
50   kSharedMemorySpace = 3
51 };
52 
53 /// Convert gpu dialect shfl mode enum to the equivalent nvvm one.
54 static NVVM::ShflKind convertShflKind(gpu::ShuffleMode mode) {
55   switch (mode) {
56   case gpu::ShuffleMode::XOR:
57     return NVVM::ShflKind::bfly;
58   case gpu::ShuffleMode::UP:
59     return NVVM::ShflKind::up;
60   case gpu::ShuffleMode::DOWN:
61     return NVVM::ShflKind::down;
62   case gpu::ShuffleMode::IDX:
63     return NVVM::ShflKind::idx;
64   }
65   llvm_unreachable("unknown shuffle mode");
66 }
67 
68 struct GPUShuffleOpLowering : public ConvertOpToLLVMPattern<gpu::ShuffleOp> {
69   using ConvertOpToLLVMPattern<gpu::ShuffleOp>::ConvertOpToLLVMPattern;
70 
71   /// Lowers a shuffle to the corresponding NVVM op.
72   ///
73   /// Convert the `width` argument into an activeMask (a bitmask which specifies
74   /// which threads participate in the shuffle) and a maskAndClamp (specifying
75   /// the highest lane which participates in the shuffle).
76   ///
77   ///     %one = llvm.constant(1 : i32) : i32
78   ///     %minus_one = llvm.constant(-1 : i32) : i32
79   ///     %thirty_two = llvm.constant(32 : i32) : i32
80   ///     %num_lanes = llvm.sub %thirty_two, %width : i32
81   ///     %active_mask = llvm.lshr %minus_one, %num_lanes : i32
82   ///     %mask_and_clamp = llvm.sub %width, %one : i32
83   ///     %shfl = nvvm.shfl.sync.bfly %active_mask, %value, %offset,
84   ///         %mask_and_clamp : !llvm<"{ float, i1 }">
85   ///     %shfl_value = llvm.extractvalue %shfl[0 : index] :
86   ///         !llvm<"{ float, i1 }">
87   ///     %shfl_pred = llvm.extractvalue %shfl[1 : index] :
88   ///         !llvm<"{ float, i1 }">
89   LogicalResult
90   matchAndRewrite(gpu::ShuffleOp op, OpAdaptor adaptor,
91                   ConversionPatternRewriter &rewriter) const override {
92     Location loc = op->getLoc();
93 
94     auto valueTy = adaptor.value().getType();
95     auto int32Type = IntegerType::get(rewriter.getContext(), 32);
96     auto predTy = IntegerType::get(rewriter.getContext(), 1);
97     auto resultTy = LLVM::LLVMStructType::getLiteral(rewriter.getContext(),
98                                                      {valueTy, predTy});
99 
100     Value one = rewriter.create<LLVM::ConstantOp>(
101         loc, int32Type, rewriter.getI32IntegerAttr(1));
102     Value minusOne = rewriter.create<LLVM::ConstantOp>(
103         loc, int32Type, rewriter.getI32IntegerAttr(-1));
104     Value thirtyTwo = rewriter.create<LLVM::ConstantOp>(
105         loc, int32Type, rewriter.getI32IntegerAttr(32));
106     Value numLeadInactiveLane = rewriter.create<LLVM::SubOp>(
107         loc, int32Type, thirtyTwo, adaptor.width());
108     // Bit mask of active lanes: `(-1) >> (32 - activeWidth)`.
109     Value activeMask = rewriter.create<LLVM::LShrOp>(loc, int32Type, minusOne,
110                                                      numLeadInactiveLane);
111     Value maskAndClamp;
112     if (op.mode() == gpu::ShuffleMode::UP) {
113       // Clamp lane: `32 - activeWidth`
114       maskAndClamp = numLeadInactiveLane;
115     } else {
116       // Clamp lane: `activeWidth - 1`
117       maskAndClamp =
118           rewriter.create<LLVM::SubOp>(loc, int32Type, adaptor.width(), one);
119     }
120 
121     auto returnValueAndIsValidAttr = rewriter.getUnitAttr();
122     Value shfl = rewriter.create<NVVM::ShflOp>(
123         loc, resultTy, activeMask, adaptor.value(), adaptor.offset(),
124         maskAndClamp, convertShflKind(op.mode()), returnValueAndIsValidAttr);
125     Value shflValue = rewriter.create<LLVM::ExtractValueOp>(
126         loc, valueTy, shfl, rewriter.getIndexArrayAttr(0));
127     Value isActiveSrcLane = rewriter.create<LLVM::ExtractValueOp>(
128         loc, predTy, shfl, rewriter.getIndexArrayAttr(1));
129 
130     rewriter.replaceOp(op, {shflValue, isActiveSrcLane});
131     return success();
132   }
133 };
134 
135 struct GPUAsyncCopyLowering
136     : public ConvertOpToLLVMPattern<gpu::DeviceAsyncCopyOp> {
137   using ConvertOpToLLVMPattern<gpu::DeviceAsyncCopyOp>::ConvertOpToLLVMPattern;
138 
139   LogicalResult
140   matchAndRewrite(gpu::DeviceAsyncCopyOp op, OpAdaptor adaptor,
141                   ConversionPatternRewriter &rewriter) const override {
142     Location loc = op->getLoc();
143     auto dstMemrefType = op.dst().getType().cast<MemRefType>();
144     Value dstPtr = getStridedElementPtr(loc, dstMemrefType, adaptor.dst(),
145                                         adaptor.dstIndices(), rewriter);
146     auto i8Ty = IntegerType::get(op.getContext(), 8);
147     auto dstPointerType =
148         LLVM::LLVMPointerType::get(i8Ty, dstMemrefType.getMemorySpaceAsInt());
149     dstPtr = rewriter.create<LLVM::BitcastOp>(loc, dstPointerType, dstPtr);
150 
151     auto srcMemrefType = op.src().getType().cast<MemRefType>();
152 
153     Value scrPtr = getStridedElementPtr(loc, srcMemrefType, adaptor.src(),
154                                         adaptor.srcIndices(), rewriter);
155     auto srcPointerType =
156         LLVM::LLVMPointerType::get(i8Ty, srcMemrefType.getMemorySpaceAsInt());
157     scrPtr = rewriter.create<LLVM::BitcastOp>(loc, srcPointerType, scrPtr);
158     // Intrinsics takes a global pointer so we need an address space cast.
159     auto srcPointerGlobalType =
160         LLVM::LLVMPointerType::get(i8Ty, NVVMMemorySpace::kGlobalMemorySpace);
161     scrPtr = rewriter.create<LLVM::AddrSpaceCastOp>(loc, srcPointerGlobalType,
162                                                     scrPtr);
163     int64_t numElements = adaptor.numElements().getZExtValue();
164     int64_t sizeInBytes =
165         (dstMemrefType.getElementTypeBitWidth() / 8) * numElements;
166     rewriter.create<NVVM::CpAsyncOp>(loc, dstPtr, scrPtr,
167                                      rewriter.getI32IntegerAttr(sizeInBytes),
168                                      /*bypassL1=*/UnitAttr());
169 
170     // Drop the result token.
171     Value zero = rewriter.create<LLVM::ConstantOp>(
172         op->getLoc(), IntegerType::get(op.getContext(), 32),
173         rewriter.getI32IntegerAttr(0));
174     rewriter.replaceOp(op, zero);
175     return success();
176   }
177 };
178 
179 struct GPUAsyncCreateGroupLowering
180     : public ConvertOpToLLVMPattern<gpu::DeviceAsyncCreateGroupOp> {
181   using ConvertOpToLLVMPattern<
182       gpu::DeviceAsyncCreateGroupOp>::ConvertOpToLLVMPattern;
183 
184   LogicalResult
185   matchAndRewrite(gpu::DeviceAsyncCreateGroupOp op, OpAdaptor adaptor,
186                   ConversionPatternRewriter &rewriter) const override {
187     rewriter.create<NVVM::CpAsyncCommitGroupOp>(op.getLoc());
188     // Drop the result token.
189     Value zero = rewriter.create<LLVM::ConstantOp>(
190         op->getLoc(), IntegerType::get(op.getContext(), 32),
191         rewriter.getI32IntegerAttr(0));
192     rewriter.replaceOp(op, zero);
193     return success();
194   }
195 };
196 
197 struct GPUAsyncWaitLowering
198     : public ConvertOpToLLVMPattern<gpu::DeviceAsyncWaitOp> {
199   using ConvertOpToLLVMPattern<gpu::DeviceAsyncWaitOp>::ConvertOpToLLVMPattern;
200 
201   LogicalResult
202   matchAndRewrite(gpu::DeviceAsyncWaitOp op, OpAdaptor adaptor,
203                   ConversionPatternRewriter &rewriter) const override {
204     // If numGroup is not present pick 0 as a conservative correct value.
205     int32_t numGroups = adaptor.numGroups() ? *adaptor.numGroups() : 0;
206     rewriter.create<NVVM::CpAsyncWaitGroupOp>(op.getLoc(), numGroups);
207     rewriter.eraseOp(op);
208     return success();
209   }
210 };
211 
212 struct GPULaneIdOpToNVVM : ConvertOpToLLVMPattern<gpu::LaneIdOp> {
213   using ConvertOpToLLVMPattern<gpu::LaneIdOp>::ConvertOpToLLVMPattern;
214 
215   LogicalResult
216   matchAndRewrite(gpu::LaneIdOp op, gpu::LaneIdOp::Adaptor adaptor,
217                   ConversionPatternRewriter &rewriter) const override {
218     auto loc = op->getLoc();
219     MLIRContext *context = rewriter.getContext();
220     Value newOp = rewriter.create<NVVM::LaneIdOp>(loc, rewriter.getI32Type());
221     // Truncate or extend the result depending on the index bitwidth specified
222     // by the LLVMTypeConverter options.
223     const unsigned indexBitwidth = getTypeConverter()->getIndexTypeBitwidth();
224     if (indexBitwidth > 32) {
225       newOp = rewriter.create<LLVM::SExtOp>(
226           loc, IntegerType::get(context, indexBitwidth), newOp);
227     } else if (indexBitwidth < 32) {
228       newOp = rewriter.create<LLVM::TruncOp>(
229           loc, IntegerType::get(context, indexBitwidth), newOp);
230     }
231     rewriter.replaceOp(op, {newOp});
232     return success();
233   }
234 };
235 
236 /// Import the GPU Ops to NVVM Patterns.
237 #include "GPUToNVVM.cpp.inc"
238 
239 /// A pass that replaces all occurrences of GPU device operations with their
240 /// corresponding NVVM equivalent.
241 ///
242 /// This pass only handles device code and is not meant to be run on GPU host
243 /// code.
244 struct LowerGpuOpsToNVVMOpsPass
245     : public ConvertGpuOpsToNVVMOpsBase<LowerGpuOpsToNVVMOpsPass> {
246   LowerGpuOpsToNVVMOpsPass() = default;
247   LowerGpuOpsToNVVMOpsPass(unsigned indexBitwidth) {
248     this->indexBitwidth = indexBitwidth;
249   }
250 
251   void runOnOperation() override {
252     gpu::GPUModuleOp m = getOperation();
253 
254     /// Customize the bitwidth used for the device side index computations.
255     LowerToLLVMOptions options(
256         m.getContext(),
257         DataLayout(cast<DataLayoutOpInterface>(m.getOperation())));
258     options.emitCWrappers = true;
259     if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)
260       options.overrideIndexBitwidth(indexBitwidth);
261 
262     /// MemRef conversion for GPU to NVVM lowering. The GPU dialect uses memory
263     /// space 5 for private memory attributions, but NVVM represents private
264     /// memory allocations as local `alloca`s in the default address space. This
265     /// converter drops the private memory space to support the use case above.
266     LLVMTypeConverter converter(m.getContext(), options);
267     converter.addConversion([&](MemRefType type) -> Optional<Type> {
268       if (type.getMemorySpaceAsInt() !=
269           gpu::GPUDialect::getPrivateAddressSpace())
270         return llvm::None;
271       return converter.convertType(MemRefType::Builder(type).setMemorySpace(0));
272     });
273     /// device-side async tokens cannot be materialized in nvvm. We just convert
274     /// them to a dummy i32 type in order to easily drop them during conversion.
275     converter.addConversion([&](gpu::DeviceAsyncTokenType type) -> Type {
276       return converter.convertType(IntegerType::get(type.getContext(), 32));
277     });
278     // Lowering for MMAMatrixType.
279     converter.addConversion([&](gpu::MMAMatrixType type) -> Type {
280       return convertMMAToLLVMType(type);
281     });
282     RewritePatternSet patterns(m.getContext());
283     RewritePatternSet llvmPatterns(m.getContext());
284 
285     // Apply in-dialect lowering first. In-dialect lowering will replace ops
286     // which need to be lowered further, which is not supported by a single
287     // conversion pass.
288     populateGpuRewritePatterns(patterns);
289     (void)applyPatternsAndFoldGreedily(m, std::move(patterns));
290 
291     arith::populateArithmeticToLLVMConversionPatterns(converter, llvmPatterns);
292     cf::populateControlFlowToLLVMConversionPatterns(converter, llvmPatterns);
293     populateFuncToLLVMConversionPatterns(converter, llvmPatterns);
294     populateMemRefToLLVMConversionPatterns(converter, llvmPatterns);
295     populateGpuToNVVMConversionPatterns(converter, llvmPatterns);
296     populateGpuWMMAToNVVMConversionPatterns(converter, llvmPatterns);
297     LLVMConversionTarget target(getContext());
298     configureGpuToNVVMConversionLegality(target);
299     if (failed(applyPartialConversion(m, target, std::move(llvmPatterns))))
300       signalPassFailure();
301   }
302 };
303 
304 } // namespace
305 
306 void mlir::configureGpuToNVVMConversionLegality(ConversionTarget &target) {
307   target.addIllegalOp<func::FuncOp>();
308   target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
309   target.addLegalDialect<::mlir::NVVM::NVVMDialect>();
310   target.addIllegalDialect<gpu::GPUDialect>();
311   target.addIllegalOp<LLVM::CosOp, LLVM::ExpOp, LLVM::Exp2Op, LLVM::FAbsOp,
312                       LLVM::FCeilOp, LLVM::FFloorOp, LLVM::LogOp, LLVM::Log10Op,
313                       LLVM::Log2Op, LLVM::PowOp, LLVM::SinOp, LLVM::SqrtOp>();
314 
315   // TODO: Remove once we support replacing non-root ops.
316   target.addLegalOp<gpu::YieldOp, gpu::GPUModuleOp, gpu::ModuleEndOp>();
317 }
318 
319 void mlir::populateGpuToNVVMConversionPatterns(LLVMTypeConverter &converter,
320                                                RewritePatternSet &patterns) {
321   populateWithGenerated(patterns);
322   patterns
323       .add<GPUIndexIntrinsicOpLowering<gpu::ThreadIdOp, NVVM::ThreadIdXOp,
324                                        NVVM::ThreadIdYOp, NVVM::ThreadIdZOp>,
325            GPUIndexIntrinsicOpLowering<gpu::BlockDimOp, NVVM::BlockDimXOp,
326                                        NVVM::BlockDimYOp, NVVM::BlockDimZOp>,
327            GPUIndexIntrinsicOpLowering<gpu::BlockIdOp, NVVM::BlockIdXOp,
328                                        NVVM::BlockIdYOp, NVVM::BlockIdZOp>,
329            GPUIndexIntrinsicOpLowering<gpu::GridDimOp, NVVM::GridDimXOp,
330                                        NVVM::GridDimYOp, NVVM::GridDimZOp>,
331            GPULaneIdOpToNVVM, GPUShuffleOpLowering, GPUReturnOpLowering>(
332           converter);
333 
334   // Explicitly drop memory space when lowering private memory
335   // attributions since NVVM models it as `alloca`s in the default
336   // memory space and does not support `alloca`s with addrspace(5).
337   patterns.add<GPUFuncOpLowering>(
338       converter, /*allocaAddrSpace=*/0,
339       StringAttr::get(&converter.getContext(),
340                       NVVM::NVVMDialect::getKernelFuncAttrName()));
341 
342   patterns.add<OpToFuncCallLowering<math::AbsOp>>(converter, "__nv_fabsf",
343                                                   "__nv_fabs");
344   patterns.add<OpToFuncCallLowering<math::AtanOp>>(converter, "__nv_atanf",
345                                                    "__nv_atan");
346   patterns.add<OpToFuncCallLowering<math::Atan2Op>>(converter, "__nv_atan2f",
347                                                     "__nv_atan2");
348   patterns.add<OpToFuncCallLowering<math::CeilOp>>(converter, "__nv_ceilf",
349                                                    "__nv_ceil");
350   patterns.add<OpToFuncCallLowering<math::CosOp>>(converter, "__nv_cosf",
351                                                   "__nv_cos");
352   patterns.add<OpToFuncCallLowering<math::ExpOp>>(converter, "__nv_expf",
353                                                   "__nv_exp");
354   patterns.add<OpToFuncCallLowering<math::Exp2Op>>(converter, "__nv_exp2f",
355                                                    "__nv_exp2");
356   patterns.add<OpToFuncCallLowering<math::ExpM1Op>>(converter, "__nv_expm1f",
357                                                     "__nv_expm1");
358   patterns.add<OpToFuncCallLowering<math::FloorOp>>(converter, "__nv_floorf",
359                                                     "__nv_floor");
360   patterns.add<OpToFuncCallLowering<math::LogOp>>(converter, "__nv_logf",
361                                                   "__nv_log");
362   patterns.add<OpToFuncCallLowering<math::Log1pOp>>(converter, "__nv_log1pf",
363                                                     "__nv_log1p");
364   patterns.add<OpToFuncCallLowering<math::Log10Op>>(converter, "__nv_log10f",
365                                                     "__nv_log10");
366   patterns.add<OpToFuncCallLowering<math::Log2Op>>(converter, "__nv_log2f",
367                                                    "__nv_log2");
368   patterns.add<OpToFuncCallLowering<math::PowFOp>>(converter, "__nv_powf",
369                                                    "__nv_pow");
370   patterns.add<OpToFuncCallLowering<math::RsqrtOp>>(converter, "__nv_rsqrtf",
371                                                     "__nv_rsqrt");
372   patterns.add<OpToFuncCallLowering<math::SinOp>>(converter, "__nv_sinf",
373                                                   "__nv_sin");
374   patterns.add<OpToFuncCallLowering<math::SqrtOp>>(converter, "__nv_sqrtf",
375                                                    "__nv_sqrt");
376   patterns.add<OpToFuncCallLowering<math::TanhOp>>(converter, "__nv_tanhf",
377                                                    "__nv_tanh");
378   patterns.add<GPUAsyncCopyLowering, GPUAsyncCreateGroupLowering,
379                GPUAsyncWaitLowering>(converter);
380 }
381 
382 std::unique_ptr<OperationPass<gpu::GPUModuleOp>>
383 mlir::createLowerGpuOpsToNVVMOpsPass(unsigned indexBitwidth) {
384   return std::make_unique<LowerGpuOpsToNVVMOpsPass>(indexBitwidth);
385 }
386