MLIR 24.0.0git
NVVMToLLVMIRTranslation.cpp
Go to the documentation of this file.
1//===- NVVMToLLVMIRTranslation.cpp - Translate NVVM to LLVM IR ------------===//
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 implements a translation between the MLIR NVVM dialect and
10// LLVM IR.
11//
12//===----------------------------------------------------------------------===//
13
16#include "mlir/IR/Operation.h"
18
19#include "llvm/ADT/StringExtras.h"
20#include "llvm/ADT/iterator_range.h"
21#include "llvm/IR/IRBuilder.h"
22#include "llvm/IR/IntrinsicsNVPTX.h"
23#include "llvm/Support/FormatVariadic.h"
24#include "llvm/Support/NVVMAttributes.h"
25#include <cmath>
26
27using namespace mlir;
28using namespace mlir::LLVM;
30
31#define REDUX_F32_ID_IMPL(op, abs, hasNaN) \
32 hasNaN ? llvm::Intrinsic::nvvm_redux_sync_f##op##abs##_NaN \
33 : llvm::Intrinsic::nvvm_redux_sync_f##op##abs
34
35#define GET_REDUX_F32_ID(op, hasAbs, hasNaN) \
36 hasAbs ? REDUX_F32_ID_IMPL(op, _abs, hasNaN) : REDUX_F32_ID_IMPL(op, , hasNaN)
37
38static llvm::Intrinsic::ID getReduxIntrinsicId(llvm::Type *resultType,
39 NVVM::ReductionKind kind,
40 bool hasAbs, bool hasNaN) {
41 switch (kind) {
42 case NVVM::ReductionKind::ADD:
43 return llvm::Intrinsic::nvvm_redux_sync_add;
44 case NVVM::ReductionKind::UMAX:
45 return llvm::Intrinsic::nvvm_redux_sync_umax;
46 case NVVM::ReductionKind::UMIN:
47 return llvm::Intrinsic::nvvm_redux_sync_umin;
48 case NVVM::ReductionKind::AND:
49 return llvm::Intrinsic::nvvm_redux_sync_and;
50 case NVVM::ReductionKind::OR:
51 return llvm::Intrinsic::nvvm_redux_sync_or;
52 case NVVM::ReductionKind::XOR:
53 return llvm::Intrinsic::nvvm_redux_sync_xor;
54 case NVVM::ReductionKind::MAX:
55 return llvm::Intrinsic::nvvm_redux_sync_max;
56 case NVVM::ReductionKind::MIN:
57 return llvm::Intrinsic::nvvm_redux_sync_min;
58 case NVVM::ReductionKind::FMIN:
59 return GET_REDUX_F32_ID(min, hasAbs, hasNaN);
60 case NVVM::ReductionKind::FMAX:
61 return GET_REDUX_F32_ID(max, hasAbs, hasNaN);
62 }
63 llvm_unreachable("unknown reduction kind");
64}
65
66static llvm::Intrinsic::ID getShflIntrinsicId(llvm::Type *resultType,
67 NVVM::ShflKind kind,
68 bool withPredicate) {
69
70 if (withPredicate) {
71 resultType = cast<llvm::StructType>(resultType)->getElementType(0);
72 switch (kind) {
73 case NVVM::ShflKind::bfly:
74 return resultType->isFloatTy()
75 ? llvm::Intrinsic::nvvm_shfl_sync_bfly_f32p
76 : llvm::Intrinsic::nvvm_shfl_sync_bfly_i32p;
77 case NVVM::ShflKind::up:
78 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_up_f32p
79 : llvm::Intrinsic::nvvm_shfl_sync_up_i32p;
80 case NVVM::ShflKind::down:
81 return resultType->isFloatTy()
82 ? llvm::Intrinsic::nvvm_shfl_sync_down_f32p
83 : llvm::Intrinsic::nvvm_shfl_sync_down_i32p;
84 case NVVM::ShflKind::idx:
85 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_idx_f32p
86 : llvm::Intrinsic::nvvm_shfl_sync_idx_i32p;
87 }
88 } else {
89 switch (kind) {
90 case NVVM::ShflKind::bfly:
91 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_bfly_f32
92 : llvm::Intrinsic::nvvm_shfl_sync_bfly_i32;
93 case NVVM::ShflKind::up:
94 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_up_f32
95 : llvm::Intrinsic::nvvm_shfl_sync_up_i32;
96 case NVVM::ShflKind::down:
97 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_down_f32
98 : llvm::Intrinsic::nvvm_shfl_sync_down_i32;
99 case NVVM::ShflKind::idx:
100 return resultType->isFloatTy() ? llvm::Intrinsic::nvvm_shfl_sync_idx_f32
101 : llvm::Intrinsic::nvvm_shfl_sync_idx_i32;
102 }
103 }
104 llvm_unreachable("unknown shuffle kind");
105}
106
107static llvm::Intrinsic::ID getMatchSyncIntrinsicId(Type valType,
108 NVVM::MatchSyncKind kind) {
109 switch (kind) {
110 case NVVM::MatchSyncKind::any:
111 return valType.isInteger(32) ? llvm::Intrinsic::nvvm_match_any_sync_i32
112 : llvm::Intrinsic::nvvm_match_any_sync_i64;
113 case NVVM::MatchSyncKind::all:
114 // match.all instruction has two variants -- one returns a single value,
115 // another returns a pair {value, predicate}. We currently only implement
116 // the latter as that's the variant exposed by CUDA API.
117 return valType.isInteger(32) ? llvm::Intrinsic::nvvm_match_all_sync_i32p
118 : llvm::Intrinsic::nvvm_match_all_sync_i64p;
119 }
120 llvm_unreachable("unsupported match sync kind");
121}
122
123static llvm::Intrinsic::ID getVoteSyncIntrinsicId(NVVM::VoteSyncKind kind) {
124 switch (kind) {
125 case NVVM::VoteSyncKind::any:
126 return llvm::Intrinsic::nvvm_vote_any_sync;
127 case NVVM::VoteSyncKind::all:
128 return llvm::Intrinsic::nvvm_vote_all_sync;
129 case NVVM::VoteSyncKind::ballot:
130 return llvm::Intrinsic::nvvm_vote_ballot_sync;
131 case NVVM::VoteSyncKind::uni:
132 return llvm::Intrinsic::nvvm_vote_uni_sync;
133 }
134 llvm_unreachable("unsupported vote kind");
135}
136
137static llvm::Intrinsic::ID
138getLdMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num,
139 NVVM::LdStMatrixShapeAttr shape,
140 NVVM::LdStMatrixEltType eltType) {
141 if (shape.getM() == 8 && shape.getN() == 8) {
142 switch (num) {
143 case 1:
144 return (layout == NVVM::MMALayout::row)
145 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_b16
146 : llvm::Intrinsic::
147 nvvm_ldmatrix_sync_aligned_m8n8_x1_trans_b16;
148 case 2:
149 return (layout == NVVM::MMALayout::row)
150 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x2_b16
151 : llvm::Intrinsic::
152 nvvm_ldmatrix_sync_aligned_m8n8_x2_trans_b16;
153 case 4:
154 return (layout == NVVM::MMALayout::row)
155 ? llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x4_b16
156 : llvm::Intrinsic::
157 nvvm_ldmatrix_sync_aligned_m8n8_x4_trans_b16;
158 }
159 } else if (shape.getM() == 8 && shape.getN() == 16) {
160 if (eltType == NVVM::LdStMatrixEltType::B8X16_B6X16_P32) {
161 switch (num) {
162 case 1:
163 return llvm::Intrinsic::
164 nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b6x16_p32;
165 case 2:
166 return llvm::Intrinsic::
167 nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b6x16_p32;
168 case 4:
169 return llvm::Intrinsic::
170 nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b6x16_p32;
171 }
172 } else if (eltType == NVVM::LdStMatrixEltType::B8X16_B4X16_P64) {
173 switch (num) {
174 case 1:
175 return llvm::Intrinsic::
176 nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b4x16_p64;
177 case 2:
178 return llvm::Intrinsic::
179 nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b4x16_p64;
180 case 4:
181 return llvm::Intrinsic::
182 nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b4x16_p64;
183 }
184 }
185 } else if (shape.getM() == 16 && shape.getN() == 16) {
186 if (eltType == NVVM::LdStMatrixEltType::B8) {
187 switch (num) {
188 case 1:
189 return llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8;
190 case 2:
191 return llvm::Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8;
192 }
193 } else if (eltType == NVVM::LdStMatrixEltType::B8X16_B6X16_P32) {
194 switch (num) {
195 case 1:
196 return llvm::Intrinsic::
197 nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b6x16_p32;
198 case 2:
199 return llvm::Intrinsic::
200 nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b6x16_p32;
201 }
202 } else if (eltType == NVVM::LdStMatrixEltType::B8X16_B4X16_P64) {
203 switch (num) {
204 case 1:
205 return llvm::Intrinsic::
206 nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b4x16_p64;
207 case 2:
208 return llvm::Intrinsic::
209 nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b4x16_p64;
210 }
211 }
212 }
213 llvm_unreachable("unknown ldmatrix kind");
214}
215
216/// Return the intrinsic ID associated with stmatrix for the given paramters.
217static llvm::Intrinsic::ID
218getStMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num,
219 NVVM::LdStMatrixShapeAttr shape,
220 NVVM::LdStMatrixEltType eltType) {
221 if (shape.getM() == 8 && shape.getN() == 8) {
222 switch (num) {
223 case 1:
224 return (layout == NVVM::MMALayout::row)
225 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x1_b16
226 : llvm::Intrinsic::
227 nvvm_stmatrix_sync_aligned_m8n8_x1_trans_b16;
228 case 2:
229 return (layout == NVVM::MMALayout::row)
230 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x2_b16
231 : llvm::Intrinsic::
232 nvvm_stmatrix_sync_aligned_m8n8_x2_trans_b16;
233 case 4:
234 return (layout == NVVM::MMALayout::row)
235 ? llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x4_b16
236 : llvm::Intrinsic::
237 nvvm_stmatrix_sync_aligned_m8n8_x4_trans_b16;
238 }
239 } else if (shape.getM() == 16 && shape.getN() == 8) {
240 switch (num) {
241 case 1:
242 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x1_trans_b8;
243 case 2:
244 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x2_trans_b8;
245 case 4:
246 return llvm::Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x4_trans_b8;
247 }
248 }
249 llvm_unreachable("unknown stmatrix kind");
250}
251
252static unsigned getUnidirectionalFenceProxyID(NVVM::ProxyKind fromProxy,
253 NVVM::ProxyKind toProxy,
254 NVVM::MemScopeKind scope,
255 bool isRelease) {
256 if (fromProxy == NVVM::ProxyKind::GENERIC &&
257 toProxy == NVVM::ProxyKind::TENSORMAP) {
258 switch (scope) {
259 case NVVM::MemScopeKind::CTA: {
260 if (isRelease)
261 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_cta;
262 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_cta;
263 }
264 case NVVM::MemScopeKind::CLUSTER: {
265 if (isRelease)
266 return llvm::Intrinsic::
267 nvvm_fence_proxy_tensormap_generic_release_cluster;
268 return llvm::Intrinsic::
269 nvvm_fence_proxy_tensormap_generic_acquire_cluster;
270 }
271 case NVVM::MemScopeKind::GPU: {
272 if (isRelease)
273 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_gpu;
274 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_gpu;
275 }
276 case NVVM::MemScopeKind::SYS: {
277 if (isRelease)
278 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_release_sys;
279 return llvm::Intrinsic::nvvm_fence_proxy_tensormap_generic_acquire_sys;
280 }
281 }
282 llvm_unreachable("Unknown scope for uni-directional fence.proxy operation");
283 }
284 llvm_unreachable("Unsupported proxy kinds");
285}
286
287static unsigned getMembarIntrinsicID(NVVM::MemScopeKind scope) {
288 switch (scope) {
289 case NVVM::MemScopeKind::CTA:
290 return llvm::Intrinsic::nvvm_membar_cta;
291 case NVVM::MemScopeKind::CLUSTER:
292 return llvm::Intrinsic::nvvm_fence_sc_cluster;
293 case NVVM::MemScopeKind::GPU:
294 return llvm::Intrinsic::nvvm_membar_gl;
295 case NVVM::MemScopeKind::SYS:
296 return llvm::Intrinsic::nvvm_membar_sys;
297 }
298 llvm_unreachable("Unknown scope for memory barrier");
299}
300
301#define TCGEN05LD(SHAPE, NUM) llvm::Intrinsic::nvvm_tcgen05_ld_##SHAPE##_##NUM
302
303static llvm::Intrinsic::ID
304getTcgen05LdIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num) {
305 llvm::Intrinsic::ID Shape16x64b[] = {
306 TCGEN05LD(16x64b, x1), TCGEN05LD(16x64b, x2), TCGEN05LD(16x64b, x4),
307 TCGEN05LD(16x64b, x8), TCGEN05LD(16x64b, x16), TCGEN05LD(16x64b, x32),
308 TCGEN05LD(16x64b, x64), TCGEN05LD(16x64b, x128),
309 };
310
311 llvm::Intrinsic::ID Shape16x128b[] = {
312 TCGEN05LD(16x128b, x1), TCGEN05LD(16x128b, x2), TCGEN05LD(16x128b, x4),
313 TCGEN05LD(16x128b, x8), TCGEN05LD(16x128b, x16), TCGEN05LD(16x128b, x32),
314 TCGEN05LD(16x128b, x64),
315 };
316
317 llvm::Intrinsic::ID Shape16x256b[] = {
318 TCGEN05LD(16x256b, x1), TCGEN05LD(16x256b, x2), TCGEN05LD(16x256b, x4),
319 TCGEN05LD(16x256b, x8), TCGEN05LD(16x256b, x16), TCGEN05LD(16x256b, x32),
320 };
321
322 llvm::Intrinsic::ID Shape16x32bx2[] = {
323 TCGEN05LD(16x32bx2, x1), TCGEN05LD(16x32bx2, x2),
324 TCGEN05LD(16x32bx2, x4), TCGEN05LD(16x32bx2, x8),
325 TCGEN05LD(16x32bx2, x16), TCGEN05LD(16x32bx2, x32),
326 TCGEN05LD(16x32bx2, x64), TCGEN05LD(16x32bx2, x128),
327 };
328
329 llvm::Intrinsic::ID Shape32x32b[] = {
330 TCGEN05LD(32x32b, x1), TCGEN05LD(32x32b, x2), TCGEN05LD(32x32b, x4),
331 TCGEN05LD(32x32b, x8), TCGEN05LD(32x32b, x16), TCGEN05LD(32x32b, x32),
332 TCGEN05LD(32x32b, x64), TCGEN05LD(32x32b, x128),
333 };
334
335 // `num` contains the length of vector and log2 of `num` returns the index
336 // into the shape array
337 unsigned Idx = std::log2(num);
338
339 switch (shape) {
340 case NVVM::Tcgen05LdStShape::SHAPE_16X64B:
341 return Shape16x64b[Idx];
342 case NVVM::Tcgen05LdStShape::SHAPE_16X128B:
343 return Shape16x128b[Idx - 1];
344 case NVVM::Tcgen05LdStShape::SHAPE_16X256B:
345 return Shape16x256b[Idx - 2];
346 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
347 return Shape32x32b[Idx];
348 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
349 return Shape16x32bx2[Idx];
350 }
351 llvm_unreachable("unhandled tcgen05.ld lowering");
352}
353
354#define TCGEN05ST(SHAPE, NUM) llvm::Intrinsic::nvvm_tcgen05_st_##SHAPE##_##NUM
355
356static llvm::Intrinsic::ID
357getTcgen05StIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num) {
358 llvm::Intrinsic::ID Shape16x64b[] = {
359 TCGEN05ST(16x64b, x1), TCGEN05ST(16x64b, x2), TCGEN05ST(16x64b, x4),
360 TCGEN05ST(16x64b, x8), TCGEN05ST(16x64b, x16), TCGEN05ST(16x64b, x32),
361 TCGEN05ST(16x64b, x64), TCGEN05ST(16x64b, x128),
362 };
363
364 llvm::Intrinsic::ID Shape16x128b[] = {
365 TCGEN05ST(16x128b, x1), TCGEN05ST(16x128b, x2), TCGEN05ST(16x128b, x4),
366 TCGEN05ST(16x128b, x8), TCGEN05ST(16x128b, x16), TCGEN05ST(16x128b, x32),
367 TCGEN05ST(16x128b, x64),
368 };
369
370 llvm::Intrinsic::ID Shape16x256b[] = {
371 TCGEN05ST(16x256b, x1), TCGEN05ST(16x256b, x2), TCGEN05ST(16x256b, x4),
372 TCGEN05ST(16x256b, x8), TCGEN05ST(16x256b, x16), TCGEN05ST(16x256b, x32),
373 };
374
375 llvm::Intrinsic::ID Shape16x32bx2[] = {
376 TCGEN05ST(16x32bx2, x1), TCGEN05ST(16x32bx2, x2),
377 TCGEN05ST(16x32bx2, x4), TCGEN05ST(16x32bx2, x8),
378 TCGEN05ST(16x32bx2, x16), TCGEN05ST(16x32bx2, x32),
379 TCGEN05ST(16x32bx2, x64), TCGEN05ST(16x32bx2, x128),
380 };
381
382 llvm::Intrinsic::ID Shape32x32b[] = {
383 TCGEN05ST(32x32b, x1), TCGEN05ST(32x32b, x2), TCGEN05ST(32x32b, x4),
384 TCGEN05ST(32x32b, x8), TCGEN05ST(32x32b, x16), TCGEN05ST(32x32b, x32),
385 TCGEN05ST(32x32b, x64), TCGEN05ST(32x32b, x128),
386 };
387
388 // `num` contains the length of vector and log2 of `num` returns the index
389 // into the shape array
390 unsigned Idx = std::log2(num);
391
392 switch (shape) {
393 case NVVM::Tcgen05LdStShape::SHAPE_16X64B:
394 return Shape16x64b[Idx];
395 case NVVM::Tcgen05LdStShape::SHAPE_16X128B:
396 return Shape16x128b[Idx - 1];
397 case NVVM::Tcgen05LdStShape::SHAPE_16X256B:
398 return Shape16x256b[Idx - 2];
399 case NVVM::Tcgen05LdStShape::SHAPE_32X32B:
400 return Shape32x32b[Idx];
401 case NVVM::Tcgen05LdStShape::SHAPE_16X32BX2:
402 return Shape16x32bx2[Idx];
403 }
404 llvm_unreachable("unhandled tcgen05.st lowering");
405}
406
407static llvm::Intrinsic::ID getFenceSyncRestrictID(NVVM::MemOrderKind order) {
408 return order == NVVM::MemOrderKind::ACQUIRE
409 ? llvm::Intrinsic::
410 nvvm_fence_acquire_sync_restrict_space_cluster_scope_cluster
411 : llvm::Intrinsic::
412 nvvm_fence_release_sync_restrict_space_cta_scope_cluster;
413}
414
415static llvm::Intrinsic::ID
416getFenceProxyID(NVVM::ProxyKind kind, std::optional<NVVM::SharedSpace> space) {
417 switch (kind) {
418 case NVVM::ProxyKind::alias:
419 return llvm::Intrinsic::nvvm_fence_proxy_alias;
420 case NVVM::ProxyKind::async:
421 return llvm::Intrinsic::nvvm_fence_proxy_async;
422 case NVVM::ProxyKind::async_global:
423 return llvm::Intrinsic::nvvm_fence_proxy_async_global;
424 case NVVM::ProxyKind::async_shared:
425 return *space == NVVM::SharedSpace::shared_cta
426 ? llvm::Intrinsic::nvvm_fence_proxy_async_shared_cta
427 : llvm::Intrinsic::nvvm_fence_proxy_async_shared_cluster;
428 default:
429 llvm_unreachable("unsupported proxy kind");
430 }
431}
432
433static llvm::Intrinsic::ID
434getFenceProxySyncRestrictID(NVVM::MemOrderKind order) {
435 return order == NVVM::MemOrderKind::ACQUIRE
436 ? llvm::Intrinsic::
437 nvvm_fence_proxy_async_generic_acquire_sync_restrict_space_cluster_scope_cluster
438 : llvm::Intrinsic::
439 nvvm_fence_proxy_async_generic_release_sync_restrict_space_cta_scope_cluster;
440}
441
442static llvm::RoundingMode
443getLLVMRoundingModeForFPArith(NVVM::FPRoundingMode rndMode) {
444 switch (rndMode) {
445 case NVVM::FPRoundingMode::RN:
446 return llvm::RoundingMode::NearestTiesToEven;
447 case NVVM::FPRoundingMode::RM:
448 return llvm::RoundingMode::TowardNegative;
449 case NVVM::FPRoundingMode::RP:
450 return llvm::RoundingMode::TowardPositive;
451 case NVVM::FPRoundingMode::RZ:
452 return llvm::RoundingMode::TowardZero;
453 default:
454 // default rounding mode is RN
455 assert(rndMode == NVVM::FPRoundingMode::NONE &&
456 "unsupported rounding mode for nvvm fp arithmetic");
457 return llvm::RoundingMode::NearestTiesToEven;
458 }
459}
460
461// Calls an LLVM intrinsic on the given operands. For f32/f64 vector types,
462// the intrinsic is called per-element and the results are packed back into a
463// vector. If retType is non-null, it is forwarded as the return-type
464// overload to `createIntrinsicCall`.
465static llvm::Value *
466createScalarizedIntrinsicCall(llvm::IRBuilderBase &builder,
467 llvm::Intrinsic::ID IID, llvm::Type *opTypeLLVM,
469 llvm::Type *retType) {
470 if (opTypeLLVM->isVectorTy() && (opTypeLLVM->getScalarType()->isFloatTy() ||
471 opTypeLLVM->getScalarType()->isDoubleTy())) {
472 llvm::Value *result = llvm::PoisonValue::get(
473 llvm::FixedVectorType::get(opTypeLLVM->getScalarType(), 2));
474 for (int64_t i = 0; i < 2; ++i) {
476 for (llvm::Value *op : operands)
477 scalarArgs.push_back(
478 op->getType()->isVectorTy()
479 ? builder.CreateExtractElement(op, builder.getInt32(i))
480 : op);
481 llvm::Value *res = createIntrinsicCall(builder, IID, retType, scalarArgs);
482 result = builder.CreateInsertElement(result, res, builder.getInt32(i));
483 }
484 return result;
485 }
486
487 return createIntrinsicCall(builder, IID, retType, operands);
488}
489
490void NVVM::AddFOp::lowerAddFToLLVMIR(llvm::Value *argLHS, llvm::Value *argRHS,
491 Value res, NVVM::FPRoundingMode rndMode,
492 NVVM::SaturationMode satMode, bool isFTZ,
494 llvm::IRBuilderBase &builder) {
495 llvm::Type *opTypeLLVM = argLHS->getType();
496 bool isSat = satMode != NVVM::SaturationMode::NONE;
497
498 static constexpr llvm::Intrinsic::ID addIDs[2][2] = {
499 {llvm::Intrinsic::nvvm_fadd, llvm::Intrinsic::nvvm_fadd_sat},
500 {llvm::Intrinsic::nvvm_fadd_ftz, llvm::Intrinsic::nvvm_fadd_ftz_sat}};
501
502 llvm::Intrinsic::ID id = addIDs[isFTZ][isSat];
503 llvm::Value *rnd = builder.getInt32(
504 static_cast<int>(getLLVMRoundingModeForFPArith(rndMode)));
505
506 // For f64 vector addition, and f32 vector addition with saturation,
507 // we need to scalarize the intrinsic call.
508 llvm::Type *scalarTypeLLVM = opTypeLLVM->getScalarType();
509 if (opTypeLLVM->isVectorTy() && (scalarTypeLLVM->isDoubleTy() ||
510 (isSat && scalarTypeLLVM->isFloatTy()))) {
511 mt.mapValue(res, createScalarizedIntrinsicCall(builder, id, opTypeLLVM,
512 {argLHS, argRHS, rnd},
513 scalarTypeLLVM));
514 return;
515 }
516
517 mt.mapValue(
518 res, createIntrinsicCall(builder, id, opTypeLLVM, {argLHS, argRHS, rnd}));
519}
520
521void NVVM::FmaOp::lowerFmaToLLVMIR(Operation &op, LLVM::ModuleTranslation &mt,
522 llvm::IRBuilderBase &builder) {
523 auto thisOp = cast<NVVM::FmaOp>(op);
524 mlir::NVVM::FPRoundingMode rndMode = thisOp.getRnd();
525 unsigned rndIndex = static_cast<unsigned>(rndMode) - 1; // 1-4 mapped to 0-3
526 mlir::NVVM::SaturationMode satMode = thisOp.getSat();
527 bool isFTZ = thisOp.getFtz();
528 bool isRelu = thisOp.getRelu();
529 bool isSat = satMode == NVVM::SaturationMode::SAT;
530 bool isOOB = thisOp.getOob();
531
532 mlir::Type opType = thisOp.getRes().getType();
533 llvm::Type *opTypeLLVM = mt.convertType(opType);
534 bool isVectorFma = opTypeLLVM->isVectorTy();
535
536 llvm::Value *argA = mt.lookupValue(thisOp.getA());
537 llvm::Value *argB = mt.lookupValue(thisOp.getB());
538 llvm::Value *argC = mt.lookupValue(thisOp.getC());
539
540 static constexpr llvm::Intrinsic::ID f16IDs[] = {
541 llvm::Intrinsic::nvvm_fma_rn_f16,
542 llvm::Intrinsic::nvvm_fma_rn_f16x2,
543 llvm::Intrinsic::nvvm_fma_rn_ftz_f16,
544 llvm::Intrinsic::nvvm_fma_rn_ftz_f16x2,
545 llvm::Intrinsic::nvvm_fma_rn_sat_f16,
546 llvm::Intrinsic::nvvm_fma_rn_sat_f16x2,
547 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f16,
548 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f16x2,
549 llvm::Intrinsic::nvvm_fma_rn_relu_f16,
550 llvm::Intrinsic::nvvm_fma_rn_relu_f16x2,
551 llvm::Intrinsic::nvvm_fma_rn_ftz_relu_f16,
552 llvm::Intrinsic::nvvm_fma_rn_ftz_relu_f16x2};
553
554 static constexpr llvm::Intrinsic::ID bf16IDs[] = {
555 llvm::Intrinsic::nvvm_fma_rn_bf16, llvm::Intrinsic::nvvm_fma_rn_bf16x2,
556 llvm::Intrinsic::nvvm_fma_rn_relu_bf16,
557 llvm::Intrinsic::nvvm_fma_rn_relu_bf16x2};
558
559 static constexpr llvm::Intrinsic::ID f32IDs[] = {
560 llvm::Intrinsic::nvvm_fma_rn_f,
561 llvm::Intrinsic::nvvm_fma_rm_f,
562 llvm::Intrinsic::nvvm_fma_rp_f,
563 llvm::Intrinsic::nvvm_fma_rz_f,
564 llvm::Intrinsic::nvvm_fma_rn_sat_f,
565 llvm::Intrinsic::nvvm_fma_rm_sat_f,
566 llvm::Intrinsic::nvvm_fma_rp_sat_f,
567 llvm::Intrinsic::nvvm_fma_rz_sat_f,
568 llvm::Intrinsic::nvvm_fma_rn_ftz_f,
569 llvm::Intrinsic::nvvm_fma_rm_ftz_f,
570 llvm::Intrinsic::nvvm_fma_rp_ftz_f,
571 llvm::Intrinsic::nvvm_fma_rz_ftz_f,
572 llvm::Intrinsic::nvvm_fma_rn_ftz_sat_f,
573 llvm::Intrinsic::nvvm_fma_rm_ftz_sat_f,
574 llvm::Intrinsic::nvvm_fma_rp_ftz_sat_f,
575 llvm::Intrinsic::nvvm_fma_rz_ftz_sat_f,
576 };
577
578 static constexpr llvm::Intrinsic::ID f64IDs[] = {
579 llvm::Intrinsic::nvvm_fma_rn_d, llvm::Intrinsic::nvvm_fma_rm_d,
580 llvm::Intrinsic::nvvm_fma_rp_d, llvm::Intrinsic::nvvm_fma_rz_d};
581
582 auto fmaIntrinsic = [&](llvm::Intrinsic::ID IID,
583 llvm::Type *retType) -> llvm::Value * {
585 builder, IID, opTypeLLVM, {argA, argB, argC}, /*retType=*/retType);
586 };
587
588 // f16 + f16 -> f16 / vector<2xf16> + vector<2xf16> -> vector<2xf16>
589 if (opTypeLLVM->getScalarType()->isHalfTy()) {
590 llvm::Value *result;
591 if (isOOB) {
592 result = fmaIntrinsic(isRelu ? llvm::Intrinsic::nvvm_fma_rn_oob_relu
593 : llvm::Intrinsic::nvvm_fma_rn_oob,
594 opTypeLLVM);
595 } else {
596 unsigned index =
597 (isRelu << 3) | (isSat << 2) | (isFTZ << 1) |
598 isVectorFma; // Op verifier ensures that this index is valid
599 result = fmaIntrinsic(f16IDs[index], opTypeLLVM);
600 }
601 mt.mapValue(thisOp.getRes(), result);
602 return;
603 }
604
605 // bf16 + bf16 -> bf16 / vector<2xbf16> + vector<2xbf16> -> vector<2xbf16>
606 if (opTypeLLVM->getScalarType()->isBFloatTy()) {
607 llvm::Value *result;
608 if (isOOB) {
609 result = fmaIntrinsic(isRelu ? llvm::Intrinsic::nvvm_fma_rn_oob_relu
610 : llvm::Intrinsic::nvvm_fma_rn_oob,
611 opTypeLLVM);
612 } else {
613 unsigned index = (isRelu << 1) | isVectorFma;
614 result = fmaIntrinsic(bf16IDs[index], opTypeLLVM);
615 }
616 mt.mapValue(thisOp.getRes(), result);
617 return;
618 }
619
620 // f64 + f64 -> f64 / vector<2xf64> + vector<2xf64> -> vector<2xf64>
621 if (opTypeLLVM->getScalarType()->isDoubleTy()) {
622 mt.mapValue(thisOp.getRes(),
623 fmaIntrinsic(f64IDs[rndIndex], opTypeLLVM->getScalarType()));
624 return;
625 }
626
627 // f32 + f32 -> f32 / vector<2xf32> + vector<2xf32> -> vector<2xf32>
628 const unsigned numRndModes = 4; // RN, RM, RP, RZ
629 if (opTypeLLVM->getScalarType()->isFloatTy()) {
630 unsigned index = ((isFTZ << 1) | isSat) * numRndModes + rndIndex;
631 mt.mapValue(thisOp.getRes(),
632 fmaIntrinsic(f32IDs[index], opTypeLLVM->getScalarType()));
633 return;
634 }
635}
636
637namespace {
638/// Implementation of the dialect interface that converts operations belonging
639/// to the NVVM dialect to LLVM IR.
640class NVVMDialectLLVMIRTranslationInterface
641 : public LLVMTranslationDialectInterface {
642public:
643 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
644
645 /// Translates the given operation to LLVM IR using the provided IR builder
646 /// and saving the state in `moduleTranslation`.
647 LogicalResult
648 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
649 LLVM::ModuleTranslation &moduleTranslation) const final {
650 // All NVVM ops are instruction-level and require an active insertion point.
651 // A null insert block means the op is misplaced (e.g., at module scope),
652 // which would otherwise cause a null dereference in createIntrinsicCall.
653 if (!builder.GetInsertBlock())
654 return op->emitOpError(
655 "cannot be translated to LLVM IR without an active insertion "
656 "point; make sure the op is inside a function");
657 Operation &opInst = *op;
658#include "mlir/Dialect/LLVMIR/NVVMConversions.inc"
659
660 return failure();
661 }
662
663 /// Attaches module-level metadata for functions marked as kernels
664 /// and managed annotations for global variables.
665 LogicalResult
666 amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
667 NamedAttribute attribute,
668 LLVM::ModuleTranslation &moduleTranslation) const final {
669 if (auto globalOp = dyn_cast<LLVM::GlobalOp>(op)) {
670 if (attribute.getName() == NVVM::NVVMDialect::getManagedAttrName()) {
671 auto *gv = cast<llvm::GlobalVariable>(
672 moduleTranslation.lookupGlobal(globalOp));
673 llvm::Module *m = gv->getParent();
674 llvm::LLVMContext &ctx = m->getContext();
675 llvm::NamedMDNode *md = m->getOrInsertNamedMetadata("nvvm.annotations");
676 md->addOperand(llvm::MDNode::get(
677 ctx, {llvm::ConstantAsMetadata::get(gv),
678 llvm::MDString::get(ctx, "managed"),
679 llvm::ConstantAsMetadata::get(llvm::ConstantInt::get(
680 llvm::Type::getInt32Ty(ctx), 1))}));
681 }
682 return success();
683 }
684
685 auto func = dyn_cast<LLVM::LLVMFuncOp>(op);
686 if (!func)
687 return failure();
688 llvm::Function *llvmFunc = moduleTranslation.lookupFunction(func.getName());
689
690 if (attribute.getName() == NVVM::NVVMDialect::getMaxntidAttrName()) {
691 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
692 return failure();
693 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
694 const std::string attr = llvm::formatv(
695 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
696 values.asArrayRef().end()));
697 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxNTID, attr);
698 } else if (attribute.getName() == NVVM::NVVMDialect::getReqntidAttrName()) {
699 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
700 return failure();
701 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
702 const std::string attr = llvm::formatv(
703 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
704 values.asArrayRef().end()));
705 llvmFunc->addFnAttr(llvm::NVVMAttr::ReqNTID, attr);
706 } else if (attribute.getName() ==
707 NVVM::NVVMDialect::getClusterDimAttrName()) {
708 if (!isa<DenseI32ArrayAttr>(attribute.getValue()))
709 return failure();
710 auto values = cast<DenseI32ArrayAttr>(attribute.getValue());
711 const std::string attr = llvm::formatv(
712 "{0:$[,]}", llvm::make_range(values.asArrayRef().begin(),
713 values.asArrayRef().end()));
714 llvmFunc->addFnAttr(llvm::NVVMAttr::ClusterDim, attr);
715 } else if (attribute.getName() ==
716 NVVM::NVVMDialect::getClusterMaxBlocksAttrName()) {
717 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
718 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxClusterRank,
719 llvm::utostr(value.getInt()));
720 } else if (attribute.getName() ==
721 NVVM::NVVMDialect::getMinctasmAttrName()) {
722 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
723 llvmFunc->addFnAttr(llvm::NVVMAttr::MinCTASm,
724 llvm::utostr(value.getInt()));
725 } else if (attribute.getName() == NVVM::NVVMDialect::getMaxnregAttrName()) {
726 auto value = dyn_cast<IntegerAttr>(attribute.getValue());
727 llvmFunc->addFnAttr(llvm::NVVMAttr::MaxNReg,
728 llvm::utostr(value.getInt()));
729 } else if (attribute.getName() ==
730 NVVM::NVVMDialect::getKernelFuncAttrName()) {
731 llvmFunc->setCallingConv(llvm::CallingConv::PTX_Kernel);
732 } else if (attribute.getName() ==
733 NVVM::NVVMDialect::getBlocksAreClustersAttrName()) {
734 llvmFunc->addFnAttr(llvm::NVVMAttr::BlocksAreClusters);
735 }
736
737 return success();
738 }
739
740 LogicalResult
741 convertParameterAttr(LLVMFuncOp funcOp, int argIdx, NamedAttribute attribute,
742 LLVM::ModuleTranslation &moduleTranslation) const final {
743
744 llvm::LLVMContext &llvmContext = moduleTranslation.getLLVMContext();
745 llvm::Function *llvmFunc =
746 moduleTranslation.lookupFunction(funcOp.getName());
747
748 if (attribute.getName() == NVVM::NVVMDialect::getGridConstantAttrName()) {
749 llvmFunc->addParamAttr(
750 argIdx,
751 llvm::Attribute::get(llvmContext, llvm::NVVMAttr::GridConstant));
752 }
753 return success();
754 }
755};
756} // namespace
757
759 registry.insert<NVVM::NVVMDialect>();
760 registry.addExtension(+[](MLIRContext *ctx, NVVM::NVVMDialect *dialect) {
761 dialect->addInterfaces<NVVMDialectLLVMIRTranslationInterface>();
762 });
763}
764
766 DialectRegistry registry;
768 context.appendDialectRegistry(registry);
769}
return success()
static LogicalResult convertParameterAttr(llvm::AttrBuilder &attrBuilder, llvm::Attribute::AttrKind llvmKind, NamedAttribute namedAttr, ModuleTranslation &moduleTranslation, Location loc)
static llvm::Intrinsic::ID getLdMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num, NVVM::LdStMatrixShapeAttr shape, NVVM::LdStMatrixEltType eltType)
static llvm::Intrinsic::ID getFenceProxyID(NVVM::ProxyKind kind, std::optional< NVVM::SharedSpace > space)
#define GET_REDUX_F32_ID(op, hasAbs, hasNaN)
static llvm::Intrinsic::ID getStMatrixIntrinsicId(NVVM::MMALayout layout, int32_t num, NVVM::LdStMatrixShapeAttr shape, NVVM::LdStMatrixEltType eltType)
Return the intrinsic ID associated with stmatrix for the given paramters.
static llvm::Intrinsic::ID getTcgen05StIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num)
static llvm::Intrinsic::ID getTcgen05LdIntrinsicID(mlir::NVVM::Tcgen05LdStShape shape, uint32_t num)
static unsigned getMembarIntrinsicID(NVVM::MemScopeKind scope)
static unsigned getUnidirectionalFenceProxyID(NVVM::ProxyKind fromProxy, NVVM::ProxyKind toProxy, NVVM::MemScopeKind scope, bool isRelease)
llvm::CallInst * createIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID intrinsic, ArrayRef< llvm::Value * > args={}, ArrayRef< llvm::Type * > tys={})
Creates a call to an LLVM IR intrinsic function with the given arguments.
static llvm::Intrinsic::ID getFenceProxySyncRestrictID(NVVM::MemOrderKind order)
#define TCGEN05ST(SHAPE, NUM)
static llvm::Intrinsic::ID getReduxIntrinsicId(llvm::Type *resultType, NVVM::ReductionKind kind, bool hasAbs, bool hasNaN)
static llvm::RoundingMode getLLVMRoundingModeForFPArith(NVVM::FPRoundingMode rndMode)
static llvm::Value * createScalarizedIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID IID, llvm::Type *opTypeLLVM, ArrayRef< llvm::Value * > operands, llvm::Type *retType)
#define TCGEN05LD(SHAPE, NUM)
static llvm::Intrinsic::ID getFenceSyncRestrictID(NVVM::MemOrderKind order)
static llvm::Intrinsic::ID getShflIntrinsicId(llvm::Type *resultType, NVVM::ShflKind kind, bool withPredicate)
static llvm::Intrinsic::ID getVoteSyncIntrinsicId(NVVM::VoteSyncKind kind)
static llvm::Intrinsic::ID getMatchSyncIntrinsicId(Type valType, NVVM::MatchSyncKind kind)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Value min(ImplicitLocOpBuilder &builder, Value value, Value bound)
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Implementation class for module translation.
llvm::Value * lookupValue(Value value) const
Finds an LLVM IR value corresponding to the given MLIR value.
llvm::Type * convertType(Type type)
Converts the type from MLIR LLVM dialect to LLVM.
void mapValue(Value mlir, llvm::Value *llvm)
Stores the mapping between an MLIR value and its LLVM IR counterpart.
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
void appendDialectRegistry(const DialectRegistry &registry)
Append the contents of the given dialect registry to the registry associated with this context.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
llvm::CallInst * createIntrinsicCall(llvm::IRBuilderBase &builder, llvm::Intrinsic::ID intrinsic, ArrayRef< llvm::Value * > args={}, ArrayRef< llvm::Type * > tys={})
Creates a call to an LLVM IR intrinsic function with the given arguments.
Include the generated interface declarations.
void registerNVVMDialectTranslation(DialectRegistry &registry)
Register the NVVM dialect and the translation from it to the LLVM IR in the given registry;.