1//===-- NVPTXInstPrinter.cpp - PTX assembly instruction printing ----------===//
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// Print MCInst instructions to .ptx format.
10//
11//===----------------------------------------------------------------------===//
12
13#include "MCTargetDesc/NVPTXInstPrinter.h"
14#include "MCTargetDesc/NVPTXBaseInfo.h"
15#include "NVPTX.h"
16#include "NVPTXUtilities.h"
17#include "llvm/ADT/StringRef.h"
18#include "llvm/IR/NVVMIntrinsicUtils.h"
19#include "llvm/MC/MCAsmInfo.h"
20#include "llvm/MC/MCExpr.h"
21#include "llvm/MC/MCInst.h"
22#include "llvm/MC/MCInstrInfo.h"
23#include "llvm/MC/MCSubtargetInfo.h"
24#include "llvm/MC/MCSymbol.h"
25#include "llvm/Support/ErrorHandling.h"
26#include "llvm/Support/FormatVariadic.h"
27#include <cctype>
28using namespace llvm;
29
30#define DEBUG_TYPE "asm-printer"
31
32#include "NVPTXGenAsmWriter.inc"
33
34static bool hasParamSubqualifiers(const MCSubtargetInfo &STI) {
35 return STI.hasFeature(Feature: NVPTX::PTX83);
36}
37
38NVPTXInstPrinter::NVPTXInstPrinter(const MCAsmInfo &MAI, const MCInstrInfo &MII,
39 const MCRegisterInfo &MRI)
40 : MCInstPrinter(MAI, MII, MRI) {}
41
42void NVPTXInstPrinter::printRegName(raw_ostream &OS, MCRegister Reg) {
43 // Decode a register packed by NVPTXAsmPrinter::encodeVirtualRegister.
44 const auto Kind = static_cast<NVPTX::VirtualRegisterKind>(
45 Reg.id() >> NVPTX::VirtualRegisterKindShift);
46
47 if (Kind == NVPTX::VirtualRegisterKind::Physical) {
48 // This is actually a physical register, so defer to the autogenerated
49 // register printer
50 OS << getRegisterName(Reg);
51 return;
52 }
53
54 OS << NVPTX::getVirtualRegisterPrefix(Kind)
55 << (Reg.id() & NVPTX::VirtualRegisterNumMask);
56}
57
58void NVPTXInstPrinter::printInst(const MCInst *MI, uint64_t Address,
59 StringRef Annot, const MCSubtargetInfo &STI,
60 raw_ostream &OS) {
61 printInstruction(MI, Address, STI, O&: OS);
62
63 // Next always print the annotation.
64 printAnnotation(OS, Annot);
65}
66
67void NVPTXInstPrinter::printOperand(const MCInst *MI, unsigned OpNo,
68 const MCSubtargetInfo &, raw_ostream &O) {
69 const MCOperand &Op = MI->getOperand(i: OpNo);
70 if (Op.isReg()) {
71 MCRegister Reg = Op.getReg();
72 printRegName(OS&: O, Reg);
73 } else if (Op.isImm()) {
74 markup(OS&: O, M: Markup::Immediate) << formatImm(Value: Op.getImm());
75 } else {
76 assert(Op.isExpr() && "Unknown operand kind in printOperand");
77 MAI.printExpr(O, *Op.getExpr());
78 }
79}
80
81void NVPTXInstPrinter::printCvtMode(const MCInst *MI, int OpNum,
82 const MCSubtargetInfo &, raw_ostream &O,
83 StringRef Modifier) {
84 const MCOperand &MO = MI->getOperand(i: OpNum);
85 int64_t Imm = MO.getImm();
86
87 if (Modifier == "ftz") {
88 // FTZ flag
89 if (Imm & NVPTX::PTXCvtMode::FTZ_FLAG)
90 O << ".ftz";
91 return;
92 } else if (Modifier == "sat") {
93 // SAT flag
94 if (Imm & NVPTX::PTXCvtMode::SAT_FLAG)
95 O << ".sat";
96 return;
97 } else if (Modifier == "satfinite") {
98 // SATFINITE flag
99 if (Imm & NVPTX::PTXCvtMode::SATFINITE_FLAG)
100 O << ".satfinite";
101 return;
102 } else if (Modifier == "relu") {
103 // RELU flag
104 if (Imm & NVPTX::PTXCvtMode::RELU_FLAG)
105 O << ".relu";
106 return;
107 } else if (Modifier == "base") {
108 // Default operand
109 switch (Imm & NVPTX::PTXCvtMode::BASE_MASK) {
110 default:
111 return;
112 case NVPTX::PTXCvtMode::NONE:
113 return;
114 case NVPTX::PTXCvtMode::RNI:
115 O << ".rni";
116 return;
117 case NVPTX::PTXCvtMode::RZI:
118 O << ".rzi";
119 return;
120 case NVPTX::PTXCvtMode::RMI:
121 O << ".rmi";
122 return;
123 case NVPTX::PTXCvtMode::RPI:
124 O << ".rpi";
125 return;
126 case NVPTX::PTXCvtMode::RN:
127 O << ".rn";
128 return;
129 case NVPTX::PTXCvtMode::RZ:
130 O << ".rz";
131 return;
132 case NVPTX::PTXCvtMode::RM:
133 O << ".rm";
134 return;
135 case NVPTX::PTXCvtMode::RP:
136 O << ".rp";
137 return;
138 case NVPTX::PTXCvtMode::RNA:
139 O << ".rna";
140 return;
141 case NVPTX::PTXCvtMode::RS:
142 O << ".rs";
143 return;
144 }
145 }
146 llvm_unreachable("Invalid conversion modifier");
147}
148
149void NVPTXInstPrinter::printFTZFlag(const MCInst *MI, int OpNum,
150 const MCSubtargetInfo &, raw_ostream &O) {
151 const MCOperand &MO = MI->getOperand(i: OpNum);
152 const int Imm = MO.getImm();
153 if (Imm)
154 O << ".ftz";
155}
156
157void NVPTXInstPrinter::printMultimem(const MCInst *MI, int OpNum,
158 const MCSubtargetInfo &, raw_ostream &O) {
159 const MCOperand &MO = MI->getOperand(i: OpNum);
160 if (MO.getImm())
161 O << "multimem.";
162}
163
164void NVPTXInstPrinter::printNegatedPredicate(const MCInst *MI, int OpNum,
165 const MCSubtargetInfo &,
166 raw_ostream &O) {
167 if (MI->getOperand(i: OpNum).getImm())
168 O << "!";
169}
170
171void NVPTXInstPrinter::printCmpMode(const MCInst *MI, int OpNum,
172 const MCSubtargetInfo &, raw_ostream &O,
173 StringRef Modifier) {
174 const MCOperand &MO = MI->getOperand(i: OpNum);
175 int64_t Imm = MO.getImm();
176
177 if (Modifier == "FCmp") {
178 switch (Imm) {
179 default:
180 return;
181 case NVPTX::PTXCmpMode::EQ:
182 O << "eq";
183 return;
184 case NVPTX::PTXCmpMode::NE:
185 O << "ne";
186 return;
187 case NVPTX::PTXCmpMode::LT:
188 O << "lt";
189 return;
190 case NVPTX::PTXCmpMode::LE:
191 O << "le";
192 return;
193 case NVPTX::PTXCmpMode::GT:
194 O << "gt";
195 return;
196 case NVPTX::PTXCmpMode::GE:
197 O << "ge";
198 return;
199 case NVPTX::PTXCmpMode::EQU:
200 O << "equ";
201 return;
202 case NVPTX::PTXCmpMode::NEU:
203 O << "neu";
204 return;
205 case NVPTX::PTXCmpMode::LTU:
206 O << "ltu";
207 return;
208 case NVPTX::PTXCmpMode::LEU:
209 O << "leu";
210 return;
211 case NVPTX::PTXCmpMode::GTU:
212 O << "gtu";
213 return;
214 case NVPTX::PTXCmpMode::GEU:
215 O << "geu";
216 return;
217 case NVPTX::PTXCmpMode::NUM:
218 O << "num";
219 return;
220 case NVPTX::PTXCmpMode::NotANumber:
221 O << "nan";
222 return;
223 }
224 }
225 if (Modifier == "ICmp") {
226 switch (Imm) {
227 default:
228 llvm_unreachable("Invalid ICmp mode");
229 case NVPTX::PTXCmpMode::EQ:
230 O << "eq";
231 return;
232 case NVPTX::PTXCmpMode::NE:
233 O << "ne";
234 return;
235 case NVPTX::PTXCmpMode::LT:
236 case NVPTX::PTXCmpMode::LTU:
237 O << "lt";
238 return;
239 case NVPTX::PTXCmpMode::LE:
240 case NVPTX::PTXCmpMode::LEU:
241 O << "le";
242 return;
243 case NVPTX::PTXCmpMode::GT:
244 case NVPTX::PTXCmpMode::GTU:
245 O << "gt";
246 return;
247 case NVPTX::PTXCmpMode::GE:
248 case NVPTX::PTXCmpMode::GEU:
249 O << "ge";
250 return;
251 }
252 }
253 if (Modifier == "IType") {
254 switch (Imm) {
255 default:
256 llvm_unreachable("Invalid IType");
257 case NVPTX::PTXCmpMode::EQ:
258 case NVPTX::PTXCmpMode::NE:
259 O << "b";
260 return;
261 case NVPTX::PTXCmpMode::LT:
262 case NVPTX::PTXCmpMode::LE:
263 case NVPTX::PTXCmpMode::GT:
264 case NVPTX::PTXCmpMode::GE:
265 O << "s";
266 return;
267 case NVPTX::PTXCmpMode::LTU:
268 case NVPTX::PTXCmpMode::LEU:
269 case NVPTX::PTXCmpMode::GTU:
270 case NVPTX::PTXCmpMode::GEU:
271 O << "u";
272 return;
273 }
274 }
275 llvm_unreachable("Empty Modifier");
276}
277
278void NVPTXInstPrinter::printAtomicCode(const MCInst *MI, int OpNum,
279 const MCSubtargetInfo &STI,
280 raw_ostream &O, StringRef Modifier) {
281 const MCOperand &MO = MI->getOperand(i: OpNum);
282 int Imm = (int)MO.getImm();
283 if (Modifier == "sem") {
284 auto Ordering = NVPTX::Ordering(Imm);
285 switch (Ordering) {
286 case NVPTX::Ordering::NotAtomic:
287 return;
288 case NVPTX::Ordering::Relaxed:
289 O << ".relaxed";
290 return;
291 case NVPTX::Ordering::Acquire:
292 O << ".acquire";
293 return;
294 case NVPTX::Ordering::Release:
295 O << ".release";
296 return;
297 case NVPTX::Ordering::AcquireRelease:
298 O << ".acq_rel";
299 return;
300 case NVPTX::Ordering::SequentiallyConsistent:
301 report_fatal_error(
302 reason: "NVPTX AtomicCode Printer does not support \"seq_cst\" ordering.");
303 return;
304 case NVPTX::Ordering::Volatile:
305 O << ".volatile";
306 return;
307 case NVPTX::Ordering::RelaxedMMIO:
308 O << ".mmio.relaxed";
309 return;
310 }
311 } else if (Modifier == "scope") {
312 auto S = NVPTX::Scope(Imm);
313 switch (S) {
314 case NVPTX::Scope::Thread:
315 case NVPTX::Scope::DefaultDevice:
316 return;
317 case NVPTX::Scope::System:
318 O << ".sys";
319 return;
320 case NVPTX::Scope::Block:
321 O << ".cta";
322 return;
323 case NVPTX::Scope::Cluster:
324 O << ".cluster";
325 return;
326 case NVPTX::Scope::Device:
327 O << ".gpu";
328 return;
329 }
330 report_fatal_error(reason: formatv(
331 Fmt: "NVPTX AtomicCode Printer does not support \"{}\" scope modifier.",
332 Vals: ScopeToString(S)));
333 } else if (Modifier == "addsp") {
334 auto A = NVPTX::AddressSpace(Imm);
335 switch (A) {
336 case NVPTX::AddressSpace::Generic:
337 return;
338 case NVPTX::AddressSpace::Global:
339 case NVPTX::AddressSpace::Const:
340 case NVPTX::AddressSpace::Shared:
341 case NVPTX::AddressSpace::SharedCluster:
342 case NVPTX::AddressSpace::EntryParam:
343 case NVPTX::AddressSpace::DeviceParam:
344 case NVPTX::AddressSpace::Local:
345 O << "." << addressSpaceToString(A, UseParamSubqualifiers: hasParamSubqualifiers(STI));
346 return;
347 }
348 report_fatal_error(reason: formatv(
349 Fmt: "NVPTX AtomicCode Printer does not support \"{}\" addsp modifier.",
350 Vals: addressSpaceToString(A)));
351 } else if (Modifier == "sign") {
352 switch (Imm) {
353 case NVPTX::PTXLdStInstCode::Signed:
354 O << "s";
355 return;
356 case NVPTX::PTXLdStInstCode::Unsigned:
357 O << "u";
358 return;
359 case NVPTX::PTXLdStInstCode::Untyped:
360 O << "b";
361 return;
362 case NVPTX::PTXLdStInstCode::Float:
363 O << "f";
364 return;
365 default:
366 llvm_unreachable("Unknown register type");
367 }
368 }
369 llvm_unreachable(formatv("Unknown Modifier: {}", Modifier).str().c_str());
370}
371
372void NVPTXInstPrinter::printEvictionAndPrefetchHint(const MCInst *MI, int OpNum,
373 const MCSubtargetInfo &,
374 raw_ostream &O,
375 StringRef Modifier) {
376 const MCOperand &MO = MI->getOperand(i: OpNum);
377 unsigned Hint = MO.getImm();
378
379 // If no hint is set, print nothing.
380 if (Hint == 0)
381 return;
382
383 // Check if L2::cache_hint mode is active.
384 bool IsCacheHintMode = NVPTX::isL2CacheHintMode(Hint);
385
386 if (Modifier == "l1") {
387 switch (NVPTX::decodeL1Eviction(Hint)) {
388 case NVPTX::L1Eviction::Normal:
389 return;
390 case NVPTX::L1Eviction::Unchanged:
391 O << ".L1::evict_unchanged";
392 return;
393 case NVPTX::L1Eviction::First:
394 O << ".L1::evict_first";
395 return;
396 case NVPTX::L1Eviction::Last:
397 O << ".L1::evict_last";
398 return;
399 case NVPTX::L1Eviction::NoAllocate:
400 O << ".L1::no_allocate";
401 return;
402 }
403 } else if (Modifier == "l2") {
404 switch (NVPTX::decodeL2Eviction(Hint)) {
405 case NVPTX::L2Eviction::Normal:
406 break;
407 case NVPTX::L2Eviction::First:
408 O << ".L2::evict_first";
409 break;
410 case NVPTX::L2Eviction::Last:
411 O << ".L2::evict_last";
412 break;
413 }
414 if (IsCacheHintMode)
415 O << ".L2::cache_hint";
416 return;
417 } else if (Modifier == "prefetch") {
418 switch (NVPTX::decodeL2Prefetch(Hint)) {
419 case NVPTX::L2Prefetch::None:
420 return;
421 case NVPTX::L2Prefetch::Bytes64:
422 O << ".L2::64B";
423 return;
424 case NVPTX::L2Prefetch::Bytes128:
425 O << ".L2::128B";
426 return;
427 case NVPTX::L2Prefetch::Bytes256:
428 O << ".L2::256B";
429 return;
430 }
431 }
432 llvm_unreachable(formatv("Unknown Modifier: {}", Modifier).str().c_str());
433}
434
435void NVPTXInstPrinter::printCachePolicy(const MCInst *MI, int OpNum,
436 const MCSubtargetInfo &,
437 raw_ostream &O) {
438 const MCOperand &MO = MI->getOperand(i: OpNum);
439 // If the operand is a register and valid, print ", $reg"
440 if (MO.isReg() && MO.getReg().isValid()) {
441 O << ", ";
442 printRegName(OS&: O, Reg: MO.getReg());
443 }
444}
445
446void NVPTXInstPrinter::printMmaCode(const MCInst *MI, int OpNum,
447 const MCSubtargetInfo &, raw_ostream &O,
448 StringRef Modifier) {
449 const MCOperand &MO = MI->getOperand(i: OpNum);
450 int Imm = (int)MO.getImm();
451 if (Modifier.empty() || Modifier == "version") {
452 O << Imm; // Just print out PTX version
453 return;
454 } else if (Modifier == "aligned") {
455 // PTX63 requires '.aligned' in the name of the instruction.
456 if (Imm >= 63)
457 O << ".aligned";
458 return;
459 }
460 llvm_unreachable("Unknown Modifier");
461}
462
463void NVPTXInstPrinter::printMemOperand(const MCInst *MI, int OpNum,
464 const MCSubtargetInfo &STI,
465 raw_ostream &O, StringRef Modifier) {
466 printOperand(MI, OpNo: OpNum, STI, O);
467
468 if (Modifier == "add") {
469 O << ", ";
470 printOperand(MI, OpNo: OpNum + 1, STI, O);
471 } else {
472 if (MI->getOperand(i: OpNum + 1).isImm() &&
473 MI->getOperand(i: OpNum + 1).getImm() == 0)
474 return; // don't print ',0' or '+0'
475 O << "+";
476 printOperand(MI, OpNo: OpNum + 1, STI, O);
477 }
478}
479
480void NVPTXInstPrinter::printUsedBytesMaskPragma(const MCInst *MI, int OpNum,
481 const MCSubtargetInfo &,
482 raw_ostream &O) {
483 auto &Op = MI->getOperand(i: OpNum);
484 assert(Op.isImm() && "Invalid operand");
485 uint32_t Imm = (uint32_t)Op.getImm();
486 if (Imm != UINT32_MAX) {
487 O << ".pragma \"used_bytes_mask " << format_hex(N: Imm, Width: 1) << "\";\n\t";
488 }
489}
490
491void NVPTXInstPrinter::printRegisterOrSinkSymbol(const MCInst *MI, int OpNum,
492 const MCSubtargetInfo &STI,
493 raw_ostream &O) {
494 const MCOperand &Op = MI->getOperand(i: OpNum);
495 if (Op.isReg() && Op.getReg() == MCRegister::NoRegister)
496 O << "_";
497 else
498 printOperand(MI, OpNo: OpNum, STI, O);
499}
500
501void NVPTXInstPrinter::printHexu32imm(const MCInst *MI, int OpNum,
502 const MCSubtargetInfo &, raw_ostream &O) {
503 int64_t Imm = MI->getOperand(i: OpNum).getImm();
504 O << formatHex(Value: Imm) << "U";
505}
506
507void NVPTXInstPrinter::printPrmtMode(const MCInst *MI, int OpNum,
508 const MCSubtargetInfo &, raw_ostream &O) {
509 const MCOperand &MO = MI->getOperand(i: OpNum);
510 int64_t Imm = MO.getImm();
511
512 switch (Imm) {
513 default:
514 return;
515 case NVPTX::PTXPrmtMode::NONE:
516 return;
517 case NVPTX::PTXPrmtMode::F4E:
518 O << ".f4e";
519 return;
520 case NVPTX::PTXPrmtMode::B4E:
521 O << ".b4e";
522 return;
523 case NVPTX::PTXPrmtMode::RC8:
524 O << ".rc8";
525 return;
526 case NVPTX::PTXPrmtMode::ECL:
527 O << ".ecl";
528 return;
529 case NVPTX::PTXPrmtMode::ECR:
530 O << ".ecr";
531 return;
532 case NVPTX::PTXPrmtMode::RC16:
533 O << ".rc16";
534 return;
535 }
536}
537
538void NVPTXInstPrinter::printTmaReductionMode(const MCInst *MI, int OpNum,
539 const MCSubtargetInfo &,
540 raw_ostream &O) {
541 const MCOperand &MO = MI->getOperand(i: OpNum);
542 O << '.'
543 << nvvm::getTMATensorReductionOpName(
544 Op: static_cast<nvvm::TMAReductionOp>(MO.getImm()));
545}
546
547void NVPTXInstPrinter::printCTAGroup(const MCInst *MI, int OpNum,
548 const MCSubtargetInfo &, raw_ostream &O) {
549 const MCOperand &MO = MI->getOperand(i: OpNum);
550 using CGTy = nvvm::CTAGroupKind;
551
552 switch (static_cast<CGTy>(MO.getImm())) {
553 case CGTy::CG_NONE:
554 O << "";
555 return;
556 case CGTy::CG_1:
557 O << ".cta_group::1";
558 return;
559 case CGTy::CG_2:
560 O << ".cta_group::2";
561 return;
562 }
563 llvm_unreachable("Invalid cta_group in printCTAGroup");
564}
565
566void NVPTXInstPrinter::printEvictPolicy(const MCInst *MI, int OpNum,
567 const MCSubtargetInfo &, raw_ostream &O,
568 StringRef Modifier) {
569 const auto Policy =
570 static_cast<nvvm::EvictPolicyType>(MI->getOperand(i: OpNum).getImm());
571 // Evict normal is the default priority policy for prefetch and does not print
572 // a qualifier.
573 if (Policy == nvvm::EvictPolicyType::EVICT_NORMAL)
574 return;
575 O << "." << nvvm::getEvictPolicyName(Policy);
576}
577
578void NVPTXInstPrinter::printCallOperand(const MCInst *MI, int OpNum,
579 const MCSubtargetInfo &, raw_ostream &O,
580 StringRef Modifier) {
581 const MCOperand &MO = MI->getOperand(i: OpNum);
582 assert(MO.isImm() && "Invalid operand");
583 const auto Imm = MO.getImm();
584
585 if (Modifier == "RetList") {
586 assert((Imm == 1 || Imm == 0) && "Invalid return list");
587 if (Imm)
588 O << " (retval0),";
589 return;
590 }
591
592 if (Modifier == "ParamList") {
593 assert(Imm >= 0 && "Invalid parameter list");
594 interleaveComma(c: llvm::seq(Size: Imm), os&: O,
595 each_fn: [&](const auto &I) { O << "param" << I; });
596 return;
597 }
598 llvm_unreachable("Invalid modifier");
599}
600
601template <unsigned Bits>
602void NVPTXInstPrinter::printHexUImm(const MCInst *MI, int OpNum,
603 const MCSubtargetInfo &, raw_ostream &O) {
604 const MCOperand &MO = MI->getOperand(i: OpNum);
605 assert(MO.isImm() && "Expected immediate operand");
606 assert(isInt<Bits>(MO.getImm()) &&
607 "Immediate value does not fit in specified bits");
608 uint64_t Imm = MO.getImm();
609 Imm &= maskTrailingOnes<uint64_t>(N: Bits);
610 O << formatHex(Value: Imm) << "U";
611}
612