MLIR 24.0.0git
NVGPUToNVVM.cpp
Go to the documentation of this file.
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
10
27#include "mlir/IR/Value.h"
28#include "mlir/Pass/Pass.h"
29#include "llvm/Support/Debug.h"
30#include "llvm/Support/DebugLog.h"
31#include "llvm/Support/ErrorHandling.h"
32#include "llvm/Support/raw_ostream.h"
33#include <optional>
34
35#define DEBUG_TYPE "nvgpu-to-nvvm"
36
37namespace mlir {
38#define GEN_PASS_DEF_CONVERTNVGPUTONVVMPASS
39#include "mlir/Conversion/Passes.h.inc"
40} // namespace mlir
41
42using namespace mlir;
43
44/// Number of bits that needs to be excluded when building matrix descriptor for
45/// wgmma operations.
46constexpr int exclude4LSB = 4;
47
48/// GPU has 32 bit registers, this function truncates values when larger width
49/// is not needed.
51 Type type = value.getType();
52 assert(llvm::isa<IntegerType>(type) && "expected an integer Value");
53 if (type.getIntOrFloatBitWidth() <= 32)
54 return value;
55 return LLVM::TruncOp::create(b, b.getI32Type(), value);
56}
57
58/// Returns the type for the intrinsic given the vectorResultType of the
59/// `gpu.mma.sync` operation.
60static Type inferIntrinsicResultType(Type vectorResultType) {
61 MLIRContext *ctx = vectorResultType.getContext();
62 auto a = cast<LLVM::LLVMArrayType>(vectorResultType);
63 auto f16x2Ty = VectorType::get(2, Float16Type::get(ctx));
64 auto i32Ty = IntegerType::get(ctx, 32);
65 auto i32x2Ty = VectorType::get(2, i32Ty);
66 Type f64Ty = Float64Type::get(ctx);
67 Type f64x2Ty = VectorType::get(2, f64Ty);
68 Type f32Ty = Float32Type::get(ctx);
69 Type f32x2Ty = VectorType::get(2, f32Ty);
70 if (a.getElementType() == f16x2Ty) {
71 return LLVM::LLVMStructType::getLiteral(
72 ctx, SmallVector<Type>(a.getNumElements(), f16x2Ty));
73 }
74 if (a.getElementType() == i32x2Ty) {
75 return LLVM::LLVMStructType::getLiteral(
76 ctx,
77 SmallVector<Type>(static_cast<size_t>(a.getNumElements()) * 2, i32Ty));
78 }
79 if (a.getElementType() == f64x2Ty) {
80 return LLVM::LLVMStructType::getLiteral(ctx, {f64Ty, f64Ty});
81 }
82 if (a.getElementType() == f32x2Ty) {
83 return LLVM::LLVMStructType::getLiteral(
84 ctx,
85 SmallVector<Type>(static_cast<size_t>(a.getNumElements()) * 2, f32Ty));
86 }
87 if (a.getElementType() == VectorType::get(1, f32Ty)) {
88 return LLVM::LLVMStructType::getLiteral(
89 ctx, SmallVector<Type>(static_cast<size_t>(a.getNumElements()), f32Ty));
90 }
91 return vectorResultType;
92}
93
94/// Convert the SSA result of the NVVM intrinsic `nvvm.mma.sync` (which is
95/// always an LLVM struct) into a fragment that is compatible with the vector
96/// type of this operation. This involves extracting elements from the struct
97/// and inserting them into an LLVM array. These extra data-movement
98/// operations should be canonicalized away by the LLVM backend.
99static Value convertIntrinsicResult(Location loc, Type intrinsicResultType,
100 Type resultType, Value intrinsicResult,
101 RewriterBase &rewriter) {
102 MLIRContext *ctx = rewriter.getContext();
103 auto structType = dyn_cast<LLVM::LLVMStructType>(intrinsicResultType);
104 auto arrayType = dyn_cast<LLVM::LLVMArrayType>(resultType);
105 Type i32Ty = rewriter.getI32Type();
106 Type f32Ty = rewriter.getF32Type();
107 Type f64Ty = rewriter.getF64Type();
108 Type f16x2Ty = VectorType::get(2, rewriter.getF16Type());
109 Type i32x2Ty = VectorType::get(2, i32Ty);
110 Type f64x2Ty = VectorType::get(2, f64Ty);
111 Type f32x2Ty = VectorType::get(2, f32Ty);
112 Type f32x1Ty = VectorType::get(1, f32Ty);
113
114 auto makeConst = [&](int32_t index) -> Value {
115 return LLVM::ConstantOp::create(rewriter, loc, IntegerType::get(ctx, 32),
116 rewriter.getI32IntegerAttr(index));
117 };
118
119 if (arrayType) {
120 SmallVector<Value, 4> elements;
121
122 // The intrinsic returns 32-bit wide elements in a form which can be
123 // directly bitcasted and inserted into the result vector.
124 if (arrayType.getElementType() == f16x2Ty ||
125 arrayType.getElementType() == f32x1Ty) {
126 for (unsigned i = 0; i < structType.getBody().size(); i++) {
127 Value el =
128 LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult, i);
129 el = rewriter.createOrFold<LLVM::BitcastOp>(
130 loc, arrayType.getElementType(), el);
131 elements.push_back(el);
132 }
133 }
134
135 // The intrinsic returns i32, f64, and f32 values as individual scalars,
136 // even when the result is notionally a 64-bit wide element (e.g. f32x2). We
137 // need to extract them from the struct and pack them into the 64-bit wide
138 // rows of the vector result.
139 if (arrayType.getElementType() == i32x2Ty ||
140 arrayType.getElementType() == f64x2Ty ||
141 arrayType.getElementType() == f32x2Ty) {
142
143 for (unsigned i = 0, e = structType.getBody().size() / 2; i < e; i++) {
144 Value vec =
145 LLVM::PoisonOp::create(rewriter, loc, arrayType.getElementType());
146 Value x1 =
147 LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult, i * 2);
148 Value x2 = LLVM::ExtractValueOp::create(rewriter, loc, intrinsicResult,
149 i * 2 + 1);
150 vec = LLVM::InsertElementOp::create(rewriter, loc, vec.getType(), vec,
151 x1, makeConst(0));
152 vec = LLVM::InsertElementOp::create(rewriter, loc, vec.getType(), vec,
153 x2, makeConst(1));
154 elements.push_back(vec);
155 }
156 }
157
158 // Create the final vectorized result.
159 Value result = LLVM::PoisonOp::create(rewriter, loc, arrayType);
160 for (const auto &el : llvm::enumerate(elements)) {
161 result = LLVM::InsertValueOp::create(rewriter, loc, result, el.value(),
162 el.index());
163 }
164 return result;
165 }
166
167 return intrinsicResult;
168}
169
170/// The `gpu.mma.sync` converter below expects matrix fragment operands to be
171/// given as 2D `vectors` where the rows are 32b or 64b wide. The
172/// `nvvm.mma.sync` op expects these argments to be a given in a long list of
173/// scalars of certain types. This function helps unpack the `vector` arguments
174/// and cast them to the types expected by `nvvm.mma.sync`.
176 Value operand,
177 NVVM::MMATypes operandPtxType) {
179 Type i32Ty = b.getI32Type();
180 Type f64Ty = b.getF64Type();
181 Type f32Ty = b.getF32Type();
182 Type i64Ty = b.getI64Type();
183 Type bf16x2Ty = VectorType::get(2, b.getBF16Type());
184 Type i8x4Ty = VectorType::get(4, b.getI8Type());
185 Type i4x8Ty = VectorType::get(8, b.getIntegerType(4));
186 Type f32x1Ty = VectorType::get(1, f32Ty);
187 auto arrayTy = cast<LLVM::LLVMArrayType>(operand.getType());
188
189 for (unsigned i = 0, e = arrayTy.getNumElements(); i < e; ++i) {
190 Value toUse = LLVM::ExtractValueOp::create(b, operand, i);
191
192 // For 4xi8 vectors, the intrinsic expects these to be provided as i32
193 // scalar types.
194 if (arrayTy.getElementType() == i8x4Ty ||
195 arrayTy.getElementType() == i4x8Ty ||
196 (arrayTy.getElementType() == bf16x2Ty &&
197 operandPtxType == NVVM::MMATypes::bf16) ||
198 (arrayTy.getElementType() == f32x1Ty &&
199 operandPtxType == NVVM::MMATypes::tf32)) {
200 result.push_back(LLVM::BitcastOp::create(b, i32Ty, toUse));
201 continue;
202 }
203
204 // For some element types (i32, f32, f64), we need to unpack the inner
205 // vector/array type as well because the intrinsic expects individual
206 // scalars to be provided.
207 VectorType innerArrayTy = dyn_cast<VectorType>(arrayTy.getElementType());
208 if (innerArrayTy && (innerArrayTy.getElementType() == i32Ty ||
209 innerArrayTy.getElementType() == f64Ty ||
210 innerArrayTy.getElementType() == f32Ty)) {
211 for (unsigned idx = 0, innerSize = innerArrayTy.getNumElements();
212 idx < innerSize; idx++) {
213 result.push_back(LLVM::ExtractElementOp::create(
214 b, toUse,
215 LLVM::ConstantOp::create(b, i64Ty, b.getI64IntegerAttr(idx))));
216 }
217 continue;
218 }
219 result.push_back(toUse);
220 }
221 return result;
222}
223
224/// Returns whether mbarrier object has shared memory address space.
225static bool isMbarrierShared(nvgpu::MBarrierGroupType barrierType) {
226 return (mlir::nvgpu::NVGPUDialect::isSharedMemoryAddressSpace(
227 barrierType.getMemorySpace()));
228}
229
230/// Returns the memory space attribute of the mbarrier object.
232 nvgpu::MBarrierGroupType barrierType) {
233 Attribute memorySpace = {};
234 if (isMbarrierShared(barrierType)) {
235 memorySpace =
236 IntegerAttr::get(IntegerType::get(context, 64),
237 nvgpu::NVGPUDialect::kSharedMemoryAddressSpace);
238 }
239 return memorySpace;
240}
241
242/// Returns memref type of the mbarrier object. The type is defined in the
243/// MBarrierGroupType.
244MemRefType nvgpu::getMBarrierMemrefType(MLIRContext *context,
245 nvgpu::MBarrierGroupType barrierType) {
246 Attribute memorySpace = nvgpu::getMbarrierMemorySpace(context, barrierType);
247 MemRefLayoutAttrInterface layout;
248 return MemRefType::get({barrierType.getNumBarriers()},
249 IntegerType::get(context, 64), layout, memorySpace);
250}
251
252namespace {
253
254struct MmaLdMatrixOpToNVVM : public ConvertOpToLLVMPattern<nvgpu::LdMatrixOp> {
255 using ConvertOpToLLVMPattern<nvgpu::LdMatrixOp>::ConvertOpToLLVMPattern;
256
257 LogicalResult
258 matchAndRewrite(nvgpu::LdMatrixOp op, OpAdaptor adaptor,
259 ConversionPatternRewriter &rewriter) const override {
260 MLIRContext *ctx = getContext();
261 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
262
263 // The result type of ldmatrix will always be a struct of 32bit integer
264 // registers if more than one 32bit value is returned. Otherwise, the result
265 // is a single i32. The result type of the GPU operation is always a vector
266 // of shape (NumRegisters, VectorRegister) where VectorRegister is the
267 // vector type of the result and always 32 bits long. We bitcast the result
268 // of the NVVM::LdMatrix to this vector type.
269 auto vectorResultType = dyn_cast<VectorType>(op->getResultTypes()[0]);
270 if (!vectorResultType) {
271 return failure();
272 }
273 Type innerVectorType = VectorType::get(vectorResultType.getDimSize(1),
274 vectorResultType.getElementType());
275
276 int64_t num32BitRegs = vectorResultType.getDimSize(0);
277
278 Type ldMatrixResultType;
279 if (num32BitRegs > 1) {
280 ldMatrixResultType = LLVM::LLVMStructType::getLiteral(
281 ctx, SmallVector<Type>(num32BitRegs, rewriter.getI32Type()));
282 } else {
283 ldMatrixResultType = rewriter.getI32Type();
284 }
285
286 auto srcMemrefType = cast<MemRefType>(op.getSrcMemref().getType());
287 Value srcPtr =
288 getStridedElementPtr(rewriter, b.getLoc(), srcMemrefType,
289 adaptor.getSrcMemref(), adaptor.getIndices());
290 auto shape = NVVM::LdStMatrixShapeAttr::get(rewriter.getContext(), 8, 8);
291 Value ldMatrixResult = NVVM::LdMatrixOp::create(
292 b, ldMatrixResultType, srcPtr,
293 /*num=*/op.getNumTiles(),
294 /*layout=*/op.getTranspose() ? NVVM::MMALayout::col
295 : NVVM::MMALayout::row,
296 /*shape=*/shape, /*eltType=*/NVVM::LdStMatrixEltType::B16);
297
298 // The ldmatrix operation returns either a single i32 value or a struct of
299 // i32 values. Here we unpack those values and cast them back to their
300 // actual vector type (still of width 32b) and repack them into a result
301 // struct.
302 Type finalResultType = typeConverter->convertType(vectorResultType);
303 Value result = LLVM::PoisonOp::create(b, finalResultType);
304 for (int64_t i = 0, e = vectorResultType.getDimSize(0); i < e; i++) {
305 Value i32Register =
306 num32BitRegs > 1 ? LLVM::ExtractValueOp::create(b, ldMatrixResult, i)
307 : ldMatrixResult;
308 Value casted = LLVM::BitcastOp::create(b, innerVectorType, i32Register);
309 result = LLVM::InsertValueOp::create(b, result, casted, i);
310 }
311
312 rewriter.replaceOp(op, result);
313 return success();
314 }
315};
316
317/// Convert the given type into the corresponding PTX type (NVVM::MMATypes
318/// enum).
319static FailureOr<NVVM::MMATypes> getNvvmMmaType(Type t) {
320 Type elType = getElementTypeOrSelf(t);
321 if (elType.isInteger(8))
322 return NVVM::MMATypes::s8;
323 if (elType.isInteger(4))
324 return NVVM::MMATypes::s4;
325 if (elType.isF16())
326 return NVVM::MMATypes::f16;
327 if (elType.isBF16())
328 return NVVM::MMATypes::bf16;
329 if (elType.isF64())
330 return NVVM::MMATypes::f64;
331 if (elType.isF32())
332 return NVVM::MMATypes::tf32;
333 if (elType.isF8E4M3FN())
334 return NVVM::MMATypes::e4m3;
335 if (elType.isF8E5M2())
336 return NVVM::MMATypes::e5m2;
337 return failure();
338}
339
340struct MmaSyncOptoNVVM : public ConvertOpToLLVMPattern<nvgpu::MmaSyncOp> {
341 using ConvertOpToLLVMPattern<nvgpu::MmaSyncOp>::ConvertOpToLLVMPattern;
342
343 LogicalResult
344 matchAndRewrite(nvgpu::MmaSyncOp op, OpAdaptor adaptor,
345 ConversionPatternRewriter &rewriter) const override {
346 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
347 // Get the shapes of the MMAMatrix type being used. The shapes will
348 // choose which intrinsic this op will be lowered to.
349 VectorType aType = op.getMatrixA().getType();
350 VectorType bType = op.getMatrixA().getType();
351 VectorType cType = op.getMatrixC().getType();
352
353 std::array<int64_t, 3> gemmShape = op.getMmaShapeAsArray();
354
355 // Tensor Cores (mma.sync) on F32 works only with TensorFloat32 (TF32).
356 bool tf32Enabled = op->hasAttr(op.getTf32EnabledAttrName());
357 if (aType.getElementType().isF32() && !tf32Enabled)
358 return failure();
359
360 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
361 if (failed(ptxTypeA))
362 return op->emitOpError("failed to deduce operand PTX types");
363 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
364 if (failed(ptxTypeB))
365 return op->emitOpError("failed to deduce operand PTX types");
366 std::optional<NVVM::MMATypes> ptxTypeC =
367 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
368 /*isAccumulator=*/true);
369 if (!ptxTypeC)
370 return op->emitError(
371 "could not infer the PTX type for the accumulator/result");
372
373 // TODO: add an attribute to the op to customize this behavior.
374 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
375 if (isa<IntegerType>(aType.getElementType()))
376 overflow = NVVM::MMAIntOverflow::satfinite;
377
378 SmallVector<Value> matA =
379 unpackOperandVector(b, adaptor.getMatrixA(), *ptxTypeA);
380 SmallVector<Value> matB =
381 unpackOperandVector(b, adaptor.getMatrixB(), *ptxTypeB);
382 SmallVector<Value> matC =
383 unpackOperandVector(b, adaptor.getMatrixC(), *ptxTypeC);
384
385 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
386 Type intrinsicResTy = inferIntrinsicResultType(
387 typeConverter->convertType(op->getResultTypes()[0]));
388 Value intrinsicResult =
389 NVVM::MmaOp::create(b, intrinsicResTy, matA, matB, matC,
390 /*shape=*/gemmShape,
391 /*b1Op=*/std::nullopt,
392 /*intOverflow=*/overflow,
393 /*multiplicandPtxTypes=*/
394 std::array<NVVM::MMATypes, 2>{*ptxTypeA, *ptxTypeB},
395 /*multiplicandLayouts=*/
396 std::array<NVVM::MMALayout, 2>{
397 NVVM::MMALayout::row, NVVM::MMALayout::col});
398 rewriter.replaceOp(op, convertIntrinsicResult(op.getLoc(), intrinsicResTy,
399 desiredRetTy, intrinsicResult,
400 rewriter));
401 return success();
402 }
403};
404
405struct ConvertNVGPUToNVVMPass
406 : public impl::ConvertNVGPUToNVVMPassBase<ConvertNVGPUToNVVMPass> {
407 using Base::Base;
408
409 void runOnOperation() override {
410 LowerToLLVMOptions options(&getContext());
411 RewritePatternSet patterns(&getContext());
412 LLVMTypeConverter converter(&getContext(), options);
413 IRRewriter rewriter(&getContext());
415
416 /// device-side async tokens cannot be materialized in nvvm. We just
417 /// convert them to a dummy i32 type in order to easily drop them during
418 /// conversion.
419 converter.addConversion([&](nvgpu::DeviceAsyncTokenType type) -> Type {
420 return converter.convertType(IntegerType::get(type.getContext(), 32));
421 });
422 converter.addConversion([&](nvgpu::WarpgroupAccumulatorType type) -> Type {
423 Type elemType = type.getFragmented().getElementType();
424 int64_t sizeM = type.getFragmented().getDimSize(0);
425 int64_t sizeN = type.getFragmented().getDimSize(1);
426
427 unsigned numMembers;
428 if (elemType.isF32() || elemType.isInteger(32))
429 numMembers = sizeN / 2;
430 else if (elemType.isF16())
431 numMembers = sizeN / 4;
432 else
433 llvm_unreachable("unsupported type for warpgroup accumulator");
434
435 SmallVector<Type> innerStructBody;
436 for (unsigned i = 0; i < numMembers; i++)
437 innerStructBody.push_back(elemType);
438 auto innerStructType =
439 LLVM::LLVMStructType::getLiteral(type.getContext(), innerStructBody);
440
441 SmallVector<Type> structBody;
442 for (int i = 0; i < sizeM; i += kWgmmaSizeM)
443 structBody.push_back(innerStructType);
444
445 auto convertedType =
446 LLVM::LLVMStructType::getLiteral(type.getContext(), structBody);
447 return converter.convertType(convertedType);
448 });
449 converter.addConversion([&](nvgpu::MBarrierTokenType type) -> Type {
450 return converter.convertType(IntegerType::get(type.getContext(), 64));
451 });
452 converter.addConversion(
453 [&](nvgpu::WarpgroupMatrixDescriptorType type) -> Type {
454 return converter.convertType(IntegerType::get(type.getContext(), 64));
455 });
456 converter.addConversion([&](nvgpu::MBarrierGroupType type) -> Type {
457 return converter.convertType(
458 nvgpu::getMBarrierMemrefType(rewriter.getContext(), type));
459 });
460 converter.addConversion([&](nvgpu::TensorMapDescriptorType type) -> Type {
461 return LLVM::LLVMPointerType::get(type.getContext());
462 });
463 populateNVGPUToNVVMConversionPatterns(converter, patterns);
464 LLVMConversionTarget target(getContext());
465 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
466 target.addLegalDialect<::mlir::arith::ArithDialect>();
467 target.addLegalDialect<::mlir::memref::MemRefDialect>();
468 target.addLegalDialect<::mlir::NVVM::NVVMDialect>();
469 target.addLegalDialect<::mlir::vector::VectorDialect>();
471 converter, patterns, target);
472 if (failed(applyPartialConversion(getOperation(), target,
473 std::move(patterns))))
474 signalPassFailure();
475 }
476};
477
478/// Returns the constraints for the sparse MMA inline assembly instruction.
479static std::string buildMmaSparseAsmConstraintString(unsigned matASize,
480 unsigned matBSize,
481 unsigned matCSize) {
482 std::string str;
483 llvm::raw_string_ostream ss(str);
484 for (unsigned i = 0; i < matCSize; i++)
485 ss << "=r,";
486 for (unsigned i = 0; i < matASize + matBSize + matCSize; i++)
487 ss << "r,";
488 // The final operand is for the sparsity metadata.
489 // The sparsity selector appears as direct literal.
490 ss << "r";
491 return str;
492}
493
494/// Returns the string for the `mma.sp.sync` instruction that corresponds to
495/// the given parameters. Note that this function doesn't do any validation,
496/// it's expected that the provided parameters correspond to a valid
497/// instruction.
498static std::string buildMmaSparseAsmString(
499 const std::array<int64_t, 3> &shape, unsigned matASize, unsigned matBSize,
500 unsigned matCSize, NVVM::MMATypes ptxTypeA, NVVM::MMATypes ptxTypeB,
501 NVVM::MMATypes ptxTypeC, NVVM::MMATypes ptxTypeD,
502 std::optional<NVVM::MMAIntOverflow> overflow, unsigned metaDataSelector) {
503 auto ptxTypeStr = [](NVVM::MMATypes ptxType) {
504 return NVVM::stringifyMMATypes(ptxType);
505 };
506
507 std::string asmStr;
508 llvm::raw_string_ostream ss(asmStr);
509 ss << "mma.sp.sync.aligned.m" << shape[0] << "n" << shape[1] << "k"
510 << shape[2] << ".row.col.";
511
512 if (overflow)
513 ss << NVVM::stringifyMMAIntOverflow(*overflow) << ".";
514
515 ss << ptxTypeStr(ptxTypeD) << "." << ptxTypeStr(ptxTypeA) << "."
516 << ptxTypeStr(ptxTypeB) << "." << ptxTypeStr(ptxTypeC) << " ";
517 unsigned asmArgIdx = 0;
518
519 // The operand string is structured into sections `{matC elements...},
520 // {matA elements...}, {matB elements...}, {matC elements}`.
521 for (const auto arrSize : {matCSize, matASize, matBSize, matCSize}) {
522 ss << "{";
523 for (unsigned i = 0; i < arrSize; i++)
524 ss << "$" << asmArgIdx++ << (i < arrSize - 1 ? "," : "");
525 ss << "},";
526 }
527 ss << "$" << asmArgIdx++ << ",";
528 assert(metaDataSelector <= 1);
529 ss << "0x" << metaDataSelector << ";";
530 return asmStr;
531}
532
533/// Builds an inline assembly operation corresponding to the specified MMA
534/// sparse sync operation.
535static FailureOr<LLVM::InlineAsmOp> emitMmaSparseSyncOpAsm(
536 ImplicitLocOpBuilder &b, NVVM::MMATypes ptxTypeA, NVVM::MMATypes ptxTypeB,
537 NVVM::MMATypes ptxTypeC, NVVM::MMATypes ptxTypeD,
538 std::optional<NVVM::MMAIntOverflow> overflow, ArrayRef<Value> unpackedAData,
539 ArrayRef<Value> unpackedB, ArrayRef<Value> unpackedC, Value indexData,
540 int64_t metadataSelector, const std::array<int64_t, 3> &shape,
541 Type intrinsicResultType) {
542 auto asmDialectAttr =
543 LLVM::AsmDialectAttr::get(b.getContext(), LLVM::AsmDialect::AD_ATT);
544
545 const unsigned matASize = unpackedAData.size();
546 const unsigned matBSize = unpackedB.size();
547 const unsigned matCSize = unpackedC.size();
548
549 std::string asmStr = buildMmaSparseAsmString(
550 shape, matASize, matBSize, matCSize, ptxTypeA, ptxTypeB, ptxTypeC,
551 ptxTypeD, overflow, metadataSelector);
552 std::string constraintStr =
553 buildMmaSparseAsmConstraintString(matASize, matBSize, matCSize);
554
555 SmallVector<Value> asmVals;
556 asmVals.reserve(matASize + matBSize + matCSize + 1);
557 for (ArrayRef<Value> args : {unpackedAData, unpackedB, unpackedC})
558 llvm::append_range(asmVals, args);
559 asmVals.push_back(indexData);
560
561 return LLVM::InlineAsmOp::create(b,
562 /*resultTypes=*/intrinsicResultType,
563 /*operands=*/asmVals,
564 /*asm_string=*/asmStr,
565 /*constraints=*/constraintStr,
566 /*has_side_effects=*/true,
567 /*is_align_stack=*/false,
568 LLVM::TailCallKind::None,
569 /*asm_dialect=*/asmDialectAttr,
570 /*operand_attrs=*/ArrayAttr());
571}
572
573/// Lowers `nvgpu.mma.sp.sync` to inline assembly.
574struct NVGPUMmaSparseSyncLowering
575 : public ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp> {
576 using ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp>::ConvertOpToLLVMPattern;
577
578 LogicalResult
579 matchAndRewrite(nvgpu::MmaSparseSyncOp op, OpAdaptor adaptor,
580 ConversionPatternRewriter &rewriter) const override {
581 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
582 // Get the shapes of the MMAMatrix type being used. The shapes will
583 // choose which intrinsic this op will be lowered to.
584 VectorType aType = op.getMatrixA().getType();
585 VectorType bType = op.getMatrixB().getType();
586 VectorType cType = op.getMatrixC().getType();
587
588 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
589 if (failed(ptxTypeA))
590 return op->emitOpError("failed to deduce operand PTX types");
591 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
592 if (failed(ptxTypeB))
593 return op->emitOpError("failed to deduce operand PTX types");
594 std::optional<NVVM::MMATypes> ptxTypeC =
595 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
596 /*isAccumulator=*/true);
597 if (!ptxTypeC)
598 return op->emitError(
599 "could not infer the PTX type for the accumulator/result");
600
601 // Same as `mma.sync`, F32 works only with TensorFloat32 (TF32).
602 bool tf32Enabled = op->hasAttr(op.getTf32EnabledAttrName());
603 if (aType.getElementType().isF32() && !tf32Enabled)
604 return failure();
605
606 // TODO: add an attribute to the op to customize this behavior.
607 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
608 if (isa<IntegerType>(aType.getElementType()))
609 overflow = NVVM::MMAIntOverflow::satfinite;
610
611 SmallVector<Value> matA =
612 unpackOperandVector(b, adaptor.getMatrixA(), *ptxTypeA);
613 SmallVector<Value> matB =
614 unpackOperandVector(b, adaptor.getMatrixB(), *ptxTypeB);
615 SmallVector<Value> matC =
616 unpackOperandVector(b, adaptor.getMatrixC(), *ptxTypeC);
617
618 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
619 Type intrinsicResTy = inferIntrinsicResultType(
620 typeConverter->convertType(op->getResultTypes()[0]));
621
622 // Bitcast the sparse metadata from vector<2xf16> to an i32.
623 Value sparseMetadata = adaptor.getSparseMetadata();
624 if (sparseMetadata.getType() != VectorType::get(2, rewriter.getI16Type()))
625 return op->emitOpError() << "Expected metadata type to be LLVM "
626 "VectorType of 2 i16 elements";
627 sparseMetadata =
628 LLVM::BitcastOp::create(b, rewriter.getI32Type(), sparseMetadata);
629
630 FailureOr<LLVM::InlineAsmOp> intrinsicResult = emitMmaSparseSyncOpAsm(
631 b, *ptxTypeA, *ptxTypeB, *ptxTypeC, *ptxTypeC, overflow, matA, matB,
632 matC, sparseMetadata, op.getSparsitySelector(), op.getMmaShapeAsArray(),
633 intrinsicResTy);
634 if (failed(intrinsicResult))
635 return failure();
636
637 assert((*intrinsicResult).getNumResults() == 1 &&
638 "expected inline asm op returns a single LLVM struct type");
639 rewriter.replaceOp(
640 op, convertIntrinsicResult(op.getLoc(), intrinsicResTy, desiredRetTy,
641 (*intrinsicResult)->getResult(0), rewriter));
642 return success();
643 }
644};
645
646struct NVGPUAsyncCopyLowering
647 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncCopyOp> {
648 using ConvertOpToLLVMPattern<
649 nvgpu::DeviceAsyncCopyOp>::ConvertOpToLLVMPattern;
650
651 LogicalResult
652 matchAndRewrite(nvgpu::DeviceAsyncCopyOp op, OpAdaptor adaptor,
653 ConversionPatternRewriter &rewriter) const override {
654 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
655 Location loc = op.getLoc();
656 auto dstMemrefType = cast<MemRefType>(op.getDst().getType());
657 Value dstPtr =
658 getStridedElementPtr(rewriter, b.getLoc(), dstMemrefType,
659 adaptor.getDst(), adaptor.getDstIndices());
660 FailureOr<unsigned> dstAddressSpace =
661 getTypeConverter()->getMemRefAddressSpace(dstMemrefType);
662 if (failed(dstAddressSpace))
663 return rewriter.notifyMatchFailure(
664 loc, "destination memref address space not convertible to integer");
665
666 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
667 FailureOr<unsigned> srcAddressSpace =
668 getTypeConverter()->getMemRefAddressSpace(srcMemrefType);
669 if (failed(srcAddressSpace))
670 return rewriter.notifyMatchFailure(
671 loc, "source memref address space not convertible to integer");
672
673 Value scrPtr =
674 getStridedElementPtr(rewriter, loc, srcMemrefType, adaptor.getSrc(),
675 adaptor.getSrcIndices());
676 // Intrinsics takes a global pointer so we need an address space cast.
677 auto srcPointerGlobalType = LLVM::LLVMPointerType::get(
678 op->getContext(), static_cast<unsigned>(NVVM::NVVMMemorySpace::Global));
679 scrPtr = LLVM::AddrSpaceCastOp::create(b, srcPointerGlobalType, scrPtr);
680 int64_t dstElements = adaptor.getDstElements().getZExtValue();
681 int64_t sizeInBytes =
682 (dstMemrefType.getElementTypeBitWidth() * dstElements) / 8;
683 // When the optional SrcElements argument is *not* present, the regular
684 // CpAsyncOp is generated. CopyAsyncOp reads bytes from source (global
685 // memory) to fill DstElements number of elements in the destination
686 // (shared memory).
687 Value srcBytes = adaptor.getSrcElements();
688 if (srcBytes) {
689 // When the optional SrcElements argument is present, the source (global
690 // memory) of CpAsyncOp is read only for SrcElements number of elements.
691 // The rest of the DstElements in the destination (shared memory) are
692 // filled with zeros.
693 Value c3I32 =
694 LLVM::ConstantOp::create(b, b.getI32Type(), b.getI32IntegerAttr(3));
695 Value bitwidth = LLVM::ConstantOp::create(
696 b, b.getI32Type(),
697 b.getI32IntegerAttr(srcMemrefType.getElementTypeBitWidth()));
698 Value srcElementsI32 = LLVM::TruncOp::create(b, b.getI32Type(), srcBytes);
699 srcBytes = LLVM::LShrOp::create(
700 b, LLVM::MulOp::create(b, bitwidth, srcElementsI32), c3I32);
701 }
702 // Cache global (.cg) for 16 dst bytes, Cache all (.ca) for sizes other than
703 // 16 dst bytes.
704 NVVM::LoadCacheModifierKind cacheModifier =
705 (op.getBypassL1().value_or(false) && sizeInBytes == 16)
706 ? NVVM::LoadCacheModifierKind::CG
707 : NVVM::LoadCacheModifierKind::CA;
708
709 NVVM::CpAsyncOp::create(
710 b, dstPtr, scrPtr, rewriter.getI32IntegerAttr(sizeInBytes),
711 NVVM::LoadCacheModifierKindAttr::get(op->getContext(), cacheModifier),
712 srcBytes);
713
714 // Drop the result token.
715 Value zero =
716 LLVM::ConstantOp::create(b, IntegerType::get(op.getContext(), 32),
717 rewriter.getI32IntegerAttr(0));
718 rewriter.replaceOp(op, zero);
719 return success();
720 }
721};
722
723struct NVGPUAsyncCreateGroupLowering
724 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncCreateGroupOp> {
725 using ConvertOpToLLVMPattern<
726 nvgpu::DeviceAsyncCreateGroupOp>::ConvertOpToLLVMPattern;
727
728 LogicalResult
729 matchAndRewrite(nvgpu::DeviceAsyncCreateGroupOp op, OpAdaptor adaptor,
730 ConversionPatternRewriter &rewriter) const override {
731 NVVM::CpAsyncCommitGroupOp::create(rewriter, op.getLoc());
732 // Drop the result token.
733 Value zero = LLVM::ConstantOp::create(rewriter, op->getLoc(),
734 IntegerType::get(op.getContext(), 32),
735 rewriter.getI32IntegerAttr(0));
736 rewriter.replaceOp(op, zero);
737 return success();
738 }
739};
740
741struct NVGPUAsyncWaitLowering
742 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncWaitOp> {
743 using ConvertOpToLLVMPattern<
744 nvgpu::DeviceAsyncWaitOp>::ConvertOpToLLVMPattern;
745
746 LogicalResult
747 matchAndRewrite(nvgpu::DeviceAsyncWaitOp op, OpAdaptor adaptor,
748 ConversionPatternRewriter &rewriter) const override {
749 // If numGroup is not present pick 0 as a conservative correct value.
750 int32_t numGroups = adaptor.getNumGroups().value_or(0);
751 NVVM::CpAsyncWaitGroupOp::create(rewriter, op.getLoc(), numGroups);
752 rewriter.eraseOp(op);
753 return success();
754 }
755};
756
757/// Creates mbarrier object in shared memory
758struct NVGPUMBarrierCreateLowering
759 : public ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp> {
760 using ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp>::ConvertOpToLLVMPattern;
761
762 template <typename moduleT>
763 memref::GlobalOp generateGlobalBarrier(ConversionPatternRewriter &rewriter,
764 Operation *funcOp, moduleT moduleOp,
765 MemRefType barrierType) const {
766 SymbolTable symbolTable(moduleOp);
767 OpBuilder::InsertionGuard guard(rewriter);
768 rewriter.setInsertionPoint(&moduleOp.front());
769 auto global = memref::GlobalOp::create(
770 rewriter, funcOp->getLoc(), "__mbarrier",
771 /*sym_visibility=*/rewriter.getStringAttr("private"),
772 /*type=*/barrierType,
773 /*initial_value=*/ElementsAttr(),
774 /*constant=*/false,
775 /*alignment=*/rewriter.getI64IntegerAttr(8));
776 symbolTable.insert(global);
777 return global;
778 }
779
780 LogicalResult
781 matchAndRewrite(nvgpu::MBarrierCreateOp op, OpAdaptor adaptor,
782 ConversionPatternRewriter &rewriter) const override {
783 Operation *funcOp = op->getParentOp();
784 MemRefType barrierType = nvgpu::getMBarrierMemrefType(
785 rewriter.getContext(), op.getBarriers().getType());
786
787 memref::GlobalOp global;
788 if (auto moduleOp = funcOp->getParentOfType<gpu::GPUModuleOp>())
789 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
790 else if (auto moduleOp = funcOp->getParentOfType<ModuleOp>())
791 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
792
793 rewriter.setInsertionPoint(op);
794 rewriter.replaceOpWithNewOp<memref::GetGlobalOp>(op, barrierType,
795 global.getName());
796 return success();
797 }
798};
799
800/// Base class for lowering mbarrier operations to nvvm intrinsics.
801template <typename SourceOp>
802struct MBarrierBasePattern : public ConvertOpToLLVMPattern<SourceOp> {
803public:
804 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
805 /// Returns the base pointer of the mbarrier object.
806 Value getMbarrierPtr(ImplicitLocOpBuilder &b,
807 nvgpu::MBarrierGroupType mbarType, Value memrefDesc,
808 Value mbarId,
809 ConversionPatternRewriter &rewriter) const {
810 MemRefType mbarrierMemrefType =
811 nvgpu::getMBarrierMemrefType(rewriter.getContext(), mbarType);
813 rewriter, b.getLoc(), mbarrierMemrefType, memrefDesc, {mbarId});
814 }
815};
816
817struct NVGPUMBarrierGetLowering
818 : public MBarrierBasePattern<nvgpu::MBarrierGetOp> {
819 using MBarrierBasePattern<nvgpu::MBarrierGetOp>::MBarrierBasePattern;
820
821 LogicalResult
822 matchAndRewrite(nvgpu::MBarrierGetOp op, OpAdaptor adaptor,
823 ConversionPatternRewriter &rewriter) const override {
824 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
825 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
826 rewriter.setInsertionPoint(op);
827 Value barrier = getMbarrierPtr(b, mbarrierType, adaptor.getBarriers(),
828 adaptor.getMbarId(), rewriter);
829 Type resType = op.getMbarrierPointer().getType();
830 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(op, resType, barrier);
831 return success();
832 }
833};
834
835/// Lowers `nvgpu.mbarrier.init` to `nvvm.mbarrier.init`
836struct NVGPUMBarrierInitLowering
837 : public MBarrierBasePattern<nvgpu::MBarrierInitOp> {
838 using MBarrierBasePattern<nvgpu::MBarrierInitOp>::MBarrierBasePattern;
839
840 LogicalResult
841 matchAndRewrite(nvgpu::MBarrierInitOp op, OpAdaptor adaptor,
842 ConversionPatternRewriter &rewriter) const override {
843 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
844 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
845 rewriter.setInsertionPoint(op);
846 Value barrier = getMbarrierPtr(b, mbarrierType, adaptor.getBarriers(),
847 adaptor.getMbarId(), rewriter);
848 Value count = truncToI32(b, adaptor.getCount());
849 rewriter.replaceOpWithNewOp<NVVM::MBarrierInitOp>(op, barrier, count,
850 adaptor.getPredicate());
851 return success();
852 }
853};
854
855/// Lowers `nvgpu.mbarrier.arrive` to `nvvm.mbarrier.arrive`
856struct NVGPUMBarrierArriveLowering
857 : public MBarrierBasePattern<nvgpu::MBarrierArriveOp> {
858 using MBarrierBasePattern<nvgpu::MBarrierArriveOp>::MBarrierBasePattern;
859 LogicalResult
860 matchAndRewrite(nvgpu::MBarrierArriveOp op, OpAdaptor adaptor,
861 ConversionPatternRewriter &rewriter) const override {
862 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
863 Value barrier =
864 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
865 adaptor.getMbarId(), rewriter);
866 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveOp>(op, barrier);
867 return success();
868 }
869};
870
871/// Lowers `nvgpu.mbarrier.arrive.nocomplete` to
872/// `nvvm.mbarrier.arrive.nocomplete`
873struct NVGPUMBarrierArriveNoCompleteLowering
874 : public MBarrierBasePattern<nvgpu::MBarrierArriveNoCompleteOp> {
875 using MBarrierBasePattern<
876 nvgpu::MBarrierArriveNoCompleteOp>::MBarrierBasePattern;
877 LogicalResult
878 matchAndRewrite(nvgpu::MBarrierArriveNoCompleteOp op, OpAdaptor adaptor,
879 ConversionPatternRewriter &rewriter) const override {
880 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
881 Value barrier =
882 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
883 adaptor.getMbarId(), rewriter);
884 Type tokenType = getTypeConverter()->convertType(
885 nvgpu::MBarrierTokenType::get(op->getContext()));
886 Value count = truncToI32(b, adaptor.getCount());
887 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveNocompleteOp>(
888 op, tokenType, barrier, count);
889 return success();
890 }
891};
892
893/// Lowers `nvgpu.mbarrier.test.wait` to `nvvm.mbarrier.test.wait`
894struct NVGPUMBarrierTestWaitLowering
895 : public MBarrierBasePattern<nvgpu::MBarrierTestWaitOp> {
896 using MBarrierBasePattern<nvgpu::MBarrierTestWaitOp>::MBarrierBasePattern;
897 LogicalResult
898 matchAndRewrite(nvgpu::MBarrierTestWaitOp op, OpAdaptor adaptor,
899 ConversionPatternRewriter &rewriter) const override {
900 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
901 Value barrier =
902 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
903 adaptor.getMbarId(), rewriter);
904 Type retType = rewriter.getI1Type();
905 rewriter.replaceOpWithNewOp<NVVM::MBarrierTestWaitOp>(op, retType, barrier,
906 adaptor.getToken());
907 return success();
908 }
909};
910
911struct NVGPUMBarrierArriveExpectTxLowering
912 : public MBarrierBasePattern<nvgpu::MBarrierArriveExpectTxOp> {
913 using MBarrierBasePattern<
914 nvgpu::MBarrierArriveExpectTxOp>::MBarrierBasePattern;
915 LogicalResult
916 matchAndRewrite(nvgpu::MBarrierArriveExpectTxOp op, OpAdaptor adaptor,
917 ConversionPatternRewriter &rewriter) const override {
918 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
919 Value barrier =
920 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
921 adaptor.getMbarId(), rewriter);
922 Value txcount = truncToI32(b, adaptor.getTxcount());
923 NVVM::MBarrierArriveExpectTxOp::create(
924 rewriter, op->getLoc(), barrier, txcount, // barrier and txcount
925 NVVM::MemScopeKind::CTA, // default scope is CTA
926 false, // relaxed-semantics is false
927 adaptor.getPredicate());
928 rewriter.eraseOp(op);
929 return success();
930 }
931};
932
933struct NVGPUMBarrierTryWaitParityLowering
934 : public MBarrierBasePattern<nvgpu::MBarrierTryWaitParityOp> {
935 using MBarrierBasePattern<
936 nvgpu::MBarrierTryWaitParityOp>::MBarrierBasePattern;
937 LogicalResult
938 matchAndRewrite(nvgpu::MBarrierTryWaitParityOp op, OpAdaptor adaptor,
939 ConversionPatternRewriter &rewriter) const override {
940 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
941 Value barrier =
942 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
943 adaptor.getMbarId(), rewriter);
944 Value ticks = truncToI32(b, adaptor.getTicks());
945 Value phase =
946 LLVM::ZExtOp::create(b, b.getI32Type(), adaptor.getPhaseParity());
947 rewriter.replaceOpWithNewOp<NVVM::MBarrierTryWaitParityOp>(op, barrier,
948 phase, ticks);
949 return success();
950 }
951};
952
953struct NVGPUTmaAsyncLoadOpLowering
954 : public MBarrierBasePattern<nvgpu::TmaAsyncLoadOp> {
955 using MBarrierBasePattern<nvgpu::TmaAsyncLoadOp>::MBarrierBasePattern;
956 LogicalResult
957 matchAndRewrite(nvgpu::TmaAsyncLoadOp op, OpAdaptor adaptor,
958 ConversionPatternRewriter &rewriter) const override {
959 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
960 auto srcMemrefType = cast<MemRefType>(op.getDst().getType());
961 Value dest = getStridedElementPtr(rewriter, op->getLoc(), srcMemrefType,
962 adaptor.getDst(), {});
963 // Intrinsics takes a shared-cluster pointer so we need an
964 // address space cast from 3 to 7.
965 // TODO: Introduce AS(7) in NVGPU.
966 auto ptrSharedClusterType = LLVM::LLVMPointerType::get(
967 op->getContext(),
968 static_cast<unsigned>(NVVM::NVVMMemorySpace::SharedCluster));
969 dest = LLVM::AddrSpaceCastOp::create(b, ptrSharedClusterType, dest);
970
971 Value barrier =
972 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
973 adaptor.getMbarId(), rewriter);
974
975 SmallVector<Value> coords = adaptor.getCoordinates();
976 for (auto [index, value] : llvm::enumerate(coords)) {
977 coords[index] = truncToI32(b, value);
978 }
979
980 // TODO: Enhance the NVGPU Op for other modes too
981 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(
982 op, dest, adaptor.getTensorMapDescriptor(), coords, barrier,
983 ValueRange{}, adaptor.getMulticastMask(), Value{},
984 NVVM::TMALoadMode::TILE, // default is TILE mode
985 false, // default is cluster-scope
986 nullptr, // default is no cta-group
987 adaptor.getPredicate());
988 return success();
989 }
990};
991
992struct NVGPUTmaAsyncStoreOpLowering
993 : public MBarrierBasePattern<nvgpu::TmaAsyncStoreOp> {
994 using MBarrierBasePattern<nvgpu::TmaAsyncStoreOp>::MBarrierBasePattern;
995 LogicalResult
996 matchAndRewrite(nvgpu::TmaAsyncStoreOp op, OpAdaptor adaptor,
997 ConversionPatternRewriter &rewriter) const override {
998 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
999 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
1000 Value dest = getStridedElementPtr(rewriter, op->getLoc(), srcMemrefType,
1001 adaptor.getSrc(), {});
1002 SmallVector<Value> coords = adaptor.getCoordinates();
1003 for (auto [index, value] : llvm::enumerate(coords)) {
1004 coords[index] = truncToI32(b, value);
1005 }
1006
1007 // TODO: Enhance the NVGPU Op for other modes too
1008 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(
1009 op, adaptor.getTensorMapDescriptor(), dest, coords, Value{},
1010 NVVM::TMAStoreMode::TILE, // default is TILE mode
1011 adaptor.getPredicate());
1012 return success();
1013 }
1014};
1015
1016struct NVGPUGenerateWarpgroupDescriptorLowering
1017 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupGenerateDescriptorOp> {
1018 using ConvertOpToLLVMPattern<
1019 nvgpu::WarpgroupGenerateDescriptorOp>::ConvertOpToLLVMPattern;
1020
1021 LogicalResult
1022 matchAndRewrite(nvgpu::WarpgroupGenerateDescriptorOp op, OpAdaptor adaptor,
1023 ConversionPatternRewriter &rewriter) const override {
1024
1025 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1026
1027 nvgpu::TensorMapSwizzleKind swizzleKind =
1028 op.getTensorMap().getType().getSwizzle();
1029
1030 unsigned layout =
1031 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 128
1032 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 64
1033 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 32
1034 : 1;
1035 unsigned swizzle =
1036 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 1
1037 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 2
1038 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 3
1039 : 0;
1040
1041 auto ti64 = b.getIntegerType(64);
1042 auto makeConst = [&](uint64_t index) -> Value {
1043 return LLVM::ConstantOp::create(b, ti64, b.getI64IntegerAttr(index));
1044 };
1045 auto shiftLeft = [&](Value value, unsigned shift) -> Value {
1046 return LLVM::ShlOp::create(b, ti64, value, makeConst(shift));
1047 };
1048 auto shiftRight = [&](Value value, unsigned shift) -> Value {
1049 return LLVM::LShrOp::create(b, ti64, value, makeConst(shift));
1050 };
1051 auto insertBit = [&](Value desc, Value val, int startBit) {
1052 return LLVM::OrOp::create(b, ti64, desc, shiftLeft(val, startBit));
1053 };
1054
1055 int64_t sizeN = op.getTensorMap().getType().getTensor().getDimSize(0);
1056 uint64_t strideDimVal = (layout << 3) >> exclude4LSB;
1057 uint64_t leadDimVal = (sizeN * layout) >> exclude4LSB;
1058 uint64_t offsetVal = 0;
1059
1060 Value strideDim = makeConst(strideDimVal);
1061 Value leadDim = makeConst(leadDimVal);
1062
1063 Value baseAddr = getStridedElementPtr(
1064 rewriter, op->getLoc(), cast<MemRefType>(op.getTensor().getType()),
1065 adaptor.getTensor(), {});
1066 Value basePtr = LLVM::PtrToIntOp::create(b, ti64, baseAddr);
1067 // Just use 14 bits for base address
1068 Value basePtr14bit = shiftRight(shiftLeft(basePtr, 46), 50);
1069
1070 int startSwizzleBit = 62, startOffsetBit = 49, startStrideBit = 32,
1071 startLeadBit = 16, startBaseAddrBit = 0;
1072 Value dsc = makeConst(0);
1073 // // [62,64) swizzle type
1074 dsc = insertBit(dsc, makeConst(swizzle), startSwizzleBit);
1075 // // [49,52) base_offset
1076 dsc = insertBit(dsc, makeConst(offsetVal), startOffsetBit);
1077 // // [32,46) stride
1078 dsc = insertBit(dsc, strideDim, startStrideBit);
1079 // // [16,30) leading dimension
1080 dsc = insertBit(dsc, leadDim, startLeadBit);
1081 // // [0,14) start_address
1082 dsc = insertBit(dsc, basePtr14bit, startBaseAddrBit);
1083
1084 LDBG() << "Generating warpgroup.descriptor: " << "leading_off:"
1085 << leadDimVal << "\t" << "stride_off :" << strideDimVal << "\t"
1086 << "base_offset:" << offsetVal << "\t" << "layout_type:" << swizzle
1087 << " (" << nvgpu::stringifyTensorMapSwizzleKind(swizzleKind)
1088 << ")\n start_addr : " << baseAddr;
1089
1090 rewriter.replaceOp(op, dsc);
1091 return success();
1092 }
1093};
1094
1095static Value makeI64Const(ImplicitLocOpBuilder &b, int32_t index) {
1096 return LLVM::ConstantOp::create(b, b.getIntegerType(64),
1097 b.getI32IntegerAttr(index));
1098}
1099
1100/// Returns a Value that holds data type enum that is expected by CUDA driver.
1101static Value elementTypeAsLLVMConstant(ImplicitLocOpBuilder &b, Type type) {
1102 // Enum is from CUDA driver API
1103 // https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__TYPES.html
1104 enum CUtensorMapDataTypeEnum {
1105 CU_TENSOR_MAP_DATA_TYPE_UINT8 = 0,
1106 CU_TENSOR_MAP_DATA_TYPE_UINT16,
1107 CU_TENSOR_MAP_DATA_TYPE_UINT32,
1108 CU_TENSOR_MAP_DATA_TYPE_INT32,
1109 CU_TENSOR_MAP_DATA_TYPE_UINT64,
1110 CU_TENSOR_MAP_DATA_TYPE_INT64,
1111 CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
1112 CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
1113 CU_TENSOR_MAP_DATA_TYPE_FLOAT64,
1114 CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
1115 CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ,
1116 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32,
1117 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ
1118 };
1119
1120 if (type.isUnsignedInteger(8))
1121 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT8);
1122 if (type.isUnsignedInteger(16))
1123 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT16);
1124 if (type.isUnsignedInteger(32))
1125 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT32);
1126 if (type.isUnsignedInteger(64))
1127 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT64);
1128 if (type.isSignlessInteger(32))
1129 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_INT32);
1130 if (type.isSignlessInteger(64))
1131 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_INT64);
1132 if (type.isF16())
1133 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT16);
1134 if (type.isF32())
1135 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT32);
1136 if (type.isF64())
1137 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT64);
1138 if (type.isBF16())
1139 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16);
1140
1141 llvm_unreachable("Not supported data type");
1142}
1143
1144struct NVGPUTmaCreateDescriptorOpLowering
1145 : public ConvertOpToLLVMPattern<nvgpu::TmaCreateDescriptorOp> {
1146 using ConvertOpToLLVMPattern<
1147 nvgpu::TmaCreateDescriptorOp>::ConvertOpToLLVMPattern;
1148 LogicalResult
1149 matchAndRewrite(nvgpu::TmaCreateDescriptorOp op, OpAdaptor adaptor,
1150 ConversionPatternRewriter &rewriter) const override {
1151 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1152 auto llvmPointerType = LLVM::LLVMPointerType::get(op->getContext());
1153 Type llvmInt64Type = IntegerType::get(op->getContext(), 64);
1154
1155 Value tensorElementType =
1156 elementTypeAsLLVMConstant(b, op.getTensor().getType().getElementType());
1157 auto promotedOperands = getTypeConverter()->promoteOperands(
1158 b.getLoc(), op->getOperands(), adaptor.getOperands(), b);
1159
1160 Value boxArrayPtr = LLVM::AllocaOp::create(
1161 b, llvmPointerType, llvmInt64Type, makeI64Const(b, 5));
1162 for (auto [index, value] : llvm::enumerate(adaptor.getBoxDimensions())) {
1163 Value gep = LLVM::GEPOp::create(b, llvmPointerType, llvmPointerType,
1164 boxArrayPtr, makeI64Const(b, index));
1165 LLVM::StoreOp::create(b, value, gep);
1166 }
1167
1168 nvgpu::TensorMapDescriptorType desc = op.getTensorMap().getType();
1169 // Set Arguments for the function call
1170 SmallVector<Value> arguments;
1171 arguments.push_back(promotedOperands[0]); // rank
1172 arguments.push_back(promotedOperands[1]); // descriptor
1173 arguments.push_back(tensorElementType); // data type
1174 arguments.push_back(
1175 makeI64Const(b, (int)desc.getInterleave())); // interleave
1176 arguments.push_back(makeI64Const(b, (int)desc.getSwizzle())); // swizzle
1177 arguments.push_back(makeI64Const(b, (int)desc.getL2promo())); // l2promo
1178 arguments.push_back(makeI64Const(b, (int)desc.getOob())); // oob
1179 arguments.push_back(boxArrayPtr); // box dimensions
1180
1181 // Set data types of the arguments
1182 SmallVector<Type> argTypes = {
1183 llvmInt64Type, /* int64_t tensorRank */
1184 llvmPointerType, /* ptr */
1185 llvmInt64Type, /* int64_t */
1186 llvmInt64Type, /* int64_t */
1187 llvmInt64Type, /* int64_t */
1188 llvmInt64Type, /* int64_t */
1189 llvmInt64Type, /* int64_t */
1190 llvmPointerType /* ptr */
1191 };
1192 FunctionCallBuilder hostRegisterCallBuilder = {
1193 "mgpuTensorMapEncodeTiledMemref", llvmPointerType, argTypes};
1194 Value tensorMap =
1195 hostRegisterCallBuilder.create(b.getLoc(), b, arguments).getResult();
1196
1197 rewriter.replaceOp(op, tensorMap);
1198 return success();
1199 }
1200};
1201
1202struct NVGPUWarpgroupMmaOpLowering
1203 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp> {
1204 using ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp>::ConvertOpToLLVMPattern;
1205
1206 /// This is a helper class to generate required NVVM Ops for warp-group level
1207 /// matrix multiplication.
1208 /// When the given GEMM shape is larger than the shape of
1209 /// a wgmma instrution in PTX, it can generate multiple NVVM::WgmmaMmaAsyncOp
1210 /// Op(s), group and execute them asynchronously. The class also handles
1211 /// waiting for completion and iterates through WarpgroupMatrixDescriptor to
1212 /// create descriptors for each instruction.
1213 ///
1214 /// For example this is the case when the shape of GEMM is 128x128x128
1215 ///
1216 /// nvvm.wgmma.fence.aligned
1217 ///
1218 /// nvvm.wgmma.mma.async descA, descB
1219 /// iterate(descA, descB)
1220 /// nvvm.wgmma.mma.async descA, descB
1221 /// [6x times more]
1222 ///
1223 /// nvvm.wgmma.group.sync.aligned
1224 /// nvvm.wgmma.wait.group.sync [groupId]
1225 ///
1226 class WarpgroupGemm {
1227 nvgpu::WarpgroupMmaOp op;
1228 ImplicitLocOpBuilder b;
1229 OpAdaptor adaptor;
1230
1231 // Entire shape of the given Op
1232 int64_t totalM, totalN, totalK;
1233
1234 // Shape of one wgmma instruction
1235 int wgmmaM = 0, wgmmaN = 0, wgmmaK = 0;
1236
1237 // Iteration counts for GEMM
1238 int iterationM = 0, iterationN = 0, iterationK = 0;
1239
1240 /// The function returns the shape of wgmma instruction that is defined in
1241 /// PTX programming guide.
1242 /// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#asynchronous-warpgroup-level-matrix-shape
1243 void findWgmmaShape(int64_t sizeM, int64_t sizeN, Type inputElemType) {
1244 wgmmaM = 64;
1245 wgmmaN = sizeN;
1246 if (inputElemType.isTF32()) {
1247 wgmmaK = 8;
1248 } else if (inputElemType.isF16() || inputElemType.isBF16()) {
1249 wgmmaK = 16;
1250 } else if (isa<Float8E4M3FNType, Float8E5M2Type>(inputElemType) ||
1251 inputElemType.isInteger(16)) {
1252 wgmmaK = 32;
1253 } else if (inputElemType.isInteger(1)) {
1254 wgmmaK = 256;
1255 } else {
1256 llvm_unreachable("msg: not supported K shape");
1257 }
1258 LDBG() << "Generating WgmmaMmaAsyncOp shape[m = " << wgmmaM
1259 << ", n = " << wgmmaN << ", k = " << wgmmaK << "]";
1260 }
1261
1262 /// Generates WGMMATypesAttr from MLIR Type
1263 NVVM::WGMMATypesAttr generateWgmmaType(Type type,
1264 bool useF32 = false) const {
1265 auto getWgmmaType = [=](Type elemType) {
1266 if (elemType.isF32() || elemType.isTF32())
1267 return useF32 ? NVVM::WGMMATypes::f32 : NVVM::WGMMATypes::tf32;
1268 if (elemType.isF16())
1269 return NVVM::WGMMATypes::f16;
1270 if (elemType.isBF16())
1271 return NVVM::WGMMATypes::bf16;
1272 if (isa<Float8E4M3FNType>(elemType))
1273 return NVVM::WGMMATypes::e4m3;
1274 if (isa<Float8E5M2Type>(elemType))
1275 return NVVM::WGMMATypes::e5m2;
1276 if (elemType.isInteger(1))
1277 return NVVM::WGMMATypes::b1;
1278 if (elemType.isInteger(8))
1279 return NVVM::WGMMATypes::s8;
1280 if (elemType.isUnsignedInteger(8))
1281 return NVVM::WGMMATypes::u8;
1282 if (elemType.isInteger(32))
1283 return NVVM::WGMMATypes::s32;
1284 llvm_unreachable("unsupported type");
1285 };
1286 return NVVM::WGMMATypesAttr::get(op->getContext(), getWgmmaType(type));
1287 }
1288
1289 /// Generates layout attribute for the input matrix for wgmma instruction
1290 NVVM::MMALayoutAttr
1291 generateWgmmaLayout(std::optional<bool> transpose) const {
1292 if (transpose.value_or(false))
1293 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::col);
1294 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::row);
1295 }
1296
1297 /// Generates shape attribute for wgmma instruction
1298 NVVM::MMAShapeAttr generateWgmmaShape() const {
1299 return NVVM::MMAShapeAttr::get(op->getContext(), wgmmaM, wgmmaN, wgmmaK);
1300 }
1301
1302 /// Generates scale attributes of output matrix for wgmma instruction
1303 NVVM::WGMMAScaleOutAttr generateScaleOut() const {
1304 return NVVM::WGMMAScaleOutAttr::get(op->getContext(),
1305 NVVM::WGMMAScaleOut::one);
1306 }
1307 /// Generates scale attributes of input matrix for wgmma instruction
1308 NVVM::WGMMAScaleInAttr generateScaleIn() const {
1309 return NVVM::WGMMAScaleInAttr::get(op->getContext(),
1310 NVVM::WGMMAScaleIn::one);
1311 }
1312
1313 /// Basic function to generate Add
1314 Value makeAdd(Value lhs, Value rhs) {
1315 return LLVM::AddOp::create(b, lhs.getType(), lhs, rhs);
1316 };
1317
1318 /// Moves the descriptor pointer of matrix-A for the next wgmma instruction.
1319 /// Currently, it only handles row-major.
1320 ///
1321 /// It moves the pointer like below for [128][64] size:
1322 /// +2 +4 +6
1323 /// ↓ ↓ ↓
1324 /// descA ---> +--+--+--+--+
1325 /// |->|->|->|->|
1326 /// | | | | |
1327 /// | | | | |
1328 /// | | | | |
1329 /// descA+512---> +-----------+
1330 /// | | | | |
1331 /// | | | | |
1332 /// | | | | |
1333 /// | | | | |
1334 /// +-----------+
1335 ///
1336 Value iterateDescriptorA(Value desc, int i, int j, int k) {
1337 MemRefType matrixTypeA = op.getDescriptorA().getType().getTensor();
1338 Type elemA = matrixTypeA.getElementType();
1339 int byte = elemA.getIntOrFloatBitWidth() / 8;
1340 int tileShapeA = matrixTypeA.getDimSize(1);
1341 int incrementVal = ((wgmmaK * k) + (totalK * tileShapeA * i)) * byte;
1342 incrementVal = incrementVal >> exclude4LSB;
1343 LDBG() << "\t\t[m: " << i << " n: " << j << " k: " << k
1344 << "] [wgmma descriptors] Descriptor A + " << incrementVal
1345 << " | \t ";
1346 if (!incrementVal)
1347 return desc;
1348 return makeAdd(desc, makeI64Const(b, incrementVal));
1349 }
1350
1351 /// Moves the descriptor pointer of matrix-B for the next wgmma instruction.
1352 /// Currently, it only handles column-major.
1353 ///
1354 /// It moves the pointer like below for [128][64] size:
1355 /// descB ---> +--+--+--+--+--+--+--+--+
1356 /// |↓ | | | | | | | |
1357 /// |↓ | | | | | | | |
1358 /// |↓ | | | | | | | |
1359 /// |↓ | | | | | | | |
1360 /// +--+--+--+--+--+--+--+--+
1361 ///
1362 Value iterateDescriptorB(Value desc, int i, int j, int k) {
1363 MemRefType matrixTypeB = op.getDescriptorB().getType().getTensor();
1364 Type elemB = matrixTypeB.getElementType();
1365 int byte = elemB.getIntOrFloatBitWidth() / 8;
1366 int incrementVal = matrixTypeB.getDimSize(0) * wgmmaK * k * byte;
1367 incrementVal = incrementVal >> exclude4LSB;
1368 LDBG() << "Descriptor B + " << incrementVal;
1369 if (!incrementVal)
1370 return desc;
1371 return makeAdd(desc, makeI64Const(b, incrementVal));
1372 }
1373
1374 /// This function generates a WgmmaMmaAsyncOp using provided GMMA matrix
1375 /// descriptors and arranges them based on induction variables: i, j, and k.
1376 Value generateWgmma(int i, int j, int k, Value matrixC) {
1377 LDBG() << "\t wgmma." << "m" << wgmmaM << "n" << wgmmaN << "k" << wgmmaK
1378 << "(A[" << (iterationM * wgmmaM) << ":"
1379 << (iterationM * wgmmaM) + wgmmaM << "][" << (iterationK * wgmmaK)
1380 << ":" << (iterationK * wgmmaK + wgmmaK) << "] * " << " B["
1381 << (iterationK * wgmmaK) << ":" << (iterationK * wgmmaK + wgmmaK)
1382 << "][" << 0 << ":" << wgmmaN << "])";
1383
1384 Value descriptorA = iterateDescriptorA(adaptor.getDescriptorA(), i, j, k);
1385 Value descriptorB = iterateDescriptorB(adaptor.getDescriptorB(), i, j, k);
1386
1387 Type elemA = op.getDescriptorA().getType().getTensor().getElementType();
1388 NVVM::WGMMATypesAttr itypeA = generateWgmmaType(elemA);
1389
1390 Type elemB = op.getDescriptorB().getType().getTensor().getElementType();
1391 NVVM::WGMMATypesAttr itypeB = generateWgmmaType(elemB);
1392
1393 Type elemD = op.getMatrixC().getType().getFragmented().getElementType();
1394 NVVM::WGMMATypesAttr itypeD = generateWgmmaType(elemD, true);
1395
1396 NVVM::MMAShapeAttr shape = generateWgmmaShape();
1397 NVVM::WGMMAScaleOutAttr scaleOut = generateScaleOut();
1398 NVVM::WGMMAScaleInAttr scaleIn = generateScaleIn();
1399 NVVM::MMALayoutAttr layoutA = generateWgmmaLayout(op.getTransposeA());
1400 NVVM::MMALayoutAttr layoutB = generateWgmmaLayout(!op.getTransposeB());
1401
1402 auto overflow = NVVM::MMAIntOverflowAttr::get(
1403 op->getContext(), NVVM::MMAIntOverflow::wrapped);
1404
1405 return NVVM::WgmmaMmaAsyncOp::create(
1406 b, matrixC.getType(), matrixC, descriptorA, descriptorB, shape,
1407 itypeA, itypeB, itypeD, scaleOut, scaleIn, scaleIn, layoutA, layoutB,
1408 overflow);
1409 }
1410
1411 /// Generates multiple wgmma instructions to complete the given GEMM shape
1412 Value generateWgmmaGroup() {
1413 Value wgmmaResult =
1414 LLVM::PoisonOp::create(b, adaptor.getMatrixC().getType());
1415
1416 // Perform GEMM
1417 SmallVector<Value> wgmmaResults;
1418 for (int i = 0; i < iterationM; ++i) {
1419 Value matrixC =
1420 LLVM::ExtractValueOp::create(b, adaptor.getMatrixC(), i);
1421 for (int j = 0; j < iterationN; ++j)
1422 for (int k = 0; k < iterationK; ++k)
1423 matrixC = generateWgmma(i, j, k, matrixC);
1424 wgmmaResults.push_back(matrixC);
1425 }
1426 for (auto [idx, matrix] : llvm::enumerate(wgmmaResults)) {
1427 wgmmaResult = LLVM::InsertValueOp::create(b, wgmmaResult.getType(),
1428 wgmmaResult, matrix, idx);
1429 }
1430 return wgmmaResult;
1431 }
1432
1433 public:
1434 WarpgroupGemm(nvgpu::WarpgroupMmaOp op, ImplicitLocOpBuilder &b,
1435 OpAdaptor adaptor)
1436 : op(op), b(b), adaptor(adaptor) {
1437 // Find the entire GEMM Shape
1438 totalM = op.getDescriptorA().getType().getTensor().getDimSize(0);
1439 totalN = op.getDescriptorB().getType().getTensor().getDimSize(1);
1440 totalK = op.getDescriptorA().getType().getTensor().getDimSize(1);
1441 LDBG() << "===--- GEMM D[" << totalM << "][" << totalN << "] += A["
1442 << totalM << "][" << totalK << "] * B[" << totalK << "][" << totalN
1443 << "] ---===";
1444
1445 // Find the shape for one wgmma instruction
1446 findWgmmaShape(
1447 totalM, totalN,
1448 op.getDescriptorA().getType().getTensor().getElementType());
1449
1450 // Iterations counts to complete the given shape with wgmma shape
1451 iterationM = totalM / wgmmaM;
1452 iterationN = totalN / wgmmaN;
1453 iterationK = totalK / wgmmaK;
1454 }
1455
1456 /// Generates WgmmaMmaAsync Ops to complete the specified GEMM shape. It
1457 /// includes generating a fence Op (WgmmaFenceAlignedOp) before the
1458 /// instructions and group synchronization, as well as waiting
1459 /// (WgmmaGroupSyncAlignedOp) for group synchronization
1460 /// (WgmmaWaitGroupSyncOp) after the instructions.
1461 Value generateWarpgroupMma() {
1462 NVVM::WgmmaFenceAlignedOp::create(b);
1463 Value wgmmaResult = generateWgmmaGroup();
1464 NVVM::WgmmaGroupSyncAlignedOp::create(b);
1465 NVVM::WgmmaWaitGroupSyncOp::create(b, op.getWaitGroup());
1466 return wgmmaResult;
1467 }
1468 };
1469 LogicalResult
1470 matchAndRewrite(nvgpu::WarpgroupMmaOp op, OpAdaptor adaptor,
1471 ConversionPatternRewriter &rewriter) const override {
1472 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1473
1474 // Step 1. Build a helper class
1475 WarpgroupGemm warpgroupGemm(op, b, adaptor);
1476
1477 // Step 2. Get the entire GEMM Shape
1478 Value wgmmaResult = warpgroupGemm.generateWarpgroupMma();
1479
1480 // Step 3. Replace fragmented result struct with the op results
1481 rewriter.replaceOp(op, wgmmaResult);
1482 return success();
1483 }
1484};
1485
1486struct NVGPUWarpgroupMmaStoreOpLowering
1487 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaStoreOp> {
1488 using ConvertOpToLLVMPattern<
1489 nvgpu::WarpgroupMmaStoreOp>::ConvertOpToLLVMPattern;
1490
1491 /// This function stores a fragmented register matrix owned by a warp group
1492 /// (128 threads) into a memref. Each thread has 64 registers, each the size
1493 /// of a struct.
1494 /// Here is what each threads (T) holds, each `d` is struct value with a
1495 /// number.
1496 ///
1497 /// Threads in warp-group (128 threads) and what they owns in the matrixD:
1498 /// 0-31 Warp-0 -> MatrixD[0:15 ][0:N]
1499 /// 32-63 Warp-1 -> MatrixD[16:31][0:N]
1500 /// 64-95 Warp-2 -> MatrixD[32:47][0:N]
1501 /// 96-127 Warp-3 -> MatrixD[48:64][0:N]
1502 ///
1503 /// Matrix-D:
1504 /// +______________________________________________________________________+
1505 /// | 0-1 | 2-3 | 4-5 | 6-7 | 8-9 | 10-11|..|N-8,N-7 |
1506 /// 0 | T0:d0-d1 |T1:d0-d1 |T2:d0-d1 |T3:d0-d1 |T0:d4-d5| T1:d4-d5..|T0:dX-dY|
1507 /// 1 | T4:d0-d1 |T5:d0-d1 |T6:d0-d1 |T7:d0-d1 |T4:d4-d5| T5:d4-d5..|T4:dX-dY|
1508 /// ..| .........|.........|.........|.........|........|...........|........|
1509 /// 8 | T0:d2-d3 |T1:d2-d3 |T2:d2-d3 |T3:d2-d3 |T0:d6-d7|T1:d6-d7,..|T0:dZ-dW|
1510 /// 9 | T4:d2-d3 |T5:d2-d3 |T6:d2-d3 |T7:d2-d3 |T4:d6-d7| T5:d6-d7..|T4:dZ-dW|
1511 /// ..| .........|.........|.........|.........|........|...........|........|
1512 /// 15| T28:d2-d3|T29:d2-d3|T30:d2-d3|T31:d2-d3|........|...........|........|
1513 /// 16| T32:d2-d3|T33:d2-d3|T34:d2-d3|T35:d2-d3|........|...........|........|
1514 /// ..| .........|.........|.........|.........|........|...........|........|
1515 /// 32| T64:d2-d3|T65:d2-d3|T66:d2-d3|T67:d2-d3|........|...........|........|
1516 /// ..| .........|.........|.........|.........|........|...........|........|
1517 /// 48| T96:d2-d3|T97:d2-d3|T98:d2-d3|T99:d2-d3|........|...........|........|
1518 /// ..| .........|.........|.........|.........|........|...........|........|
1519 /// +______________________________________________________________________+
1520 ///
1521 /// \param rewriter: The pattern rewriter.
1522 /// \param matrixD: Result of the warp-group MMA operation (fragmented
1523 /// matrix). It is holded by a thread and a struct with 64 elements.
1524 /// \param dstMemref: The memref where the registers will be stored.
1525 /// \param offset: the offset within the memref where the registers will be
1526 /// stored.
1527 void storeFragmentedMatrix(ImplicitLocOpBuilder &b, Value matrixD,
1528 TypedValue<MemRefType> dstMemref,
1529 int offset) const {
1530 Type i32 = b.getI32Type();
1531
1532 auto makeConst = [&](int32_t index) -> Value {
1533 return LLVM::ConstantOp::create(b, i32, b.getI32IntegerAttr(index));
1534 };
1535 Value c1 = makeConst(1);
1536 Value c2 = makeConst(2);
1537 Value c4 = makeConst(4);
1538 Value c8 = makeConst(8);
1539 Value c16 = makeConst(16);
1540 Value warpSize = makeConst(kWarpSize);
1541
1542 auto makeMul = [&](Value lhs, Value rhs) -> Value {
1543 return LLVM::MulOp::create(b, lhs.getType(), lhs, rhs);
1544 };
1545 auto makeAdd = [&](Value lhs, Value rhs) -> Value {
1546 return LLVM::AddOp::create(b, lhs.getType(), lhs, rhs);
1547 };
1548
1549 auto makeExtractAndStore = [&](int i, Value wgmmaResult, Value x, Value y,
1551 Type it = b.getIndexType();
1552 Value idx = arith::IndexCastOp::create(b, it, x);
1553 Value idy0 = arith::IndexCastOp::create(b, it, y);
1554 Value idy1 = arith::IndexCastOp::create(b, it, makeAdd(y, c1));
1555 Value d0 = LLVM::ExtractValueOp::create(b, wgmmaResult, i);
1556 Value d1 = LLVM::ExtractValueOp::create(b, wgmmaResult, i + 1);
1557 memref::StoreOp::create(b, d0, memref, ValueRange{idx, idy0});
1558 memref::StoreOp::create(b, d1, memref, ValueRange{idx, idy1});
1559 };
1560
1561 Value tidx = NVVM::ThreadIdXOp::create(b, i32);
1562 Value laneId = LLVM::URemOp::create(b, i32, tidx, warpSize);
1563 Value warpId = LLVM::UDivOp::create(b, i32, tidx, warpSize);
1564 Value lane4Id = LLVM::UDivOp::create(b, i32, laneId, c4);
1565 Value lane4modId = LLVM::URemOp::create(b, i32, laneId, c4);
1566
1567 Value tj = makeMul(lane4modId, c2);
1568 Value ti = makeAdd(lane4Id, makeMul(warpId, c16));
1569 if (offset)
1570 ti = makeAdd(ti, makeConst(offset));
1571
1572 auto structType = cast<LLVM::LLVMStructType>(matrixD.getType());
1573
1574 // Number of 32-bit registers owns per thread
1575 constexpr unsigned numAdjacentRegisters = 2;
1576 // Number of 8x8 matrices one below another per warp
1577 constexpr unsigned numStackedMatrices = 2;
1578
1579 size_t storeCount = (structType.getBody().size() /
1580 (numStackedMatrices * numAdjacentRegisters));
1581
1582 for (size_t i = 0; i < numStackedMatrices; ++i) {
1583 Value idx = makeAdd(ti, makeMul(makeConst(i), c8));
1584 for (size_t j = 0; j < storeCount; ++j) {
1585 Value idy = makeAdd(tj, makeMul(makeConst(j), c8));
1586 size_t structIndex = (i * numAdjacentRegisters) +
1587 (j * (numStackedMatrices * numAdjacentRegisters));
1588 makeExtractAndStore(structIndex, matrixD, idx, idy, dstMemref);
1589 }
1590 }
1591 }
1592
1593 LogicalResult
1594 matchAndRewrite(nvgpu::WarpgroupMmaStoreOp op, OpAdaptor adaptor,
1595 ConversionPatternRewriter &rewriter) const override {
1596 int offset = 0;
1597 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1598 Value matriDValue = adaptor.getMatrixD();
1599 auto stype = cast<LLVM::LLVMStructType>(matriDValue.getType());
1600 for (auto [idx, matrixD] : llvm::enumerate(stype.getBody())) {
1601 auto structType = cast<LLVM::LLVMStructType>(matrixD);
1602 Value innerStructValue =
1603 LLVM::ExtractValueOp::create(b, matriDValue, idx);
1604 storeFragmentedMatrix(b, innerStructValue, op.getDstMemref(), offset);
1605 offset += structType.getBody().size();
1606 }
1607 rewriter.eraseOp(op);
1608 return success();
1609 }
1610};
1611
1612struct NVGPUWarpgroupMmaInitAccumulatorOpLowering
1613 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaInitAccumulatorOp> {
1614 using ConvertOpToLLVMPattern<
1615 nvgpu::WarpgroupMmaInitAccumulatorOp>::ConvertOpToLLVMPattern;
1616 LogicalResult
1617 matchAndRewrite(nvgpu::WarpgroupMmaInitAccumulatorOp op, OpAdaptor adaptor,
1618 ConversionPatternRewriter &rewriter) const override {
1619 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1620 LLVM::LLVMStructType packStructType = cast<LLVM::LLVMStructType>(
1621 getTypeConverter()->convertType(op.getMatrixC().getType()));
1622 Type elemType = cast<LLVM::LLVMStructType>(packStructType.getBody().front())
1623 .getBody()
1624 .front();
1625 Value zero = LLVM::ConstantOp::create(b, elemType, b.getZeroAttr(elemType));
1626 Value packStruct = LLVM::PoisonOp::create(b, packStructType);
1627 SmallVector<Value> innerStructs;
1628 // Unpack the structs and set all values to zero
1629 for (auto [idx, s] : llvm::enumerate(packStructType.getBody())) {
1630 auto structType = cast<LLVM::LLVMStructType>(s);
1631 Value structValue = LLVM::ExtractValueOp::create(b, packStruct, idx);
1632 for (unsigned i = 0; i < structType.getBody().size(); ++i) {
1633 structValue = LLVM::InsertValueOp::create(b, structType, structValue,
1634 zero, ArrayRef<int64_t>({i}));
1635 }
1636 innerStructs.push_back(structValue);
1637 }
1638 // Pack the inner structs into a single struct
1639 for (auto [idx, matrix] : llvm::enumerate(innerStructs)) {
1640 packStruct = LLVM::InsertValueOp::create(b, packStruct.getType(),
1641 packStruct, matrix, idx);
1642 }
1643 rewriter.replaceOp(op, packStruct);
1644 return success();
1645 }
1646};
1647
1648struct NVGPUTmaFenceOpLowering
1649 : public ConvertOpToLLVMPattern<nvgpu::TmaFenceOp> {
1650 using ConvertOpToLLVMPattern<nvgpu::TmaFenceOp>::ConvertOpToLLVMPattern;
1651 LogicalResult
1652 matchAndRewrite(nvgpu::TmaFenceOp op, OpAdaptor adaptor,
1653 ConversionPatternRewriter &rewriter) const override {
1654 MLIRContext *ctx = op.getContext();
1655 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1656 auto i32Ty = b.getI32Type();
1657 Value tensormapSize =
1658 LLVM::ConstantOp::create(b, i32Ty, rewriter.getI32IntegerAttr(128));
1659
1660 auto memscope =
1661 NVVM::MemScopeKindAttr::get(ctx, ::mlir::NVVM::MemScopeKind::SYS);
1662
1663 rewriter.replaceOpWithNewOp<NVVM::FenceProxyAcquireOp>(
1664 op, memscope, adaptor.getTensorMapDescriptor(), tensormapSize);
1665
1666 return success();
1667 }
1668};
1669
1670struct NVGPUTmaPrefetchOpLowering
1671 : public ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp> {
1672 using ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp>::ConvertOpToLLVMPattern;
1673 LogicalResult
1674 matchAndRewrite(nvgpu::TmaPrefetchOp op, OpAdaptor adaptor,
1675 ConversionPatternRewriter &rewriter) const override {
1676 rewriter.replaceOpWithNewOp<NVVM::PrefetchOp>(
1677 op, /* CacheLevel */ nullptr, /* Cache Eviction Priority */ nullptr,
1678 adaptor.getTensorMapDescriptor(), adaptor.getPredicate(),
1679 /* Tensormap UnitAttr */ mlir::UnitAttr::get(op.getContext()));
1680 return success();
1681 }
1682};
1683
1684struct NVGPURcpOpLowering : public ConvertOpToLLVMPattern<nvgpu::RcpOp> {
1685 using ConvertOpToLLVMPattern<nvgpu::RcpOp>::ConvertOpToLLVMPattern;
1686 LogicalResult
1687 matchAndRewrite(nvgpu::RcpOp op, OpAdaptor adaptor,
1688 ConversionPatternRewriter &rewriter) const override {
1689 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1690 auto i64Ty = b.getI64Type();
1691 auto f32Ty = b.getF32Type();
1692 VectorType inTy = op.getIn().getType();
1693 // apply rcp.approx.ftz.f on each element in vector.
1694 auto convert1DVec = [&](Type llvm1DVectorTy, Value inVec) {
1695 Value ret1DVec = LLVM::PoisonOp::create(b, llvm1DVectorTy);
1696 int numElems = llvm::cast<VectorType>(llvm1DVectorTy).getNumElements();
1697 for (int i = 0; i < numElems; i++) {
1698 Value idx = LLVM::ConstantOp::create(b, i64Ty, b.getI64IntegerAttr(i));
1699 Value elem = LLVM::ExtractElementOp::create(b, inVec, idx);
1700 Value dst = NVVM::RcpApproxFtzF32Op::create(b, f32Ty, elem);
1701 ret1DVec = LLVM::InsertElementOp::create(b, ret1DVec, dst, idx);
1702 }
1703 return ret1DVec;
1704 };
1705 if (inTy.getRank() == 1) {
1706 rewriter.replaceOp(op, convert1DVec(inTy, adaptor.getIn()));
1707 return success();
1708 }
1710 op.getOperation(), adaptor.getOperands(), *(this->getTypeConverter()),
1711 [&](Type llvm1DVectorTy, ValueRange operands) -> Value {
1712 OpAdaptor adaptor(operands);
1713 return convert1DVec(llvm1DVectorTy, adaptor.getIn());
1714 },
1715 rewriter);
1716 }
1717};
1718
1719//===----------------------------------------------------------------------===//
1720// NVGPUTruncfOp Lowering
1721//===----------------------------------------------------------------------===//
1722
1723enum class FPKind { F32, BF16, F16, F8, F6, F4 };
1724
1725/// Get the effective bit width of a floating-point type.
1726/// f6 types are 6-bit but NVVM Ops expect 8-bit (i8) containers.
1727static int getEffectiveBitWidth(int bitWidth) {
1728 return bitWidth == 6 ? 8 : bitWidth;
1729}
1730
1731static std::optional<FPKind> classifyFPType(Type t) {
1732 static constexpr auto isConvertibleF8Type = [](Type t) {
1733 return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
1734 };
1735 static constexpr auto isConvertibleF6Type = [](Type t) {
1736 return isa<Float6E2M3FNType, Float6E3M2FNType>(t);
1737 };
1738 static constexpr auto isConvertibleF4Type = [](Type t) {
1739 return isa<Float4E2M1FNType>(t);
1740 };
1741
1742 if (t.isF32())
1743 return FPKind::F32;
1744 if (t.isBF16())
1745 return FPKind::BF16;
1746 if (t.isF16())
1747 return FPKind::F16;
1748 if (isConvertibleF8Type(t))
1749 return FPKind::F8;
1750 if (isConvertibleF6Type(t))
1751 return FPKind::F6;
1752 if (isConvertibleF4Type(t))
1753 return FPKind::F4;
1754
1755 return std::nullopt;
1756}
1757
1758/// Conversion op identifier for nvgpu.truncf lowering dispatch table.
1759enum class FPTruncConvOp {
1760 F32x2_TO_F16x2,
1761 F32x2_TO_BF16x2,
1762 F32x2_TO_F8x2,
1763 F32x2_TO_F6x2,
1764 F32x2_TO_F4x2,
1765 F16x2_TO_F8x2,
1766 F16x2_TO_F6x2,
1767 F16x2_TO_F4x2,
1768 BF16x2_TO_F8x2,
1769 BF16x2_TO_F6x2,
1770 BF16x2_TO_F4x2,
1771};
1772
1773struct FPTruncTableEntry {
1774 FPKind src;
1775 FPKind dst;
1776 FPTruncConvOp convOp;
1777};
1778
1779static constexpr FPTruncTableEntry kFPTruncTable[] = {
1780 // f32 source
1781 {FPKind::F32, FPKind::F16, FPTruncConvOp::F32x2_TO_F16x2},
1782 {FPKind::F32, FPKind::BF16, FPTruncConvOp::F32x2_TO_BF16x2},
1783 {FPKind::F32, FPKind::F8, FPTruncConvOp::F32x2_TO_F8x2},
1784 {FPKind::F32, FPKind::F6, FPTruncConvOp::F32x2_TO_F6x2},
1785 {FPKind::F32, FPKind::F4, FPTruncConvOp::F32x2_TO_F4x2},
1786 // f16 source
1787 {FPKind::F16, FPKind::F8, FPTruncConvOp::F16x2_TO_F8x2},
1788 {FPKind::F16, FPKind::F6, FPTruncConvOp::F16x2_TO_F6x2},
1789 {FPKind::F16, FPKind::F4, FPTruncConvOp::F16x2_TO_F4x2},
1790 // bf16 source
1791 {FPKind::BF16, FPKind::F8, FPTruncConvOp::BF16x2_TO_F8x2},
1792 {FPKind::BF16, FPKind::F6, FPTruncConvOp::BF16x2_TO_F6x2},
1793 {FPKind::BF16, FPKind::F4, FPTruncConvOp::BF16x2_TO_F4x2},
1794};
1795
1796/// Find the conversion table entry whose source/destination `FPKind`s match the
1797/// given element types.
1798template <typename TableEntry, size_t N>
1799static std::optional<TableEntry>
1800lookupConvOp(const TableEntry (&table)[N], Type srcElemType, Type dstElemType) {
1801 std::optional<FPKind> srcKind = classifyFPType(srcElemType);
1802 std::optional<FPKind> dstKind = classifyFPType(dstElemType);
1803 if (!srcKind || !dstKind)
1804 return std::nullopt;
1805 for (const TableEntry &entry : table) {
1806 if (entry.src == *srcKind && entry.dst == *dstKind)
1807 return entry;
1808 }
1809 return std::nullopt;
1810}
1811
1812/// Extract a single element from a vector.
1813static Value extractElement(ImplicitLocOpBuilder &b, Value srcVec, int idx) {
1814 assert(idx >= 0 &&
1815 idx < cast<VectorType>(srcVec.getType()).getNumElements() &&
1816 "extractElement: index out of bounds");
1817 IntegerType i64Ty = b.getI64Type();
1818 return b.create<LLVM::ExtractElementOp>(
1819 srcVec, b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(idx)));
1820}
1821
1822/// Extract a pair of f32 values from an i32 vector at the given base index.
1823static std::pair<Value, Value> extractF32Pair(ImplicitLocOpBuilder &b,
1824 Value srcI32Vec, int baseIdx) {
1825 FloatType f32Ty = b.getF32Type();
1826 Value elem0 = extractElement(b, srcI32Vec, baseIdx);
1827 Value elem1 = extractElement(b, srcI32Vec, baseIdx + 1);
1828 return {b.create<LLVM::BitcastOp>(f32Ty, elem0),
1829 b.create<LLVM::BitcastOp>(f32Ty, elem1)};
1830}
1831
1832/// Extract a vector of elements of size i32 from an i32 vector and bitcast to
1833/// the specified vector type.
1834static Value extractAndBitcast(ImplicitLocOpBuilder &b, Value srcI32Vec,
1835 int idx, VectorType vecTy) {
1836 Value elem = extractElement(b, srcI32Vec, idx);
1837 return b.create<LLVM::BitcastOp>(vecTy, elem);
1838}
1839
1840/// Create a sub-byte conversion from an f32 pair source and return the native
1841/// result.
1842template <typename ConvertOp, typename... Args>
1843static Value convertFromF32Pair(ImplicitLocOpBuilder &b, Value srcI32Vec,
1844 int srcBaseIdx, Type resultTy, Args &&...args) {
1845 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1846 return b.create<ConvertOp>(resultTy, hi, lo, std::forward<Args>(args)...);
1847}
1848
1849/// Create a sub-byte conversion from a packed f16x2/bf16x2 source and return
1850/// the native result.
1851template <typename ConvertOp, typename... Args>
1852static Value convertFromPacked(ImplicitLocOpBuilder &b, Value srcI32Vec,
1853 int srcBaseIdx, Type srcElemTy, Type resultTy,
1854 Args &&...args) {
1855 Value src = extractAndBitcast(b, srcI32Vec, srcBaseIdx,
1856 VectorType::get(2, srcElemTy));
1857 return b.create<ConvertOp>(resultTy, src, std::forward<Args>(args)...);
1858}
1859
1860/// Create a typed NVVM truncation conversion.
1861static Value createTruncConversion(
1862 ImplicitLocOpBuilder &b, MLIRContext *ctx, FPTruncConvOp convOp,
1863 Value srcI32Vec, int srcBaseIdx, NVVM::FPRoundingModeAttr rndAttr,
1864 NVVM::SaturationModeAttr satAttr, BoolAttr reluAttr, Type dstElemType,
1865 Type actualDstFloatType, Value randomBits = Value()) {
1866 IntegerType i8Ty = b.getI8Type();
1867 IntegerType i16Ty = b.getI16Type();
1868 IntegerType i32Ty = b.getI32Type();
1869 TypeAttr dstTyAttr = TypeAttr::get(dstElemType);
1870 TypeAttr actualDstTyAttr = TypeAttr::get(actualDstFloatType);
1871
1872 switch (convOp) {
1873 case FPTruncConvOp::F32x2_TO_F16x2: {
1874 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1875 Value r = b.create<NVVM::ConvertF32x2ToF16x2Op>(
1876 VectorType::get(2, b.getF16Type()), hi, lo, randomBits, rndAttr,
1877 satAttr, reluAttr);
1878 return b.create<LLVM::BitcastOp>(i32Ty, r);
1879 }
1880 case FPTruncConvOp::F32x2_TO_BF16x2: {
1881 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1882 Value r = b.create<NVVM::ConvertF32x2ToBF16x2Op>(
1883 VectorType::get(2, b.getBF16Type()), hi, lo, randomBits, rndAttr,
1884 satAttr, reluAttr);
1885 return b.create<LLVM::BitcastOp>(i32Ty, r);
1886 }
1887 case FPTruncConvOp::F32x2_TO_F8x2:
1888 return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
1889 b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
1890 case FPTruncConvOp::F32x2_TO_F6x2:
1891 return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
1892 b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
1893 case FPTruncConvOp::F32x2_TO_F4x2:
1894 return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
1895 b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
1896 case FPTruncConvOp::F16x2_TO_F8x2:
1897 return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
1898 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
1899 case FPTruncConvOp::F16x2_TO_F6x2:
1900 return convertFromPacked<NVVM::ConvertF16x2ToF6x2Op>(
1901 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr,
1902 actualDstTyAttr);
1903 case FPTruncConvOp::F16x2_TO_F4x2:
1904 return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
1905 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i8Ty, reluAttr,
1906 actualDstTyAttr);
1907 case FPTruncConvOp::BF16x2_TO_F8x2:
1908 return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
1909 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, rndAttr, satAttr,
1910 reluAttr, dstTyAttr);
1911 case FPTruncConvOp::BF16x2_TO_F6x2:
1912 return convertFromPacked<NVVM::ConvertBF16x2ToF6x2Op>(
1913 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, reluAttr,
1914 actualDstTyAttr);
1915 case FPTruncConvOp::BF16x2_TO_F4x2:
1916 return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
1917 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i8Ty, reluAttr,
1918 actualDstTyAttr);
1919 }
1920 llvm_unreachable("unhandled FPTruncConvOp");
1921}
1922
1923static LogicalResult lowerTruncf(nvgpu::TruncfOp op,
1924 nvgpu::TruncfOp::Adaptor adaptor,
1925 ConversionPatternRewriter &rewriter,
1926 const LLVMTypeConverter *typeConverter) {
1927 MLIRContext *ctx = op.getContext();
1928 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1929 IntegerType i32Ty = b.getI32Type();
1930 IntegerType i64Ty = b.getI64Type();
1931 static constexpr int regBits = 32;
1932
1933 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
1934 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
1935 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
1936 return rewriter.notifyMatchFailure(
1937 op, "expected 1-D vector; canonicalize pattern handles other shapes");
1938
1939 auto srcElemType = srcType.getElementType();
1940 auto dstElemType = dstType.getElementType();
1941 int srcBW = srcType.getElementTypeBitWidth();
1942 int dstBW = dstType.getElementTypeBitWidth();
1943 int numElems = srcType.getNumElements();
1944
1945 NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
1946 NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
1947 auto reluBoolAttr = op.getReluAttr();
1948 Value randomBits = adaptor.getRandomBits();
1949 Type actualDstFloatType = dstElemType;
1950
1951 // STEP 1: bitcast input vector to i32 vector type.
1952 // f64 -> f32/f16/bf16 lowers to a single direct LLVM fptrunc
1953 // f64 -> f8/f6/f4 first truncates to f32 and then reuses the narrow
1954 // conversion path below.
1955 Value input = adaptor.getIn();
1956 if (srcBW == 64) {
1957 if (dstBW >= 16) {
1958 Type convertedType = typeConverter->convertType(dstType);
1959 assert(convertedType && "failed to convert type");
1960 Value result = b.create<LLVM::FPTruncOp>(convertedType, input);
1961 rewriter.replaceOp(op, result);
1962 return success();
1963 }
1964 auto f32VecTy = VectorType::get(srcType.getShape(), b.getF32Type());
1965 input = b.create<LLVM::FPTruncOp>(f32VecTy, input);
1966 srcType = f32VecTy;
1967 srcElemType = b.getF32Type();
1968 srcBW = 32;
1969 }
1970
1971 // f6 types are 6-bit in MLIR but NVVM uses 8-bit containers.
1972 int effectiveDstBW = getEffectiveBitWidth(dstBW);
1973
1974 int srcI32Elems = numElems * srcBW / regBits;
1975 int dstI32Elems = numElems * effectiveDstBW / regBits;
1976 Value srcI32Vec =
1977 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), input);
1978 Value dstI32Vec =
1979 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
1980
1981 // STEP 2: look up the conversion op from the (srcType, dstType) table.
1982 auto convEntry = lookupConvOp(kFPTruncTable, srcElemType, dstElemType);
1983 if (!convEntry)
1984 return rewriter.notifyMatchFailure(
1985 op, "unsupported type combination for truncation");
1986 FPTruncConvOp convOp = convEntry->convOp;
1987
1988 // Number of source-side i32 register slots consumed by each NVVM convert Op.
1989 auto getNumSrcI32PerConvert = [](FPKind src) {
1990 return src == FPKind::F32 ? 2 : 1;
1991 };
1992 int numSrcI32PerConv = getNumSrcI32PerConvert(convEntry->src);
1993
1994 // STEP 3: pack conversion results into destination i32 vector.
1995 const int srcStep = srcBW / effectiveDstBW;
1996 const int resultBW =
1997 effectiveDstBW * 2; // each conversion produces 2 (packed) elements
1998 const int numConvsPerI32 = regBits / resultBW;
1999
2000 for (int srcIdx = 0, dstIdx = 0; dstIdx < dstI32Elems;
2001 srcIdx += srcStep, dstIdx++) {
2002 Value dstIdxConst =
2003 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2004 Value dstValue;
2005
2006 if (numConvsPerI32 == 1) {
2007 // f16/bf16 destinations
2008 dstValue = createTruncConversion(
2009 b, ctx, convOp, srcI32Vec, srcIdx, rndModeAttr, satModeAttr,
2010 reluBoolAttr, dstElemType, actualDstFloatType, randomBits);
2011 } else {
2012 // f8/f6/f4 destinations: pack sub-results via vector insert + bitcast.
2013 auto subResultType = IntegerType::get(ctx, resultBW);
2014 auto subVecTy = VectorType::get(numConvsPerI32, subResultType);
2015 Value subVec = b.create<LLVM::UndefOp>(subVecTy);
2016
2017 int insertIdx = numConvsPerI32 - 1;
2018 int curStep = srcStep;
2019 while (curStep > 0) {
2020 curStep -= numSrcI32PerConv;
2021 Value subResult = createTruncConversion(
2022 b, ctx, convOp, srcI32Vec, srcIdx + curStep, rndModeAttr,
2023 satModeAttr, reluBoolAttr, dstElemType, actualDstFloatType,
2024 /*randomBits=*/Value());
2025 subVec = b.create<LLVM::InsertElementOp>(
2026 subVec, subResult,
2027 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(insertIdx)));
2028 insertIdx--;
2029 }
2030
2031 dstValue = b.create<LLVM::BitcastOp>(i32Ty, subVec);
2032 }
2033
2034 dstI32Vec =
2035 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2036 }
2037
2038 // STEP 4: produce final result.
2039 Type convertedType = typeConverter->convertType(dstType);
2040 assert(convertedType && "failed to convert type");
2041 if (convEntry->dst == FPKind::F6) {
2042 IntegerType i8Ty = b.getI8Type();
2043 auto i8VecTy = VectorType::get(numElems, i8Ty);
2044 Value i8Vec = b.create<LLVM::BitcastOp>(i8VecTy, dstI32Vec);
2045 Value truncVec = b.create<LLVM::TruncOp>(convertedType, i8Vec);
2046 rewriter.replaceOp(op, truncVec);
2047 } else {
2048 auto dstVec = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2049 rewriter.replaceOp(op, dstVec);
2050 }
2051 return success();
2052}
2053
2054struct NVGPUTruncfOpLowering : public ConvertOpToLLVMPattern<nvgpu::TruncfOp> {
2055 using ConvertOpToLLVMPattern<nvgpu::TruncfOp>::ConvertOpToLLVMPattern;
2056
2057 LogicalResult
2058 matchAndRewrite(nvgpu::TruncfOp op, OpAdaptor adaptor,
2059 ConversionPatternRewriter &rewriter) const override {
2060 return lowerTruncf(op, adaptor, rewriter, getTypeConverter());
2061 }
2062};
2063
2064//===----------------------------------------------------------------------===//
2065// NVGPUExtfOp Lowering
2066//===----------------------------------------------------------------------===//
2067
2068/// Conversion op identifier for nvgpu.extf lowering dispatch table.
2069enum class FPExtConvOp {
2070 F8x2_TO_F16x2,
2071 F8x2_TO_BF16x2,
2072 F6x2_TO_F16x2,
2073 F6x2_TO_BF16x2,
2074 F4x2_TO_F16x2,
2075 F4x2_TO_BF16x2,
2076};
2077
2078struct FPExtTableEntry {
2079 FPKind src;
2080 FPKind dst;
2081 FPExtConvOp convOp;
2082};
2083
2084static constexpr FPExtTableEntry kFPExtTable[] = {
2085 {FPKind::F8, FPKind::F16, FPExtConvOp::F8x2_TO_F16x2},
2086 {FPKind::F8, FPKind::BF16, FPExtConvOp::F8x2_TO_BF16x2},
2087 {FPKind::F6, FPKind::F16, FPExtConvOp::F6x2_TO_F16x2},
2088 {FPKind::F6, FPKind::BF16, FPExtConvOp::F6x2_TO_BF16x2},
2089 {FPKind::F4, FPKind::F16, FPExtConvOp::F4x2_TO_F16x2},
2090 {FPKind::F4, FPKind::BF16, FPExtConvOp::F4x2_TO_BF16x2},
2091};
2092
2093/// Create a typed NVVM extension conversion.
2094/// For f8/f6: src is vector<2xi8>. For f4: src is i8.
2095/// Returns i32 (bitcast from vector<2xf16> or vector<2xbf16>).
2096static Value createExtConversion(ImplicitLocOpBuilder &b, MLIRContext *ctx,
2097 FPExtConvOp convOp, Value src,
2098 BoolAttr reluAttr, Type actualSrcFloatType,
2099 Value extScaleFactor = Value()) {
2100 IntegerType i32Ty = b.getI32Type();
2101 auto srcTyAttr = TypeAttr::get(actualSrcFloatType);
2102
2103 switch (convOp) {
2104 case FPExtConvOp::F8x2_TO_F16x2: {
2105 Value r = NVVM::ConvertF8x2ToF16x2Op::create(
2106 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2107 return b.create<LLVM::BitcastOp>(i32Ty, r);
2108 }
2109 case FPExtConvOp::F8x2_TO_BF16x2: {
2110 Value r = NVVM::ConvertF8x2ToBF16x2Op::create(
2111 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2112 return b.create<LLVM::BitcastOp>(i32Ty, r);
2113 }
2114 case FPExtConvOp::F6x2_TO_F16x2: {
2115 Value r = NVVM::ConvertF6x2ToF16x2Op::create(
2116 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2117 return b.create<LLVM::BitcastOp>(i32Ty, r);
2118 }
2119 case FPExtConvOp::F6x2_TO_BF16x2: {
2120 Value r = NVVM::ConvertF6x2ToBF16x2Op::create(
2121 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2122 return b.create<LLVM::BitcastOp>(i32Ty, r);
2123 }
2124 case FPExtConvOp::F4x2_TO_F16x2: {
2125 Value r = NVVM::ConvertF4x2ToF16x2Op::create(
2126 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2127 return b.create<LLVM::BitcastOp>(i32Ty, r);
2128 }
2129 case FPExtConvOp::F4x2_TO_BF16x2: {
2130 Value r = NVVM::ConvertF4x2ToBF16x2Op::create(
2131 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2132 return b.create<LLVM::BitcastOp>(i32Ty, r);
2133 }
2134 }
2135 llvm_unreachable("unhandled FPExtConvOp");
2136}
2137
2138static LogicalResult lowerExtf(nvgpu::ExtfOp op, nvgpu::ExtfOp::Adaptor adaptor,
2139 ConversionPatternRewriter &rewriter,
2140 const LLVMTypeConverter *typeConverter) {
2141 MLIRContext *ctx = op.getContext();
2142 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
2143 IntegerType i8Ty = b.getI8Type();
2144 IntegerType i16Ty = b.getI16Type();
2145 IntegerType i32Ty = b.getI32Type();
2146 IntegerType i64Ty = b.getI64Type();
2147
2148 static constexpr int regBits = 32;
2149 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
2150 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
2151 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
2152 return rewriter.notifyMatchFailure(
2153 op, "expected 1-D vector; canonicalize pattern handles other shapes");
2154
2155 auto srcElemType = srcType.getElementType();
2156 auto dstElemType = dstType.getElementType();
2157 int srcBW = srcType.getElementTypeBitWidth();
2158 int dstBW = dstType.getElementTypeBitWidth();
2159 int numElems = srcType.getNumElements();
2160
2161 auto reluBoolAttr = op.getReluAttr();
2162 Type actualSrcFloatType = srcElemType;
2163
2164 assert(dstBW == 16 || dstBW == 32 || dstBW == 64);
2165
2166 // Wide source (f16/bf16/f32) to wide destination (f32/f64): single FPExt.
2167 if (srcBW >= 16 && dstBW >= 32) {
2168 Value result = adaptor.getIn();
2169 if (srcElemType != dstElemType) {
2170 Type convertedType = typeConverter->convertType(dstType);
2171 assert(convertedType && "failed to convert type");
2172 result = b.create<LLVM::FPExtOp>(convertedType, result);
2173 }
2174 rewriter.replaceOp(op, result);
2175 return success();
2176 }
2177
2178 // Narrow source (f8/f6/f4): NVVM typed op produces f16/bf16; optionally
2179 // followed by FPExt to the final f32/f64 destination.
2180 bool needsFinalFPExt = (dstBW >= 32);
2181 Type intermediateDstElem = dstElemType;
2182 if (needsFinalFPExt && llvm::isa<Float8E8M0FNUType>(srcElemType))
2183 intermediateDstElem = b.getBF16Type();
2184 else if (needsFinalFPExt)
2185 intermediateDstElem = b.getF16Type();
2186 int intermediateDstBW = needsFinalFPExt ? 16 : dstBW;
2187
2188 // f6 types are 6-bit in MLIR but NVVM uses 8-bit containers.
2189 int effectiveSrcBW = getEffectiveBitWidth(srcBW);
2190
2191 // STEP 1: prepare input as i32 register vector.
2192 // For f6: zext from vector<Nxi6> to vector<Nxi8>, then bitcast to i32s.
2193 Value inputVec = adaptor.getIn();
2194 if (srcBW == 6) {
2195 auto i8VecTy = VectorType::get(numElems, i8Ty);
2196 inputVec = b.create<LLVM::ZExtOp>(i8VecTy, inputVec);
2197 }
2198
2199 int srcI32Elems = numElems * effectiveSrcBW / regBits;
2200 int dstI32Elems = numElems * intermediateDstBW / regBits;
2201 Value srcI32Vec =
2202 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), inputVec);
2203 Value dstI32Vec =
2204 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
2205
2206 // STEP 2: look up the conversion op from the (srcType, dstType) table.
2207 auto convEntry = lookupConvOp(kFPExtTable, srcElemType, intermediateDstElem);
2208 if (!convEntry)
2209 return rewriter.notifyMatchFailure(
2210 op, "unsupported type combination for extension");
2211 FPExtConvOp convOp = convEntry->convOp;
2212 Value extScaleFactor;
2213
2214 // STEP 3: iterate over source i32 elements, producing destination i32s.
2215 for (int srcIdx = 0, dstIdx = 0; srcIdx < srcI32Elems; srcIdx++) {
2216 Value srcI32 = b.create<LLVM::ExtractElementOp>(
2217 srcI32Vec,
2218 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(srcIdx)));
2219
2220 if (effectiveSrcBW == 8) {
2221 // f8/f6: one i32 holds 4 bytes -> split into 2 pairs of i16 -> 2 convs.
2222 Value i16Vec =
2223 b.create<LLVM::BitcastOp>(VectorType::get(2, i16Ty), srcI32);
2224 for (int half = 0; half < 2; half++) {
2225 Value halfI16 = b.create<LLVM::ExtractElementOp>(
2226 i16Vec,
2227 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(half)));
2228 Value src =
2229 b.create<LLVM::BitcastOp>(VectorType::get(2, i8Ty), halfI16);
2230 Value dstValue =
2231 createExtConversion(b, ctx, convOp, src, reluBoolAttr,
2232 actualSrcFloatType, extScaleFactor);
2233 Value dstIdxConst =
2234 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2235 dstI32Vec =
2236 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2237 dstIdx++;
2238 }
2239 } else {
2240 // f4: one i32 holds 4 bytes -> each byte is one conversion input.
2241 Value i8Vec = b.create<LLVM::BitcastOp>(VectorType::get(4, i8Ty), srcI32);
2242 for (int byteIdx = 0; byteIdx < 4; byteIdx++) {
2243 Value src = b.create<LLVM::ExtractElementOp>(
2244 i8Vec,
2245 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(byteIdx)));
2246 Value dstValue =
2247 createExtConversion(b, ctx, convOp, src, reluBoolAttr,
2248 actualSrcFloatType, extScaleFactor);
2249 Value dstIdxConst =
2250 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2251 dstI32Vec =
2252 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2253 dstIdx++;
2254 }
2255 }
2256 }
2257
2258 // STEP 4: produce final result.
2259 Type convertedType = typeConverter->convertType(dstType);
2260 assert(convertedType && "failed to convert type");
2261 Value result;
2262 if (needsFinalFPExt) {
2263 auto intermediateVecTy = VectorType::get(numElems, intermediateDstElem);
2264 Value intermediateVec =
2265 b.create<LLVM::BitcastOp>(intermediateVecTy, dstI32Vec);
2266 result = b.create<LLVM::FPExtOp>(convertedType, intermediateVec);
2267 } else {
2268 result = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2269 }
2270 rewriter.replaceOp(op, result);
2271 return success();
2272}
2273
2274struct NVGPUExtfOpLowering : public ConvertOpToLLVMPattern<nvgpu::ExtfOp> {
2275 using ConvertOpToLLVMPattern<nvgpu::ExtfOp>::ConvertOpToLLVMPattern;
2276
2277 LogicalResult
2278 matchAndRewrite(nvgpu::ExtfOp op, OpAdaptor adaptor,
2279 ConversionPatternRewriter &rewriter) const override {
2280 return lowerExtf(op, adaptor, rewriter, getTypeConverter());
2281 }
2282};
2283
2284static int64_t computePaddedElems(int64_t numElems, int srcBW, int dstBW,
2285 int step) {
2286 static constexpr int regBits = 32;
2287 int effSrcBW = getEffectiveBitWidth(srcBW);
2288 int effDstBW = getEffectiveBitWidth(dstBW);
2289 auto ceilDiv = [](int64_t x, int64_t y) { return (x + y - 1) / y; };
2290 int64_t padded =
2291 std::max(ceilDiv(numElems * effSrcBW, regBits) * regBits / effSrcBW,
2292 ceilDiv(numElems * effDstBW, regBits) * regBits / effDstBW);
2293 return ceilDiv(padded, step) * step;
2294}
2295
2296/// Canonicalization pattern for nvgpu.truncf / nvgpu.extf:
2297/// handles scalar inputs, non-32-bit-aligned vectors, and multi-rank vectors.
2298/// Runs as an OpRewritePattern on MLIR types before LLVM type conversion.
2299template <typename CvtOp, bool IsTrunc>
2300struct NVGPUFPCanonicalizePattern : public OpRewritePattern<CvtOp> {
2301 using OpRewritePattern<CvtOp>::OpRewritePattern;
2302
2303 LogicalResult matchAndRewrite(CvtOp op,
2304 PatternRewriter &rewriter) const override {
2305 Type inType = op.getIn().getType();
2306 Type outType = op.getOut().getType();
2307
2308 Type srcElemTy = getElementTypeOrSelf(inType);
2309 Type dstElemTy = getElementTypeOrSelf(outType);
2310 int srcBW = srcElemTy.getIntOrFloatBitWidth();
2311 int dstBW = dstElemTy.getIntOrFloatBitWidth();
2312 int effSrcBW = getEffectiveBitWidth(srcBW);
2313 int effDstBW = getEffectiveBitWidth(dstBW);
2314
2315 bool isScalar = !isa<VectorType>(inType);
2316 auto srcVecTy = dyn_cast<VectorType>(inType);
2317 bool isMultiRank = srcVecTy && srcVecTy.getRank() > 1;
2318 int64_t numElems = isScalar ? 1 : srcVecTy.getNumElements();
2319 int step = IsTrunc ? effSrcBW / effDstBW : effDstBW / effSrcBW;
2320 int64_t paddedElems = computePaddedElems(numElems, srcBW, dstBW, step);
2321 bool needsPad = (paddedElems != numElems);
2322
2323 if (!isScalar && !isMultiRank && !needsPad)
2324 return failure();
2325
2326 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
2327 Value input = op.getIn();
2328
2329 if (isScalar)
2330 input = vector::BroadcastOp::create(b, VectorType::get({1}, srcElemTy),
2331 input);
2332 if (isMultiRank)
2333 input = vector::ShapeCastOp::create(
2334 b, VectorType::get({numElems}, srcElemTy), input);
2335
2336 if (needsPad) {
2337 auto paddedTy = VectorType::get({paddedElems}, srcElemTy);
2338 Value zero = arith::ConstantOp::create(
2339 b, DenseElementsAttr::get(paddedTy, b.getZeroAttr(srcElemTy)));
2340 input = vector::InsertStridedSliceOp::create(
2341 b, input, zero, SmallVector<int64_t>{0}, SmallVector<int64_t>{1});
2342 }
2343
2344 auto cvtDstTy =
2345 VectorType::get({needsPad ? paddedElems : numElems}, dstElemTy);
2346 Value cvt;
2347 if constexpr (IsTrunc) {
2348 cvt = CvtOp::create(b, cvtDstTy, input, op.getRndAttr(), op.getSatAttr(),
2349 op.getReluAttr(), op.getRandomBits());
2350 } else {
2351 cvt =
2352 CvtOp::create(b, cvtDstTy, input, op.getRndAttr(), op.getReluAttr());
2353 }
2354 Value result = cvt;
2355
2356 if (needsPad) {
2357 result = vector::ExtractStridedSliceOp::create(
2358 b, result, SmallVector<int64_t>{0}, SmallVector<int64_t>{numElems},
2359 SmallVector<int64_t>{1});
2360 }
2361
2362 if (isMultiRank) {
2363 result =
2364 vector::ShapeCastOp::create(b, cast<VectorType>(outType), result);
2365 }
2366
2367 if (isScalar) {
2368 result = vector::ExtractOp::create(b, result, SmallVector<int64_t>{0});
2369 }
2370
2371 rewriter.replaceOp(op, result);
2372 return success();
2373 }
2374};
2375
2376using NVGPUTruncfCanonicalizePattern =
2377 NVGPUFPCanonicalizePattern<nvgpu::TruncfOp, true>;
2378using NVGPUExtfCanonicalizePattern =
2379 NVGPUFPCanonicalizePattern<nvgpu::ExtfOp, false>;
2380} // namespace
2381
2383 TypeConverter &typeConverter) {
2384 // NVVM uses alloca in the default address space to represent private
2385 // memory allocations, so drop private annotations. NVVM uses address
2386 // space 3 for shared memory. NVVM uses the default address space to
2387 // represent global memory.
2389 typeConverter, [](gpu::AddressSpace space) -> unsigned {
2390 switch (space) {
2391 case gpu::AddressSpace::Global:
2392 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Global);
2393 case gpu::AddressSpace::Workgroup:
2394 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared);
2395 case gpu::AddressSpace::Private:
2396 return 0;
2397 case gpu::AddressSpace::Constant:
2398 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Constant);
2399 }
2400 llvm_unreachable("unknown address space enum value");
2401 });
2402}
2403
2405 const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
2406 patterns.add<
2407 NVGPUMBarrierCreateLowering, // nvgpu.mbarrier.create
2408 NVGPUMBarrierInitLowering, // nvgpu.mbarrier.init
2409 NVGPUMBarrierGetLowering, // nvgpu.mbarrier.get
2410 NVGPUMBarrierArriveLowering, // nvgpu.mbarrier.arrive
2411 NVGPUMBarrierArriveNoCompleteLowering, // nvgpu.mbarrier.arrive.no_complete
2412 NVGPUMBarrierTestWaitLowering, // nvgpu.mbarrier.test_wait_parity
2413 NVGPUMBarrierTryWaitParityLowering, // nvgpu.mbarrier.try_wait_parity
2414 NVGPUTmaAsyncLoadOpLowering, // nvgpu.tma.async.load
2415 NVGPUTmaAsyncStoreOpLowering, // nvgpu.tma.async.store
2416 NVGPUTmaCreateDescriptorOpLowering, // nvgpu.tma.create.descriptor
2417 NVGPUTmaPrefetchOpLowering, // nvgpu.tma.prefetch.descriptor
2418 NVGPUTmaFenceOpLowering, // nvgpu.tma.fence.descriptor
2419 NVGPUMBarrierArriveExpectTxLowering, // nvgpu.mbarrier.arrive.expect_tx
2420 NVGPUGenerateWarpgroupDescriptorLowering, // nvgpu.warpgroup.generate.descriptor
2421 NVGPUWarpgroupMmaOpLowering, // nvgpu.warpgroup.mma
2422 NVGPUWarpgroupMmaStoreOpLowering, // nvgpu.warpgroup.mma.store
2423 NVGPUWarpgroupMmaInitAccumulatorOpLowering, // nvgpu.warpgroup.mma.init.accumulator
2424 NVGPUTruncfOpLowering, // nvgpu.truncf
2425 NVGPUExtfOpLowering, // nvgpu.extf
2426 MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM, NVGPUAsyncCopyLowering,
2427 NVGPUAsyncCreateGroupLowering, NVGPUAsyncWaitLowering,
2428 NVGPUMmaSparseSyncLowering, NVGPURcpOpLowering>(converter);
2429
2430 patterns.add<NVGPUTruncfCanonicalizePattern, NVGPUExtfCanonicalizePattern>(
2431 patterns.getContext());
2432}
return success()
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
constexpr int kWgmmaSizeM
M size of wgmma.mma_async instruction.
constexpr int kWarpSize
static Value truncToI32(ImplicitLocOpBuilder &b, Value value)
GPU has 32 bit registers, this function truncates values when larger width is not needed.
static SmallVector< Value > unpackOperandVector(ImplicitLocOpBuilder &b, Value operand, NVVM::MMATypes operandPtxType)
The gpu.mma.sync converter below expects matrix fragment operands to be given as 2D vectors where the...
static Type inferIntrinsicResultType(Type vectorResultType)
Returns the type for the intrinsic given the vectorResultType of the gpu.mma.sync operation.
constexpr int exclude4LSB
Number of bits that needs to be excluded when building matrix descriptor for wgmma operations.
static bool isMbarrierShared(nvgpu::MBarrierGroupType barrierType)
Returns whether mbarrier object has shared memory address space.
static Value convertIntrinsicResult(Location loc, Type intrinsicResultType, Type resultType, Value intrinsicResult, RewriterBase &rewriter)
Convert the SSA result of the NVVM intrinsic nvvm.mma.sync (which is always an LLVM struct) into a fr...
static llvm::ManagedStatic< PassManagerOptions > options
Attributes are known-constant values of operations.
Definition Attributes.h:25
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:204
FloatType getF32Type()
Definition Builders.cpp:47
IntegerType getI32Type()
Definition Builders.cpp:67
FloatType getF16Type()
Definition Builders.cpp:43
MLIRContext * getContext() const
Definition Builders.h:56
FloatType getF64Type()
Definition Builders.cpp:49
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:227
Value getStridedElementPtr(ConversionPatternRewriter &rewriter, Location loc, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none) const
Convenience wrapper for the corresponding helper utility.
Definition Pattern.cpp:66
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Definition Builders.h:632
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:528
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isF64() const
Definition Types.cpp:41
bool isTF32() const
Definition Types.cpp:39
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isF8E5M2() const
Definition Types.cpp:45
bool isSignlessInteger() const
Return true if this is a signless integer type (with the specified width).
Definition Types.cpp:66
bool isF8E4M3FN() const
Definition Types.cpp:44
bool isF32() const
Definition Types.cpp:40
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
Definition Types.cpp:90
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isF16() const
Definition Types.cpp:38
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isBF16() const
Definition Types.cpp:37
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
LogicalResult handleMultidimensionalVectors(Operation *op, ValueRange operands, const LLVMTypeConverter &typeConverter, std::function< Value(Type, ValueRange)> createOperand, ConversionPatternRewriter &rewriter)
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
Definition Pattern.cpp:603
void populateCommonGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap common GPU memory spaces (Workgroup, Private, etc) to LLVM address spaces.
MemRefType getMBarrierMemrefType(MLIRContext *context, MBarrierGroupType barrierType)
Return the memref type that can be used to represent an mbarrier object.
Attribute getMbarrierMemorySpace(MLIRContext *context, MBarrierGroupType barrierType)
Returns the memory space attribute of the mbarrier object.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
void populateNVGPUToNVVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
void populateGpuMemorySpaceAttributeConversions(TypeConverter &typeConverter, const MemorySpaceMapping &mapping)
Populates memory space attribute conversion rules for lowering gpu.address_space to integer values.
LLVM::CallOp create(Location loc, OpBuilder &builder, ArrayRef< Value > arguments) const
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...