blob: ce6ab63eebb64664aecf3ec4a77069e66c620db7 [file] [edit]
//===-- NVPTXInstPrinter.cpp - PTX assembly instruction printing ----------===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
//
// Print MCInst instructions to .ptx format.
//
//===----------------------------------------------------------------------===//
#include "MCTargetDesc/NVPTXInstPrinter.h"
#include "MCTargetDesc/NVPTXBaseInfo.h"
#include "NVPTX.h"
#include "NVPTXUtilities.h"
#include "llvm/ADT/StringRef.h"
#include "llvm/IR/NVVMIntrinsicUtils.h"
#include "llvm/MC/MCAsmInfo.h"
#include "llvm/MC/MCExpr.h"
#include "llvm/MC/MCInst.h"
#include "llvm/MC/MCInstrInfo.h"
#include "llvm/MC/MCSubtargetInfo.h"
#include "llvm/MC/MCSymbol.h"
#include "llvm/Support/ErrorHandling.h"
#include "llvm/Support/FormatVariadic.h"
using namespace llvm;
#define DEBUG_TYPE "asm-printer"
#include "NVPTXGenAsmWriter.inc"
static bool hasParamSubqualifiers(const MCSubtargetInfo &STI) {
return STI.hasFeature(NVPTX::PTX83);
}
NVPTXInstPrinter::NVPTXInstPrinter(const MCAsmInfo &MAI, const MCInstrInfo &MII,
const MCRegisterInfo &MRI)
: MCInstPrinter(MAI, MII, MRI) {}
void NVPTXInstPrinter::printRegName(raw_ostream &OS, MCRegister Reg) {
// Decode a register packed by NVPTXAsmPrinter::encodeVirtualRegister.
const auto Kind = static_cast<NVPTX::VirtualRegisterKind>(
Reg.id() >> NVPTX::VirtualRegisterKindShift);
if (Kind == NVPTX::VirtualRegisterKind::Physical) {
// This is actually a physical register, so defer to the autogenerated
// register printer
OS << getRegisterName(Reg);
return;
}
OS << NVPTX::getVirtualRegisterPrefix(Kind)
<< (Reg.id() & NVPTX::VirtualRegisterNumMask);
}
void NVPTXInstPrinter::printInst(const MCInst *MI, uint64_t Address,
StringRef Annot, const MCSubtargetInfo &STI,
raw_ostream &OS) {
printInstruction(MI, Address, STI, OS);
// Next always print the annotation.
printAnnotation(OS, Annot);
}
void NVPTXInstPrinter::printOperand(const MCInst *MI, unsigned OpNo,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &Op = MI->getOperand(OpNo);
if (Op.isReg()) {
MCRegister Reg = Op.getReg();
printRegName(O, Reg);
} else if (Op.isImm()) {
markup(O, Markup::Immediate) << formatImm(Op.getImm());
} else {
assert(Op.isExpr() && "Unknown operand kind in printOperand");
MAI.printExpr(O, *Op.getExpr());
}
}
void NVPTXInstPrinter::printCvtMode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O,
StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
int64_t Imm = MO.getImm();
if (Modifier == "ftz") {
// FTZ flag
if (Imm & NVPTX::PTXCvtMode::FTZ_FLAG)
O << ".ftz";
return;
} else if (Modifier == "sat") {
// SAT flag
if (Imm & NVPTX::PTXCvtMode::SAT_FLAG)
O << ".sat";
return;
} else if (Modifier == "satfinite") {
// SATFINITE flag
if (Imm & NVPTX::PTXCvtMode::SATFINITE_FLAG)
O << ".satfinite";
return;
} else if (Modifier == "pzo") {
// PZO flag
if (Imm & NVPTX::PTXCvtMode::PZO_FLAG)
O << ".pzo";
return;
} else if (Modifier == "relu") {
// RELU flag
if (Imm & NVPTX::PTXCvtMode::RELU_FLAG)
O << ".relu";
return;
} else if (Modifier == "base") {
// Default operand
switch (Imm & NVPTX::PTXCvtMode::BASE_MASK) {
default:
return;
case NVPTX::PTXCvtMode::NONE:
return;
case NVPTX::PTXCvtMode::RNI:
O << ".rni";
return;
case NVPTX::PTXCvtMode::RZI:
O << ".rzi";
return;
case NVPTX::PTXCvtMode::RMI:
O << ".rmi";
return;
case NVPTX::PTXCvtMode::RPI:
O << ".rpi";
return;
case NVPTX::PTXCvtMode::RN:
O << ".rn";
return;
case NVPTX::PTXCvtMode::RZ:
O << ".rz";
return;
case NVPTX::PTXCvtMode::RM:
O << ".rm";
return;
case NVPTX::PTXCvtMode::RP:
O << ".rp";
return;
case NVPTX::PTXCvtMode::RNA:
O << ".rna";
return;
case NVPTX::PTXCvtMode::RS:
O << ".rs";
return;
}
}
llvm_unreachable("Invalid conversion modifier");
}
void NVPTXInstPrinter::printFPRoundingMode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
const auto RM =
static_cast<APFloat::roundingMode>(MI->getOperand(OpNum).getImm());
const StringRef Name = nvvm::GetRoundingModeName(RM);
assert(!Name.empty() && "Invalid FP rounding mode");
O << Name;
}
void NVPTXInstPrinter::printFTZFlag(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
const int Imm = MO.getImm();
if (Imm)
O << ".ftz";
}
void NVPTXInstPrinter::printMultimem(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
if (MO.getImm())
O << "multimem.";
}
void NVPTXInstPrinter::printNegatedPredicate(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
if (MI->getOperand(OpNum).getImm())
O << "!";
}
void NVPTXInstPrinter::printCmpMode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O,
StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
int64_t Imm = MO.getImm();
if (Modifier == "FCmp") {
switch (Imm) {
default:
return;
case NVPTX::PTXCmpMode::EQ:
O << "eq";
return;
case NVPTX::PTXCmpMode::NE:
O << "ne";
return;
case NVPTX::PTXCmpMode::LT:
O << "lt";
return;
case NVPTX::PTXCmpMode::LE:
O << "le";
return;
case NVPTX::PTXCmpMode::GT:
O << "gt";
return;
case NVPTX::PTXCmpMode::GE:
O << "ge";
return;
case NVPTX::PTXCmpMode::EQU:
O << "equ";
return;
case NVPTX::PTXCmpMode::NEU:
O << "neu";
return;
case NVPTX::PTXCmpMode::LTU:
O << "ltu";
return;
case NVPTX::PTXCmpMode::LEU:
O << "leu";
return;
case NVPTX::PTXCmpMode::GTU:
O << "gtu";
return;
case NVPTX::PTXCmpMode::GEU:
O << "geu";
return;
case NVPTX::PTXCmpMode::NUM:
O << "num";
return;
case NVPTX::PTXCmpMode::NotANumber:
O << "nan";
return;
}
}
if (Modifier == "ICmp") {
switch (Imm) {
default:
llvm_unreachable("Invalid ICmp mode");
case NVPTX::PTXCmpMode::EQ:
O << "eq";
return;
case NVPTX::PTXCmpMode::NE:
O << "ne";
return;
case NVPTX::PTXCmpMode::LT:
case NVPTX::PTXCmpMode::LTU:
O << "lt";
return;
case NVPTX::PTXCmpMode::LE:
case NVPTX::PTXCmpMode::LEU:
O << "le";
return;
case NVPTX::PTXCmpMode::GT:
case NVPTX::PTXCmpMode::GTU:
O << "gt";
return;
case NVPTX::PTXCmpMode::GE:
case NVPTX::PTXCmpMode::GEU:
O << "ge";
return;
}
}
if (Modifier == "IType") {
switch (Imm) {
default:
llvm_unreachable("Invalid IType");
case NVPTX::PTXCmpMode::EQ:
case NVPTX::PTXCmpMode::NE:
O << "b";
return;
case NVPTX::PTXCmpMode::LT:
case NVPTX::PTXCmpMode::LE:
case NVPTX::PTXCmpMode::GT:
case NVPTX::PTXCmpMode::GE:
O << "s";
return;
case NVPTX::PTXCmpMode::LTU:
case NVPTX::PTXCmpMode::LEU:
case NVPTX::PTXCmpMode::GTU:
case NVPTX::PTXCmpMode::GEU:
O << "u";
return;
}
}
llvm_unreachable("Empty Modifier");
}
void NVPTXInstPrinter::printAtomicCode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &STI,
raw_ostream &O, StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
int Imm = (int)MO.getImm();
if (Modifier == "sem") {
auto Ordering = NVPTX::Ordering(Imm);
switch (Ordering) {
case NVPTX::Ordering::NotAtomic:
return;
case NVPTX::Ordering::Relaxed:
O << ".relaxed";
return;
case NVPTX::Ordering::Acquire:
O << ".acquire";
return;
case NVPTX::Ordering::Release:
O << ".release";
return;
case NVPTX::Ordering::AcquireRelease:
O << ".acq_rel";
return;
case NVPTX::Ordering::SequentiallyConsistent:
report_fatal_error(
"NVPTX AtomicCode Printer does not support \"seq_cst\" ordering.");
return;
case NVPTX::Ordering::Volatile:
O << ".volatile";
return;
case NVPTX::Ordering::RelaxedMMIO:
O << ".mmio.relaxed";
return;
}
} else if (Modifier == "scope") {
auto S = NVPTX::Scope(Imm);
switch (S) {
case NVPTX::Scope::Thread:
case NVPTX::Scope::DefaultDevice:
return;
case NVPTX::Scope::System:
O << ".sys";
return;
case NVPTX::Scope::Block:
O << ".cta";
return;
case NVPTX::Scope::Cluster:
O << ".cluster";
return;
case NVPTX::Scope::Device:
O << ".gpu";
return;
}
report_fatal_error(formatv(
"NVPTX AtomicCode Printer does not support \"{}\" scope modifier.",
ScopeToString(S)));
} else if (Modifier == "addsp") {
auto A = NVPTX::AddressSpace(Imm);
switch (A) {
case NVPTX::AddressSpace::Generic:
return;
case NVPTX::AddressSpace::Global:
case NVPTX::AddressSpace::Const:
case NVPTX::AddressSpace::Shared:
case NVPTX::AddressSpace::SharedCluster:
case NVPTX::AddressSpace::EntryParam:
case NVPTX::AddressSpace::DeviceParam:
case NVPTX::AddressSpace::Local:
O << "." << addressSpaceToString(A, hasParamSubqualifiers(STI));
return;
}
report_fatal_error(formatv(
"NVPTX AtomicCode Printer does not support \"{}\" addsp modifier.",
addressSpaceToString(A)));
} else if (Modifier == "sign") {
switch (Imm) {
case NVPTX::PTXLdStInstCode::Signed:
O << "s";
return;
case NVPTX::PTXLdStInstCode::Unsigned:
O << "u";
return;
case NVPTX::PTXLdStInstCode::Untyped:
O << "b";
return;
case NVPTX::PTXLdStInstCode::Float:
O << "f";
return;
default:
llvm_unreachable("Unknown register type");
}
}
llvm_unreachable(formatv("Unknown Modifier: {}", Modifier).str().c_str());
}
void NVPTXInstPrinter::printEvictionAndPrefetchHint(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O,
StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
unsigned Hint = MO.getImm();
// If no hint is set, print nothing.
if (Hint == 0)
return;
// Check if L2::cache_hint mode is active.
bool IsCacheHintMode = NVPTX::isL2CacheHintMode(Hint);
if (Modifier == "l1") {
switch (NVPTX::decodeL1Eviction(Hint)) {
case NVPTX::L1Eviction::Normal:
return;
case NVPTX::L1Eviction::Unchanged:
O << ".L1::evict_unchanged";
return;
case NVPTX::L1Eviction::First:
O << ".L1::evict_first";
return;
case NVPTX::L1Eviction::Last:
O << ".L1::evict_last";
return;
case NVPTX::L1Eviction::NoAllocate:
O << ".L1::no_allocate";
return;
}
} else if (Modifier == "l2") {
switch (NVPTX::decodeL2Eviction(Hint)) {
case NVPTX::L2Eviction::Normal:
break;
case NVPTX::L2Eviction::First:
O << ".L2::evict_first";
break;
case NVPTX::L2Eviction::Last:
O << ".L2::evict_last";
break;
}
if (IsCacheHintMode)
O << ".L2::cache_hint";
return;
} else if (Modifier == "prefetch") {
switch (NVPTX::decodeL2Prefetch(Hint)) {
case NVPTX::L2Prefetch::None:
return;
case NVPTX::L2Prefetch::Bytes64:
O << ".L2::64B";
return;
case NVPTX::L2Prefetch::Bytes128:
O << ".L2::128B";
return;
case NVPTX::L2Prefetch::Bytes256:
O << ".L2::256B";
return;
}
}
llvm_unreachable(formatv("Unknown Modifier: {}", Modifier).str().c_str());
}
void NVPTXInstPrinter::printCachePolicy(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
// If the operand is a register and valid, print ", $reg"
if (MO.isReg() && MO.getReg().isValid()) {
O << ", ";
printRegName(O, MO.getReg());
}
}
void NVPTXInstPrinter::printMmaCode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O,
StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
int Imm = (int)MO.getImm();
if (Modifier.empty() || Modifier == "version") {
O << Imm; // Just print out PTX version
return;
} else if (Modifier == "aligned") {
// PTX63 requires '.aligned' in the name of the instruction.
if (Imm >= 63)
O << ".aligned";
return;
}
llvm_unreachable("Unknown Modifier");
}
void NVPTXInstPrinter::printMemOperand(const MCInst *MI, int OpNum,
const MCSubtargetInfo &STI,
raw_ostream &O, StringRef Modifier) {
printOperand(MI, OpNum, STI, O);
if (Modifier == "add") {
O << ", ";
printOperand(MI, OpNum + 1, STI, O);
} else {
if (MI->getOperand(OpNum + 1).isImm() &&
MI->getOperand(OpNum + 1).getImm() == 0)
return; // don't print ',0' or '+0'
O << "+";
printOperand(MI, OpNum + 1, STI, O);
}
}
void NVPTXInstPrinter::printUsedBytesMaskPragma(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
auto &Op = MI->getOperand(OpNum);
assert(Op.isImm() && "Invalid operand");
uint32_t Imm = (uint32_t)Op.getImm();
if (Imm != UINT32_MAX) {
O << ".pragma \"used_bytes_mask " << format_hex(Imm, 1) << "\";\n\t";
}
}
void NVPTXInstPrinter::printRegisterOrSinkSymbol(const MCInst *MI, int OpNum,
const MCSubtargetInfo &STI,
raw_ostream &O) {
const MCOperand &Op = MI->getOperand(OpNum);
if (Op.isReg() && Op.getReg() == MCRegister::NoRegister)
O << "_";
else
printOperand(MI, OpNum, STI, O);
}
void NVPTXInstPrinter::printHexu32imm(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
int64_t Imm = MI->getOperand(OpNum).getImm();
O << formatHex(Imm) << "U";
}
void NVPTXInstPrinter::printPrmtMode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
int64_t Imm = MO.getImm();
switch (Imm) {
default:
return;
case NVPTX::PTXPrmtMode::NONE:
return;
case NVPTX::PTXPrmtMode::F4E:
O << ".f4e";
return;
case NVPTX::PTXPrmtMode::B4E:
O << ".b4e";
return;
case NVPTX::PTXPrmtMode::RC8:
O << ".rc8";
return;
case NVPTX::PTXPrmtMode::ECL:
O << ".ecl";
return;
case NVPTX::PTXPrmtMode::ECR:
O << ".ecr";
return;
case NVPTX::PTXPrmtMode::RC16:
O << ".rc16";
return;
}
}
void NVPTXInstPrinter::printTmaReductionMode(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
O << '.'
<< nvvm::getTMATensorReductionOpName(
static_cast<nvvm::TMAReductionOp>(MO.getImm()));
}
void NVPTXInstPrinter::printCTAGroup(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
using CGTy = nvvm::CTAGroupKind;
switch (static_cast<CGTy>(MO.getImm())) {
case CGTy::CG_NONE:
O << "";
return;
case CGTy::CG_1:
O << ".cta_group::1";
return;
case CGTy::CG_2:
O << ".cta_group::2";
return;
}
llvm_unreachable("Invalid cta_group in printCTAGroup");
}
void NVPTXInstPrinter::printTMAValidateDataFlags(const MCInst *MI, int OpNum,
const MCSubtargetInfo &,
raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
using VDTy = nvvm::TMAValidateDataPattern;
const VDTy Pattern = static_cast<VDTy>(MO.getImm());
// Qualifier omitted for disabled pattern
if (Pattern == VDTy::DISABLED)
return;
O << ".mbarrier::report::validity::"
<< nvvm::getTMAValidateDataPatternName(Pattern);
}
void NVPTXInstPrinter::printEvictPolicy(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O,
StringRef Modifier) {
const auto Policy =
static_cast<nvvm::EvictPolicyType>(MI->getOperand(OpNum).getImm());
// Evict normal is the default priority policy for prefetch and does not print
// a qualifier.
if (Policy == nvvm::EvictPolicyType::EVICT_NORMAL)
return;
O << "." << nvvm::getEvictPolicyName(Policy);
}
void NVPTXInstPrinter::printCallOperand(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O,
StringRef Modifier) {
const MCOperand &MO = MI->getOperand(OpNum);
assert(MO.isImm() && "Invalid operand");
const auto Imm = MO.getImm();
if (Modifier == "RetList") {
assert((Imm == 1 || Imm == 0) && "Invalid return list");
if (Imm)
O << " (retval0),";
return;
}
if (Modifier == "ParamList") {
assert(Imm >= 0 && "Invalid parameter list");
interleaveComma(llvm::seq(Imm), O,
[&](const auto &I) { O << "param" << I; });
return;
}
llvm_unreachable("Invalid modifier");
}
template <unsigned Bits>
void NVPTXInstPrinter::printHexUImm(const MCInst *MI, int OpNum,
const MCSubtargetInfo &, raw_ostream &O) {
const MCOperand &MO = MI->getOperand(OpNum);
assert(MO.isImm() && "Expected immediate operand");
assert(isInt<Bits>(MO.getImm()) &&
"Immediate value does not fit in specified bits");
uint64_t Imm = MO.getImm();
Imm &= maskTrailingOnes<uint64_t>(Bits);
O << formatHex(Imm) << "U";
}