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