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.getTf32Enabled().value_or(false);
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 /*convergent=*/false,
570 /*asm_dialect=*/asmDialectAttr,
571 /*operand_attrs=*/ArrayAttr());
572}
573
574/// Lowers `nvgpu.mma.sp.sync` to inline assembly.
575struct NVGPUMmaSparseSyncLowering
576 : public ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp> {
577 using ConvertOpToLLVMPattern<nvgpu::MmaSparseSyncOp>::ConvertOpToLLVMPattern;
578
579 LogicalResult
580 matchAndRewrite(nvgpu::MmaSparseSyncOp op, OpAdaptor adaptor,
581 ConversionPatternRewriter &rewriter) const override {
582 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
583 // Get the shapes of the MMAMatrix type being used. The shapes will
584 // choose which intrinsic this op will be lowered to.
585 VectorType aType = op.getMatrixA().getType();
586 VectorType bType = op.getMatrixB().getType();
587 VectorType cType = op.getMatrixC().getType();
588
589 FailureOr<NVVM::MMATypes> ptxTypeA = getNvvmMmaType(aType);
590 if (failed(ptxTypeA))
591 return op->emitOpError("failed to deduce operand PTX types");
592 FailureOr<NVVM::MMATypes> ptxTypeB = getNvvmMmaType(bType);
593 if (failed(ptxTypeB))
594 return op->emitOpError("failed to deduce operand PTX types");
595 std::optional<NVVM::MMATypes> ptxTypeC =
596 NVVM::MmaOp::inferOperandMMAType(cType.getElementType(),
597 /*isAccumulator=*/true);
598 if (!ptxTypeC)
599 return op->emitError(
600 "could not infer the PTX type for the accumulator/result");
601
602 // Same as `mma.sync`, F32 works only with TensorFloat32 (TF32).
603 bool tf32Enabled = op.getTf32Enabled().value_or(false);
604 if (aType.getElementType().isF32() && !tf32Enabled)
605 return failure();
606
607 // TODO: add an attribute to the op to customize this behavior.
608 std::optional<NVVM::MMAIntOverflow> overflow(std::nullopt);
609 if (isa<IntegerType>(aType.getElementType()))
610 overflow = NVVM::MMAIntOverflow::satfinite;
611
612 SmallVector<Value> matA =
613 unpackOperandVector(b, adaptor.getMatrixA(), *ptxTypeA);
614 SmallVector<Value> matB =
615 unpackOperandVector(b, adaptor.getMatrixB(), *ptxTypeB);
616 SmallVector<Value> matC =
617 unpackOperandVector(b, adaptor.getMatrixC(), *ptxTypeC);
618
619 Type desiredRetTy = typeConverter->convertType(op->getResultTypes()[0]);
620 Type intrinsicResTy = inferIntrinsicResultType(
621 typeConverter->convertType(op->getResultTypes()[0]));
622
623 // Bitcast the sparse metadata from vector<2xf16> to an i32.
624 Value sparseMetadata = adaptor.getSparseMetadata();
625 if (sparseMetadata.getType() != VectorType::get(2, rewriter.getI16Type()))
626 return op->emitOpError() << "Expected metadata type to be LLVM "
627 "VectorType of 2 i16 elements";
628 sparseMetadata =
629 LLVM::BitcastOp::create(b, rewriter.getI32Type(), sparseMetadata);
630
631 FailureOr<LLVM::InlineAsmOp> intrinsicResult = emitMmaSparseSyncOpAsm(
632 b, *ptxTypeA, *ptxTypeB, *ptxTypeC, *ptxTypeC, overflow, matA, matB,
633 matC, sparseMetadata, op.getSparsitySelector(), op.getMmaShapeAsArray(),
634 intrinsicResTy);
635 if (failed(intrinsicResult))
636 return failure();
637
638 assert((*intrinsicResult).getNumResults() == 1 &&
639 "expected inline asm op returns a single LLVM struct type");
640 rewriter.replaceOp(
641 op, convertIntrinsicResult(op.getLoc(), intrinsicResTy, desiredRetTy,
642 (*intrinsicResult)->getResult(0), rewriter));
643 return success();
644 }
645};
646
647struct NVGPUAsyncCopyLowering
648 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncCopyOp> {
649 using ConvertOpToLLVMPattern<
650 nvgpu::DeviceAsyncCopyOp>::ConvertOpToLLVMPattern;
651
652 LogicalResult
653 matchAndRewrite(nvgpu::DeviceAsyncCopyOp op, OpAdaptor adaptor,
654 ConversionPatternRewriter &rewriter) const override {
655 ImplicitLocOpBuilder b(op.getLoc(), rewriter);
656 Location loc = op.getLoc();
657 auto dstMemrefType = cast<MemRefType>(op.getDst().getType());
658 Value dstPtr =
659 getStridedElementPtr(rewriter, b.getLoc(), dstMemrefType,
660 adaptor.getDst(), adaptor.getDstIndices());
661 FailureOr<unsigned> dstAddressSpace =
662 getTypeConverter()->getMemRefAddressSpace(dstMemrefType);
663 if (failed(dstAddressSpace))
664 return rewriter.notifyMatchFailure(
665 loc, "destination memref address space not convertible to integer");
666
667 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
668 FailureOr<unsigned> srcAddressSpace =
669 getTypeConverter()->getMemRefAddressSpace(srcMemrefType);
670 if (failed(srcAddressSpace))
671 return rewriter.notifyMatchFailure(
672 loc, "source memref address space not convertible to integer");
673
674 Value scrPtr =
675 getStridedElementPtr(rewriter, loc, srcMemrefType, adaptor.getSrc(),
676 adaptor.getSrcIndices());
677 // Intrinsics takes a global pointer so we need an address space cast.
678 auto srcPointerGlobalType = LLVM::LLVMPointerType::get(
679 op->getContext(), static_cast<unsigned>(NVVM::NVVMMemorySpace::Global));
680 scrPtr = LLVM::AddrSpaceCastOp::create(b, srcPointerGlobalType, scrPtr);
681 int64_t dstElements = adaptor.getDstElements().getZExtValue();
682 int64_t sizeInBytes =
683 (dstMemrefType.getElementTypeBitWidth() * dstElements) / 8;
684 // When the optional SrcElements argument is *not* present, the regular
685 // CpAsyncOp is generated. CopyAsyncOp reads bytes from source (global
686 // memory) to fill DstElements number of elements in the destination
687 // (shared memory).
688 Value srcBytes = adaptor.getSrcElements();
689 if (srcBytes) {
690 // When the optional SrcElements argument is present, the source (global
691 // memory) of CpAsyncOp is read only for SrcElements number of elements.
692 // The rest of the DstElements in the destination (shared memory) are
693 // filled with zeros.
694 Value c3I32 =
695 LLVM::ConstantOp::create(b, b.getI32Type(), b.getI32IntegerAttr(3));
696 Value bitwidth = LLVM::ConstantOp::create(
697 b, b.getI32Type(),
698 b.getI32IntegerAttr(srcMemrefType.getElementTypeBitWidth()));
699 Value srcElementsI32 = LLVM::TruncOp::create(b, b.getI32Type(), srcBytes);
700 srcBytes = LLVM::LShrOp::create(
701 b, LLVM::MulOp::create(b, bitwidth, srcElementsI32), c3I32);
702 }
703 // Cache global (.cg) for 16 dst bytes, Cache all (.ca) for sizes other than
704 // 16 dst bytes.
705 NVVM::LoadCacheModifierKind cacheModifier =
706 (op.getBypassL1().value_or(false) && sizeInBytes == 16)
707 ? NVVM::LoadCacheModifierKind::CG
708 : NVVM::LoadCacheModifierKind::CA;
709
710 NVVM::CpAsyncOp::create(
711 b, dstPtr, scrPtr, rewriter.getI32IntegerAttr(sizeInBytes),
712 NVVM::LoadCacheModifierKindAttr::get(op->getContext(), cacheModifier),
713 srcBytes);
714
715 // Drop the result token.
716 Value zero =
717 LLVM::ConstantOp::create(b, IntegerType::get(op.getContext(), 32),
718 rewriter.getI32IntegerAttr(0));
719 rewriter.replaceOp(op, zero);
720 return success();
721 }
722};
723
724struct NVGPUAsyncCreateGroupLowering
725 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncCreateGroupOp> {
726 using ConvertOpToLLVMPattern<
727 nvgpu::DeviceAsyncCreateGroupOp>::ConvertOpToLLVMPattern;
728
729 LogicalResult
730 matchAndRewrite(nvgpu::DeviceAsyncCreateGroupOp op, OpAdaptor adaptor,
731 ConversionPatternRewriter &rewriter) const override {
732 NVVM::CpAsyncCommitGroupOp::create(rewriter, op.getLoc());
733 // Drop the result token.
734 Value zero = LLVM::ConstantOp::create(rewriter, op->getLoc(),
735 IntegerType::get(op.getContext(), 32),
736 rewriter.getI32IntegerAttr(0));
737 rewriter.replaceOp(op, zero);
738 return success();
739 }
740};
741
742struct NVGPUAsyncWaitLowering
743 : public ConvertOpToLLVMPattern<nvgpu::DeviceAsyncWaitOp> {
744 using ConvertOpToLLVMPattern<
745 nvgpu::DeviceAsyncWaitOp>::ConvertOpToLLVMPattern;
746
747 LogicalResult
748 matchAndRewrite(nvgpu::DeviceAsyncWaitOp op, OpAdaptor adaptor,
749 ConversionPatternRewriter &rewriter) const override {
750 // If numGroup is not present pick 0 as a conservative correct value.
751 int32_t numGroups = adaptor.getNumGroups().value_or(0);
752 NVVM::CpAsyncWaitGroupOp::create(rewriter, op.getLoc(), numGroups);
753 rewriter.eraseOp(op);
754 return success();
755 }
756};
757
758/// Creates mbarrier object in shared memory
759struct NVGPUMBarrierCreateLowering
760 : public ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp> {
761 using ConvertOpToLLVMPattern<nvgpu::MBarrierCreateOp>::ConvertOpToLLVMPattern;
762
763 template <typename moduleT>
764 memref::GlobalOp generateGlobalBarrier(ConversionPatternRewriter &rewriter,
765 Operation *funcOp, moduleT moduleOp,
766 MemRefType barrierType) const {
767 SymbolTable symbolTable(moduleOp);
768 OpBuilder::InsertionGuard guard(rewriter);
769 rewriter.setInsertionPoint(&moduleOp.front());
770 auto global = memref::GlobalOp::create(
771 rewriter, funcOp->getLoc(), "__mbarrier",
772 /*sym_visibility=*/rewriter.getStringAttr("private"),
773 /*type=*/barrierType,
774 /*initial_value=*/ElementsAttr(),
775 /*constant=*/false,
776 /*alignment=*/rewriter.getI64IntegerAttr(8));
777 symbolTable.insert(global);
778 return global;
779 }
780
781 LogicalResult
782 matchAndRewrite(nvgpu::MBarrierCreateOp op, OpAdaptor adaptor,
783 ConversionPatternRewriter &rewriter) const override {
784 Operation *funcOp = op->getParentOp();
785 MemRefType barrierType = nvgpu::getMBarrierMemrefType(
786 rewriter.getContext(), op.getBarriers().getType());
787
788 memref::GlobalOp global;
789 if (auto moduleOp = funcOp->getParentOfType<gpu::GPUModuleOp>())
790 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
791 else if (auto moduleOp = funcOp->getParentOfType<ModuleOp>())
792 global = generateGlobalBarrier(rewriter, funcOp, moduleOp, barrierType);
793
794 rewriter.setInsertionPoint(op);
795 rewriter.replaceOpWithNewOp<memref::GetGlobalOp>(op, barrierType,
796 global.getName());
797 return success();
798 }
799};
800
801/// Base class for lowering mbarrier operations to nvvm intrinsics.
802template <typename SourceOp>
803struct MBarrierBasePattern : public ConvertOpToLLVMPattern<SourceOp> {
804public:
805 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
806 /// Returns the base pointer of the mbarrier object.
807 Value getMbarrierPtr(ImplicitLocOpBuilder &b,
808 nvgpu::MBarrierGroupType mbarType, Value memrefDesc,
809 Value mbarId,
810 ConversionPatternRewriter &rewriter) const {
811 MemRefType mbarrierMemrefType =
812 nvgpu::getMBarrierMemrefType(rewriter.getContext(), mbarType);
814 rewriter, b.getLoc(), mbarrierMemrefType, memrefDesc, {mbarId});
815 }
816};
817
818struct NVGPUMBarrierGetLowering
819 : public MBarrierBasePattern<nvgpu::MBarrierGetOp> {
820 using MBarrierBasePattern<nvgpu::MBarrierGetOp>::MBarrierBasePattern;
821
822 LogicalResult
823 matchAndRewrite(nvgpu::MBarrierGetOp op, OpAdaptor adaptor,
824 ConversionPatternRewriter &rewriter) const override {
825 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
826 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
827 rewriter.setInsertionPoint(op);
828 Value barrier = getMbarrierPtr(b, mbarrierType, adaptor.getBarriers(),
829 adaptor.getMbarId(), rewriter);
830 Type resType = op.getMbarrierPointer().getType();
831 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(op, resType, barrier);
832 return success();
833 }
834};
835
836/// Lowers `nvgpu.mbarrier.init` to `nvvm.mbarrier.init`
837struct NVGPUMBarrierInitLowering
838 : public MBarrierBasePattern<nvgpu::MBarrierInitOp> {
839 using MBarrierBasePattern<nvgpu::MBarrierInitOp>::MBarrierBasePattern;
840
841 LogicalResult
842 matchAndRewrite(nvgpu::MBarrierInitOp op, OpAdaptor adaptor,
843 ConversionPatternRewriter &rewriter) const override {
844 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
845 nvgpu::MBarrierGroupType mbarrierType = op.getBarriers().getType();
846 rewriter.setInsertionPoint(op);
847 Value barrier = getMbarrierPtr(b, mbarrierType, adaptor.getBarriers(),
848 adaptor.getMbarId(), rewriter);
849 Value count = truncToI32(b, adaptor.getCount());
850 rewriter.replaceOpWithNewOp<NVVM::MBarrierInitOp>(op, barrier, count, 0,
851 adaptor.getPredicate());
852 return success();
853 }
854};
855
856/// Lowers `nvgpu.mbarrier.arrive` to `nvvm.mbarrier.arrive`
857struct NVGPUMBarrierArriveLowering
858 : public MBarrierBasePattern<nvgpu::MBarrierArriveOp> {
859 using MBarrierBasePattern<nvgpu::MBarrierArriveOp>::MBarrierBasePattern;
860 LogicalResult
861 matchAndRewrite(nvgpu::MBarrierArriveOp op, OpAdaptor adaptor,
862 ConversionPatternRewriter &rewriter) const override {
863 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
864 Value barrier =
865 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
866 adaptor.getMbarId(), rewriter);
867 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveOp>(op, barrier, Value{});
868 return success();
869 }
870};
871
872/// Lowers `nvgpu.mbarrier.arrive.nocomplete` to
873/// `nvvm.mbarrier.arrive.nocomplete`
874struct NVGPUMBarrierArriveNoCompleteLowering
875 : public MBarrierBasePattern<nvgpu::MBarrierArriveNoCompleteOp> {
876 using MBarrierBasePattern<
877 nvgpu::MBarrierArriveNoCompleteOp>::MBarrierBasePattern;
878 LogicalResult
879 matchAndRewrite(nvgpu::MBarrierArriveNoCompleteOp op, OpAdaptor adaptor,
880 ConversionPatternRewriter &rewriter) const override {
881 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
882 Value barrier =
883 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
884 adaptor.getMbarId(), rewriter);
885 Type tokenType = getTypeConverter()->convertType(
886 nvgpu::MBarrierTokenType::get(op->getContext()));
887 Value count = truncToI32(b, adaptor.getCount());
888 rewriter.replaceOpWithNewOp<NVVM::MBarrierArriveNocompleteOp>(
889 op, tokenType, barrier, count);
890 return success();
891 }
892};
893
894/// Lowers `nvgpu.mbarrier.test.wait` to `nvvm.mbarrier.test.wait`
895struct NVGPUMBarrierTestWaitLowering
896 : public MBarrierBasePattern<nvgpu::MBarrierTestWaitOp> {
897 using MBarrierBasePattern<nvgpu::MBarrierTestWaitOp>::MBarrierBasePattern;
898 LogicalResult
899 matchAndRewrite(nvgpu::MBarrierTestWaitOp op, OpAdaptor adaptor,
900 ConversionPatternRewriter &rewriter) const override {
901 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
902 Value barrier =
903 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
904 adaptor.getMbarId(), rewriter);
905 Type retType = rewriter.getI1Type();
906 rewriter.replaceOpWithNewOp<NVVM::MBarrierTestWaitOp>(op, retType, barrier,
907 adaptor.getToken());
908 return success();
909 }
910};
911
912struct NVGPUMBarrierArriveExpectTxLowering
913 : public MBarrierBasePattern<nvgpu::MBarrierArriveExpectTxOp> {
914 using MBarrierBasePattern<
915 nvgpu::MBarrierArriveExpectTxOp>::MBarrierBasePattern;
916 LogicalResult
917 matchAndRewrite(nvgpu::MBarrierArriveExpectTxOp op, OpAdaptor adaptor,
918 ConversionPatternRewriter &rewriter) const override {
919 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
920 Value barrier =
921 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
922 adaptor.getMbarId(), rewriter);
923 Value txcount = truncToI32(b, adaptor.getTxcount());
924 NVVM::MBarrierArriveExpectTxOp::create(
925 rewriter, op->getLoc(), barrier, txcount, // barrier and txcount
926 NVVM::MemScopeKind::CTA, // default scope is CTA
927 false, // relaxed-semantics is false
928 adaptor.getPredicate());
929 rewriter.eraseOp(op);
930 return success();
931 }
932};
933
934struct NVGPUMBarrierTryWaitParityLowering
935 : public MBarrierBasePattern<nvgpu::MBarrierTryWaitParityOp> {
936 using MBarrierBasePattern<
937 nvgpu::MBarrierTryWaitParityOp>::MBarrierBasePattern;
938 LogicalResult
939 matchAndRewrite(nvgpu::MBarrierTryWaitParityOp op, OpAdaptor adaptor,
940 ConversionPatternRewriter &rewriter) const override {
941 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
942 Value barrier =
943 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
944 adaptor.getMbarId(), rewriter);
945 Value ticks = truncToI32(b, adaptor.getTicks());
946 Value phase =
947 LLVM::ZExtOp::create(b, b.getI32Type(), adaptor.getPhaseParity());
948 rewriter.replaceOpWithNewOp<NVVM::MBarrierTryWaitParityOp>(op, barrier,
949 phase, ticks);
950 return success();
951 }
952};
953
954struct NVGPUTmaAsyncLoadOpLowering
955 : public MBarrierBasePattern<nvgpu::TmaAsyncLoadOp> {
956 using MBarrierBasePattern<nvgpu::TmaAsyncLoadOp>::MBarrierBasePattern;
957 LogicalResult
958 matchAndRewrite(nvgpu::TmaAsyncLoadOp op, OpAdaptor adaptor,
959 ConversionPatternRewriter &rewriter) const override {
960 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
961 auto srcMemrefType = cast<MemRefType>(op.getDst().getType());
962 Value dest = getStridedElementPtr(rewriter, op->getLoc(), srcMemrefType,
963 adaptor.getDst(), {});
964 // Intrinsics takes a shared-cluster pointer so we need an
965 // address space cast from 3 to 7.
966 // TODO: Introduce AS(7) in NVGPU.
967 auto ptrSharedClusterType = LLVM::LLVMPointerType::get(
968 op->getContext(),
969 static_cast<unsigned>(NVVM::NVVMMemorySpace::SharedCluster));
970 dest = LLVM::AddrSpaceCastOp::create(b, ptrSharedClusterType, dest);
971
972 Value barrier =
973 getMbarrierPtr(b, op.getBarriers().getType(), adaptor.getBarriers(),
974 adaptor.getMbarId(), rewriter);
975
976 SmallVector<Value> coords = adaptor.getCoordinates();
977 for (auto [index, value] : llvm::enumerate(coords)) {
978 coords[index] = truncToI32(b, value);
979 }
980
981 // TODO: Enhance the NVGPU Op for other modes too
982 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(
983 op, dest, adaptor.getTensorMapDescriptor(), coords, barrier,
984 ValueRange{}, adaptor.getMulticastMask(), Value{},
985 NVVM::TMALoadMode::TILE, // default is TILE mode
986 false, // default is cluster-scope
987 nullptr, // default is no cta-group
988 adaptor.getPredicate());
989 return success();
990 }
991};
992
993struct NVGPUTmaAsyncStoreOpLowering
994 : public MBarrierBasePattern<nvgpu::TmaAsyncStoreOp> {
995 using MBarrierBasePattern<nvgpu::TmaAsyncStoreOp>::MBarrierBasePattern;
996 LogicalResult
997 matchAndRewrite(nvgpu::TmaAsyncStoreOp op, OpAdaptor adaptor,
998 ConversionPatternRewriter &rewriter) const override {
999 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1000 auto srcMemrefType = cast<MemRefType>(op.getSrc().getType());
1001 Value dest = getStridedElementPtr(rewriter, op->getLoc(), srcMemrefType,
1002 adaptor.getSrc(), {});
1003 SmallVector<Value> coords = adaptor.getCoordinates();
1004 for (auto [index, value] : llvm::enumerate(coords)) {
1005 coords[index] = truncToI32(b, value);
1006 }
1007
1008 // TODO: Enhance the NVGPU Op for other modes too
1009 rewriter.replaceOpWithNewOp<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(
1010 op, adaptor.getTensorMapDescriptor(), dest, coords, Value{},
1011 NVVM::TMAStoreMode::TILE, // default is TILE mode
1012 adaptor.getPredicate());
1013 return success();
1014 }
1015};
1016
1017struct NVGPUGenerateWarpgroupDescriptorLowering
1018 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupGenerateDescriptorOp> {
1019 using ConvertOpToLLVMPattern<
1020 nvgpu::WarpgroupGenerateDescriptorOp>::ConvertOpToLLVMPattern;
1021
1022 LogicalResult
1023 matchAndRewrite(nvgpu::WarpgroupGenerateDescriptorOp op, OpAdaptor adaptor,
1024 ConversionPatternRewriter &rewriter) const override {
1025
1026 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1027
1028 nvgpu::TensorMapSwizzleKind swizzleKind =
1029 op.getTensorMap().getType().getSwizzle();
1030
1031 unsigned layout =
1032 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 128
1033 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 64
1034 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 32
1035 : 1;
1036 unsigned swizzle =
1037 (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_128B) ? 1
1038 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_64B) ? 2
1039 : (swizzleKind == nvgpu::TensorMapSwizzleKind::SWIZZLE_32B) ? 3
1040 : 0;
1041
1042 auto ti64 = b.getIntegerType(64);
1043 auto makeConst = [&](uint64_t index) -> Value {
1044 return LLVM::ConstantOp::create(b, ti64, b.getI64IntegerAttr(index));
1045 };
1046 auto shiftLeft = [&](Value value, unsigned shift) -> Value {
1047 return LLVM::ShlOp::create(b, ti64, value, makeConst(shift));
1048 };
1049 auto shiftRight = [&](Value value, unsigned shift) -> Value {
1050 return LLVM::LShrOp::create(b, ti64, value, makeConst(shift));
1051 };
1052 auto insertBit = [&](Value desc, Value val, int startBit) {
1053 return LLVM::OrOp::create(b, ti64, desc, shiftLeft(val, startBit));
1054 };
1055
1056 int64_t sizeN = op.getTensorMap().getType().getTensor().getDimSize(0);
1057 uint64_t strideDimVal = (layout << 3) >> exclude4LSB;
1058 uint64_t leadDimVal = (sizeN * layout) >> exclude4LSB;
1059 uint64_t offsetVal = 0;
1060
1061 Value strideDim = makeConst(strideDimVal);
1062 Value leadDim = makeConst(leadDimVal);
1063
1064 Value baseAddr = getStridedElementPtr(
1065 rewriter, op->getLoc(), cast<MemRefType>(op.getTensor().getType()),
1066 adaptor.getTensor(), {});
1067 Value basePtr = LLVM::PtrToIntOp::create(b, ti64, baseAddr);
1068 // Just use 14 bits for base address
1069 Value basePtr14bit = shiftRight(shiftLeft(basePtr, 46), 50);
1070
1071 int startSwizzleBit = 62, startOffsetBit = 49, startStrideBit = 32,
1072 startLeadBit = 16, startBaseAddrBit = 0;
1073 Value dsc = makeConst(0);
1074 // // [62,64) swizzle type
1075 dsc = insertBit(dsc, makeConst(swizzle), startSwizzleBit);
1076 // // [49,52) base_offset
1077 dsc = insertBit(dsc, makeConst(offsetVal), startOffsetBit);
1078 // // [32,46) stride
1079 dsc = insertBit(dsc, strideDim, startStrideBit);
1080 // // [16,30) leading dimension
1081 dsc = insertBit(dsc, leadDim, startLeadBit);
1082 // // [0,14) start_address
1083 dsc = insertBit(dsc, basePtr14bit, startBaseAddrBit);
1084
1085 LDBG() << "Generating warpgroup.descriptor: " << "leading_off:"
1086 << leadDimVal << "\t" << "stride_off :" << strideDimVal << "\t"
1087 << "base_offset:" << offsetVal << "\t" << "layout_type:" << swizzle
1088 << " (" << nvgpu::stringifyTensorMapSwizzleKind(swizzleKind)
1089 << ")\n start_addr : " << baseAddr;
1090
1091 rewriter.replaceOp(op, dsc);
1092 return success();
1093 }
1094};
1095
1096static Value makeI64Const(ImplicitLocOpBuilder &b, int32_t index) {
1097 return LLVM::ConstantOp::create(b, b.getIntegerType(64),
1098 b.getI64IntegerAttr(index));
1099}
1100
1101/// Returns a Value that holds data type enum that is expected by CUDA driver.
1102static Value elementTypeAsLLVMConstant(ImplicitLocOpBuilder &b, Type type) {
1103 // Enum is from CUDA driver API
1104 // https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__TYPES.html
1105 enum CUtensorMapDataTypeEnum {
1106 CU_TENSOR_MAP_DATA_TYPE_UINT8 = 0,
1107 CU_TENSOR_MAP_DATA_TYPE_UINT16,
1108 CU_TENSOR_MAP_DATA_TYPE_UINT32,
1109 CU_TENSOR_MAP_DATA_TYPE_INT32,
1110 CU_TENSOR_MAP_DATA_TYPE_UINT64,
1111 CU_TENSOR_MAP_DATA_TYPE_INT64,
1112 CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
1113 CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
1114 CU_TENSOR_MAP_DATA_TYPE_FLOAT64,
1115 CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
1116 CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ,
1117 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32,
1118 CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ
1119 };
1120
1121 if (type.isUnsignedInteger(8))
1122 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT8);
1123 if (type.isUnsignedInteger(16))
1124 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT16);
1125 if (type.isUnsignedInteger(32))
1126 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT32);
1127 if (type.isUnsignedInteger(64))
1128 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_UINT64);
1129 if (type.isSignlessInteger(32))
1130 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_INT32);
1131 if (type.isSignlessInteger(64))
1132 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_INT64);
1133 if (type.isF16())
1134 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT16);
1135 if (type.isF32())
1136 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT32);
1137 if (type.isF64())
1138 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_FLOAT64);
1139 if (type.isBF16())
1140 return makeI64Const(b, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16);
1141
1142 llvm_unreachable("Not supported data type");
1143}
1144
1145struct NVGPUTmaCreateDescriptorOpLowering
1146 : public ConvertOpToLLVMPattern<nvgpu::TmaCreateDescriptorOp> {
1147 using ConvertOpToLLVMPattern<
1148 nvgpu::TmaCreateDescriptorOp>::ConvertOpToLLVMPattern;
1149 LogicalResult
1150 matchAndRewrite(nvgpu::TmaCreateDescriptorOp op, OpAdaptor adaptor,
1151 ConversionPatternRewriter &rewriter) const override {
1152 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1153 auto llvmPointerType = LLVM::LLVMPointerType::get(op->getContext());
1154 Type llvmInt64Type = IntegerType::get(op->getContext(), 64);
1155
1156 Value tensorElementType =
1157 elementTypeAsLLVMConstant(b, op.getTensor().getType().getElementType());
1158 auto promotedOperands = getTypeConverter()->promoteOperands(
1159 b.getLoc(), op->getOperands(), adaptor.getOperands(), b);
1160
1161 Value boxArrayPtr = LLVM::AllocaOp::create(
1162 b, llvmPointerType, llvmInt64Type, makeI64Const(b, 5));
1163 for (auto [index, value] : llvm::enumerate(adaptor.getBoxDimensions())) {
1164 Value gep = LLVM::GEPOp::create(b, llvmPointerType, llvmPointerType,
1165 boxArrayPtr, makeI64Const(b, index));
1166 LLVM::StoreOp::create(b, value, gep);
1167 }
1168
1169 nvgpu::TensorMapDescriptorType desc = op.getTensorMap().getType();
1170 // Set Arguments for the function call
1171 SmallVector<Value> arguments;
1172 arguments.push_back(promotedOperands[0]); // rank
1173 arguments.push_back(promotedOperands[1]); // descriptor
1174 arguments.push_back(tensorElementType); // data type
1175 arguments.push_back(
1176 makeI64Const(b, (int)desc.getInterleave())); // interleave
1177 arguments.push_back(makeI64Const(b, (int)desc.getSwizzle())); // swizzle
1178 arguments.push_back(makeI64Const(b, (int)desc.getL2promo())); // l2promo
1179 arguments.push_back(makeI64Const(b, (int)desc.getOob())); // oob
1180 arguments.push_back(boxArrayPtr); // box dimensions
1181
1182 // Set data types of the arguments
1183 SmallVector<Type> argTypes = {
1184 llvmInt64Type, /* int64_t tensorRank */
1185 llvmPointerType, /* ptr */
1186 llvmInt64Type, /* int64_t */
1187 llvmInt64Type, /* int64_t */
1188 llvmInt64Type, /* int64_t */
1189 llvmInt64Type, /* int64_t */
1190 llvmInt64Type, /* int64_t */
1191 llvmPointerType /* ptr */
1192 };
1193 FunctionCallBuilder hostRegisterCallBuilder = {
1194 "mgpuTensorMapEncodeTiledMemref", llvmPointerType, argTypes};
1195 Value tensorMap =
1196 hostRegisterCallBuilder.create(b.getLoc(), b, arguments).getResult();
1197
1198 rewriter.replaceOp(op, tensorMap);
1199 return success();
1200 }
1201};
1202
1203struct NVGPUWarpgroupMmaOpLowering
1204 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp> {
1205 using ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaOp>::ConvertOpToLLVMPattern;
1206
1207 /// This is a helper class to generate required NVVM Ops for warp-group level
1208 /// matrix multiplication.
1209 /// When the given GEMM shape is larger than the shape of
1210 /// a wgmma instrution in PTX, it can generate multiple NVVM::WgmmaMmaAsyncOp
1211 /// Op(s), group and execute them asynchronously. The class also handles
1212 /// waiting for completion and iterates through WarpgroupMatrixDescriptor to
1213 /// create descriptors for each instruction.
1214 ///
1215 /// For example this is the case when the shape of GEMM is 128x128x128
1216 ///
1217 /// nvvm.wgmma.fence.aligned
1218 ///
1219 /// nvvm.wgmma.mma.async descA, descB
1220 /// iterate(descA, descB)
1221 /// nvvm.wgmma.mma.async descA, descB
1222 /// [6x times more]
1223 ///
1224 /// nvvm.wgmma.group.sync.aligned
1225 /// nvvm.wgmma.wait.group.sync [groupId]
1226 ///
1227 class WarpgroupGemm {
1228 nvgpu::WarpgroupMmaOp op;
1229 ImplicitLocOpBuilder b;
1230 OpAdaptor adaptor;
1231
1232 // Entire shape of the given Op
1233 int64_t totalM, totalN, totalK;
1234
1235 // Shape of one wgmma instruction
1236 int wgmmaM = 0, wgmmaN = 0, wgmmaK = 0;
1237
1238 // Iteration counts for GEMM
1239 int iterationM = 0, iterationN = 0, iterationK = 0;
1240
1241 /// The function returns the shape of wgmma instruction that is defined in
1242 /// PTX programming guide.
1243 /// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#asynchronous-warpgroup-level-matrix-shape
1244 void findWgmmaShape(int64_t sizeM, int64_t sizeN, Type inputElemType) {
1245 wgmmaM = 64;
1246 wgmmaN = sizeN;
1247 if (inputElemType.isTF32()) {
1248 wgmmaK = 8;
1249 } else if (inputElemType.isF16() || inputElemType.isBF16()) {
1250 wgmmaK = 16;
1251 } else if (isa<Float8E4M3FNType, Float8E5M2Type>(inputElemType) ||
1252 inputElemType.isInteger(16)) {
1253 wgmmaK = 32;
1254 } else if (inputElemType.isInteger(1)) {
1255 wgmmaK = 256;
1256 } else {
1257 llvm_unreachable("msg: not supported K shape");
1258 }
1259 LDBG() << "Generating WgmmaMmaAsyncOp shape[m = " << wgmmaM
1260 << ", n = " << wgmmaN << ", k = " << wgmmaK << "]";
1261 }
1262
1263 /// Generates WGMMATypesAttr from MLIR Type
1264 NVVM::WGMMATypesAttr generateWgmmaType(Type type,
1265 bool useF32 = false) const {
1266 auto getWgmmaType = [=](Type elemType) {
1267 if (elemType.isF32() || elemType.isTF32())
1268 return useF32 ? NVVM::WGMMATypes::f32 : NVVM::WGMMATypes::tf32;
1269 if (elemType.isF16())
1270 return NVVM::WGMMATypes::f16;
1271 if (elemType.isBF16())
1272 return NVVM::WGMMATypes::bf16;
1273 if (isa<Float8E4M3FNType>(elemType))
1274 return NVVM::WGMMATypes::e4m3;
1275 if (isa<Float8E5M2Type>(elemType))
1276 return NVVM::WGMMATypes::e5m2;
1277 if (elemType.isInteger(1))
1278 return NVVM::WGMMATypes::b1;
1279 if (elemType.isInteger(8))
1280 return NVVM::WGMMATypes::s8;
1281 if (elemType.isUnsignedInteger(8))
1282 return NVVM::WGMMATypes::u8;
1283 if (elemType.isInteger(32))
1284 return NVVM::WGMMATypes::s32;
1285 llvm_unreachable("unsupported type");
1286 };
1287 return NVVM::WGMMATypesAttr::get(op->getContext(), getWgmmaType(type));
1288 }
1289
1290 /// Generates layout attribute for the input matrix for wgmma instruction
1291 NVVM::MMALayoutAttr
1292 generateWgmmaLayout(std::optional<bool> transpose) const {
1293 if (transpose.value_or(false))
1294 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::col);
1295 return NVVM::MMALayoutAttr::get(op->getContext(), NVVM::MMALayout::row);
1296 }
1297
1298 /// Generates shape attribute for wgmma instruction
1299 NVVM::MMAShapeAttr generateWgmmaShape() const {
1300 return NVVM::MMAShapeAttr::get(op->getContext(), wgmmaM, wgmmaN, wgmmaK);
1301 }
1302
1303 /// Generates scale attributes of output matrix for wgmma instruction
1304 NVVM::WGMMAScaleOutAttr generateScaleOut() const {
1305 return NVVM::WGMMAScaleOutAttr::get(op->getContext(),
1306 NVVM::WGMMAScaleOut::one);
1307 }
1308 /// Generates scale attributes of input matrix for wgmma instruction
1309 NVVM::WGMMAScaleInAttr generateScaleIn() const {
1310 return NVVM::WGMMAScaleInAttr::get(op->getContext(),
1311 NVVM::WGMMAScaleIn::one);
1312 }
1313
1314 /// Basic function to generate Add
1315 Value makeAdd(Value lhs, Value rhs) {
1316 return LLVM::AddOp::create(b, lhs.getType(), lhs, rhs);
1317 };
1318
1319 /// Moves the descriptor pointer of matrix-A for the next wgmma instruction.
1320 /// Currently, it only handles row-major.
1321 ///
1322 /// It moves the pointer like below for [128][64] size:
1323 /// +2 +4 +6
1324 /// ↓ ↓ ↓
1325 /// descA ---> +--+--+--+--+
1326 /// |->|->|->|->|
1327 /// | | | | |
1328 /// | | | | |
1329 /// | | | | |
1330 /// descA+512---> +-----------+
1331 /// | | | | |
1332 /// | | | | |
1333 /// | | | | |
1334 /// | | | | |
1335 /// +-----------+
1336 ///
1337 Value iterateDescriptorA(Value desc, int i, int j, int k) {
1338 MemRefType matrixTypeA = op.getDescriptorA().getType().getTensor();
1339 Type elemA = matrixTypeA.getElementType();
1340 int byte = elemA.getIntOrFloatBitWidth() / 8;
1341 int tileShapeA = matrixTypeA.getDimSize(1);
1342 int incrementVal = ((wgmmaK * k) + (totalK * tileShapeA * i)) * byte;
1343 incrementVal = incrementVal >> exclude4LSB;
1344 LDBG() << "\t\t[m: " << i << " n: " << j << " k: " << k
1345 << "] [wgmma descriptors] Descriptor A + " << incrementVal
1346 << " | \t ";
1347 if (!incrementVal)
1348 return desc;
1349 return makeAdd(desc, makeI64Const(b, incrementVal));
1350 }
1351
1352 /// Moves the descriptor pointer of matrix-B for the next wgmma instruction.
1353 /// Currently, it only handles column-major.
1354 ///
1355 /// It moves the pointer like below for [128][64] size:
1356 /// descB ---> +--+--+--+--+--+--+--+--+
1357 /// |↓ | | | | | | | |
1358 /// |↓ | | | | | | | |
1359 /// |↓ | | | | | | | |
1360 /// |↓ | | | | | | | |
1361 /// +--+--+--+--+--+--+--+--+
1362 ///
1363 Value iterateDescriptorB(Value desc, int i, int j, int k) {
1364 MemRefType matrixTypeB = op.getDescriptorB().getType().getTensor();
1365 Type elemB = matrixTypeB.getElementType();
1366 int byte = elemB.getIntOrFloatBitWidth() / 8;
1367 int incrementVal = matrixTypeB.getDimSize(0) * wgmmaK * k * byte;
1368 incrementVal = incrementVal >> exclude4LSB;
1369 LDBG() << "Descriptor B + " << incrementVal;
1370 if (!incrementVal)
1371 return desc;
1372 return makeAdd(desc, makeI64Const(b, incrementVal));
1373 }
1374
1375 /// This function generates a WgmmaMmaAsyncOp using provided GMMA matrix
1376 /// descriptors and arranges them based on induction variables: i, j, and k.
1377 Value generateWgmma(int i, int j, int k, Value matrixC) {
1378 LDBG() << "\t wgmma." << "m" << wgmmaM << "n" << wgmmaN << "k" << wgmmaK
1379 << "(A[" << (iterationM * wgmmaM) << ":"
1380 << (iterationM * wgmmaM) + wgmmaM << "][" << (iterationK * wgmmaK)
1381 << ":" << (iterationK * wgmmaK + wgmmaK) << "] * " << " B["
1382 << (iterationK * wgmmaK) << ":" << (iterationK * wgmmaK + wgmmaK)
1383 << "][" << 0 << ":" << wgmmaN << "])";
1384
1385 Value descriptorA = iterateDescriptorA(adaptor.getDescriptorA(), i, j, k);
1386 Value descriptorB = iterateDescriptorB(adaptor.getDescriptorB(), i, j, k);
1387
1388 Type elemA = op.getDescriptorA().getType().getTensor().getElementType();
1389 NVVM::WGMMATypesAttr itypeA = generateWgmmaType(elemA);
1390
1391 Type elemB = op.getDescriptorB().getType().getTensor().getElementType();
1392 NVVM::WGMMATypesAttr itypeB = generateWgmmaType(elemB);
1393
1394 Type elemD = op.getMatrixC().getType().getFragmented().getElementType();
1395 NVVM::WGMMATypesAttr itypeD = generateWgmmaType(elemD, true);
1396
1397 NVVM::MMAShapeAttr shape = generateWgmmaShape();
1398 NVVM::WGMMAScaleOutAttr scaleOut = generateScaleOut();
1399 NVVM::WGMMAScaleInAttr scaleIn = generateScaleIn();
1400 NVVM::MMALayoutAttr layoutA = generateWgmmaLayout(op.getTransposeA());
1401 NVVM::MMALayoutAttr layoutB = generateWgmmaLayout(!op.getTransposeB());
1402
1403 auto overflow = NVVM::MMAIntOverflowAttr::get(
1404 op->getContext(), NVVM::MMAIntOverflow::wrapped);
1405
1406 return NVVM::WgmmaMmaAsyncOp::create(
1407 b, matrixC.getType(), matrixC, descriptorA, descriptorB, shape,
1408 itypeA, itypeB, itypeD, scaleOut, scaleIn, scaleIn, layoutA, layoutB,
1409 overflow);
1410 }
1411
1412 /// Generates multiple wgmma instructions to complete the given GEMM shape
1413 Value generateWgmmaGroup() {
1414 Value wgmmaResult =
1415 LLVM::PoisonOp::create(b, adaptor.getMatrixC().getType());
1416
1417 // Perform GEMM
1418 SmallVector<Value> wgmmaResults;
1419 for (int i = 0; i < iterationM; ++i) {
1420 Value matrixC =
1421 LLVM::ExtractValueOp::create(b, adaptor.getMatrixC(), i);
1422 for (int j = 0; j < iterationN; ++j)
1423 for (int k = 0; k < iterationK; ++k)
1424 matrixC = generateWgmma(i, j, k, matrixC);
1425 wgmmaResults.push_back(matrixC);
1426 }
1427 for (auto [idx, matrix] : llvm::enumerate(wgmmaResults)) {
1428 wgmmaResult = LLVM::InsertValueOp::create(b, wgmmaResult.getType(),
1429 wgmmaResult, matrix, idx);
1430 }
1431 return wgmmaResult;
1432 }
1433
1434 public:
1435 WarpgroupGemm(nvgpu::WarpgroupMmaOp op, ImplicitLocOpBuilder &b,
1436 OpAdaptor adaptor)
1437 : op(op), b(b), adaptor(adaptor) {
1438 // Find the entire GEMM Shape
1439 totalM = op.getDescriptorA().getType().getTensor().getDimSize(0);
1440 totalN = op.getDescriptorB().getType().getTensor().getDimSize(1);
1441 totalK = op.getDescriptorA().getType().getTensor().getDimSize(1);
1442 LDBG() << "===--- GEMM D[" << totalM << "][" << totalN << "] += A["
1443 << totalM << "][" << totalK << "] * B[" << totalK << "][" << totalN
1444 << "] ---===";
1445
1446 // Find the shape for one wgmma instruction
1447 findWgmmaShape(
1448 totalM, totalN,
1449 op.getDescriptorA().getType().getTensor().getElementType());
1450
1451 // Iterations counts to complete the given shape with wgmma shape
1452 iterationM = totalM / wgmmaM;
1453 iterationN = totalN / wgmmaN;
1454 iterationK = totalK / wgmmaK;
1455 }
1456
1457 /// Generates WgmmaMmaAsync Ops to complete the specified GEMM shape. It
1458 /// includes generating a fence Op (WgmmaFenceAlignedOp) before the
1459 /// instructions and group synchronization, as well as waiting
1460 /// (WgmmaGroupSyncAlignedOp) for group synchronization
1461 /// (WgmmaWaitGroupSyncOp) after the instructions.
1462 Value generateWarpgroupMma() {
1463 NVVM::WgmmaFenceAlignedOp::create(b);
1464 Value wgmmaResult = generateWgmmaGroup();
1465 NVVM::WgmmaGroupSyncAlignedOp::create(b);
1466 NVVM::WgmmaWaitGroupSyncOp::create(b, op.getWaitGroup());
1467 return wgmmaResult;
1468 }
1469 };
1470 LogicalResult
1471 matchAndRewrite(nvgpu::WarpgroupMmaOp op, OpAdaptor adaptor,
1472 ConversionPatternRewriter &rewriter) const override {
1473 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1474
1475 // Step 1. Build a helper class
1476 WarpgroupGemm warpgroupGemm(op, b, adaptor);
1477
1478 // Step 2. Get the entire GEMM Shape
1479 Value wgmmaResult = warpgroupGemm.generateWarpgroupMma();
1480
1481 // Step 3. Replace fragmented result struct with the op results
1482 rewriter.replaceOp(op, wgmmaResult);
1483 return success();
1484 }
1485};
1486
1487struct NVGPUWarpgroupMmaStoreOpLowering
1488 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaStoreOp> {
1489 using ConvertOpToLLVMPattern<
1490 nvgpu::WarpgroupMmaStoreOp>::ConvertOpToLLVMPattern;
1491
1492 /// This function stores a fragmented register matrix owned by a warp group
1493 /// (128 threads) into a memref. Each thread has 64 registers, each the size
1494 /// of a struct.
1495 /// Here is what each threads (T) holds, each `d` is struct value with a
1496 /// number.
1497 ///
1498 /// Threads in warp-group (128 threads) and what they owns in the matrixD:
1499 /// 0-31 Warp-0 -> MatrixD[0:15 ][0:N]
1500 /// 32-63 Warp-1 -> MatrixD[16:31][0:N]
1501 /// 64-95 Warp-2 -> MatrixD[32:47][0:N]
1502 /// 96-127 Warp-3 -> MatrixD[48:64][0:N]
1503 ///
1504 /// Matrix-D:
1505 /// +______________________________________________________________________+
1506 /// | 0-1 | 2-3 | 4-5 | 6-7 | 8-9 | 10-11|..|N-8,N-7 |
1507 /// 0 | T0:d0-d1 |T1:d0-d1 |T2:d0-d1 |T3:d0-d1 |T0:d4-d5| T1:d4-d5..|T0:dX-dY|
1508 /// 1 | T4:d0-d1 |T5:d0-d1 |T6:d0-d1 |T7:d0-d1 |T4:d4-d5| T5:d4-d5..|T4:dX-dY|
1509 /// ..| .........|.........|.........|.........|........|...........|........|
1510 /// 8 | T0:d2-d3 |T1:d2-d3 |T2:d2-d3 |T3:d2-d3 |T0:d6-d7|T1:d6-d7,..|T0:dZ-dW|
1511 /// 9 | T4:d2-d3 |T5:d2-d3 |T6:d2-d3 |T7:d2-d3 |T4:d6-d7| T5:d6-d7..|T4:dZ-dW|
1512 /// ..| .........|.........|.........|.........|........|...........|........|
1513 /// 15| T28:d2-d3|T29:d2-d3|T30:d2-d3|T31:d2-d3|........|...........|........|
1514 /// 16| T32:d2-d3|T33:d2-d3|T34:d2-d3|T35:d2-d3|........|...........|........|
1515 /// ..| .........|.........|.........|.........|........|...........|........|
1516 /// 32| T64:d2-d3|T65:d2-d3|T66:d2-d3|T67:d2-d3|........|...........|........|
1517 /// ..| .........|.........|.........|.........|........|...........|........|
1518 /// 48| T96:d2-d3|T97:d2-d3|T98:d2-d3|T99:d2-d3|........|...........|........|
1519 /// ..| .........|.........|.........|.........|........|...........|........|
1520 /// +______________________________________________________________________+
1521 ///
1522 /// \param rewriter: The pattern rewriter.
1523 /// \param matrixD: Result of the warp-group MMA operation (fragmented
1524 /// matrix). It is holded by a thread and a struct with 64 elements.
1525 /// \param dstMemref: The memref where the registers will be stored.
1526 /// \param offset: the offset within the memref where the registers will be
1527 /// stored.
1528 void storeFragmentedMatrix(ImplicitLocOpBuilder &b, Value matrixD,
1529 TypedValue<MemRefType> dstMemref,
1530 int offset) const {
1531 Type i32 = b.getI32Type();
1532
1533 auto makeConst = [&](int32_t index) -> Value {
1534 return LLVM::ConstantOp::create(b, i32, b.getI32IntegerAttr(index));
1535 };
1536 Value c1 = makeConst(1);
1537 Value c2 = makeConst(2);
1538 Value c4 = makeConst(4);
1539 Value c8 = makeConst(8);
1540 Value c16 = makeConst(16);
1541 Value warpSize = makeConst(kWarpSize);
1542
1543 auto makeMul = [&](Value lhs, Value rhs) -> Value {
1544 return LLVM::MulOp::create(b, lhs.getType(), lhs, rhs);
1545 };
1546 auto makeAdd = [&](Value lhs, Value rhs) -> Value {
1547 return LLVM::AddOp::create(b, lhs.getType(), lhs, rhs);
1548 };
1549
1550 auto makeExtractAndStore = [&](int i, Value wgmmaResult, Value x, Value y,
1552 Type it = b.getIndexType();
1553 Value idx = arith::IndexCastOp::create(b, it, x);
1554 Value idy0 = arith::IndexCastOp::create(b, it, y);
1555 Value idy1 = arith::IndexCastOp::create(b, it, makeAdd(y, c1));
1556 Value d0 = LLVM::ExtractValueOp::create(b, wgmmaResult, i);
1557 Value d1 = LLVM::ExtractValueOp::create(b, wgmmaResult, i + 1);
1558 memref::StoreOp::create(b, d0, memref, ValueRange{idx, idy0});
1559 memref::StoreOp::create(b, d1, memref, ValueRange{idx, idy1});
1560 };
1561
1562 Value tidx = NVVM::ThreadIdXOp::create(b, i32);
1563 Value laneId = LLVM::URemOp::create(b, i32, tidx, warpSize);
1564 Value warpId = LLVM::UDivOp::create(b, i32, tidx, warpSize);
1565 Value lane4Id = LLVM::UDivOp::create(b, i32, laneId, c4);
1566 Value lane4modId = LLVM::URemOp::create(b, i32, laneId, c4);
1567
1568 Value tj = makeMul(lane4modId, c2);
1569 Value ti = makeAdd(lane4Id, makeMul(warpId, c16));
1570 if (offset)
1571 ti = makeAdd(ti, makeConst(offset));
1572
1573 auto structType = cast<LLVM::LLVMStructType>(matrixD.getType());
1574
1575 // Number of 32-bit registers owns per thread
1576 constexpr unsigned numAdjacentRegisters = 2;
1577 // Number of 8x8 matrices one below another per warp
1578 constexpr unsigned numStackedMatrices = 2;
1579
1580 size_t storeCount = (structType.getBody().size() /
1581 (numStackedMatrices * numAdjacentRegisters));
1582
1583 for (size_t i = 0; i < numStackedMatrices; ++i) {
1584 Value idx = makeAdd(ti, makeMul(makeConst(i), c8));
1585 for (size_t j = 0; j < storeCount; ++j) {
1586 Value idy = makeAdd(tj, makeMul(makeConst(j), c8));
1587 size_t structIndex = (i * numAdjacentRegisters) +
1588 (j * (numStackedMatrices * numAdjacentRegisters));
1589 makeExtractAndStore(structIndex, matrixD, idx, idy, dstMemref);
1590 }
1591 }
1592 }
1593
1594 LogicalResult
1595 matchAndRewrite(nvgpu::WarpgroupMmaStoreOp op, OpAdaptor adaptor,
1596 ConversionPatternRewriter &rewriter) const override {
1597 int offset = 0;
1598 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1599 Value matriDValue = adaptor.getMatrixD();
1600 auto stype = cast<LLVM::LLVMStructType>(matriDValue.getType());
1601 for (auto [idx, matrixD] : llvm::enumerate(stype.getBody())) {
1602 auto structType = cast<LLVM::LLVMStructType>(matrixD);
1603 Value innerStructValue =
1604 LLVM::ExtractValueOp::create(b, matriDValue, idx);
1605 storeFragmentedMatrix(b, innerStructValue, op.getDstMemref(), offset);
1606 offset += structType.getBody().size();
1607 }
1608 rewriter.eraseOp(op);
1609 return success();
1610 }
1611};
1612
1613struct NVGPUWarpgroupMmaInitAccumulatorOpLowering
1614 : public ConvertOpToLLVMPattern<nvgpu::WarpgroupMmaInitAccumulatorOp> {
1615 using ConvertOpToLLVMPattern<
1616 nvgpu::WarpgroupMmaInitAccumulatorOp>::ConvertOpToLLVMPattern;
1617 LogicalResult
1618 matchAndRewrite(nvgpu::WarpgroupMmaInitAccumulatorOp op, OpAdaptor adaptor,
1619 ConversionPatternRewriter &rewriter) const override {
1620 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1621 LLVM::LLVMStructType packStructType = cast<LLVM::LLVMStructType>(
1622 getTypeConverter()->convertType(op.getMatrixC().getType()));
1623 Type elemType = cast<LLVM::LLVMStructType>(packStructType.getBody().front())
1624 .getBody()
1625 .front();
1626 Value zero = LLVM::ConstantOp::create(b, elemType, b.getZeroAttr(elemType));
1627 Value packStruct = LLVM::PoisonOp::create(b, packStructType);
1628 SmallVector<Value> innerStructs;
1629 // Unpack the structs and set all values to zero
1630 for (auto [idx, s] : llvm::enumerate(packStructType.getBody())) {
1631 auto structType = cast<LLVM::LLVMStructType>(s);
1632 Value structValue = LLVM::ExtractValueOp::create(b, packStruct, idx);
1633 for (unsigned i = 0; i < structType.getBody().size(); ++i) {
1634 structValue = LLVM::InsertValueOp::create(b, structType, structValue,
1635 zero, ArrayRef<int64_t>({i}));
1636 }
1637 innerStructs.push_back(structValue);
1638 }
1639 // Pack the inner structs into a single struct
1640 for (auto [idx, matrix] : llvm::enumerate(innerStructs)) {
1641 packStruct = LLVM::InsertValueOp::create(b, packStruct.getType(),
1642 packStruct, matrix, idx);
1643 }
1644 rewriter.replaceOp(op, packStruct);
1645 return success();
1646 }
1647};
1648
1649struct NVGPUTmaFenceOpLowering
1650 : public ConvertOpToLLVMPattern<nvgpu::TmaFenceOp> {
1651 using ConvertOpToLLVMPattern<nvgpu::TmaFenceOp>::ConvertOpToLLVMPattern;
1652 LogicalResult
1653 matchAndRewrite(nvgpu::TmaFenceOp op, OpAdaptor adaptor,
1654 ConversionPatternRewriter &rewriter) const override {
1655 MLIRContext *ctx = op.getContext();
1656 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1657 auto i32Ty = b.getI32Type();
1658 Value tensormapSize =
1659 LLVM::ConstantOp::create(b, i32Ty, rewriter.getI32IntegerAttr(128));
1660
1661 auto memscope =
1662 NVVM::MemScopeKindAttr::get(ctx, ::mlir::NVVM::MemScopeKind::SYS);
1663
1664 rewriter.replaceOpWithNewOp<NVVM::FenceProxyAcquireOp>(
1665 op, memscope, adaptor.getTensorMapDescriptor(), tensormapSize);
1666
1667 return success();
1668 }
1669};
1670
1671struct NVGPUTmaPrefetchOpLowering
1672 : public ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp> {
1673 using ConvertOpToLLVMPattern<nvgpu::TmaPrefetchOp>::ConvertOpToLLVMPattern;
1674 LogicalResult
1675 matchAndRewrite(nvgpu::TmaPrefetchOp op, OpAdaptor adaptor,
1676 ConversionPatternRewriter &rewriter) const override {
1677 rewriter.replaceOpWithNewOp<NVVM::PrefetchOp>(
1678 op, /* CacheLevel */ nullptr, /* Cache Eviction Priority */ nullptr,
1679 adaptor.getTensorMapDescriptor(), adaptor.getPredicate(),
1680 /* Tensormap UnitAttr */ mlir::UnitAttr::get(op.getContext()));
1681 return success();
1682 }
1683};
1684
1685struct NVGPURcpOpLowering : public ConvertOpToLLVMPattern<nvgpu::RcpOp> {
1686 using ConvertOpToLLVMPattern<nvgpu::RcpOp>::ConvertOpToLLVMPattern;
1687 LogicalResult
1688 matchAndRewrite(nvgpu::RcpOp op, OpAdaptor adaptor,
1689 ConversionPatternRewriter &rewriter) const override {
1690 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1691 auto i64Ty = b.getI64Type();
1692 auto f32Ty = b.getF32Type();
1693 VectorType inTy = op.getIn().getType();
1694 // apply rcp.approx.ftz.f on each element in vector.
1695 auto convert1DVec = [&](Type llvm1DVectorTy, Value inVec) {
1696 Value ret1DVec = LLVM::PoisonOp::create(b, llvm1DVectorTy);
1697 int numElems = llvm::cast<VectorType>(llvm1DVectorTy).getNumElements();
1698 for (int i = 0; i < numElems; i++) {
1699 Value idx = LLVM::ConstantOp::create(b, i64Ty, b.getI64IntegerAttr(i));
1700 Value elem = LLVM::ExtractElementOp::create(b, inVec, idx);
1701 Value dst = NVVM::RcpApproxFtzF32Op::create(b, f32Ty, elem);
1702 ret1DVec = LLVM::InsertElementOp::create(b, ret1DVec, dst, idx);
1703 }
1704 return ret1DVec;
1705 };
1706 if (inTy.getRank() == 1) {
1707 rewriter.replaceOp(op, convert1DVec(inTy, adaptor.getIn()));
1708 return success();
1709 }
1711 op.getOperation(), adaptor.getOperands(), *(this->getTypeConverter()),
1712 [&](Type llvm1DVectorTy, ValueRange operands) -> Value {
1713 OpAdaptor adaptor(operands);
1714 return convert1DVec(llvm1DVectorTy, adaptor.getIn());
1715 },
1716 rewriter);
1717 }
1718};
1719
1720//===----------------------------------------------------------------------===//
1721// NVGPUTruncfOp Lowering
1722//===----------------------------------------------------------------------===//
1723
1724enum class FPKind { F32, BF16, F16, F8, F6, F4 };
1725
1726/// Get the effective bit width of a floating-point type.
1727/// f6 types are 6-bit but NVVM Ops expect 8-bit (i8) containers.
1728static int getEffectiveBitWidth(int bitWidth) {
1729 return bitWidth == 6 ? 8 : bitWidth;
1730}
1731
1732static std::optional<FPKind> classifyFPType(Type t) {
1733 static constexpr auto isConvertibleF8Type = [](Type t) {
1734 return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
1735 };
1736 static constexpr auto isConvertibleF6Type = [](Type t) {
1737 return isa<Float6E2M3FNType, Float6E3M2FNType>(t);
1738 };
1739 static constexpr auto isConvertibleF4Type = [](Type t) {
1740 return isa<Float4E2M1FNType>(t);
1741 };
1742
1743 if (t.isF32())
1744 return FPKind::F32;
1745 if (t.isBF16())
1746 return FPKind::BF16;
1747 if (t.isF16())
1748 return FPKind::F16;
1749 if (isConvertibleF8Type(t))
1750 return FPKind::F8;
1751 if (isConvertibleF6Type(t))
1752 return FPKind::F6;
1753 if (isConvertibleF4Type(t))
1754 return FPKind::F4;
1755
1756 return std::nullopt;
1757}
1758
1759/// Conversion op identifier for nvgpu.truncf lowering dispatch table.
1760enum class FPTruncConvOp {
1761 F32x2_TO_F16x2,
1762 F32x2_TO_BF16x2,
1763 F32x2_TO_F8x2,
1764 F32x2_TO_F6x2,
1765 F32x2_TO_F4x2,
1766 F16x2_TO_F8x2,
1767 F16x2_TO_F6x2,
1768 F16x2_TO_F4x2,
1769 BF16x2_TO_F8x2,
1770 BF16x2_TO_F6x2,
1771 BF16x2_TO_F4x2,
1772};
1773
1774struct FPTruncTableEntry {
1775 FPKind src;
1776 FPKind dst;
1777 FPTruncConvOp convOp;
1778};
1779
1780static constexpr FPTruncTableEntry kFPTruncTable[] = {
1781 // f32 source
1782 {FPKind::F32, FPKind::F16, FPTruncConvOp::F32x2_TO_F16x2},
1783 {FPKind::F32, FPKind::BF16, FPTruncConvOp::F32x2_TO_BF16x2},
1784 {FPKind::F32, FPKind::F8, FPTruncConvOp::F32x2_TO_F8x2},
1785 {FPKind::F32, FPKind::F6, FPTruncConvOp::F32x2_TO_F6x2},
1786 {FPKind::F32, FPKind::F4, FPTruncConvOp::F32x2_TO_F4x2},
1787 // f16 source
1788 {FPKind::F16, FPKind::F8, FPTruncConvOp::F16x2_TO_F8x2},
1789 {FPKind::F16, FPKind::F6, FPTruncConvOp::F16x2_TO_F6x2},
1790 {FPKind::F16, FPKind::F4, FPTruncConvOp::F16x2_TO_F4x2},
1791 // bf16 source
1792 {FPKind::BF16, FPKind::F8, FPTruncConvOp::BF16x2_TO_F8x2},
1793 {FPKind::BF16, FPKind::F6, FPTruncConvOp::BF16x2_TO_F6x2},
1794 {FPKind::BF16, FPKind::F4, FPTruncConvOp::BF16x2_TO_F4x2},
1795};
1796
1797/// Find the conversion table entry whose source/destination `FPKind`s match the
1798/// given element types.
1799template <typename TableEntry, size_t N>
1800static std::optional<TableEntry>
1801lookupConvOp(const TableEntry (&table)[N], Type srcElemType, Type dstElemType) {
1802 std::optional<FPKind> srcKind = classifyFPType(srcElemType);
1803 std::optional<FPKind> dstKind = classifyFPType(dstElemType);
1804 if (!srcKind || !dstKind)
1805 return std::nullopt;
1806 for (const TableEntry &entry : table) {
1807 if (entry.src == *srcKind && entry.dst == *dstKind)
1808 return entry;
1809 }
1810 return std::nullopt;
1811}
1812
1813/// Extract a single element from a vector.
1814static Value extractElement(ImplicitLocOpBuilder &b, Value srcVec, int idx) {
1815 assert(idx >= 0 &&
1816 idx < cast<VectorType>(srcVec.getType()).getNumElements() &&
1817 "extractElement: index out of bounds");
1818 IntegerType i64Ty = b.getI64Type();
1819 return b.create<LLVM::ExtractElementOp>(
1820 srcVec, b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(idx)));
1821}
1822
1823/// Extract a pair of f32 values from an i32 vector at the given base index.
1824static std::pair<Value, Value> extractF32Pair(ImplicitLocOpBuilder &b,
1825 Value srcI32Vec, int baseIdx) {
1826 FloatType f32Ty = b.getF32Type();
1827 Value elem0 = extractElement(b, srcI32Vec, baseIdx);
1828 Value elem1 = extractElement(b, srcI32Vec, baseIdx + 1);
1829 return {b.create<LLVM::BitcastOp>(f32Ty, elem0),
1830 b.create<LLVM::BitcastOp>(f32Ty, elem1)};
1831}
1832
1833/// Extract a vector of elements of size i32 from an i32 vector and bitcast to
1834/// the specified vector type.
1835static Value extractAndBitcast(ImplicitLocOpBuilder &b, Value srcI32Vec,
1836 int idx, VectorType vecTy) {
1837 Value elem = extractElement(b, srcI32Vec, idx);
1838 return b.create<LLVM::BitcastOp>(vecTy, elem);
1839}
1840
1841/// Create a sub-byte conversion from an f32 pair source and return the native
1842/// result.
1843template <typename ConvertOp, typename... Args>
1844static Value convertFromF32Pair(ImplicitLocOpBuilder &b, Value srcI32Vec,
1845 int srcBaseIdx, Type resultTy, Args &&...args) {
1846 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1847 return b.create<ConvertOp>(resultTy, hi, lo, std::forward<Args>(args)...);
1848}
1849
1850/// Create a sub-byte conversion from a packed f16x2/bf16x2 source and return
1851/// the native result.
1852template <typename ConvertOp, typename... Args>
1853static Value convertFromPacked(ImplicitLocOpBuilder &b, Value srcI32Vec,
1854 int srcBaseIdx, Type srcElemTy, Type resultTy,
1855 Args &&...args) {
1856 Value src = extractAndBitcast(b, srcI32Vec, srcBaseIdx,
1857 VectorType::get(2, srcElemTy));
1858 return b.create<ConvertOp>(resultTy, src, std::forward<Args>(args)...);
1859}
1860
1861/// Create a typed NVVM truncation conversion.
1862static Value createTruncConversion(
1863 ImplicitLocOpBuilder &b, MLIRContext *ctx, FPTruncConvOp convOp,
1864 Value srcI32Vec, int srcBaseIdx, NVVM::FPRoundingModeAttr rndAttr,
1865 NVVM::SaturationModeAttr satAttr, BoolAttr reluAttr, Type dstElemType,
1866 Type actualDstFloatType, Value randomBits = Value()) {
1867 IntegerType i8Ty = b.getI8Type();
1868 IntegerType i16Ty = b.getI16Type();
1869 IntegerType i32Ty = b.getI32Type();
1870 TypeAttr dstTyAttr = TypeAttr::get(dstElemType);
1871 TypeAttr actualDstTyAttr = TypeAttr::get(actualDstFloatType);
1872
1873 switch (convOp) {
1874 case FPTruncConvOp::F32x2_TO_F16x2: {
1875 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1876 Value r = b.create<NVVM::ConvertF32x2ToF16x2Op>(
1877 VectorType::get(2, b.getF16Type()), hi, lo, randomBits, rndAttr,
1878 satAttr, reluAttr);
1879 return b.create<LLVM::BitcastOp>(i32Ty, r);
1880 }
1881 case FPTruncConvOp::F32x2_TO_BF16x2: {
1882 auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
1883 Value r = b.create<NVVM::ConvertF32x2ToBF16x2Op>(
1884 VectorType::get(2, b.getBF16Type()), hi, lo, randomBits, rndAttr,
1885 satAttr, reluAttr);
1886 return b.create<LLVM::BitcastOp>(i32Ty, r);
1887 }
1888 case FPTruncConvOp::F32x2_TO_F8x2:
1889 return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
1890 b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
1891 case FPTruncConvOp::F32x2_TO_F6x2:
1892 return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
1893 b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
1894 case FPTruncConvOp::F32x2_TO_F4x2:
1895 return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
1896 b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
1897 case FPTruncConvOp::F16x2_TO_F8x2:
1898 return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
1899 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
1900 case FPTruncConvOp::F16x2_TO_F6x2:
1901 return convertFromPacked<NVVM::ConvertF16x2ToF6x2Op>(
1902 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr,
1903 actualDstTyAttr);
1904 case FPTruncConvOp::F16x2_TO_F4x2:
1905 return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
1906 b, srcI32Vec, srcBaseIdx, b.getF16Type(), i8Ty, reluAttr,
1907 actualDstTyAttr);
1908 case FPTruncConvOp::BF16x2_TO_F8x2:
1909 return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
1910 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, rndAttr, satAttr,
1911 reluAttr, dstTyAttr);
1912 case FPTruncConvOp::BF16x2_TO_F6x2:
1913 return convertFromPacked<NVVM::ConvertBF16x2ToF6x2Op>(
1914 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, reluAttr,
1915 actualDstTyAttr);
1916 case FPTruncConvOp::BF16x2_TO_F4x2:
1917 return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
1918 b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i8Ty, reluAttr,
1919 actualDstTyAttr);
1920 }
1921 llvm_unreachable("unhandled FPTruncConvOp");
1922}
1923
1924static LogicalResult lowerTruncf(nvgpu::TruncfOp op,
1925 nvgpu::TruncfOp::Adaptor adaptor,
1926 ConversionPatternRewriter &rewriter,
1927 const LLVMTypeConverter *typeConverter) {
1928 MLIRContext *ctx = op.getContext();
1929 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
1930 IntegerType i32Ty = b.getI32Type();
1931 IntegerType i64Ty = b.getI64Type();
1932 static constexpr int regBits = 32;
1933
1934 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
1935 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
1936 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
1937 return rewriter.notifyMatchFailure(
1938 op, "expected 1-D vector; canonicalize pattern handles other shapes");
1939
1940 auto srcElemType = srcType.getElementType();
1941 auto dstElemType = dstType.getElementType();
1942 int srcBW = srcType.getElementTypeBitWidth();
1943 int dstBW = dstType.getElementTypeBitWidth();
1944 int numElems = srcType.getNumElements();
1945
1946 NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
1947 NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
1948 auto reluBoolAttr = op.getReluAttr();
1949 Value randomBits = adaptor.getRandomBits();
1950 Type actualDstFloatType = dstElemType;
1951
1952 // STEP 1: bitcast input vector to i32 vector type.
1953 // f64 -> f32/f16/bf16 lowers to a single direct LLVM fptrunc
1954 // f64 -> f8/f6/f4 first truncates to f32 and then reuses the narrow
1955 // conversion path below.
1956 Value input = adaptor.getIn();
1957 if (srcBW == 64) {
1958 if (dstBW >= 16) {
1959 Type convertedType = typeConverter->convertType(dstType);
1960 assert(convertedType && "failed to convert type");
1961 Value result = b.create<LLVM::FPTruncOp>(convertedType, input);
1962 rewriter.replaceOp(op, result);
1963 return success();
1964 }
1965 auto f32VecTy = VectorType::get(srcType.getShape(), b.getF32Type());
1966 input = b.create<LLVM::FPTruncOp>(f32VecTy, input);
1967 srcType = f32VecTy;
1968 srcElemType = b.getF32Type();
1969 srcBW = 32;
1970 }
1971
1972 // f6 types are 6-bit in MLIR but NVVM uses 8-bit containers.
1973 int effectiveDstBW = getEffectiveBitWidth(dstBW);
1974
1975 int srcI32Elems = numElems * srcBW / regBits;
1976 int dstI32Elems = numElems * effectiveDstBW / regBits;
1977 Value srcI32Vec =
1978 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), input);
1979 Value dstI32Vec =
1980 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
1981
1982 // STEP 2: look up the conversion op from the (srcType, dstType) table.
1983 auto convEntry = lookupConvOp(kFPTruncTable, srcElemType, dstElemType);
1984 if (!convEntry)
1985 return rewriter.notifyMatchFailure(
1986 op, "unsupported type combination for truncation");
1987 FPTruncConvOp convOp = convEntry->convOp;
1988
1989 // Number of source-side i32 register slots consumed by each NVVM convert Op.
1990 auto getNumSrcI32PerConvert = [](FPKind src) {
1991 return src == FPKind::F32 ? 2 : 1;
1992 };
1993 int numSrcI32PerConv = getNumSrcI32PerConvert(convEntry->src);
1994
1995 // STEP 3: pack conversion results into destination i32 vector.
1996 const int srcStep = srcBW / effectiveDstBW;
1997 const int resultBW =
1998 effectiveDstBW * 2; // each conversion produces 2 (packed) elements
1999 const int numConvsPerI32 = regBits / resultBW;
2000
2001 for (int srcIdx = 0, dstIdx = 0; dstIdx < dstI32Elems;
2002 srcIdx += srcStep, dstIdx++) {
2003 Value dstIdxConst =
2004 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2005 Value dstValue;
2006
2007 if (numConvsPerI32 == 1) {
2008 // f16/bf16 destinations
2009 dstValue = createTruncConversion(
2010 b, ctx, convOp, srcI32Vec, srcIdx, rndModeAttr, satModeAttr,
2011 reluBoolAttr, dstElemType, actualDstFloatType, randomBits);
2012 } else {
2013 // f8/f6/f4 destinations: pack sub-results via vector insert + bitcast.
2014 auto subResultType = IntegerType::get(ctx, resultBW);
2015 auto subVecTy = VectorType::get(numConvsPerI32, subResultType);
2016 Value subVec = b.create<LLVM::UndefOp>(subVecTy);
2017
2018 int insertIdx = numConvsPerI32 - 1;
2019 int curStep = srcStep;
2020 while (curStep > 0) {
2021 curStep -= numSrcI32PerConv;
2022 Value subResult = createTruncConversion(
2023 b, ctx, convOp, srcI32Vec, srcIdx + curStep, rndModeAttr,
2024 satModeAttr, reluBoolAttr, dstElemType, actualDstFloatType,
2025 /*randomBits=*/Value());
2026 subVec = b.create<LLVM::InsertElementOp>(
2027 subVec, subResult,
2028 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(insertIdx)));
2029 insertIdx--;
2030 }
2031
2032 dstValue = b.create<LLVM::BitcastOp>(i32Ty, subVec);
2033 }
2034
2035 dstI32Vec =
2036 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2037 }
2038
2039 // STEP 4: produce final result.
2040 Type convertedType = typeConverter->convertType(dstType);
2041 assert(convertedType && "failed to convert type");
2042 if (convEntry->dst == FPKind::F6) {
2043 IntegerType i8Ty = b.getI8Type();
2044 auto i8VecTy = VectorType::get(numElems, i8Ty);
2045 Value i8Vec = b.create<LLVM::BitcastOp>(i8VecTy, dstI32Vec);
2046 Value truncVec = b.create<LLVM::TruncOp>(convertedType, i8Vec);
2047 rewriter.replaceOp(op, truncVec);
2048 } else {
2049 auto dstVec = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2050 rewriter.replaceOp(op, dstVec);
2051 }
2052 return success();
2053}
2054
2055struct NVGPUTruncfOpLowering : public ConvertOpToLLVMPattern<nvgpu::TruncfOp> {
2056 using ConvertOpToLLVMPattern<nvgpu::TruncfOp>::ConvertOpToLLVMPattern;
2057
2058 LogicalResult
2059 matchAndRewrite(nvgpu::TruncfOp op, OpAdaptor adaptor,
2060 ConversionPatternRewriter &rewriter) const override {
2061 return lowerTruncf(op, adaptor, rewriter, getTypeConverter());
2062 }
2063};
2064
2065//===----------------------------------------------------------------------===//
2066// NVGPUExtfOp Lowering
2067//===----------------------------------------------------------------------===//
2068
2069/// Conversion op identifier for nvgpu.extf lowering dispatch table.
2070enum class FPExtConvOp {
2071 F8x2_TO_F16x2,
2072 F8x2_TO_BF16x2,
2073 F6x2_TO_F16x2,
2074 F6x2_TO_BF16x2,
2075 F4x2_TO_F16x2,
2076 F4x2_TO_BF16x2,
2077};
2078
2079struct FPExtTableEntry {
2080 FPKind src;
2081 FPKind dst;
2082 FPExtConvOp convOp;
2083};
2084
2085static constexpr FPExtTableEntry kFPExtTable[] = {
2086 {FPKind::F8, FPKind::F16, FPExtConvOp::F8x2_TO_F16x2},
2087 {FPKind::F8, FPKind::BF16, FPExtConvOp::F8x2_TO_BF16x2},
2088 {FPKind::F6, FPKind::F16, FPExtConvOp::F6x2_TO_F16x2},
2089 {FPKind::F6, FPKind::BF16, FPExtConvOp::F6x2_TO_BF16x2},
2090 {FPKind::F4, FPKind::F16, FPExtConvOp::F4x2_TO_F16x2},
2091 {FPKind::F4, FPKind::BF16, FPExtConvOp::F4x2_TO_BF16x2},
2092};
2093
2094/// Create a typed NVVM extension conversion.
2095/// For f8/f6: src is vector<2xi8>. For f4: src is i8.
2096/// Returns i32 (bitcast from vector<2xf16> or vector<2xbf16>).
2097static Value createExtConversion(ImplicitLocOpBuilder &b, MLIRContext *ctx,
2098 FPExtConvOp convOp, Value src,
2099 BoolAttr reluAttr, Type actualSrcFloatType,
2100 Value extScaleFactor = Value()) {
2101 IntegerType i32Ty = b.getI32Type();
2102 auto srcTyAttr = TypeAttr::get(actualSrcFloatType);
2103
2104 switch (convOp) {
2105 case FPExtConvOp::F8x2_TO_F16x2: {
2106 Value r = NVVM::ConvertF8x2ToF16x2Op::create(
2107 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2108 return b.create<LLVM::BitcastOp>(i32Ty, r);
2109 }
2110 case FPExtConvOp::F8x2_TO_BF16x2: {
2111 Value r = NVVM::ConvertF8x2ToBF16x2Op::create(
2112 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2113 return b.create<LLVM::BitcastOp>(i32Ty, r);
2114 }
2115 case FPExtConvOp::F6x2_TO_F16x2: {
2116 Value r = NVVM::ConvertF6x2ToF16x2Op::create(
2117 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2118 return b.create<LLVM::BitcastOp>(i32Ty, r);
2119 }
2120 case FPExtConvOp::F6x2_TO_BF16x2: {
2121 Value r = NVVM::ConvertF6x2ToBF16x2Op::create(
2122 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2123 return b.create<LLVM::BitcastOp>(i32Ty, r);
2124 }
2125 case FPExtConvOp::F4x2_TO_F16x2: {
2126 Value r = NVVM::ConvertF4x2ToF16x2Op::create(
2127 b, VectorType::get(2, b.getF16Type()), src, srcTyAttr, reluAttr);
2128 return b.create<LLVM::BitcastOp>(i32Ty, r);
2129 }
2130 case FPExtConvOp::F4x2_TO_BF16x2: {
2131 Value r = NVVM::ConvertF4x2ToBF16x2Op::create(
2132 b, VectorType::get(2, b.getBF16Type()), src, extScaleFactor, srcTyAttr);
2133 return b.create<LLVM::BitcastOp>(i32Ty, r);
2134 }
2135 }
2136 llvm_unreachable("unhandled FPExtConvOp");
2137}
2138
2139static LogicalResult lowerExtf(nvgpu::ExtfOp op, nvgpu::ExtfOp::Adaptor adaptor,
2140 ConversionPatternRewriter &rewriter,
2141 const LLVMTypeConverter *typeConverter) {
2142 MLIRContext *ctx = op.getContext();
2143 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
2144 IntegerType i8Ty = b.getI8Type();
2145 IntegerType i16Ty = b.getI16Type();
2146 IntegerType i32Ty = b.getI32Type();
2147 IntegerType i64Ty = b.getI64Type();
2148
2149 static constexpr int regBits = 32;
2150 auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
2151 auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
2152 if (!srcType || srcType.getRank() != 1 || !dstType || dstType.getRank() != 1)
2153 return rewriter.notifyMatchFailure(
2154 op, "expected 1-D vector; canonicalize pattern handles other shapes");
2155
2156 auto srcElemType = srcType.getElementType();
2157 auto dstElemType = dstType.getElementType();
2158 int srcBW = srcType.getElementTypeBitWidth();
2159 int dstBW = dstType.getElementTypeBitWidth();
2160 int numElems = srcType.getNumElements();
2161
2162 auto reluBoolAttr = op.getReluAttr();
2163 Type actualSrcFloatType = srcElemType;
2164
2165 assert(dstBW == 16 || dstBW == 32 || dstBW == 64);
2166
2167 // Wide source (f16/bf16/f32) to wide destination (f32/f64): single FPExt.
2168 if (srcBW >= 16 && dstBW >= 32) {
2169 Value result = adaptor.getIn();
2170 if (srcElemType != dstElemType) {
2171 Type convertedType = typeConverter->convertType(dstType);
2172 assert(convertedType && "failed to convert type");
2173 result = b.create<LLVM::FPExtOp>(convertedType, result);
2174 }
2175 rewriter.replaceOp(op, result);
2176 return success();
2177 }
2178
2179 // Narrow source (f8/f6/f4): NVVM typed op produces f16/bf16; optionally
2180 // followed by FPExt to the final f32/f64 destination.
2181 bool needsFinalFPExt = (dstBW >= 32);
2182 Type intermediateDstElem = dstElemType;
2183 if (needsFinalFPExt && llvm::isa<Float8E8M0FNUType>(srcElemType))
2184 intermediateDstElem = b.getBF16Type();
2185 else if (needsFinalFPExt)
2186 intermediateDstElem = b.getF16Type();
2187 int intermediateDstBW = needsFinalFPExt ? 16 : dstBW;
2188
2189 // f6 types are 6-bit in MLIR but NVVM uses 8-bit containers.
2190 int effectiveSrcBW = getEffectiveBitWidth(srcBW);
2191
2192 // STEP 1: prepare input as i32 register vector.
2193 // For f6: zext from vector<Nxi6> to vector<Nxi8>, then bitcast to i32s.
2194 Value inputVec = adaptor.getIn();
2195 if (srcBW == 6) {
2196 auto i8VecTy = VectorType::get(numElems, i8Ty);
2197 inputVec = b.create<LLVM::ZExtOp>(i8VecTy, inputVec);
2198 }
2199
2200 int srcI32Elems = numElems * effectiveSrcBW / regBits;
2201 int dstI32Elems = numElems * intermediateDstBW / regBits;
2202 Value srcI32Vec =
2203 b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), inputVec);
2204 Value dstI32Vec =
2205 b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
2206
2207 // STEP 2: look up the conversion op from the (srcType, dstType) table.
2208 auto convEntry = lookupConvOp(kFPExtTable, srcElemType, intermediateDstElem);
2209 if (!convEntry)
2210 return rewriter.notifyMatchFailure(
2211 op, "unsupported type combination for extension");
2212 FPExtConvOp convOp = convEntry->convOp;
2213 Value extScaleFactor;
2214
2215 // STEP 3: iterate over source i32 elements, producing destination i32s.
2216 for (int srcIdx = 0, dstIdx = 0; srcIdx < srcI32Elems; srcIdx++) {
2217 Value srcI32 = b.create<LLVM::ExtractElementOp>(
2218 srcI32Vec,
2219 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(srcIdx)));
2220
2221 if (effectiveSrcBW == 8) {
2222 // f8/f6: one i32 holds 4 bytes -> split into 2 pairs of i16 -> 2 convs.
2223 Value i16Vec =
2224 b.create<LLVM::BitcastOp>(VectorType::get(2, i16Ty), srcI32);
2225 for (int half = 0; half < 2; half++) {
2226 Value halfI16 = b.create<LLVM::ExtractElementOp>(
2227 i16Vec,
2228 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(half)));
2229 Value src =
2230 b.create<LLVM::BitcastOp>(VectorType::get(2, i8Ty), halfI16);
2231 Value dstValue =
2232 createExtConversion(b, ctx, convOp, src, reluBoolAttr,
2233 actualSrcFloatType, extScaleFactor);
2234 Value dstIdxConst =
2235 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2236 dstI32Vec =
2237 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2238 dstIdx++;
2239 }
2240 } else {
2241 // f4: one i32 holds 4 bytes -> each byte is one conversion input.
2242 Value i8Vec = b.create<LLVM::BitcastOp>(VectorType::get(4, i8Ty), srcI32);
2243 for (int byteIdx = 0; byteIdx < 4; byteIdx++) {
2244 Value src = b.create<LLVM::ExtractElementOp>(
2245 i8Vec,
2246 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(byteIdx)));
2247 Value dstValue =
2248 createExtConversion(b, ctx, convOp, src, reluBoolAttr,
2249 actualSrcFloatType, extScaleFactor);
2250 Value dstIdxConst =
2251 b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
2252 dstI32Vec =
2253 b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
2254 dstIdx++;
2255 }
2256 }
2257 }
2258
2259 // STEP 4: produce final result.
2260 Type convertedType = typeConverter->convertType(dstType);
2261 assert(convertedType && "failed to convert type");
2262 Value result;
2263 if (needsFinalFPExt) {
2264 auto intermediateVecTy = VectorType::get(numElems, intermediateDstElem);
2265 Value intermediateVec =
2266 b.create<LLVM::BitcastOp>(intermediateVecTy, dstI32Vec);
2267 result = b.create<LLVM::FPExtOp>(convertedType, intermediateVec);
2268 } else {
2269 result = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
2270 }
2271 rewriter.replaceOp(op, result);
2272 return success();
2273}
2274
2275struct NVGPUExtfOpLowering : public ConvertOpToLLVMPattern<nvgpu::ExtfOp> {
2276 using ConvertOpToLLVMPattern<nvgpu::ExtfOp>::ConvertOpToLLVMPattern;
2277
2278 LogicalResult
2279 matchAndRewrite(nvgpu::ExtfOp op, OpAdaptor adaptor,
2280 ConversionPatternRewriter &rewriter) const override {
2281 return lowerExtf(op, adaptor, rewriter, getTypeConverter());
2282 }
2283};
2284
2285static int64_t computePaddedElems(int64_t numElems, int srcBW, int dstBW,
2286 int step) {
2287 static constexpr int regBits = 32;
2288 int effSrcBW = getEffectiveBitWidth(srcBW);
2289 int effDstBW = getEffectiveBitWidth(dstBW);
2290 auto ceilDiv = [](int64_t x, int64_t y) { return (x + y - 1) / y; };
2291 int64_t padded =
2292 std::max(ceilDiv(numElems * effSrcBW, regBits) * regBits / effSrcBW,
2293 ceilDiv(numElems * effDstBW, regBits) * regBits / effDstBW);
2294 return ceilDiv(padded, step) * step;
2295}
2296
2297/// Canonicalization pattern for nvgpu.truncf / nvgpu.extf:
2298/// handles scalar inputs, non-32-bit-aligned vectors, and multi-rank vectors.
2299/// Runs as an OpRewritePattern on MLIR types before LLVM type conversion.
2300template <typename CvtOp, bool IsTrunc>
2301struct NVGPUFPCanonicalizePattern : public OpRewritePattern<CvtOp> {
2302 using OpRewritePattern<CvtOp>::OpRewritePattern;
2303
2304 LogicalResult matchAndRewrite(CvtOp op,
2305 PatternRewriter &rewriter) const override {
2306 Type inType = op.getIn().getType();
2307 Type outType = op.getOut().getType();
2308
2309 Type srcElemTy = getElementTypeOrSelf(inType);
2310 Type dstElemTy = getElementTypeOrSelf(outType);
2311 int srcBW = srcElemTy.getIntOrFloatBitWidth();
2312 int dstBW = dstElemTy.getIntOrFloatBitWidth();
2313 int effSrcBW = getEffectiveBitWidth(srcBW);
2314 int effDstBW = getEffectiveBitWidth(dstBW);
2315
2316 bool isScalar = !isa<VectorType>(inType);
2317 auto srcVecTy = dyn_cast<VectorType>(inType);
2318 bool isMultiRank = srcVecTy && srcVecTy.getRank() > 1;
2319 int64_t numElems = isScalar ? 1 : srcVecTy.getNumElements();
2320 int step = IsTrunc ? effSrcBW / effDstBW : effDstBW / effSrcBW;
2321 int64_t paddedElems = computePaddedElems(numElems, srcBW, dstBW, step);
2322 bool needsPad = (paddedElems != numElems);
2323
2324 if (!isScalar && !isMultiRank && !needsPad)
2325 return failure();
2326
2327 ImplicitLocOpBuilder b(op->getLoc(), rewriter);
2328 Value input = op.getIn();
2329
2330 if (isScalar)
2331 input = vector::BroadcastOp::create(b, VectorType::get({1}, srcElemTy),
2332 input);
2333 if (isMultiRank)
2334 input = vector::ShapeCastOp::create(
2335 b, VectorType::get({numElems}, srcElemTy), input);
2336
2337 if (needsPad) {
2338 auto paddedTy = VectorType::get({paddedElems}, srcElemTy);
2339 Value zero = arith::ConstantOp::create(
2340 b, DenseElementsAttr::get(paddedTy, b.getZeroAttr(srcElemTy)));
2341 input = vector::InsertStridedSliceOp::create(
2342 b, input, zero, SmallVector<int64_t>{0}, SmallVector<int64_t>{1});
2343 }
2344
2345 auto cvtDstTy =
2346 VectorType::get({needsPad ? paddedElems : numElems}, dstElemTy);
2347 Value cvt;
2348 if constexpr (IsTrunc) {
2349 cvt = CvtOp::create(b, cvtDstTy, input, op.getRndAttr(), op.getSatAttr(),
2350 op.getReluAttr(), op.getRandomBits());
2351 } else {
2352 cvt =
2353 CvtOp::create(b, cvtDstTy, input, op.getRndAttr(), op.getReluAttr());
2354 }
2355 Value result = cvt;
2356
2357 if (needsPad) {
2358 result = vector::ExtractStridedSliceOp::create(
2359 b, result, SmallVector<int64_t>{0}, SmallVector<int64_t>{numElems},
2360 SmallVector<int64_t>{1});
2361 }
2362
2363 if (isMultiRank) {
2364 result =
2365 vector::ShapeCastOp::create(b, cast<VectorType>(outType), result);
2366 }
2367
2368 if (isScalar) {
2369 result = vector::ExtractOp::create(b, result, SmallVector<int64_t>{0});
2370 }
2371
2372 rewriter.replaceOp(op, result);
2373 return success();
2374 }
2375};
2376
2377using NVGPUTruncfCanonicalizePattern =
2378 NVGPUFPCanonicalizePattern<nvgpu::TruncfOp, true>;
2379using NVGPUExtfCanonicalizePattern =
2380 NVGPUFPCanonicalizePattern<nvgpu::ExtfOp, false>;
2381} // namespace
2382
2384 TypeConverter &typeConverter) {
2385 // NVVM uses alloca in the default address space to represent private
2386 // memory allocations, so drop private annotations. NVVM uses address
2387 // space 3 for shared memory. NVVM uses the default address space to
2388 // represent global memory.
2390 typeConverter, [](gpu::AddressSpace space) -> unsigned {
2391 switch (space) {
2392 case gpu::AddressSpace::Global:
2393 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Global);
2394 case gpu::AddressSpace::Workgroup:
2395 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Shared);
2396 case gpu::AddressSpace::Private:
2397 return 0;
2398 case gpu::AddressSpace::Constant:
2399 return static_cast<unsigned>(NVVM::NVVMMemorySpace::Constant);
2400 }
2401 llvm_unreachable("unknown address space enum value");
2402 });
2403}
2404
2406 const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
2407 patterns.add<
2408 NVGPUMBarrierCreateLowering, // nvgpu.mbarrier.create
2409 NVGPUMBarrierInitLowering, // nvgpu.mbarrier.init
2410 NVGPUMBarrierGetLowering, // nvgpu.mbarrier.get
2411 NVGPUMBarrierArriveLowering, // nvgpu.mbarrier.arrive
2412 NVGPUMBarrierArriveNoCompleteLowering, // nvgpu.mbarrier.arrive.no_complete
2413 NVGPUMBarrierTestWaitLowering, // nvgpu.mbarrier.test_wait_parity
2414 NVGPUMBarrierTryWaitParityLowering, // nvgpu.mbarrier.try_wait_parity
2415 NVGPUTmaAsyncLoadOpLowering, // nvgpu.tma.async.load
2416 NVGPUTmaAsyncStoreOpLowering, // nvgpu.tma.async.store
2417 NVGPUTmaCreateDescriptorOpLowering, // nvgpu.tma.create.descriptor
2418 NVGPUTmaPrefetchOpLowering, // nvgpu.tma.prefetch.descriptor
2419 NVGPUTmaFenceOpLowering, // nvgpu.tma.fence.descriptor
2420 NVGPUMBarrierArriveExpectTxLowering, // nvgpu.mbarrier.arrive.expect_tx
2421 NVGPUGenerateWarpgroupDescriptorLowering, // nvgpu.warpgroup.generate.descriptor
2422 NVGPUWarpgroupMmaOpLowering, // nvgpu.warpgroup.mma
2423 NVGPUWarpgroupMmaStoreOpLowering, // nvgpu.warpgroup.mma.store
2424 NVGPUWarpgroupMmaInitAccumulatorOpLowering, // nvgpu.warpgroup.mma.init.accumulator
2425 NVGPUTruncfOpLowering, // nvgpu.truncf
2426 NVGPUExtfOpLowering, // nvgpu.extf
2427 MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM, NVGPUAsyncCopyLowering,
2428 NVGPUAsyncCreateGroupLowering, NVGPUAsyncWaitLowering,
2429 NVGPUMmaSparseSyncLowering, NVGPURcpOpLowering>(converter);
2430
2431 patterns.add<NVGPUTruncfCanonicalizePattern, NVGPUExtfCanonicalizePattern>(
2432 patterns.getContext());
2433}
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:208
FloatType getF32Type()
Definition Builders.cpp:51
IntegerType getI32Type()
Definition Builders.cpp:71
FloatType getF16Type()
Definition Builders.cpp:47
MLIRContext * getContext() const
Definition Builders.h:56
FloatType getF64Type()
Definition Builders.cpp:53
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
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:71
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:620
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:733
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:310
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...