clang 24.0.0git
CIRGenBuiltinNVPTX.cpp
Go to the documentation of this file.
1//===---- CIRGenBuiltinNVPTX.cpp - Emit CIR for NVPTX builtins ------------===//
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 contains code to emit NVPTX Builtin calls.
10//
11//===----------------------------------------------------------------------===//
12
13#include "CIRGenFunction.h"
14
15#include "mlir/IR/Value.h"
19#include "llvm/Support/NVPTXAddrSpace.h"
20
21using namespace clang;
22using namespace clang::CIRGen;
23
24static mlir::Value makeLdu(CIRGenFunction &cgf, const CallExpr *expr,
25 llvm::StringRef intrinsicName) {
26 auto &builder = cgf.getBuilder();
27 Address ptr = cgf.emitPointerWithAlignment(expr->getArg(0));
28 QualType argType = expr->getArg(0)->getType();
29 mlir::Type elemTy = cgf.convertTypeForMem(argType->getPointeeType());
30 clang::CharUnits align = ptr.getAlignment();
31 mlir::Location loc = cgf.getLoc(expr->getExprLoc());
32 mlir::Value alignVal =
33 builder.getConstantInt(loc, builder.getSInt32Ty(), align.getQuantity());
34 return cir::LLVMIntrinsicCallOp::create(
35 builder, loc, builder.getStringAttr(intrinsicName), elemTy,
36 {ptr.emitRawPointer(), alignVal})
37 .getResult();
38}
39
40static mlir::Value makeLdg(CIRGenFunction &cgf, const CallExpr *expr) {
41 auto &builder = cgf.getBuilder();
42 Address ptr = cgf.emitPointerWithAlignment(expr->getArg(0));
43 QualType argType = expr->getArg(0)->getType();
44 mlir::Type elemTy = cgf.convertTypeForMem(argType->getPointeeType());
45 mlir::Location loc = cgf.getLoc(expr->getExprLoc());
46
47 // Use addrspace(1) for NVPTX ADDRESS_SPACE_GLOBAL.
48 mlir::Type globalPtrTy = cir::PointerType::get(
49 elemTy, cir::TargetAddressSpaceAttr::get(
50 builder.getContext(), llvm::NVPTXAS::ADDRESS_SPACE_GLOBAL));
51 mlir::Value asc =
52 builder.createAddrSpaceCast(loc, ptr.getPointer(), globalPtrTy);
53 cir::LoadOp load =
54 builder.createAlignedLoad(loc, elemTy, asc, ptr.getAlignment());
55 load.setInvariant(true);
56 return load.getResult();
57}
58
59/// Emit a CIR LLVMIntrinsicCallOp for a unary NVVM intrinsic.
60/// The result type is inferred from the single argument.
61static mlir::Value emitUnaryNVVMIntrinsic(CIRGenFunction &cgf,
62 const CallExpr *expr,
63 llvm::StringRef intrinsicName) {
64 auto &builder = cgf.getBuilder();
65 mlir::Value arg = cgf.emitScalarExpr(expr->getArg(0));
66 return cir::LLVMIntrinsicCallOp::create(
67 builder, cgf.getLoc(expr->getExprLoc()),
68 builder.getStringAttr(intrinsicName), arg.getType(), {arg})
69 .getResult();
70}
71
72/// Emit a CIR LLVMIntrinsicCallOp for an NVVM fadd intrinsic, which takes the
73/// rounding mode as a trailing operand.
74static mlir::Value emitNVVMFAdd(CIRGenFunction &cgf, const CallExpr *expr,
75 llvm::StringRef intrinsicName,
76 llvm::APFloat::roundingMode rm) {
77 auto &builder = cgf.getBuilder();
78 mlir::Location loc = cgf.getLoc(expr->getExprLoc());
79 mlir::Value lhs = cgf.emitScalarExpr(expr->getArg(0));
80 mlir::Value rhs = cgf.emitScalarExpr(expr->getArg(1));
81 mlir::Value rnd =
82 builder.getConstInt(loc, builder.getSInt32Ty(), static_cast<int>(rm));
83 return cir::LLVMIntrinsicCallOp::create(builder, loc,
84 builder.getStringAttr(intrinsicName),
85 lhs.getType(), {lhs, rhs, rnd})
86 .getResult();
87}
88
89static mlir::Value emitBar0Reduction(CIRGenFunction &cgf, const CallExpr *expr,
90 llvm::StringRef intrinsicName,
91 bool returnsPred) {
92 CIRGenBuilderTy &builder = cgf.getBuilder();
93 mlir::Location loc = cgf.getLoc(expr->getExprLoc());
94 mlir::Type si32Ty = builder.getSInt32Ty();
95 mlir::Value zero = builder.getNullValue(si32Ty, loc);
96 mlir::Value pred = builder.createCompare(
97 loc, cir::CmpOpKind::ne, cgf.emitScalarExpr(expr->getArg(0)), zero);
98 mlir::Type resultTy = returnsPred ? mlir::Type(builder.getBoolTy()) : si32Ty;
99 mlir::Value result = builder.emitIntrinsicCallOp(
100 loc, intrinsicName, resultTy, mlir::ValueRange{zero, pred});
101 if (returnsPred)
102 result = builder.createBoolToInt(result, si32Ty);
103 return result;
104}
105
106static mlir::Value makeScopedAtomicRMW(CIRGenFunction &cgf,
107 const CallExpr *expr,
108 cir::AtomicFetchKind kind,
109 cir::SyncScopeKind scope) {
110 auto &builder = cgf.getBuilder();
111 Address destAddr = cgf.emitPointerWithAlignment(expr->getArg(0));
112 mlir::Value destValue = destAddr.emitRawPointer();
113 mlir::Value val = cgf.emitScalarExpr(expr->getArg(1));
114 auto rmwi = cir::AtomicFetchOp::create(
115 builder, cgf.getLoc(expr->getSourceRange()), destValue, val, kind,
116 cir::MemOrder::Relaxed, scope, /*is_volatile=*/false,
117 /*fetch_first=*/true);
118 return rmwi->getResult(0);
119}
120
121static mlir::Value makeScopedAtomicXchg(CIRGenFunction &cgf,
122 const CallExpr *expr,
123 cir::SyncScopeKind scope) {
124 auto &builder = cgf.getBuilder();
125 Address destAddr = cgf.emitPointerWithAlignment(expr->getArg(0));
126 mlir::Value destValue = destAddr.emitRawPointer();
127 mlir::Value val = cgf.emitScalarExpr(expr->getArg(1));
128 auto xchg = cir::AtomicXchgOp::create(
129 builder, cgf.getLoc(expr->getSourceRange()), destValue, val,
130 cir::MemOrder::Relaxed, scope, /*is_volatile=*/false);
131 return xchg.getResult();
132}
133
134std::optional<mlir::Value>
136 switch (builtinId) {
137 case NVPTX::BI__nvvm_atom_add_gen_i:
138 case NVPTX::BI__nvvm_atom_add_gen_l:
139 case NVPTX::BI__nvvm_atom_add_gen_ll:
140 return makeBinaryAtomicValue(cir::AtomicFetchKind::Add, expr,
141 /*originalArgType=*/nullptr,
142 /*emittedArgValue=*/nullptr,
143 cir::MemOrder::Relaxed);
144 case NVPTX::BI__nvvm_atom_sub_gen_i:
145 case NVPTX::BI__nvvm_atom_sub_gen_l:
146 case NVPTX::BI__nvvm_atom_sub_gen_ll:
147 return makeBinaryAtomicValue(cir::AtomicFetchKind::Sub, expr,
148 /*originalArgType=*/nullptr,
149 /*emittedArgValue=*/nullptr,
150 cir::MemOrder::Relaxed);
151 case NVPTX::BI__nvvm_atom_and_gen_i:
152 case NVPTX::BI__nvvm_atom_and_gen_l:
153 case NVPTX::BI__nvvm_atom_and_gen_ll:
154 return makeBinaryAtomicValue(cir::AtomicFetchKind::And, expr,
155 /*originalArgType=*/nullptr,
156 /*emittedArgValue=*/nullptr,
157 cir::MemOrder::Relaxed);
158 case NVPTX::BI__nvvm_atom_or_gen_i:
159 case NVPTX::BI__nvvm_atom_or_gen_l:
160 case NVPTX::BI__nvvm_atom_or_gen_ll:
161 return makeBinaryAtomicValue(cir::AtomicFetchKind::Or, expr,
162 /*originalArgType=*/nullptr,
163 /*emittedArgValue=*/nullptr,
164 cir::MemOrder::Relaxed);
165 case NVPTX::BI__nvvm_atom_xor_gen_i:
166 case NVPTX::BI__nvvm_atom_xor_gen_l:
167 case NVPTX::BI__nvvm_atom_xor_gen_ll:
168 return makeBinaryAtomicValue(cir::AtomicFetchKind::Xor, expr,
169 /*originalArgType=*/nullptr,
170 /*emittedArgValue=*/nullptr,
171 cir::MemOrder::Relaxed);
172 case NVPTX::BI__nvvm_atom_xchg_gen_i:
173 case NVPTX::BI__nvvm_atom_xchg_gen_l:
174 case NVPTX::BI__nvvm_atom_xchg_gen_ll:
175 return makeScopedAtomicXchg(*this, expr, cir::SyncScopeKind::System);
176 case NVPTX::BI__nvvm_atom_max_gen_i:
177 case NVPTX::BI__nvvm_atom_max_gen_l:
178 case NVPTX::BI__nvvm_atom_max_gen_ll:
179 return makeBinaryAtomicValue(cir::AtomicFetchKind::Max, expr,
180 /*originalArgType=*/nullptr,
181 /*emittedArgValue=*/nullptr,
182 cir::MemOrder::Relaxed);
183 case NVPTX::BI__nvvm_atom_max_gen_ui:
184 case NVPTX::BI__nvvm_atom_max_gen_ul:
185 case NVPTX::BI__nvvm_atom_max_gen_ull:
186 return makeBinaryAtomicValue(cir::AtomicFetchKind::Max, expr,
187 /*originalArgType=*/nullptr,
188 /*emittedArgValue=*/nullptr,
189 cir::MemOrder::Relaxed);
190 case NVPTX::BI__nvvm_atom_min_gen_i:
191 case NVPTX::BI__nvvm_atom_min_gen_l:
192 case NVPTX::BI__nvvm_atom_min_gen_ll:
193 return makeBinaryAtomicValue(cir::AtomicFetchKind::Min, expr,
194 /*originalArgType=*/nullptr,
195 /*emittedArgValue=*/nullptr,
196 cir::MemOrder::Relaxed);
197 case NVPTX::BI__nvvm_atom_min_gen_ui:
198 case NVPTX::BI__nvvm_atom_min_gen_ul:
199 case NVPTX::BI__nvvm_atom_min_gen_ull:
200 return makeBinaryAtomicValue(cir::AtomicFetchKind::Min, expr,
201 /*originalArgType=*/nullptr,
202 /*emittedArgValue=*/nullptr,
203 cir::MemOrder::Relaxed);
204 case NVPTX::BI__nvvm_atom_cas_gen_us:
205 case NVPTX::BI__nvvm_atom_cas_gen_i:
206 case NVPTX::BI__nvvm_atom_cas_gen_l:
207 case NVPTX::BI__nvvm_atom_cas_gen_ll:
208 return emitAtomicCmpXchg(expr, /*returnBool=*/false, cir::MemOrder::Relaxed,
209 cir::MemOrder::Relaxed,
210 cir::SyncScopeKind::System);
211 case NVPTX::BI__nvvm_atom_add_gen_f:
212 case NVPTX::BI__nvvm_atom_add_gen_d:
213 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Add,
214 cir::SyncScopeKind::System);
215 case NVPTX::BI__nvvm_atom_inc_gen_ui:
216 return makeBinaryAtomicValue(cir::AtomicFetchKind::UIncWrap, expr,
217 /*originalArgType=*/nullptr,
218 /*emittedArgValue=*/nullptr,
219 cir::MemOrder::Relaxed);
220 case NVPTX::BI__nvvm_atom_dec_gen_ui:
221 return makeBinaryAtomicValue(cir::AtomicFetchKind::UDecWrap, expr,
222 /*originalArgType=*/nullptr,
223 /*emittedArgValue=*/nullptr,
224 cir::MemOrder::Relaxed);
225 case NVPTX::BI__nvvm_ldg_c:
226 case NVPTX::BI__nvvm_ldg_sc:
227 case NVPTX::BI__nvvm_ldg_c2:
228 case NVPTX::BI__nvvm_ldg_sc2:
229 case NVPTX::BI__nvvm_ldg_c4:
230 case NVPTX::BI__nvvm_ldg_sc4:
231 case NVPTX::BI__nvvm_ldg_s:
232 case NVPTX::BI__nvvm_ldg_s2:
233 case NVPTX::BI__nvvm_ldg_s4:
234 case NVPTX::BI__nvvm_ldg_i:
235 case NVPTX::BI__nvvm_ldg_i2:
236 case NVPTX::BI__nvvm_ldg_i4:
237 case NVPTX::BI__nvvm_ldg_l:
238 case NVPTX::BI__nvvm_ldg_l2:
239 case NVPTX::BI__nvvm_ldg_ll:
240 case NVPTX::BI__nvvm_ldg_ll2:
241 case NVPTX::BI__nvvm_ldg_uc:
242 case NVPTX::BI__nvvm_ldg_uc2:
243 case NVPTX::BI__nvvm_ldg_uc4:
244 case NVPTX::BI__nvvm_ldg_us:
245 case NVPTX::BI__nvvm_ldg_us2:
246 case NVPTX::BI__nvvm_ldg_us4:
247 case NVPTX::BI__nvvm_ldg_ui:
248 case NVPTX::BI__nvvm_ldg_ui2:
249 case NVPTX::BI__nvvm_ldg_ui4:
250 case NVPTX::BI__nvvm_ldg_ul:
251 case NVPTX::BI__nvvm_ldg_ul2:
252 case NVPTX::BI__nvvm_ldg_ull:
253 case NVPTX::BI__nvvm_ldg_ull2:
254 case NVPTX::BI__nvvm_ldg_f:
255 case NVPTX::BI__nvvm_ldg_f2:
256 case NVPTX::BI__nvvm_ldg_f4:
257 case NVPTX::BI__nvvm_ldg_d:
258 case NVPTX::BI__nvvm_ldg_d2:
259 return makeLdg(*this, expr);
260 case NVPTX::BI__nvvm_ldu_c:
261 case NVPTX::BI__nvvm_ldu_sc:
262 case NVPTX::BI__nvvm_ldu_c2:
263 case NVPTX::BI__nvvm_ldu_sc2:
264 case NVPTX::BI__nvvm_ldu_c4:
265 case NVPTX::BI__nvvm_ldu_sc4:
266 case NVPTX::BI__nvvm_ldu_s:
267 case NVPTX::BI__nvvm_ldu_s2:
268 case NVPTX::BI__nvvm_ldu_s4:
269 case NVPTX::BI__nvvm_ldu_i:
270 case NVPTX::BI__nvvm_ldu_i2:
271 case NVPTX::BI__nvvm_ldu_i4:
272 case NVPTX::BI__nvvm_ldu_l:
273 case NVPTX::BI__nvvm_ldu_l2:
274 case NVPTX::BI__nvvm_ldu_ll:
275 case NVPTX::BI__nvvm_ldu_ll2:
276 case NVPTX::BI__nvvm_ldu_uc:
277 case NVPTX::BI__nvvm_ldu_uc2:
278 case NVPTX::BI__nvvm_ldu_uc4:
279 case NVPTX::BI__nvvm_ldu_us:
280 case NVPTX::BI__nvvm_ldu_us2:
281 case NVPTX::BI__nvvm_ldu_us4:
282 case NVPTX::BI__nvvm_ldu_ui:
283 case NVPTX::BI__nvvm_ldu_ui2:
284 case NVPTX::BI__nvvm_ldu_ui4:
285 case NVPTX::BI__nvvm_ldu_ul:
286 case NVPTX::BI__nvvm_ldu_ul2:
287 case NVPTX::BI__nvvm_ldu_ull:
288 case NVPTX::BI__nvvm_ldu_ull2:
289 return makeLdu(*this, expr, "nvvm.ldu.global.i");
290 case NVPTX::BI__nvvm_ldu_f:
291 case NVPTX::BI__nvvm_ldu_f2:
292 case NVPTX::BI__nvvm_ldu_f4:
293 case NVPTX::BI__nvvm_ldu_d:
294 case NVPTX::BI__nvvm_ldu_d2:
295 return makeLdu(*this, expr, "nvvm.ldu.global.f");
296 case NVPTX::BI__nvvm_atom_cta_add_gen_i:
297 case NVPTX::BI__nvvm_atom_cta_add_gen_l:
298 case NVPTX::BI__nvvm_atom_cta_add_gen_ll:
299 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Add,
300 cir::SyncScopeKind::Workgroup);
301 case NVPTX::BI__nvvm_atom_sys_add_gen_i:
302 case NVPTX::BI__nvvm_atom_sys_add_gen_l:
303 case NVPTX::BI__nvvm_atom_sys_add_gen_ll:
304 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Add,
305 cir::SyncScopeKind::System);
306 case NVPTX::BI__nvvm_atom_cta_add_gen_f:
307 case NVPTX::BI__nvvm_atom_cta_add_gen_d:
308 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Add,
309 cir::SyncScopeKind::Workgroup);
310 case NVPTX::BI__nvvm_atom_sys_add_gen_f:
311 case NVPTX::BI__nvvm_atom_sys_add_gen_d:
312 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Add,
313 cir::SyncScopeKind::System);
314 case NVPTX::BI__nvvm_atom_cta_xchg_gen_i:
315 case NVPTX::BI__nvvm_atom_cta_xchg_gen_l:
316 case NVPTX::BI__nvvm_atom_cta_xchg_gen_ll:
317 return makeScopedAtomicXchg(*this, expr, cir::SyncScopeKind::Workgroup);
318 case NVPTX::BI__nvvm_atom_sys_xchg_gen_i:
319 case NVPTX::BI__nvvm_atom_sys_xchg_gen_l:
320 case NVPTX::BI__nvvm_atom_sys_xchg_gen_ll:
321 return makeScopedAtomicXchg(*this, expr, cir::SyncScopeKind::System);
322 case NVPTX::BI__nvvm_atom_cta_max_gen_i:
323 case NVPTX::BI__nvvm_atom_cta_max_gen_ui:
324 case NVPTX::BI__nvvm_atom_cta_max_gen_l:
325 case NVPTX::BI__nvvm_atom_cta_max_gen_ul:
326 case NVPTX::BI__nvvm_atom_cta_max_gen_ll:
327 case NVPTX::BI__nvvm_atom_cta_max_gen_ull:
328 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Max,
329 cir::SyncScopeKind::Workgroup);
330 case NVPTX::BI__nvvm_atom_sys_max_gen_i:
331 case NVPTX::BI__nvvm_atom_sys_max_gen_ui:
332 case NVPTX::BI__nvvm_atom_sys_max_gen_l:
333 case NVPTX::BI__nvvm_atom_sys_max_gen_ul:
334 case NVPTX::BI__nvvm_atom_sys_max_gen_ll:
335 case NVPTX::BI__nvvm_atom_sys_max_gen_ull:
336 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Max,
337 cir::SyncScopeKind::System);
338 case NVPTX::BI__nvvm_atom_cta_min_gen_i:
339 case NVPTX::BI__nvvm_atom_cta_min_gen_ui:
340 case NVPTX::BI__nvvm_atom_cta_min_gen_l:
341 case NVPTX::BI__nvvm_atom_cta_min_gen_ul:
342 case NVPTX::BI__nvvm_atom_cta_min_gen_ll:
343 case NVPTX::BI__nvvm_atom_cta_min_gen_ull:
344 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Min,
345 cir::SyncScopeKind::Workgroup);
346 case NVPTX::BI__nvvm_atom_sys_min_gen_i:
347 case NVPTX::BI__nvvm_atom_sys_min_gen_ui:
348 case NVPTX::BI__nvvm_atom_sys_min_gen_l:
349 case NVPTX::BI__nvvm_atom_sys_min_gen_ul:
350 case NVPTX::BI__nvvm_atom_sys_min_gen_ll:
351 case NVPTX::BI__nvvm_atom_sys_min_gen_ull:
352 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Min,
353 cir::SyncScopeKind::System);
354 case NVPTX::BI__nvvm_atom_cta_inc_gen_ui:
355 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::UIncWrap,
356 cir::SyncScopeKind::Workgroup);
357 case NVPTX::BI__nvvm_atom_cta_dec_gen_ui:
358 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::UDecWrap,
359 cir::SyncScopeKind::Workgroup);
360 case NVPTX::BI__nvvm_atom_sys_inc_gen_ui:
361 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::UIncWrap,
362 cir::SyncScopeKind::System);
363 case NVPTX::BI__nvvm_atom_sys_dec_gen_ui:
364 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::UDecWrap,
365 cir::SyncScopeKind::System);
366 case NVPTX::BI__nvvm_atom_cta_and_gen_i:
367 case NVPTX::BI__nvvm_atom_cta_and_gen_l:
368 case NVPTX::BI__nvvm_atom_cta_and_gen_ll:
369 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::And,
370 cir::SyncScopeKind::Workgroup);
371 case NVPTX::BI__nvvm_atom_sys_and_gen_i:
372 case NVPTX::BI__nvvm_atom_sys_and_gen_l:
373 case NVPTX::BI__nvvm_atom_sys_and_gen_ll:
374 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::And,
375 cir::SyncScopeKind::System);
376 case NVPTX::BI__nvvm_atom_cta_or_gen_i:
377 case NVPTX::BI__nvvm_atom_cta_or_gen_l:
378 case NVPTX::BI__nvvm_atom_cta_or_gen_ll:
379 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Or,
380 cir::SyncScopeKind::Workgroup);
381 case NVPTX::BI__nvvm_atom_sys_or_gen_i:
382 case NVPTX::BI__nvvm_atom_sys_or_gen_l:
383 case NVPTX::BI__nvvm_atom_sys_or_gen_ll:
384 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Or,
385 cir::SyncScopeKind::System);
386 case NVPTX::BI__nvvm_atom_cta_xor_gen_i:
387 case NVPTX::BI__nvvm_atom_cta_xor_gen_l:
388 case NVPTX::BI__nvvm_atom_cta_xor_gen_ll:
389 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Xor,
390 cir::SyncScopeKind::Workgroup);
391 case NVPTX::BI__nvvm_atom_sys_xor_gen_i:
392 case NVPTX::BI__nvvm_atom_sys_xor_gen_l:
393 case NVPTX::BI__nvvm_atom_sys_xor_gen_ll:
394 return makeScopedAtomicRMW(*this, expr, cir::AtomicFetchKind::Xor,
395 cir::SyncScopeKind::System);
396 case NVPTX::BI__nvvm_atom_cta_cas_gen_us:
397 case NVPTX::BI__nvvm_atom_cta_cas_gen_i:
398 case NVPTX::BI__nvvm_atom_cta_cas_gen_l:
399 case NVPTX::BI__nvvm_atom_cta_cas_gen_ll:
400 return emitAtomicCmpXchg(expr, /*returnBool=*/false, cir::MemOrder::Relaxed,
401 cir::MemOrder::Relaxed,
402 cir::SyncScopeKind::Workgroup);
403 case NVPTX::BI__nvvm_atom_sys_cas_gen_us:
404 case NVPTX::BI__nvvm_atom_sys_cas_gen_i:
405 case NVPTX::BI__nvvm_atom_sys_cas_gen_l:
406 case NVPTX::BI__nvvm_atom_sys_cas_gen_ll:
407 return emitAtomicCmpXchg(expr, /*returnBool=*/false, cir::MemOrder::Relaxed,
408 cir::MemOrder::Relaxed,
409 cir::SyncScopeKind::System);
410 case NVPTX::BI__nvvm_match_all_sync_i32p:
411 case NVPTX::BI__nvvm_match_all_sync_i64p:
412 cgm.errorNYI(expr->getSourceRange(),
413 std::string("unimplemented NVPTX builtin call: ") +
414 getContext().BuiltinInfo.getName(builtinId));
415 return mlir::Value{};
416 // FP MMA loads
417 case NVPTX::BI__hmma_m16n16k16_ld_a:
418 case NVPTX::BI__hmma_m16n16k16_ld_b:
419 case NVPTX::BI__hmma_m16n16k16_ld_c_f16:
420 case NVPTX::BI__hmma_m16n16k16_ld_c_f32:
421 case NVPTX::BI__hmma_m32n8k16_ld_a:
422 case NVPTX::BI__hmma_m32n8k16_ld_b:
423 case NVPTX::BI__hmma_m32n8k16_ld_c_f16:
424 case NVPTX::BI__hmma_m32n8k16_ld_c_f32:
425 case NVPTX::BI__hmma_m8n32k16_ld_a:
426 case NVPTX::BI__hmma_m8n32k16_ld_b:
427 case NVPTX::BI__hmma_m8n32k16_ld_c_f16:
428 case NVPTX::BI__hmma_m8n32k16_ld_c_f32:
429 cgm.errorNYI(expr->getSourceRange(),
430 std::string("unimplemented NVPTX builtin call: ") +
431 getContext().BuiltinInfo.getName(builtinId));
432 return mlir::Value{};
433 case NVPTX::BI__imma_m16n16k16_ld_a_s8:
434 case NVPTX::BI__imma_m16n16k16_ld_a_u8:
435 case NVPTX::BI__imma_m16n16k16_ld_b_s8:
436 case NVPTX::BI__imma_m16n16k16_ld_b_u8:
437 case NVPTX::BI__imma_m16n16k16_ld_c:
438 case NVPTX::BI__imma_m32n8k16_ld_a_s8:
439 case NVPTX::BI__imma_m32n8k16_ld_a_u8:
440 case NVPTX::BI__imma_m32n8k16_ld_b_s8:
441 case NVPTX::BI__imma_m32n8k16_ld_b_u8:
442 case NVPTX::BI__imma_m32n8k16_ld_c:
443 case NVPTX::BI__imma_m8n32k16_ld_a_s8:
444 case NVPTX::BI__imma_m8n32k16_ld_a_u8:
445 case NVPTX::BI__imma_m8n32k16_ld_b_s8:
446 case NVPTX::BI__imma_m8n32k16_ld_b_u8:
447 case NVPTX::BI__imma_m8n32k16_ld_c:
448 cgm.errorNYI(expr->getSourceRange(),
449 std::string("unimplemented NVPTX builtin call: ") +
450 getContext().BuiltinInfo.getName(builtinId));
451 return mlir::Value{};
452 case NVPTX::BI__imma_m8n8k32_ld_a_s4:
453 case NVPTX::BI__imma_m8n8k32_ld_a_u4:
454 case NVPTX::BI__imma_m8n8k32_ld_b_s4:
455 case NVPTX::BI__imma_m8n8k32_ld_b_u4:
456 case NVPTX::BI__imma_m8n8k32_ld_c:
457 case NVPTX::BI__bmma_m8n8k128_ld_a_b1:
458 case NVPTX::BI__bmma_m8n8k128_ld_b_b1:
459 case NVPTX::BI__bmma_m8n8k128_ld_c:
460 cgm.errorNYI(expr->getSourceRange(),
461 std::string("unimplemented NVPTX builtin call: ") +
462 getContext().BuiltinInfo.getName(builtinId));
463 return mlir::Value{};
464 case NVPTX::BI__dmma_m8n8k4_ld_a:
465 case NVPTX::BI__dmma_m8n8k4_ld_b:
466 case NVPTX::BI__dmma_m8n8k4_ld_c:
467 cgm.errorNYI(expr->getSourceRange(),
468 std::string("unimplemented NVPTX builtin call: ") +
469 getContext().BuiltinInfo.getName(builtinId));
470 return mlir::Value{};
471 case NVPTX::BI__mma_bf16_m16n16k16_ld_a:
472 case NVPTX::BI__mma_bf16_m16n16k16_ld_b:
473 case NVPTX::BI__mma_bf16_m8n32k16_ld_a:
474 case NVPTX::BI__mma_bf16_m8n32k16_ld_b:
475 case NVPTX::BI__mma_bf16_m32n8k16_ld_a:
476 case NVPTX::BI__mma_bf16_m32n8k16_ld_b:
477 case NVPTX::BI__mma_tf32_m16n16k8_ld_a:
478 case NVPTX::BI__mma_tf32_m16n16k8_ld_b:
479 case NVPTX::BI__mma_tf32_m16n16k8_ld_c:
480 cgm.errorNYI(expr->getSourceRange(),
481 std::string("unimplemented NVPTX builtin call: ") +
482 getContext().BuiltinInfo.getName(builtinId));
483 return mlir::Value{};
484 case NVPTX::BI__hmma_m16n16k16_st_c_f16:
485 case NVPTX::BI__hmma_m16n16k16_st_c_f32:
486 case NVPTX::BI__hmma_m32n8k16_st_c_f16:
487 case NVPTX::BI__hmma_m32n8k16_st_c_f32:
488 case NVPTX::BI__hmma_m8n32k16_st_c_f16:
489 case NVPTX::BI__hmma_m8n32k16_st_c_f32:
490 case NVPTX::BI__imma_m16n16k16_st_c_i32:
491 case NVPTX::BI__imma_m32n8k16_st_c_i32:
492 case NVPTX::BI__imma_m8n32k16_st_c_i32:
493 case NVPTX::BI__imma_m8n8k32_st_c_i32:
494 case NVPTX::BI__bmma_m8n8k128_st_c_i32:
495 case NVPTX::BI__dmma_m8n8k4_st_c_f64:
496 case NVPTX::BI__mma_m16n16k8_st_c_f32:
497 cgm.errorNYI(expr->getSourceRange(),
498 std::string("unimplemented NVPTX builtin call: ") +
499 getContext().BuiltinInfo.getName(builtinId));
500 return mlir::Value{};
501 // BI__hmma_m16n16k16_mma_<Dtype><CType>(d, a, b, c, layout, satf) -->
502 // Intrinsic::nvvm_wmma_m16n16k16_mma_sync<layout A,B><DType><CType><Satf>
503 case NVPTX::BI__hmma_m16n16k16_mma_f16f16:
504 case NVPTX::BI__hmma_m16n16k16_mma_f32f16:
505 case NVPTX::BI__hmma_m16n16k16_mma_f32f32:
506 case NVPTX::BI__hmma_m16n16k16_mma_f16f32:
507 case NVPTX::BI__hmma_m32n8k16_mma_f16f16:
508 case NVPTX::BI__hmma_m32n8k16_mma_f32f16:
509 case NVPTX::BI__hmma_m32n8k16_mma_f32f32:
510 case NVPTX::BI__hmma_m32n8k16_mma_f16f32:
511 case NVPTX::BI__hmma_m8n32k16_mma_f16f16:
512 case NVPTX::BI__hmma_m8n32k16_mma_f32f16:
513 case NVPTX::BI__hmma_m8n32k16_mma_f32f32:
514 case NVPTX::BI__hmma_m8n32k16_mma_f16f32:
515 case NVPTX::BI__imma_m16n16k16_mma_s8:
516 case NVPTX::BI__imma_m16n16k16_mma_u8:
517 case NVPTX::BI__imma_m32n8k16_mma_s8:
518 case NVPTX::BI__imma_m32n8k16_mma_u8:
519 case NVPTX::BI__imma_m8n32k16_mma_s8:
520 case NVPTX::BI__imma_m8n32k16_mma_u8:
521 case NVPTX::BI__imma_m8n8k32_mma_s4:
522 case NVPTX::BI__imma_m8n8k32_mma_u4:
523 case NVPTX::BI__bmma_m8n8k128_mma_xor_popc_b1:
524 case NVPTX::BI__bmma_m8n8k128_mma_and_popc_b1:
525 case NVPTX::BI__dmma_m8n8k4_mma_f64:
526 case NVPTX::BI__mma_bf16_m16n16k16_mma_f32:
527 case NVPTX::BI__mma_bf16_m8n32k16_mma_f32:
528 case NVPTX::BI__mma_bf16_m32n8k16_mma_f32:
529 case NVPTX::BI__mma_tf32_m16n16k8_mma_f32:
530 cgm.errorNYI(expr->getSourceRange(),
531 std::string("unimplemented NVPTX builtin call: ") +
532 getContext().BuiltinInfo.getName(builtinId));
533 return mlir::Value{};
534 // The following builtins require half type support
535 case NVPTX::BI__nvvm_ff2f16x2_rn:
536 cgm.errorNYI(expr->getSourceRange(),
537 std::string("unimplemented NVPTX builtin call: ") +
538 getContext().BuiltinInfo.getName(builtinId));
539 return mlir::Value{};
540 case NVPTX::BI__nvvm_ff2f16x2_rn_relu:
541 cgm.errorNYI(expr->getSourceRange(),
542 std::string("unimplemented NVPTX builtin call: ") +
543 getContext().BuiltinInfo.getName(builtinId));
544 return mlir::Value{};
545 case NVPTX::BI__nvvm_ff2f16x2_rz:
546 cgm.errorNYI(expr->getSourceRange(),
547 std::string("unimplemented NVPTX builtin call: ") +
548 getContext().BuiltinInfo.getName(builtinId));
549 return mlir::Value{};
550 case NVPTX::BI__nvvm_ff2f16x2_rz_relu:
551 cgm.errorNYI(expr->getSourceRange(),
552 std::string("unimplemented NVPTX builtin call: ") +
553 getContext().BuiltinInfo.getName(builtinId));
554 return mlir::Value{};
555 case NVPTX::BI__nvvm_fma_rn_f16:
556 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.f16")
557 .getValue();
558 case NVPTX::BI__nvvm_fma_rn_f16x2:
559 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.f16x2")
560 .getValue();
561 case NVPTX::BI__nvvm_fma_rn_ftz_f16:
562 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.ftz.f16")
563 .getValue();
564 case NVPTX::BI__nvvm_fma_rn_ftz_f16x2:
565 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.ftz.f16x2")
566 .getValue();
567 case NVPTX::BI__nvvm_fma_rn_ftz_relu_f16:
568 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.ftz.relu.f16")
569 .getValue();
570 case NVPTX::BI__nvvm_fma_rn_ftz_relu_f16x2:
572 "nvvm.fma.rn.ftz.relu.f16x2")
573 .getValue();
574 case NVPTX::BI__nvvm_fma_rn_ftz_sat_f16:
575 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.ftz.sat.f16")
576 .getValue();
577 case NVPTX::BI__nvvm_fma_rn_ftz_sat_f16x2:
579 "nvvm.fma.rn.ftz.sat.f16x2")
580 .getValue();
581 case NVPTX::BI__nvvm_fma_rn_relu_f16:
582 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.relu.f16")
583 .getValue();
584 case NVPTX::BI__nvvm_fma_rn_relu_f16x2:
585 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.relu.f16x2")
586 .getValue();
587 case NVPTX::BI__nvvm_fma_rn_sat_f16:
588 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.sat.f16")
589 .getValue();
590 case NVPTX::BI__nvvm_fma_rn_sat_f16x2:
591 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.sat.f16x2")
592 .getValue();
593 case NVPTX::BI__nvvm_fma_rn_oob_f16:
594 case NVPTX::BI__nvvm_fma_rn_oob_f16x2:
595 case NVPTX::BI__nvvm_fma_rn_oob_bf16:
596 case NVPTX::BI__nvvm_fma_rn_oob_bf16x2:
597 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.oob")
598 .getValue();
599 case NVPTX::BI__nvvm_fma_rn_oob_relu_f16:
600 case NVPTX::BI__nvvm_fma_rn_oob_relu_f16x2:
601 case NVPTX::BI__nvvm_fma_rn_oob_relu_bf16:
602 case NVPTX::BI__nvvm_fma_rn_oob_relu_bf16x2:
603 return emitBuiltinWithOneOverloadedType<3>(expr, "nvvm.fma.rn.oob.relu")
604 .getValue();
605 case NVPTX::BI__nvvm_fmax_f16:
606 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.f16")
607 .getValue();
608 case NVPTX::BI__nvvm_fmax_f16x2:
609 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.f16x2")
610 .getValue();
611 case NVPTX::BI__nvvm_fmax_ftz_f16:
612 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.ftz.f16")
613 .getValue();
614 case NVPTX::BI__nvvm_fmax_ftz_f16x2:
615 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.ftz.f16x2")
616 .getValue();
617 case NVPTX::BI__nvvm_fmax_ftz_nan_f16:
618 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.ftz.nan.f16")
619 .getValue();
620 case NVPTX::BI__nvvm_fmax_ftz_nan_f16x2:
621 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.ftz.nan.f16x2")
622 .getValue();
623 case NVPTX::BI__nvvm_fmax_ftz_nan_xorsign_abs_f16:
625 expr, "nvvm.fmax.ftz.nan.xorsign.abs.f16")
626 .getValue();
627 case NVPTX::BI__nvvm_fmax_ftz_nan_xorsign_abs_f16x2:
629 expr, "nvvm.fmax.ftz.nan.xorsign.abs.f16x2")
630 .getValue();
631 case NVPTX::BI__nvvm_fmax_ftz_xorsign_abs_f16:
633 "nvvm.fmax.ftz.xorsign.abs.f16")
634 .getValue();
635 case NVPTX::BI__nvvm_fmax_ftz_xorsign_abs_f16x2:
637 expr, "nvvm.fmax.ftz.xorsign.abs.f16x2")
638 .getValue();
639 case NVPTX::BI__nvvm_fmax_nan_f16:
640 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.nan.f16")
641 .getValue();
642 case NVPTX::BI__nvvm_fmax_nan_f16x2:
643 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmax.nan.f16x2")
644 .getValue();
645 case NVPTX::BI__nvvm_fmax_nan_xorsign_abs_f16:
647 "nvvm.fmax.nan.xorsign.abs.f16")
648 .getValue();
649 case NVPTX::BI__nvvm_fmax_nan_xorsign_abs_f16x2:
651 expr, "nvvm.fmax.nan.xorsign.abs.f16x2")
652 .getValue();
653 case NVPTX::BI__nvvm_fmax_xorsign_abs_f16:
655 "nvvm.fmax.xorsign.abs.f16")
656 .getValue();
657 case NVPTX::BI__nvvm_fmax_xorsign_abs_f16x2:
659 "nvvm.fmax.xorsign.abs.f16x2")
660 .getValue();
661 case NVPTX::BI__nvvm_fmin_f16:
662 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.f16")
663 .getValue();
664 case NVPTX::BI__nvvm_fmin_f16x2:
665 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.f16x2")
666 .getValue();
667 case NVPTX::BI__nvvm_fmin_ftz_f16:
668 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.ftz.f16")
669 .getValue();
670 case NVPTX::BI__nvvm_fmin_ftz_f16x2:
671 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.ftz.f16x2")
672 .getValue();
673 case NVPTX::BI__nvvm_fmin_ftz_nan_f16:
674 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.ftz.nan.f16")
675 .getValue();
676 case NVPTX::BI__nvvm_fmin_ftz_nan_f16x2:
677 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.ftz.nan.f16x2")
678 .getValue();
679 case NVPTX::BI__nvvm_fmin_ftz_nan_xorsign_abs_f16:
681 expr, "nvvm.fmin.ftz.nan.xorsign.abs.f16")
682 .getValue();
683 case NVPTX::BI__nvvm_fmin_ftz_nan_xorsign_abs_f16x2:
685 expr, "nvvm.fmin.ftz.nan.xorsign.abs.f16x2")
686 .getValue();
687 case NVPTX::BI__nvvm_fmin_ftz_xorsign_abs_f16:
689 "nvvm.fmin.ftz.xorsign.abs.f16")
690 .getValue();
691 case NVPTX::BI__nvvm_fmin_ftz_xorsign_abs_f16x2:
693 expr, "nvvm.fmin.ftz.xorsign.abs.f16x2")
694 .getValue();
695 case NVPTX::BI__nvvm_fmin_nan_f16:
696 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.nan.f16")
697 .getValue();
698 case NVPTX::BI__nvvm_fmin_nan_f16x2:
699 return emitBuiltinWithOneOverloadedType<2>(expr, "nvvm.fmin.nan.f16x2")
700 .getValue();
701 case NVPTX::BI__nvvm_fmin_nan_xorsign_abs_f16:
703 "nvvm.fmin.nan.xorsign.abs.f16")
704 .getValue();
705 case NVPTX::BI__nvvm_fmin_nan_xorsign_abs_f16x2:
707 expr, "nvvm.fmin.nan.xorsign.abs.f16x2")
708 .getValue();
709 case NVPTX::BI__nvvm_fmin_xorsign_abs_f16:
711 "nvvm.fmin.xorsign.abs.f16")
712 .getValue();
713 case NVPTX::BI__nvvm_fmin_xorsign_abs_f16x2:
715 "nvvm.fmin.xorsign.abs.f16x2")
716 .getValue();
717 case NVPTX::BI__nvvm_fabs_f:
718 case NVPTX::BI__nvvm_abs_bf16:
719 case NVPTX::BI__nvvm_abs_bf16x2:
720 case NVPTX::BI__nvvm_fabs_f16:
721 case NVPTX::BI__nvvm_fabs_f16x2:
722 return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.fabs");
723 case NVPTX::BI__nvvm_fabs_ftz_f:
724 case NVPTX::BI__nvvm_fabs_ftz_f16:
725 case NVPTX::BI__nvvm_fabs_ftz_f16x2:
726 return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.fabs.ftz");
727 case NVPTX::BI__nvvm_fabs_d:
728 return emitUnaryNVVMIntrinsic(*this, expr, "fabs");
729 case NVPTX::BI__nvvm_ex2_approx_d:
730 case NVPTX::BI__nvvm_ex2_approx_f:
731 case NVPTX::BI__nvvm_ex2_approx_f16:
732 case NVPTX::BI__nvvm_ex2_approx_f16x2:
733 return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.ex2.approx");
734 case NVPTX::BI__nvvm_ex2_approx_ftz_f:
735 return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.ex2.approx.ftz");
736 case NVPTX::BI__nvvm_add_rn_f:
737 case NVPTX::BI__nvvm_add_rn_d:
738 return emitNVVMFAdd(*this, expr, "nvvm.fadd",
739 llvm::APFloat::rmNearestTiesToEven);
740 case NVPTX::BI__nvvm_add_rz_f:
741 case NVPTX::BI__nvvm_add_rz_d:
742 return emitNVVMFAdd(*this, expr, "nvvm.fadd", llvm::APFloat::rmTowardZero);
743 case NVPTX::BI__nvvm_add_rm_f:
744 case NVPTX::BI__nvvm_add_rm_d:
745 return emitNVVMFAdd(*this, expr, "nvvm.fadd",
746 llvm::APFloat::rmTowardNegative);
747 case NVPTX::BI__nvvm_add_rp_f:
748 case NVPTX::BI__nvvm_add_rp_d:
749 return emitNVVMFAdd(*this, expr, "nvvm.fadd",
750 llvm::APFloat::rmTowardPositive);
751 case NVPTX::BI__nvvm_add_rn_ftz_f:
752 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
753 llvm::APFloat::rmNearestTiesToEven);
754 case NVPTX::BI__nvvm_add_rz_ftz_f:
755 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
756 llvm::APFloat::rmTowardZero);
757 case NVPTX::BI__nvvm_add_rm_ftz_f:
758 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
759 llvm::APFloat::rmTowardNegative);
760 case NVPTX::BI__nvvm_add_rp_ftz_f:
761 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
762 llvm::APFloat::rmTowardPositive);
763 case NVPTX::BI__nvvm_add_rn_sat_f:
764 case NVPTX::BI__nvvm_add_rn_sat_f16:
765 case NVPTX::BI__nvvm_add_rn_sat_v2f16:
766 return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
767 llvm::APFloat::rmNearestTiesToEven);
768 case NVPTX::BI__nvvm_add_rz_sat_f:
769 return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
770 llvm::APFloat::rmTowardZero);
771 case NVPTX::BI__nvvm_add_rm_sat_f:
772 return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
773 llvm::APFloat::rmTowardNegative);
774 case NVPTX::BI__nvvm_add_rp_sat_f:
775 return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
776 llvm::APFloat::rmTowardPositive);
777 case NVPTX::BI__nvvm_add_rn_ftz_sat_f:
778 case NVPTX::BI__nvvm_add_rn_ftz_sat_f16:
779 case NVPTX::BI__nvvm_add_rn_ftz_sat_v2f16:
780 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
781 llvm::APFloat::rmNearestTiesToEven);
782 case NVPTX::BI__nvvm_add_rz_ftz_sat_f:
783 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
784 llvm::APFloat::rmTowardZero);
785 case NVPTX::BI__nvvm_add_rm_ftz_sat_f:
786 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
787 llvm::APFloat::rmTowardNegative);
788 case NVPTX::BI__nvvm_add_rp_ftz_sat_f:
789 return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
790 llvm::APFloat::rmTowardPositive);
791 case NVPTX::BI__nvvm_ldg_h:
792 case NVPTX::BI__nvvm_ldg_h2:
793 cgm.errorNYI(expr->getSourceRange(),
794 std::string("unimplemented NVPTX builtin call: ") +
795 getContext().BuiltinInfo.getName(builtinId));
796 return mlir::Value{};
797 case NVPTX::BI__nvvm_ldu_h:
798 case NVPTX::BI__nvvm_ldu_h2:
799 return makeLdu(*this, expr, "nvvm.ldu.global.f");
800 case NVPTX::BI__nvvm_cp_async_ca_shared_global_4:
801 cgm.errorNYI(expr->getSourceRange(),
802 std::string("unimplemented NVPTX builtin call: ") +
803 getContext().BuiltinInfo.getName(builtinId));
804 return mlir::Value{};
805 case NVPTX::BI__nvvm_cp_async_ca_shared_global_8:
806 cgm.errorNYI(expr->getSourceRange(),
807 std::string("unimplemented NVPTX builtin call: ") +
808 getContext().BuiltinInfo.getName(builtinId));
809 return mlir::Value{};
810 case NVPTX::BI__nvvm_cp_async_ca_shared_global_16:
811 cgm.errorNYI(expr->getSourceRange(),
812 std::string("unimplemented NVPTX builtin call: ") +
813 getContext().BuiltinInfo.getName(builtinId));
814 return mlir::Value{};
815 case NVPTX::BI__nvvm_cp_async_cg_shared_global_16:
816 cgm.errorNYI(expr->getSourceRange(),
817 std::string("unimplemented NVPTX builtin call: ") +
818 getContext().BuiltinInfo.getName(builtinId));
819 return mlir::Value{};
820 case NVPTX::BI__nvvm_read_ptx_sreg_clusterid_x:
821 cgm.errorNYI(expr->getSourceRange(),
822 std::string("unimplemented NVPTX builtin call: ") +
823 getContext().BuiltinInfo.getName(builtinId));
824 return mlir::Value{};
825 case NVPTX::BI__nvvm_read_ptx_sreg_clusterid_y:
826 cgm.errorNYI(expr->getSourceRange(),
827 std::string("unimplemented NVPTX builtin call: ") +
828 getContext().BuiltinInfo.getName(builtinId));
829 return mlir::Value{};
830 case NVPTX::BI__nvvm_read_ptx_sreg_clusterid_z:
831 cgm.errorNYI(expr->getSourceRange(),
832 std::string("unimplemented NVPTX builtin call: ") +
833 getContext().BuiltinInfo.getName(builtinId));
834 return mlir::Value{};
835 case NVPTX::BI__nvvm_read_ptx_sreg_clusterid_w:
836 cgm.errorNYI(expr->getSourceRange(),
837 std::string("unimplemented NVPTX builtin call: ") +
838 getContext().BuiltinInfo.getName(builtinId));
839 return mlir::Value{};
840 case NVPTX::BI__nvvm_read_ptx_sreg_nclusterid_x:
841 cgm.errorNYI(expr->getSourceRange(),
842 std::string("unimplemented NVPTX builtin call: ") +
843 getContext().BuiltinInfo.getName(builtinId));
844 return mlir::Value{};
845 case NVPTX::BI__nvvm_read_ptx_sreg_nclusterid_y:
846 cgm.errorNYI(expr->getSourceRange(),
847 std::string("unimplemented NVPTX builtin call: ") +
848 getContext().BuiltinInfo.getName(builtinId));
849 return mlir::Value{};
850 case NVPTX::BI__nvvm_read_ptx_sreg_nclusterid_z:
851 cgm.errorNYI(expr->getSourceRange(),
852 std::string("unimplemented NVPTX builtin call: ") +
853 getContext().BuiltinInfo.getName(builtinId));
854 return mlir::Value{};
855 case NVPTX::BI__nvvm_read_ptx_sreg_nclusterid_w:
856 cgm.errorNYI(expr->getSourceRange(),
857 std::string("unimplemented NVPTX builtin call: ") +
858 getContext().BuiltinInfo.getName(builtinId));
859 return mlir::Value{};
860 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_ctaid_x:
861 cgm.errorNYI(expr->getSourceRange(),
862 std::string("unimplemented NVPTX builtin call: ") +
863 getContext().BuiltinInfo.getName(builtinId));
864 return mlir::Value{};
865 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_ctaid_y:
866 cgm.errorNYI(expr->getSourceRange(),
867 std::string("unimplemented NVPTX builtin call: ") +
868 getContext().BuiltinInfo.getName(builtinId));
869 return mlir::Value{};
870 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_ctaid_z:
871 cgm.errorNYI(expr->getSourceRange(),
872 std::string("unimplemented NVPTX builtin call: ") +
873 getContext().BuiltinInfo.getName(builtinId));
874 return mlir::Value{};
875 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_ctaid_w:
876 cgm.errorNYI(expr->getSourceRange(),
877 std::string("unimplemented NVPTX builtin call: ") +
878 getContext().BuiltinInfo.getName(builtinId));
879 return mlir::Value{};
880 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_nctaid_x:
881 cgm.errorNYI(expr->getSourceRange(),
882 std::string("unimplemented NVPTX builtin call: ") +
883 getContext().BuiltinInfo.getName(builtinId));
884 return mlir::Value{};
885 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_nctaid_y:
886 cgm.errorNYI(expr->getSourceRange(),
887 std::string("unimplemented NVPTX builtin call: ") +
888 getContext().BuiltinInfo.getName(builtinId));
889 return mlir::Value{};
890 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_nctaid_z:
891 cgm.errorNYI(expr->getSourceRange(),
892 std::string("unimplemented NVPTX builtin call: ") +
893 getContext().BuiltinInfo.getName(builtinId));
894 return mlir::Value{};
895 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_nctaid_w:
896 cgm.errorNYI(expr->getSourceRange(),
897 std::string("unimplemented NVPTX builtin call: ") +
898 getContext().BuiltinInfo.getName(builtinId));
899 return mlir::Value{};
900 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_ctarank:
901 cgm.errorNYI(expr->getSourceRange(),
902 std::string("unimplemented NVPTX builtin call: ") +
903 getContext().BuiltinInfo.getName(builtinId));
904 return mlir::Value{};
905 case NVPTX::BI__nvvm_read_ptx_sreg_cluster_nctarank:
906 cgm.errorNYI(expr->getSourceRange(),
907 std::string("unimplemented NVPTX builtin call: ") +
908 getContext().BuiltinInfo.getName(builtinId));
909 return mlir::Value{};
910 case NVPTX::BI__nvvm_is_explicit_cluster:
911 cgm.errorNYI(expr->getSourceRange(),
912 std::string("unimplemented NVPTX builtin call: ") +
913 getContext().BuiltinInfo.getName(builtinId));
914 return mlir::Value{};
915 case NVPTX::BI__nvvm_isspacep_shared_cluster:
916 cgm.errorNYI(expr->getSourceRange(),
917 std::string("unimplemented NVPTX builtin call: ") +
918 getContext().BuiltinInfo.getName(builtinId));
919 return mlir::Value{};
920 case NVPTX::BI__nvvm_mapa:
921 cgm.errorNYI(expr->getSourceRange(),
922 std::string("unimplemented NVPTX builtin call: ") +
923 getContext().BuiltinInfo.getName(builtinId));
924 return mlir::Value{};
925 case NVPTX::BI__nvvm_mapa_shared_cluster:
926 cgm.errorNYI(expr->getSourceRange(),
927 std::string("unimplemented NVPTX builtin call: ") +
928 getContext().BuiltinInfo.getName(builtinId));
929 return mlir::Value{};
930 case NVPTX::BI__nvvm_getctarank:
931 cgm.errorNYI(expr->getSourceRange(),
932 std::string("unimplemented NVPTX builtin call: ") +
933 getContext().BuiltinInfo.getName(builtinId));
934 return mlir::Value{};
935 case NVPTX::BI__nvvm_getctarank_shared_cluster:
936 cgm.errorNYI(expr->getSourceRange(),
937 std::string("unimplemented NVPTX builtin call: ") +
938 getContext().BuiltinInfo.getName(builtinId));
939 return mlir::Value{};
940 case NVPTX::BI__nvvm_barrier_cluster_arrive:
941 return builder.emitIntrinsicCallOp(getLoc(expr->getExprLoc()),
942 "nvvm.barrier.cluster.arrive",
943 builder.getVoidTy());
944 case NVPTX::BI__nvvm_barrier_cluster_arrive_relaxed:
945 return builder.emitIntrinsicCallOp(getLoc(expr->getExprLoc()),
946 "nvvm.barrier.cluster.arrive.relaxed",
947 builder.getVoidTy());
948 case NVPTX::BI__nvvm_barrier_cluster_wait:
949 return builder.emitIntrinsicCallOp(getLoc(expr->getExprLoc()),
950 "nvvm.barrier.cluster.wait",
951 builder.getVoidTy());
952 case NVPTX::BI__nvvm_fence_sc_cluster:
953 return builder.emitIntrinsicCallOp(getLoc(expr->getExprLoc()),
954 "nvvm.fence.sc.cluster",
955 builder.getVoidTy());
956 case NVPTX::BI__nvvm_bar_sync:
957 return builder.emitIntrinsicCallOp(
958 getLoc(expr->getExprLoc()), "nvvm.barrier.cta.sync.aligned.all",
959 builder.getVoidTy(), mlir::ValueRange{emitScalarExpr(expr->getArg(0))});
960 case NVPTX::BI__syncthreads:
961 return builder.emitIntrinsicCallOp(
962 getLoc(expr->getExprLoc()), "nvvm.barrier.cta.sync.aligned.all",
963 builder.getVoidTy(),
964 mlir::ValueRange{builder.getConstInt(getLoc(expr->getExprLoc()),
965 builder.getSInt32Ty(), 0)});
966 case NVPTX::BI__nvvm_barrier_sync:
967 return builder.emitIntrinsicCallOp(
968 getLoc(expr->getExprLoc()), "nvvm.barrier.cta.sync.all",
969 builder.getVoidTy(), mlir::ValueRange{emitScalarExpr(expr->getArg(0))});
970 case NVPTX::BI__nvvm_barrier_sync_cnt:
971 return builder.emitIntrinsicCallOp(
972 getLoc(expr->getExprLoc()), "nvvm.barrier.cta.sync.count",
973 builder.getVoidTy(),
974 mlir::ValueRange{emitScalarExpr(expr->getArg(0)),
975 emitScalarExpr(expr->getArg(1))});
976 case NVPTX::BI__nvvm_bar0_and:
977 return emitBar0Reduction(*this, expr,
978 "nvvm.barrier.cta.red.and.aligned.all",
979 /*returnsPred=*/true);
980 case NVPTX::BI__nvvm_bar0_or:
981 return emitBar0Reduction(*this, expr, "nvvm.barrier.cta.red.or.aligned.all",
982 /*returnsPred=*/true);
983 case NVPTX::BI__nvvm_bar0_popc:
984 return emitBar0Reduction(*this, expr,
985 "nvvm.barrier.cta.red.popc.aligned.all",
986 /*returnsPred=*/false);
987
988 default:
989 return std::nullopt;
990 }
991}
992
993// vprintf takes two args: A format string, and a pointer to a buffer containing
994// the varargs.
995//
996// For example, the call
997//
998// printf("format string", arg1, arg2, arg3);
999//
1000// is converted into something resembling
1001//
1002// struct Tmp {
1003// Arg1 a1;
1004// Arg2 a2;
1005// Arg3 a3;
1006// };
1007// char* buf = alloca(sizeof(Tmp));
1008// *(Tmp*)buf = {a1, a2, a3};
1009// vprintf("format string", buf);
1010//
1011// `buf` is aligned to the max of {alignof(Arg1), ...}. Furthermore, each of
1012// the args is itself aligned to its preferred alignment.
1013//
1014// Note that by the time this function runs, the arguments have already
1015// undergone the standard C vararg promotion (short -> int, float -> double
1016// etc). In this function we pack the arguments into the buffer described above.
1018 const CallArgList &args,
1019 mlir::Location loc) {
1020 const cir::CIRDataLayout dataLayout = cgf.cgm.getDataLayout();
1021 CIRGenBuilderTy &builder = cgf.getBuilder();
1022
1023 if (args.size() <= 1)
1024 // If there are no arguments other than the format string,
1025 // pass a nullptr to vprintf.
1026 return builder.getNullPtr(builder.getVoidPtrTy(), loc);
1027
1029 for (const auto &arg : llvm::drop_begin(args))
1030 argTypes.push_back(arg.getKnownRValue().getValue().getType());
1031
1032 // We can directly store the arguments into a struct, and the alignment
1033 // would automatically be correct. That's because vprintf does not
1034 // accept aggregates.
1035 mlir::Type allocaTy = builder.getAnonRecordTy(
1036 argTypes, /*packed=*/false, cir::RecordType::getAllDataKinds(argTypes));
1037 auto allocaAlign = clang::CharUnits::fromQuantity(
1038 dataLayout.getABITypeAlign(allocaTy).value());
1039 Address allocaAddr =
1040 cgf.createTempAlloca(allocaTy, allocaAlign, loc, "printf_args");
1041 mlir::Value alloca = allocaAddr.getPointer();
1042
1043 for (auto [i, arg] : llvm::enumerate(llvm::drop_begin(args))) {
1044 mlir::Value member = builder.createGetMember(
1045 loc, cir::PointerType::get(argTypes[i]), alloca, /*name=*/"",
1046 /*index=*/i);
1047 auto abiAlign = clang::CharUnits::fromQuantity(
1048 dataLayout.getABITypeAlign(argTypes[i]).value());
1049 cir::StoreOp::create(builder, loc, arg.getKnownRValue().getValue(), member,
1050 /*is_volatile=*/false,
1051 /*isNontemporal=*/false,
1052 builder.getAlignmentAttr(abiAlign),
1053 /*sync_scope=*/cir::SyncScopeKindAttr{},
1054 /*mem_order=*/cir::MemOrderAttr{});
1055 }
1056
1057 return builder.createBitcast(alloca, builder.getVoidPtrTy());
1058}
1059
1060mlir::Value
1062 assert(cgm.getTriple().isNVPTX());
1063 assert(expr->getBuiltinCallee() == Builtin::BIprintf ||
1064 expr->getBuiltinCallee() == Builtin::BI__builtin_printf);
1065 assert(expr->getNumArgs() >= 1); // printf always has at least one arg.
1066 CallArgList args;
1067 emitCallArgs(args,
1068 expr->getDirectCallee()->getType()->getAs<FunctionProtoType>(),
1069 expr->arguments(), expr->getDirectCallee());
1070
1071 mlir::Location loc = getLoc(expr->getBeginLoc());
1072
1073 // We don't know how to emit non-scalar varargs.
1074 bool hasNonScalar =
1075 llvm::any_of(llvm::drop_begin(args), [&](const CallArg &a) {
1076 return a.hasLValue() || !a.getKnownRValue().isScalar();
1077 });
1078 if (hasNonScalar) {
1079 cgm.errorUnsupported(expr, "non-scalar args to printf");
1080 return builder.getConstInt(loc, builder.getSInt32Ty(), 0);
1081 }
1082
1083 mlir::Value packedData = packArgsIntoNVPTXFormatBuffer(*this, args, loc);
1084
1085 // int vprintf(char *format, void *packedData);
1086 auto vprintf = cgm.createRuntimeFunction(
1087 cir::FuncType::get(
1088 {cir::PointerType::get(builder.getSInt8Ty()), builder.getVoidPtrTy()},
1089 builder.getSInt32Ty()),
1090 "vprintf");
1091 auto formatString = args[0].getKnownRValue().getValue();
1092 return builder
1093 .createCallOp(loc, vprintf, mlir::ValueRange{formatString, packedData})
1094 .getResult();
1095}
static mlir::Value makeScopedAtomicRMW(CIRGenFunction &cgf, const CallExpr *expr, cir::AtomicFetchKind kind, cir::SyncScopeKind scope)
static mlir::Value packArgsIntoNVPTXFormatBuffer(CIRGenFunction &cgf, const CallArgList &args, mlir::Location loc)
static mlir::Value makeLdg(CIRGenFunction &cgf, const CallExpr *expr)
static mlir::Value emitBar0Reduction(CIRGenFunction &cgf, const CallExpr *expr, llvm::StringRef intrinsicName, bool returnsPred)
static mlir::Value emitNVVMFAdd(CIRGenFunction &cgf, const CallExpr *expr, llvm::StringRef intrinsicName, llvm::APFloat::roundingMode rm)
Emit a CIR LLVMIntrinsicCallOp for an NVVM fadd intrinsic, which takes the rounding mode as a trailin...
static mlir::Value makeScopedAtomicXchg(CIRGenFunction &cgf, const CallExpr *expr, cir::SyncScopeKind scope)
static mlir::Value makeLdu(CIRGenFunction &cgf, const CallExpr *expr, llvm::StringRef intrinsicName)
static mlir::Value emitUnaryNVVMIntrinsic(CIRGenFunction &cgf, const CallExpr *expr, llvm::StringRef intrinsicName)
Emit a CIR LLVMIntrinsicCallOp for a unary NVVM intrinsic.
*collection of selector each with an associated kind and an ordered *collection of selectors A selector has a kind
Enumerates target-specific builtins in their own namespaces within namespace clang.
cir::ConstantOp getNullValue(mlir::Type ty, mlir::Location loc)
mlir::Value createBoolToInt(mlir::Value src, mlir::Type newTy)
cir::ConstantOp getNullPtr(mlir::Type ty, mlir::Location loc)
mlir::Value createBitcast(mlir::Value src, mlir::Type newTy)
cir::CmpOp createCompare(mlir::Location loc, cir::CmpOpKind kind, mlir::Value lhs, mlir::Value rhs)
mlir::IntegerAttr getAlignmentAttr(clang::CharUnits alignment)
cir::PointerType getVoidPtrTy(clang::LangAS langAS=clang::LangAS::Default)
cir::BoolType getBoolTy()
llvm::Align getABITypeAlign(mlir::Type ty) const
static llvm::SmallVector< RecordMemberKind > getAllDataKinds(llvm::ArrayRef< mlir::Type > members)
One Data kind per member.
Definition CIRTypes.cpp:162
mlir::Value getPointer() const
Definition Address.h:98
clang::CharUnits getAlignment() const
Definition Address.h:138
mlir::Value emitRawPointer() const
Return the pointer contained in this class after authenticating it and adding offset to it if necessa...
Definition Address.h:112
mlir::Value emitIntrinsicCallOp(mlir::Location loc, const llvm::StringRef str, const mlir::Type &resTy, Operands &&...op)
Address createGetMember(mlir::Location loc, Address base, llvm::StringRef name, unsigned index)
cir::StructType getAnonRecordTy(llvm::ArrayRef< mlir::Type > members, bool packed, llvm::ArrayRef< cir::RecordMemberKind > memberKinds)
Get a CIR anonymous struct type.
void emitCallArgs(CallArgList &args, PrototypeWrapper prototype, llvm::iterator_range< clang::CallExpr::const_arg_iterator > argRange, AbstractCallee callee=AbstractCallee(), unsigned paramsToSkip=0)
Address emitPointerWithAlignment(const clang::Expr *expr, LValueBaseInfo *baseInfo=nullptr)
Given an expression with a pointer type, emit the value and compute our best estimate of the alignmen...
cir::AllocaOp createTempAlloca(mlir::Type ty, mlir::Location loc, const Twine &name="tmp", mlir::Value arraySize=nullptr, bool insertIntoFnEntryBlock=false)
This creates an alloca and inserts it into the entry block if ArraySize is nullptr,...
mlir::Value emitNVPTXDevicePrintfCallExpr(const CallExpr *expr)
Emit a device-side printf call for NVPTX targets.
mlir::Location getLoc(clang::SourceLocation srcLoc)
Helpers to convert Clang's SourceLocation to a MLIR Location.
mlir::Value makeBinaryAtomicValue(cir::AtomicFetchKind kind, const clang::CallExpr *expr, mlir::Type *originalArgType=nullptr, mlir::Value *emittedArgValue=nullptr, cir::MemOrder ordering=cir::MemOrder::SequentiallyConsistent)
Utility to insert an atomic instruction based on Intrinsic::ID and the expression node.
mlir::Type convertTypeForMem(QualType t)
RValue emitBuiltinWithOneOverloadedType(const CallExpr *e, llvm::StringRef intrinName, mlir::Type resultType={})
Emit a simple LLVM intrinsic that takes N scalar arguments.
mlir::Value emitScalarExpr(const clang::Expr *e, bool ignoreResultAssign=false)
Emit the computation of the specified expression of scalar type.
mlir::Value emitAtomicCmpXchg(const clang::CallExpr *expr, bool returnBool, cir::MemOrder successOrder=cir::MemOrder::SequentiallyConsistent, cir::MemOrder failureOrder=cir::MemOrder::SequentiallyConsistent, cir::SyncScopeKind scope=cir::SyncScopeKind::System)
Emit cir.atomic.cmpxchg.
CIRGenBuilderTy & getBuilder()
std::optional< mlir::Value > emitNVPTXBuiltinExpr(unsigned builtinID, const CallExpr *expr)
Emit a call to an NVPTX builtin function.
clang::ASTContext & getContext() const
const cir::CIRDataLayout getDataLayout() const
mlir::Value getValue() const
Return the value of this scalar value.
Definition CIRGenValue.h:57
bool isScalar() const
Definition CIRGenValue.h:49
CallExpr - Represents a function call (C99 6.5.2.2, C++ [expr.call]).
Definition Expr.h:2987
CharUnits - This is an opaque type for sizes expressed in character units.
Definition CharUnits.h:38
static CharUnits fromQuantity(QuantityType Quantity)
fromQuantity - Construct a CharUnits quantity from a raw integer type.
Definition CharUnits.h:63
Represents a prototype with parameter type info, e.g.
Definition TypeBase.h:5398
A (possibly-)qualified type.
Definition TypeBase.h:938
QualType getPointeeType() const
If this is a pointer, ObjC object pointer, or block pointer, this returns the respective pointee.
Definition Type.cpp:881
const internal::VariadicDynCastAllOfMatcher< Stmt, Expr > expr
Matches expressions.
Top level wrappers for InstallAPI frontend operations.
bool hasLValue() const
Definition CIRGenCall.h:216
RValue getKnownRValue() const
Definition CIRGenCall.h:227