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