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