MLIR 24.0.0git
NVVMDialect.cpp
Go to the documentation of this file.
1//===- NVVMDialect.cpp - NVVM IR Ops and Dialect registration -------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file defines the types and operation details for the NVVM IR dialect in
10// MLIR, and the LLVM IR dialect. It also registers the dialect.
11//
12// The NVVM dialect only contains GPU specific additions on top of the general
13// LLVM dialect.
14//
15//===----------------------------------------------------------------------===//
16
18
22#include "mlir/IR/Builders.h"
25#include "mlir/IR/Diagnostics.h"
27#include "mlir/IR/MLIRContext.h"
28#include "mlir/IR/Operation.h"
30#include "mlir/IR/Types.h"
32#include "llvm/ADT/STLExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/IR/IRBuilder.h"
35#include "llvm/IR/NVVMIntrinsicUtils.h"
36#include "llvm/Support/Casting.h"
37#include "llvm/Support/FormatVariadic.h"
38#include "llvm/Support/NVPTXAddrSpace.h"
39#include "llvm/Support/raw_ostream.h"
40#include <cassert>
41#include <cmath>
42#include <optional>
43#include <string>
44
45using namespace mlir;
46using namespace NVVM;
47
48#include "mlir/Dialect/LLVMIR/NVVMOpsDialect.cpp.inc"
49#include "mlir/Dialect/LLVMIR/NVVMOpsEnums.cpp.inc"
50
51static constexpr unsigned notIntrinsic = llvm::Intrinsic::not_intrinsic;
52
53//===----------------------------------------------------------------------===//
54// Helper/Utility methods
55//===----------------------------------------------------------------------===//
56
57static bool isPtrInAddrSpace(mlir::Value ptr, NVVMMemorySpace targetAS) {
58 auto ptrTy = llvm::cast<LLVM::LLVMPointerType>(ptr.getType());
59 return ptrTy.getAddressSpace() == static_cast<unsigned>(targetAS);
60}
61
63 return isPtrInAddrSpace(ptr, NVVMMemorySpace::Generic);
64}
65
67 return isPtrInAddrSpace(ptr, NVVMMemorySpace::Shared);
68}
69
71 return isPtrInAddrSpace(ptr, NVVMMemorySpace::SharedCluster);
72}
73
74static llvm::Value *castPtrToAddrSpace(llvm::IRBuilderBase &builder,
75 llvm::Value *ptr,
76 NVVMMemorySpace targetAS) {
77 unsigned AS = static_cast<unsigned>(targetAS);
78 return builder.CreateAddrSpaceCast(
79 ptr, llvm::PointerType::get(builder.getContext(), AS));
80}
81
82// Helper method to convert CtaGroupKind in NVVM Dialect to CtaGroupKind in LLVM
83static llvm::nvvm::CTAGroupKind
84getNVVMCtaGroupKind(NVVM::CTAGroupKind ctaGroup) {
85 switch (ctaGroup) {
86 case NVVM::CTAGroupKind::CTA_1:
87 return llvm::nvvm::CTAGroupKind::CG_1;
88 case NVVM::CTAGroupKind::CTA_2:
89 return llvm::nvvm::CTAGroupKind::CG_2;
90 }
91 llvm_unreachable("unsupported cta_group value");
92}
93
94static ParseResult parseCTAGroup(OpAsmParser &parser,
95 NVVM::CTAGroupKindAttr &groupAttr) {
96 StringRef keyword;
97 if (parser.parseKeyword(&keyword))
98 return failure();
99 std::optional<NVVM::CTAGroupKind> group =
100 NVVM::symbolizeCTAGroupKind(keyword);
101 if (!group)
102 return parser.emitError(parser.getNameLoc()) << " expected CTA group";
103 groupAttr = NVVM::CTAGroupKindAttr::get(parser.getContext(), *group);
104 return success();
105}
106
107static void printCTAGroup(OpAsmPrinter &printer, Operation *,
108 NVVM::CTAGroupKindAttr groupAttr) {
109 printer << NVVM::stringifyCTAGroupKind(groupAttr.getValue());
110}
111
112template <typename AttrTy>
113static ParseResult parseEnumKeyword(OpAsmParser &parser, AttrTy &attr) {
114 SMLoc loc = parser.getCurrentLocation();
115 std::string keyword;
116 if (parser.parseKeywordOrString(&keyword))
117 return failure();
118
119 using EnumTy = decltype(attr.getValue());
120 std::optional<EnumTy> value = NVVM::symbolizeEnum<EnumTy>(keyword);
121 if (!value)
122 return parser.emitError(loc) << "unknown enum value '" << keyword << "'";
123
124 attr = AttrTy::get(parser.getContext(), *value);
125 return success();
126}
127
128//===----------------------------------------------------------------------===//
129// Verifier methods
130//===----------------------------------------------------------------------===//
131
132// This verifier is shared among the following Ops:
133// CpAsyncBulkTensorSharedCTAToGlobalOp (TMA Store)
134// CpAsyncBulkTensorReduceOp (TMA Store-Reduce)
135static LogicalResult cpAsyncBulkTensorCommonVerifier(size_t tensorDims,
136 bool isIm2Col,
137 size_t numIm2ColOffsets,
138 Location loc) {
139 if (tensorDims < 1 || tensorDims > 5)
140 return emitError(loc, "expects coordinates between 1 to 5 dimension");
141
142 // For Im2Col mode, there are two constraints:
143 if (isIm2Col) {
144 // 1. Tensor must always be at least 3-d.
145 if (tensorDims < 3)
146 return emitError(
147 loc,
148 "to use im2col mode, the tensor has to be at least 3-dimensional");
149 // 2. When there are Im2ColOffsets, they must be (Dims - 2) in number.
150 if (numIm2ColOffsets && (tensorDims != (numIm2ColOffsets + 2)))
151 return emitError(
152 loc, "im2col offsets must be 2 less than number of coordinates");
153 }
154 return success();
155}
156
158 OperandRange coordinates, OperandRange tensorSize, OperandRange lowerStride,
159 Value upperStride, bool isTile, Location loc) {
160 LogicalResult res = success();
161 if (!tensorSize.empty() && coordinates.size() != tensorSize.size()) {
162 res =
163 emitError(loc, "Expected coordinates size to be equal to tensor size");
164 }
165
166 if (!lowerStride.empty() && tensorSize.empty()) {
167 res = emitError(
168 loc,
169 "Expected tensor_size to be present when lower_stride is provided");
170 } else if (!lowerStride.empty() &&
171 lowerStride.size() != tensorSize.size() - 1) {
172 res = emitError(
173 loc,
174 "Expected lower_stride size to be equal to one less than tensor size");
175 }
176
177 if (!lowerStride.empty() != static_cast<bool>(upperStride)) {
178 res = emitError(loc,
179 "Expected lower_stride and upper_stride to be either both "
180 "present or both absent");
181 }
182
183 bool isDimStride = tensorSize.size() > 0;
184 if (!isTile && isDimStride) {
185 res = emitError(
186 loc, "Only tile mode supports override address with dim and stride");
187 }
188
189 return res;
190}
191
192LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {
193 TMAStoreMode mode = getMode();
194 // We lower through inline-ptx when getPredicate() is true.
195 // a) Only TILE mode is supported
196 // b) Cache-hint is not supported
197 if (getPredicate()) {
198 if (mode != TMAStoreMode::TILE)
199 return emitError("Inline-ptx lowering supported only for Tile mode.");
200 if (getL2CacheHint())
201 return emitError("Inline-ptx lowering unsupported with L2 cache-hint.");
202 }
203
204 size_t dims = getCoordinates().size();
205 switch (mode) {
206 case TMAStoreMode::TILE:
207 return cpAsyncBulkTensorCommonVerifier(dims, false, 0, getLoc());
208 case TMAStoreMode::IM2COL:
209 case TMAStoreMode::IM2COL_W:
210 return cpAsyncBulkTensorCommonVerifier(dims, true, 0, getLoc());
211 case TMAStoreMode::TILE_SCATTER4:
212 if (dims != 5)
213 return emitError("Scatter4 mode expects 5 coordinates");
214 }
215 return success();
216}
217
218LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::verify() {
219 TMAStoreMode mode = getMode();
220 bool isIm2Col =
221 mode == TMAStoreMode::IM2COL || mode == TMAStoreMode::IM2COL_W;
222 bool isTile = mode == TMAStoreMode::TILE;
223
224 LogicalResult commonRes = cpAsyncBulkTensorCommonVerifier(
225 getCoordinates().size(), isIm2Col, 0, getLoc());
226
227 LogicalResult overrideAddrRes = CpAsyncBulkTensorOverrideAddrCommonVerifier(
228 getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
229 isTile, getLoc());
230
231 if (mode == TMAStoreMode::TILE_SCATTER4 && getCoordinates().size() != 5)
232 overrideAddrRes = emitError("Mode tile scatter4 expects 5 coordinates");
233
234 return failed(commonRes) || failed(overrideAddrRes) ? failure() : success();
235}
236
237LogicalResult CpAsyncOp::verify() {
238 if (getModifier() != LoadCacheModifierKind::CG &&
239 getModifier() != LoadCacheModifierKind::CA)
240 return emitError("Only CG and CA cache modifiers are supported.");
241 if (getSize() != 4 && getSize() != 8 && getSize() != 16)
242 return emitError("expected byte size to be either 4, 8 or 16.");
243 if (getModifier() == LoadCacheModifierKind::CG && getSize() != 16)
244 return emitError("CG cache modifier is only support for 16 bytes copy.");
245 return success();
246}
247
248// This verify params can be shared across TMA Load and Prefetch Ops.
249static LogicalResult verifyTMALoadParams(size_t tensorDims, size_t numIm2colOff,
250 TMALoadMode mode, Location loc) {
251 if (tensorDims < 1 || tensorDims > 5)
252 return emitError(loc, "expects coordinates between 1 to 5 dimension");
253
254 auto checkTMALoadParams = [&](TMALoadMode mode, bool isIm2col,
255 size_t expectedIm2colOff) -> LogicalResult {
256 if (isIm2col && (tensorDims < 3))
257 return emitError(loc)
258 << "to use " << mode
259 << " mode, the tensor has to be at least 3-dimensional";
260
261 if (numIm2colOff != expectedIm2colOff)
262 return emitError(loc) << " im2col offsets expected " << expectedIm2colOff
263 << " (provided " << numIm2colOff << ")";
264
265 return success();
266 };
267
268 switch (mode) {
269 case TMALoadMode::TILE:
270 return checkTMALoadParams(mode, false, 0);
271 case TMALoadMode::IM2COL:
272 return checkTMALoadParams(mode, true, tensorDims - 2);
273 case TMALoadMode::IM2COL_W:
274 case TMALoadMode::IM2COL_W_128:
275 return checkTMALoadParams(mode, true, 2);
276 case TMALoadMode::TILE_GATHER4:
277 return (tensorDims == 5)
278 ? checkTMALoadParams(mode, false, 0)
279 : emitError(loc, "Gather4 mode expects 5 coordinates");
280 }
281 return success();
282}
283
284LogicalResult CpAsyncBulkTensorPrefetchOp::verify() {
285 return verifyTMALoadParams(getCoordinates().size(), getIm2colOffsets().size(),
286 getMode(), getLoc());
287}
288
289LogicalResult CpAsyncBulkTensorGlobalToSharedClusterOp::verify() {
290 TMALoadMode mode = getMode();
291 bool isCTAOnly = getIsCTAOnly();
292 if (getPredicate()) { // Inline-asm based lowering
293 if (isCTAOnly)
294 return emitError("Predicate is supported only for shared::cluster mode.");
295 if (mode != TMALoadMode::TILE && mode != TMALoadMode::IM2COL)
296 return emitError(
297 "Predicate is supported only for Tile and Im2col modes.");
298 } else { // Intrinsics-based lowering
299 NVVMMemorySpace expectedAS =
300 isCTAOnly ? NVVMMemorySpace::Shared : NVVMMemorySpace::SharedCluster;
301 unsigned AS = llvm::cast<LLVM::LLVMPointerType>(getDstMem().getType())
302 .getAddressSpace();
303 if (AS != expectedAS)
304 return emitError()
305 << (isCTAOnly
306 ? "Shared::cta destination requires address-space 3."
307 : "Shared::cluster destination requires address-space 7.");
308 // Checks specific to shared::cta mode
309 if (isCTAOnly) {
310 if (getMulticastMask())
311 return emitError("Multicast is not supported with shared::cta mode.");
312 if (getGroup())
313 return emitError("CTAGroup is not supported with shared::cta mode.");
314 }
315 }
316
317 return verifyTMALoadParams(getCoordinates().size(), getIm2colOffsets().size(),
318 getMode(), getLoc());
319}
320
321LogicalResult CpAsyncBulkTensorReduceOp::verify() {
322 TMAStoreMode mode = getMode();
323 size_t dims = getCoordinates().size();
324 switch (mode) {
325 case TMAStoreMode::TILE:
326 return cpAsyncBulkTensorCommonVerifier(dims, false, 0, getLoc());
327 case TMAStoreMode::IM2COL:
328 case TMAStoreMode::IM2COL_W:
329 return cpAsyncBulkTensorCommonVerifier(dims, true, 0, getLoc());
330 case TMAStoreMode::TILE_SCATTER4:
331 return emitError("Scatter mode unsupported for CpAsyncBulkTensorReduceOp");
332 }
333 return success();
334}
335
336LogicalResult CpAsyncBulkTensorReduceOverrideAddrOp::verify() {
337 bool isIm2Col =
338 getMode() == TMAStoreMode::IM2COL || getMode() == TMAStoreMode::IM2COL_W;
339 bool isTile = getMode() == TMAStoreMode::TILE;
340
341 LogicalResult commonRes = cpAsyncBulkTensorCommonVerifier(
342 getCoordinates().size(), isIm2Col, 0, getLoc());
343
344 LogicalResult overrideAddrRes = CpAsyncBulkTensorOverrideAddrCommonVerifier(
345 getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
346 isTile, getLoc());
347
348 if (getMode() == TMAStoreMode::TILE_SCATTER4)
349 overrideAddrRes = emitError(
350 "Scatter mode unsupported for CpAsyncBulkTensorReduceOverrideAddrOp");
351
352 return failed(commonRes) || failed(overrideAddrRes) ? failure() : success();
353}
354
355LogicalResult CpAsyncBulkGlobalToSharedClusterOp::verify() {
356 bool isSharedCTA = isPtrInSharedCTASpace(getDstMem());
357 if (isSharedCTA && getMulticastMask())
358 return emitError("Multicast is not supported with shared::cta mode.");
359
360 return success();
361}
362
363static LogicalResult verifyMBarrierArriveLikeOp(Operation *op, Value addr,
364 NVVM::MemScopeKind scope,
365 Value retVal = nullptr) {
366 if (scope != NVVM::MemScopeKind::CTA && scope != NVVM::MemScopeKind::CLUSTER)
367 return op->emitError("mbarrier scope must be either CTA or Cluster");
368
369 bool isSharedCluster = isPtrInSharedClusterSpace(addr);
370 bool hasRetValue = static_cast<bool>(retVal);
371 if (isSharedCluster && hasRetValue)
372 return op->emitError(
373 "mbarrier in shared_cluster space cannot return any value");
374
375 return success();
376}
377
378LogicalResult MBarrierArriveOp::verify() {
379 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),
380 getRes());
381}
382
383LogicalResult MBarrierArriveDropOp::verify() {
384 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),
385 getRes());
386}
387
388LogicalResult MBarrierArriveExpectTxOp::verify() {
389 // The inline-ptx version of this Op does not support all features.
390 // With predicate, this Op lowers to inline-ptx. So, verify and
391 // error-out if there are unsupported features.
392 if (getPredicate()) {
393 if (getScope() != NVVM::MemScopeKind::CTA)
394 return emitError("mbarrier scope must be CTA when using predicate");
395
396 if (isPtrInSharedClusterSpace(getAddr()))
397 return emitError("mbarrier in shared_cluster space is not supported when "
398 "using predicate");
399
400 if (getRes())
401 return emitError("return-value is not supported when using predicate");
402
403 if (getRelaxed() == true)
404 return emitError("mbarrier with relaxed semantics is not supported when "
405 "using predicate");
406 }
407 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),
408 getRes());
409}
410
411LogicalResult MBarrierArriveDropExpectTxOp::verify() {
412 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope(),
413 getRes());
414}
415
416//===----------------------------------------------------------------------===//
417// inferReturnTypes for mbarrier arrive-like ops
418//===----------------------------------------------------------------------===//
419
420/// Only shared_cluster (ptr<7>) produces zero results; all other address
421/// spaces (including generic) return i64.
422static LogicalResult
424 SmallVectorImpl<Type> &inferredReturnTypes) {
425 if (!isPtrInSharedClusterSpace(addr))
426 inferredReturnTypes.push_back(IntegerType::get(context, 64));
427 return success();
428}
429
430LogicalResult
431MBarrierArriveOp::inferReturnTypes(MLIRContext *context,
432 std::optional<Location> location,
433 MBarrierArriveOp::Adaptor adaptor,
434 SmallVectorImpl<Type> &inferredReturnTypes) {
435 return inferMBarrierArriveResultTypes(context, adaptor.getAddr(),
436 inferredReturnTypes);
437}
438
439LogicalResult MBarrierArriveDropOp::inferReturnTypes(
440 MLIRContext *context, std::optional<Location> location,
441 MBarrierArriveDropOp::Adaptor adaptor,
442 SmallVectorImpl<Type> &inferredReturnTypes) {
443 return inferMBarrierArriveResultTypes(context, adaptor.getAddr(),
444 inferredReturnTypes);
445}
446
447LogicalResult MBarrierArriveExpectTxOp::inferReturnTypes(
448 MLIRContext *context, std::optional<Location> location,
449 MBarrierArriveExpectTxOp::Adaptor adaptor,
450 SmallVectorImpl<Type> &inferredReturnTypes) {
451 // Predicate forces no return value (inline PTX path).
452 // Note: predicate + shared_cluster is rejected by the verifier separately.
453 if (adaptor.getPredicate())
454 return success();
455 return inferMBarrierArriveResultTypes(context, adaptor.getAddr(),
456 inferredReturnTypes);
457}
458
459LogicalResult MBarrierArriveDropExpectTxOp::inferReturnTypes(
460 MLIRContext *context, std::optional<Location> location,
461 MBarrierArriveDropExpectTxOp::Adaptor adaptor,
462 SmallVectorImpl<Type> &inferredReturnTypes) {
463 return inferMBarrierArriveResultTypes(context, adaptor.getAddr(),
464 inferredReturnTypes);
465}
466
467/// For ops with optional results, allow the user to omit the result even when
468/// inference would produce one. This preserves backward compatibility: the
469/// result can be silently discarded (e.g., for fire-and-forget arrive ops).
471 TypeRange actual) {
472 if (actual.empty())
473 return true;
474 return inferred == actual;
475}
476
477bool MBarrierArriveOp::isCompatibleReturnTypes(TypeRange l, TypeRange r) {
479}
480bool MBarrierArriveDropOp::isCompatibleReturnTypes(TypeRange l, TypeRange r) {
482}
483bool MBarrierArriveExpectTxOp::isCompatibleReturnTypes(TypeRange l,
484 TypeRange r) {
486}
487bool MBarrierArriveDropExpectTxOp::isCompatibleReturnTypes(TypeRange l,
488 TypeRange r) {
490}
491
492LogicalResult MBarrierExpectTxOp::verify() {
493 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());
494}
495
496LogicalResult MBarrierCompleteTxOp::verify() {
497 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());
498}
499
500LogicalResult MBarrierTestWaitOp::verify() {
501 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());
502}
503
504LogicalResult MBarrierTryWaitOp::verify() {
505 return verifyMBarrierArriveLikeOp(getOperation(), getAddr(), getScope());
506}
507
508LogicalResult ConvertFloatToTF32Op::verify() {
509 using RndMode = NVVM::FPRoundingMode;
510 switch (getRnd()) {
511 case RndMode::RNA:
512 if (getRelu())
513 return emitError("Relu not supported with rna rounding mode.");
514 break;
515 case RndMode::RN:
516 case RndMode::RZ:
517 break;
518 default:
519 return emitError(
520 "Only {rn,rz,rna} rounding modes supported for ConvertFloatToTF32Op.");
521 }
522 return success();
523}
524
525LogicalResult ConvertF32x2ToF6x2Op::verify() {
527
528 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy())) {
529 return emitOpError("Only ")
530 << mlir::Float6E2M3FNType::get(ctx) << " and "
531 << mlir::Float6E3M2FNType::get(ctx)
532 << " types are supported for conversions from f32x2 to f6x2.";
533 }
534 return success();
535}
536
537LogicalResult ConvertF32x2ToF8x2Op::verify() {
538 using RndMode = NVVM::FPRoundingMode;
539 using SatMode = NVVM::SaturationMode;
540
541 bool isRoundingModeRN = getRnd() == RndMode::RN;
542 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
543 bool isRoundingModeRP = getRnd() == RndMode::RP;
544 bool isSatFinite = getSat() == SatMode::SATFINITE;
545
546 bool hasRelu = getRelu();
547
549
551 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
552 [&](mlir::Type) -> LogicalResult {
553 if (!isRoundingModeRN) {
554 return emitOpError("Only RN rounding mode is supported for "
555 "conversions from f32x2 to ")
556 << mlir::Float8E4M3FNType::get(ctx) << " and "
557 << mlir::Float8E5M2Type::get(ctx) << " types";
558 }
559 if (!isSatFinite) {
560 return emitOpError("Only SATFINITE saturation mode is supported "
561 "for conversions "
562 "from f32x2 to ")
563 << mlir::Float8E4M3FNType::get(ctx) << " and "
564 << mlir::Float8E5M2Type::get(ctx) << " types";
565 }
566 return success();
567 })
568 .Case<mlir::Float8E8M0FNUType>([&](mlir::Type) -> LogicalResult {
569 if (!(isRoundingModeRZ || isRoundingModeRP)) {
570 return emitOpError("Only RZ and RP rounding modes are supported for "
571 "conversions from f32x2 to ")
572 << mlir::Float8E8M0FNUType::get(ctx) << " type";
573 }
574 if (hasRelu) {
575 return emitOpError("relu not supported for conversions to ")
576 << mlir::Float8E8M0FNUType::get(ctx) << " type";
577 }
578 return success();
579 })
580 .Default([&](mlir::Type) {
581 return emitOpError("Only ")
582 << mlir::Float8E4M3FNType::get(ctx) << ", "
583 << mlir::Float8E5M2Type::get(ctx) << ", and "
584 << mlir::Float8E8M0FNUType::get(ctx)
585 << " types are "
586 "supported for conversions from f32x2 to f8x2";
587 });
588}
589
590LogicalResult ConvertF16x2ToF8x2Op::verify() {
592
593 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy())) {
594 return emitOpError("Only ")
595 << mlir::Float8E4M3FNType::get(ctx) << " and "
596 << mlir::Float8E5M2Type::get(ctx)
597 << " types are supported for conversions from f16x2 to f8x2.";
598 }
599 return success();
600}
601
602LogicalResult ConvertBF16x2ToF8x2Op::verify() {
603 using RndMode = NVVM::FPRoundingMode;
604 using SatMode = NVVM::SaturationMode;
605
606 bool isRoundingModeRN = getRnd() == RndMode::RN;
607 bool isRoundingModeRZ = getRnd() == RndMode::RZ;
608 bool isRoundingModeRP = getRnd() == RndMode::RP;
609 bool isSatFinite = getSat() == SatMode::SATFINITE;
610 bool hasRelu = getRelu();
611
613
615 .Case<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(
616 [&](mlir::Type) -> LogicalResult {
617 if (!isRoundingModeRN)
618 return emitOpError("Only RN rounding mode is supported for "
619 "conversions from bf16x2 to ")
620 << mlir::Float8E4M3FNType::get(ctx) << " and "
621 << mlir::Float8E5M2Type::get(ctx) << " types";
622 if (!isSatFinite)
623 return emitOpError("Only SATFINITE saturation mode is supported "
624 "for conversions from bf16x2 to ")
625 << mlir::Float8E4M3FNType::get(ctx) << " and "
626 << mlir::Float8E5M2Type::get(ctx) << " types";
627 return success();
628 })
629 .Case<mlir::Float8E8M0FNUType>([&](mlir::Type) -> LogicalResult {
630 if (!(isRoundingModeRZ || isRoundingModeRP))
631 return emitOpError("Only RZ and RP rounding modes are supported for "
632 "conversions from bf16x2 to ")
633 << mlir::Float8E8M0FNUType::get(ctx) << " type";
634 if (hasRelu)
635 return emitOpError("relu not supported for conversions to ")
636 << mlir::Float8E8M0FNUType::get(ctx) << " type";
637 return success();
638 })
639 .Default([&](mlir::Type) -> LogicalResult {
640 llvm_unreachable("Invalid conversion in ConvertBF16x2ToF8x2Op");
641 return failure();
642 });
643}
644
645LogicalResult ConvertF32x2ToF4x2Op::verify() {
647
648 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
649 return emitOpError("Only ")
650 << mlir::Float4E2M1FNType::get(ctx)
651 << " type is supported for conversions from f32x2 to f4x2.";
652
653 return success();
654}
655
656LogicalResult ConvertF8x2ToBF16x2Op::verify() {
658 if (llvm::isa<Float8E8M0FNUType>(getSrcType())) {
659 if (getSat() != SaturationMode::NONE)
660 return emitOpError(
661 "Only NONE saturation mode is supported for conversions from ")
662 << Float8E8M0FNUType::get(ctx) << " type";
663 if (getScaleFactor())
664 return emitOpError("scaleFactor not supported for conversions from ")
665 << Float8E8M0FNUType::get(ctx) << " type";
666 if (getRelu())
667 return emitOpError("relu not supported for conversions from ")
668 << Float8E8M0FNUType::get(ctx) << " type";
669 }
670
671 return success();
672}
673
674LogicalResult PermuteOp::verify() {
675 using Mode = NVVM::PermuteMode;
676 bool hasHi = static_cast<bool>(getHi());
677
678 switch (getMode()) {
679 case Mode::DEFAULT:
680 case Mode::F4E:
681 case Mode::B4E:
682 if (!hasHi)
683 return emitError("mode '") << getMode() << "' requires 'hi' operand.";
684 break;
685 case Mode::RC8:
686 case Mode::ECL:
687 case Mode::ECR:
688 case Mode::RC16:
689 if (hasHi)
690 return emitError("mode '")
691 << getMode() << "' does not accept 'hi' operand.";
692 break;
693 }
694
695 return success();
696}
697
698//===----------------------------------------------------------------------===//
699// Stochastic Rounding Conversion Ops
700//===----------------------------------------------------------------------===//
701
702static LogicalResult verifyConvertF32x2ToFP16x2Op(Twine dstType,
703 FPRoundingMode rnd,
704 bool hasRandomBits,
705 Operation *op) {
706 static constexpr FPRoundingMode validRndModes[] = {
707 FPRoundingMode::RN, FPRoundingMode::RZ, FPRoundingMode::RS};
708
709 if (!llvm::is_contained(validRndModes, rnd)) {
710 return op->emitOpError(
711 "Only RN, RZ, and RS rounding modes are supported for "
712 "conversions from f32x2 to ")
713 << dstType << ".";
714 }
715
716 if (rnd == FPRoundingMode::RS) {
717 if (!hasRandomBits) {
718 return op->emitOpError("random_bits is required for RS rounding mode.");
719 }
720 } else {
721 if (hasRandomBits) {
722 return op->emitOpError(
723 "random_bits not supported for RN and RZ rounding modes.");
724 }
725 }
726
727 return success();
728}
729
730LogicalResult ConvertF32x2ToF16x2Op::verify() {
731 return verifyConvertF32x2ToFP16x2Op("f16x2", getRnd(),
732 getRandomBits() ? true : false, *this);
733}
734
735LogicalResult ConvertF32x2ToBF16x2Op::verify() {
736 return verifyConvertF32x2ToFP16x2Op("bf16x2", getRnd(),
737 getRandomBits() ? true : false, *this);
738}
739
740LogicalResult ConvertF32x4ToF8x4Op::verify() {
742
743 if (!llvm::isa<mlir::Float8E4M3FNType, mlir::Float8E5M2Type>(getDstTy()))
744 return emitOpError("Only ")
745 << mlir::Float8E4M3FNType::get(ctx) << " and "
746 << mlir::Float8E5M2Type::get(ctx)
747 << " types are supported for conversions from f32x4 to f8x4.";
748
749 return success();
750}
751
752LogicalResult ConvertF32x4ToF6x4Op::verify() {
754
755 if (!llvm::isa<mlir::Float6E2M3FNType, mlir::Float6E3M2FNType>(getDstTy()))
756 return emitOpError("Only ")
757 << mlir::Float6E2M3FNType::get(ctx) << " and "
758 << mlir::Float6E3M2FNType::get(ctx)
759 << " types are supported for conversions from f32x4 to f6x4.";
760
761 return success();
762}
763
764LogicalResult ConvertF32x4ToF4x4Op::verify() {
766
767 if (!llvm::isa<mlir::Float4E2M1FNType>(getDstTy()))
768 return emitOpError("Only ") << mlir::Float4E2M1FNType::get(ctx)
769 << " type is supported for conversions from "
770 "f32x4 to f4x4.";
771
772 return success();
773}
774
775LogicalResult BulkStoreOp::verify() {
776 if (getInitVal() != 0)
777 return emitOpError("only 0 is supported for initVal, got ") << getInitVal();
778 return success();
779}
780
781LogicalResult AsyncStoreGlobalOp::verify() {
782 NVVM::MemScopeKind scope = getScope();
783 bool isMmio = getMmio();
784 bool isMultimem = getMultimem();
785
786 if (scope != MemScopeKind::SYS && scope != MemScopeKind::GPU)
787 return emitOpError("scope must be either SYS or GPU");
788
789 if (isMmio && scope != MemScopeKind::SYS)
790 return emitOpError("mmio is only supported for SYS scope");
791
792 if (isMmio && isMultimem)
793 return emitOpError("multimem is not supported with mmio");
794
795 return success();
796}
797
798LogicalResult PMEventOp::verify() {
799 auto eventId = getEventId();
800 auto maskedEventId = getMaskedEventId();
801 if (!maskedEventId && !eventId) {
802 return emitOpError() << "either `id` or `mask` must be set";
803 }
804
805 if (maskedEventId && eventId) {
806 return emitOpError() << "`id` and `mask` cannot be set at the same time";
807 }
808
809 if (eventId) {
810 if (eventId < 0 || eventId > 15) {
811 return emitOpError() << "`id` must be between 0 and 15";
812 }
813 }
814
815 return llvm::success();
816}
817
818// Given the element type of an operand and whether or not it is an accumulator,
819// this function returns the PTX type (`NVVM::MMATypes`) that corresponds to the
820// operand's element type.
821std::optional<mlir::NVVM::MMATypes>
822MmaOp::inferOperandMMAType(Type operandElType, bool isAccumulator) {
823 auto half2Type =
824 VectorType::get(2, Float16Type::get(operandElType.getContext()));
825 if (operandElType.isF64())
826 return NVVM::MMATypes::f64;
827 if (operandElType.isF16() || operandElType == half2Type)
828 return NVVM::MMATypes::f16;
829 if (operandElType.isF32() && isAccumulator)
830 return NVVM::MMATypes::f32;
831 if (operandElType.isF32() && !isAccumulator)
832 return NVVM::MMATypes::tf32;
833 if (llvm::isa<IntegerType>(operandElType)) {
834 if (isAccumulator)
835 return NVVM::MMATypes::s32;
836 return std::nullopt;
837 }
838
839 if (auto structType = llvm::dyn_cast<LLVM::LLVMStructType>(operandElType)) {
840 if (structType.getBody().empty())
841 return std::nullopt;
842 return inferOperandMMAType(structType.getBody()[0], isAccumulator);
843 }
844
845 return std::nullopt;
846}
847
848static bool isInt4PtxType(MMATypes type) {
849 return (type == MMATypes::u4 || type == MMATypes::s4);
850}
851
852static bool isInt8PtxType(MMATypes type) {
853 return (type == MMATypes::u8 || type == MMATypes::s8);
854}
855
856static bool isIntegerPtxType(MMATypes type) {
857 return isInt4PtxType(type) || isInt8PtxType(type) || type == MMATypes::b1 ||
858 type == MMATypes::s32;
859}
860
861MMATypes MmaOp::accumPtxType() {
862 std::optional<mlir::NVVM::MMATypes> val = inferOperandMMAType(
863 getODSOperands(2).getTypes().front(), /*isAccumulator=*/true);
864 assert(val.has_value() && "accumulator PTX type should always be inferrable");
865 return val.value();
866}
867
868MMATypes MmaOp::resultPtxType() {
869 std::optional<mlir::NVVM::MMATypes> val =
870 inferOperandMMAType(getResult().getType(), /*isAccumulator=*/true);
871 assert(val.has_value() && "result PTX type should always be inferrable");
872 return val.value();
873}
874
875template <typename AttrTy>
876static void printMmaProperty(OpAsmPrinter &printer, bool &isFirst,
877 StringRef keyword, AttrTy value) {
878 printer << (isFirst ? " " : ", ") << keyword << " = ";
879 printer.printStrippedAttrOrType(value);
880 isFirst = false;
881}
882
883template <typename AttrTy>
884static void printMmaEnumProperty(OpAsmPrinter &printer, bool &isFirst,
885 StringRef keyword, AttrTy value) {
886 printer << (isFirst ? " " : ", ") << keyword << " = "
887 << NVVM::stringifyEnum(value.getValue());
888 isFirst = false;
889}
890
891static void printMmaUnitProperty(OpAsmPrinter &printer, bool &isFirst,
892 StringRef keyword) {
893 printer << (isFirst ? " " : ", ") << keyword;
894 isFirst = false;
895}
896
897template <typename AttrTy>
898static ParseResult parseMmaPropertyValue(OpAsmParser &parser,
899 NamedAttrList &attributes,
900 StringRef name) {
901 if (attributes.get(name))
902 return parser.emitError(parser.getCurrentLocation(),
903 "duplicate property '" + name + "'");
904 AttrTy value;
905 if (parser.parseEqual() || parser.parseCustomAttributeWithFallback(value))
906 return failure();
907 attributes.append(name, value);
908 return success();
909}
910
911template <typename AttrTy>
912static ParseResult parseMmaEnumPropertyValue(OpAsmParser &parser,
913 NamedAttrList &attributes,
914 StringRef name) {
915 if (attributes.get(name))
916 return parser.emitError(parser.getCurrentLocation(),
917 "duplicate property '" + name + "'");
918 AttrTy value;
919 if (parser.parseEqual() || parseEnumKeyword(parser, value))
920 return failure();
921 attributes.append(name, value);
922 return success();
923}
924
925static bool isMmaPropertyName(StringRef name) {
926 return llvm::is_contained(
928 "shape", "b1Op", "intOverflowBehavior", "layoutA", "layoutB",
929 "multiplicandAPtxType", "multiplicandBPtxType", "orderedMetadata",
930 "kind", "scaleVecSize", "blockScaleFormat", "operandSegmentSizes"},
931 name);
932}
933
934static ParseResult parseMmaProperties(OpAsmParser &parser,
935 NamedAttrList &attributes,
936 ArrayRef<StringRef> allowedKeywords,
937 ArrayRef<StringRef> requiredProperties) {
938 while (true) {
939 StringRef keyword;
940 if (parser.parseKeyword(&keyword))
941 return failure();
942 if (!llvm::is_contained(allowedKeywords, keyword))
943 return parser.emitError(parser.getCurrentLocation(),
944 "unknown MMA property '" + keyword + "'");
945
946 ParseResult parseResult = success();
947 if (keyword == "shape")
948 parseResult =
949 parseMmaPropertyValue<MMAShapeAttr>(parser, attributes, "shape");
950 else if (keyword == "b1_op")
951 parseResult =
952 parseMmaEnumPropertyValue<MMAB1OpAttr>(parser, attributes, "b1Op");
953 else if (keyword == "int_overflow")
955 parser, attributes, "intOverflowBehavior");
956 else if (keyword == "layout_a")
957 parseResult = parseMmaEnumPropertyValue<MMALayoutAttr>(parser, attributes,
958 "layoutA");
959 else if (keyword == "layout_b")
960 parseResult = parseMmaEnumPropertyValue<MMALayoutAttr>(parser, attributes,
961 "layoutB");
962 else if (keyword == "multiplicand_a_ptx_type")
964 parser, attributes, "multiplicandAPtxType");
965 else if (keyword == "multiplicand_b_ptx_type")
967 parser, attributes, "multiplicandBPtxType");
968 else if (keyword == "kind") {
969 if (llvm::is_contained(allowedKeywords, "block_scale_format"))
971 parser, attributes, "kind");
972 else
973 parseResult =
974 parseMmaEnumPropertyValue<MMAKindAttr>(parser, attributes, "kind");
975 } else if (keyword == "scale_vec_size")
977 parser, attributes, "scaleVecSize");
978 else if (keyword == "block_scale_format")
980 parser, attributes, "blockScaleFormat");
981 else if (keyword == "ordered_metadata") {
982 if (attributes.get("orderedMetadata"))
983 return parser.emitError(parser.getCurrentLocation(),
984 "duplicate property 'orderedMetadata'");
985 attributes.append("orderedMetadata", parser.getBuilder().getUnitAttr());
986 } else {
987 return parser.emitError(parser.getCurrentLocation(),
988 "unknown MMA property '" + keyword + "'");
989 }
990 if (failed(parseResult))
991 return failure();
992 if (failed(parser.parseOptionalComma()))
993 break;
994 }
995
996 for (StringRef property : requiredProperties) {
997 if (!attributes.get(property))
998 return parser.emitError(parser.getCurrentLocation(),
999 "missing required property '" + property + "'");
1000 }
1001
1002 NamedAttrList discardableAttributes;
1003 if (parser.parseOptionalAttrDict(discardableAttributes))
1004 return failure();
1005 for (NamedAttribute attribute : discardableAttributes) {
1006 if (isMmaPropertyName(attribute.getName()))
1007 return parser.emitError(
1008 parser.getCurrentLocation(),
1009 "inherent property '" + attribute.getName().getValue() +
1010 "' must be spelled directly in the operation syntax");
1011 attributes.append(attribute);
1012 }
1013 return success();
1014}
1015
1016void MmaOp::print(OpAsmPrinter &p) {
1017 SmallVector<Type, 4> regTypes;
1018 struct MMAOperandFragment {
1019 StringRef operandName;
1020 StringRef ptxTypeAttr;
1021 SmallVector<Value, 4> regs;
1022 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1023 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1024 };
1025
1026 std::array<MMAOperandFragment, 3> frags{
1027 MMAOperandFragment("A", getMultiplicandAPtxTypeAttrName()),
1028 MMAOperandFragment("B", getMultiplicandBPtxTypeAttrName()),
1029 MMAOperandFragment("C", "")};
1030 SmallVector<StringRef, 4> ignoreAttrNames{
1031 mlir::NVVM::MmaOp::getOperandSegmentSizeAttr()};
1032
1033 for (unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
1034 auto &frag = frags[fragIdx];
1035 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
1036 for (auto operandIdx = varOperandSpec.first;
1037 operandIdx < varOperandSpec.first + varOperandSpec.second;
1038 operandIdx++) {
1039 frag.regs.push_back(this->getOperand(operandIdx));
1040 if (operandIdx == 0) {
1041 regTypes.push_back(this->getOperand(operandIdx).getType());
1042 }
1043 }
1044 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
1045 regTypes.back(), /*isAccumulator=*/fragIdx >= 2);
1046 if (inferredType)
1047 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1048 }
1049
1050 auto printMmaOperand = [&](const MMAOperandFragment &frag) -> void {
1051 p << " " << frag.operandName;
1052 p << "[";
1053 p.printOperands(frag.regs);
1054 p << "] ";
1055 };
1056
1057 for (const auto &frag : frags) {
1058 printMmaOperand(frag);
1059 }
1060
1061 bool isFirstProperty = true;
1062 printMmaProperty(p, isFirstProperty, "shape", getShapeAttr());
1063 if (getB1OpAttr())
1064 printMmaEnumProperty(p, isFirstProperty, "b1_op", getB1OpAttr());
1065 if (getIntOverflowBehaviorAttr())
1066 printMmaEnumProperty(p, isFirstProperty, "int_overflow",
1067 getIntOverflowBehaviorAttr());
1068 printMmaEnumProperty(p, isFirstProperty, "layout_a", getLayoutAAttr());
1069 printMmaEnumProperty(p, isFirstProperty, "layout_b", getLayoutBAttr());
1070 if (getMultiplicandAPtxTypeAttr() &&
1071 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
1072 printMmaEnumProperty(p, isFirstProperty, "multiplicand_a_ptx_type",
1073 getMultiplicandAPtxTypeAttr());
1074 if (getMultiplicandBPtxTypeAttr() &&
1075 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
1076 printMmaEnumProperty(p, isFirstProperty, "multiplicand_b_ptx_type",
1077 getMultiplicandBPtxTypeAttr());
1078 llvm::append_range(ignoreAttrNames,
1079 ArrayRef<StringRef>{getShapeAttrName(), getB1OpAttrName(),
1080 getIntOverflowBehaviorAttrName(),
1081 getLayoutAAttrName(),
1082 getLayoutBAttrName(),
1083 getMultiplicandAPtxTypeAttrName(),
1084 getMultiplicandBPtxTypeAttrName()});
1085 p.printOptionalAttrDict(this->getOperation()->getAttrs(), ignoreAttrNames);
1086
1087 // Print the types of the operands and result.
1088 p << " : " << "(";
1089 llvm::interleaveComma(SmallVector<Type, 3>{frags[0].regs[0].getType(),
1090 frags[1].regs[0].getType(),
1091 frags[2].regs[0].getType()},
1092 p);
1093 p << ")";
1094 p.printArrowTypeList(TypeRange{this->getRes().getType()});
1095}
1096
1097void MmaOp::build(OpBuilder &builder, OperationState &result, Type resultType,
1098 ValueRange operandA, ValueRange operandB, ValueRange operandC,
1099 ArrayRef<int64_t> shape, std::optional<MMAB1Op> b1Op,
1100 std::optional<MMAIntOverflow> intOverflow,
1101 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
1102 std::optional<std::array<MMALayout, 2>> multiplicandLayouts) {
1103
1104 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");
1105 MLIRContext *ctx = builder.getContext();
1106 result.addAttribute(
1107 "shape", builder.getAttr<MMAShapeAttr>(shape[0], shape[1], shape[2]));
1108
1109 result.addOperands(operandA);
1110 result.addOperands(operandB);
1111 result.addOperands(operandC);
1112
1113 if (multiplicandPtxTypes) {
1114 result.addAttribute("multiplicandAPtxType",
1115 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1116 result.addAttribute("multiplicandBPtxType",
1117 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1118 } else {
1119 if (auto res = inferOperandMMAType(operandA[0].getType(), false))
1120 result.addAttribute("multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1121 if (auto res = inferOperandMMAType(operandB[0].getType(), false))
1122 result.addAttribute("multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1123 }
1124
1125 if (multiplicandLayouts) {
1126 result.addAttribute("layoutA",
1127 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[0]));
1128 result.addAttribute("layoutB",
1129 MMALayoutAttr::get(ctx, (*multiplicandLayouts)[1]));
1130 } else {
1131 result.addAttribute("layoutA", MMALayoutAttr::get(ctx, MMALayout::row));
1132 result.addAttribute("layoutB", MMALayoutAttr::get(ctx, MMALayout::col));
1133 }
1134
1135 if (intOverflow.has_value())
1136 result.addAttribute("intOverflowBehavior",
1137 MMAIntOverflowAttr::get(ctx, *intOverflow));
1138 if (b1Op.has_value())
1139 result.addAttribute("b1Op", MMAB1OpAttr::get(ctx, *b1Op));
1140
1141 result.addTypes(resultType);
1142 result.addAttribute(
1143 MmaOp::getOperandSegmentSizeAttr(),
1144 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operandA.size()),
1145 static_cast<int32_t>(operandB.size()),
1146 static_cast<int32_t>(operandC.size())}));
1147}
1148
1149// <operation> :=
1150// A `[` $operandA `]` B `[` $operandB `]` C `[` $operandC `]`
1151// properties attr-dict
1152// : (type($operandA[0]), type($operandB[0]), type($operandC[0]))
1153// `->` type($res)
1154ParseResult MmaOp::parse(OpAsmParser &parser, OperationState &result) {
1155 struct MMAOperandFragment {
1156 std::optional<MMATypes> elemtype;
1157 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1158 SmallVector<Type> regTypes;
1159 };
1160
1161 Builder &builder = parser.getBuilder();
1162 std::array<MMAOperandFragment, 4> frags;
1163
1164 NamedAttrList namedAttributes;
1165
1166 // A helper to parse the operand segments.
1167 auto parseMmaOperand = [&](StringRef operandName,
1168 MMAOperandFragment &frag) -> LogicalResult {
1169 if (parser.parseKeyword(operandName).failed())
1170 return failure();
1171 if (parser
1172 .parseOperandList(frag.regs, OpAsmParser::Delimiter::OptionalSquare)
1173 .failed())
1174 return failure();
1175 return success();
1176 };
1177
1178 // Parse the operand segments.
1179 if (parseMmaOperand("A", frags[0]).failed())
1180 return failure();
1181 if (parseMmaOperand("B", frags[1]).failed())
1182 return failure();
1183 if (parseMmaOperand("C", frags[2]).failed())
1184 return failure();
1185
1186 if (parseMmaProperties(parser, namedAttributes,
1187 {"shape", "b1_op", "int_overflow", "layout_a",
1188 "layout_b", "multiplicand_a_ptx_type",
1189 "multiplicand_b_ptx_type"},
1190 {"shape", "layoutA", "layoutB"}))
1191 return failure();
1192
1193 // Parse the type specification and resolve operands.
1194 SmallVector<Type, 3> operandTypes;
1195 if (failed(parser.parseColon()))
1196 return failure();
1197 if (failed(parser.parseLParen()))
1198 return failure();
1199 if (failed(parser.parseTypeList(operandTypes)))
1200 return failure();
1201 if (failed(parser.parseRParen()))
1202 if (operandTypes.size() != 3)
1203 return parser.emitError(
1204 parser.getNameLoc(),
1205 "expected one type for each operand segment but got " +
1206 Twine(operandTypes.size()) + " types");
1207 for (const auto &iter : llvm::enumerate(operandTypes)) {
1208 auto &frag = frags[iter.index()];
1209 frag.regTypes.resize(frag.regs.size(), iter.value());
1210 if (failed(parser.resolveOperands(frag.regs, frag.regTypes,
1211 parser.getNameLoc(), result.operands)))
1212 return failure();
1213 frag.elemtype = inferOperandMMAType(frag.regTypes[0],
1214 /*isAccumulator*/ iter.index() < 2);
1215 }
1216
1217 Type resultType;
1218 if (parser.parseArrow() || parser.parseType(resultType))
1219 return failure();
1220 frags[3].elemtype = inferOperandMMAType(resultType, /*isAccumulator*/ true);
1221
1222 std::array<StringRef, 2> names{"multiplicandAPtxType",
1223 "multiplicandBPtxType"};
1224 for (unsigned idx = 0; idx < names.size(); idx++) {
1225 const auto &frag = frags[idx];
1226 std::optional<NamedAttribute> attr = namedAttributes.getNamed(names[idx]);
1227 if (!frag.elemtype.has_value() && !attr.has_value()) {
1228 return parser.emitError(
1229 parser.getNameLoc(),
1230 "attribute " + names[idx] +
1231 " is not provided explicitly and cannot be inferred");
1232 }
1233 if (!attr.has_value())
1234 result.addAttribute(
1235 names[idx], MMATypesAttr::get(parser.getContext(), *frag.elemtype));
1236 }
1237
1238 result.addTypes(resultType);
1239 if (!namedAttributes.empty())
1240 result.addAttributes(namedAttributes);
1241 result.addAttribute(MmaOp::getOperandSegmentSizeAttr(),
1242 builder.getDenseI32ArrayAttr({
1243 static_cast<int32_t>(frags[0].regs.size()),
1244 static_cast<int32_t>(frags[1].regs.size()),
1245 static_cast<int32_t>(frags[2].regs.size()),
1246 }));
1247 return success();
1248}
1249
1250LogicalResult MmaOp::verify() {
1251 MLIRContext *context = getContext();
1252 auto f16Ty = Float16Type::get(context);
1253 auto i32Ty = IntegerType::get(context, 32);
1254 auto f16x2Ty = VectorType::get(2, f16Ty);
1255 auto f32Ty = Float32Type::get(context);
1256 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
1257 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
1258
1259 auto s32x4StructTy =
1260 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
1261 auto f32x8StructTy =
1262 LLVM::LLVMStructType::getLiteral(context, SmallVector<Type>(8, f32Ty));
1263 auto f16x2x2StructTy =
1264 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
1265 auto f32x4StructTy =
1266 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
1267 auto s32x2StructTy =
1268 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
1269
1270 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
1271 getShapeAttr().getK()};
1272
1273 // These variables define the set of allowed data types for matrices A, B, C,
1274 // and result.
1275 using AllowedShapes = SmallVector<std::array<int64_t, 3>, 2>;
1276 using AllowedTypes = SmallVector<SmallVector<Type, 4>, 2>;
1277 AllowedShapes allowedShapes;
1278 AllowedTypes expectedA;
1279 AllowedTypes expectedB;
1280 AllowedTypes expectedC;
1281 SmallVector<Type> expectedResult;
1282
1283 // When M = 16, we just need to calculate the number of 8xk tiles, where
1284 // k is a factor that depends on the data type.
1285 if (mmaShape[0] == 16) {
1286 int64_t kFactor;
1287 Type multiplicandFragType;
1288 switch (*getMultiplicandAPtxType()) {
1289 case MMATypes::tf32:
1290 kFactor = 4;
1291 multiplicandFragType = i32Ty;
1292 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1293 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1294 break;
1295 case MMATypes::bf16:
1296 kFactor = 8;
1297 multiplicandFragType = i32Ty;
1298 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1299 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1300 break;
1301 case MMATypes::f16:
1302 kFactor = 8;
1303 multiplicandFragType = f16x2Ty;
1304 expectedResult.push_back(f16x2x2StructTy);
1305 expectedResult.push_back(f32x4StructTy);
1306 break;
1307 case MMATypes::e4m3:
1308 case MMATypes::e5m2:
1309 // FP8 (m16n8k16 / m16n8k32) packs 4 values per 32-bit register, same
1310 // as s8/u8, but the accumulator is f16 or f32 (not integer).
1311 kFactor = 16;
1312 multiplicandFragType = i32Ty;
1313 expectedResult.push_back(f16x2x2StructTy);
1314 expectedResult.push_back(f32x4StructTy);
1315 break;
1316 case MMATypes::s4:
1317 case MMATypes::u4:
1318 kFactor = 32;
1319 break;
1320 case MMATypes::b1:
1321 kFactor = 128;
1322 break;
1323 case MMATypes::s8:
1324 case MMATypes::u8:
1325 kFactor = 16;
1326 break;
1327 default:
1328 return emitError("invalid shape or multiplicand type: ")
1329 << getMultiplicandAPtxType().value();
1330 }
1331
1332 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {
1333 expectedResult.push_back(s32x4StructTy);
1334 expectedC.emplace_back(4, i32Ty);
1335 multiplicandFragType = i32Ty;
1336 } else {
1337 expectedC.emplace_back(2, f16x2Ty);
1338 expectedC.emplace_back(4, f32Ty);
1339 }
1340
1341 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor);
1342 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1343 expectedA.emplace_back(unitA, multiplicandFragType);
1344 expectedB.emplace_back(unitB, multiplicandFragType);
1345 allowedShapes.push_back({16, 8, kFactor});
1346 allowedShapes.push_back({16, 8, kFactor * 2});
1347
1348 if (resultPtxType() != accumPtxType())
1349 return emitOpError("ctype does not match dtype");
1350 }
1351
1352 // In the M=8 case, there is only 1 possible case per data type.
1353 if (mmaShape[0] == 8) {
1354 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1355 expectedA.emplace_back(2, f16x2Ty);
1356 expectedB.emplace_back(2, f16x2Ty);
1357 expectedResult.push_back(f16x2x4StructTy);
1358 expectedResult.push_back(f32x8StructTy);
1359 expectedC.emplace_back(4, f16x2Ty);
1360 expectedC.emplace_back(8, f32Ty);
1361 allowedShapes.push_back({8, 8, 4});
1362 }
1363 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1364 Type f64Ty = Float64Type::get(context);
1365 expectedA.emplace_back(1, f64Ty);
1366 expectedB.emplace_back(1, f64Ty);
1367 expectedC.emplace_back(2, f64Ty);
1368 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1369 context, SmallVector<Type>(2, f64Ty)));
1370 allowedShapes.push_back({8, 8, 4});
1371 }
1372 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {
1373 expectedA.push_back({i32Ty});
1374 expectedB.push_back({i32Ty});
1375 expectedC.push_back({i32Ty, i32Ty});
1376 expectedResult.push_back(s32x2StructTy);
1377 if (isInt4PtxType(getMultiplicandAPtxType().value()))
1378 allowedShapes.push_back({8, 8, 32});
1379 if (isInt8PtxType(getMultiplicandAPtxType().value()))
1380 allowedShapes.push_back({8, 8, 16});
1381 if (getMultiplicandAPtxType().value() == MMATypes::b1)
1382 allowedShapes.push_back({8, 8, 128});
1383 }
1384 }
1385
1386 std::string errorMessage;
1387 llvm::raw_string_ostream errorStream(errorMessage);
1388
1389 // Check that we matched an existing shape/dtype combination.
1390 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1391 !llvm::is_contained(allowedShapes, mmaShape)) {
1392 errorStream << "unimplemented variant for MMA shape <";
1393 llvm::interleaveComma(mmaShape, errorStream);
1394 errorStream << ">";
1395 return emitOpError(errorMessage);
1396 }
1397
1398 // Verify the operand types for segments of A, B, and C operands.
1399 std::array<StringRef, 3> operandNames{"A", "B", "C"};
1400 for (const auto &iter : llvm::enumerate(
1401 SmallVector<AllowedTypes, 3>{expectedA, expectedB, expectedC})) {
1402 auto spec = this->getODSOperandIndexAndLength(iter.index());
1403 SmallVector<Type, 4> operandTySeg(operand_type_begin() + spec.first,
1404 operand_type_begin() + spec.first +
1405 spec.second);
1406 bool match = llvm::is_contained(iter.value(), operandTySeg);
1407
1408 if (!match) {
1409 errorStream << "Could not match types for the "
1410 << operandNames[iter.index()]
1411 << " operands; expected one of ";
1412 for (const auto &x : iter.value()) {
1413 errorStream << x.size() << "x" << x[0] << " ";
1414 }
1415 errorStream << "but got ";
1416 llvm::interleaveComma(operandTySeg, errorStream);
1417 return emitOpError(errorMessage);
1418 }
1419 }
1420
1421 // Check the result type
1422 if (!llvm::any_of(expectedResult, [&](Type expectedResultType) {
1423 return expectedResultType == getResult().getType();
1424 })) {
1425 errorStream
1426 << "Could not match allowed types for the result; expected one of ";
1427 llvm::interleaveComma(expectedResult, errorStream);
1428 errorStream << " but got " << getResult().getType();
1429 return emitOpError(errorMessage);
1430 }
1431
1432 // Ensure that binary MMA variants have a b1 MMA operation defined.
1433 if (getMultiplicandAPtxType() == MMATypes::b1 && !getB1Op()) {
1434 return emitOpError("op requires " + getB1OpAttrName().strref() +
1435 " attribute");
1436 }
1437
1438 // Ensure int4/int8 MMA variants specify the accum overflow behavior
1439 // attribute.
1440 if (isInt4PtxType(*getMultiplicandAPtxType()) ||
1441 isInt8PtxType(*getMultiplicandAPtxType())) {
1442 if (!getIntOverflowBehavior())
1443 return emitOpError("op requires " +
1444 getIntOverflowBehaviorAttrName().strref() +
1445 " attribute");
1446 }
1447
1448 // Validate layout combinations. According to the operation description, most
1449 // MMA operations require layoutA=row and layoutB=col. Only m8n8k4 with f16
1450 // can use other layout combinations.
1451 bool isM8N8K4_F16 =
1452 (mmaShape[0] == 8 && mmaShape[1] == 8 && mmaShape[2] == 4 &&
1453 getMultiplicandAPtxType() == MMATypes::f16);
1454
1455 if (!isM8N8K4_F16) {
1456 // For all other shapes/types, layoutA must be row and layoutB must be col
1457 if (getLayoutA() != MMALayout::row || getLayoutB() != MMALayout::col) {
1458 return emitOpError("requires layoutA = #nvvm.mma_layout<row> and "
1459 "layoutB = #nvvm.mma_layout<col> for shape <")
1460 << mmaShape[0] << ", " << mmaShape[1] << ", " << mmaShape[2]
1461 << "> with element types " << *getMultiplicandAPtxType() << " and "
1462 << *getMultiplicandBPtxType()
1463 << ". Only m8n8k4 with f16 supports other layouts.";
1464 }
1465 }
1466
1467 return success();
1468}
1469
1470MMATypes MmaSpOp::accumPtxType() {
1471 std::optional<mlir::NVVM::MMATypes> val = MmaOp::inferOperandMMAType(
1472 getODSOperands(2).getTypes().front(), /*isAccumulator=*/true);
1473 assert(val.has_value() && "accumulator PTX type should always be inferrable");
1474 return val.value();
1475}
1476
1477MMATypes MmaSpOp::resultPtxType() {
1478 std::optional<mlir::NVVM::MMATypes> val =
1479 MmaOp::inferOperandMMAType(getResult().getType(), /*isAccumulator=*/true);
1480 assert(val.has_value() && "result PTX type should always be inferrable");
1481 return val.value();
1482}
1483
1485MmaSpOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
1486 llvm::IRBuilderBase &builder) {
1487 auto thisOp = cast<NVVM::MmaSpOp>(op);
1488
1489 // Get operands
1491 for (mlir::Value v : thisOp.getOperands())
1492 args.push_back(mt.lookupValue(v));
1493
1494 // Get intrinsic ID using the existing getIntrinsicID method
1495 auto intId = MmaSpOp::getIntrinsicID(
1496 thisOp.getShape().getM(), thisOp.getShape().getN(),
1497 thisOp.getShape().getK(), thisOp.getIntOverflowBehavior(),
1498 thisOp.getOrderedMetadata(), thisOp.getKind(),
1499 *thisOp.getMultiplicandAPtxType(), *thisOp.getMultiplicandBPtxType(),
1500 thisOp.accumPtxType(), thisOp.resultPtxType());
1501
1502 return {intId, args};
1503}
1504
1505void MmaSpOp::print(OpAsmPrinter &p) {
1506 SmallVector<Type, 4> regTypes;
1507 struct MMAOperandFragment {
1508 StringRef operandName;
1509 StringRef ptxTypeAttr;
1510 SmallVector<Value, 4> regs;
1511 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1512 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1513 };
1514
1515 std::array<MMAOperandFragment, 5> frags{
1516 MMAOperandFragment("A", getMultiplicandAPtxTypeAttrName()),
1517 MMAOperandFragment("B", getMultiplicandBPtxTypeAttrName()),
1518 MMAOperandFragment("C", ""), MMAOperandFragment("sparseMetadata", ""),
1519 MMAOperandFragment("selector", "")};
1520 SmallVector<StringRef, 4> ignoreAttrNames{
1521 mlir::NVVM::MmaSpOp::getOperandSegmentSizeAttr()};
1522
1523 // Handle variadic operands A, B, C
1524 for (unsigned fragIdx = 0; fragIdx < 3; fragIdx++) {
1525 auto &frag = frags[fragIdx];
1526 auto varOperandSpec = getODSOperandIndexAndLength(fragIdx);
1527 for (auto operandIdx = varOperandSpec.first;
1528 operandIdx < varOperandSpec.first + varOperandSpec.second;
1529 operandIdx++) {
1530 frag.regs.push_back(this->getOperand(operandIdx));
1531 if (operandIdx == varOperandSpec.first) {
1532 regTypes.push_back(this->getOperand(operandIdx).getType());
1533 }
1534 }
1535 std::optional<MMATypes> inferredType = MmaOp::inferOperandMMAType(
1536 regTypes.back(), /*isAccumulator=*/fragIdx >= 2);
1537 if (inferredType)
1538 ignoreAttrNames.push_back(frag.ptxTypeAttr);
1539 }
1540
1541 // Handle sparse metadata and selector (single operands)
1542 frags[3].regs.push_back(getSparseMetadata());
1543 frags[4].regs.push_back(getSparsitySelector());
1544
1545 auto printMmaSpOperand = [&](const MMAOperandFragment &frag) -> void {
1546 p << " " << frag.operandName;
1547 p << "[";
1548 p.printOperands(frag.regs);
1549 p << "]";
1550 };
1551
1552 for (const auto &frag : frags)
1553 printMmaSpOperand(frag);
1554
1555 bool isFirstProperty = true;
1556 printMmaProperty(p, isFirstProperty, "shape", getShapeAttr());
1557 if (getIntOverflowBehaviorAttr())
1558 printMmaEnumProperty(p, isFirstProperty, "int_overflow",
1559 getIntOverflowBehaviorAttr());
1560 if (getMultiplicandAPtxTypeAttr() &&
1561 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
1562 printMmaEnumProperty(p, isFirstProperty, "multiplicand_a_ptx_type",
1563 getMultiplicandAPtxTypeAttr());
1564 if (getMultiplicandBPtxTypeAttr() &&
1565 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
1566 printMmaEnumProperty(p, isFirstProperty, "multiplicand_b_ptx_type",
1567 getMultiplicandBPtxTypeAttr());
1568 if (getOrderedMetadata())
1569 printMmaUnitProperty(p, isFirstProperty, "ordered_metadata");
1570 if (getKindAttr())
1571 printMmaEnumProperty(p, isFirstProperty, "kind", getKindAttr());
1572 llvm::append_range(
1573 ignoreAttrNames,
1574 ArrayRef<StringRef>{getShapeAttrName(), getIntOverflowBehaviorAttrName(),
1575 getMultiplicandAPtxTypeAttrName(),
1576 getMultiplicandBPtxTypeAttrName(),
1577 getOrderedMetadataAttrName(), getKindAttrName()});
1578 p.printOptionalAttrDict((*this)->getAttrs(), ignoreAttrNames);
1579 p << " : ";
1580 p << "(";
1581 for (int i = 0; i < 3; ++i) {
1582 p << regTypes[i];
1583 if (i < 2)
1584 p << ", ";
1585 }
1586 p << ") -> " << getResult().getType();
1587}
1588
1589void MmaSpOp::build(
1590 OpBuilder &builder, OperationState &result, Type resultType,
1591 ValueRange operandA, ValueRange operandB, ValueRange operandC,
1592 Value sparseMetadata, Value sparsitySelector, ArrayRef<int64_t> shape,
1593 std::optional<MMAIntOverflow> intOverflow,
1594 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
1595
1596 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");
1597 MLIRContext *ctx = builder.getContext();
1598 result.addAttribute(
1599 "shape", builder.getAttr<MMAShapeAttr>(shape[0], shape[1], shape[2]));
1600
1601 result.addOperands(operandA);
1602 result.addOperands(operandB);
1603 result.addOperands(operandC);
1604 result.addOperands(sparseMetadata);
1605 result.addOperands(sparsitySelector);
1606
1607 if (multiplicandPtxTypes) {
1608 result.addAttribute("multiplicandAPtxType",
1609 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
1610 result.addAttribute("multiplicandBPtxType",
1611 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
1612 } else {
1613 if (auto res = MmaOp::inferOperandMMAType(operandA[0].getType(), false))
1614 result.addAttribute("multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
1615 if (auto res = MmaOp::inferOperandMMAType(operandB[0].getType(), false))
1616 result.addAttribute("multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
1617 }
1618
1619 if (intOverflow.has_value())
1620 result.addAttribute("intOverflowBehavior",
1621 MMAIntOverflowAttr::get(ctx, *intOverflow));
1622
1623 result.addTypes(resultType);
1624 result.addAttribute(
1625 MmaSpOp::getOperandSegmentSizeAttr(),
1626 builder.getDenseI32ArrayAttr({static_cast<int32_t>(operandA.size()),
1627 static_cast<int32_t>(operandB.size()),
1628 static_cast<int32_t>(operandC.size()), 1,
1629 1})); // sparseMetadata and sparsitySelector
1630}
1631
1632ParseResult MmaSpOp::parse(OpAsmParser &parser, OperationState &result) {
1633 struct MMAOperandFragment {
1634 std::optional<MMATypes> elemtype;
1635 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
1636 SmallVector<Type> regTypes;
1637 };
1638
1639 Builder &builder = parser.getBuilder();
1640 std::array<MMAOperandFragment, 6> frags; // A, B, C, sparseMetadata, selector
1641
1642 NamedAttrList namedAttributes;
1643
1644 // A helper to parse the operand segments.
1645 auto parseMmaSpOperand = [&](StringRef operandName,
1646 MMAOperandFragment &frag) -> LogicalResult {
1647 if (parser.parseKeyword(operandName).failed())
1648 return failure();
1649 if (parser
1650 .parseOperandList(frag.regs, OpAsmParser::Delimiter::OptionalSquare)
1651 .failed())
1652 return failure();
1653 return success();
1654 };
1655
1656 // Parse the operand segments.
1657 if (parseMmaSpOperand("A", frags[0]).failed())
1658 return failure();
1659 if (parseMmaSpOperand("B", frags[1]).failed())
1660 return failure();
1661 if (parseMmaSpOperand("C", frags[2]).failed())
1662 return failure();
1663 if (parseMmaSpOperand("sparseMetadata", frags[3]).failed())
1664 return failure();
1665 if (parseMmaSpOperand("selector", frags[4]).failed())
1666 return failure();
1667
1668 if (parseMmaProperties(parser, namedAttributes,
1669 {"shape", "int_overflow", "multiplicand_a_ptx_type",
1670 "multiplicand_b_ptx_type", "ordered_metadata",
1671 "kind"},
1672 {"shape"}))
1673 return failure();
1674
1675 // Parse the type specification and resolve operands.
1676 SmallVector<Type, 3> operandTypes;
1677 if (failed(parser.parseColon()))
1678 return failure();
1679 if (failed(parser.parseLParen()))
1680 return failure();
1681 if (failed(parser.parseTypeList(operandTypes)))
1682 return failure();
1683 if (failed(parser.parseRParen()))
1684 return failure();
1685 if (operandTypes.size() != 3)
1686 return parser.emitError(
1687 parser.getNameLoc(),
1688 "expected one type for each operand segment but got " +
1689 Twine(operandTypes.size()) + " types");
1690 for (const auto &iter : llvm::enumerate(operandTypes)) {
1691 auto &frag = frags[iter.index()];
1692 frag.regTypes.resize(frag.regs.size(), iter.value());
1693 if (failed(parser.resolveOperands(frag.regs, frag.regTypes,
1694 parser.getNameLoc(), result.operands)))
1695 return failure();
1696 frag.elemtype =
1697 MmaOp::inferOperandMMAType(frag.regTypes[0],
1698 /*isAccumulator*/ iter.index() >= 2);
1699 }
1700
1701 Type resultType;
1702 if (parser.parseArrow() || parser.parseType(resultType))
1703 return failure();
1704 frags[5].elemtype =
1705 MmaOp::inferOperandMMAType(resultType, /*isAccumulator*/ true);
1706
1707 // Resolve sparse metadata and selector (assume i32 type)
1708 Type i32Type = builder.getIntegerType(32);
1709 if (parser
1710 .resolveOperands(frags[3].regs, i32Type, parser.getCurrentLocation(),
1711 result.operands)
1712 .failed())
1713 return failure();
1714 if (parser
1715 .resolveOperands(frags[4].regs, i32Type, parser.getCurrentLocation(),
1716 result.operands)
1717 .failed())
1718 return failure();
1719
1720 std::array<StringRef, 2> names{"multiplicandAPtxType",
1721 "multiplicandBPtxType"};
1722 for (unsigned idx = 0; idx < names.size(); idx++) {
1723 const auto &frag = frags[idx];
1724 std::optional<NamedAttribute> attr = namedAttributes.getNamed(names[idx]);
1725 if (!frag.elemtype.has_value() && !attr.has_value()) {
1726 return parser.emitError(
1727 parser.getNameLoc(),
1728 "attribute " + names[idx] +
1729 " is not provided explicitly and cannot be inferred");
1730 }
1731 if (!attr.has_value())
1732 result.addAttribute(
1733 names[idx], MMATypesAttr::get(parser.getContext(), *frag.elemtype));
1734 }
1735
1736 result.addTypes(resultType);
1737 if (!namedAttributes.empty())
1738 result.addAttributes(namedAttributes);
1739 result.addAttribute(MmaSpOp::getOperandSegmentSizeAttr(),
1740 builder.getDenseI32ArrayAttr({
1741 static_cast<int32_t>(frags[0].regs.size()),
1742 static_cast<int32_t>(frags[1].regs.size()),
1743 static_cast<int32_t>(frags[2].regs.size()),
1744 1, // sparseMetadata
1745 1 // sparsitySelector
1746 }));
1747 return success();
1748}
1749
1750LogicalResult MmaSpOp::verify() {
1751 MLIRContext *context = getContext();
1752 auto f16Ty = Float16Type::get(context);
1753 auto i32Ty = IntegerType::get(context, 32);
1754 auto f16x2Ty = VectorType::get(2, f16Ty);
1755 auto f32Ty = Float32Type::get(context);
1756 auto f16x2x4StructTy = LLVM::LLVMStructType::getLiteral(
1757 context, {f16x2Ty, f16x2Ty, f16x2Ty, f16x2Ty});
1758
1759 auto s32x4StructTy =
1760 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty, i32Ty, i32Ty});
1761 auto f32x8StructTy =
1762 LLVM::LLVMStructType::getLiteral(context, SmallVector<Type>(8, f32Ty));
1763 auto f16x2x2StructTy =
1764 LLVM::LLVMStructType::getLiteral(context, {f16x2Ty, f16x2Ty});
1765 auto f32x4StructTy =
1766 LLVM::LLVMStructType::getLiteral(context, {f32Ty, f32Ty, f32Ty, f32Ty});
1767 auto s32x2StructTy =
1768 LLVM::LLVMStructType::getLiteral(context, {i32Ty, i32Ty});
1769
1770 std::array<int64_t, 3> mmaShape{getShapeAttr().getM(), getShapeAttr().getN(),
1771 getShapeAttr().getK()};
1772
1773 // These variables define the set of allowed data types for matrices A, B, C,
1774 // and result.
1775 using AllowedShapes = SmallVector<std::array<int64_t, 3>, 2>;
1776 using AllowedTypes = SmallVector<SmallVector<Type, 4>, 2>;
1777 AllowedShapes allowedShapes;
1778 AllowedTypes expectedA;
1779 AllowedTypes expectedB;
1780 AllowedTypes expectedC;
1781 SmallVector<Type> expectedResult;
1782
1783 // When M = 16, we just need to calculate the number of 8xk tiles, where
1784 // k is a factor that depends on the data type.
1785 if (mmaShape[0] == 16) {
1786 int64_t kFactor;
1787 Type multiplicandFragType;
1788 switch (*getMultiplicandAPtxType()) {
1789 case MMATypes::tf32:
1790 kFactor = 4;
1791 multiplicandFragType = i32Ty;
1792 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1793 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1794 // Sparse MMA supports m16n8k8 and m16n8k16 for tf32
1795 allowedShapes.push_back({16, 8, 8});
1796 allowedShapes.push_back({16, 8, 16});
1797 break;
1798 case MMATypes::bf16:
1799 kFactor = 8;
1800 multiplicandFragType = i32Ty;
1801 expectedResult.push_back(LLVM::LLVMStructType::getLiteral(
1802 context, {f32Ty, f32Ty, f32Ty, f32Ty}));
1803 // Sparse MMA supports m16n8k16 and m16n8k32 for bf16
1804 allowedShapes.push_back({16, 8, 16});
1805 allowedShapes.push_back({16, 8, 32});
1806 break;
1807 case MMATypes::f16:
1808 kFactor = 8;
1809 multiplicandFragType = f16x2Ty;
1810 expectedResult.push_back(f16x2x2StructTy);
1811 expectedResult.push_back(f32x4StructTy);
1812 // Sparse MMA supports m16n8k16 and m16n8k32 for f16
1813 allowedShapes.push_back({16, 8, 16});
1814 allowedShapes.push_back({16, 8, 32});
1815 break;
1816 case MMATypes::s4:
1817 case MMATypes::u4:
1818 kFactor = 32;
1819 // Sparse MMA supports m16n8k64 and m16n8k128 for s4/u4
1820 allowedShapes.push_back({16, 8, 64});
1821 allowedShapes.push_back({16, 8, 128});
1822 break;
1823 case MMATypes::s8:
1824 case MMATypes::u8:
1825 kFactor = 16;
1826 // Sparse MMA supports m16n8k32 and m16n8k64 for s8/u8
1827 allowedShapes.push_back({16, 8, 32});
1828 allowedShapes.push_back({16, 8, 64});
1829 break;
1830 case MMATypes::e4m3:
1831 case MMATypes::e5m2:
1832 case MMATypes::e3m2:
1833 case MMATypes::e2m3:
1834 case MMATypes::e2m1:
1835 kFactor = 16;
1836 multiplicandFragType = i32Ty;
1837 expectedResult.push_back(f16x2x2StructTy);
1838 expectedResult.push_back(f32x4StructTy);
1839 // Sparse MMA supports m16n8k64 for FP8 types
1840 allowedShapes.push_back({16, 8, 64});
1841 break;
1842 default:
1843 return emitError("invalid shape or multiplicand type: ")
1844 << getMultiplicandAPtxType().value();
1845 }
1846
1847 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {
1848 expectedResult.push_back(s32x4StructTy);
1849 expectedC.emplace_back(4, i32Ty);
1850 multiplicandFragType = i32Ty;
1851 } else if (*getMultiplicandAPtxType() >= MMATypes::e4m3 &&
1852 *getMultiplicandAPtxType() <= MMATypes::e2m1) {
1853 // FP8 types
1854 expectedC.emplace_back(2, f16x2Ty);
1855 expectedC.emplace_back(4, f32Ty);
1856 } else {
1857 expectedC.emplace_back(2, f16x2Ty);
1858 expectedC.emplace_back(4, f32Ty);
1859 }
1860
1861 // For sparse MMA, A operand is compressed (2:4 sparsity means half the
1862 // elements)
1863 int64_t unitA = (mmaShape[0] / 8) * (mmaShape[2] / kFactor) / 2;
1864 int64_t unitB = (mmaShape[1] / 8) * (mmaShape[2] / kFactor);
1865 expectedA.emplace_back(unitA, multiplicandFragType);
1866 expectedB.emplace_back(unitB, multiplicandFragType);
1867
1868 if (resultPtxType() != accumPtxType())
1869 return emitOpError("ctype does not match dtype");
1870 }
1871
1872 // In the M=8 case, there is only 1 possible case per data type.
1873 if (mmaShape[0] == 8) {
1874 if (*getMultiplicandAPtxType() == MMATypes::f16) {
1875 expectedA.emplace_back(2, f16x2Ty);
1876 expectedB.emplace_back(2, f16x2Ty);
1877 expectedResult.push_back(f16x2x4StructTy);
1878 expectedResult.push_back(f32x8StructTy);
1879 expectedC.emplace_back(4, f16x2Ty);
1880 expectedC.emplace_back(8, f32Ty);
1881 allowedShapes.push_back({8, 8, 4});
1882 }
1883 if (*getMultiplicandAPtxType() == MMATypes::f64) {
1884 Type f64Ty = Float64Type::get(context);
1885 expectedA.emplace_back(1, f64Ty);
1886 expectedB.emplace_back(1, f64Ty);
1887 expectedC.emplace_back(2, f64Ty);
1888 expectedResult.emplace_back(LLVM::LLVMStructType::getLiteral(
1889 context, SmallVector<Type>(2, f64Ty)));
1890 allowedShapes.push_back({8, 8, 4});
1891 }
1892 if (isIntegerPtxType(getMultiplicandAPtxType().value())) {
1893 expectedA.push_back({i32Ty});
1894 expectedB.push_back({i32Ty});
1895 expectedC.push_back({i32Ty, i32Ty});
1896 expectedResult.push_back(s32x2StructTy);
1897 if (isInt4PtxType(getMultiplicandAPtxType().value()))
1898 allowedShapes.push_back({8, 8, 32});
1899 if (isInt8PtxType(getMultiplicandAPtxType().value()))
1900 allowedShapes.push_back({8, 8, 16});
1901 }
1902 }
1903
1904 std::string errorMessage;
1905 llvm::raw_string_ostream errorStream(errorMessage);
1906
1907 // Check that we matched an existing shape/dtype combination.
1908 if (expectedA.empty() || expectedB.empty() || expectedC.empty() ||
1909 !llvm::is_contained(allowedShapes, mmaShape)) {
1910 errorStream << "unimplemented variant for MMA shape <";
1911 llvm::interleaveComma(mmaShape, errorStream);
1912 errorStream << ">";
1913 return emitOpError(errorMessage);
1914 }
1915
1916 // Verify the operand types for segments of A, B, and C operands.
1917 std::array<StringRef, 3> operandNames{"A", "B", "C"};
1918 for (const auto &iter : llvm::enumerate(
1919 SmallVector<AllowedTypes, 3>{expectedA, expectedB, expectedC})) {
1920 auto spec = this->getODSOperandIndexAndLength(iter.index());
1921 SmallVector<Type, 4> operandTySeg(operand_type_begin() + spec.first,
1922 operand_type_begin() + spec.first +
1923 spec.second);
1924 bool match = llvm::is_contained(iter.value(), operandTySeg);
1925
1926 if (!match) {
1927 errorStream << "Could not match types for the "
1928 << operandNames[iter.index()]
1929 << " operands; expected one of ";
1930 for (const auto &x : iter.value()) {
1931 errorStream << x.size() << "x" << x[0] << " ";
1932 }
1933 errorStream << "but got ";
1934 llvm::interleaveComma(operandTySeg, errorStream);
1935 return emitOpError(errorMessage);
1936 }
1937 }
1938
1939 // Check the result type
1940 if (!llvm::any_of(expectedResult, [&](Type expectedResultType) {
1941 return expectedResultType == getResult().getType();
1942 })) {
1943 errorStream
1944 << "Could not match allowed types for the result; expected one of ";
1945 llvm::interleaveComma(expectedResult, errorStream);
1946 errorStream << " but got " << getResult().getType();
1947 return emitOpError(errorMessage);
1948 }
1949
1950 // Ensure int4/int8 MMA variants specify the accum overflow behavior
1951 // attribute.
1952 if (isInt4PtxType(*getMultiplicandAPtxType()) ||
1953 isInt8PtxType(*getMultiplicandAPtxType())) {
1954 if (!getIntOverflowBehavior())
1955 return emitOpError("op requires " +
1956 getIntOverflowBehaviorAttrName().strref() +
1957 " attribute");
1958 }
1959
1960 // Validate sparse metadata type (should be i32)
1961 if (!getSparseMetadata().getType().isInteger(32)) {
1962 return emitOpError() << "sparse metadata must be i32 type";
1963 }
1964
1965 // Validate sparsity selector type (should be i32)
1966 if (!getSparsitySelector().getType().isInteger(32)) {
1967 return emitOpError() << "sparsity selector must be i32 type";
1968 }
1969
1970 return success();
1971}
1972
1973//===----------------------------------------------------------------------===//
1974// MMA Block Scale Operations - Shared Helpers
1975//===----------------------------------------------------------------------===//
1976
1977namespace {
1978// Shared structure for MMA operand fragments (A, B, C)
1979struct MMAOperandFragment {
1980 StringRef operandName;
1981 StringRef ptxTypeAttr;
1982 SmallVector<Value, 4> regs;
1983 explicit MMAOperandFragment(StringRef name, StringRef ptxTypeName)
1984 : operandName(name), ptxTypeAttr(ptxTypeName) {}
1985};
1986} // namespace
1987
1988// Helper to print operand list in the format: name[operands]
1989static void printOperandList(OpAsmPrinter &p, StringRef name,
1990 ArrayRef<Value> operands) {
1991 p << " " << name << "[";
1992 p.printOperands(operands);
1993 p << "]";
1994}
1995
1996// Helper to parse operand list in the format: name[operands]
1997static LogicalResult
1998parseMmaOperand(OpAsmParser &parser, StringRef operandName,
2000 if (parser.parseKeyword(operandName).failed())
2001 return failure();
2003 .failed())
2004 return failure();
2005 return success();
2006}
2007
2008// Helper to process operand fragments and determine which attributes can be
2009// inferred
2010template <typename Op>
2011static void
2012processOperandFragments(Op &op, std::array<MMAOperandFragment, 3> &frags,
2013 SmallVectorImpl<Type> &regTypes,
2014 SmallVectorImpl<StringRef> &ignoreAttrNames) {
2015 for (unsigned fragIdx = 0; fragIdx < frags.size(); fragIdx++) {
2016 auto &frag = frags[fragIdx];
2017 auto varOperandSpec = op.getODSOperandIndexAndLength(fragIdx);
2018 for (auto operandIdx = varOperandSpec.first;
2019 operandIdx < varOperandSpec.first + varOperandSpec.second;
2020 operandIdx++) {
2021 frag.regs.push_back(op.getOperand(operandIdx));
2022 if (fragIdx == 0 && operandIdx == varOperandSpec.first) {
2023 regTypes.push_back(op.getOperand(operandIdx).getType());
2024 }
2025 }
2026 if (fragIdx < 2) {
2027 regTypes.push_back(frag.regs[0].getType());
2028 }
2029 std::optional<MMATypes> inferredType =
2030 MmaOp::inferOperandMMAType(regTypes.back(),
2031 /*isAccumulator=*/fragIdx >= 2);
2032 if (inferredType)
2033 ignoreAttrNames.push_back(frag.ptxTypeAttr);
2034 }
2035}
2036
2037// Helper to parse type signature: (A_type, B_type, C_type)
2038static LogicalResult
2040 SmallVectorImpl<Type> &operandTypes) {
2041 if (parser.parseColon().failed() || parser.parseLParen().failed())
2042 return failure();
2043
2044 auto typeParser = [&]() {
2045 Type ty;
2046 if (parser.parseType(ty).failed())
2047 return failure();
2048 operandTypes.push_back(ty);
2049 return success();
2050 };
2051 if (parser.parseCommaSeparatedList(typeParser))
2052 return failure();
2053
2054 if (operandTypes.size() != 3)
2055 return parser.emitError(parser.getCurrentLocation(),
2056 "expected exactly 3 types");
2057
2058 return parser.parseRParen();
2059}
2060
2061// Helper to infer and set multiplicand PTX type attributes
2062static void
2064 const SmallVectorImpl<Type> &operandTypes) {
2065 if (!attrs.get("multiplicandAPtxType")) {
2066 if (auto inferredType =
2067 MmaOp::inferOperandMMAType(operandTypes[0], false)) {
2068 attrs.set("multiplicandAPtxType", MMATypesAttr::get(ctx, *inferredType));
2069 }
2070 }
2071 if (!attrs.get("multiplicandBPtxType")) {
2072 if (auto inferredType =
2073 MmaOp::inferOperandMMAType(operandTypes[1], false)) {
2074 attrs.set("multiplicandBPtxType", MMATypesAttr::get(ctx, *inferredType));
2075 }
2076 }
2077}
2078
2079// Helper to add common block scale properties
2080template <typename OpType>
2083 ScaleVecSize scaleVecSize,
2084 BlockScaleFormat blockScaleFormat,
2085 MMABlockScaleKind kind) {
2086 MLIRContext *ctx = builder.getContext();
2087 auto &properties = result.getOrAddProperties<typename OpType::Properties>();
2088 properties.setShape(
2089 builder.getAttr<MMAShapeAttr>(shape[0], shape[1], shape[2]));
2090 properties.setScaleVecSize(ScaleVecSizeAttr::get(ctx, scaleVecSize));
2091 properties.setBlockScaleFormat(
2092 BlockScaleFormatAttr::get(ctx, blockScaleFormat));
2093 properties.setKind(MMABlockScaleKindAttr::get(ctx, kind));
2094}
2095
2096// Helper to infer and add multiplicand PTX types to builder
2099 ValueRange operandB,
2100 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes) {
2101 if (multiplicandPtxTypes) {
2102 result.addAttribute("multiplicandAPtxType",
2103 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[0]));
2104 result.addAttribute("multiplicandBPtxType",
2105 MMATypesAttr::get(ctx, (*multiplicandPtxTypes)[1]));
2106 } else {
2107 if (auto res = MmaOp::inferOperandMMAType(operandA[0].getType(), false))
2108 result.addAttribute("multiplicandAPtxType", MMATypesAttr::get(ctx, *res));
2109 if (auto res = MmaOp::inferOperandMMAType(operandB[0].getType(), false))
2110 result.addAttribute("multiplicandBPtxType", MMATypesAttr::get(ctx, *res));
2111 }
2112}
2113
2114// Template helper for common accumPtxType/resultPtxType implementation
2115template <typename OpTy>
2116static MMATypes inferPtxTypeFromResult(OpTy op) {
2117 return *MmaOp::inferOperandMMAType(
2118 cast<LLVM::LLVMStructType>(op.getRes().getType()).getBody()[0],
2119 /*isAccumulator=*/true);
2120}
2121
2122//===----------------------------------------------------------------------===//
2123// MmaBlockScaleOp
2124//===----------------------------------------------------------------------===//
2125
2126void MmaBlockScaleOp::print(OpAsmPrinter &p) {
2127 SmallVector<Type, 4> regTypes;
2128 std::array<MMAOperandFragment, 3> frags{
2129 MMAOperandFragment("A", getMultiplicandAPtxTypeAttrName()),
2130 MMAOperandFragment("B", getMultiplicandBPtxTypeAttrName()),
2131 MMAOperandFragment("C", "")};
2132 SmallVector<StringRef, 4> ignoreAttrNames{
2133 mlir::NVVM::MmaBlockScaleOp::getOperandSegmentSizeAttr()};
2134
2135 processOperandFragments(*this, frags, regTypes, ignoreAttrNames);
2136
2137 // Print A, B, C operands
2138 for (const auto &frag : frags)
2139 printOperandList(p, frag.operandName, frag.regs);
2140
2141 // Print scale operands
2142 printOperandList(p, "scaleA",
2143 {getScaleAData(), getByteIdA(), getThreadIdA()});
2144 printOperandList(p, "scaleB",
2145 {getScaleBData(), getByteIdB(), getThreadIdB()});
2146
2147 bool isFirstProperty = true;
2148 printMmaProperty(p, isFirstProperty, "shape", getShapeAttr());
2149 if (getMultiplicandAPtxTypeAttr() &&
2150 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
2151 printMmaEnumProperty(p, isFirstProperty, "multiplicand_a_ptx_type",
2152 getMultiplicandAPtxTypeAttr());
2153 if (getMultiplicandBPtxTypeAttr() &&
2154 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
2155 printMmaEnumProperty(p, isFirstProperty, "multiplicand_b_ptx_type",
2156 getMultiplicandBPtxTypeAttr());
2157 printMmaEnumProperty(p, isFirstProperty, "scale_vec_size",
2158 getScaleVecSizeAttr());
2159 printMmaEnumProperty(p, isFirstProperty, "block_scale_format",
2160 getBlockScaleFormatAttr());
2161 printMmaEnumProperty(p, isFirstProperty, "kind", getKindAttr());
2162 llvm::append_range(
2163 ignoreAttrNames,
2164 ArrayRef<StringRef>{getShapeAttrName(), getMultiplicandAPtxTypeAttrName(),
2165 getMultiplicandBPtxTypeAttrName(),
2166 getScaleVecSizeAttrName(),
2167 getBlockScaleFormatAttrName(), getKindAttrName()});
2168 p.printOptionalAttrDict(this->getOperation()->getAttrs(), ignoreAttrNames);
2169
2170 // Print type signature
2171 p << " : (";
2172 llvm::interleaveComma(SmallVector<Type, 3>{frags[0].regs[0].getType(),
2173 frags[1].regs[0].getType(),
2174 frags[2].regs[0].getType()},
2175 p);
2176 p << ")";
2177 p.printArrowTypeList(TypeRange{this->getRes().getType()});
2178}
2179
2180ParseResult MmaBlockScaleOp::parse(OpAsmParser &parser,
2182 struct LocalOperandFragment {
2183 std::optional<MMATypes> elemtype;
2184 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
2185 };
2186
2187 Builder &builder = parser.getBuilder();
2188 std::array<LocalOperandFragment, 3> frags;
2189 NamedAttrList namedAttributes;
2190
2191 // Parse A[...] B[...] C[...]
2192 if (parseMmaOperand(parser, "A", frags[0].regs).failed() ||
2193 parseMmaOperand(parser, "B", frags[1].regs).failed() ||
2194 parseMmaOperand(parser, "C", frags[2].regs).failed())
2195 return failure();
2196
2197 // Parse scale operands: scaleA[...] scaleB[...]
2198 SmallVector<OpAsmParser::UnresolvedOperand, 3> scaleAOperands, scaleBOperands;
2199 if (parseMmaOperand(parser, "scaleA", scaleAOperands).failed() ||
2200 parseMmaOperand(parser, "scaleB", scaleBOperands).failed())
2201 return failure();
2202
2203 if (parseMmaProperties(parser, namedAttributes,
2204 {"shape", "multiplicand_a_ptx_type",
2205 "multiplicand_b_ptx_type", "scale_vec_size",
2206 "block_scale_format", "kind"},
2207 {"shape", "scaleVecSize", "blockScaleFormat", "kind"}))
2208 return failure();
2209
2210 // Parse type signature
2211 SmallVector<Type, 3> operandTypes;
2212 if (parseMmaTypeSignature(parser, operandTypes).failed())
2213 return failure();
2214
2215 // Parse result type
2216 SmallVector<Type, 1> resultTypes;
2217 if (parser.parseArrowTypeList(resultTypes).failed())
2218 return failure();
2219
2220 // Infer element types and resolve operands
2221 for (const auto &[idx, frag] : llvm::enumerate(frags)) {
2222 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
2223 /*isAccumulator=*/idx >= 2);
2224 if (parser
2225 .resolveOperands(frag.regs, operandTypes[idx], parser.getNameLoc(),
2226 result.operands)
2227 .failed())
2228 return failure();
2229 }
2230
2231 // Resolve scale operands
2232 SmallVector<Type, 3> scaleTypes = {builder.getI32Type(), builder.getI16Type(),
2233 builder.getI16Type()};
2234 if (parser
2235 .resolveOperands(scaleAOperands, scaleTypes, parser.getNameLoc(),
2236 result.operands)
2237 .failed() ||
2238 parser
2239 .resolveOperands(scaleBOperands, scaleTypes, parser.getNameLoc(),
2240 result.operands)
2241 .failed())
2242 return failure();
2243
2244 // Add attributes
2245 result.addAttributes(namedAttributes);
2246 inferAndSetMultiplicandTypes(parser.getContext(), result.attributes,
2247 operandTypes);
2248
2249 result.addTypes(resultTypes);
2250 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
2251 builder.getDenseI32ArrayAttr({
2252 static_cast<int32_t>(frags[0].regs.size()),
2253 static_cast<int32_t>(frags[1].regs.size()),
2254 static_cast<int32_t>(frags[2].regs.size()),
2255 1, // scaleAData
2256 1, // byteIdA
2257 1, // threadIdA
2258 1, // scaleBData
2259 1, // byteIdB
2260 1 // threadIdB
2261 }));
2262 return success();
2263}
2264
2265void MmaBlockScaleOp::build(
2266 OpBuilder &builder, OperationState &result, Type resultType,
2267 ValueRange operandA, ValueRange operandB, ValueRange operandC,
2268 Value scaleAData, Value byteIdA, Value threadIdA, Value scaleBData,
2269 Value byteIdB, Value threadIdB, ArrayRef<int64_t> shape,
2270 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
2271 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
2272 MMABlockScaleKind kind) {
2273 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");
2274
2276 blockScaleFormat, kind);
2277
2278 result.addOperands(operandA);
2279 result.addOperands(operandB);
2280 result.addOperands(operandC);
2281 result.addOperands(
2282 {scaleAData, byteIdA, threadIdA, scaleBData, byteIdB, threadIdB});
2283
2284 addInferredMultiplicandTypes(builder.getContext(), result, operandA, operandB,
2285 multiplicandPtxTypes);
2286
2287 result.addTypes(resultType);
2288 result.addAttribute(MmaBlockScaleOp::getOperandSegmentSizeAttr(),
2289 builder.getDenseI32ArrayAttr({
2290 static_cast<int32_t>(operandA.size()),
2291 static_cast<int32_t>(operandB.size()),
2292 static_cast<int32_t>(operandC.size()),
2293 1, // scaleAData
2294 1, // byteIdA
2295 1, // threadIdA
2296 1, // scaleBData
2297 1, // byteIdB
2298 1 // threadIdB
2299 }));
2300}
2301
2302NVVM::IDArgPair MmaBlockScaleOp::getIntrinsicIDAndArgs(
2303 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
2304 auto curOp = cast<NVVM::MmaBlockScaleOp>(op);
2305
2307 // Add A, B, C operands
2308 for (Value operand : curOp.getOperandA())
2309 args.push_back(mt.lookupValue(operand));
2310 for (Value operand : curOp.getOperandB())
2311 args.push_back(mt.lookupValue(operand));
2312 for (Value operand : curOp.getOperandC())
2313 args.push_back(mt.lookupValue(operand));
2314
2315 // Add scale operands
2316 args.push_back(mt.lookupValue(curOp.getScaleAData()));
2317 args.push_back(mt.lookupValue(curOp.getByteIdA()));
2318 args.push_back(mt.lookupValue(curOp.getThreadIdA()));
2319 args.push_back(mt.lookupValue(curOp.getScaleBData()));
2320 args.push_back(mt.lookupValue(curOp.getByteIdB()));
2321 args.push_back(mt.lookupValue(curOp.getThreadIdB()));
2322
2323 unsigned intId = MmaBlockScaleOp::getIntrinsicID(
2324 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
2325 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
2326 inferPtxTypeFromResult(curOp), curOp.getScaleVecSize(),
2327 curOp.getBlockScaleFormat(), curOp.getKind());
2328
2329 return {intId, args};
2330}
2331
2332LogicalResult MmaBlockScaleOp::verify() {
2333 LogicalResult result = success();
2334 int m = getShape().getM();
2335 int n = getShape().getN();
2336 int k = getShape().getK();
2337
2338 if (m == 16 && n == 8 && k == 64) {
2339 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
2340 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
2342 "unsupported MMATypes attribute for mma.m16n8k64.(mxf4nvf4|mxf4)");
2343 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
2344 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2346 "unsupported ScaleVecSize attribute for mma.m16n8k64.mxf4");
2347 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2349 "unsupported BlockScaleFormat attribute for mma.m16n8k64.mxf4");
2350 } else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2351 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2352 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2353 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2354 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2355 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2356 result = emitOpError("unsupported ScaleVecSize and BlockScaleFormat "
2357 "attributes for mma.m16n8k64.mxf4nvf4");
2358 } else {
2359 result = emitOpError("unsupported Kind attribute for mma.m16n8k64");
2360 }
2361 } else if (m == 16 && n == 8 && k == 32) {
2362 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2363 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2364 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2365 result =
2366 emitOpError("unsupported Kind, ScaleVecSize and BlockScaleFormat "
2367 "attributes for mma.m16n8k32");
2368 } else {
2369 result = emitOpError("unsupported Geom for mma with block scaling");
2370 }
2371 return result;
2372}
2373
2374//===----------------------------------------------------------------------===//
2375// MmaSpBlockScaleOp
2376//===----------------------------------------------------------------------===//
2377
2378void MmaSpBlockScaleOp::print(OpAsmPrinter &p) {
2379 SmallVector<Type, 4> regTypes;
2380 std::array<MMAOperandFragment, 3> frags{
2381 MMAOperandFragment("A", getMultiplicandAPtxTypeAttrName()),
2382 MMAOperandFragment("B", getMultiplicandBPtxTypeAttrName()),
2383 MMAOperandFragment("C", "")};
2384 SmallVector<StringRef, 4> ignoreAttrNames{
2385 mlir::NVVM::MmaSpBlockScaleOp::getOperandSegmentSizeAttr()};
2386
2387 processOperandFragments(*this, frags, regTypes, ignoreAttrNames);
2388
2389 // Print A, B, C operands
2390 for (const auto &frag : frags)
2391 printOperandList(p, frag.operandName, frag.regs);
2392
2393 // Print sparse-specific operands
2394 printOperandList(p, "sparseMetadata", {getSparseMetadata()});
2395 printOperandList(p, "selector", {getSparsitySelector()});
2396
2397 // Print scale operands
2398 printOperandList(p, "scaleA",
2399 {getScaleAData(), getByteIdA(), getThreadIdA()});
2400 printOperandList(p, "scaleB",
2401 {getScaleBData(), getByteIdB(), getThreadIdB()});
2402
2403 bool isFirstProperty = true;
2404 printMmaProperty(p, isFirstProperty, "shape", getShapeAttr());
2405 if (getMultiplicandAPtxTypeAttr() &&
2406 !llvm::is_contained(ignoreAttrNames, getMultiplicandAPtxTypeAttrName()))
2407 printMmaEnumProperty(p, isFirstProperty, "multiplicand_a_ptx_type",
2408 getMultiplicandAPtxTypeAttr());
2409 if (getMultiplicandBPtxTypeAttr() &&
2410 !llvm::is_contained(ignoreAttrNames, getMultiplicandBPtxTypeAttrName()))
2411 printMmaEnumProperty(p, isFirstProperty, "multiplicand_b_ptx_type",
2412 getMultiplicandBPtxTypeAttr());
2413 printMmaUnitProperty(p, isFirstProperty, "ordered_metadata");
2414 printMmaEnumProperty(p, isFirstProperty, "scale_vec_size",
2415 getScaleVecSizeAttr());
2416 printMmaEnumProperty(p, isFirstProperty, "block_scale_format",
2417 getBlockScaleFormatAttr());
2418 printMmaEnumProperty(p, isFirstProperty, "kind", getKindAttr());
2419 llvm::append_range(
2420 ignoreAttrNames,
2421 ArrayRef<StringRef>{getShapeAttrName(), getMultiplicandAPtxTypeAttrName(),
2422 getMultiplicandBPtxTypeAttrName(),
2423 getOrderedMetadataAttrName(),
2424 getScaleVecSizeAttrName(),
2425 getBlockScaleFormatAttrName(), getKindAttrName()});
2426 p.printOptionalAttrDict(this->getOperation()->getAttrs(), ignoreAttrNames);
2427
2428 // Print type signature
2429 p << " : (";
2430 llvm::interleaveComma(SmallVector<Type, 3>{frags[0].regs[0].getType(),
2431 frags[1].regs[0].getType(),
2432 frags[2].regs[0].getType()},
2433 p);
2434 p << ")";
2435 p.printArrowTypeList(TypeRange{this->getRes().getType()});
2436}
2437
2438ParseResult MmaSpBlockScaleOp::parse(OpAsmParser &parser,
2440 struct LocalOperandFragment {
2441 std::optional<MMATypes> elemtype;
2442 SmallVector<OpAsmParser::UnresolvedOperand, 4> regs;
2443 };
2444
2445 Builder &builder = parser.getBuilder();
2446 std::array<LocalOperandFragment, 3> frags;
2447 NamedAttrList namedAttributes;
2448
2449 // Parse A[...] B[...] C[...]
2450 if (parseMmaOperand(parser, "A", frags[0].regs).failed() ||
2451 parseMmaOperand(parser, "B", frags[1].regs).failed() ||
2452 parseMmaOperand(parser, "C", frags[2].regs).failed())
2453 return failure();
2454
2455 // Parse sparse-specific operands
2457 selectorOperands;
2458 if (parseMmaOperand(parser, "sparseMetadata", metadataOperands).failed() ||
2459 parseMmaOperand(parser, "selector", selectorOperands).failed())
2460 return failure();
2461
2462 // Parse scale operands
2463 SmallVector<OpAsmParser::UnresolvedOperand, 3> scaleAOperands, scaleBOperands;
2464 if (parseMmaOperand(parser, "scaleA", scaleAOperands).failed() ||
2465 parseMmaOperand(parser, "scaleB", scaleBOperands).failed())
2466 return failure();
2467
2468 if (parseMmaProperties(parser, namedAttributes,
2469 {"shape", "multiplicand_a_ptx_type",
2470 "multiplicand_b_ptx_type", "ordered_metadata",
2471 "scale_vec_size", "block_scale_format", "kind"},
2472 {"shape", "scaleVecSize", "blockScaleFormat", "kind"}))
2473 return failure();
2474
2475 // Parse type signature
2476 SmallVector<Type, 3> operandTypes;
2477 if (parseMmaTypeSignature(parser, operandTypes).failed())
2478 return failure();
2479
2480 // Parse result type
2481 SmallVector<Type, 1> resultTypes;
2482 if (parser.parseArrowTypeList(resultTypes).failed())
2483 return failure();
2484
2485 // Infer element types and resolve operands
2486 for (const auto &[idx, frag] : llvm::enumerate(frags)) {
2487 frag.elemtype = MmaOp::inferOperandMMAType(operandTypes[idx],
2488 /*isAccumulator=*/idx >= 2);
2489 if (parser
2490 .resolveOperands(frag.regs, operandTypes[idx], parser.getNameLoc(),
2491 result.operands)
2492 .failed())
2493 return failure();
2494 }
2495
2496 // Resolve sparse metadata and selector
2497 Type i32Type = builder.getI32Type();
2498 if (parser
2499 .resolveOperands(metadataOperands, i32Type, parser.getNameLoc(),
2500 result.operands)
2501 .failed() ||
2502 parser
2503 .resolveOperands(selectorOperands, i32Type, parser.getNameLoc(),
2504 result.operands)
2505 .failed())
2506 return failure();
2507
2508 // Resolve scale operands
2509 SmallVector<Type, 3> scaleTypes = {i32Type, builder.getI16Type(),
2510 builder.getI16Type()};
2511 if (parser
2512 .resolveOperands(scaleAOperands, scaleTypes, parser.getNameLoc(),
2513 result.operands)
2514 .failed() ||
2515 parser
2516 .resolveOperands(scaleBOperands, scaleTypes, parser.getNameLoc(),
2517 result.operands)
2518 .failed())
2519 return failure();
2520
2521 // Add attributes
2522 result.addAttributes(namedAttributes);
2523 inferAndSetMultiplicandTypes(parser.getContext(), result.attributes,
2524 operandTypes);
2525
2526 // orderedMetadata is mandatory
2527 if (!result.attributes.get("orderedMetadata"))
2528 result.addAttribute("orderedMetadata", builder.getUnitAttr());
2529
2530 result.addTypes(resultTypes);
2531 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2532 builder.getDenseI32ArrayAttr({
2533 static_cast<int32_t>(frags[0].regs.size()),
2534 static_cast<int32_t>(frags[1].regs.size()),
2535 static_cast<int32_t>(frags[2].regs.size()),
2536 1, // sparseMetadata
2537 1, // sparsitySelector
2538 1, // scaleAData
2539 1, // byteIdA
2540 1, // threadIdA
2541 1, // scaleBData
2542 1, // byteIdB
2543 1 // threadIdB
2544 }));
2545 return success();
2546}
2547
2548void MmaSpBlockScaleOp::build(
2549 OpBuilder &builder, OperationState &result, Type resultType,
2550 ValueRange operandA, ValueRange operandB, ValueRange operandC,
2551 Value sparseMetadata, Value sparsitySelector, Value scaleAData,
2552 Value byteIdA, Value threadIdA, Value scaleBData, Value byteIdB,
2553 Value threadIdB, ArrayRef<int64_t> shape,
2554 std::optional<std::array<MMATypes, 2>> multiplicandPtxTypes,
2555 ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat,
2556 MMABlockScaleKind kind) {
2557 assert(shape.size() == 3 && "expected shape to have size 3 (m, n, k)");
2558
2560 builder, result, shape, scaleVecSize, blockScaleFormat, kind);
2561 result.addAttribute("orderedMetadata", builder.getUnitAttr());
2562
2563 result.addOperands(operandA);
2564 result.addOperands(operandB);
2565 result.addOperands(operandC);
2566 result.addOperands({sparseMetadata, sparsitySelector, scaleAData, byteIdA,
2567 threadIdA, scaleBData, byteIdB, threadIdB});
2568
2569 addInferredMultiplicandTypes(builder.getContext(), result, operandA, operandB,
2570 multiplicandPtxTypes);
2571
2572 result.addTypes(resultType);
2573 result.addAttribute(MmaSpBlockScaleOp::getOperandSegmentSizeAttr(),
2574 builder.getDenseI32ArrayAttr({
2575 static_cast<int32_t>(operandA.size()),
2576 static_cast<int32_t>(operandB.size()),
2577 static_cast<int32_t>(operandC.size()),
2578 1, // sparseMetadata
2579 1, // sparsitySelector
2580 1, // scaleAData
2581 1, // byteIdA
2582 1, // threadIdA
2583 1, // scaleBData
2584 1, // byteIdB
2585 1 // threadIdB
2586 }));
2587}
2588
2589NVVM::IDArgPair MmaSpBlockScaleOp::getIntrinsicIDAndArgs(
2590 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
2591 auto curOp = cast<NVVM::MmaSpBlockScaleOp>(op);
2592
2594 // Add A, B, C operands
2595 for (Value operand : curOp.getOperandA())
2596 args.push_back(mt.lookupValue(operand));
2597 for (Value operand : curOp.getOperandB())
2598 args.push_back(mt.lookupValue(operand));
2599 for (Value operand : curOp.getOperandC())
2600 args.push_back(mt.lookupValue(operand));
2601
2602 // Add sparse metadata and selector
2603 args.push_back(mt.lookupValue(curOp.getSparseMetadata()));
2604 args.push_back(mt.lookupValue(curOp.getSparsitySelector()));
2605
2606 // Add scale operands
2607 args.push_back(mt.lookupValue(curOp.getScaleAData()));
2608 args.push_back(mt.lookupValue(curOp.getByteIdA()));
2609 args.push_back(mt.lookupValue(curOp.getThreadIdA()));
2610 args.push_back(mt.lookupValue(curOp.getScaleBData()));
2611 args.push_back(mt.lookupValue(curOp.getByteIdB()));
2612 args.push_back(mt.lookupValue(curOp.getThreadIdB()));
2613
2614 unsigned intId = MmaSpBlockScaleOp::getIntrinsicID(
2615 curOp.getShape().getM(), curOp.getShape().getN(), curOp.getShape().getK(),
2616 *curOp.getMultiplicandAPtxType(), *curOp.getMultiplicandBPtxType(),
2617 inferPtxTypeFromResult(curOp), curOp.getScaleVecSize(),
2618 curOp.getBlockScaleFormat(), curOp.getKind());
2619
2620 return {intId, args};
2621}
2622
2623LogicalResult MmaSpBlockScaleOp::verify() {
2624 // Check that orderedMetadata is present
2625 if (!getOrderedMetadata()) {
2626 return emitOpError("'orderedMetadata' attribute is mandatory");
2627 }
2628
2629 LogicalResult result = success();
2630 int m = getShape().getM();
2631 int n = getShape().getN();
2632 int k = getShape().getK();
2633
2634 if (m == 16 && n == 8 && k == 128) {
2635 if (getMultiplicandAPtxType() != NVVM::MMATypes::e2m1 ||
2636 getMultiplicandBPtxType() != NVVM::MMATypes::e2m1)
2638 "unsupported MMATypes attribute for mma.m16n8k128.(mxf4nvf4|mxf4)");
2639 if (getKind() == NVVM::MMABlockScaleKind::MXF4) {
2640 if (getScaleVecSize() != NVVM::ScaleVecSize::X2)
2642 "unsupported ScaleVecSize attribute for mma.m16n8k128.mxf4");
2643 if (getBlockScaleFormat() != NVVM::BlockScaleFormat::UE8M0)
2645 "unsupported BlockScaleFormat attribute for mma.m16n8k128.mxf4");
2646 } else if (getKind() == NVVM::MMABlockScaleKind::MXF4NVF4) {
2647 if (!((getScaleVecSize() == NVVM::ScaleVecSize::X2 &&
2648 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0) ||
2649 (getScaleVecSize() == NVVM::ScaleVecSize::X4 &&
2650 (getBlockScaleFormat() == NVVM::BlockScaleFormat::UE4M3 ||
2651 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))))
2652 result = emitOpError("unsupported ScaleVecSize and BlockScaleFormat "
2653 "attributes for mma.m16n8k128.mxf4nvf4");
2654 } else {
2655 result = emitOpError("unsupported Kind attribute for mma.m16n8k128");
2656 }
2657 } else if (m == 16 && n == 8 && k == 64) {
2658 if (!(getKind() == NVVM::MMABlockScaleKind::MXF8F6F4 &&
2659 getScaleVecSize() == NVVM::ScaleVecSize::X1 &&
2660 getBlockScaleFormat() == NVVM::BlockScaleFormat::UE8M0))
2661 result =
2662 emitOpError("unsupported Kind, ScaleVecSize and BlockScaleFormat "
2663 "attributes for mma.m16n8k64");
2664 } else {
2665 result = emitOpError("unsupported Geom for sparse mma with block scaling");
2666 }
2667 return result;
2668}
2669
2670LogicalResult ShflOp::verify() {
2671 auto returnStructType = llvm::dyn_cast<LLVM::LLVMStructType>(getType());
2672
2673 auto verifyTypeError = [&](Twine desc, Type expectedType,
2674 Type actualType) -> LogicalResult {
2675 return emitOpError("expected " + desc + " to be of type ")
2676 << expectedType << " but got " << actualType << " instead";
2677 };
2678
2679 if (returnStructType) {
2680 if (!getReturnValueAndIsValid())
2681 return emitOpError("\"return_value_and_is_valid\" attribute must be "
2682 "specified when the return type is a struct type");
2683
2684 if (returnStructType.getBody().size() != 2)
2685 return emitOpError("expected return type to be a two-element struct");
2686
2687 llvm::ArrayRef<Type> returnStruct = returnStructType.getBody();
2688 auto resultType = returnStruct[0];
2689 if (resultType != getVal().getType())
2690 return verifyTypeError("first element in the returned struct",
2691 getVal().getType(), resultType);
2692
2693 auto predicateType = returnStruct[1];
2694 if (!predicateType.isInteger(1))
2695 return verifyTypeError("second element in the returned struct",
2696 mlir::IntegerType::get(getContext(), 1),
2697 predicateType);
2698 } else {
2699 if (getReturnValueAndIsValid())
2700 return emitOpError("expected return type to be a two-element struct");
2701
2702 if (getType() != getVal().getType())
2703 return verifyTypeError("return type", getVal().getType(), getType());
2704 }
2705 return success();
2706}
2707
2708LogicalResult
2709ShflOp::inferReturnTypes(MLIRContext *context, std::optional<Location> location,
2710 ShflOp::Adaptor adaptor,
2711 SmallVectorImpl<Type> &inferredReturnTypes) {
2712 Type valType = adaptor.getVal().getType();
2713 if (adaptor.getReturnValueAndIsValid())
2714 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2715 context, {valType, IntegerType::get(context, 1)}));
2716 else
2717 inferredReturnTypes.push_back(valType);
2718 return success();
2719}
2720
2721std::pair<mlir::Type, unsigned> NVVM::inferMMAType(NVVM::MMATypes type,
2722 NVVM::MMAFrag frag, int nRow,
2723 int nCol,
2724 MLIRContext *context) {
2725 unsigned numberElements = 0;
2726 Type elementType;
2727 OpBuilder builder(context);
2728 Type f16x2 = VectorType::get(2, builder.getF16Type());
2729 if (type == NVVM::MMATypes::f16) {
2730 elementType = f16x2;
2731 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2732 numberElements = 8;
2733 else
2734 numberElements = 4;
2735 } else if (type == NVVM::MMATypes::f32) {
2736 elementType = builder.getF32Type();
2737 numberElements = 8;
2738 } else if (type == NVVM::MMATypes::f64) {
2739 elementType = builder.getF64Type();
2740 if (frag == NVVM::MMAFrag::a || frag == NVVM::MMAFrag::b)
2741 numberElements = 1;
2742 else
2743 numberElements = 2;
2744 } else if (type == NVVM::MMATypes::tf32) {
2745 elementType = builder.getI32Type();
2746 numberElements = 4;
2747 } else if (type == NVVM::MMATypes::s8 || type == NVVM::MMATypes::u8) {
2748 elementType = builder.getI32Type();
2749 int parallelSize = 0;
2750 if (frag == NVVM::MMAFrag::a)
2751 parallelSize = nRow;
2752 if (frag == NVVM::MMAFrag::b)
2753 parallelSize = nCol;
2754
2755 // m == 16 && n == 16 && k == 16
2756 if (parallelSize == 16)
2757 numberElements = 2;
2758 // m == 8 && n == 32 && k == 16 or m == 32 && n == 8 && k == 16
2759 else if (parallelSize == 8)
2760 numberElements = 1;
2761 else if (parallelSize == 32)
2762 numberElements = 4;
2763 } else if (type == NVVM::MMATypes::s32) {
2764 elementType = builder.getI32Type();
2765 numberElements = 8;
2766 }
2767 assert(numberElements != 0 && elementType != nullptr);
2768 return std::make_pair(elementType, numberElements);
2769}
2770
2771static std::pair<mlir::Type, unsigned>
2772inferMMATypeFromMNK(NVVM::MMATypes type, NVVM::MMAFrag frag, int m, int n,
2773 int k, MLIRContext *context) {
2774 int nRow, nCol;
2775 if (frag == NVVM::MMAFrag::a) {
2776 nRow = m;
2777 nCol = k;
2778 } else if (frag == NVVM::MMAFrag::b) {
2779 nRow = k;
2780 nCol = n;
2781 } else {
2782 nRow = m;
2783 nCol = n;
2784 }
2785 assert(nRow && nCol);
2786 return inferMMAType(type, frag, nRow, nCol, context);
2787}
2788
2789LogicalResult NVVM::WMMALoadOp::verify() {
2790 unsigned addressSpace =
2791 llvm::cast<LLVM::LLVMPointerType>(getPtr().getType()).getAddressSpace();
2792 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2793 addressSpace != NVVMMemorySpace::Shared)
2794 return emitOpError("expected source pointer in memory "
2795 "space 0, 1, 3");
2796
2797 if (NVVM::WMMALoadOp::getIntrinsicID(getM(), getN(), getK(), getLayout(),
2798 getEltype(), getFrag()) == 0)
2799 return emitOpError() << "invalid attribute combination";
2800 std::pair<Type, unsigned> typeInfo = inferMMATypeFromMNK(
2801 getEltype(), getFrag(), getM(), getN(), getK(), getContext());
2802 // Special case for f64 fragments
2803 Type f64Ty = Float64Type::get(getContext());
2804 if (typeInfo.first == f64Ty && typeInfo.second == 1) {
2805 if (getType() != f64Ty)
2806 return emitOpError("expected destination type to be f64");
2807 return success();
2808 }
2809 // Everything else is a struct
2810 Type dstType = LLVM::LLVMStructType::getLiteral(
2811 getContext(), SmallVector<Type, 8>(typeInfo.second, typeInfo.first));
2812 if (getType() != dstType)
2813 return emitOpError("expected destination type is a structure of ")
2814 << typeInfo.second << " elements of type " << typeInfo.first;
2815 return success();
2816}
2817
2818LogicalResult NVVM::WMMAStoreOp::verify() {
2819 unsigned addressSpace =
2820 llvm::cast<LLVM::LLVMPointerType>(getPtr().getType()).getAddressSpace();
2821 if (addressSpace != 0 && addressSpace != NVVMMemorySpace::Global &&
2822 addressSpace != NVVMMemorySpace::Shared)
2823 return emitOpError("expected operands to be a source pointer in memory "
2824 "space 0, 1, 3");
2825
2826 if (NVVM::WMMAStoreOp::getIntrinsicID(getM(), getN(), getK(), getLayout(),
2827 getEltype()) == 0)
2828 return emitOpError() << "invalid attribute combination";
2829 std::pair<Type, unsigned> typeInfo = inferMMATypeFromMNK(
2830 getEltype(), NVVM::MMAFrag::c, getM(), getN(), getK(), getContext());
2831 if (getArgs().size() != typeInfo.second)
2832 return emitOpError() << "expected " << typeInfo.second << " data operands";
2833 if (llvm::any_of(getArgs(), [&typeInfo](Value operands) {
2834 return operands.getType() != typeInfo.first;
2835 }))
2836 return emitOpError() << "expected data operands of type " << typeInfo.first;
2837 return success();
2838}
2839
2840LogicalResult NVVM::WMMAMmaOp::verify() {
2841 if (NVVM::WMMAMmaOp::getIntrinsicID(getM(), getN(), getK(), getLayoutA(),
2842 getLayoutB(), getEltypeA(),
2843 getEltypeB()) == 0)
2844 return emitOpError() << "invalid attribute combination";
2845 std::pair<Type, unsigned> typeInfoA = inferMMATypeFromMNK(
2846 getEltypeA(), NVVM::MMAFrag::a, getM(), getN(), getK(), getContext());
2847 std::pair<Type, unsigned> typeInfoB = inferMMATypeFromMNK(
2848 getEltypeA(), NVVM::MMAFrag::b, getM(), getN(), getK(), getContext());
2849 std::pair<Type, unsigned> typeInfoC = inferMMATypeFromMNK(
2850 getEltypeB(), NVVM::MMAFrag::c, getM(), getN(), getK(), getContext());
2851 SmallVector<Type, 32> arguments;
2852 arguments.append(typeInfoA.second, typeInfoA.first);
2853 arguments.append(typeInfoB.second, typeInfoB.first);
2854 arguments.append(typeInfoC.second, typeInfoC.first);
2855 unsigned numArgs = arguments.size();
2856 if (getArgs().size() != numArgs)
2857 return emitOpError() << "expected " << numArgs << " arguments";
2858 for (unsigned i = 0; i < numArgs; i++) {
2859 if (getArgs()[i].getType() != arguments[i])
2860 return emitOpError() << "expected argument " << i << " to be of type "
2861 << arguments[i];
2862 }
2863 Type dstType = LLVM::LLVMStructType::getLiteral(
2864 getContext(), SmallVector<Type, 8>(typeInfoC.second, typeInfoC.first));
2865 if (getType() != dstType)
2866 return emitOpError("expected destination type is a structure of ")
2867 << typeInfoC.second << " elements of type " << typeInfoC.first;
2868 return success();
2869}
2870
2871LogicalResult NVVM::LdMatrixOp::verify() {
2872 uint32_t num = getNum(), m = getShape().getM(), n = getShape().getN();
2873 if (m == 8 && n == 8) {
2874 if (num != 1 && num != 2 && num != 4) {
2875 return emitOpError("expected num attribute to be 1, 2 or 4 for 8x8 "
2876 "matrix");
2877 }
2878 if (getEltType() != LdStMatrixEltType::B16) {
2879 return emitOpError("expected element type to be b16 for 8x8 matrix");
2880 }
2881 } else if (m == 8 && n == 16) {
2882 if (num != 1 && num != 2 && num != 4) {
2883 return emitOpError("expected num attribute to be 1, 2 or 4 for 8x16 "
2884 "matrix");
2885 }
2886 if (getLayout() != MMALayout::row) {
2887 return emitOpError("expected layout to be row for 8x16 matrix");
2888 }
2889 if (getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2890 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2891 return emitOpError("expected element type to be b8x16.b4x16_p64 or "
2892 "b8x16.b6x16_p32 for 8x16 matrix");
2893 }
2894 } else if (m == 16 && n == 16) {
2895 if (num != 1 && num != 2) {
2896 return emitOpError("expected num attribute to be 1 or 2 for 16x16 "
2897 "matrix");
2898 }
2899 if (getLayout() != MMALayout::col) {
2900 return emitOpError("expected layout to be col for 16x16 matrix");
2901 }
2902 if (getEltType() != LdStMatrixEltType::B8 &&
2903 getEltType() != LdStMatrixEltType::B8X16_B4X16_P64 &&
2904 getEltType() != LdStMatrixEltType::B8X16_B6X16_P32) {
2905 return emitOpError("expected element type to be b8, b8x16.b4x16_p64 or "
2906 "b8x16.b6x16_p32 for 16x16 matrix");
2907 }
2908 } else {
2909 return emitOpError("expected shape to be 8x8, 8x16 or 16x16");
2910 }
2911
2912 Type i32 = IntegerType::get(getContext(), 32);
2913 uint32_t numElements = (m == 16 && n == 16 ? num * 2 : num);
2914 if (numElements == 1 && getType() != i32)
2915 return emitOpError("expected destination type is i32");
2916 if (numElements == 2 || numElements == 4) {
2917 Type dstType = LLVM::LLVMStructType::getLiteral(
2918 getContext(), SmallVector<Type>(numElements, i32));
2919 if (getType() != dstType)
2920 return emitOpError("expected destination type is a structure of ")
2921 << numElements << " elements of type i32";
2922 }
2923
2924 return success();
2925}
2926
2927LogicalResult LdMatrixOp::inferReturnTypes(
2928 MLIRContext *context, std::optional<Location> location,
2929 LdMatrixOp::Adaptor adaptor, SmallVectorImpl<Type> &inferredReturnTypes) {
2930 uint32_t num = adaptor.getNum();
2931 uint32_t m = adaptor.getShape().getM();
2932 uint32_t n = adaptor.getShape().getN();
2933 uint32_t numElements = (m == 16 && n == 16) ? num * 2 : num;
2934
2935 Type i32 = IntegerType::get(context, 32);
2936 if (numElements == 1)
2937 inferredReturnTypes.push_back(i32);
2938 else
2939 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
2940 context, SmallVector<Type>(numElements, i32)));
2941 return success();
2942}
2943
2944LogicalResult NVVM::StMatrixOp::verify() {
2945 int numMatrix = getSources().size();
2946 if (numMatrix != 1 && numMatrix != 2 && numMatrix != 4)
2947 return emitOpError("expected num attribute to be 1, 2 or 4");
2948
2949 int m = getShape().getM(), n = getShape().getN();
2950 if (m == 8 && n == 8) {
2951 if (getEltType() != NVVM::LdStMatrixEltType::B16) {
2952 return emitOpError("expected element type to be B16 for 8x8 matrix");
2953 }
2954 } else if (m == 16 && n == 8) {
2955 if (getEltType() != NVVM::LdStMatrixEltType::B8) {
2956 return emitOpError("expected element type to be B8 for 16x8 matrix");
2957 }
2958 if (getLayout() != NVVM::MMALayout::col) {
2959 return emitOpError("expected layout to be col for 16x8 matrix");
2960 }
2961 } else {
2962 return emitOpError("expected shape to be 8x8 or 16x8");
2963 }
2964
2965 return success();
2966}
2967
2968LogicalResult NVVM::MovMatrixOp::verify() {
2969 int m = getShape().getM(), n = getShape().getN();
2970 if (m != 8 || n != 8)
2971 return emitOpError("expected shape to be 8x8");
2972 if (getLayout() != NVVM::MMALayout::col)
2973 return emitOpError("expected layout to be col");
2974 if (getEltType() != NVVM::LdStMatrixEltType::B16)
2975 return emitOpError("expected element type to be b16");
2976 return success();
2977}
2978
2979static FailureOr<int> getAllowedSizeK(NVVM::WGMMATypes typeA) {
2980 if (typeA == NVVM::WGMMATypes::tf32)
2981 return 8;
2982 if (typeA == NVVM::WGMMATypes::f16 || typeA == NVVM::WGMMATypes::bf16)
2983 return 16;
2984 if (typeA == NVVM::WGMMATypes::s8 || typeA == NVVM::WGMMATypes::u8)
2985 return 32;
2986 if (typeA == NVVM::WGMMATypes::e4m3 || typeA == NVVM::WGMMATypes::e5m2)
2987 return 32;
2988 if (typeA == NVVM::WGMMATypes::b1)
2989 return 256;
2990 return failure();
2991}
2992
2993static LogicalResult isAllowedWGMMADataType(NVVM::WGMMATypes typeD,
2994 NVVM::WGMMATypes typeA,
2995 NVVM::WGMMATypes typeB) {
2996 switch (typeA) {
2997 case NVVM::WGMMATypes::f16:
2998 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
2999 typeB == NVVM::WGMMATypes::f16)
3000 return success();
3001 break;
3002 case NVVM::WGMMATypes::tf32:
3003 if (typeD == NVVM::WGMMATypes::f32 && typeB == NVVM::WGMMATypes::tf32)
3004 return success();
3005 break;
3006 case NVVM::WGMMATypes::u8:
3007 case NVVM::WGMMATypes::s8:
3008 if (typeD == NVVM::WGMMATypes::s32 &&
3009 (typeB == NVVM::WGMMATypes::u8 || typeB == NVVM::WGMMATypes::s8))
3010 return success();
3011 break;
3012 case NVVM::WGMMATypes::b1:
3013 if (typeD == NVVM::WGMMATypes::s32 && typeB == NVVM::WGMMATypes::b1)
3014 return success();
3015 break;
3016 case NVVM::WGMMATypes::bf16:
3017 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
3018 typeB == NVVM::WGMMATypes::bf16)
3019 return success();
3020 break;
3021 case NVVM::WGMMATypes::e4m3:
3022 case NVVM::WGMMATypes::e5m2:
3023 if ((typeD == NVVM::WGMMATypes::f32 || typeD == NVVM::WGMMATypes::f16) &&
3024 (typeB == NVVM::WGMMATypes::e5m2 || typeB == NVVM::WGMMATypes::e4m3))
3025 return success();
3026 break;
3027 case WGMMATypes::f32:
3028 case WGMMATypes::s32:
3029 llvm_unreachable("unsupported input types");
3030 break;
3031 }
3032 return failure();
3033}
3034
3035static LogicalResult isAllowedSizeN(int sizeN, NVVM::WGMMATypes typeA) {
3036 SmallVector<int> allowedN = {8, 16, 24, 32, 40, 48, 56, 64,
3037 72, 80, 88, 96, 104, 112, 120, 128,
3038 136, 144, 152, 160, 168, 176, 184, 192,
3039 200, 208, 216, 224, 232, 240, 248, 256};
3040 SmallVector<int> allowedNshort = {8, 16, 24, 32, 48, 64,
3041 80, 96, 112, 128, 144, 160,
3042 176, 192, 208, 224, 240, 256};
3043 switch (typeA) {
3044 case WGMMATypes::f16:
3045 case WGMMATypes::tf32:
3046 case WGMMATypes::bf16:
3047 case WGMMATypes::e4m3:
3048 case WGMMATypes::e5m2:
3049 if (llvm::is_contained(allowedN, sizeN))
3050 return success();
3051 break;
3052 case WGMMATypes::u8:
3053 case WGMMATypes::s8:
3054 case WGMMATypes::b1:
3055 if (llvm::is_contained(allowedNshort, sizeN))
3056 return success();
3057 break;
3058 case WGMMATypes::f32:
3059 case WGMMATypes::s32:
3060 llvm_unreachable("unsupported input types");
3061 break;
3062 }
3063 return failure();
3064}
3065
3066LogicalResult NVVM::WgmmaMmaAsyncOp::verify() {
3067 Value outValue = getResults();
3068 auto stype = dyn_cast<LLVM::LLVMStructType>(outValue.getType());
3069 if (!stype)
3070 return emitOpError() << "expected results to be struct";
3071 int outputSize = stype.getBody().size();
3072 WGMMATypes typeD = getTypeD();
3073 WGMMATypes typeA = getTypeA();
3074 WGMMATypes typeB = getTypeB();
3075
3076 for (Type t : stype.getBody()) {
3077 if (t != stype.getBody().front())
3078 return emitOpError()
3079 << "all elements in struct must be same type but there is " << t;
3080 }
3081
3082 if (typeD != WGMMATypes::f32 && typeD != WGMMATypes::f16 &&
3083 typeD != WGMMATypes::s32) {
3084 return emitOpError() << "does not support the given output type " << typeD;
3085 }
3086 if (typeD == WGMMATypes::s32 &&
3087 (getScaleA() == WGMMAScaleIn::neg || getScaleB() == WGMMAScaleIn::neg)) {
3088 return emitOpError() << "has s32 output, scaleA and scaleB cannot be neg";
3089 }
3090
3091 if (failed(isAllowedWGMMADataType(typeD, typeA, typeB))) {
3092 return emitOpError() << typeD << " += " << typeA << " * " << typeB
3093 << ", it is not supported.";
3094 }
3095
3096 // Check M
3097 if (getShape().getM() != 64)
3098 return emitOpError() << "shape 'm' must be 64";
3099
3100 // Check K
3101 FailureOr<int> allowedK = getAllowedSizeK(typeA);
3102 if (failed(allowedK) || allowedK.value() != getShape().getK())
3103 return emitOpError() << "shape 'k' must be " << allowedK.value()
3104 << " for input type " << typeA;
3105
3106 // Check N
3107 if (failed(isAllowedSizeN(getShape().getN(), typeA))) {
3108 return emitOpError() << "has input type " << typeA << " n is set to "
3109 << getShape().getN() << ", it is not supported.";
3110 }
3111
3112 // Check transpose (only available for f16/bf16)
3113 // Matrices A should be stored in row-major and B in column-major.
3114 // Only f16/bf16 matrices can be stored in either column-major or row-major
3115 // by setting the transpose value(imm-trans-a,imm-trans-b) in PTX code.
3116 if ((typeA != WGMMATypes::f16 && typeA != WGMMATypes::bf16) &&
3117 (getLayoutA() == mlir::NVVM::MMALayout::col ||
3118 getLayoutB() == mlir::NVVM::MMALayout::row)) {
3119 return emitOpError()
3120 << "given layouts layout_a = " << getLayoutA()
3121 << " and layout_b = " << getLayoutB() << " for input types " << typeA
3122 << " and " << typeB
3123 << " requires transpose. However, this is only supported for: "
3124 << MMATypes::f16 << " and " << MMATypes::bf16;
3125 }
3126
3127 // Check result registers
3128 int expectedOutput = 0;
3129 if (typeD == WGMMATypes::f32 || typeD == WGMMATypes::s32)
3130 expectedOutput = getShape().getN() / 2;
3131 if (typeD == WGMMATypes::f16)
3132 expectedOutput = getShape().getN() / 4;
3133 if (outputSize != expectedOutput) {
3134 return emitOpError() << "results " << expectedOutput
3135 << ", however output struct has " << outputSize
3136 << " elements";
3137 }
3138 // Check satfinite (only available for s32 accumulator)
3139 if (typeD != WGMMATypes::s32 &&
3140 getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
3141 NVVM::MMAIntOverflow::satfinite) {
3142 return emitOpError()
3143 << " `satfinite` can be only used with s32 accumulator, however "
3144 "the current accumulator is "
3145 << typeD;
3146 }
3147
3148 return success();
3149}
3150
3151std::string NVVM::WgmmaMmaAsyncOp::getPtx() {
3152
3153 int m = getShape().getM(), n = getShape().getN(), k = getShape().getK();
3154 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
3155
3156 StringRef outputTypeName = stringifyWGMMATypes(getTypeD());
3157
3158 int expectedOutputRegisters = 0;
3159 if (getTypeD() == WGMMATypes::f16)
3160 expectedOutputRegisters = getShape().getN() / 4;
3161 else
3162 expectedOutputRegisters = getShape().getN() / 2;
3163
3164 std::string ptx;
3165 llvm::raw_string_ostream ss(ptx);
3166
3167 ss << "{\n"
3168 ".reg .pred p;\n"
3169 "setp.ne.b32 p, $"
3170 << ((expectedOutputRegisters * 2) + 2)
3171 << ", 0;\n"
3172 "wgmma.mma_async.sync.aligned.m"
3173 << m << "n" << n << "k" << k << "." << outputTypeName << "." << getTypeA()
3174 << "." << getTypeB();
3175 if (getSatfinite().value_or(NVVM::MMAIntOverflow::wrapped) ==
3176 NVVM::MMAIntOverflow::satfinite)
3177 ss << ".satfinite";
3178 ss << " {";
3179 int regCnt = 0;
3180 for (; regCnt < expectedOutputRegisters; ++regCnt) {
3181 ss << "$" << regCnt;
3182 if (regCnt != expectedOutputRegisters - 1)
3183 ss << ", ";
3184 }
3185
3186 ss << "},";
3187 // Need to map read/write registers correctly.
3188 regCnt = (regCnt * 2);
3189 ss << " $" << (regCnt) << "," << " $" << (regCnt + 1) << "," << " p";
3190 if (getTypeD() != WGMMATypes::s32) {
3191 ss << ", $" << (regCnt + 3) << ", $" << (regCnt + 4);
3192 }
3193 // Don't add transpose parameters unless needed.
3194 if (isF16) {
3195 ss << ", $" << (regCnt + 5) << ", $" << (regCnt + 6);
3196 }
3197 ss << ";\n"
3198 << "}\n";
3199 return ptx;
3200}
3201
3202bool NVVM::WgmmaMmaAsyncOp::getAsmValues(
3203 RewriterBase &rewriter,
3204 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>
3205 &asmValues) {
3206 bool isF16 = getTypeA() == WGMMATypes::f16 || getTypeA() == WGMMATypes::bf16;
3207 if (getResults())
3208 asmValues.push_back({getResults(), mlir::NVVM::PTXRegisterMod::Write});
3209 if (getInouts())
3210 asmValues.push_back({getInouts(), mlir::NVVM::PTXRegisterMod::ReadWrite});
3211 asmValues.push_back({getDescriptorA(), mlir::NVVM::PTXRegisterMod::Read});
3212 asmValues.push_back({getDescriptorB(), mlir::NVVM::PTXRegisterMod::Read});
3213 asmValues.push_back({makeConstantI32(rewriter, static_cast<int>(getScaleD())),
3215 if (getTypeD() != WGMMATypes::s32) {
3216 asmValues.push_back(
3217 {makeConstantI32(rewriter,
3218 getScaleA() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
3220 asmValues.push_back(
3221 {makeConstantI32(rewriter,
3222 getScaleB() == NVVM::WGMMAScaleIn::neg ? -1 : 1),
3224 }
3225 if (isF16) {
3226 asmValues.push_back(
3227 {makeConstantI32(rewriter, static_cast<int>(getLayoutA())),
3229 asmValues.push_back(
3230 {makeConstantI32(rewriter, 1 - static_cast<int>(getLayoutB())),
3232 }
3233 return true; // Has manual mapping
3234}
3235
3236LogicalResult NVVM::FenceProxyOp::verify() {
3237 if (getKind() == NVVM::ProxyKind::async_shared && !getSpace().has_value()) {
3238 return emitOpError() << "async_shared fence requires space attribute";
3239 }
3240 if (getKind() != NVVM::ProxyKind::async_shared && getSpace().has_value()) {
3241 return emitOpError() << "only async_shared fence can have space attribute";
3242 }
3243 return success();
3244}
3245
3246LogicalResult NVVM::FenceProxyAcquireOp::verify() {
3247 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3248 return emitOpError("uni-directional proxies only support generic for "
3249 "from_proxy attribute");
3250
3251 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
3252 return emitOpError("uni-directional proxies only support tensormap "
3253 "for to_proxy attribute");
3254 return success();
3255}
3256
3257LogicalResult NVVM::FenceProxyReleaseOp::verify() {
3258 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3259 return emitOpError("uni-directional proxies only support generic for "
3260 "from_proxy attribute");
3261
3262 if (getToProxy() != NVVM::ProxyKind::TENSORMAP)
3263 return emitOpError("uni-directional proxies only support tensormap "
3264 "for to_proxy attribute");
3265 return success();
3266}
3267
3268LogicalResult NVVM::FenceProxySyncRestrictOp::verify() {
3269 if (getFromProxy() != NVVM::ProxyKind::GENERIC)
3270 return emitOpError("only generic is support for from_proxy attribute");
3271
3272 if (getToProxy() != NVVM::ProxyKind::async)
3273 return emitOpError("only async is supported for to_proxy attribute");
3274 return success();
3275}
3276
3277LogicalResult NVVM::SetMaxRegisterOp::verify() {
3278 if (getRegCount() % 8)
3279 return emitOpError("new register size must be multiple of 8");
3280 if (getRegCount() < 24 || getRegCount() > 256)
3281 return emitOpError("new register size must be in between 24 to 256");
3282 return success();
3283}
3284
3285LogicalResult NVVM::Tcgen05CpOp::verify() {
3286 auto mc = getMulticast();
3287
3288 using SH = Tcgen05CpShape;
3289 using MC = Tcgen05CpMulticast;
3290 switch (getShape()) {
3291 case SH::SHAPE_128x256b:
3292 case SH::SHAPE_128x128b:
3293 case SH::SHAPE_4x256b:
3294 if (mc != MC::NONE)
3295 return emitError("Invalid multicast type for tcgen05.cp Op");
3296 break;
3297 case SH::SHAPE_64x128b:
3298 if (mc != MC::WARPX2_01_23 && mc != MC::WARPX2_02_13)
3299 return emitError("Shape 64x128b requires multicast warpx2_01_23 or "
3300 "warpx2_02_13 for tcgen05.cp Op");
3301 break;
3302 case SH::SHAPE_32x128b:
3303 if (mc != MC::WARPX4)
3304 return emitError(
3305 "Shape 32x128b requires multicast warpx4 for tcgen05.cp Op");
3306 break;
3307 }
3308 return success();
3309}
3310
3311LogicalResult NVVM::MatchSyncOp::verify() {
3312 if (getKind() == NVVM::MatchSyncKind::all) {
3313 auto type = llvm::dyn_cast<LLVM::LLVMStructType>(getType());
3314 if (!type || type.getBody().size() != 2 ||
3315 !type.getBody()[0].isInteger(32) || !type.getBody()[1].isInteger(1)) {
3316 return emitOpError("match.sync 'all' returns a two element struct with "
3317 "first element as i32 and second element as i1");
3318 }
3319 } else {
3320 if (!getType().isInteger(32)) {
3321 return emitOpError("match.sync 'any' returns an i32");
3322 }
3323 }
3324 return success();
3325}
3326
3327LogicalResult MatchSyncOp::inferReturnTypes(
3328 MLIRContext *context, std::optional<Location> location,
3329 MatchSyncOp::Adaptor adaptor, SmallVectorImpl<Type> &inferredReturnTypes) {
3330 if (adaptor.getKind() == NVVM::MatchSyncKind::all)
3331 inferredReturnTypes.push_back(LLVM::LLVMStructType::getLiteral(
3332 context,
3333 {IntegerType::get(context, 32), IntegerType::get(context, 1)}));
3334 else
3335 inferredReturnTypes.push_back(IntegerType::get(context, 32));
3336 return success();
3337}
3338
3339LogicalResult NVVM::VoteSyncOp::verify() {
3340 if (getKind() == NVVM::VoteSyncKind::ballot) {
3341 if (!getType().isInteger(32)) {
3342 return emitOpError("vote.sync 'ballot' returns an i32");
3343 }
3344 } else {
3345 if (!getType().isInteger(1)) {
3346 return emitOpError("vote.sync 'any', 'all' and 'uni' returns an i1");
3347 }
3348 }
3349 return success();
3350}
3351
3352LogicalResult VoteSyncOp::inferReturnTypes(
3353 MLIRContext *context, std::optional<Location> location,
3354 VoteSyncOp::Adaptor adaptor, SmallVectorImpl<Type> &inferredReturnTypes) {
3355 unsigned width = adaptor.getKind() == NVVM::VoteSyncKind::ballot ? 32 : 1;
3356 inferredReturnTypes.push_back(IntegerType::get(context, width));
3357 return success();
3358}
3359
3360LogicalResult NVVM::PrefetchOp::verify() {
3361 using MemSpace = NVVM::NVVMMemorySpace;
3362 using CacheLevel = NVVM::PrefetchCacheLevel;
3363
3364 unsigned addressSpace =
3365 llvm::cast<LLVM::LLVMPointerType>(getAddr().getType()).getAddressSpace();
3366 std::optional<NVVM::CacheEvictionPriority> evictPriority = getEvictPriority();
3367 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = getCacheLevel();
3368
3369 if (getTensormap() && cacheLevel)
3370 return emitOpError("cannot specify both tensormap and cache level");
3371
3372 if (getTensormap()) {
3373 if (addressSpace != MemSpace::Generic &&
3374 addressSpace != MemSpace::Constant) {
3375 return emitOpError(
3376 "prefetch tensormap requires a generic or constant pointer");
3377 }
3378
3379 if (evictPriority) {
3380 return emitOpError(
3381 "prefetch tensormap does not support eviction priority");
3382 }
3383
3384 if (getInParamSpace() && addressSpace != MemSpace::Generic) {
3385 return emitOpError(
3386 "in_param_space can only be specified for a generic pointer");
3387 }
3388
3389 } else if (cacheLevel) {
3390 if (addressSpace != MemSpace::Generic && addressSpace != MemSpace::Global &&
3391 addressSpace != MemSpace::Local) {
3392 return emitOpError("prefetch to cache level requires a generic, global, "
3393 "or local pointer");
3394 }
3395
3396 if (getUniform()) {
3397 if (*cacheLevel != CacheLevel::L1) {
3398 return emitOpError(
3399 "unsupported cache level, the only supported uniform "
3400 "cache level is L1");
3401 }
3402
3403 if (addressSpace != MemSpace::Generic) {
3404 return emitOpError(
3405 "prefetch to uniform cache requires a generic pointer");
3406 }
3407 }
3408
3409 if (evictPriority) {
3410 if (*cacheLevel != CacheLevel::L2)
3411 return emitOpError(
3412 "cache eviction priority supported only for cache level L2");
3413
3414 if (addressSpace != MemSpace::Global)
3415 return emitOpError("cache eviction priority requires a global pointer");
3416
3417 if (*evictPriority != NVVM::CacheEvictionPriority::EvictNormal &&
3418 *evictPriority != NVVM::CacheEvictionPriority::EvictLast)
3419 return emitOpError(
3420 "unsupported cache eviction priority, only evict_last and "
3421 "evict_normal are supported");
3422 }
3423
3424 if (getPredicate())
3425 return emitOpError("predicate supported only on prefetch tensormap");
3426
3427 } else {
3428 return emitOpError(
3429 "requires specification of either cache level or tensormap");
3430 }
3431
3432 return success();
3433}
3434
3435LogicalResult NVVM::ClusterLaunchControlQueryCancelOp::verify() {
3436 switch (getQueryType()) {
3437 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
3438 if (!getType().isInteger(1))
3439 return emitOpError("is_canceled query type returns an i1");
3440 break;
3441 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
3442 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
3443 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
3444 if (!getType().isInteger(32)) {
3445 return emitOpError("get_first_cta_id_x, get_first_cta_id_y, "
3446 "get_first_cta_id_z query types return an i32");
3447 }
3448 break;
3449 }
3450 return success();
3451}
3452
3453LogicalResult ClusterLaunchControlQueryCancelOp::inferReturnTypes(
3454 MLIRContext *context, std::optional<Location> location,
3455 ClusterLaunchControlQueryCancelOp::Adaptor adaptor,
3456 SmallVectorImpl<Type> &inferredReturnTypes) {
3457 unsigned width =
3458 adaptor.getQueryType() == NVVM::ClusterLaunchControlQueryType::IS_CANCELED
3459 ? 1
3460 : 32;
3461 inferredReturnTypes.push_back(IntegerType::get(context, width));
3462 return success();
3463}
3464
3465LogicalResult NVVM::ReduxOp::verify() {
3466 mlir::Type reduxType = getType();
3467
3468 if (!reduxType.isF32()) {
3469 if (getAbs())
3470 return emitOpError("abs attribute is supported only for f32 type");
3471 if (getNan())
3472 return emitOpError("nan attribute is supported only for f32 type");
3473 }
3474
3475 NVVM::ReductionKind kind = getKind();
3476 switch (kind) {
3477 case NVVM::ReductionKind::ADD:
3478 case NVVM::ReductionKind::AND:
3479 case NVVM::ReductionKind::OR:
3480 case NVVM::ReductionKind::XOR:
3481 case NVVM::ReductionKind::MAX:
3482 case NVVM::ReductionKind::MIN:
3483 case NVVM::ReductionKind::UMAX:
3484 case NVVM::ReductionKind::UMIN:
3485 if (!reduxType.isInteger(32))
3486 return emitOpError("'")
3487 << kind << "' reduction kind unsupported with " << reduxType
3488 << " type. Only supported type is 'i32'.";
3489 break;
3490 case NVVM::ReductionKind::FMIN:
3491 case NVVM::ReductionKind::FMAX:
3492 if (!reduxType.isF32())
3493 return emitOpError("'")
3494 << kind << "' reduction kind unsupported with " << reduxType
3495 << " type. Only supported type is 'f32'.";
3496 break;
3497 }
3498
3499 return success();
3500}
3501
3502LogicalResult NVVM::TensormapReplaceOp::verify() {
3503 auto ord = getOrd();
3504 Value newVal = getNewValue();
3505 auto newValAttr = getNewValueAttr();
3506 auto fieldName = stringifyEnum(getField());
3507
3508 if (ord && !llvm::is_contained({NVVM::TensormapField::BOX_DIM,
3509 NVVM::TensormapField::GLOBAL_DIM,
3510 NVVM::TensormapField::GLOBAL_STRIDE,
3511 NVVM::TensormapField::ELEMENT_STRIDE},
3512 getField()))
3513 return emitOpError("ordinal is not supported for ")
3514 << fieldName << " field";
3515
3516 auto invalidNewVal = [&](llvm::Twine type) -> std::string {
3517 return llvm::Twine("new_value must be specified and must be an " + type +
3518 " for " + llvm::Twine(fieldName) + " field")
3519 .str();
3520 };
3521
3522 auto invalidNewValAttr = [&]() -> std::string {
3523 return (llvm::Twine(
3524 "new_value_attr must be specified and must be a valid ") +
3525 llvm::Twine(fieldName) + " attribute for " + fieldName + " field")
3526 .str();
3527 };
3528
3529 switch (getField()) {
3530 case NVVM::TensormapField::GLOBAL_ADDRESS:
3531 if (!(newVal && newVal.getType().isInteger(64)))
3532 return emitOpError(invalidNewVal("i64"));
3533 break;
3534 case NVVM::TensormapField::RANK:
3535 if (!(newVal && newVal.getType().isInteger(32)))
3536 return emitOpError(invalidNewVal("i32"));
3537 break;
3538 case NVVM::TensormapField::GLOBAL_STRIDE:
3539 if (!ord)
3540 return emitOpError("ordinal is required for global_stride field");
3541 if (!(newVal && newVal.getType().isInteger(64)))
3542 return emitOpError(invalidNewVal("i64"));
3543 break;
3544 case NVVM::TensormapField::BOX_DIM:
3545 case NVVM::TensormapField::GLOBAL_DIM:
3546 case NVVM::TensormapField::ELEMENT_STRIDE:
3547 if (!ord)
3548 return emitOpError("ordinal is required for ")
3549 << stringifyEnum(getField()) << " field";
3550 if (!(newVal && newVal.getType().isInteger(32)))
3551 return emitOpError(invalidNewVal("i32"));
3552 break;
3553 case NVVM::TensormapField::ELEMTYPE:
3554 if (!(newValAttr && llvm::isa<TensormapElemtypeAttr>(*newValAttr)))
3555 return emitOpError(invalidNewValAttr());
3556 break;
3557 case NVVM::TensormapField::INTERLEAVE_LAYOUT:
3558 if (!(newValAttr && llvm::isa<TensormapInterleaveLayoutAttr>(*newValAttr)))
3559 return emitOpError(invalidNewValAttr());
3560 break;
3561 case NVVM::TensormapField::SWIZZLE_MODE:
3562 if (!(newValAttr && llvm::isa<TensormapSwizzleModeAttr>(*newValAttr)))
3563 return emitOpError(invalidNewValAttr());
3564 break;
3565 case NVVM::TensormapField::SWIZZLE_ATOMICITY:
3566 if (!(newValAttr && llvm::isa<TensormapSwizzleAtomicityAttr>(*newValAttr)))
3567 return emitOpError(invalidNewValAttr());
3568 break;
3569 case NVVM::TensormapField::FILL_MODE:
3570 if (!(newValAttr && llvm::isa<TensormapFillModeAttr>(*newValAttr)))
3571 return emitOpError(invalidNewValAttr());
3572 break;
3573 }
3574
3575 return success();
3576}
3577
3578template <typename OpType>
3579static LogicalResult verifyAddSubFOp(OpType op) {
3580 mlir::NVVM::FPRoundingMode rndMode = op.getRnd();
3581 mlir::NVVM::SaturationMode satMode = op.getSat();
3582 bool isFTZ = op.getFtz();
3583
3584 mlir::Type opType = op.getRes().getType();
3585 mlir::Type opBaseType = isa<VectorType>(opType)
3586 ? cast<VectorType>(opType).getElementType()
3587 : opType;
3588
3589 if (opBaseType.isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3590 return op.emitOpError("FTZ and saturation are not supported for "
3591 "additions/subtractions involving f64 type");
3592
3593 if (opBaseType.isF16() && !(rndMode == NVVM::FPRoundingMode::RN ||
3594 rndMode == NVVM::FPRoundingMode::NONE))
3595 return op.emitOpError("only RN rounding mode is supported for f16 and "
3596 "vector<2xf16> additions/subtractions");
3597
3598 if (opBaseType.isBF16()) {
3599 if (rndMode != NVVM::FPRoundingMode::RN &&
3600 rndMode != NVVM::FPRoundingMode::NONE)
3601 return op.emitOpError("only RN rounding mode is supported for bf16 and "
3602 "vector<2xbf16> additions/subtractions");
3603 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3604 return op.emitOpError("FTZ and saturation are not supported for bf16 and "
3605 "vector<2xbf16> additions/subtractions");
3606 }
3607
3608 // FIXME: This is a temporary check disallowing lowering to add.rn.ftz.f16(x2)
3609 // PTX instructions since the corresponding LLVM intrinsic is missing. This
3610 // should be removed once the intrinsics for f16 addition (with FTZ only) are
3611 // available.
3612 if (opBaseType.isF16() && isFTZ && satMode == NVVM::SaturationMode::NONE)
3613 return op.emitOpError("FTZ with no saturation is not supported for f16 and "
3614 "vector<2xf16> additions/subtractions");
3615
3616 return success();
3617}
3618
3619LogicalResult NVVM::AddFOp::verify() { return verifyAddSubFOp<AddFOp>(*this); }
3620
3621LogicalResult NVVM::SubFOp::verify() { return verifyAddSubFOp<SubFOp>(*this); }
3622
3623LogicalResult NVVM::FmaOp::verify() {
3624 auto opType = getRes().getType();
3625 mlir::NVVM::FPRoundingMode rndMode = getRnd();
3626 mlir::NVVM::SaturationMode satMode = getSat();
3627 bool isFTZ = getFtz();
3628 bool isRelu = getRelu();
3629 bool hasOOB = getOob();
3630
3631 auto getBaseFType = [](Type type) -> Type {
3632 if (isa<VectorType>(type))
3633 return cast<VectorType>(type).getElementType();
3634 return type;
3635 };
3636
3637 auto opBaseType = getBaseFType(opType);
3638
3639 if (rndMode == NVVM::FPRoundingMode::NONE)
3640 return emitOpError("rounding mode must be specified");
3641
3642 if (isRelu && satMode == NVVM::SaturationMode::SAT)
3643 return emitOpError("relu and saturation are not supported together");
3644
3645 if (hasOOB && (satMode == NVVM::SaturationMode::SAT || isFTZ))
3646 return emitOpError("oob is not supported with saturation or FTZ");
3647
3648 if (!(opBaseType.isF16() || opBaseType.isBF16()) && (isRelu || hasOOB))
3649 return emitOpError("relu and oob are only supported for f16 and bf16");
3650
3651 if (opBaseType.isF64() && (satMode != NVVM::SaturationMode::NONE || isFTZ))
3652 return emitOpError("FTZ and saturation are not supported for f64 type");
3653
3654 if (opBaseType.isF16() && rndMode != NVVM::FPRoundingMode::RN)
3655 return emitOpError(
3656 "only RN rounding mode is supported for f16 and vector<2xf16>");
3657
3658 if (opBaseType.isBF16()) {
3659 if (rndMode != NVVM::FPRoundingMode::RN)
3660 return emitOpError(
3661 "only RN rounding mode is supported for bf16 and vector<2xbf16>");
3662 if (satMode != NVVM::SaturationMode::NONE || isFTZ)
3663 return emitOpError(
3664 "FTZ and saturation are not supported for bf16 and vector<2xbf16>");
3665 }
3666
3667 return success();
3668}
3669
3670LogicalResult NVVM::SqrtOp::verify() {
3671 if (getRnd() == NVVM::FPRoundingMode::NONE)
3672 return emitOpError("rounding mode cannot be None");
3673
3674 if (getRes().getType().isF64() && getFtz())
3675 return emitOpError("FTZ is not supported for f64");
3676
3677 return success();
3678}
3679
3680LogicalResult NVVM::DivFOp::verify() {
3681 bool isApprox = getApprox();
3682 bool isFull = getFull();
3683 bool isF64 = getRes().getType().isF64();
3684 bool isFtz = getFtz();
3685 NVVM::FPRoundingMode rndMode = getRnd();
3686
3687 if (isApprox && isFull)
3688 return emitOpError("'approx' and 'full' are mutually exclusive");
3689
3690 if (isApprox || isFull) {
3691 if (isF64)
3692 return emitOpError("'approx' and 'full' forms are f32-only");
3693 if (rndMode != NVVM::FPRoundingMode::NONE)
3694 return emitOpError(
3695 "'approx' and 'full' forms do not accept a rounding mode");
3696 return success();
3697 }
3698
3699 // Rounded form below.
3700 if (rndMode == NVVM::FPRoundingMode::NONE)
3701 return emitOpError("rounding mode cannot be None for the rounded divide");
3702 if (isF64 && isFtz)
3703 return emitOpError("FTZ is not supported for f64");
3704
3705 return success();
3706}
3707
3708/// Packs the given `field` into the `result`.
3709/// The `result` is 64-bits and each `field` can be 32-bits or narrower.
3710static llvm::Value *
3711packValInto64Bits(llvm::IRBuilderBase &builder,
3712 llvm::Value *result, // the `result` (unset bits are zero)
3713 llvm::Value *field, // `field` to pack into `result`
3714 unsigned sizeInBits, // Size of `field` in bits
3715 unsigned start) { // Starting bit within `result`
3716 field = builder.CreateZExtOrBitCast(field, builder.getInt32Ty());
3717
3718 unsigned mask = (sizeInBits < 32 ? ((1u << sizeInBits) - 1) : 0xffffffffu);
3719 if (mask != 0xffffffffu)
3720 field = builder.CreateAnd(field, builder.getInt32(mask));
3721
3722 field = builder.CreateZExtOrBitCast(field, builder.getInt64Ty());
3723 field = builder.CreateShl(field, start);
3724
3725 return builder.CreateOr(result, field);
3726}
3727
3728void Tcgen05MmaSmemDescOp::createSmemDescriptor(Operation &op,
3730 llvm::IRBuilderBase &builder) {
3731 auto thisOp = cast<NVVM::Tcgen05MmaSmemDescOp>(op);
3732 llvm::Value *smemDesc = builder.getInt64(0);
3733
3734 smemDesc = packValInto64Bits(builder, smemDesc,
3735 mt.lookupValue(thisOp.getStartAddr()), 14, 0);
3736 smemDesc = packValInto64Bits(
3737 builder, smemDesc, mt.lookupValue(thisOp.getLeadingDimOffset()), 14, 16);
3738 smemDesc = packValInto64Bits(
3739 builder, smemDesc, mt.lookupValue(thisOp.getStrideDimOffset()), 14, 32);
3740
3741 smemDesc = packValInto64Bits(builder, smemDesc, builder.getInt32(1), 3, 46);
3742 smemDesc = packValInto64Bits(builder, smemDesc,
3743 mt.lookupValue(thisOp.getBaseOffset()), 3, 49);
3744 smemDesc = packValInto64Bits(
3745 builder, smemDesc, mt.lookupValue(thisOp.getLeadingDimMode()), 1, 52);
3746 smemDesc = packValInto64Bits(builder, smemDesc,
3747 mt.lookupValue(thisOp.getSwizzleMode()), 3, 61);
3748
3749 mt.mapValue(thisOp.getRes()) = smemDesc;
3750}
3751
3752//===----------------------------------------------------------------------===//
3753// getPtx methods
3754//===----------------------------------------------------------------------===//
3755
3756std::string NVVM::MBarrierInitOp::getPtx() {
3757 bool isShared = isPtrInSharedCTASpace(getAddr());
3758 return isShared ? std::string("mbarrier.init.shared.b64 [%0], %1;")
3759 : std::string("mbarrier.init.b64 [%0], %1;");
3760}
3761
3762std::string NVVM::MBarrierArriveExpectTxOp::getPtx() {
3763 bool isShared = isPtrInSharedCTASpace(getAddr());
3764 return isShared
3765 ? std::string("mbarrier.arrive.expect_tx.shared.b64 _, [%0], %1;")
3766 : std::string("mbarrier.arrive.expect_tx.b64 _, [%0], %1;");
3767}
3768
3769std::string NVVM::MBarrierTryWaitParityOp::getPtx() {
3770 bool isShared = isPtrInSharedCTASpace(getAddr());
3771 llvm::StringRef space = isShared ? ".shared" : "";
3772
3773 return llvm::formatv("{\n\t"
3774 ".reg .pred P1; \n\t"
3775 "LAB_WAIT: \n\t"
3776 "mbarrier.try_wait.parity{0}.b64 P1, [%0], %1, %2; \n\t"
3777 "@P1 bra.uni DONE; \n\t"
3778 "bra.uni LAB_WAIT; \n\t"
3779 "DONE: \n\t"
3780 "}",
3781 space);
3782}
3783
3784//===----------------------------------------------------------------------===//
3785// Canonicalization patterns
3786//===----------------------------------------------------------------------===//
3787
3790
3791 LogicalResult matchAndRewrite(SubFOp op,
3792 PatternRewriter &rewriter) const override {
3793 Location loc = op.getLoc();
3794 Value negRhs =
3795 LLVM::FNegOp::create(rewriter, loc, op.getRhs().getType(), op.getRhs());
3796
3797 rewriter.replaceOpWithNewOp<AddFOp>(op, op.getType(), op.getLhs(), negRhs,
3798 op.getRnd(), op.getSat(), op.getFtz());
3799 return success();
3800 }
3801};
3802
3803void SubFOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
3804 MLIRContext *context) {
3805 patterns.add<ConvertFsubToFnegFadd>(context);
3806}
3807
3808//===----------------------------------------------------------------------===//
3809// getIntrinsicID/getIntrinsicIDAndArgs methods
3810//===----------------------------------------------------------------------===//
3811
3812/// Maps the (aligned, hasCount) pair to the `@llvm.nvvm.barrier.cta.sync.*`
3813/// intrinsic ID.
3814static llvm::Intrinsic::ID getBarrierSyncIntrinsic(bool aligned,
3815 bool hasCount) {
3816 if (hasCount) {
3817 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_count
3818 : llvm::Intrinsic::nvvm_barrier_cta_sync_count;
3819 }
3820 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_sync_aligned_all
3821 : llvm::Intrinsic::nvvm_barrier_cta_sync_all;
3822}
3823
3824/// Maps the (aligned, kind) pair to the `@llvm.nvvm.barrier.cta.red.*`
3825/// intrinsic ID.
3826static llvm::Intrinsic::ID
3827getBarrierReductionIntrinsic(bool aligned, NVVM::BarrierReduction kind) {
3828 switch (kind) {
3829 case NVVM::BarrierReduction::AND:
3830 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_and_aligned_all
3831 : llvm::Intrinsic::nvvm_barrier_cta_red_and_all;
3832 case NVVM::BarrierReduction::OR:
3833 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_or_aligned_all
3834 : llvm::Intrinsic::nvvm_barrier_cta_red_or_all;
3835 case NVVM::BarrierReduction::POPC:
3836 return aligned ? llvm::Intrinsic::nvvm_barrier_cta_red_popc_aligned_all
3837 : llvm::Intrinsic::nvvm_barrier_cta_red_popc_all;
3838 }
3839 llvm_unreachable("unknown BarrierReduction kind");
3840}
3841
3842mlir::NVVM::IDArgPair NVVM::BarrierOp::getIntrinsicIDAndArgs(
3843 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
3844 auto thisOp = cast<NVVM::BarrierOp>(op);
3845 llvm::Value *barrierId = thisOp.getBarrierId()
3846 ? mt.lookupValue(thisOp.getBarrierId())
3847 : builder.getInt32(0);
3848 bool hasCount = static_cast<bool>(thisOp.getNumberOfThreads());
3849 llvm::Intrinsic::ID id =
3850 getBarrierSyncIntrinsic(thisOp.getAligned(), hasCount);
3851 llvm::SmallVector<llvm::Value *> args = {barrierId};
3852 if (hasCount)
3853 args.push_back(mt.lookupValue(thisOp.getNumberOfThreads()));
3854 return {id, std::move(args)};
3855}
3856
3857mlir::NVVM::IDArgPair NVVM::BarrierArriveOp::getIntrinsicIDAndArgs(
3858 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
3859 auto thisOp = cast<NVVM::BarrierArriveOp>(op);
3860 llvm::Value *barrierId = thisOp.getBarrierId()
3861 ? mt.lookupValue(thisOp.getBarrierId())
3862 : builder.getInt32(0);
3863 llvm::Value *numThreads = mt.lookupValue(thisOp.getNumberOfThreads());
3864 llvm::Intrinsic::ID id =
3865 thisOp.getAligned()
3866 ? llvm::Intrinsic::nvvm_barrier_cta_arrive_aligned_count
3867 : llvm::Intrinsic::nvvm_barrier_cta_arrive_count;
3868 return {id, {barrierId, numThreads}};
3869}
3870
3871mlir::NVVM::IDArgPair NVVM::BarrierReductionOp::getIntrinsicIDAndArgs(
3872 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
3873 auto thisOp = cast<NVVM::BarrierReductionOp>(op);
3874 llvm::Intrinsic::ID id = getBarrierReductionIntrinsic(
3875 thisOp.getAligned(), thisOp.getReductionOp());
3876 llvm::Value *barrierId = thisOp.getBarrierId()
3877 ? mt.lookupValue(thisOp.getBarrierId())
3878 : builder.getInt32(0);
3880 barrierId,
3881 builder.CreateICmpNE(mt.lookupValue(thisOp.getReductionPredicate()),
3882 builder.getInt32(0))};
3883 return {id, std::move(args)};
3884}
3885
3887CosOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3888 llvm::IRBuilderBase &builder) {
3889 auto thisOp = cast<NVVM::CosOp>(op);
3890 llvm::Intrinsic::ID id = thisOp.getFtz()
3891 ? llvm::Intrinsic::nvvm_cos_approx_ftz_f
3892 : llvm::Intrinsic::nvvm_cos_approx_f;
3893 return {id, {mt.lookupValue(thisOp.getSrc())}};
3894}
3895
3897SinOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3898 llvm::IRBuilderBase &builder) {
3899 auto thisOp = cast<NVVM::SinOp>(op);
3900 llvm::Intrinsic::ID id = thisOp.getFtz()
3901 ? llvm::Intrinsic::nvvm_sin_approx_ftz_f
3902 : llvm::Intrinsic::nvvm_sin_approx_f;
3903 return {id, {mt.lookupValue(thisOp.getSrc())}};
3904}
3905
3907Log2Op::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3908 llvm::IRBuilderBase &builder) {
3909 auto thisOp = cast<NVVM::Log2Op>(op);
3910 llvm::Intrinsic::ID id = thisOp.getFtz()
3911 ? llvm::Intrinsic::nvvm_lg2_approx_ftz_f
3912 : llvm::Intrinsic::nvvm_lg2_approx_f;
3913 return {id, {mt.lookupValue(thisOp.getSrc())}};
3914}
3915
3917Ex2Op::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3918 llvm::IRBuilderBase &builder) {
3919 auto thisOp = cast<NVVM::Ex2Op>(op);
3920 llvm::Intrinsic::ID id = thisOp.getFtz()
3921 ? llvm::Intrinsic::nvvm_ex2_approx_ftz
3922 : llvm::Intrinsic::nvvm_ex2_approx;
3923 return {id, {mt.lookupValue(thisOp.getSrc())}};
3924}
3925
3927RsqrtOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3928 llvm::IRBuilderBase &builder) {
3929 auto thisOp = cast<NVVM::RsqrtOp>(op);
3930 Type t = thisOp.getRes().getType();
3931 bool isFtz = thisOp.getFtz();
3932
3933 llvm::Intrinsic::ID id = [&] {
3934 if (t.isF32()) {
3935 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_f
3936 : llvm::Intrinsic::nvvm_rsqrt_approx_f;
3937 }
3938 // f64
3939 return isFtz ? llvm::Intrinsic::nvvm_rsqrt_approx_ftz_d
3940 : llvm::Intrinsic::nvvm_rsqrt_approx_d;
3941 }();
3942
3943 return {id, {mt.lookupValue(thisOp.getSrc())}};
3944}
3945
3947SqrtOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3948 llvm::IRBuilderBase &builder) {
3949 auto thisOp = cast<NVVM::SqrtOp>(op);
3950 Type t = thisOp.getRes().getType();
3951 NVVM::FPRoundingMode rndMode = thisOp.getRnd();
3952 bool isFtz = thisOp.getFtz();
3953
3954 // RM is one of RN/RM/RP/RZ (verifier rejects NONE).
3955 // Subtracting 1 maps RN=1..RZ=4 to 0..3.
3956 unsigned rndIndex = static_cast<unsigned>(rndMode) - 1;
3957
3958 static constexpr llvm::Intrinsic::ID f32IDs[] = {
3959 llvm::Intrinsic::nvvm_sqrt_rn_f,
3960 llvm::Intrinsic::nvvm_sqrt_rm_f,
3961 llvm::Intrinsic::nvvm_sqrt_rp_f,
3962 llvm::Intrinsic::nvvm_sqrt_rz_f,
3963 };
3964 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
3965 llvm::Intrinsic::nvvm_sqrt_rn_ftz_f,
3966 llvm::Intrinsic::nvvm_sqrt_rm_ftz_f,
3967 llvm::Intrinsic::nvvm_sqrt_rp_ftz_f,
3968 llvm::Intrinsic::nvvm_sqrt_rz_ftz_f,
3969 };
3970 static constexpr llvm::Intrinsic::ID f64IDs[] = {
3971 llvm::Intrinsic::nvvm_sqrt_rn_d,
3972 llvm::Intrinsic::nvvm_sqrt_rm_d,
3973 llvm::Intrinsic::nvvm_sqrt_rp_d,
3974 llvm::Intrinsic::nvvm_sqrt_rz_d,
3975 };
3976
3977 llvm::Intrinsic::ID id =
3978 t.isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
3979 : f64IDs[rndIndex];
3980
3981 return {id, {mt.lookupValue(thisOp.getSrc())}};
3982}
3983
3985SqrtApproxOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3986 llvm::IRBuilderBase &builder) {
3987 auto thisOp = cast<NVVM::SqrtApproxOp>(op);
3988 llvm::Intrinsic::ID id = thisOp.getFtz()
3989 ? llvm::Intrinsic::nvvm_sqrt_approx_ftz_f
3990 : llvm::Intrinsic::nvvm_sqrt_approx_f;
3991 return {id, {mt.lookupValue(thisOp.getSrc())}};
3992}
3993
3995DivFOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
3996 llvm::IRBuilderBase &builder) {
3997 auto thisOp = cast<NVVM::DivFOp>(op);
3998 bool isFtz = thisOp.getFtz();
3999
4000 llvm::Intrinsic::ID id;
4001
4002 if (thisOp.getApprox()) {
4003 id = isFtz ? llvm::Intrinsic::nvvm_div_approx_ftz_f
4004 : llvm::Intrinsic::nvvm_div_approx_f;
4005 } else if (thisOp.getFull()) {
4006 // Intrinsic Naming quirk: int_nvvm_div_full has no `_f` suffix (unlike
4007 // approx).
4008 id = isFtz ? llvm::Intrinsic::nvvm_div_full_ftz
4009 : llvm::Intrinsic::nvvm_div_full;
4010 } else {
4011 // Rounded form — three 4-entry tables indexed by (rndMode - 1).
4012 unsigned rndIndex = static_cast<unsigned>(thisOp.getRnd()) - 1;
4013
4014 static constexpr llvm::Intrinsic::ID f32IDs[] = {
4015 llvm::Intrinsic::nvvm_div_rn_f,
4016 llvm::Intrinsic::nvvm_div_rm_f,
4017 llvm::Intrinsic::nvvm_div_rp_f,
4018 llvm::Intrinsic::nvvm_div_rz_f,
4019 };
4020 static constexpr llvm::Intrinsic::ID f32FTZIDs[] = {
4021 llvm::Intrinsic::nvvm_div_rn_ftz_f,
4022 llvm::Intrinsic::nvvm_div_rm_ftz_f,
4023 llvm::Intrinsic::nvvm_div_rp_ftz_f,
4024 llvm::Intrinsic::nvvm_div_rz_ftz_f,
4025 };
4026 static constexpr llvm::Intrinsic::ID f64IDs[] = {
4027 llvm::Intrinsic::nvvm_div_rn_d,
4028 llvm::Intrinsic::nvvm_div_rm_d,
4029 llvm::Intrinsic::nvvm_div_rp_d,
4030 llvm::Intrinsic::nvvm_div_rz_d,
4031 };
4032 Type t = thisOp.getRes().getType();
4033 id = t.isF32() ? (isFtz ? f32FTZIDs[rndIndex] : f32IDs[rndIndex])
4034 : f64IDs[rndIndex];
4035 }
4036
4037 return {id,
4038 {mt.lookupValue(thisOp.getLhs()), mt.lookupValue(thisOp.getRhs())}};
4039}
4040
4041mlir::NVVM::IDArgPair AsyncStoreGlobalOp::getIntrinsicIDAndArgs(
4042 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4044 auto thisOp = cast<NVVM::AsyncStoreGlobalOp>(op);
4045 mlir::NVVM::MemScopeKind scope = thisOp.getScope();
4046 bool isMmio = thisOp.getMmio();
4047
4048 llvm::Value *addr = mt.lookupValue(thisOp.getAddr());
4049 llvm::Value *value = mt.lookupValue(thisOp.getValue());
4050 llvm::Value *isMultimem = builder.getInt1(thisOp.getMultimem());
4051
4052 if (scope == MemScopeKind::SYS) {
4053 return isMmio ? IDArgPair(llvm::Intrinsic::nvvm_st_async_mmio_sys,
4054 {addr, value})
4055 : IDArgPair(llvm::Intrinsic::nvvm_st_async_sys,
4056 {addr, value, isMultimem});
4057 } else if (scope == MemScopeKind::GPU) {
4058 return IDArgPair(llvm::Intrinsic::nvvm_st_async_gpu,
4059 {addr, value, isMultimem});
4060 }
4061 llvm_unreachable("unsupported scope for AsyncStoreGlobalOp");
4062}
4063
4065PMEventOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
4066 llvm::IRBuilderBase &builder) {
4067 auto thisOp = cast<NVVM::PMEventOp>(op);
4068 llvm::Type *i16Ty = llvm::Type::getInt16Ty(mt.getLLVMContext());
4069
4070 // With event-id, mask is generated as (1 << event-id)
4071 llvm::Value *maskVal;
4072 if (auto eventAttr = thisOp.getEventIdAttr()) {
4073 uint16_t mask = static_cast<uint16_t>(1u << eventAttr.getInt());
4074 maskVal = llvm::ConstantInt::get(i16Ty, mask);
4075 } else {
4076 maskVal =
4077 llvm::ConstantInt::get(i16Ty, thisOp.getMaskedEventIdAttr().getValue());
4078 }
4079
4080 return {llvm::Intrinsic::nvvm_pm_event_mask, {maskVal}};
4081}
4082
4083mlir::NVVM::IDArgPair MBarrierInitOp::getIntrinsicIDAndArgs(
4084 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4085 auto thisOp = cast<NVVM::MBarrierInitOp>(op);
4086 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());
4087 llvm::Intrinsic::ID id = isShared ? llvm::Intrinsic::nvvm_mbarrier_init_shared
4088 : llvm::Intrinsic::nvvm_mbarrier_init;
4089
4090 // Fill the Intrinsic Args
4092 args.push_back(mt.lookupValue(thisOp.getAddr()));
4093 args.push_back(mt.lookupValue(thisOp.getCount()));
4094
4095 return {id, std::move(args)};
4096}
4097
4098mlir::NVVM::IDArgPair MBarrierInvalOp::getIntrinsicIDAndArgs(
4099 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4100 auto thisOp = cast<NVVM::MBarrierInvalOp>(op);
4101 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());
4102 llvm::Intrinsic::ID id = isShared
4103 ? llvm::Intrinsic::nvvm_mbarrier_inval_shared
4104 : llvm::Intrinsic::nvvm_mbarrier_inval;
4105
4106 return {id, {mt.lookupValue(thisOp.getAddr())}};
4107}
4108
4109mlir::NVVM::IDArgPair MBarrierExpectTxOp::getIntrinsicIDAndArgs(
4110 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4111 auto thisOp = cast<NVVM::MBarrierExpectTxOp>(op);
4112
4113 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4114 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4115 // bit-0: Space
4116 // bit-1: Scope
4117 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4118
4119 static constexpr llvm::Intrinsic::ID IDs[] = {
4120 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cta,
4121 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cta_space_cluster,
4122 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cta,
4123 llvm::Intrinsic::nvvm_mbarrier_expect_tx_scope_cluster_space_cluster};
4124
4125 // Fill the Intrinsic Args
4127 args.push_back(mt.lookupValue(thisOp.getAddr()));
4128 args.push_back(mt.lookupValue(thisOp.getTxcount()));
4129
4130 return {IDs[index], std::move(args)};
4131}
4132
4133mlir::NVVM::IDArgPair MBarrierCompleteTxOp::getIntrinsicIDAndArgs(
4134 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4135 auto thisOp = cast<NVVM::MBarrierCompleteTxOp>(op);
4136
4137 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4138 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4139 // bit-0: Space
4140 // bit-1: Scope
4141 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4142
4143 static constexpr llvm::Intrinsic::ID IDs[] = {
4144 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cta,
4145 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cta_space_cluster,
4146 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cta,
4147 llvm::Intrinsic::nvvm_mbarrier_complete_tx_scope_cluster_space_cluster};
4148
4149 // Fill the Intrinsic Args
4151 args.push_back(mt.lookupValue(thisOp.getAddr()));
4152 args.push_back(mt.lookupValue(thisOp.getTxcount()));
4153
4154 return {IDs[index], std::move(args)};
4155}
4156
4157mlir::NVVM::IDArgPair MBarrierArriveOp::getIntrinsicIDAndArgs(
4158 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4159 auto thisOp = cast<NVVM::MBarrierArriveOp>(op);
4160
4161 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4162 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4163 // bit-0: Space
4164 // bit-1: Scope
4165 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4166
4167 static constexpr llvm::Intrinsic::ID IDs[] = {
4168 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta,
4169 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cluster,
4170 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cta,
4171 llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cluster_space_cluster};
4172 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4173 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cta,
4174 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cta_space_cluster,
4175 llvm::Intrinsic::nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cta,
4176 llvm::Intrinsic::
4177 nvvm_mbarrier_arrive_relaxed_scope_cluster_space_cluster};
4178 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4179
4180 // Tidy-up the Intrinsic Args
4181 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4182 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4183 if (needCast)
4184 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4185
4186 // We have the most basic mbarrier.arrive supported on sm_80.
4187 // It supports: Space=cta, scope=cta, No relaxed, No explicit count.
4188 // So, only for this combination use the legacy intrinsic.
4189 bool hasCount = static_cast<bool>(thisOp.getCount());
4190 if (!hasCount &&
4191 (id == llvm::Intrinsic::nvvm_mbarrier_arrive_scope_cta_space_cta))
4192 return {llvm::Intrinsic::nvvm_mbarrier_arrive_shared, {mbar}};
4193
4194 // When count is not explicitly specified, the default is 1.
4195 llvm::LLVMContext &ctx = mt.getLLVMContext();
4196 llvm::Value *count =
4197 hasCount ? mt.lookupValue(thisOp.getCount())
4198 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
4199 return {id, {mbar, count}};
4200}
4201
4202mlir::NVVM::IDArgPair MBarrierArriveDropOp::getIntrinsicIDAndArgs(
4203 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4204 auto thisOp = cast<NVVM::MBarrierArriveDropOp>(op);
4205
4206 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4207 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4208 // bit-0: Space
4209 // bit-1: Scope
4210 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4211
4212 static constexpr llvm::Intrinsic::ID IDs[] = {
4213 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cta,
4214 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cta_space_cluster,
4215 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cta,
4216 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_scope_cluster_space_cluster};
4217 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4218 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cta,
4219 llvm::Intrinsic::
4220 nvvm_mbarrier_arrive_drop_relaxed_scope_cta_space_cluster,
4221 llvm::Intrinsic::
4222 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cta,
4223 llvm::Intrinsic::
4224 nvvm_mbarrier_arrive_drop_relaxed_scope_cluster_space_cluster};
4225 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4226
4227 // Tidy-up the Intrinsic Args
4228 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4229 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4230 if (needCast)
4231 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4232
4233 // When count is not explicitly specified, the default is 1.
4234 llvm::LLVMContext &ctx = mt.getLLVMContext();
4235 bool hasCount = static_cast<bool>(thisOp.getCount());
4236 llvm::Value *count =
4237 hasCount ? mt.lookupValue(thisOp.getCount())
4238 : llvm::ConstantInt::get(llvm::Type::getInt32Ty(ctx), 1);
4239
4240 return {id, {mbar, count}};
4241}
4242
4243bool MBarrierArriveExpectTxOp::getAsmValues(
4244 RewriterBase &rewriter,
4245 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>
4246 &asmValues) {
4247 // Add all the operands but not the attrs to the asmValues list.
4248 // The attrs here are used to generate the right variants for
4249 // intrinsics-lowering. So, we ignore them while generating inline-PTX.
4250 for (auto val : getOperands())
4251 asmValues.push_back({val, mlir::NVVM::PTXRegisterMod::Read});
4252
4253 return false;
4254}
4255
4256mlir::NVVM::IDArgPair MBarrierArriveExpectTxOp::getIntrinsicIDAndArgs(
4257 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4258 auto thisOp = cast<NVVM::MBarrierArriveExpectTxOp>(op);
4259
4260 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4261 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4262 // bit-0: Space
4263 // bit-1: Scope
4264 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4265
4266 // clang-format off
4267 static constexpr llvm::Intrinsic::ID IDs[] = {
4268 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cta,
4269 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cta_space_cluster,
4270 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cta,
4271 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_scope_cluster_space_cluster};
4272 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4273 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cta,
4274 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cta_space_cluster,
4275 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cta,
4276 llvm::Intrinsic::nvvm_mbarrier_arrive_expect_tx_relaxed_scope_cluster_space_cluster};
4277 // clang-format on
4278 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4279
4280 // Tidy-up the Intrinsic Args
4281 llvm::Value *txcount = mt.lookupValue(thisOp.getTxcount());
4282 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4283 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4284 if (needCast)
4285 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4286
4287 return {id, {mbar, txcount}};
4288}
4289
4290mlir::NVVM::IDArgPair MBarrierArriveDropExpectTxOp::getIntrinsicIDAndArgs(
4291 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4292 auto thisOp = cast<NVVM::MBarrierArriveDropExpectTxOp>(op);
4293
4294 bool isClusterSpace = isPtrInSharedClusterSpace(thisOp.getAddr());
4295 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4296 // bit-0: Space
4297 // bit-1: Scope
4298 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isClusterSpace ? 1 : 0);
4299
4300 // clang-format off
4301 static constexpr llvm::Intrinsic::ID IDs[] = {
4302 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cta,
4303 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cta_space_cluster,
4304 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cta,
4305 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_scope_cluster_space_cluster};
4306 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4307 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cta,
4308 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cta_space_cluster,
4309 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cta,
4310 llvm::Intrinsic::nvvm_mbarrier_arrive_drop_expect_tx_relaxed_scope_cluster_space_cluster};
4311 // clang-format on
4312 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4313
4314 // Tidy-up the Intrinsic Args
4315 llvm::Value *txcount = mt.lookupValue(thisOp.getTxcount());
4316 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4317 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4318 if (needCast)
4319 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4320
4321 return {id, {mbar, txcount}};
4322}
4323
4324mlir::NVVM::IDArgPair MBarrierArriveNocompleteOp::getIntrinsicIDAndArgs(
4325 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4326 auto thisOp = cast<NVVM::MBarrierArriveNocompleteOp>(op);
4327 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());
4328 llvm::Intrinsic::ID id =
4329 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete_shared
4330 : llvm::Intrinsic::nvvm_mbarrier_arrive_noComplete;
4331 // Fill the Intrinsic Args
4333 args.push_back(mt.lookupValue(thisOp.getAddr()));
4334 args.push_back(mt.lookupValue(thisOp.getCount()));
4335
4336 return {id, std::move(args)};
4337}
4338
4339mlir::NVVM::IDArgPair MBarrierArriveDropNocompleteOp::getIntrinsicIDAndArgs(
4340 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4341 auto thisOp = cast<NVVM::MBarrierArriveDropNocompleteOp>(op);
4342 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());
4343 llvm::Intrinsic::ID id =
4344 isShared ? llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete_shared
4345 : llvm::Intrinsic::nvvm_mbarrier_arrive_drop_noComplete;
4346 // Fill the Intrinsic Args
4348 args.push_back(mt.lookupValue(thisOp.getAddr()));
4349 args.push_back(mt.lookupValue(thisOp.getCount()));
4350
4351 return {id, std::move(args)};
4352}
4353
4354mlir::NVVM::IDArgPair MBarrierTestWaitOp::getIntrinsicIDAndArgs(
4355 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4356 auto thisOp = cast<NVVM::MBarrierTestWaitOp>(op);
4357 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
4358 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4359 // bit-0: isPhaseParity
4360 // bit-1: Scope
4361 size_t index = ((isClusterScope ? 1 : 0) << 1) | (isPhaseParity ? 1 : 0);
4362
4363 // clang-format off
4364 static constexpr llvm::Intrinsic::ID IDs[] = {
4365 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cta_space_cta,
4366 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cta_space_cta,
4367 llvm::Intrinsic::nvvm_mbarrier_test_wait_scope_cluster_space_cta,
4368 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_scope_cluster_space_cta};
4369 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4370 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cta_space_cta,
4371 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cta_space_cta,
4372 llvm::Intrinsic::nvvm_mbarrier_test_wait_relaxed_scope_cluster_space_cta,
4373 llvm::Intrinsic::nvvm_mbarrier_test_wait_parity_relaxed_scope_cluster_space_cta};
4374 // clang-format on
4375 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4376
4377 // Tidy-up the Intrinsic Args
4378 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4379 llvm::Value *input = mt.lookupValue(thisOp.getStateOrPhase());
4380 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4381 if (needCast)
4382 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4383
4384 return {id, {mbar, input}};
4385}
4386
4387mlir::NVVM::IDArgPair MBarrierTryWaitOp::getIntrinsicIDAndArgs(
4388 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4389 auto thisOp = cast<NVVM::MBarrierTryWaitOp>(op);
4390 bool isPhaseParity = thisOp.getStateOrPhase().getType().isInteger(32);
4391 bool isClusterScope = thisOp.getScope() == NVVM::MemScopeKind::CLUSTER;
4392 bool hasTicks = static_cast<bool>(thisOp.getTicks());
4393 // bit-0: isPhaseParity
4394 // bit-1: Scope
4395 // bit-2: hasTicks
4396 size_t index = ((hasTicks ? 1 : 0) << 2) | ((isClusterScope ? 1 : 0) << 1) |
4397 (isPhaseParity ? 1 : 0);
4398
4399 // clang-format off
4400 static constexpr llvm::Intrinsic::ID IDs[] = {
4401 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cta_space_cta,
4402 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cta_space_cta,
4403 llvm::Intrinsic::nvvm_mbarrier_try_wait_scope_cluster_space_cta,
4404 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_scope_cluster_space_cta,
4405 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cta_space_cta,
4406 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cta_space_cta,
4407 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_scope_cluster_space_cta,
4408 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_scope_cluster_space_cta};
4409 static constexpr llvm::Intrinsic::ID relaxedIDs[] = {
4410 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cta_space_cta,
4411 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cta_space_cta,
4412 llvm::Intrinsic::nvvm_mbarrier_try_wait_relaxed_scope_cluster_space_cta,
4413 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_relaxed_scope_cluster_space_cta,
4414 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cta_space_cta,
4415 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cta_space_cta,
4416 llvm::Intrinsic::nvvm_mbarrier_try_wait_tl_relaxed_scope_cluster_space_cta,
4417 llvm::Intrinsic::nvvm_mbarrier_try_wait_parity_tl_relaxed_scope_cluster_space_cta};
4418 // clang-format on
4419 auto id = thisOp.getRelaxed() ? relaxedIDs[index] : IDs[index];
4420
4421 // Tidy-up the mbarrier pointer
4422 llvm::Value *mbar = mt.lookupValue(thisOp.getAddr());
4423 bool needCast = isPtrInGenericSpace(thisOp.getAddr());
4424 if (needCast)
4425 mbar = castPtrToAddrSpace(builder, mbar, NVVMMemorySpace::Shared);
4426
4427 // Fill the Intrinsic Args
4429 args.push_back(mbar);
4430 args.push_back(mt.lookupValue(thisOp.getStateOrPhase()));
4431 if (hasTicks)
4432 args.push_back(mt.lookupValue(thisOp.getTicks()));
4433
4434 return {id, std::move(args)};
4435}
4436
4437mlir::NVVM::IDArgPair CpAsyncMBarrierArriveOp::getIntrinsicIDAndArgs(
4438 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4439 auto thisOp = cast<NVVM::CpAsyncMBarrierArriveOp>(op);
4440 bool isShared = isPtrInSharedCTASpace(thisOp.getAddr());
4441
4442 llvm::Intrinsic::ID id;
4443 if (thisOp.getNoinc()) {
4444 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc_shared
4445 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_noinc;
4446 } else {
4447 id = isShared ? llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive_shared
4448 : llvm::Intrinsic::nvvm_cp_async_mbarrier_arrive;
4449 }
4450
4451 return {id, {mt.lookupValue(thisOp.getAddr())}};
4452}
4453
4455MovMatrixOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
4456 llvm::IRBuilderBase &builder) {
4457 auto thisOp = cast<NVVM::MovMatrixOp>(op);
4458 return {llvm::Intrinsic::nvvm_movmatrix_sync_aligned_m8n8_trans_b16,
4459 {mt.lookupValue(thisOp.getSrc())}};
4460}
4461
4462#define CP_ASYNC_ID_IMPL(mod, size, suffix) \
4463 llvm::Intrinsic::nvvm_cp_async_##mod##_shared_global_##size##suffix
4464
4465#define GET_CP_ASYNC_ID(mod, size, has_cpsize) \
4466 has_cpsize ? CP_ASYNC_ID_IMPL(mod, size, _s) : CP_ASYNC_ID_IMPL(mod, size, )
4467
4468llvm::Intrinsic::ID
4469CpAsyncOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
4471 llvm::Intrinsic::ID id;
4472
4473 auto cpAsyncOp = cast<NVVM::CpAsyncOp>(op);
4474 bool hasCpSize = static_cast<bool>(cpAsyncOp.getCpSize());
4475 switch (cpAsyncOp.getSize()) {
4476 case 4:
4477 id = GET_CP_ASYNC_ID(ca, 4, hasCpSize);
4478 break;
4479 case 8:
4480 id = GET_CP_ASYNC_ID(ca, 8, hasCpSize);
4481 break;
4482 case 16:
4483 id = (cpAsyncOp.getModifier() == NVVM::LoadCacheModifierKind::CG)
4484 ? GET_CP_ASYNC_ID(cg, 16, hasCpSize)
4485 : GET_CP_ASYNC_ID(ca, 16, hasCpSize);
4486 break;
4487 default:
4488 llvm_unreachable("Invalid copy size in CpAsyncOp.");
4489 }
4490
4491 // Fill the Intrinsic Args
4492 args.push_back(mt.lookupValue(cpAsyncOp.getDst()));
4493 args.push_back(mt.lookupValue(cpAsyncOp.getSrc()));
4494 if (hasCpSize)
4495 args.push_back(mt.lookupValue(cpAsyncOp.getCpSize()));
4496
4497 return id;
4498}
4499
4500mlir::NVVM::IDArgPair CpAsyncBulkPrefetchOp::getIntrinsicIDAndArgs(
4501 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4502 auto thisOp = cast<NVVM::CpAsyncBulkPrefetchOp>(op);
4504 llvm::Intrinsic::ID id = llvm::Intrinsic::nvvm_cp_async_bulk_prefetch_L2;
4505
4506 // Fill the Intrinsic Args
4507 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4508 args.push_back(mt.lookupValue(thisOp.getSize()));
4509
4510 mlir::Value cacheHint = thisOp.getL2CacheHint();
4511 const bool hasCacheHint = static_cast<bool>(cacheHint);
4512 llvm::Value *i64Unused =
4513 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);
4514 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);
4515 args.push_back(builder.getInt1(hasCacheHint));
4516
4517 return {id, std::move(args)};
4518}
4519
4520mlir::NVVM::IDArgPair CpAsyncBulkGlobalToSharedClusterOp::getIntrinsicIDAndArgs(
4521 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4522 auto thisOp = cast<NVVM::CpAsyncBulkGlobalToSharedClusterOp>(op);
4524
4525 // Fill the Intrinsic Args: dst, mbar, src, size.
4526 args.push_back(mt.lookupValue(thisOp.getDstMem()));
4527 args.push_back(mt.lookupValue(thisOp.getMbar()));
4528 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4529 args.push_back(mt.lookupValue(thisOp.getSize()));
4530
4531 // Multicast mask for shared::cluster only, if available.
4532 mlir::Value multicastMask = thisOp.getMulticastMask();
4533 const bool hasMulticastMask = static_cast<bool>(multicastMask);
4534 const bool isSharedCTA = isPtrInSharedCTASpace(thisOp.getDstMem());
4535 if (!isSharedCTA) {
4536 llvm::Value *i16Unused = llvm::ConstantInt::get(builder.getInt16Ty(), 0);
4537 args.push_back(hasMulticastMask ? mt.lookupValue(multicastMask)
4538 : i16Unused);
4539 }
4540
4541 // Cache hint, if available.
4542 mlir::Value cacheHint = thisOp.getL2CacheHint();
4543 const bool hasCacheHint = static_cast<bool>(cacheHint);
4544 llvm::Value *i64Unused = llvm::ConstantInt::get(builder.getInt64Ty(), 0);
4545 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);
4546
4547 // Flag arguments for multicast and cachehint.
4548 if (!isSharedCTA)
4549 args.push_back(builder.getInt1(hasMulticastMask));
4550 args.push_back(builder.getInt1(hasCacheHint));
4551
4552 llvm::Intrinsic::ID id =
4553 isSharedCTA
4554 ? llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cta
4555 : llvm::Intrinsic::nvvm_cp_async_bulk_global_to_shared_cluster;
4556
4557 return {id, std::move(args)};
4558}
4559
4560mlir::NVVM::IDArgPair CpAsyncBulkSharedCTAToGlobalOp::getIntrinsicIDAndArgs(
4561 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4562 auto thisOp = cast<NVVM::CpAsyncBulkSharedCTAToGlobalOp>(op);
4564 llvm::Intrinsic::ID id =
4565 llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global;
4566
4567 // Fill the Intrinsic Args
4568 args.push_back(mt.lookupValue(thisOp.getDstMem()));
4569 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4570 args.push_back(mt.lookupValue(thisOp.getSize()));
4571
4572 mlir::Value cacheHint = thisOp.getL2CacheHint();
4573 const bool hasCacheHint = static_cast<bool>(cacheHint);
4574 llvm::Value *i64Unused =
4575 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);
4576 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);
4577 args.push_back(builder.getInt1(hasCacheHint));
4578
4579 // Choose the bytemask variant
4580 if (mlir::Value byteMask = thisOp.getByteMask()) {
4581 args.push_back(mt.lookupValue(byteMask));
4582 id = llvm::Intrinsic::nvvm_cp_async_bulk_shared_cta_to_global_bytemask;
4583 }
4584
4585 return {id, std::move(args)};
4586}
4587
4588bool CpAsyncBulkTensorGlobalToSharedClusterOp::getAsmValues(
4589 RewriterBase &rewriter,
4590 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>
4591 &asmValues) {
4592 // Add all the operands but not the attrs to the asmValues list.
4593 // The attrs here are used to generate the right variants for
4594 // intrinsics-lowering. So, we ignore them while generating inline-PTX.
4595 for (auto val : getOperands())
4596 asmValues.push_back({val, mlir::NVVM::PTXRegisterMod::Read});
4597
4598 return false;
4599}
4600
4602CpAsyncBulkTensorGlobalToSharedClusterOp::getIntrinsicIDAndArgs(
4603 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4604 auto thisOp = cast<NVVM::CpAsyncBulkTensorGlobalToSharedClusterOp>(op);
4605 const bool isCTAOnly = thisOp.getIsCTAOnly();
4607
4608 // Fill the Intrinsic Args
4609 args.push_back(mt.lookupValue(thisOp.getDstMem()));
4610 args.push_back(mt.lookupValue(thisOp.getMbar()));
4611 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4612
4613 // Coordinates and im2col-offsets
4614 for (mlir::Value v : thisOp.getCoordinates())
4615 args.push_back(mt.lookupValue(v));
4616 for (mlir::Value v : thisOp.getIm2colOffsets())
4617 args.push_back(mt.lookupValue(v));
4618
4619 // MulticastMask, if available
4620 mlir::Value mcMask = thisOp.getMulticastMask();
4621 const bool hasMC = static_cast<bool>(mcMask);
4622 llvm::Value *i16Zero =
4623 llvm::ConstantInt::get(llvm::Type::getInt16Ty(mt.getLLVMContext()), 0);
4624
4625 // CacheHint, if available
4626 mlir::Value cacheHint = thisOp.getL2CacheHint();
4627 const bool hasCacheHint = static_cast<bool>(cacheHint);
4628 llvm::Value *i64Zero =
4629 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);
4630
4631 // Flag argument CTAGroup
4632 // CTA_1/2 is mapped to values 1 and 2 for the intrinsics.
4633 // Hence, the +1 to getGroup().
4634 const int32_t val =
4635 thisOp.getGroup() ? (static_cast<int32_t>(*thisOp.getGroup()) + 1) : 0;
4636 llvm::Value *cg =
4637 llvm::ConstantInt::get(llvm::Type::getInt32Ty(mt.getLLVMContext()), val);
4638
4639 if (!isCTAOnly) {
4640 // For shared::cluster, all the arguments that we build are applicable.
4641 args.push_back(hasMC ? mt.lookupValue(mcMask) : i16Zero);
4642 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Zero);
4643 args.push_back(builder.getInt1(hasMC));
4644 args.push_back(builder.getInt1(hasCacheHint));
4645 args.push_back(cg);
4646 } else {
4647 // For shared::cta, only cache-hint is applicable.
4648 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Zero);
4649 args.push_back(builder.getInt1(hasCacheHint));
4650 }
4651
4652 constexpr size_t numDims = 5; // 1D to 5D
4653 constexpr size_t numModes = 5; // Tile, Im2col, w, w_128, gather4
4654 using rowTy = std::array<llvm::Intrinsic::ID, numDims + 1>;
4655 using TableTy = std::array<rowTy, numModes>;
4656 static constexpr TableTy IDTable{
4657 {{notIntrinsic, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_1d,
4658 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_2d,
4659 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_3d,
4660 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_4d,
4661 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_5d},
4663 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_3d,
4664 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_4d,
4665 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_5d},
4667 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_3d,
4668 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_4d,
4669 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_5d},
4671 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_3d,
4672 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_4d,
4673 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_w_128_5d},
4675 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_gather4_2d}}};
4676
4677 static constexpr TableTy IDTableCTA{
4678 {{notIntrinsic,
4679 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_1d,
4680 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_2d,
4681 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_3d,
4682 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_4d,
4683 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_5d},
4685 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_3d,
4686 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_4d,
4687 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_5d},
4689 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_3d,
4690 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_4d,
4691 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_5d},
4693 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_3d,
4694 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_4d,
4695 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_im2col_w_128_5d},
4697 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_tile_gather4_2d}}};
4698
4699 static_assert(
4700 (getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1) &&
4701 (getMaxEnumValForTMALoadMode() == std::size(IDTableCTA) - 1),
4702 "TMALoadModes must match number of rows in IDTable and IDTableCTA");
4703 size_t mode = static_cast<size_t>(thisOp.getMode());
4704 size_t dim = thisOp.getCoordinates().size();
4705 auto id = isCTAOnly ? IDTableCTA[mode][dim] : IDTable[mode][dim];
4706 assert(id != notIntrinsic &&
4707 "Invalid intrinsic for CpAsyncBulkTensorGlobalToSharedClusterOp.");
4708
4709 return {id, std::move(args)};
4710}
4711
4712mlir::NVVM::IDArgPair CpAsyncBulkTensorPrefetchOp::getIntrinsicIDAndArgs(
4713 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4714 auto thisOp = cast<NVVM::CpAsyncBulkTensorPrefetchOp>(op);
4716
4717 // Fill the Intrinsic Args
4718 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4719
4720 for (auto v : thisOp.getCoordinates())
4721 args.push_back(mt.lookupValue(v));
4722 for (auto v : thisOp.getIm2colOffsets())
4723 args.push_back(mt.lookupValue(v));
4724
4725 mlir::Value cacheHint = thisOp.getL2CacheHint();
4726 const bool hasCacheHint = static_cast<bool>(cacheHint);
4727 llvm::Value *i64Unused =
4728 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);
4729 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);
4730 args.push_back(builder.getInt1(hasCacheHint));
4731
4732 const unsigned NI = llvm::Intrinsic::not_intrinsic;
4733 static constexpr llvm::Intrinsic::ID IDTable[][6] = {
4734 {NI, llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_1d,
4735 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_2d,
4736 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_3d,
4737 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_4d,
4738 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_5d},
4739 {NI, NI, NI,
4740 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_3d,
4741 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_4d,
4742 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_5d},
4743 {NI, NI, NI,
4744 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_3d,
4745 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_4d,
4746 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_5d},
4747 {NI, NI, NI,
4748 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_3d,
4749 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_4d,
4750 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_im2col_w_128_5d},
4751 {NI, NI, NI, NI, NI,
4752 llvm::Intrinsic::nvvm_cp_async_bulk_tensor_prefetch_tile_gather4_2d}};
4753
4754 static_assert(getMaxEnumValForTMALoadMode() == std::size(IDTable) - 1,
4755 "TMALoadModes must match number of rows in IDTable");
4756 size_t mode = static_cast<size_t>(thisOp.getMode());
4757 size_t dim = thisOp.getCoordinates().size();
4758 llvm::Intrinsic::ID id = IDTable[mode][dim];
4759 if (id == llvm::Intrinsic::not_intrinsic)
4760 llvm_unreachable("Invalid intrinsic for CpAsyncBulkTensorPrefetchOp.");
4761
4762 return {id, std::move(args)};
4763}
4764
4766CpAsyncBulkTensorSharedCTAToGlobalOp::getIntrinsicIDAndArgs(
4767 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4768 auto thisOp = cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOp>(op);
4770
4771 // Fill the Intrinsic Args
4772 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4773 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4774
4775 for (auto v : thisOp.getCoordinates())
4776 args.push_back(mt.lookupValue(v));
4777
4778 mlir::Value cacheHint = thisOp.getL2CacheHint();
4779 const bool hasCacheHint = static_cast<bool>(cacheHint);
4780 llvm::Value *i64Unused =
4781 llvm::ConstantInt::get(llvm::Type::getInt64Ty(mt.getLLVMContext()), 0);
4782 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint) : i64Unused);
4783 args.push_back(builder.getInt1(hasCacheHint));
4784
4785 using namespace llvm::Intrinsic;
4786 const unsigned NI = not_intrinsic;
4787 static constexpr ID IDTable[][6] = {
4788 {NI, nvvm_cp_async_bulk_tensor_s2g_tile_1d,
4789 nvvm_cp_async_bulk_tensor_s2g_tile_2d,
4790 nvvm_cp_async_bulk_tensor_s2g_tile_3d,
4791 nvvm_cp_async_bulk_tensor_s2g_tile_4d,
4792 nvvm_cp_async_bulk_tensor_s2g_tile_5d},
4793 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_3d,
4794 nvvm_cp_async_bulk_tensor_s2g_im2col_4d,
4795 nvvm_cp_async_bulk_tensor_s2g_im2col_5d},
4796 {NI, NI, NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_2d},
4797 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_3d,
4798 nvvm_cp_async_bulk_tensor_s2g_im2col_w_4d,
4799 nvvm_cp_async_bulk_tensor_s2g_im2col_w_5d}};
4800
4801 static_assert(getMaxEnumValForTMAStoreMode() == std::size(IDTable) - 1,
4802 "TMAStoreModes must match number of rows in IDTable");
4803 size_t mode = static_cast<size_t>(thisOp.getMode());
4804 size_t dim = thisOp.getCoordinates().size();
4805 ID id = IDTable[mode][dim];
4806 if (id == llvm::Intrinsic::not_intrinsic)
4807 llvm_unreachable(
4808 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOp.");
4809
4810 return {id, std::move(args)};
4811}
4812
4814CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::getIntrinsicIDAndArgs(
4815 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4816 auto thisOp =
4817 cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp>(op);
4818
4820 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4821 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4822 args.push_back(mt.lookupValue(thisOp.getOverrideAddr()));
4823 for (Value v : thisOp.getTensorSize())
4824 args.push_back(mt.lookupValue(v));
4825 for (Value v : thisOp.getLowerStride())
4826 args.push_back(mt.lookupValue(v));
4827 if (thisOp.getUpperStride())
4828 args.push_back(mt.lookupValue(thisOp.getUpperStride()));
4829 for (Value v : thisOp.getCoordinates())
4830 args.push_back(mt.lookupValue(v));
4831
4832 mlir::Value cacheHint = thisOp.getL2CacheHint();
4833 const bool hasCacheHint = static_cast<bool>(cacheHint);
4834 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint)
4835 : builder.getInt64(0));
4836 args.push_back(builder.getInt1(hasCacheHint));
4837
4838 using namespace llvm::Intrinsic;
4839 const unsigned NI = not_intrinsic;
4840 // clang-format off
4841 // override_addr variants, indexed [mode][dim].
4842 static constexpr ID IDTable[][6] = {
4843 {NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_1d,
4844 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_2d,
4845 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_3d,
4846 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_4d,
4847 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_5d},
4848 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_3d,
4849 nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_4d,
4850 nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_5d},
4851 {NI, NI, NI, NI, NI,
4852 nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_override_addr_2d},
4853 {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_3d,
4854 nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_4d,
4855 nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_5d}};
4856
4857 // Tile-only override_addr_dim (1D) / override_addr_dim_stride (2D-5D)
4858 // variants, indexed [dim].
4859 static constexpr ID dimStrideIDTable[] = {
4860 NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_1d,
4861 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_2d,
4862 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_3d,
4863 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_4d,
4864 nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_5d};
4865 // clang-format on
4866
4867 size_t mode = static_cast<size_t>(thisOp.getMode());
4868 size_t dim = thisOp.getCoordinates().size();
4869 bool isDimStride = !thisOp.getTensorSize().empty();
4870
4871 assert(mode < std::size(IDTable) &&
4872 "Invalid mode for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4873 assert(dim < std::size(IDTable[mode]) && dim < std::size(dimStrideIDTable) &&
4874 "Invalid dim for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4875
4876 ID intrinsicID = isDimStride ? dimStrideIDTable[dim] : IDTable[mode][dim];
4877 assert(
4878 intrinsicID != NI &&
4879 "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
4880 return {intrinsicID, std::move(args)};
4881}
4882
4883NVVM::IDArgPair CpAsyncBulkTensorReduceOp::getIntrinsicIDAndArgs(
4884 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4885 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOp>(op);
4886
4888 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4889 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4890 for (Value v : thisOp.getCoordinates())
4891 args.push_back(mt.lookupValue(v));
4892
4893 mlir::Value cacheHint = thisOp.getL2CacheHint();
4894 const bool hasCacheHint = static_cast<bool>(cacheHint);
4895 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint)
4896 : builder.getInt64(0));
4897 args.push_back(builder.getInt32(static_cast<uint32_t>(thisOp.getRedKind())));
4898 args.push_back(builder.getInt1(hasCacheHint));
4899
4900 using namespace llvm::Intrinsic;
4901 const unsigned NI = not_intrinsic;
4902 static constexpr ID IDTable[][6] = {
4903 {NI, nvvm_cp_async_bulk_tensor_reduce_tile_1d,
4904 nvvm_cp_async_bulk_tensor_reduce_tile_2d,
4905 nvvm_cp_async_bulk_tensor_reduce_tile_3d,
4906 nvvm_cp_async_bulk_tensor_reduce_tile_4d,
4907 nvvm_cp_async_bulk_tensor_reduce_tile_5d},
4908 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_3d,
4909 nvvm_cp_async_bulk_tensor_reduce_im2col_4d,
4910 nvvm_cp_async_bulk_tensor_reduce_im2col_5d},
4911 {NI, NI, NI, NI, NI, NI}, // scatter4 not supported for reduce
4912 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_3d,
4913 nvvm_cp_async_bulk_tensor_reduce_im2col_w_4d,
4914 nvvm_cp_async_bulk_tensor_reduce_im2col_w_5d}};
4915
4916 size_t mode = static_cast<size_t>(thisOp.getMode());
4917 size_t dim = thisOp.getCoordinates().size();
4918 assert(mode < std::size(IDTable) &&
4919 "Invalid mode for CpAsyncBulkTensorReduceOp");
4920 assert(dim < std::size(IDTable[mode]) &&
4921 "Invalid dim for CpAsyncBulkTensorReduceOp");
4922
4923 ID intrinsicID = IDTable[mode][dim];
4924 assert(intrinsicID != NI &&
4925 "Invalid intrinsic for CpAsyncBulkTensorReduceOp");
4926 return {intrinsicID, std::move(args)};
4927}
4928
4929NVVM::IDArgPair CpAsyncBulkTensorReduceOverrideAddrOp::getIntrinsicIDAndArgs(
4930 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
4931 auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOverrideAddrOp>(op);
4932
4934 args.push_back(mt.lookupValue(thisOp.getSrcMem()));
4935 args.push_back(mt.lookupValue(thisOp.getTmaDescriptor()));
4936 args.push_back(mt.lookupValue(thisOp.getOverrideAddr()));
4937
4938 for (Value v : thisOp.getTensorSize())
4939 args.push_back(mt.lookupValue(v));
4940 for (Value v : thisOp.getLowerStride())
4941 args.push_back(mt.lookupValue(v));
4942 if (thisOp.getUpperStride())
4943 args.push_back(mt.lookupValue(thisOp.getUpperStride()));
4944 for (Value v : thisOp.getCoordinates())
4945 args.push_back(mt.lookupValue(v));
4946
4947 mlir::Value cacheHint = thisOp.getL2CacheHint();
4948 const bool hasCacheHint = static_cast<bool>(cacheHint);
4949 args.push_back(hasCacheHint ? mt.lookupValue(cacheHint)
4950 : builder.getInt64(0));
4951 args.push_back(builder.getInt32(static_cast<uint32_t>(thisOp.getRedKind())));
4952 args.push_back(builder.getInt1(hasCacheHint));
4953
4954 using namespace llvm::Intrinsic;
4955 const unsigned NI = not_intrinsic;
4956 // clang-format off
4957// override_addr variants, indexed [mode][dim].
4958static constexpr ID IDTable[][6] = {
4959 {NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_1d,
4960 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_2d,
4961 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_3d,
4962 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_4d,
4963 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_5d},
4964 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_3d,
4965 nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_4d,
4966 nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_5d},
4967 {NI, NI, NI, NI, NI, NI}, // scatter4 not supported for reduce
4968 {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_3d,
4969 nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_4d,
4970 nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_5d}};
4971
4972// Tile-only override_addr_dim (1D) / override_addr_dim_stride (2D-5D)
4973// variants, indexed [dim].
4974static constexpr ID dimStrideIDTable[] = {
4975 NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_1d,
4976 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_2d,
4977 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_3d,
4978 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_4d,
4979 nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_stride_5d};
4980 // clang-format on
4981
4982 size_t mode = static_cast<size_t>(thisOp.getMode());
4983 size_t dim = thisOp.getCoordinates().size();
4984 bool isDimStride = !thisOp.getTensorSize().empty();
4985
4986 assert(mode < std::size(IDTable) &&
4987 "Invalid mode for CpAsyncBulkTensorReduceOverrideAddrOp");
4988 assert(dim < std::size(IDTable[mode]) && dim < std::size(dimStrideIDTable) &&
4989 "Invalid dim for CpAsyncBulkTensorReduceOverrideAddrOp");
4990
4991 ID intrinsicID = isDimStride ? dimStrideIDTable[dim] : IDTable[mode][dim];
4992 assert(intrinsicID != NI &&
4993 "Invalid intrinsic for CpAsyncBulkTensorReduceOverrideAddrOp");
4994 return {intrinsicID, std::move(args)};
4995}
4996
4997#define _none
4998
4999#define CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
5000 hasRelu ? llvm::Intrinsic::nvvm_f2tf32_##rnd##relu##sf \
5001 : llvm::Intrinsic::nvvm_f2tf32_##rnd##sf
5002
5003#define GET_CVT_F2TF32_ID(rnd, relu, sf) \
5004 hasSatFinite ? CVT_F2TF32_ID_IMPL(rnd, relu, sf) \
5005 : CVT_F2TF32_ID_IMPL(rnd, relu, )
5006
5007llvm::Intrinsic::ID
5008ConvertFloatToTF32Op::getIntrinsicID(NVVM::FPRoundingMode rnd,
5009 NVVM::SaturationMode sat, bool hasRelu) {
5010 using RndMode = NVVM::FPRoundingMode;
5011 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5012 switch (rnd) {
5013 case RndMode::RN:
5014 return GET_CVT_F2TF32_ID(rn, _relu, _satfinite);
5015 case RndMode::RZ:
5016 return GET_CVT_F2TF32_ID(rz, _relu, _satfinite);
5017 case RndMode::RNA:
5018 return GET_CVT_F2TF32_ID(rna, _none, _satfinite);
5019 default:
5020 llvm_unreachable("Invalid RoundingMode for CvtFloatToTF32Op");
5021 }
5022}
5023
5025ConvertF32x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF4x2Op op,
5027 llvm::IRBuilderBase &builder) {
5029 args.push_back(mt.lookupValue(op.getA()));
5030 args.push_back(mt.lookupValue(op.getB()));
5031
5032 bool hasRelu = op.getRelu();
5033
5034 llvm::Intrinsic::ID intId =
5035 hasRelu ? llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_relu_satfinite
5036 : llvm::Intrinsic::nvvm_ff_to_e2m1x2_rn_satfinite;
5037
5038 return {intId, std::move(args)};
5039}
5040
5041#define GET_F32x2_TO_F6x2_ID(type, has_relu) \
5042 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu_satfinite \
5043 : llvm::Intrinsic::nvvm_ff_to_##type##_rn_satfinite
5044
5045llvm::Intrinsic::ID ConvertF32x2ToF6x2Op::getIntrinsicID(mlir::Type dstTy,
5046 bool hasRelu) {
5048 .Case([&](mlir::Float6E2M3FNType) {
5049 return GET_F32x2_TO_F6x2_ID(e2m3x2, hasRelu);
5050 })
5051 .Case([&](mlir::Float6E3M2FNType) {
5052 return GET_F32x2_TO_F6x2_ID(e3m2x2, hasRelu);
5053 })
5054 .Default([](mlir::Type) {
5055 llvm_unreachable("Invalid conversion in ConvertF32x2ToF6x2Op");
5056 return llvm::Intrinsic::not_intrinsic;
5057 });
5058}
5059
5061ConvertF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF16x2ToF4x2Op &op,
5063 llvm::IRBuilderBase &builder) {
5064 mlir::Type dstTy = op.getDstTy();
5065 bool hasRelu = op.getRelu();
5066
5067 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
5068
5069 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
5070 intId = hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_relu_satfinite
5071 : llvm::Intrinsic::nvvm_f16x2_to_e2m1x2_rn_satfinite;
5072
5074 args.push_back(mt.lookupValue(op.getSrc()));
5075
5076 return {intId, std::move(args)};
5077}
5078
5080ConvertBF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertBF16x2ToF4x2Op &op,
5082 llvm::IRBuilderBase &builder) {
5083 mlir::Type dstTy = op.getDstTy();
5084 bool hasRelu = op.getRelu();
5085
5086 llvm::Intrinsic::ID intId = llvm::Intrinsic::not_intrinsic;
5087
5088 if (llvm::isa<mlir::Float4E2M1FNType>(dstTy))
5089 intId = hasRelu ? llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_relu_satfinite
5090 : llvm::Intrinsic::nvvm_bf16x2_to_e2m1x2_rn_satfinite;
5091
5093 args.push_back(mt.lookupValue(op.getSrc()));
5094
5095 return {intId, std::move(args)};
5096}
5097
5098llvm::Intrinsic::ID ConvertF16x2ToF6x2Op::getIntrinsicID(mlir::Type dstTy,
5099 bool hasRelu) {
5101 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
5102 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_relu_satfinite
5103 : llvm::Intrinsic::nvvm_f16x2_to_e2m3x2_rn_satfinite;
5104 })
5105 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
5106 return hasRelu ? llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_relu_satfinite
5107 : llvm::Intrinsic::nvvm_f16x2_to_e3m2x2_rn_satfinite;
5108 })
5109 .Default([](mlir::Type) {
5110 llvm_unreachable("Invalid conversion in ConvertF16x2ToF6x2Op");
5111 return llvm::Intrinsic::not_intrinsic;
5112 });
5113}
5114
5115llvm::Intrinsic::ID ConvertBF16x2ToF6x2Op::getIntrinsicID(mlir::Type dstTy,
5116 bool hasRelu) {
5118 .Case<mlir::Float6E2M3FNType>([&](mlir::Float6E2M3FNType) {
5119 return hasRelu
5120 ? llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_relu_satfinite
5121 : llvm::Intrinsic::nvvm_bf16x2_to_e2m3x2_rn_satfinite;
5122 })
5123 .Case<mlir::Float6E3M2FNType>([&](mlir::Float6E3M2FNType) {
5124 return hasRelu
5125 ? llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_relu_satfinite
5126 : llvm::Intrinsic::nvvm_bf16x2_to_e3m2x2_rn_satfinite;
5127 })
5128 .Default([](mlir::Type) {
5129 llvm_unreachable("Invalid conversion in ConvertBF16x2ToF6x2Op");
5130 return llvm::Intrinsic::not_intrinsic;
5131 });
5132}
5133
5134#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf) \
5135 has_satf ? llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd##_satfinite \
5136 : llvm::Intrinsic::nvvm_ff_to_ue8m0x2_##rnd
5137
5138#define GET_F32x2_TO_F8X2_S_ID(type, has_relu) \
5139 has_relu ? llvm::Intrinsic::nvvm_ff_to_##type##_rn_relu \
5140 : llvm::Intrinsic::nvvm_ff_to_##type##_rn
5141
5142llvm::Intrinsic::ID
5143ConvertF32x2ToF8x2Op::getIntrinsicID(mlir::Type dstTy, NVVM::FPRoundingMode rnd,
5144 NVVM::SaturationMode sat, bool hasRelu) {
5145 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5146 bool hasRoundingModeRZ = (rnd == NVVM::FPRoundingMode::RZ);
5147 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
5148
5150 .Case([&](mlir::Float8E4M3FNType) {
5151 return GET_F32x2_TO_F8X2_S_ID(e4m3x2, hasRelu);
5152 })
5153 .Case([&](mlir::Float8E5M2Type) {
5154 return GET_F32x2_TO_F8X2_S_ID(e5m2x2, hasRelu);
5155 })
5156 .Case([&](mlir::Float8E8M0FNUType) {
5157 if (hasRoundingModeRZ)
5158 return GET_F32x2_TO_F8X2_US_ID(rz, hasSatFinite);
5159 else if (hasRoundingModeRP)
5160 return GET_F32x2_TO_F8X2_US_ID(rp, hasSatFinite);
5161
5162 llvm_unreachable("Invalid conversion in ConvertF32x2ToF8x2Op");
5163 })
5164 .Default([](mlir::Type) {
5165 llvm_unreachable("Invalid conversion in ConvertF32x2ToF8x2Op");
5166 return llvm::Intrinsic::not_intrinsic;
5167 });
5168}
5169
5170#define GET_F16x2_TO_F8X2_ID(type, has_relu) \
5171 has_relu ? llvm::Intrinsic::nvvm_f16x2_to_##type##_rn_relu \
5172 : llvm::Intrinsic::nvvm_f16x2_to_##type##_rn
5173
5174llvm::Intrinsic::ID ConvertF16x2ToF8x2Op::getIntrinsicID(mlir::Type dstTy,
5175 bool hasRelu) {
5177 .Case([&](mlir::Float8E4M3FNType) {
5178 return GET_F16x2_TO_F8X2_ID(e4m3x2, hasRelu);
5179 })
5180 .Case([&](mlir::Float8E5M2Type) {
5181 return GET_F16x2_TO_F8X2_ID(e5m2x2, hasRelu);
5182 })
5183 .Default([](mlir::Type) {
5184 llvm_unreachable("Invalid conversion in ConvertF16x2ToF8x2Op");
5185 return llvm::Intrinsic::not_intrinsic;
5186 });
5187}
5188
5189llvm::Intrinsic::ID
5190ConvertBF16x2ToF8x2Op::getIntrinsicID(mlir::Type dstTy,
5191 NVVM::FPRoundingMode rnd,
5192 NVVM::SaturationMode sat, bool hasRelu) {
5193 bool hasSatFinite = (sat == NVVM::SaturationMode::SATFINITE);
5194
5195 static constexpr llvm::Intrinsic::ID ue8m0x2IDs[] = {
5196 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz,
5197 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp,
5198 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rz_satfinite,
5199 llvm::Intrinsic::nvvm_bf16x2_to_ue8m0x2_rp_satfinite,
5200 };
5201
5203 .Case<mlir::Float8E4M3FNType>([&](mlir::Float8E4M3FNType) {
5204 return hasRelu
5205 ? llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_relu_satfinite
5206 : llvm::Intrinsic::nvvm_bf16x2_to_e4m3x2_rn_satfinite;
5207 })
5208 .Case<mlir::Float8E5M2Type>([&](mlir::Float8E5M2Type) {
5209 return hasRelu
5210 ? llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_relu_satfinite
5211 : llvm::Intrinsic::nvvm_bf16x2_to_e5m2x2_rn_satfinite;
5212 })
5213 .Case<mlir::Float8E8M0FNUType>([&](mlir::Float8E8M0FNUType) {
5214 bool hasRoundingModeRP = (rnd == NVVM::FPRoundingMode::RP);
5215 unsigned index = (hasSatFinite << 1) | hasRoundingModeRP;
5216 return ue8m0x2IDs[index];
5217 })
5218 .Default([](mlir::Type) {
5219 llvm_unreachable("Invalid conversion in ConvertBF16x2ToF8x2Op");
5220 return llvm::Intrinsic::not_intrinsic;
5221 });
5222}
5223
5224NVVM::IDArgPair ConvertF8x2ToF16x2Op::getIntrinsicIDAndArgs(
5225 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5226 auto curOp = cast<NVVM::ConvertF8x2ToF16x2Op>(op);
5227
5228 bool hasRelu = curOp.getRelu();
5229
5230 llvm::Intrinsic::ID intId =
5232 .Case([&](Float8E4M3FNType type) {
5233 return hasRelu ? llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn_relu
5234 : llvm::Intrinsic::nvvm_e4m3x2_to_f16x2_rn;
5235 })
5236 .Case([&](Float8E5M2Type type) {
5237 return hasRelu ? llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn_relu
5238 : llvm::Intrinsic::nvvm_e5m2x2_to_f16x2_rn;
5239 })
5240 .Default([](mlir::Type type) {
5241 llvm_unreachable("Invalid type for ConvertF8x2ToF16x2Op");
5242 return llvm::Intrinsic::not_intrinsic;
5243 });
5244
5245 llvm::Value *packedI16 =
5246 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),
5247 llvm::Type::getInt16Ty(builder.getContext()));
5248
5249 return {intId, {packedI16}};
5250}
5251
5252NVVM::IDArgPair ConvertF8x2ToBF16x2Op::getIntrinsicIDAndArgs(
5253 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5254 auto curOp = cast<NVVM::ConvertF8x2ToBF16x2Op>(op);
5255 bool hasScale = static_cast<bool>(curOp.getScaleFactor());
5256 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5257 bool hasRelu = curOp.getRelu();
5258
5259 static constexpr llvm::Intrinsic::ID E4M3Ids[] = {
5260 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_scale_n2_ue8m0,
5261 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5262 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5263 llvm::Intrinsic::nvvm_e4m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5264 };
5265
5266 static constexpr llvm::Intrinsic::ID E5M2Ids[] = {
5267 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_scale_n2_ue8m0,
5268 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5269 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5270 llvm::Intrinsic::nvvm_e5m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5271 };
5272
5273 llvm::Intrinsic::ID intId =
5275 .Case([&](Float8E8M0FNUType type) {
5276 return llvm::Intrinsic::nvvm_ue8m0x2_to_bf16x2;
5277 })
5278 .Case([&](Float8E4M3FNType type) {
5279 return E4M3Ids[hasSatfinite << 1 | hasRelu];
5280 })
5281 .Case([&](Float8E5M2Type type) {
5282 return E5M2Ids[hasSatfinite << 1 | hasRelu];
5283 })
5284 .Default([](mlir::Type type) {
5285 llvm_unreachable("Invalid type for ConvertF8x2ToBF16x2Op");
5286 return llvm::Intrinsic::not_intrinsic;
5287 });
5288 llvm::Value *packedI16 =
5289 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),
5290 llvm::Type::getInt16Ty(builder.getContext()));
5291
5293 args.push_back(packedI16);
5294 if (!isa<Float8E8M0FNUType>(curOp.getSrcType()))
5295 args.push_back(
5296 hasScale ? mt.lookupValue(curOp.getScaleFactor())
5297 : builder.getInt16(0x7f7f)); // default scale factor (value of
5298 // 1 for both elements)
5299
5300 return {intId, std::move(args)};
5301}
5302
5303NVVM::IDArgPair ConvertF6x2ToF16x2Op::getIntrinsicIDAndArgs(
5304 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5305 auto curOp = cast<NVVM::ConvertF6x2ToF16x2Op>(op);
5306
5307 bool hasRelu = curOp.getRelu();
5308
5309 llvm::Intrinsic::ID intId =
5311 .Case([&](Float6E2M3FNType type) {
5312 return hasRelu ? llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn_relu
5313 : llvm::Intrinsic::nvvm_e2m3x2_to_f16x2_rn;
5314 })
5315 .Case([&](Float6E3M2FNType type) {
5316 return hasRelu ? llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn_relu
5317 : llvm::Intrinsic::nvvm_e3m2x2_to_f16x2_rn;
5318 })
5319 .Default([](mlir::Type type) {
5320 llvm_unreachable("Invalid type for ConvertF6x2ToF16x2Op");
5321 return llvm::Intrinsic::not_intrinsic;
5322 });
5323
5324 llvm::Value *packedI16 =
5325 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),
5326 llvm::Type::getInt16Ty(builder.getContext()));
5327
5328 return {intId, {packedI16}};
5329}
5330
5331NVVM::IDArgPair ConvertF6x2ToBF16x2Op::getIntrinsicIDAndArgs(
5332 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5333 auto curOp = cast<NVVM::ConvertF6x2ToBF16x2Op>(op);
5334 bool hasScale = static_cast<bool>(curOp.getScaleFactor());
5335 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5336 bool hasRelu = curOp.getRelu();
5337
5338 static constexpr llvm::Intrinsic::ID E2M3Ids[] = {
5339 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_scale_n2_ue8m0,
5340 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5341 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5342 llvm::Intrinsic::nvvm_e2m3x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5343 };
5344
5345 static constexpr llvm::Intrinsic::ID E3M2Ids[] = {
5346 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_scale_n2_ue8m0,
5347 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5348 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5349 llvm::Intrinsic::nvvm_e3m2x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5350 };
5351
5352 unsigned idx = (hasSatfinite << 1) | hasRelu;
5353 llvm::Intrinsic::ID intId =
5355 .Case([&](Float6E2M3FNType type) { return E2M3Ids[idx]; })
5356 .Case([&](Float6E3M2FNType type) { return E3M2Ids[idx]; })
5357 .Default([](mlir::Type type) {
5358 llvm_unreachable("Invalid type for ConvertF6x2ToBF16x2Op");
5359 return llvm::Intrinsic::not_intrinsic;
5360 });
5361
5362 llvm::Value *packedI16 =
5363 builder.CreateBitCast(mt.lookupValue(curOp.getSrc()),
5364 llvm::Type::getInt16Ty(builder.getContext()));
5365
5367 args.push_back(packedI16);
5368 args.push_back(
5369 hasScale
5370 ? mt.lookupValue(curOp.getScaleFactor())
5371 : builder.getInt16(
5372 0x7f7f)); // default scale factor (value of 1 for both elements)
5373
5374 return {intId, std::move(args)};
5375}
5376
5377NVVM::IDArgPair ConvertF4x2ToF16x2Op::getIntrinsicIDAndArgs(
5378 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5379 auto curOp = cast<NVVM::ConvertF4x2ToF16x2Op>(op);
5380
5381 bool hasRelu = curOp.getRelu();
5382
5383 llvm::Intrinsic::ID intId =
5385 .Case([&](Float4E2M1FNType type) {
5386 return hasRelu ? llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn_relu
5387 : llvm::Intrinsic::nvvm_e2m1x2_to_f16x2_rn;
5388 })
5389 .Default([](mlir::Type type) {
5390 llvm_unreachable("Invalid type for ConvertF4x2ToF16x2Op");
5391 return llvm::Intrinsic::not_intrinsic;
5392 });
5393
5394 llvm::Value *extendedI16 =
5395 builder.CreateZExt(mt.lookupValue(curOp.getSrc()),
5396 llvm::Type::getInt16Ty(builder.getContext()));
5397
5398 return {intId, {extendedI16}};
5399}
5400
5401NVVM::IDArgPair ConvertF4x2ToBF16x2Op::getIntrinsicIDAndArgs(
5402 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5403 auto curOp = cast<NVVM::ConvertF4x2ToBF16x2Op>(op);
5404 bool hasScale = static_cast<bool>(curOp.getScaleFactor());
5405 bool hasSatfinite = curOp.getSat() == NVVM::SaturationMode::SATFINITE;
5406 bool hasRelu = curOp.getRelu();
5407
5408 static constexpr llvm::Intrinsic::ID E2M1Ids[] = {
5409 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_scale_n2_ue8m0,
5410 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5411 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5412 llvm::Intrinsic::nvvm_e2m1x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5413 };
5414
5415 unsigned idx = (hasSatfinite << 1) | hasRelu;
5416 llvm::Intrinsic::ID intId =
5418 .Case([&](Float4E2M1FNType type) { return E2M1Ids[idx]; })
5419 .Default([](mlir::Type type) {
5420 llvm_unreachable("Invalid type for ConvertF4x2ToBF16x2Op");
5421 return llvm::Intrinsic::not_intrinsic;
5422 });
5423
5424 llvm::Value *extendedI16 =
5425 builder.CreateZExt(mt.lookupValue(curOp.getSrc()),
5426 llvm::Type::getInt16Ty(builder.getContext()));
5427
5429 args.push_back(extendedI16);
5430 args.push_back(
5431 hasScale
5432 ? mt.lookupValue(curOp.getScaleFactor())
5433 : builder.getInt16(
5434 0x7f7f)); // default scale factor (value of 1 for both elements)
5435
5436 return {intId, std::move(args)};
5437}
5438
5439NVVM::IDArgPair ConvertF32x2ToS2F6x2Op::getIntrinsicIDAndArgs(
5440 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5441 auto thisOp = cast<NVVM::ConvertF32x2ToS2F6x2Op>(op);
5442 bool hasRelu = thisOp.getRelu();
5443 bool hasScale = static_cast<bool>(thisOp.getScaleFactor());
5444
5445 llvm::Intrinsic::ID id =
5446 hasRelu
5447 ? llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
5448 : llvm::Intrinsic::nvvm_ff_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
5449
5450 // Fill the Intrinsic Args
5452 args.push_back(mt.lookupValue(thisOp.getA()));
5453 args.push_back(mt.lookupValue(thisOp.getB()));
5454 args.push_back(hasScale ? mt.lookupValue(thisOp.getScaleFactor())
5455 : builder.getInt16(0x7f7f));
5456 return {id, std::move(args)};
5457}
5458
5459NVVM::IDArgPair ConvertBF16x2ToS2F6x2Op::getIntrinsicIDAndArgs(
5460 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5461 auto thisOp = cast<NVVM::ConvertBF16x2ToS2F6x2Op>(op);
5462 bool hasRelu = thisOp.getRelu();
5463 bool hasScale = static_cast<bool>(thisOp.getScaleFactor());
5464
5465 llvm::Intrinsic::ID id =
5466 hasRelu
5467 ? llvm::Intrinsic::
5468 nvvm_bf16x2_to_s2f6x2_rn_relu_satfinite_scale_n2_ue8m0
5469 : llvm::Intrinsic::nvvm_bf16x2_to_s2f6x2_rn_satfinite_scale_n2_ue8m0;
5470
5471 // Fill the Intrinsic Args
5473 args.push_back(mt.lookupValue(thisOp.getSrc()));
5474 args.push_back(hasScale ? mt.lookupValue(thisOp.getScaleFactor())
5475 : builder.getInt16(0x7f7f));
5476 return {id, std::move(args)};
5477}
5478
5479NVVM::IDArgPair ConvertS2F6x2ToBF16x2Op::getIntrinsicIDAndArgs(
5480 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5481 auto thisOp = cast<NVVM::ConvertS2F6x2ToBF16x2Op>(op);
5482 bool hasRelu = thisOp.getRelu();
5483 bool hasScale = static_cast<bool>(thisOp.getScaleFactor());
5484 bool hasSat = thisOp.getSat() == NVVM::SaturationMode::SATFINITE;
5485
5486 static constexpr llvm::Intrinsic::ID ids[] = {
5487 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_scale_n2_ue8m0,
5488 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_scale_n2_ue8m0,
5489 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_satfinite_scale_n2_ue8m0,
5490 llvm::Intrinsic::nvvm_s2f6x2_to_bf16x2_rn_relu_satfinite_scale_n2_ue8m0,
5491 };
5492
5493 unsigned idx = (hasSat << 1) | hasRelu;
5494
5495 // Fill the Intrinsic Args
5497 llvm::Value *packedI16 =
5498 builder.CreateBitCast(mt.lookupValue(thisOp.getSrc()),
5499 llvm::Type::getInt16Ty(builder.getContext()));
5500 args.push_back(packedI16);
5501 args.push_back(hasScale ? mt.lookupValue(thisOp.getScaleFactor())
5502 : builder.getInt16(0x7f7f));
5503
5504 return {ids[idx], std::move(args)};
5505}
5506
5507llvm::Intrinsic::ID
5508Tcgen05AllocOp::getIntrinsicIDAndArgs(Operation &op,
5511 auto curOp = cast<NVVM::Tcgen05AllocOp>(op);
5512 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
5513
5514 llvm::Intrinsic::ID id = is2CTAMode ? llvm::Intrinsic::nvvm_tcgen05_alloc_cg2
5515 : llvm::Intrinsic::nvvm_tcgen05_alloc_cg1;
5516
5517 // Fill the Intrinsic Args
5518 args.push_back(mt.lookupValue(curOp.getAddr()));
5519 args.push_back(mt.lookupValue(curOp.getNCols()));
5520 args.push_back(llvm::ConstantInt::getFalse(mt.getLLVMContext()));
5521
5522 return id;
5523}
5524
5525llvm::Intrinsic::ID Tcgen05DeallocOp::getIntrinsicIDAndArgs(
5528 auto curOp = cast<NVVM::Tcgen05DeallocOp>(op);
5529 auto id = (curOp.getGroup() == CTAGroupKind::CTA_1)
5530 ? llvm::Intrinsic::nvvm_tcgen05_dealloc_cg1
5531 : llvm::Intrinsic::nvvm_tcgen05_dealloc_cg2;
5532
5533 // Fill the Intrinsic Args
5534 args.push_back(mt.lookupValue(curOp.getTaddr()));
5535 args.push_back(mt.lookupValue(curOp.getNCols()));
5536 args.push_back(llvm::ConstantInt::getFalse(mt.getLLVMContext()));
5537
5538 return id;
5539}
5540
5541llvm::Intrinsic::ID
5542Tcgen05CommitOp::getIntrinsicIDAndArgs(Operation &op,
5545 auto curOp = cast<NVVM::Tcgen05CommitOp>(op);
5546 bool hasMulticast = static_cast<bool>(curOp.getMulticastMask());
5547 bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
5548 bool hasSmemARead = curOp.getSmemARead();
5549 unsigned index = (static_cast<unsigned>(hasSmemARead) << 1) |
5550 static_cast<unsigned>(is2CTAMode);
5551
5552 using namespace llvm::Intrinsic;
5553 static constexpr ID IDs[] = {
5554 nvvm_tcgen05_commit_cg1,
5555 nvvm_tcgen05_commit_cg2,
5556 nvvm_tcgen05_commit_smem_a_read_cg1,
5557 nvvm_tcgen05_commit_smem_a_read_cg2,
5558 };
5559
5560 static constexpr ID multicastIDs[] = {
5561 nvvm_tcgen05_commit_mc_cg1,
5562 nvvm_tcgen05_commit_mc_cg2,
5563 nvvm_tcgen05_commit_smem_a_read_mc_cg1,
5564 nvvm_tcgen05_commit_smem_a_read_mc_cg2,
5565 };
5566
5567 ID id = hasMulticast ? multicastIDs[index] : IDs[index];
5568 // Fill the Intrinsic Args
5569 args.push_back(mt.lookupValue(curOp.getAddr()));
5570 if (hasMulticast)
5571 args.push_back(mt.lookupValue(curOp.getMulticastMask()));
5572
5573 return id;
5574}
5575
5576#define TCGEN05_CP_IMPL(shape_mc, src_fmt, cg) \
5577 llvm::Intrinsic::nvvm_tcgen05_cp##shape_mc##src_fmt##cg
5578
5579#define TCGEN05_CP_2CTA(shape_mc, src_fmt, is_2cta) \
5580 is_2cta ? TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg2) \
5581 : TCGEN05_CP_IMPL(shape_mc, src_fmt, _cg1)
5582
5583#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta) \
5584 [&]() -> auto { \
5585 if ((src_fmt) == Tcgen05CpSrcFormat::B6x16_P32) \
5586 return TCGEN05_CP_2CTA(shape_mc, _b6x16_p32, is_2cta); \
5587 if ((src_fmt) == Tcgen05CpSrcFormat::B4x16_P64) \
5588 return TCGEN05_CP_2CTA(shape_mc, _b4x16_p64, is_2cta); \
5589 return TCGEN05_CP_2CTA(shape_mc, , is_2cta); \
5590 }()
5591
5593ConvertF32x2ToF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF16x2Op &op,
5595 llvm::IRBuilderBase &builder) {
5596 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5597 llvm::Intrinsic::nvvm_ff2f16x2_rn,
5598 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu,
5599 llvm::Intrinsic::nvvm_ff2f16x2_rn_satfinite,
5600 llvm::Intrinsic::nvvm_ff2f16x2_rn_relu_satfinite,
5601 };
5602 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5603 llvm::Intrinsic::nvvm_ff2f16x2_rz,
5604 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu,
5605 llvm::Intrinsic::nvvm_ff2f16x2_rz_satfinite,
5606 llvm::Intrinsic::nvvm_ff2f16x2_rz_relu_satfinite,
5607 };
5608 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5609 llvm::Intrinsic::nvvm_ff2f16x2_rs,
5610 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu,
5611 llvm::Intrinsic::nvvm_ff2f16x2_rs_satfinite,
5612 llvm::Intrinsic::nvvm_ff2f16x2_rs_relu_satfinite,
5613 };
5614
5615 unsigned hasRelu = op.getRelu() ? 1 : 0;
5616 unsigned hasSatFinite =
5617 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5618 // idx: bit-0 - relu
5619 // bit-1 - satfinite
5620 unsigned idx = (hasSatFinite << 1) | hasRelu;
5621
5623 args.push_back(mt.lookupValue(op.getSrcHi()));
5624 args.push_back(mt.lookupValue(op.getSrcLo()));
5625 if (op.getRandomBits())
5626 args.push_back(mt.lookupValue(op.getRandomBits()));
5627
5628 // TODO: Add support for PZO modifier
5629 args.push_back(builder.getInt1(false));
5630
5631 switch (op.getRnd()) {
5632 case FPRoundingMode::RN:
5633 return {rndRNIds[idx], std::move(args)};
5634 case FPRoundingMode::RZ:
5635 return {rndRZIds[idx], std::move(args)};
5636 case FPRoundingMode::RS:
5637 return {rndRSIds[idx], std::move(args)};
5638 default:
5639 llvm_unreachable("Invalid rounding mode for ConvertF32x2ToF16x2Op");
5640 }
5641}
5642
5644ConvertF32x2ToBF16x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToBF16x2Op &op,
5646 llvm::IRBuilderBase &builder) {
5647 static constexpr llvm::Intrinsic::ID rndRNIds[] = {
5648 llvm::Intrinsic::nvvm_ff2bf16x2_rn,
5649 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu,
5650 llvm::Intrinsic::nvvm_ff2bf16x2_rn_satfinite,
5651 llvm::Intrinsic::nvvm_ff2bf16x2_rn_relu_satfinite,
5652 };
5653 static constexpr llvm::Intrinsic::ID rndRZIds[] = {
5654 llvm::Intrinsic::nvvm_ff2bf16x2_rz,
5655 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu,
5656 llvm::Intrinsic::nvvm_ff2bf16x2_rz_satfinite,
5657 llvm::Intrinsic::nvvm_ff2bf16x2_rz_relu_satfinite,
5658 };
5659 static constexpr llvm::Intrinsic::ID rndRSIds[] = {
5660 llvm::Intrinsic::nvvm_ff2bf16x2_rs,
5661 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu,
5662 llvm::Intrinsic::nvvm_ff2bf16x2_rs_satfinite,
5663 llvm::Intrinsic::nvvm_ff2bf16x2_rs_relu_satfinite,
5664 };
5665
5666 unsigned hasRelu = op.getRelu() ? 1 : 0;
5667 unsigned hasSatFinite =
5668 (op.getSat() == NVVM::SaturationMode::SATFINITE) ? 1 : 0;
5669 // idx: bit-0 - relu
5670 // bit-1 - satfinite
5671 unsigned idx = (hasSatFinite << 1) | hasRelu;
5672
5674 args.push_back(mt.lookupValue(op.getSrcHi()));
5675 args.push_back(mt.lookupValue(op.getSrcLo()));
5676 if (op.getRandomBits())
5677 args.push_back(mt.lookupValue(op.getRandomBits()));
5678
5679 // TODO: Add support for PZO modifier
5680 args.push_back(builder.getInt1(false));
5681
5682 switch (op.getRnd()) {
5683 case FPRoundingMode::RN:
5684 return {rndRNIds[idx], std::move(args)};
5685 case FPRoundingMode::RZ:
5686 return {rndRZIds[idx], std::move(args)};
5687 case FPRoundingMode::RS:
5688 return {rndRSIds[idx], std::move(args)};
5689 default:
5690 llvm_unreachable("Invalid rounding mode for ConvertF32x2ToBF16x2Op");
5691 }
5692}
5693
5694llvm::Intrinsic::ID ConvertF32x4ToF8x4Op::getIntrinsicID() {
5695 mlir::Type dstTy = getDstTy();
5696 bool hasRelu = getRelu();
5697
5699 .Case([&](mlir::Float8E4M3FNType) {
5700 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite
5701 : llvm::Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite;
5702 })
5703 .Case([&](mlir::Float8E5M2Type) {
5704 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite
5705 : llvm::Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite;
5706 })
5707 .Default([](mlir::Type) {
5708 llvm_unreachable("Invalid F8 type in ConvertF32x4ToF8x4Op");
5709 return llvm::Intrinsic::not_intrinsic;
5710 });
5711}
5712
5713llvm::Intrinsic::ID ConvertF32x4ToF6x4Op::getIntrinsicID() {
5714 mlir::Type dstTy = getDstTy();
5715 bool hasRelu = getRelu();
5716
5718 .Case([&](mlir::Float6E2M3FNType) {
5719 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite
5720 : llvm::Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite;
5721 })
5722 .Case([&](mlir::Float6E3M2FNType) {
5723 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite
5724 : llvm::Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite;
5725 })
5726 .Default([](mlir::Type) {
5727 llvm_unreachable("Invalid F6 type in ConvertF32x4ToF6x4Op");
5728 return llvm::Intrinsic::not_intrinsic;
5729 });
5730}
5731
5732llvm::Intrinsic::ID ConvertF32x4ToF4x4Op::getIntrinsicID() {
5733 mlir::Type dstTy = getDstTy();
5734 bool hasRelu = getRelu();
5735
5737 .Case([&](mlir::Float4E2M1FNType) {
5738 return hasRelu ? llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite
5739 : llvm::Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite;
5740 })
5741 .Default([](mlir::Type) {
5742 llvm_unreachable("Invalid F4 type in ConvertF32x4ToF4x4Op");
5743 return llvm::Intrinsic::not_intrinsic;
5744 });
5745}
5746
5747llvm::Intrinsic::ID Tcgen05CpOp::getIntrinsicID(Operation &op) {
5748 auto curOp = cast<NVVM::Tcgen05CpOp>(op);
5749 bool is2CTA = curOp.getGroup() == CTAGroupKind::CTA_2;
5750 auto srcFmt = curOp.getSrcFormat();
5751 auto mc = curOp.getMulticast();
5752
5753 switch (curOp.getShape()) {
5754 case Tcgen05CpShape::SHAPE_128x256b:
5755 return GET_TCGEN05_CP_ID(_128x256b, srcFmt, is2CTA);
5756 case Tcgen05CpShape::SHAPE_128x128b:
5757 return GET_TCGEN05_CP_ID(_128x128b, srcFmt, is2CTA);
5758 case Tcgen05CpShape::SHAPE_4x256b:
5759 return GET_TCGEN05_CP_ID(_4x256b, srcFmt, is2CTA);
5760 case Tcgen05CpShape::SHAPE_32x128b:
5761 return GET_TCGEN05_CP_ID(_32x128b_warpx4, srcFmt, is2CTA);
5762 case Tcgen05CpShape::SHAPE_64x128b:
5763 return (mc == Tcgen05CpMulticast::WARPX2_01_23)
5764 ? GET_TCGEN05_CP_ID(_64x128b_warpx2_01_23, srcFmt, is2CTA)
5765 : GET_TCGEN05_CP_ID(_64x128b_warpx2_02_13, srcFmt, is2CTA);
5766 }
5767 llvm_unreachable("Invalid shape in tcgen05 cp Op");
5768}
5769
5770// Returns the valid vector length for a given shape and vector length, the
5771// function models the table mentioned in the tcgen05.{ld, st} Op description
5772static unsigned isValidVectorLength(NVVM::Tcgen05LdStShape shape,
5773 unsigned vecLen) {
5774 if (shape == NVVM::Tcgen05LdStShape::SHAPE_16X128B)
5775 return vecLen >= 2;
5776 if (shape == NVVM::Tcgen05LdStShape::SHAPE_16X256B)
5777 return vecLen >= 4;
5778 return true;
5779}
5780
5781LogicalResult Tcgen05LdOp::verify() {
5782 LogicalResult result = success();
5783 if (getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5784 result = emitError("shape 16x32bx2 requires offset argument");
5785
5786 if (getShape() != NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && getOffset())
5787 result = emitError("offset argument is only supported for shape 16x32bx2");
5788
5789 auto resTy = getRes().getType();
5790 unsigned resLen = isa<VectorType>(resTy)
5791 ? llvm::cast<VectorType>(resTy).getNumElements()
5792 : 1;
5793 if (!isValidVectorLength(getShape(), resLen))
5794 result = emitError(llvm::formatv("invalid result type length {0} for shape "
5795 "{1} in tcgen05.ld Op",
5796 resLen, stringifyEnum(getShape())));
5797
5798 return result;
5799}
5800
5801LogicalResult Tcgen05StOp::verify() {
5802 LogicalResult result = success();
5803 if (getShape() == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2 && !getOffset())
5804 result = emitError("shape 16x32bx2 requires offset argument");
5805
5806 auto valTy = getVal().getType();
5807 unsigned valLen = isa<VectorType>(valTy)
5808 ? llvm::cast<VectorType>(valTy).getNumElements()
5809 : 1;
5810 if (!isValidVectorLength(getShape(), valLen))
5811 result = emitError(llvm::formatv("invalid input length {0} for shape "
5812 "{1} in tcgen05.st Op",
5813 valLen, stringifyEnum(getShape())));
5814
5815 return result;
5816}
5817
5818/// Infer the result ranges for the NVVM SpecialRangeableRegisterOp that might
5819/// have ConstantRangeAttr.
5822 SetIntRangeFn setResultRanges) {
5823 if (auto rangeAttr = op->getAttrOfType<LLVM::ConstantRangeAttr>("range")) {
5824 setResultRanges(result, {rangeAttr.getLower(), rangeAttr.getUpper(),
5825 rangeAttr.getLower(), rangeAttr.getUpper()});
5826 } else {
5827 setResultRanges(result, IntegerValueRange::getMaxRange(result).getValue());
5828 }
5829}
5830
5831/// Verify the range attribute satisfies LLVM ConstantRange constructor
5832/// requirements for NVVM SpecialRangeableRegisterOp.
5833static LogicalResult
5835 std::optional<LLVM::ConstantRangeAttr> rangeAttr) {
5836 if (!rangeAttr)
5837 return success();
5838
5839 const llvm::APInt &lower = rangeAttr->getLower();
5840 const llvm::APInt &upper = rangeAttr->getUpper();
5841
5842 // Check LLVM ConstantRange constructor condition
5843 if (lower == upper && !lower.isMaxValue() && !lower.isMinValue()) {
5844 unsigned bitWidth = lower.getBitWidth();
5845 llvm::APInt minVal = llvm::APInt::getMinValue(bitWidth);
5846 llvm::APInt maxVal = llvm::APInt::getMaxValue(bitWidth);
5847 return op->emitOpError(
5848 "invalid range attribute: Lower == Upper, but they aren't min (")
5849 << llvm::toString(minVal, 10, false) << ") or max ("
5850 << llvm::toString(maxVal, 10, false)
5851 << ") value! This is an invalid constant range.";
5852 }
5853
5854 return success();
5855}
5856
5857static llvm::Value *getAsPackedI32(llvm::Value *arg,
5858 llvm::IRBuilderBase &builder) {
5859 return builder.CreateBitCast(arg,
5860 llvm::Type::getInt32Ty(builder.getContext()));
5861}
5862
5863NVVM::IDArgPair DotAccumulate4WayOp::getIntrinsicIDAndArgs(
5864 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5865 auto curOp = cast<NVVM::DotAccumulate4WayOp>(op);
5866
5868 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getA()), builder));
5869 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getB()), builder));
5870 args.push_back(mt.lookupValue(curOp.getC()));
5871
5872 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5873 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5874 unsigned type = (isASigned << 1) | isBSigned;
5875 const llvm::Intrinsic::ID ids[] = {
5876 llvm::Intrinsic::nvvm_idp4a_u_u,
5877 llvm::Intrinsic::nvvm_idp4a_u_s,
5878 llvm::Intrinsic::nvvm_idp4a_s_u,
5879 llvm::Intrinsic::nvvm_idp4a_s_s,
5880 };
5881 return {ids[type], args};
5882}
5883
5884NVVM::IDArgPair DotAccumulate2WayOp::getIntrinsicIDAndArgs(
5885 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5886 auto curOp = cast<NVVM::DotAccumulate2WayOp>(op);
5887
5889 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getA()), builder));
5890 args.push_back(getAsPackedI32(mt.lookupValue(curOp.getB()), builder));
5891 args.push_back(builder.getInt1(curOp.getBHi()));
5892 args.push_back(mt.lookupValue(curOp.getC()));
5893
5894 bool isASigned = curOp.getAType() == NVVM::DotAccumulateType::SIGNED;
5895 bool isBSigned = curOp.getBType() == NVVM::DotAccumulateType::SIGNED;
5896 unsigned type = (isASigned << 1) | isBSigned;
5897 const llvm::Intrinsic::ID ids[] = {
5898 llvm::Intrinsic::nvvm_idp2a_u_u,
5899 llvm::Intrinsic::nvvm_idp2a_u_s,
5900 llvm::Intrinsic::nvvm_idp2a_s_u,
5901 llvm::Intrinsic::nvvm_idp2a_s_s,
5902 };
5903 return {ids[type], args};
5904}
5905
5906static llvm::Value *getParamCastedAddr(llvm::Value *addr,
5907 llvm::IRBuilderBase &builder) {
5908 return builder.CreateAddrSpaceCast(
5909 addr, builder.getPtrTy(llvm::NVPTXAS::ADDRESS_SPACE_ENTRY_PARAM));
5910}
5911
5913PrefetchOp::getIntrinsicIDAndArgs(NVVM::PrefetchOp &op,
5915 llvm::IRBuilderBase &builder) {
5916 using MemSpace = NVVM::NVVMMemorySpace;
5917 using CacheLevel = NVVM::PrefetchCacheLevel;
5918
5919 std::optional<NVVM::PrefetchCacheLevel> cacheLevel = op.getCacheLevel();
5920 std::optional<NVVM::CacheEvictionPriority> evictPriority =
5921 op.getEvictPriority();
5922 unsigned addressSpace =
5923 llvm::cast<LLVM::LLVMPointerType>(op.getAddr().getType())
5924 .getAddressSpace();
5925
5927 llvm::Value *addr = mt.lookupValue(op.getAddr());
5928 args.push_back(op.getInParamSpace() ? getParamCastedAddr(addr, builder)
5929 : addr);
5930
5931 if (op.getTensormap())
5932 return {llvm::Intrinsic::nvvm_prefetch_tensormap, args};
5933
5934 assert(cacheLevel && "expected cache level for non-tensormap prefetch");
5935
5936 if (op.getUniform() && *cacheLevel == CacheLevel::L1)
5937 return {llvm::Intrinsic::nvvm_prefetchu_L1, args};
5938
5939 if (evictPriority && *cacheLevel == CacheLevel::L2) {
5940 switch (*evictPriority) {
5941 case NVVM::CacheEvictionPriority::EvictLast:
5942 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_last, args};
5943 case NVVM::CacheEvictionPriority::EvictNormal:
5944 return {llvm::Intrinsic::nvvm_prefetch_global_L2_evict_normal, args};
5945 default:
5946 llvm_unreachable("Invalid cache eviction priority");
5947 }
5948 }
5949
5950 switch (static_cast<MemSpace>(addressSpace)) {
5951 case MemSpace::Generic:
5952 return *cacheLevel == CacheLevel::L1
5953 ? NVVM::IDArgPair({llvm::Intrinsic::nvvm_prefetch_L1, args})
5954 : NVVM::IDArgPair({llvm::Intrinsic::nvvm_prefetch_L2, args});
5955 case MemSpace::Global:
5956 return *cacheLevel == CacheLevel::L1
5958 {llvm::Intrinsic::nvvm_prefetch_global_L1, args})
5959 : NVVM::IDArgPair(
5960 {llvm::Intrinsic::nvvm_prefetch_global_L2, args});
5961 case MemSpace::Local:
5962 return *cacheLevel == CacheLevel::L1
5964 {llvm::Intrinsic::nvvm_prefetch_local_L1, args})
5965 : NVVM::IDArgPair(
5966 {llvm::Intrinsic::nvvm_prefetch_local_L2, args});
5967 default:
5968 llvm_unreachable("Invalid pointer address space");
5969 }
5970}
5971
5972bool NVVM::InlinePtxOp::getAsmValues(
5973 RewriterBase &rewriter,
5974 llvm::SmallVectorImpl<std::pair<mlir::Value, mlir::NVVM::PTXRegisterMod>>
5975 &asmValues) {
5976 for (auto arg : getReadWriteArgs())
5977 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::ReadWrite});
5978 for (auto arg : getResults())
5979 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::Write});
5980 for (auto arg : getReadOnlyArgs())
5981 asmValues.push_back({arg, mlir::NVVM::PTXRegisterMod::Read});
5982 if (getPredicate())
5983 asmValues.push_back({getPredicate(), mlir::NVVM::PTXRegisterMod::Read});
5984 return false; // No manual mapping needed
5985}
5986
5987NVVM::IDArgPair ClusterLaunchControlTryCancelOp::getIntrinsicIDAndArgs(
5988 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
5989 auto curOp = cast<NVVM::ClusterLaunchControlTryCancelOp>(op);
5991 args.push_back(mt.lookupValue(curOp.getSmemAddress()));
5992 args.push_back(mt.lookupValue(curOp.getMbarrier()));
5993
5994 llvm::Intrinsic::ID intrinsicID =
5995 curOp.getMulticast()
5996 ? llvm::Intrinsic::
5997 nvvm_clusterlaunchcontrol_try_cancel_async_multicast_shared
5998 : llvm::Intrinsic::nvvm_clusterlaunchcontrol_try_cancel_async_shared;
5999
6000 return {intrinsicID, args};
6001}
6002
6003NVVM::IDArgPair ClusterLaunchControlQueryCancelOp::getIntrinsicIDAndArgs(
6004 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6005 auto curOp = cast<NVVM::ClusterLaunchControlQueryCancelOp>(op);
6007 args.push_back(mt.lookupValue(curOp.getTryCancelResponse()));
6008
6009 llvm::Intrinsic::ID intrinsicID;
6010
6011 switch (curOp.getQueryType()) {
6012 case NVVM::ClusterLaunchControlQueryType::IS_CANCELED:
6013 intrinsicID =
6014 llvm::Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled;
6015 break;
6016 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_X:
6017 intrinsicID = llvm::Intrinsic::
6018 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x;
6019 break;
6020 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Y:
6021 intrinsicID = llvm::Intrinsic::
6022 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y;
6023 break;
6024 case NVVM::ClusterLaunchControlQueryType::GET_FIRST_CTA_ID_Z:
6025 intrinsicID = llvm::Intrinsic::
6026 nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z;
6027 break;
6028 }
6029 return {intrinsicID, args};
6030}
6031
6033PermuteOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
6034 llvm::IRBuilderBase &builder) {
6035 auto thisOp = cast<NVVM::PermuteOp>(op);
6036 NVVM::PermuteMode mode = thisOp.getMode();
6037
6038 static constexpr llvm::Intrinsic::ID IDs[] = {
6039 llvm::Intrinsic::nvvm_prmt, llvm::Intrinsic::nvvm_prmt_f4e,
6040 llvm::Intrinsic::nvvm_prmt_b4e, llvm::Intrinsic::nvvm_prmt_rc8,
6041 llvm::Intrinsic::nvvm_prmt_ecl, llvm::Intrinsic::nvvm_prmt_ecr,
6042 llvm::Intrinsic::nvvm_prmt_rc16};
6043
6044 unsigned modeIndex = static_cast<unsigned>(mode);
6046 args.push_back(mt.lookupValue(thisOp.getLo()));
6047
6048 // Only first 3 modes (Default, f4e, b4e) need the hi operand.
6049 if (modeIndex < 3)
6050 args.push_back(mt.lookupValue(thisOp.getHi()));
6051
6052 args.push_back(mt.lookupValue(thisOp.getSelector()));
6053
6054 return {IDs[modeIndex], args};
6055}
6056
6057mlir::NVVM::IDArgPair TensormapReplaceOp::getIntrinsicIDAndArgs(
6058 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6059 auto thisOp = cast<NVVM::TensormapReplaceOp>(op);
6060
6062 args.push_back(mt.lookupValue(thisOp.getAddr()));
6063 if (thisOp.getOrd())
6064 args.push_back(builder.getInt32(thisOp.getOrd().value()));
6065 if (thisOp.getNewValue())
6066 args.push_back(mt.lookupValue(thisOp.getNewValue()));
6067 if (auto attr = thisOp.getNewValueAttr()) {
6068 auto val =
6070 .Case<TensormapElemtypeAttr, TensormapInterleaveLayoutAttr,
6071 TensormapSwizzleModeAttr, TensormapSwizzleAtomicityAttr,
6072 TensormapFillModeAttr>([](auto attr) {
6073 return static_cast<unsigned>(attr.getValue());
6074 })
6075 .Default([](auto attr) {
6076 llvm_unreachable("Invalid attribute type");
6077 return 0;
6078 });
6079 args.push_back(builder.getInt32(val));
6080 }
6081
6082 static constexpr llvm::Intrinsic::ID IDs[] = {
6083 llvm::Intrinsic::nvvm_tensormap_replace_global_address,
6084 llvm::Intrinsic::nvvm_tensormap_replace_rank,
6085 llvm::Intrinsic::nvvm_tensormap_replace_box_dim,
6086 llvm::Intrinsic::nvvm_tensormap_replace_global_dim,
6087 llvm::Intrinsic::nvvm_tensormap_replace_global_stride,
6088 llvm::Intrinsic::nvvm_tensormap_replace_element_stride,
6089 llvm::Intrinsic::nvvm_tensormap_replace_elemtype,
6090 llvm::Intrinsic::nvvm_tensormap_replace_interleave_layout,
6091 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_mode,
6092 llvm::Intrinsic::nvvm_tensormap_replace_swizzle_atomicity,
6093 llvm::Intrinsic::nvvm_tensormap_replace_fill_mode,
6094 };
6095
6096 unsigned fieldIndex = static_cast<unsigned>(thisOp.getField());
6097
6098 return {IDs[fieldIndex], args};
6099}
6100
6101//===----------------------------------------------------------------------===//
6102// NVVM tcgen05.mma functions
6103//===----------------------------------------------------------------------===//
6104
6105static llvm::nvvm::Tcgen05MMAKind
6106getNVVMTcgen05MMAKind(NVVM::Tcgen05MMAKind kind) {
6107 switch (kind) {
6108 case NVVM::Tcgen05MMAKind::F16:
6109 return llvm::nvvm::Tcgen05MMAKind::F16;
6110 case NVVM::Tcgen05MMAKind::TF32:
6111 return llvm::nvvm::Tcgen05MMAKind::TF32;
6112 case NVVM::Tcgen05MMAKind::F8F6F4:
6113 return llvm::nvvm::Tcgen05MMAKind::F8F6F4;
6114 case NVVM::Tcgen05MMAKind::I8:
6115 return llvm::nvvm::Tcgen05MMAKind::I8;
6116 case NVVM::Tcgen05MMAKind::TI16:
6117 return llvm::nvvm::Tcgen05MMAKind::TI16;
6118 case NVVM::Tcgen05MMAKind::MXF8F6F4:
6119 case NVVM::Tcgen05MMAKind::MXF4:
6120 case NVVM::Tcgen05MMAKind::MXF4NVF4:
6121 // Block-scale kinds are handled by the tcgen05.mma.block_scale
6122 // lowering paths and are not valid for plain tcgen05.mma.
6123 llvm_unreachable("Unsupported tcgen05.mma kind");
6124 }
6125}
6126
6128Tcgen05MMAOp::getIntrinsicIDAndArgs(Operation &op, LLVM::ModuleTranslation &mt,
6129 llvm::IRBuilderBase &builder) {
6130
6131 auto thisOp = cast<NVVM::Tcgen05MMAOp>(op);
6133
6134 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6135
6136 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6137 const bool isATensor = isa<llvm::PointerType>(A->getType());
6138 args.push_back(A);
6139
6140 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6141 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6142 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6143
6144 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
6145 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
6146 using IsATensorArray = std::array<CtaGroupArray, 2>;
6147 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
6148 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
6149
6150 // [hasDisableOutputLane][hasScaleInputD][isATensor][CtaGroup][EnableAShift]
6151 static constexpr HasDisableOutputLaneArray tcgen05MMAIDs = {
6152 { // without diable output lane
6153 {{// without scale input D
6154 {{
6155 // shared
6156 {{// cg1
6157 {llvm::Intrinsic::nvvm_tcgen05_mma_shared, notIntrinsic},
6158 // cg2
6159 {llvm::Intrinsic::nvvm_tcgen05_mma_shared, notIntrinsic}}},
6160 {{// tensor
6161 {
6162 // cg1
6163 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
6164 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
6165 },
6166 {
6167 // cg2
6168 llvm::Intrinsic::nvvm_tcgen05_mma_tensor,
6169 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_ashift,
6170 }}},
6171 }},
6172 // with scale input D
6173 {{ // shared
6174 {{// cg1
6175 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d, notIntrinsic},
6176 // cg2
6177 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_scale_d, notIntrinsic}}},
6178 {{// tensor
6179 {
6180 // cg1
6181 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
6182 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
6183 },
6184 {
6185 // cg2
6186 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d,
6187 llvm::Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_ashift,
6188 }}}}}}},
6189 // with disable output lane
6190 {{ // without scale input D
6191 {{ // shared
6192 {{// cg1
6193 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1,
6194 notIntrinsic},
6195 // cg2
6196 {llvm::Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2,
6197 notIntrinsic}}},
6198 {{// cg1
6199 {
6200 llvm::Intrinsic::
6201 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1,
6202 llvm::Intrinsic::
6203 nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift,
6204 },
6205 // cg2
6206 {
6207 llvm::Intrinsic::
6208 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2,
6209 llvm::Intrinsic::
6210 nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift,
6211 }}}}},
6212 // with scale input D
6213 {{ // shared
6214 {{// cg1
6215 {llvm::Intrinsic::
6216 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1,
6217 notIntrinsic},
6218 // cg2
6219 {llvm::Intrinsic::
6220 nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2,
6221 notIntrinsic}}},
6222 // tensor
6223 {{// cg1
6224 {llvm::Intrinsic::
6225 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1,
6226 llvm::Intrinsic::
6227 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift},
6228 // cg2
6229 {
6230 llvm::Intrinsic::
6231 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2,
6232 llvm::Intrinsic::
6233 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift,
6234 }}}}}}}}};
6235
6236 llvm::Value *ScaleInputD = mt.lookupValue(thisOp.getScaleInputD());
6237 bool hasScaleInputD = ScaleInputD != nullptr;
6238
6239 llvm::Value *DisableOutputLane =
6240 mt.lookupValue(thisOp.getDisableOutputLane());
6241 bool hasDisableOutputLane = DisableOutputLane != nullptr;
6242
6243 const unsigned ctaGroup =
6244 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()));
6245
6246 llvm::Intrinsic::ID ID =
6247 tcgen05MMAIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
6248 [ctaGroup - 1][thisOp.getAShift()];
6249
6250 assert(ID != notIntrinsic && "Invalid intrinsic for Tcgen05MMAOp.");
6251
6252 if (hasScaleInputD)
6253 args.push_back(ScaleInputD);
6254
6255 if (hasDisableOutputLane)
6256 args.push_back(DisableOutputLane);
6257
6258 args.push_back(builder.getInt32(
6259 static_cast<unsigned>(getNVVMTcgen05MMAKind(thisOp.getKind()))));
6260
6261 if (!hasDisableOutputLane)
6262 args.push_back(builder.getInt32(ctaGroup));
6263
6264 args.push_back(
6265 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6266
6267 args.push_back(
6268 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6269
6270 return {ID, args};
6271}
6272
6273static LogicalResult
6274verifyTcgen05MMAOp(bool isATensor, mlir::Value disableOutputLane,
6275 NVVM::CTAGroupKind ctaGroup, bool hasAShift,
6276 NVVM::Tcgen05MMACollectorOp collectorOp, Location loc) {
6277
6278 if (disableOutputLane) {
6279 mlir::VectorType disableOutputLaneType =
6280 cast<mlir::VectorType>(disableOutputLane.getType());
6281 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&
6282 disableOutputLaneType.getNumElements() != 4) ||
6283 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&
6284 disableOutputLaneType.getNumElements() != 8))
6285 return emitError(loc) << "Disable Output Lane of length "
6286 << disableOutputLaneType.getNumElements()
6287 << " is incompatible with CtaGroupAttr";
6288 }
6289
6290 if (hasAShift && !isATensor)
6291 return emitError(
6292 loc, "A-shift can be applied only when matrix A is in tensor memory");
6293
6294 if (hasAShift == true && (collectorOp == Tcgen05MMACollectorOp::FILL ||
6295 collectorOp == Tcgen05MMACollectorOp::USE))
6296 return emitError(
6297 loc, "Cannot use collector buffer operation fill or use with ashift");
6298
6299 return success();
6300}
6301
6302LogicalResult Tcgen05MMAOp::verify() {
6303 return verifyTcgen05MMAOp(isa<LLVM::LLVMPointerType>(getMatrixA().getType()),
6304 getDisableOutputLane(), getCtaGroup(), getAShift(),
6305 getCollectorOp(), getLoc());
6306}
6307
6308//===----------------------------------------------------------------------===//
6309// NVVM tcgen05.mma.sp functions
6310//===----------------------------------------------------------------------===//
6311
6312mlir::NVVM::IDArgPair Tcgen05MMASparseOp::getIntrinsicIDAndArgs(
6313 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6314
6315 auto thisOp = cast<NVVM::Tcgen05MMASparseOp>(op);
6317
6318 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6319
6320 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6321 bool isATensor = isa<llvm::PointerType>(A->getType());
6322 args.push_back(A);
6323
6324 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6325 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6326 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6327 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));
6328
6329 using EnableAShiftArray = std::array<llvm::Intrinsic::ID, 2>;
6330 using CtaGroupArray = std::array<EnableAShiftArray, 2>;
6331 using IsATensorArray = std::array<CtaGroupArray, 2>;
6332 using HasScaleInputDArray = std::array<IsATensorArray, 2>;
6333 using HasDisableOutputLaneArray = std::array<HasScaleInputDArray, 2>;
6334
6335 // [hasDisableOutputLane][hasScaleInputD][isATensor][CtaGroup][EnableAShift]
6336 static constexpr HasDisableOutputLaneArray tcgen05MMASparseIDs = {
6337 { // without diable output lane
6338 {{// without scale input D
6339 {{
6340 // shared
6341 {{// cg1
6342 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared, notIntrinsic},
6343 // cg2
6344 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared, notIntrinsic}}},
6345 {{// tensor
6346 {
6347 // cg1
6348 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
6349 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
6350 },
6351 {
6352 // cg2
6353 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor,
6354 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_ashift,
6355 }}},
6356 }},
6357 // with scale input D
6358 {{ // shared
6359 {{// cg1
6360 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
6361 notIntrinsic},
6362 // cg2
6363 {llvm::Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d,
6364 notIntrinsic}}},
6365 {{// tensor
6366 {
6367 // cg1
6368 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
6369 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
6370 },
6371 {
6372 // cg2
6373 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d,
6374 llvm::Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_ashift,
6375 }}}}}}},
6376 // with disable output lane
6377 {{ // without scale input D
6378 {{ // shared
6379 {{// cg1
6380 {llvm::Intrinsic::
6381 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1,
6382 notIntrinsic},
6383 // cg2
6384 {llvm::Intrinsic::
6385 nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2,
6386 notIntrinsic}}},
6387 {{// cg1
6388 {
6389 llvm::Intrinsic::
6390 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1,
6391 llvm::Intrinsic::
6392 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift,
6393 },
6394 // cg2
6395 {
6396 llvm::Intrinsic::
6397 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2,
6398 llvm::Intrinsic::
6399 nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift,
6400 }}}}},
6401 // with scale input D
6402 {{ // shared
6403 {{// cg1
6404 {llvm::Intrinsic::
6405 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1,
6406 notIntrinsic},
6407 // cg2
6408 {llvm::Intrinsic::
6409 nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2,
6410 notIntrinsic}}},
6411 // tensor
6412 {{// cg1
6413 {llvm::Intrinsic::
6414 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1,
6415 llvm::Intrinsic::
6416 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift},
6417 // cg2
6418 {
6419 llvm::Intrinsic::
6420 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2,
6421 llvm::Intrinsic::
6422 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift,
6423 }}}}}}}}};
6424
6425 llvm::Value *ScaleInputD = mt.lookupValue(thisOp.getScaleInputD());
6426 bool hasScaleInputD = ScaleInputD != nullptr;
6427
6428 llvm::Value *DisableOutputLane =
6429 mt.lookupValue(thisOp.getDisableOutputLane());
6430 bool hasDisableOutputLane = DisableOutputLane != nullptr;
6431
6432 unsigned ctaGroup =
6433 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()));
6434
6435 llvm::Intrinsic::ID ID =
6436 tcgen05MMASparseIDs[hasDisableOutputLane][hasScaleInputD][isATensor]
6437 [ctaGroup - 1][thisOp.getAShift()];
6438
6439 assert(ID != notIntrinsic && "Invalid intrinsic for Tcgen05MMASparseOp.");
6440
6441 if (hasScaleInputD)
6442 args.push_back(ScaleInputD);
6443
6444 if (hasDisableOutputLane)
6445 args.push_back(DisableOutputLane);
6446
6447 args.push_back(builder.getInt32(
6448 static_cast<unsigned>(getNVVMTcgen05MMAKind(thisOp.getKind()))));
6449
6450 if (!hasDisableOutputLane)
6451 args.push_back(builder.getInt32(ctaGroup));
6452
6453 args.push_back(
6454 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6455
6456 args.push_back(
6457 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6458
6459 return {ID, args};
6460}
6461
6462LogicalResult Tcgen05MMASparseOp::verify() {
6463 return verifyTcgen05MMAOp(isa<LLVM::LLVMPointerType>(getMatrixA().getType()),
6464 getDisableOutputLane(), getCtaGroup(), getAShift(),
6465 getCollectorOp(), getLoc());
6466}
6467
6468//===----------------------------------------------------------------------===//
6469// NVVM tcgen05.mma.block_scale functions
6470//===----------------------------------------------------------------------===//
6471
6472mlir::NVVM::IDArgPair Tcgen05MMABlockScaleOp::getIntrinsicIDAndArgs(
6473 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6474
6475 auto thisOp = cast<NVVM::Tcgen05MMABlockScaleOp>(op);
6477
6478 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6479
6480 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6481 bool isATensor = isa<llvm::PointerType>(A->getType());
6482 args.push_back(A);
6483
6484 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6485 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6486 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6487 args.push_back(mt.lookupValue(thisOp.getScaleA()));
6488 args.push_back(mt.lookupValue(thisOp.getScaleB()));
6489 args.push_back(builder.getInt32(
6490 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()))));
6491 args.push_back(
6492 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6493 args.push_back(
6494 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6495
6496 auto kind = thisOp.getKind();
6497 auto blockScale = thisOp.getBlockScale();
6498 llvm::Intrinsic::ID ID = [&]() {
6499 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
6500 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6501 return isATensor ? llvm::Intrinsic::
6502 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale
6503 : llvm::Intrinsic::
6504 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale;
6505 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6506 return isATensor
6507 ? llvm::Intrinsic::
6508 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block32
6509 : llvm::Intrinsic::
6510 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32;
6511 }
6512 } else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
6513 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6514 return isATensor
6515 ? llvm::Intrinsic::nvvm_tcgen05_mma_tensor_mxf4_block_scale
6516 : llvm::Intrinsic::nvvm_tcgen05_mma_shared_mxf4_block_scale;
6517 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6518 return isATensor ? llvm::Intrinsic::
6519 nvvm_tcgen05_mma_tensor_mxf4_block_scale_block32
6520 : llvm::Intrinsic::
6521 nvvm_tcgen05_mma_shared_mxf4_block_scale_block32;
6522 }
6523 } else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
6524 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6525 return isATensor
6526 ? llvm::Intrinsic::
6527 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block32
6528 : llvm::Intrinsic::
6529 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block32;
6530
6531 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
6532 return isATensor
6533 ? llvm::Intrinsic::
6534 nvvm_tcgen05_mma_tensor_mxf4nvf4_block_scale_block16
6535 : llvm::Intrinsic::
6536 nvvm_tcgen05_mma_shared_mxf4nvf4_block_scale_block16;
6537 }
6538 }
6539 llvm_unreachable("Invalid tcgen05.mma.block_scale attributes");
6540 }();
6541
6542 return {ID, args};
6543}
6544
6545static LogicalResult verifyTcgen05MMABlockScaleOp(
6546 NVVM::Tcgen05MMACollectorOp collectorOp, NVVM::Tcgen05MMAKind kind,
6547 NVVM::Tcgen05MMABlockScale blockScale, Location loc) {
6548 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT &&
6549 kind == NVVM::Tcgen05MMAKind::MXF4NVF4)
6550 return emitError(loc, "mxf4nvf4 requires block scale attribute");
6551
6552 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16 &&
6553 kind != NVVM::Tcgen05MMAKind::MXF4NVF4)
6554 return emitError(loc,
6555 llvm::formatv("{} kind does not support block16 attribute",
6556 stringifyEnum(kind)));
6557
6558 return success();
6559}
6560
6561LogicalResult Tcgen05MMABlockScaleOp::verify() {
6562 return verifyTcgen05MMABlockScaleOp(getCollectorOp(), getKind(),
6563 getBlockScale(), getLoc());
6564}
6565
6566//===----------------------------------------------------------------------===//
6567// NVVM tcgen05.mma.sp.block_scale functions
6568//===----------------------------------------------------------------------===//
6569
6570mlir::NVVM::IDArgPair Tcgen05MMASparseBlockScaleOp::getIntrinsicIDAndArgs(
6571 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6572
6573 auto thisOp = cast<NVVM::Tcgen05MMASparseBlockScaleOp>(op);
6575
6576 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6577
6578 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6579 bool isATensor = isa<llvm::PointerType>(A->getType());
6580 args.push_back(A);
6581
6582 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6583 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6584 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6585 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));
6586 args.push_back(mt.lookupValue(thisOp.getScaleA()));
6587 args.push_back(mt.lookupValue(thisOp.getScaleB()));
6588 args.push_back(builder.getInt32(
6589 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()))));
6590 args.push_back(
6591 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6592 args.push_back(
6593 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6594
6595 auto kind = thisOp.getKind();
6596 auto blockScale = thisOp.getBlockScale();
6597 llvm::Intrinsic::ID ID = [&]() {
6598 if (kind == NVVM::Tcgen05MMAKind::MXF8F6F4) {
6599 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6600 return isATensor ? llvm::Intrinsic::
6601 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale
6602 : llvm::Intrinsic::
6603 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale;
6604 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6605 return isATensor
6606 ? llvm::Intrinsic::
6607 nvvm_tcgen05_mma_sp_tensor_mxf8f6f4_block_scale_block32
6608 : llvm::Intrinsic::
6609 nvvm_tcgen05_mma_sp_shared_mxf8f6f4_block_scale_block32;
6610 }
6611 } else if (kind == NVVM::Tcgen05MMAKind::MXF4) {
6612 if (blockScale == NVVM::Tcgen05MMABlockScale::DEFAULT) {
6613 return isATensor ? llvm::Intrinsic::
6614 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale
6615 : llvm::Intrinsic::
6616 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale;
6617 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6618 return isATensor
6619 ? llvm::Intrinsic::
6620 nvvm_tcgen05_mma_sp_tensor_mxf4_block_scale_block32
6621 : llvm::Intrinsic::
6622 nvvm_tcgen05_mma_sp_shared_mxf4_block_scale_block32;
6623 }
6624 } else if (kind == NVVM::Tcgen05MMAKind::MXF4NVF4) {
6625 if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK32) {
6626 return isATensor
6627 ? llvm::Intrinsic::
6628 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block32
6629 : llvm::Intrinsic::
6630 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block32;
6631
6632 } else if (blockScale == NVVM::Tcgen05MMABlockScale::BLOCK16) {
6633 return isATensor
6634 ? llvm::Intrinsic::
6635 nvvm_tcgen05_mma_sp_tensor_mxf4nvf4_block_scale_block16
6636 : llvm::Intrinsic::
6637 nvvm_tcgen05_mma_sp_shared_mxf4nvf4_block_scale_block16;
6638 }
6639 }
6640 llvm_unreachable("Invalid tcgen05.mma.sp.block_scale attributes");
6641 }();
6642
6643 return {ID, args};
6644}
6645
6646LogicalResult Tcgen05MMASparseBlockScaleOp::verify() {
6647 return verifyTcgen05MMABlockScaleOp(getCollectorOp(), getKind(),
6648 getBlockScale(), getLoc());
6649}
6650
6651//===----------------------------------------------------------------------===//
6652// NVVM tcgen05.mma.ws functions
6653//===----------------------------------------------------------------------===//
6654
6655mlir::NVVM::IDArgPair Tcgen05MMAWsOp::getIntrinsicIDAndArgs(
6656 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6657
6658 auto thisOp = cast<NVVM::Tcgen05MMAWsOp>(op);
6660
6661 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6662
6663 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6664 bool isATensor = isa<llvm::PointerType>(A->getType());
6665 args.push_back(A);
6666
6667 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6668 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6669 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6670
6671 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6672 llvm::Intrinsic::ID ID = notIntrinsic;
6673 if (ZeroColMask) {
6674 args.push_back(mt.lookupValue(ZeroColMask));
6675 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor_zero_col_mask
6676 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared_zero_col_mask;
6677 } else
6678 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_tensor
6679 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_shared;
6680
6681 args.push_back(builder.getInt32(
6682 static_cast<unsigned>(getNVVMTcgen05MMAKind(thisOp.getKind()))));
6683 args.push_back(
6684 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6685 args.push_back(
6686 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6687
6688 return {ID, args};
6689}
6690
6691//===----------------------------------------------------------------------===//
6692// NVVM tcgen05.mma.ws.sp functions
6693//===----------------------------------------------------------------------===//
6694
6695mlir::NVVM::IDArgPair Tcgen05MMAWsSparseOp::getIntrinsicIDAndArgs(
6696 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6697
6698 auto thisOp = cast<NVVM::Tcgen05MMAWsSparseOp>(op);
6700
6701 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6702
6703 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6704 bool isATensor = isa<llvm::PointerType>(A->getType());
6705 args.push_back(A);
6706
6707 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6708 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6709 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6710 args.push_back(mt.lookupValue(thisOp.getSparseMetadata()));
6711
6712 mlir::Value ZeroColMask = thisOp.getZeroColMask();
6713 llvm::Intrinsic::ID ID = notIntrinsic;
6714 if (ZeroColMask) {
6715 args.push_back(mt.lookupValue(ZeroColMask));
6716 ID = isATensor
6717 ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor_zero_col_mask
6718 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared_zero_col_mask;
6719 } else
6720 ID = isATensor ? llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_tensor
6721 : llvm::Intrinsic::nvvm_tcgen05_mma_ws_sp_shared;
6722
6723 args.push_back(builder.getInt32(
6724 static_cast<unsigned>(getNVVMTcgen05MMAKind(thisOp.getKind()))));
6725 args.push_back(
6726 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorBBuffer())));
6727 args.push_back(
6728 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOp())));
6729
6730 return {ID, args};
6731}
6732
6733//===----------------------------------------------------------------------===//
6734// NVVM tcgen05.mma.decompress_b functions
6735//===----------------------------------------------------------------------===//
6736
6737mlir::NVVM::IDArgPair Tcgen05MMADecompressBOp::getIntrinsicIDAndArgs(
6738 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6739 auto thisOp = cast<Tcgen05MMADecompressBOp>(op);
6741
6742 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6743
6744 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6745 const bool isATensor = isa<llvm::PointerType>(A->getType());
6746 args.push_back(A);
6747
6748 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6749 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6750 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6751 args.push_back(mt.lookupValue(thisOp.getDecompressBMetadata()));
6752
6753 llvm::Value *DisableOutputLane =
6754 mt.lookupValue(thisOp.getDisableOutputLane());
6755 bool hasDisableOutputLane = DisableOutputLane != nullptr;
6756
6757 NVVM::CTAGroupKind ctaGroup = thisOp.getCtaGroup();
6758
6759 using namespace llvm::Intrinsic;
6760 ID intrinsicID = not_intrinsic;
6761
6762 if (hasDisableOutputLane) {
6763 if (ctaGroup == NVVM::CTAGroupKind::CTA_1) {
6764 intrinsicID =
6765 isATensor
6766 ? nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg1_decompress_b
6767 : nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg1_decompress_b;
6768 } else if (ctaGroup == NVVM::CTAGroupKind::CTA_2) {
6769 intrinsicID =
6770 isATensor
6771 ? nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg2_decompress_b
6772 : nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg2_decompress_b;
6773 } else {
6774 llvm_unreachable("Unknown ctaGroup for tcgen05.mma.decompress_b");
6775 }
6776 } else {
6777 intrinsicID = isATensor ? nvvm_tcgen05_mma_tensor_f8f6f4_decompress_b
6778 : nvvm_tcgen05_mma_shared_f8f6f4_decompress_b;
6779 }
6780
6781 assert(intrinsicID != not_intrinsic &&
6782 "Invalid intrinsic for Tcgen05MMADecompressBOp.");
6783
6784 if (hasDisableOutputLane)
6785 args.push_back(DisableOutputLane);
6786 else
6787 args.push_back(
6788 builder.getInt32(static_cast<unsigned>(getNVVMCtaGroupKind(ctaGroup))));
6789
6790 args.push_back(
6791 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpA())));
6792 args.push_back(
6793 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6794
6795 return {intrinsicID, args};
6796}
6797
6798LogicalResult Tcgen05MMADecompressBOp::verify() {
6799 mlir::Value disableOutputLane = getDisableOutputLane();
6800
6801 if (disableOutputLane) {
6802 NVVM::CTAGroupKind ctaGroup = getCtaGroup();
6803
6804 mlir::VectorType disableOutputLaneType =
6805 cast<mlir::VectorType>(disableOutputLane.getType());
6806 if ((ctaGroup == NVVM::CTAGroupKind::CTA_1 &&
6807 disableOutputLaneType.getNumElements() != 4) ||
6808 (ctaGroup == NVVM::CTAGroupKind::CTA_2 &&
6809 disableOutputLaneType.getNumElements() != 8))
6810 return emitOpError() << "Disable Output Lane of length "
6811 << disableOutputLaneType.getNumElements()
6812 << " is incompatible with CtaGroupAttr";
6813 }
6814
6815 return success();
6816}
6817
6818//===----------------------------------------------------------------------===//
6819// NVVM tcgen05.mma.block_scale.decompress_b functions
6820//===----------------------------------------------------------------------===//
6821
6822mlir::NVVM::IDArgPair Tcgen05MMABlockScaleDecompressBOp::getIntrinsicIDAndArgs(
6823 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6824 auto thisOp = cast<Tcgen05MMABlockScaleDecompressBOp>(op);
6826
6827 args.push_back(mt.lookupValue(thisOp.getMatrixD()));
6828
6829 llvm::Value *A = mt.lookupValue(thisOp.getMatrixA());
6830 const bool isATensor = isa<llvm::PointerType>(A->getType());
6831 args.push_back(A);
6832
6833 args.push_back(mt.lookupValue(thisOp.getMatrixB()));
6834 args.push_back(mt.lookupValue(thisOp.getIdesc()));
6835 args.push_back(mt.lookupValue(thisOp.getEnableInputD()));
6836 args.push_back(mt.lookupValue(thisOp.getScaleA()));
6837 args.push_back(mt.lookupValue(thisOp.getScaleB()));
6838 args.push_back(mt.lookupValue(thisOp.getDecompressBMetadata()));
6839 args.push_back(builder.getInt32(
6840 static_cast<unsigned>(getNVVMCtaGroupKind(thisOp.getCtaGroup()))));
6841 args.push_back(
6842 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpA())));
6843 args.push_back(
6844 builder.getInt32(static_cast<unsigned>(thisOp.getCollectorOpB())));
6845
6846 llvm::Intrinsic::ID intrinsicID =
6847 isATensor
6848 ? llvm::Intrinsic::
6849 nvvm_tcgen05_mma_tensor_mxf8f6f4_block_scale_block32_decompress_b
6850 : llvm::Intrinsic::
6851 nvvm_tcgen05_mma_shared_mxf8f6f4_block_scale_block32_decompress_b;
6852
6853 return {intrinsicID, args};
6854}
6855
6856//===----------------------------------------------------------------------===//
6857// NVVM tcgen05.ld.red functions
6858//===----------------------------------------------------------------------===//
6859
6860#define TCGEN05LDRED(SHAPE, NUM, TYPE) \
6861 llvm::Intrinsic::nvvm_tcgen05_ld_red_##SHAPE##_##NUM##_##TYPE
6862
6863mlir::NVVM::IDArgPair NVVM::Tcgen05LdRedOp::getIntrinsicIDAndArgs(
6864 Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
6865 auto thisOp = cast<NVVM::Tcgen05LdRedOp>(op);
6867
6868 mlir::VectorType VecResTy =
6869 cast<mlir::VectorType>(thisOp.getData().getType());
6870 unsigned Num = VecResTy.getNumElements();
6871 bool IsFloat = thisOp.getRedVal().getType().isF32();
6872
6873 llvm::Intrinsic::ID Shape32x32b[][2] = {
6875 {TCGEN05LDRED(32x32b, x2, i32), TCGEN05LDRED(32x32b, x2, f32)},
6876 {TCGEN05LDRED(32x32b, x4, i32), TCGEN05LDRED(32x32b, x4, f32)},
6877 {TCGEN05LDRED(32x32b, x8, i32), TCGEN05LDRED(32x32b, x8, f32)},
6878 {TCGEN05LDRED(32x32b, x16, i32), TCGEN05LDRED(32x32b, x16, f32)},
6879 {TCGEN05LDRED(32x32b, x32, i32), TCGEN05LDRED(32x32b, x32, f32)},
6880 {TCGEN05LDRED(32x32b, x64, i32), TCGEN05LDRED(32x32b, x64, f32)},
6881 {TCGEN05LDRED(32x32b, x128, i32), TCGEN05LDRED(32x32b, x128, f32)},
6882 };
6883
6884 llvm::Intrinsic::ID Shape16x32bx2[][2] = {
6886 {TCGEN05LDRED(16x32bx2, x2, i32), TCGEN05LDRED(16x32bx2, x2, f32)},
6887 {TCGEN05LDRED(16x32bx2, x4, i32), TCGEN05LDRED(16x32bx2, x4, f32)},
6888 {TCGEN05LDRED(16x32bx2, x8, i32), TCGEN05LDRED(16x32bx2, x8, f32)},
6889 {TCGEN05LDRED(16x32bx2, x16, i32), TCGEN05LDRED(16x32bx2, x16, f32)},
6890 {TCGEN05LDRED(16x32bx2, x32, i32), TCGEN05LDRED(16x32bx2, x32, f32)},
6891 {TCGEN05LDRED(16x32bx2, x64, i32), TCGEN05LDRED(16x32bx2, x64, f32)},
6892 {TCGEN05LDRED(16x32bx2, x128, i32), TCGEN05LDRED(16x32bx2, x128, f32)},
6893 };
6894
6895 NVVM::Tcgen05LdStShape shape = thisOp.getShape();
6896 unsigned ID = [&]() {
6897 // `num` contains the length of vector and log2 of `num` returns the index
6898 // into the shape array
6899 unsigned idx = std::log2(Num);
6900 switch (shape) {
6901 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
6902 return Shape32x32b[idx][IsFloat];
6903 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
6904 return Shape16x32bx2[idx][IsFloat];
6905 default:
6906 llvm_unreachable("unhandled tcgen05.ld lowering");
6907 }
6908 }();
6909
6910 args.push_back(mt.lookupValue(thisOp.getAddr()));
6911
6912 if (shape == NVVM::Tcgen05LdStShape::SHAPE_16X32BX2)
6913 args.push_back(mt.lookupValue(thisOp.getOffset()));
6914
6915 args.push_back(
6916 builder.getInt32(thisOp.getOp() == NVVM::ReductionKind::MIN ? 0 : 1));
6917
6918 if (IsFloat) {
6919 args.push_back(builder.getInt1(static_cast<unsigned>(thisOp.getAbs())));
6920 args.push_back(builder.getInt1(static_cast<unsigned>(thisOp.getNan())));
6921 }
6922 return {ID, args};
6923}
6924
6925LogicalResult Tcgen05LdRedOp::verify() {
6926 VectorType data = cast<VectorType>(getData().getType());
6927 Type redVal = getRedVal().getType();
6928
6929 if (data.getElementType() != redVal)
6930 return emitError(
6931 "type of reduction value and element type of vector data should match");
6932
6933 if (getOp() != NVVM::ReductionKind::MIN &&
6934 getOp() != NVVM::ReductionKind::MAX)
6935 return emitError("only min and max reduction kinds are supported");
6936
6937 if (redVal.isInteger() && (getAbs() || getNan())) {
6938 return emitError("abs or nan is only applicable for f32 type");
6939 }
6940 return success();
6941}
6942
6943//===----------------------------------------------------------------------===//
6944// NVVMDialect initialization, type parsing, and registration.
6945//===----------------------------------------------------------------------===//
6946
6947namespace {
6948struct NVVMInlinerInterface final : DialectInlinerInterface {
6949 using DialectInlinerInterface::DialectInlinerInterface;
6950 bool isLegalToInline(Operation *, Region *, bool, IRMapping &) const final {
6951 return true;
6952 }
6953};
6954} // namespace
6955
6956// TODO: This should be the llvm.nvvm dialect once this is supported.
6957void NVVMDialect::initialize() {
6958 addOperations<
6959#define GET_OP_LIST
6960#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
6961 >();
6962 addAttributes<
6963#define GET_ATTRDEF_LIST
6964#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"
6965 >();
6966
6967 // Support unknown operations because not all NVVM operations are
6968 // registered.
6969 allowUnknownOperations();
6970 addInterfaces<NVVMInlinerInterface>();
6971 declarePromisedInterface<ConvertToLLVMPatternInterface, NVVMDialect>();
6972 declarePromisedInterface<gpu::TargetAttrInterface, NVVMTargetAttr>();
6973}
6974
6975LogicalResult NVVMDialect::verifyOperationAttribute(Operation *op,
6976 NamedAttribute attr) {
6977 StringAttr attrName = attr.getName();
6978 // Kernel function attribute should be attached to functions.
6979 if (attrName == NVVMDialect::getKernelFuncAttrName()) {
6980 if (!isa<LLVM::LLVMFuncOp>(op)) {
6981 return op->emitError() << "'" << NVVMDialect::getKernelFuncAttrName()
6982 << "' attribute attached to unexpected op";
6983 }
6984 }
6985 // If maxntid / reqntid / cluster_dim exist, it must be an array with max 3
6986 // dim
6987 if (attrName == NVVMDialect::getMaxntidAttrName() ||
6988 attrName == NVVMDialect::getReqntidAttrName() ||
6989 attrName == NVVMDialect::getClusterDimAttrName()) {
6990 auto values = llvm::dyn_cast<DenseI32ArrayAttr>(attr.getValue());
6991 if (!values || values.empty() || values.size() > 3) {
6992 return op->emitError()
6993 << "'" << attrName
6994 << "' attribute must be integer array with maximum 3 index";
6995 }
6996 }
6997 // If minctasm / maxnreg / cluster_max_blocks exist, it must be an integer
6998 // attribute
6999 if (attrName == NVVMDialect::getMinctasmAttrName() ||
7000 attrName == NVVMDialect::getMaxnregAttrName() ||
7001 attrName == NVVMDialect::getClusterMaxBlocksAttrName()) {
7002 if (!llvm::dyn_cast<IntegerAttr>(attr.getValue())) {
7003 return op->emitError()
7004 << "'" << attrName << "' attribute must be integer constant";
7005 }
7006 }
7007 // blocksareclusters must be used along with reqntid and cluster_dim
7008 if (attrName == NVVMDialect::getBlocksAreClustersAttrName()) {
7009 if (!op->hasAttr(NVVMDialect::getReqntidAttrName()) ||
7010 !op->hasAttr(NVVMDialect::getClusterDimAttrName())) {
7011 return op->emitError()
7012 << "'" << attrName << "' attribute must be used along with " << "'"
7013 << NVVMDialect::getReqntidAttrName() << "' and " << "'"
7014 << NVVMDialect::getClusterDimAttrName() << "'";
7015 }
7016 }
7017
7018 return success();
7019}
7020
7021LogicalResult NVVMDialect::verifyRegionArgAttribute(Operation *op,
7022 unsigned regionIndex,
7023 unsigned argIndex,
7024 NamedAttribute argAttr) {
7025 auto funcOp = dyn_cast<FunctionOpInterface>(op);
7026 if (!funcOp)
7027 return success();
7028
7029 bool isKernel = op->hasAttr(NVVMDialect::getKernelFuncAttrName());
7030 StringAttr attrName = argAttr.getName();
7031 if (attrName == NVVM::NVVMDialect::getGridConstantAttrName()) {
7032 if (!isKernel) {
7033 return op->emitError()
7034 << "'" << attrName
7035 << "' attribute must be present only on kernel arguments";
7036 }
7037 if (!isa<UnitAttr>(argAttr.getValue()))
7038 return op->emitError() << "'" << attrName << "' must be a unit attribute";
7039 if (!funcOp.getArgAttr(argIndex, LLVM::LLVMDialect::getByValAttrName())) {
7040 return op->emitError()
7041 << "'" << attrName
7042 << "' attribute requires the argument to also have attribute '"
7043 << LLVM::LLVMDialect::getByValAttrName() << "'";
7044 }
7045 }
7046
7047 return success();
7048}
7049
7050//===----------------------------------------------------------------------===//
7051// NVVM Address Space Attr
7052//===----------------------------------------------------------------------===//
7053
7054unsigned NVVMMemorySpaceAttr::getAddressSpace() const {
7055 return static_cast<unsigned>(getValue());
7056}
7057
7058bool NVVMMemorySpaceAttr::isValidLoad(
7059 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
7060 const ::mlir::DataLayout *dataLayout,
7062 return LLVM::detail::isValidLoadStoreImpl(type, ordering, alignment,
7063 dataLayout, emitError);
7064}
7065
7066bool NVVMMemorySpaceAttr::isValidStore(
7067 Type type, ptr::AtomicOrdering ordering, std::optional<int64_t> alignment,
7068 const ::mlir::DataLayout *dataLayout,
7070 return LLVM::detail::isValidLoadStoreImpl(type, ordering, alignment,
7071 dataLayout, emitError);
7072}
7073
7074bool NVVMMemorySpaceAttr::isValidAtomicOp(
7075 ptr::AtomicBinOp op, Type type, ptr::AtomicOrdering ordering,
7076 std::optional<int64_t> alignment, const ::mlir::DataLayout *dataLayout,
7078 // TODO: update this method once `ptr.atomic_rmw` is implemented.
7079 assert(false && "unimplemented, see TODO in the source.");
7080 return false;
7081}
7082
7083bool NVVMMemorySpaceAttr::isValidAtomicXchg(
7084 Type type, ptr::AtomicOrdering successOrdering,
7085 ptr::AtomicOrdering failureOrdering, std::optional<int64_t> alignment,
7086 const ::mlir::DataLayout *dataLayout,
7088 // TODO: update this method once `ptr.atomic_cmpxchg` is implemented.
7089 assert(false && "unimplemented, see TODO in the source.");
7090 return false;
7091}
7092
7093bool NVVMMemorySpaceAttr::isValidAddrSpaceCast(
7094 Type tgt, Type src, function_ref<InFlightDiagnostic()> emitError) const {
7095 // TODO: update this method once the `ptr.addrspace_cast` op is added to the
7096 // dialect.
7097 assert(false && "unimplemented, see TODO in the source.");
7098 return false;
7099}
7100
7101bool NVVMMemorySpaceAttr::isValidPtrIntCast(
7102 Type intLikeTy, Type ptrLikeTy,
7104 // TODO: update this method once the int-cast ops are added to the `ptr`
7105 // dialect.
7106 assert(false && "unimplemented, see TODO in the source.");
7107 return false;
7108}
7109
7110//===----------------------------------------------------------------------===//
7111// NVVM target attribute.
7112//===----------------------------------------------------------------------===//
7113LogicalResult
7114NVVMTargetAttr::verify(function_ref<InFlightDiagnostic()> emitError,
7115 int optLevel, StringRef triple, StringRef chip,
7116 StringRef features, DictionaryAttr flags,
7117 ArrayAttr files, bool verifyTarget) {
7118 if (optLevel < 0 || optLevel > 3) {
7119 emitError() << "The optimization level must be a number between 0 and 3.";
7120 return failure();
7121 }
7122 if (triple.empty()) {
7123 emitError() << "The target triple cannot be empty.";
7124 return failure();
7125 }
7126 if (chip.empty()) {
7127 emitError() << "The target chip cannot be empty.";
7128 return failure();
7129 }
7130 if (files && !llvm::all_of(files, [](::mlir::Attribute attr) {
7131 return mlir::isa_and_nonnull<StringAttr>(attr);
7132 })) {
7133 emitError() << "All the elements in the `link` array must be strings.";
7134 return failure();
7135 }
7136 return success();
7137}
7138
7139LogicalResult NVVMTargetAttr::verifyTarget(Operation *gpuModule) {
7140 if (!getVerifyTarget())
7141 return success();
7142
7143 auto gpuModuleOp = llvm::dyn_cast<gpu::GPUModuleOp>(gpuModule);
7144 if (!gpuModuleOp) {
7145 return emitError(gpuModule->getLoc(),
7146 "NVVM target attribute must be attached to a GPU module");
7147 }
7148
7149 const unsigned targetFullSmVersion =
7151 if (!NVVMCheckSMVersion::isMinimumSMVersion(targetFullSmVersion)) {
7152 return emitError(gpuModule->getLoc(),
7153 "Minimum NVVM target SM version is sm_20");
7154 }
7155
7156 if (gpuModuleOp
7157 ->walk([&](Operation *op) {
7158 if (auto reqOp = llvm::dyn_cast<NVVM::RequiresSMInterface>(op)) {
7159 const NVVMCheckSMVersion requirement =
7160 reqOp.getRequiredMinSMVersion();
7161 if (!requirement.isCompatibleWith(targetFullSmVersion)) {
7162 op->emitOpError() << "is not supported on " << getChip();
7163 return WalkResult::interrupt();
7164 }
7165 }
7166 return WalkResult::advance();
7167 })
7168 .wasInterrupted())
7169 return failure();
7170
7171 return success();
7172}
7173
7174#define GET_OP_CLASSES
7175#include "mlir/Dialect/LLVMIR/NVVMOps.cpp.inc"
7176
7177#define GET_ATTRDEF_CLASSES
7178#include "mlir/Dialect/LLVMIR/NVVMOpsAttributes.cpp.inc"
return success()
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
ArrayAttr()
b getContext())
#define GET_TCGEN05_CP_ID(shape_mc, src_fmt, is_2cta)
static LogicalResult verifyTMALoadParams(size_t tensorDims, size_t numIm2colOff, TMALoadMode mode, Location loc)
static LogicalResult verifyTcgen05MMAOp(bool isATensor, mlir::Value disableOutputLane, NVVM::CTAGroupKind ctaGroup, bool hasAShift, NVVM::Tcgen05MMACollectorOp collectorOp, Location loc)
#define _none
static bool isPtrInAddrSpace(mlir::Value ptr, NVVMMemorySpace targetAS)
static bool isCompatibleReturnTypesOptionalResult(TypeRange inferred, TypeRange actual)
For ops with optional results, allow the user to omit the result even when inference would produce on...
static bool isPtrInSharedCTASpace(mlir::Value ptr)
static LogicalResult isAllowedSizeN(int sizeN, NVVM::WGMMATypes typeA)
static llvm::nvvm::CTAGroupKind getNVVMCtaGroupKind(NVVM::CTAGroupKind ctaGroup)
static void printCTAGroup(OpAsmPrinter &printer, Operation *, NVVM::CTAGroupKindAttr groupAttr)
static void addInferredMultiplicandTypes(MLIRContext *ctx, OperationState &result, ValueRange operandA, ValueRange operandB, std::optional< std::array< MMATypes, 2 > > multiplicandPtxTypes)
#define GET_CVT_F2TF32_ID(rnd, relu, sf)
static void addBlockScaleProperties(OpBuilder &builder, OperationState &result, ArrayRef< int64_t > shape, ScaleVecSize scaleVecSize, BlockScaleFormat blockScaleFormat, MMABlockScaleKind kind)
static ParseResult parseCTAGroup(OpAsmParser &parser, NVVM::CTAGroupKindAttr &groupAttr)
static ParseResult parseEnumKeyword(OpAsmParser &parser, AttrTy &attr)
#define GET_F32x2_TO_F8X2_US_ID(rnd, has_satf)
static llvm::nvvm::Tcgen05MMAKind getNVVMTcgen05MMAKind(NVVM::Tcgen05MMAKind kind)
static llvm::Value * getParamCastedAddr(llvm::Value *addr, llvm::IRBuilderBase &builder)
static LogicalResult verifyAddSubFOp(OpType op)
static LogicalResult verifyTcgen05MMABlockScaleOp(NVVM::Tcgen05MMACollectorOp collectorOp, NVVM::Tcgen05MMAKind kind, NVVM::Tcgen05MMABlockScale blockScale, Location loc)
static void printMmaUnitProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword)
static llvm::Value * packValInto64Bits(llvm::IRBuilderBase &builder, llvm::Value *result, llvm::Value *field, unsigned sizeInBits, unsigned start)
Packs the given field into the result.
static void printOperandList(OpAsmPrinter &p, StringRef name, ArrayRef< Value > operands)
static void printMmaProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword, AttrTy value)
static bool isMmaPropertyName(StringRef name)
#define GET_F32x2_TO_F6x2_ID(type, has_relu)
static llvm::Value * getAsPackedI32(llvm::Value *arg, llvm::IRBuilderBase &builder)
static void printMmaEnumProperty(OpAsmPrinter &printer, bool &isFirst, StringRef keyword, AttrTy value)
#define GET_F16x2_TO_F8X2_ID(type, has_relu)
static LogicalResult verifyMBarrierArriveLikeOp(Operation *op, Value addr, NVVM::MemScopeKind scope, Value retVal=nullptr)
static llvm::Value * castPtrToAddrSpace(llvm::IRBuilderBase &builder, llvm::Value *ptr, NVVMMemorySpace targetAS)
static LogicalResult isAllowedWGMMADataType(NVVM::WGMMATypes typeD, NVVM::WGMMATypes typeA, NVVM::WGMMATypes typeB)
static llvm::Intrinsic::ID getBarrierReductionIntrinsic(bool aligned, NVVM::BarrierReduction kind)
Maps the (aligned, kind) pair to the @llvm.nvvm.barrier.cta.red.
static ParseResult parseMmaProperties(OpAsmParser &parser, NamedAttrList &attributes, ArrayRef< StringRef > allowedKeywords, ArrayRef< StringRef > requiredProperties)
static void inferAndSetMultiplicandTypes(MLIRContext *ctx, NamedAttrList &attrs, const SmallVectorImpl< Type > &operandTypes)
static LogicalResult parseMmaOperand(OpAsmParser &parser, StringRef operandName, SmallVectorImpl< OpAsmParser::UnresolvedOperand > &regs)
static std::pair< mlir::Type, unsigned > inferMMATypeFromMNK(NVVM::MMATypes type, NVVM::MMAFrag frag, int m, int n, int k, MLIRContext *context)
static bool isInt8PtxType(MMATypes type)
#define TCGEN05LDRED(SHAPE, NUM, TYPE)
static bool isInt4PtxType(MMATypes type)
static bool isIntegerPtxType(MMATypes type)
#define GET_F32x2_TO_F8X2_S_ID(type, has_relu)
static MMATypes inferPtxTypeFromResult(OpTy op)
static LogicalResult verifyConstantRangeAttr(Operation *op, std::optional< LLVM::ConstantRangeAttr > rangeAttr)
Verify the range attribute satisfies LLVM ConstantRange constructor requirements for NVVM SpecialRang...
static LogicalResult parseMmaTypeSignature(OpAsmParser &parser, SmallVectorImpl< Type > &operandTypes)
static FailureOr< int > getAllowedSizeK(NVVM::WGMMATypes typeA)
static bool isPtrInSharedClusterSpace(mlir::Value ptr)
LogicalResult CpAsyncBulkTensorOverrideAddrCommonVerifier(OperandRange coordinates, OperandRange tensorSize, OperandRange lowerStride, Value upperStride, bool isTile, Location loc)
static ParseResult parseMmaEnumPropertyValue(OpAsmParser &parser, NamedAttrList &attributes, StringRef name)
#define GET_CP_ASYNC_ID(mod, size, has_cpsize)
static unsigned isValidVectorLength(NVVM::Tcgen05LdStShape shape, unsigned vecLen)
static LogicalResult verifyConvertF32x2ToFP16x2Op(Twine dstType, FPRoundingMode rnd, bool hasRandomBits, Operation *op)
static void nvvmInferResultRanges(Operation *op, Value result, ArrayRef<::mlir::ConstantIntRanges > argRanges, SetIntRangeFn setResultRanges)
Infer the result ranges for the NVVM SpecialRangeableRegisterOp that might have ConstantRangeAttr.
static LogicalResult cpAsyncBulkTensorCommonVerifier(size_t tensorDims, bool isIm2Col, size_t numIm2ColOffsets, Location loc)
static ParseResult parseMmaPropertyValue(OpAsmParser &parser, NamedAttrList &attributes, StringRef name)
static bool isPtrInGenericSpace(mlir::Value ptr)
static void processOperandFragments(Op &op, std::array< MMAOperandFragment, 3 > &frags, SmallVectorImpl< Type > &regTypes, SmallVectorImpl< StringRef > &ignoreAttrNames)
static llvm::Intrinsic::ID getBarrierSyncIntrinsic(bool aligned, bool hasCount)
Maps the (aligned, hasCount) pair to the @llvm.nvvm.barrier.cta.sync.
static constexpr unsigned notIntrinsic
static LogicalResult inferMBarrierArriveResultTypes(MLIRContext *context, Value addr, SmallVectorImpl< Type > &inferredReturnTypes)
Only shared_cluster (ptr<7>) produces zero results; all other address spaces (including generic) retu...
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
@ OptionalSquare
Square brackets supporting zero or more ops, or nothing.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
ParseResult parseKeywordOrString(std::string *result)
Parse a keyword or a quoted string.
virtual ParseResult parseCustomAttributeWithFallback(Attribute &result, Type type, function_ref< ParseResult(Attribute &result, Type type)> parseAttribute)=0
Parse a custom attribute with the provided callback, unless the next token is #, in which case the ge...
virtual ParseResult parseEqual()=0
Parse a = token.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseArrow()=0
Parse a '->' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an arrow followed by a type list.
ParseResult parseTypeList(SmallVectorImpl< Type > &result)
Parse a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
void printArrowTypeList(TypeRange &&types)
void printStrippedAttrOrType(AttrOrType attrOrType)
Print the provided attribute in the context of an operation custom printer/parser: this will invoke d...
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
IntegerType getI16Type()
Definition Builders.cpp:69
UnitAttr getUnitAttr()
Definition Builders.cpp:106
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:171
IntegerType getI32Type()
Definition Builders.cpp:71
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
MLIRContext * getContext() const
Definition Builders.h:56
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Definition Builders.h:101
This class represents a diagnostic that is inflight and set to be reported.
static IntegerValueRange getMaxRange(Value value)
Create a maximal range ([0, uint_max(t)] / [int_min(t), int_max(t)]) range that is used to mark the v...
Implementation class for module translation.
llvm::Value * lookupValue(Value value) const
Finds an LLVM IR value corresponding to the given MLIR value.
void mapValue(Value mlir, llvm::Value *llvm)
Stores the mapping between an MLIR value and its LLVM IR counterpart.
llvm::LLVMContext & getLLVMContext() const
Returns the LLVM context in which the IR is being constructed.
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
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
std::optional< NamedAttribute > getNamed(StringRef name) const
Return the specified named attribute if present, std::nullopt otherwise.
Attribute get(StringAttr name) const
Return the specified attribute if present, null otherwise.
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
Attribute set(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
Definition Attributes.h:179
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
void printOperands(const ContainerType &container)
Print a comma separated list of operands.
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
This class helps build Operations.
Definition Builders.h:210
This provides public APIs that all operations should have.
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
AttrClass getAttrOfType(StringAttr name)
Definition Operation.h:602
bool hasAttr(StringAttr name)
Return true if the operation has an attribute with the provided name, false otherwise.
Definition Operation.h:612
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
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
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isF32() const
Definition Types.cpp:40
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
bool isBF16() const
Definition Types.cpp:37
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
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
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
The OpAsmOpInterface, see OpAsmInterface.td for more details.
Definition CallGraph.h:227
bool isValidLoadStoreImpl(Type type, ptr::AtomicOrdering ordering, std::optional< int64_t > alignment, const ::mlir::DataLayout *dataLayout, function_ref< InFlightDiagnostic()> emitError)
Checks whether the given type is an LLVM type that can be loaded or stored.
Definition LLVMAttrs.cpp:75
SmallVector< int64_t, 4 > getCoordinates(ArrayRef< int64_t > basis, unsigned linearIndex)
@ Write
Write register with '=' modifier.
@ ReadWrite
ReadWrite register with '+' modifier.
@ Read
Read register with no modifier.
std::pair< mlir::Type, unsigned > inferMMAType(mlir::NVVM::MMATypes type, mlir::NVVM::MMAFrag frag, int nRow, int nCol, mlir::MLIRContext *context)
Return the element type and number of elements associated with a wmma matrix of given chracteristics.
std::pair< llvm::Intrinsic::ID, llvm::SmallVector< llvm::Value * > > IDArgPair
A pair type of LLVM's Intrinsic ID and args (which are llvm values).
Definition NVVMDialect.h:53
void walk(Operation *op, function_ref< void(Region *)> callback, WalkOrder order)
Walk all of the regions, blocks, or operations nested under (and including) the given operation.
Definition Visitors.h:102
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
uint64_t getN(LevelType lt)
Definition Enums.h:442
uint64_t getM(LevelType lt)
Definition Enums.h:443
Include the generated interface declarations.
llvm::function_ref< void(Value, const ConstantIntRanges &)> SetIntRangeFn
The type of the setResultRanges callback provided to ops implementing InferIntRangeInterface.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
LogicalResult matchAndRewrite(SubFOp op, PatternRewriter &rewriter) const override
static bool isMinimumSMVersion(unsigned fullSmVersion)
static unsigned getTargetFullSmVersionFromStr(StringRef smVersionString)
bool isCompatibleWith(const unsigned &targetFullSmVersion) const
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.