1 //===------ WmmaOpsToNVVM.cpp - WMMA LD/ST/Compute to NVVM lowering -------===//
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 contains definitions of patterns to lower GPU Subgroup MMA ops to
10 // NVVM Dialect.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "mlir/Conversion/StandardToLLVM/ConvertStandardToLLVM.h"
15 #include "mlir/Dialect/GPU/GPUDialect.h"
16 #include "mlir/Dialect/LLVMIR/LLVMDialect.h"
17 #include "mlir/Dialect/LLVMIR/NVVMDialect.h"
18 
19 using namespace mlir;
20 
21 namespace {
22 
23 /// Contains all the common LLVM types which are used across the lowerings of
24 /// GPU subgroup ops to NVVM dialect.
25 struct CommonLLVMAndBuiltInMLIRTypes {
26 public:
27   CommonLLVMAndBuiltInMLIRTypes(MLIRContext *context) {
28     numHalfsInOpFrags.resize(4);
29     numHalfsInOpFrags[A] = 8;
30     numHalfsInOpFrags[B] = 8;
31     numHalfsInOpFrags[C] = 4;
32     i32Ty = IntegerType::get(context, 32);
33     f16Ty = FloatType::getF16(context);
34     f32Ty = FloatType::getF32(context);
35     f16x2Ty = VectorType::get(2, f16Ty);
36     fragArrayABTy = LLVM::LLVMStructType::getLiteral(
37         context, SmallVector<Type>(8, f16x2Ty));
38     fragArrayCDTy = LLVM::LLVMStructType::getLiteral(
39         context, SmallVector<Type>(4, f16x2Ty));
40     fragArrayCDF32Ty =
41         LLVM::LLVMStructType::getLiteral(context, SmallVector<Type>(8, f32Ty));
42   };
43 
44   Type i32Ty;
45   Type f16Ty;
46   Type f32Ty;
47   Type f16x2Ty;
48   /// Type for the fragment of A and B operands that a single thread holds for
49   /// fp16 data type in a WMMA operation of the form D = (alpha*(A*B)) +
50   /// (beta*C).
51   Type fragArrayABTy;
52   /// Type for the fragment of C and D operands that a single thread holds for
53   /// fp16 data type in a WMMA operation of the form D = (alpha*(A*B)) +
54   /// (beta*C).
55   Type fragArrayCDTy;
56   /// Type for the fragment of C and D operands that a single thread holds for
57   /// fp32 data type in a WMMA operation of the form D = (alpha*(A*B)) +
58   /// (beta*C).
59   Type fragArrayCDF32Ty;
60   /// Represents the number of f16 elements a single thread holds in a WMMA
61   /// operation of the form D = (alpha*(A*B)) + (beta*C) .
62   SmallVector<unsigned, 4> numHalfsInOpFrags;
63   /// Represents the operands of a MMA operation of the form D = (alpha*(A*B)) +
64   /// (beta*C).
65   enum OperandMap { A, B, C };
66 };
67 
68 /// Checks if all the operands of the op being lowered are of LLVM Types. The
69 /// types are expected to be converted by the `LLVMTypeConverter` before the op
70 /// is actually lowered. If the type of an operands is not already converted it
71 /// hints a missing typeConversion and failure is returned in that case.
72 static LogicalResult areAllLLVMTypes(Operation *op, ValueRange operands,
73                                      ConversionPatternRewriter &rewriter) {
74   if (!llvm::all_of(operands, [](Value value) {
75         return LLVM::isCompatibleType(value.getType());
76       })) {
77     return rewriter.notifyMatchFailure(
78         op, "cannot convert if operands aren't of LLVM type.");
79   }
80 
81   return success();
82 }
83 
84 /// Error string to emit when unimplemented WMMA variant is encountered.
85 static constexpr StringRef kInvalidCaseStr =
86     "Unimplemented WMMA variant, Only M16N16K16 version implemented.";
87 
88 /// This class implements the conversion of GPU MMA loadOp to wmma.load op
89 /// in the NVVM dialect. The conversion not only emits the NVVM op but also
90 /// emits code that is necessary to store the data in the destination memref
91 /// after it has been loaded.
92 struct WmmaLoadOpToNVVMLowering
93     : public ConvertOpToLLVMPattern<gpu::SubgroupMmaLoadMatrixOp>,
94       private CommonLLVMAndBuiltInMLIRTypes {
95 public:
96   explicit WmmaLoadOpToNVVMLowering(LLVMTypeConverter &typeConverter)
97       : ConvertOpToLLVMPattern<gpu::SubgroupMmaLoadMatrixOp>(typeConverter),
98         CommonLLVMAndBuiltInMLIRTypes(&this->getTypeConverter()->getContext()) {
99   }
100 
101   LogicalResult
102   matchAndRewrite(gpu::SubgroupMmaLoadMatrixOp subgroupMmaLoadMatrixOp,
103                   ArrayRef<Value> operands,
104                   ConversionPatternRewriter &rewriter) const override {
105     Operation *op = subgroupMmaLoadMatrixOp.getOperation();
106     if (failed(areAllLLVMTypes(op, operands, rewriter)))
107       return failure();
108 
109     unsigned indexTypeBitwidth =
110         this->getTypeConverter()->getIndexTypeBitwidth();
111 
112     // The corresponding intrinsics expects leadDimension to be a 32-bit
113     // integer, so all the calculations of linearizing the load address
114     // must also follow this restriction.
115     if (indexTypeBitwidth != 32)
116       return rewriter.notifyMatchFailure(
117           op, "Expected indices to the memref to be 32-bit wide.");
118 
119     // Source memref of the original op.
120     MemRefType srcMemrefType =
121         subgroupMmaLoadMatrixOp.srcMemref().getType().cast<MemRefType>();
122     Location loc = op->getLoc();
123 
124     auto leadDimension = subgroupMmaLoadMatrixOp.leadDimensionAttr();
125 
126     // MemRefDescriptor to extract alignedPtr and offset.
127     MemRefDescriptor promotedSrcOp(
128         gpu::SubgroupMmaLoadMatrixOpAdaptor(operands).srcMemref());
129 
130     // Emit ops which compute the load offset using `srcOffsetI`,
131     // `srcOffsetJ`. The actualOffset is (memrefOffset + (alignedPtr +
132     // ((leadDimension * srcOffsetI) + srcOffsetJ)). The memrefs here are
133     // assumed to be normalized and hence the simple conversion works.
134     SmallVector<Value> indices(subgroupMmaLoadMatrixOp.indices());
135     Value srcOffsetIVal = indices[0];
136     Value srcOffsetJVal = indices[1];
137     Value leadingDim32 =
138         rewriter.create<LLVM::ConstantOp>(loc, i32Ty, leadDimension);
139     Value numElemsLeadDim =
140         rewriter.create<LLVM::MulOp>(loc, i32Ty, leadingDim32, srcOffsetIVal);
141     Value loadOffset = rewriter.create<LLVM::AddOp>(loc, i32Ty, numElemsLeadDim,
142                                                     srcOffsetJVal);
143 
144     Value promotedSrcOpToUse;
145     promotedSrcOpToUse = promotedSrcOp.offset(rewriter, loc);
146     Value actualOffset = rewriter.create<LLVM::AddOp>(loc, i32Ty, loadOffset,
147                                                       promotedSrcOpToUse);
148     Value loadAddress = rewriter.create<LLVM::GEPOp>(
149         loc,
150         LLVM::LLVMPointerType::get(f16Ty, srcMemrefType.getMemorySpaceAsInt()),
151         promotedSrcOp.alignedPtr(rewriter, loc), ArrayRef<Value>{actualOffset});
152 
153     // Bitcast the base address pointer of the destination memref, So that
154     // values can be stored in chunks of 32-bits and semantics match with the
155     // intrinsic exposed by NVPTX backend.
156     Value loadAddressCasted = rewriter.create<LLVM::BitcastOp>(
157         loc,
158         LLVM::LLVMPointerType::get(i32Ty, srcMemrefType.getMemorySpaceAsInt()),
159         loadAddress);
160 
161     // Get the shape of the MMAMatrix type being returned. The shape will
162     // choose which intrinsic this op will be lowered to.
163     gpu::MMAMatrixType retType =
164         subgroupMmaLoadMatrixOp.res().getType().cast<gpu::MMAMatrixType>();
165     ArrayRef<int64_t> retTypeShape = retType.getShape();
166 
167     Type resType;
168     StringRef operandStr = retType.getOperand();
169     if (operandStr.equals("AOp") || operandStr.equals("BOp")) {
170       resType = fragArrayABTy;
171     } else {
172       if (srcMemrefType.getElementType().isF16())
173         resType = fragArrayCDTy;
174       else if (srcMemrefType.getElementType().isF32())
175         resType = fragArrayCDF32Ty;
176       else
177         return failure();
178     }
179 
180     // Create nvvm.mma_load op according to the operand types.
181     SmallVector<Value, 2> loadOpOperands({loadAddressCasted, leadingDim32});
182     if (operandStr.equals("AOp")) {
183       if (retTypeShape[0] == 16 && retTypeShape[1] == 16) {
184         NVVM::WMMALoadAM16N16K16Op wmmaLoadAOp =
185             rewriter.create<NVVM::WMMALoadAM16N16K16Op>(loc, resType,
186                                                         loadOpOperands);
187         rewriter.replaceOp(op, wmmaLoadAOp.getResult());
188       } else {
189         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
190       }
191     } else if (operandStr.equals("BOp")) {
192       if (retTypeShape[0] == 16 && retTypeShape[1] == 16) {
193         NVVM::WMMALoadBM16N16K16Op wmmaLoadBOp =
194             rewriter.create<NVVM::WMMALoadBM16N16K16Op>(loc, resType,
195                                                         loadOpOperands);
196         rewriter.replaceOp(op, wmmaLoadBOp.getResult());
197       } else {
198         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
199       }
200     } else {
201       if (retTypeShape[0] == 16 && retTypeShape[1] == 16) {
202         if (srcMemrefType.getElementType().isF16()) {
203           NVVM::WMMALoadCF16M16N16K16Op wmmaLoadCOp =
204               rewriter.create<NVVM::WMMALoadCF16M16N16K16Op>(loc, resType,
205                                                              loadOpOperands);
206           rewriter.replaceOp(op, wmmaLoadCOp.getResult());
207         } else if (srcMemrefType.getElementType().isF32()) {
208           NVVM::WMMALoadCF32M16N16K16Op wmmaLoadCOp =
209               rewriter.create<NVVM::WMMALoadCF32M16N16K16Op>(loc, resType,
210                                                              loadOpOperands);
211           rewriter.replaceOp(op, wmmaLoadCOp.getResult());
212         }
213       } else {
214         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
215       }
216     }
217     return success();
218   }
219 };
220 
221 /// This class implements the conversion of GPU MMA storeOp to wmma.store op
222 /// in the NVVM dialect. The conversion not only emits the NVVM op but also
223 /// emits code that is necessary to unpack the data in the source and
224 /// convert the data in the format that is needed by the NVVM op.
225 struct WmmaStoreOpToNVVMLowering
226     : public ConvertOpToLLVMPattern<gpu::SubgroupMmaStoreMatrixOp>,
227       private CommonLLVMAndBuiltInMLIRTypes {
228 public:
229   explicit WmmaStoreOpToNVVMLowering(LLVMTypeConverter &typeConverter)
230       : ConvertOpToLLVMPattern<gpu::SubgroupMmaStoreMatrixOp>(typeConverter),
231         CommonLLVMAndBuiltInMLIRTypes(&this->getTypeConverter()->getContext()) {
232   }
233 
234   LogicalResult
235   matchAndRewrite(gpu::SubgroupMmaStoreMatrixOp subgroupMmaStoreMatrixOp,
236                   ArrayRef<Value> operands,
237                   ConversionPatternRewriter &rewriter) const override {
238     Operation *op = subgroupMmaStoreMatrixOp.getOperation();
239     if (failed(areAllLLVMTypes(op, operands, rewriter)))
240       return failure();
241 
242     unsigned indexTypeBitwidth =
243         this->getTypeConverter()->getIndexTypeBitwidth();
244     // The corresponding intrinsics expects leadDimension to be a 32-bit
245     // integer, so all the calculations of linearizing the store address
246     // must also follow this restriction.
247     if (indexTypeBitwidth != 32)
248       return rewriter.notifyMatchFailure(
249           op, "expected indices to the memref to be 32-bit wide.");
250 
251     Location loc = op->getLoc();
252 
253     // Destination memref of the original op.
254     MemRefType dstMemrefType =
255         subgroupMmaStoreMatrixOp.dstMemref().getType().cast<MemRefType>();
256 
257     // MemRefDescriptor to extract alignedPtr and offset.
258     MemRefDescriptor promotedDstOp(
259         gpu::SubgroupMmaStoreMatrixOpAdaptor(operands).dstMemref());
260 
261     auto leadDimension = subgroupMmaStoreMatrixOp.leadDimensionAttr();
262 
263     // Emit ops which compute the store offset using `dstOffsetI`,
264     // `dstOffsetJ`. The actualOffset is (memrefOffset + (alignedPtr +
265     // ((leadDimension * dstOffsetI) + dstOffsetJ)).
266     SmallVector<Value> indices(subgroupMmaStoreMatrixOp.indices());
267     Value dstOffsetIVal = indices[0];
268     Value dstOffsetJVal = indices[1];
269     Value leadingDim32 =
270         rewriter.create<LLVM::ConstantOp>(loc, i32Ty, leadDimension);
271     Value numElemsLeadDim =
272         rewriter.create<LLVM::MulOp>(loc, i32Ty, leadingDim32, dstOffsetIVal);
273     Value loadOffset = rewriter.create<LLVM::AddOp>(loc, i32Ty, numElemsLeadDim,
274                                                     dstOffsetJVal);
275 
276     Value promotedDstOpToUse;
277     promotedDstOpToUse = promotedDstOp.offset(rewriter, loc);
278     Value actualOffset = rewriter.create<LLVM::AddOp>(loc, i32Ty, loadOffset,
279                                                       promotedDstOpToUse);
280     Value storeAddress = rewriter.create<LLVM::GEPOp>(
281         loc,
282         LLVM::LLVMPointerType::get(f16Ty, dstMemrefType.getMemorySpaceAsInt()),
283         promotedDstOp.alignedPtr(rewriter, loc), ArrayRef<Value>{actualOffset});
284 
285     // Bitcast the base address pointer of the destination memref, So that
286     // values can be stored in chunks of 32-bits and semantics match with the
287     // intrinsic exposed by NVPTX backend.
288     Value storeAddressCasted = rewriter.create<LLVM::BitcastOp>(
289         loc,
290         LLVM::LLVMPointerType::get(i32Ty, dstMemrefType.getMemorySpaceAsInt()),
291         storeAddress);
292 
293     SmallVector<Value, 4> storeOpOperands;
294     storeOpOperands.push_back(storeAddressCasted);
295 
296     // Get the shape of the MMAMatrix type being stored. The shape will
297     // choose which intrinsic this op will be lowered to.
298     gpu::MMAMatrixType srcType =
299         subgroupMmaStoreMatrixOp.src().getType().cast<gpu::MMAMatrixType>();
300     ArrayRef<int64_t> srcTypeShape = srcType.getShape();
301 
302     // Unpack the results from the source.
303     if (subgroupMmaStoreMatrixOp.src()
304             .getType()
305             .cast<gpu::MMAMatrixType>()
306             .getElementType() == f16Ty) {
307       for (unsigned i = 0, e = numHalfsInOpFrags[C]; i < e; ++i) {
308         Value toUse = rewriter.create<LLVM::ExtractValueOp>(
309             loc, f16x2Ty, operands[0], rewriter.getI32ArrayAttr(i));
310         storeOpOperands.push_back(toUse);
311       }
312       storeOpOperands.push_back(leadingDim32);
313 
314       // Create nvvm.mma_store op.
315       if (srcTypeShape[0] == 16 && srcTypeShape[1] == 16) {
316         rewriter.create<NVVM::WMMAStoreF16M16N16K16Op>(loc, storeOpOperands);
317       } else {
318         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
319       }
320       rewriter.eraseOp(op);
321       return success();
322     } else if (subgroupMmaStoreMatrixOp.src()
323                    .getType()
324                    .cast<gpu::MMAMatrixType>()
325                    .getElementType() == f32Ty) {
326       for (unsigned i = 0, e = 8; i < e; ++i) {
327         Value toUse = rewriter.create<LLVM::ExtractValueOp>(
328             loc, f32Ty, operands[0], rewriter.getI32ArrayAttr(i));
329         storeOpOperands.push_back(toUse);
330       }
331       storeOpOperands.push_back(leadingDim32);
332 
333       // Create nvvm.mma_store op.
334       if (srcTypeShape[0] == 16 && srcTypeShape[1] == 16)
335         rewriter.create<NVVM::WMMAStoreF32M16N16K16Op>(loc, storeOpOperands);
336       else {
337         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
338       }
339       rewriter.eraseOp(op);
340       return success();
341     }
342 
343     return failure();
344   }
345 };
346 
347 /// This class implements the conversion of GPU MMA computeOp to wmma.mma op
348 /// in the NVVM dialect.
349 struct WmmaMmaOpToNVVMLowering
350     : public ConvertOpToLLVMPattern<gpu::SubgroupMmaComputeOp>,
351       private CommonLLVMAndBuiltInMLIRTypes {
352   explicit WmmaMmaOpToNVVMLowering(LLVMTypeConverter &typeConverter)
353       : ConvertOpToLLVMPattern<gpu::SubgroupMmaComputeOp>(typeConverter),
354         CommonLLVMAndBuiltInMLIRTypes(&this->getTypeConverter()->getContext()) {
355   }
356 
357   LogicalResult
358   matchAndRewrite(gpu::SubgroupMmaComputeOp subgroupMmaComputeOp,
359                   ArrayRef<Value> operands,
360                   ConversionPatternRewriter &rewriter) const override {
361     Operation *op = subgroupMmaComputeOp.getOperation();
362     if (failed(areAllLLVMTypes(op, operands, rewriter)))
363       return failure();
364 
365     Location loc = op->getLoc();
366 
367     // The wmma.mma intrinsic in llvm requires the operands as individual
368     // values. So individual elements from the memrefs need to be extracted and
369     // then passed on to the intrinsic call. Emit llvm ops to extract individual
370     // values form lowered memrefs.
371     SmallVector<Value> unpackedOps;
372 
373     auto unpackOp = [&](CommonLLVMAndBuiltInMLIRTypes::OperandMap op,
374                         Value operand, unsigned numElems, Type elemType) {
375       for (unsigned i = 0; i < numElems; ++i) {
376         Value toUse = rewriter.create<LLVM::ExtractValueOp>(
377             loc, elemType, operand, rewriter.getI32ArrayAttr(i));
378         unpackedOps.push_back(toUse);
379       }
380     };
381 
382     // Get the shapes of the MMAMatrix type being used. The shapes will
383     // choose which intrinsic this op will be lowered to.
384     gpu::MMAMatrixType aType =
385         subgroupMmaComputeOp.opA().getType().cast<gpu::MMAMatrixType>();
386     ArrayRef<int64_t> aTypeShape = aType.getShape();
387     gpu::MMAMatrixType bType =
388         subgroupMmaComputeOp.opA().getType().cast<gpu::MMAMatrixType>();
389     ArrayRef<int64_t> bTypeShape = bType.getShape();
390     gpu::MMAMatrixType cType =
391         subgroupMmaComputeOp.opA().getType().cast<gpu::MMAMatrixType>();
392     ArrayRef<int64_t> cTypeShape = cType.getShape();
393 
394     gpu::SubgroupMmaComputeOpAdaptor transformedOperands(operands);
395     if (subgroupMmaComputeOp.opC()
396             .getType()
397             .cast<gpu::MMAMatrixType>()
398             .getElementType() == f16Ty) {
399       unpackOp(A, transformedOperands.opA(), numHalfsInOpFrags[A], f16x2Ty);
400       unpackOp(B, transformedOperands.opB(), numHalfsInOpFrags[B], f16x2Ty);
401       unpackOp(C, transformedOperands.opC(), numHalfsInOpFrags[C], f16x2Ty);
402 
403       if (aTypeShape[0] == 16 && aTypeShape[1] == 16 && bTypeShape[0] == 16 &&
404           bTypeShape[1] == 16 && cTypeShape[0] == 16 && cTypeShape[1] == 16) {
405         // Create nvvm.wmma.mma op.
406         NVVM::WMMAMmaF16F16M16N16K16Op wmmaMmaOp =
407             rewriter.create<NVVM::WMMAMmaF16F16M16N16K16Op>(loc, fragArrayCDTy,
408                                                             unpackedOps);
409 
410         rewriter.replaceOp(op, wmmaMmaOp.getResult());
411         return success();
412       } else {
413         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
414       }
415     } else if (subgroupMmaComputeOp.opC()
416                    .getType()
417                    .cast<gpu::MMAMatrixType>()
418                    .getElementType() == f32Ty) {
419       unpackOp(A, transformedOperands.opA(), numHalfsInOpFrags[A], f16x2Ty);
420       unpackOp(B, transformedOperands.opB(), numHalfsInOpFrags[B], f16x2Ty);
421       unpackOp(C, transformedOperands.opC(), 8, f32Ty);
422 
423       if (aTypeShape[0] == 16 && aTypeShape[1] == 16 && bTypeShape[0] == 16 &&
424           bTypeShape[1] == 16 && cTypeShape[0] == 16 && cTypeShape[1] == 16) {
425         // Create nvvm.wmma.mma op.
426         NVVM::WMMAMmaF32F32M16N16K16Op wmmaMmaOp =
427             rewriter.create<NVVM::WMMAMmaF32F32M16N16K16Op>(
428                 loc, fragArrayCDF32Ty, unpackedOps);
429 
430         rewriter.replaceOp(op, wmmaMmaOp.getResult());
431         return success();
432       } else {
433         return rewriter.notifyMatchFailure(op, kInvalidCaseStr);
434       }
435     }
436 
437     return failure();
438   }
439 };
440 
441 } // anonymous namespace
442 
443 namespace mlir {
444 void populateGpuWMMAToNVVMConversionPatterns(LLVMTypeConverter &converter,
445                                              RewritePatternSet &patterns) {
446   patterns.insert<WmmaLoadOpToNVVMLowering>(converter);
447   patterns.insert<WmmaMmaOpToNVVMLowering>(converter);
448   patterns.insert<WmmaStoreOpToNVVMLowering>(converter);
449 }
450 } // namespace mlir
451