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