1 //===- NVGPUToNVVM.cpp - NVGPU to NVVM dialect conversion -----------------===//
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 #include "mlir/Conversion/NVGPUToNVVM/NVGPUToNVVM.h"
10 #include "../PassDetail.h"
11 #include "mlir/Conversion/LLVMCommon/ConversionTarget.h"
12 #include "mlir/Conversion/LLVMCommon/Pattern.h"
13 #include "mlir/Dialect/LLVMIR/NVVMDialect.h"
14 #include "mlir/Dialect/NVGPU/NVGPUDialect.h"
15 
16 using namespace mlir;
17 
18 /// Returns the type for the intrinsic given the vectorResultType of the
19 /// `gpu.mma.sync` operation.
20 static Type inferIntrinsicResultType(Type vectorResultType) {
21   MLIRContext *ctx = vectorResultType.getContext();
22   auto a = vectorResultType.cast<LLVM::LLVMArrayType>();
23   auto f16x2Ty = LLVM::getFixedVectorType(Float16Type::get(ctx), 2);
24   auto i32Ty = IntegerType::get(ctx, 32);
25   auto i32x2Ty = LLVM::getFixedVectorType(i32Ty, 2);
26   Type f64Ty = Float64Type::get(ctx);
27   Type f64x2Ty = LLVM::getFixedVectorType(f64Ty, 2);
28   Type f32Ty = Float32Type::get(ctx);
29   Type f32x2Ty = LLVM::getFixedVectorType(f32Ty, 2);
30   if (a.getElementType() == f16x2Ty) {
31     return LLVM::LLVMStructType::getLiteral(
32         ctx, SmallVector<Type>(a.getNumElements(), f16x2Ty));
33   }
34   if (a.getElementType() == i32x2Ty) {
35     return LLVM::LLVMStructType::getLiteral(
36         ctx,
37         SmallVector<Type>(static_cast<size_t>(a.getNumElements()) * 2, i32Ty));
38   }
39   if (a.getElementType() == f64x2Ty) {
40     return LLVM::LLVMStructType::getLiteral(ctx, {f64Ty, f64Ty});
41   }
42   if (a.getElementType() == f32x2Ty) {
43     return LLVM::LLVMStructType::getLiteral(
44         ctx,
45         SmallVector<Type>(static_cast<size_t>(a.getNumElements()) * 2, f32Ty));
46   }
47   if (a.getElementType() == LLVM::getFixedVectorType(f32Ty, 1)) {
48     return LLVM::LLVMStructType::getLiteral(
49         ctx, SmallVector<Type>(static_cast<size_t>(a.getNumElements()), f32Ty));
50   }
51   return vectorResultType;
52 }
53 
54 /// Convert the SSA result of the NVVM intrinsic `nvvm.mma.sync` (which is
55 /// always an LLVM struct) into a fragment that is compatible with the vector
56 /// type of this operation. This involves extracting elements from the struct
57 /// and inserting them into an LLVM array. These extra data-movement
58 /// operations should be canonicalized away by the LLVM backend.
59 static Value convertIntrinsicResult(Location loc, Type intrinsicResultType,
60                                     Type resultType, Value intrinsicResult,
61                                     RewriterBase &rewriter) {
62   MLIRContext *ctx = rewriter.getContext();
63   auto structType = intrinsicResultType.dyn_cast<LLVM::LLVMStructType>();
64   auto arrayType = resultType.dyn_cast<LLVM::LLVMArrayType>();
65   Type i32Ty = rewriter.getI32Type();
66   Type f32Ty = rewriter.getF32Type();
67   Type f64Ty = rewriter.getF64Type();
68   Type f16x2Ty = LLVM::getFixedVectorType(rewriter.getF16Type(), 2);
69   Type i32x2Ty = LLVM::getFixedVectorType(i32Ty, 2);
70   Type f64x2Ty = LLVM::getFixedVectorType(f64Ty, 2);
71   Type f32x2Ty = LLVM::getFixedVectorType(f32Ty, 2);
72   Type f32x1Ty = LLVM::getFixedVectorType(f32Ty, 1);
73 
74   auto makeConst = [&](int32_t index) -> Value {
75     return rewriter.create<LLVM::ConstantOp>(loc, IntegerType::get(ctx, 32),
76                                              rewriter.getI32IntegerAttr(index));
77   };
78 
79   if (arrayType) {
80     SmallVector<Value, 4> elements;
81 
82     // The intrinsic returns 32-bit wide elements in a form which can be
83     // directly bitcasted and inserted into the result vector.
84     if (arrayType.getElementType() == f16x2Ty ||
85         arrayType.getElementType() == f32x1Ty) {
86       for (unsigned i = 0; i < structType.getBody().size(); i++) {
87         Value el = rewriter.create<LLVM::ExtractValueOp>(
88             loc, structType.getBody()[i], intrinsicResult,
89             rewriter.getI64ArrayAttr(i));
90         el = rewriter.createOrFold<LLVM::BitcastOp>(
91             loc, arrayType.getElementType(), el);
92         elements.push_back(el);
93       }
94     }
95 
96     // The intrinsic returns i32, f64, and f32 values as individual scalars,
97     // even when the result is notionally a 64-bit wide element (e.g. f32x2). We
98     // need to extract them from the struct and pack them into the 64-bit wide
99     // rows of the vector result.
100     if (arrayType.getElementType() == i32x2Ty ||
101         arrayType.getElementType() == f64x2Ty ||
102         arrayType.getElementType() == f32x2Ty) {
103 
104       for (unsigned i = 0, e = structType.getBody().size() / 2; i < e; i++) {
105         Value vec =
106             rewriter.create<LLVM::UndefOp>(loc, arrayType.getElementType());
107         Value x1 = rewriter.create<LLVM::ExtractValueOp>(
108             loc, structType.getBody()[i * 2], intrinsicResult,
109             rewriter.getI64ArrayAttr(i * 2));
110         Value x2 = rewriter.create<LLVM::ExtractValueOp>(
111             loc, structType.getBody()[i * 2 + 1], intrinsicResult,
112             rewriter.getI64ArrayAttr(i * 2 + 1));
113         vec = rewriter.create<LLVM::InsertElementOp>(loc, vec.getType(), vec,
114                                                      x1, makeConst(0));
115         vec = rewriter.create<LLVM::InsertElementOp>(loc, vec.getType(), vec,
116                                                      x2, makeConst(1));
117         elements.push_back(vec);
118       }
119     }
120 
121     // Create the final vectorized result.
122     Value result = rewriter.create<LLVM::UndefOp>(loc, arrayType);
123     for (const auto &el : llvm::enumerate(elements)) {
124       result = rewriter.create<LLVM::InsertValueOp>(
125           loc, arrayType, result, el.value(),
126           rewriter.getI64ArrayAttr(el.index()));
127     }
128     return result;
129   }
130 
131   return intrinsicResult;
132 }
133 
134 /// The `gpu.mma.sync` converter below expects matrix fragment operands to be
135 /// given as 2D `vectors` where the rows are 32b or 64b wide. The
136 /// `nvvm.mma.sync` op expects these argments to be a given in a long list of
137 /// scalars of certain types. This function helps unpack the `vector` arguments
138 /// and cast them to the types expected by `nvvm.mma.sync`.
139 static SmallVector<Value> unpackOperandVector(RewriterBase &rewriter,
140                                               Location loc, Value operand,
141                                               NVVM::MMATypes operandPtxType) {
142   SmallVector<Value> result;
143   Type i32Ty = rewriter.getI32Type();
144   Type f64Ty = rewriter.getF64Type();
145   Type f32Ty = rewriter.getF32Type();
146   Type i8Ty = rewriter.getI8Type();
147   Type i8x4Ty = LLVM::getFixedVectorType(i8Ty, 4);
148   Type f32x1Ty = LLVM::getFixedVectorType(f32Ty, 1);
149   auto arrayTy = operand.getType().cast<LLVM::LLVMArrayType>();
150 
151   for (unsigned i = 0, e = arrayTy.getNumElements(); i < e; ++i) {
152     Value toUse = rewriter.create<LLVM::ExtractValueOp>(
153         loc, arrayTy.getElementType(), operand, rewriter.getI64ArrayAttr(i));
154 
155     // For 4xi8 vectors, the intrinsic expects these to be provided as i32
156     // scalar types.
157     if (arrayTy.getElementType() == i8x4Ty ||
158         (arrayTy.getElementType() == f32x1Ty &&
159          operandPtxType == NVVM::MMATypes::tf32)) {
160       result.push_back(
161           rewriter.create<LLVM::BitcastOp>(loc, rewriter.getI32Type(), toUse));
162       continue;
163     }
164 
165     // For some element types (i32, f32, f64), we need to unpack the inner
166     // vector/array type as well because the intrinsic expects individual
167     // scalars to be provided.
168     VectorType innerArrayTy = arrayTy.getElementType().dyn_cast<VectorType>();
169     if (innerArrayTy && (innerArrayTy.getElementType() == i32Ty ||
170                          innerArrayTy.getElementType() == f64Ty ||
171                          innerArrayTy.getElementType() == f32Ty)) {
172       for (unsigned idx = 0, innerSize = innerArrayTy.getNumElements();
173            idx < innerSize; idx++) {
174         result.push_back(rewriter.create<LLVM::ExtractElementOp>(
175             loc, toUse,
176             rewriter.create<LLVM::ConstantOp>(
177                 loc, rewriter.getI64Type(), rewriter.getI64IntegerAttr(idx))));
178       }
179       continue;
180     }
181     result.push_back(toUse);
182   }
183   return result;
184 }
185 
186 namespace {
187 
188 struct MmaLdMatrixOpToNVVM : public ConvertOpToLLVMPattern<nvgpu::LdMatrixOp> {
189   using ConvertOpToLLVMPattern<nvgpu::LdMatrixOp>::ConvertOpToLLVMPattern;
190 
191   LogicalResult
192   matchAndRewrite(nvgpu::LdMatrixOp op, OpAdaptor adaptor,
193                   ConversionPatternRewriter &rewriter) const override {
194     MLIRContext *ctx = getContext();
195     Location loc = op->getLoc();
196 
197     // The result type of ldmatrix will always be a struct of 32bit integer
198     // registers if more than one 32bit value is returned. Otherwise, the result
199     // is a single i32. The result type of the GPU operation is always a vector
200     // of shape (NumRegisters, VectorRegister) where VectorRegister is the
201     // vector type of the result and always 32 bits long. We bitcast the result
202     // of the NVVM::LdMatrix to this vector type.
203     auto vectorResultType = op->getResultTypes()[0].dyn_cast<VectorType>();
204     if (!vectorResultType) {
205       return failure();
206     }
207     Type innerVectorType = LLVM::getFixedVectorType(
208         vectorResultType.getElementType(), vectorResultType.getDimSize(1));
209 
210     int64_t num32BitRegs = vectorResultType.getDimSize(0);
211 
212     Type ldMatrixResultType;
213     if (num32BitRegs > 1) {
214       ldMatrixResultType = LLVM::LLVMStructType::getLiteral(
215           ctx, SmallVector<Type>(num32BitRegs, rewriter.getI32Type()));
216     } else {
217       ldMatrixResultType = rewriter.getI32Type();
218     }
219 
220     auto srcMemrefType = op.srcMemref().getType().cast<MemRefType>();
221     Value srcPtr = getStridedElementPtr(loc, srcMemrefType, adaptor.srcMemref(),
222                                         adaptor.indices(), rewriter);
223     Value ldMatrixResult = rewriter.create<NVVM::LdMatrixOp>(
224         loc, ldMatrixResultType, srcPtr,
225         /*num=*/op.numTiles(),
226         /*layout=*/op.transpose() ? NVVM::MMALayout::col
227                                   : NVVM::MMALayout::row);
228 
229     // The ldmatrix operation returns either a single i32 value or a struct of
230     // i32 values. Here we unpack those values and cast them back to their
231     // actual vector type (still of width 32b) and repack them into a result
232     // struct.
233     Type finalResultType = typeConverter->convertType(vectorResultType);
234     Value result = rewriter.create<LLVM::UndefOp>(loc, finalResultType);
235     for (int64_t i = 0, e = vectorResultType.getDimSize(0); i < e; i++) {
236       Value i32Register = num32BitRegs > 1
237                               ? rewriter.create<LLVM::ExtractValueOp>(
238                                     loc, rewriter.getI32Type(), ldMatrixResult,
239                                     rewriter.getI64ArrayAttr(i))
240                               : ldMatrixResult;
241       Value casted =
242           rewriter.create<LLVM::BitcastOp>(loc, innerVectorType, i32Register);
243       result = rewriter.create<LLVM::InsertValueOp>(
244           loc, finalResultType, result, casted, rewriter.getI64ArrayAttr(i));
245     }
246 
247     rewriter.replaceOp(op, result);
248     return success();
249   }
250 };
251 
252 struct MmaSyncOptoNVVM : public ConvertOpToLLVMPattern<nvgpu::MmaSyncOp> {
253   using ConvertOpToLLVMPattern<nvgpu::MmaSyncOp>::ConvertOpToLLVMPattern;
254 
255   LogicalResult
256   matchAndRewrite(nvgpu::MmaSyncOp op, OpAdaptor adaptor,
257                   ConversionPatternRewriter &rewriter) const override {
258     Location loc = op->getLoc();
259     // Get the shapes of the MMAMatrix type being used. The shapes will
260     // choose which intrinsic this op will be lowered to.
261     auto aType = op.matrixA().getType().cast<VectorType>();
262     auto cType = op.matrixC().getType().cast<VectorType>();
263 
264     int64_t m = op.mmaShape()[0].cast<IntegerAttr>().getInt();
265     int64_t n = op.mmaShape()[1].cast<IntegerAttr>().getInt();
266     int64_t k = op.mmaShape()[2].cast<IntegerAttr>().getInt();
267     std::array<int64_t, 3> gemmShape{m, n, k};
268 
269     NVVM::MMATypes ptxTypeA;
270     NVVM::MMATypes ptxTypeB;
271     Optional<NVVM::MMATypes> ptxTypeC = NVVM::MmaOp::inferOperandMMAType(
272         cType.getElementType(), /*isAccumulator=*/true);
273     if (!ptxTypeC) {
274       return op->emitError(
275           "could not infer the PTX type for the accumulator/result");
276     }
277 
278     Optional<NVVM::MMAIntOverflow> overflow(llvm::None);
279     if (aType.getElementType().isInteger(8)) {
280       ptxTypeA = NVVM::MMATypes::s8;
281       ptxTypeB = NVVM::MMATypes::s8;
282       overflow = NVVM::MMAIntOverflow::satfinite;
283     } else if (aType.getElementType().isF16()) {
284       ptxTypeA = NVVM::MMATypes::f16;
285       ptxTypeB = NVVM::MMATypes::f16;
286     } else if (aType.getElementType().isF64()) {
287       ptxTypeA = NVVM::MMATypes::f64;
288       ptxTypeB = NVVM::MMATypes::f64;
289     } else if (aType.getElementType().isF32()) {
290       ptxTypeA = NVVM::MMATypes::tf32;
291       ptxTypeB = NVVM::MMATypes::tf32;
292     } else {
293       return op->emitError("could not deduce operand PTX types");
294     }
295 
296     SmallVector<Value> matA =
297         unpackOperandVector(rewriter, loc, adaptor.matrixA(), ptxTypeA);
298     SmallVector<Value> matB =
299         unpackOperandVector(rewriter, loc, adaptor.matrixB(), ptxTypeB);
300     SmallVector<Value> matC =
301         unpackOperandVector(rewriter, loc, adaptor.matrixC(), *ptxTypeC);
302 
303     Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
304     Type intrinsicResTy = inferIntrinsicResultType(
305         typeConverter->convertType(op->getResultTypes()[0]));
306     Value intrinsicResult = rewriter.create<NVVM::MmaOp>(
307         op.getLoc(), intrinsicResTy, matA, matB, matC,
308         /*shape=*/gemmShape,
309         /*b1Op=*/llvm::None,
310         /*intOverflow=*/overflow,
311         /*multiplicandPtxTypes=*/
312         std::array<NVVM::MMATypes, 2>{ptxTypeA, ptxTypeB},
313         /*multiplicandLayouts=*/
314         std::array<NVVM::MMALayout, 2>{NVVM::MMALayout::row,
315                                        NVVM::MMALayout::col});
316     rewriter.replaceOp(op, convertIntrinsicResult(op.getLoc(), intrinsicResTy,
317                                                   desiredRetTy, intrinsicResult,
318                                                   rewriter));
319     return success();
320   }
321 };
322 
323 struct ConvertNVGPUToNVVMPass
324     : public ConvertNVGPUToNVVMBase<ConvertNVGPUToNVVMPass> {
325   ConvertNVGPUToNVVMPass() = default;
326 
327   void runOnOperation() override {
328     RewritePatternSet patterns(&getContext());
329     LLVMTypeConverter converter(&getContext());
330     populateNVGPUToNVVMConversionPatterns(converter, patterns);
331     LLVMConversionTarget target(getContext());
332     target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
333     target.addLegalDialect<::mlir::NVVM::NVVMDialect>();
334     if (failed(applyPartialConversion(getOperation(), target,
335                                       std::move(patterns))))
336       signalPassFailure();
337   }
338 };
339 
340 } // namespace
341 void mlir::populateNVGPUToNVVMConversionPatterns(LLVMTypeConverter &converter,
342                                                  RewritePatternSet &patterns) {
343   patterns.add<MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM>(converter);
344 }
345 
346 std::unique_ptr<Pass> mlir::createConvertNVGPUToNVVMPass() {
347   return std::make_unique<ConvertNVGPUToNVVMPass>();
348 }
349