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