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