1//===- NVPTXInstrInfo.cpp - NVPTX Instruction Information -----------------===//
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 file contains the NVPTX implementation of the TargetInstrInfo class.
10//
11//===----------------------------------------------------------------------===//
12
13#include "NVPTXInstrInfo.h"
14#include "NVPTX.h"
15#include "NVPTXSubtarget.h"
16#include "llvm/CodeGen/MachineFunction.h"
17#include "llvm/CodeGen/MachineInstrBuilder.h"
18#include "llvm/CodeGen/MachineRegisterInfo.h"
19#include "llvm/Support/ErrorHandling.h"
20
21using namespace llvm;
22
23#define GET_INSTRINFO_CTOR_DTOR
24#include "NVPTXGenInstrInfo.inc"
25
26// Pin the vtable to this file.
27void NVPTXInstrInfo::anchor() {}
28
29NVPTXInstrInfo::NVPTXInstrInfo(const NVPTXSubtarget &STI)
30 : NVPTXGenInstrInfo(STI, RegInfo), RegInfo() {}
31
32void NVPTXInstrInfo::copyPhysReg(MachineBasicBlock &MBB,
33 MachineBasicBlock::iterator I,
34 const DebugLoc &DL, Register DestReg,
35 Register SrcReg, bool KillSrc,
36 bool RenamableDest, bool RenamableSrc) const {
37 const MachineRegisterInfo &MRI = MBB.getParent()->getRegInfo();
38 const TargetRegisterClass *DestRC = MRI.getRegClass(Reg: DestReg);
39 const TargetRegisterClass *SrcRC = MRI.getRegClass(Reg: SrcReg);
40
41 if (DestRC != SrcRC)
42 report_fatal_error(reason: "Copy one register into another with a different width");
43
44 unsigned Op;
45 if (DestRC == &NVPTX::B1RegClass)
46 Op = NVPTX::MOV_B1_r;
47 else if (DestRC == &NVPTX::B16RegClass)
48 Op = NVPTX::MOV_B16_r;
49 else if (DestRC == &NVPTX::B32RegClass)
50 Op = NVPTX::MOV_B32_r;
51 else if (DestRC == &NVPTX::B64RegClass)
52 Op = NVPTX::MOV_B64_r;
53 else if (DestRC == &NVPTX::B128RegClass)
54 Op = NVPTX::MOV_B128_r;
55 else
56 llvm_unreachable("Bad register copy");
57
58 BuildMI(BB&: MBB, I, MIMD: DL, MCID: get(Opcode: Op), DestReg)
59 .addReg(RegNo: SrcReg, Flags: getKillRegState(B: KillSrc));
60}
61
62/// analyzeBranch - Analyze the branching code at the end of MBB, returning
63/// true if it cannot be understood (e.g. it's a switch dispatch or isn't
64/// implemented for a target). Upon success, this returns false and returns
65/// with the following information in various cases:
66///
67/// 1. If this block ends with no branches (it just falls through to its succ)
68/// just return false, leaving TBB/FBB null.
69/// 2. If this block ends with only an unconditional branch, it sets TBB to be
70/// the destination block.
71/// 3. If this block ends with an conditional branch and it falls through to
72/// an successor block, it sets TBB to be the branch destination block and a
73/// list of operands that evaluate the condition. These
74/// operands can be passed to other TargetInstrInfo methods to create new
75/// branches.
76/// 4. If this block ends with an conditional branch and an unconditional
77/// block, it returns the 'true' destination in TBB, the 'false' destination
78/// in FBB, and a list of operands that evaluate the condition. These
79/// operands can be passed to other TargetInstrInfo methods to create new
80/// branches.
81///
82/// Note that removeBranch and insertBranch must be implemented to support
83/// cases where this method returns success.
84///
85bool NVPTXInstrInfo::analyzeBranch(MachineBasicBlock &MBB,
86 MachineBasicBlock *&TBB,
87 MachineBasicBlock *&FBB,
88 SmallVectorImpl<MachineOperand> &Cond,
89 bool AllowModify) const {
90 // If the block has no terminators, it just falls into the block after it.
91 MachineBasicBlock::iterator I = MBB.end();
92 if (I == MBB.begin() || !isUnpredicatedTerminator(MI: *--I))
93 return false;
94
95 // Get the last instruction in the block.
96 MachineInstr &LastInst = *I;
97
98 // If there is only one terminator instruction, process it.
99 if (I == MBB.begin() || !isUnpredicatedTerminator(MI: *--I)) {
100 if (LastInst.getOpcode() == NVPTX::GOTO) {
101 TBB = LastInst.getOperand(i: 0).getMBB();
102 return false;
103 } else if (LastInst.getOpcode() == NVPTX::CBranch) {
104 // Block ends with fall-through condbranch.
105 TBB = LastInst.getOperand(i: 1).getMBB();
106 Cond.push_back(Elt: LastInst.getOperand(i: 0));
107 Cond.push_back(Elt: LastInst.getOperand(i: 2));
108 return false;
109 }
110 // Otherwise, don't know what this is.
111 return true;
112 }
113
114 // Get the instruction before it if it's a terminator.
115 MachineInstr &SecondLastInst = *I;
116
117 // If there are three terminators, we don't know what sort of block this is.
118 if (I != MBB.begin() && isUnpredicatedTerminator(MI: *--I))
119 return true;
120
121 // If the block ends with NVPTX::GOTO and NVPTX:CBranch, handle it.
122 if (SecondLastInst.getOpcode() == NVPTX::CBranch &&
123 LastInst.getOpcode() == NVPTX::GOTO) {
124 TBB = SecondLastInst.getOperand(i: 1).getMBB();
125 Cond.push_back(Elt: SecondLastInst.getOperand(i: 0));
126 Cond.push_back(Elt: SecondLastInst.getOperand(i: 2));
127 FBB = LastInst.getOperand(i: 0).getMBB();
128 return false;
129 }
130
131 // If the block ends with two NVPTX:GOTOs, handle it. The second one is not
132 // executed, so remove it.
133 if (SecondLastInst.getOpcode() == NVPTX::GOTO &&
134 LastInst.getOpcode() == NVPTX::GOTO) {
135 TBB = SecondLastInst.getOperand(i: 0).getMBB();
136 I = LastInst;
137 if (AllowModify)
138 I->eraseFromParent();
139 return false;
140 }
141
142 // Otherwise, can't handle this.
143 return true;
144}
145
146unsigned NVPTXInstrInfo::removeBranch(MachineBasicBlock &MBB,
147 int *BytesRemoved) const {
148 assert(!BytesRemoved && "code size not handled");
149 MachineBasicBlock::iterator I = MBB.end();
150 if (I == MBB.begin())
151 return 0;
152 --I;
153 if (I->getOpcode() != NVPTX::GOTO && I->getOpcode() != NVPTX::CBranch)
154 return 0;
155
156 // Remove the branch.
157 I->eraseFromParent();
158
159 I = MBB.end();
160
161 if (I == MBB.begin())
162 return 1;
163 --I;
164 if (I->getOpcode() != NVPTX::CBranch)
165 return 1;
166
167 // Remove the branch.
168 I->eraseFromParent();
169 return 2;
170}
171
172unsigned NVPTXInstrInfo::insertBranch(MachineBasicBlock &MBB,
173 MachineBasicBlock *TBB,
174 MachineBasicBlock *FBB,
175 ArrayRef<MachineOperand> Cond,
176 const DebugLoc &DL,
177 int *BytesAdded) const {
178 assert(!BytesAdded && "code size not handled");
179
180 // Shouldn't be a fall through.
181 assert(TBB && "insertBranch must not be told to insert a fallthrough");
182 assert((Cond.size() == 2 || Cond.size() == 0) &&
183 "NVPTX branch conditions have two components!");
184
185 // One-way branch.
186 if (!FBB) {
187 if (Cond.empty()) // Unconditional branch
188 BuildMI(BB: &MBB, MIMD: DL, MCID: get(Opcode: NVPTX::GOTO)).addMBB(MBB: TBB);
189 else // Conditional branch
190 BuildMI(BB: &MBB, MIMD: DL, MCID: get(Opcode: NVPTX::CBranch))
191 .add(MO: Cond[0])
192 .addMBB(MBB: TBB)
193 .add(MO: Cond[1]);
194 return 1;
195 }
196
197 // Two-way Conditional Branch.
198 BuildMI(BB: &MBB, MIMD: DL, MCID: get(Opcode: NVPTX::CBranch)).add(MO: Cond[0]).addMBB(MBB: TBB).add(MO: Cond[1]);
199 BuildMI(BB: &MBB, MIMD: DL, MCID: get(Opcode: NVPTX::GOTO)).addMBB(MBB: FBB);
200 return 2;
201}
202
203bool NVPTXInstrInfo::reverseBranchCondition(
204 SmallVectorImpl<MachineOperand> &Cond) const {
205 assert(Cond.size() == 2 && "Invalid NVPTX branch condition!");
206 Cond[1].setImm(!Cond[1].getImm());
207 return false;
208}
209
210bool NVPTXInstrInfo::invertPredicateBranchInstr(MachineBasicBlock &MBB) const {
211 MachineBasicBlock *TBB = nullptr, *FBB = nullptr;
212 SmallVector<MachineOperand, 4> Cond;
213 if (analyzeBranch(MBB, TBB, FBB, Cond, /*AllowModify=*/false))
214 return false;
215 if (Cond.empty())
216 return false;
217 if (reverseBranchCondition(Cond))
218 return false;
219 DebugLoc DL = MBB.findBranchDebugLoc();
220 removeBranch(MBB);
221 insertBranch(MBB, TBB, FBB, Cond, DL);
222 return true;
223}
224
225static bool isIntegerSetp(const MachineInstr &MI) {
226 switch (MI.getOpcode()) {
227 case NVPTX::SETP_i16:
228 case NVPTX::SETP_i32:
229 case NVPTX::SETP_i64:
230 return true;
231 default:
232 return false;
233 }
234}
235
236static bool isScalarFloatSetp(const MachineInstr &MI) {
237 switch (MI.getOpcode()) {
238 case NVPTX::SETP_bf16:
239 case NVPTX::SETP_f16:
240 case NVPTX::SETP_f32:
241 case NVPTX::SETP_f64:
242 return true;
243 default:
244 return false;
245 }
246}
247
248static int64_t invertIntegerCmpMode(int64_t Mode) {
249 switch (Mode) {
250 case NVPTX::PTXCmpMode::EQ:
251 return NVPTX::PTXCmpMode::NE;
252 case NVPTX::PTXCmpMode::NE:
253 return NVPTX::PTXCmpMode::EQ;
254 case NVPTX::PTXCmpMode::LT:
255 return NVPTX::PTXCmpMode::GE;
256 case NVPTX::PTXCmpMode::LE:
257 return NVPTX::PTXCmpMode::GT;
258 case NVPTX::PTXCmpMode::GT:
259 return NVPTX::PTXCmpMode::LE;
260 case NVPTX::PTXCmpMode::GE:
261 return NVPTX::PTXCmpMode::LT;
262 case NVPTX::PTXCmpMode::LTU:
263 return NVPTX::PTXCmpMode::GEU;
264 case NVPTX::PTXCmpMode::LEU:
265 return NVPTX::PTXCmpMode::GTU;
266 case NVPTX::PTXCmpMode::GTU:
267 return NVPTX::PTXCmpMode::LEU;
268 case NVPTX::PTXCmpMode::GEU:
269 return NVPTX::PTXCmpMode::LTU;
270 default:
271 llvm_unreachable("Invalid integer comparison mode");
272 }
273}
274
275static int64_t invertScalarFloatCmpMode(int64_t Mode) {
276 switch (Mode) {
277 case NVPTX::PTXCmpMode::EQ:
278 return NVPTX::PTXCmpMode::NEU;
279 case NVPTX::PTXCmpMode::NE:
280 return NVPTX::PTXCmpMode::EQU;
281 case NVPTX::PTXCmpMode::EQU:
282 return NVPTX::PTXCmpMode::NE;
283 case NVPTX::PTXCmpMode::NEU:
284 return NVPTX::PTXCmpMode::EQ;
285 case NVPTX::PTXCmpMode::LT:
286 return NVPTX::PTXCmpMode::GEU;
287 case NVPTX::PTXCmpMode::LE:
288 return NVPTX::PTXCmpMode::GTU;
289 case NVPTX::PTXCmpMode::GT:
290 return NVPTX::PTXCmpMode::LEU;
291 case NVPTX::PTXCmpMode::GE:
292 return NVPTX::PTXCmpMode::LTU;
293 case NVPTX::PTXCmpMode::LTU:
294 return NVPTX::PTXCmpMode::GE;
295 case NVPTX::PTXCmpMode::LEU:
296 return NVPTX::PTXCmpMode::GT;
297 case NVPTX::PTXCmpMode::GTU:
298 return NVPTX::PTXCmpMode::LE;
299 case NVPTX::PTXCmpMode::GEU:
300 return NVPTX::PTXCmpMode::LT;
301 case NVPTX::PTXCmpMode::NUM:
302 return NVPTX::PTXCmpMode::NotANumber;
303 case NVPTX::PTXCmpMode::NotANumber:
304 return NVPTX::PTXCmpMode::NUM;
305 default:
306 llvm_unreachable("Invalid scalar float comparison mode");
307 }
308}
309
310static void invertScalarCompareInstr(MachineInstr &MI) {
311 MachineOperand &ModeOp = MI.getOperand(i: 3);
312
313 if (isIntegerSetp(MI))
314 ModeOp.setImm(invertIntegerCmpMode(Mode: ModeOp.getImm()));
315 else if (isScalarFloatSetp(MI))
316 ModeOp.setImm(invertScalarFloatCmpMode(Mode: ModeOp.getImm()));
317 else
318 llvm_unreachable("Invalid SETP instruction");
319}
320
321static unsigned getInvertedSelpOpcode(unsigned Opcode) {
322 switch (Opcode) {
323 case NVPTX::SELP_b16ri:
324 return NVPTX::SELP_b16ir;
325 case NVPTX::SELP_b16ir:
326 return NVPTX::SELP_b16ri;
327 case NVPTX::SELP_b32ri:
328 return NVPTX::SELP_b32ir;
329 case NVPTX::SELP_b32ir:
330 return NVPTX::SELP_b32ri;
331 case NVPTX::SELP_b64ri:
332 return NVPTX::SELP_b64ir;
333 case NVPTX::SELP_b64ir:
334 return NVPTX::SELP_b64ri;
335 case NVPTX::SELP_f16ri:
336 return NVPTX::SELP_f16ir;
337 case NVPTX::SELP_f16ir:
338 return NVPTX::SELP_f16ri;
339 case NVPTX::SELP_f32ri:
340 return NVPTX::SELP_f32ir;
341 case NVPTX::SELP_f32ir:
342 return NVPTX::SELP_f32ri;
343 case NVPTX::SELP_f64ri:
344 return NVPTX::SELP_f64ir;
345 case NVPTX::SELP_f64ir:
346 return NVPTX::SELP_f64ri;
347 case NVPTX::SELP_bf16ri:
348 return NVPTX::SELP_bf16ir;
349 case NVPTX::SELP_bf16ir:
350 return NVPTX::SELP_bf16ri;
351 case NVPTX::SELP_b16rr:
352 case NVPTX::SELP_b16ii:
353 case NVPTX::SELP_b32rr:
354 case NVPTX::SELP_b32ii:
355 case NVPTX::SELP_b64rr:
356 case NVPTX::SELP_b64ii:
357 case NVPTX::SELP_f16rr:
358 case NVPTX::SELP_f16ii:
359 case NVPTX::SELP_f32rr:
360 case NVPTX::SELP_f32ii:
361 case NVPTX::SELP_f64rr:
362 case NVPTX::SELP_f64ii:
363 case NVPTX::SELP_bf16rr:
364 case NVPTX::SELP_bf16ii:
365 return Opcode;
366 default:
367 llvm_unreachable("Unexpected select instruction");
368 }
369}
370
371static void invertSelpInstr(MachineInstr &MI, const NVPTXInstrInfo &TII) {
372 MI.setDesc(TII.get(Opcode: getInvertedSelpOpcode(Opcode: MI.getOpcode())));
373 MachineOperand Src0 = MI.getOperand(i: 1);
374 MI.removeOperand(OpNo: 1);
375 MI.insert(InsertBefore: MI.operands_begin() + 2, Ops: {Src0});
376}
377
378bool NVPTXInstrInfo::findCommutedOpIndices(const MachineInstr &MI,
379 unsigned &SrcOpIdx1,
380 unsigned &SrcOpIdx2) const {
381 if (isIntegerSetp(MI) || isScalarFloatSetp(MI))
382 return fixCommutedOpIndices(ResultIdx1&: SrcOpIdx1, ResultIdx2&: SrcOpIdx2, CommutableOpIdx1: 1, CommutableOpIdx2: 2);
383 return TargetInstrInfo::findCommutedOpIndices(MI, SrcOpIdx1, SrcOpIdx2);
384}
385
386MachineInstr *NVPTXInstrInfo::commuteInstructionImpl(MachineInstr &MI,
387 bool NewMI,
388 unsigned OpIdx1,
389 unsigned OpIdx2) const {
390 assert(!NewMI && "this should never be used");
391
392 if (!isIntegerSetp(MI) && !isScalarFloatSetp(MI))
393 return TargetInstrInfo::commuteInstructionImpl(MI, NewMI, OpIdx1, OpIdx2);
394
395 // For now all users must be invertible conditional branches or selects.
396 // TODO: Support other invertible predicate users.
397 MachineRegisterInfo &MRI = MI.getParent()->getParent()->getRegInfo();
398 SmallVector<MachineBasicBlock *, 4> BranchMBBs;
399 SmallVector<MachineInstr *, 4> SelectInstrs;
400 for (MachineInstr &UseMI :
401 MRI.use_nodbg_instructions(Reg: MI.getOperand(i: 0).getReg())) {
402 if (UseMI.isConditionalBranch())
403 BranchMBBs.push_back(Elt: UseMI.getParent());
404 else if (UseMI.isSelect())
405 SelectInstrs.push_back(Elt: &UseMI);
406 else
407 return nullptr;
408 }
409
410 invertScalarCompareInstr(MI);
411
412 auto *Failed = llvm::find_if(Range&: BranchMBBs, P: [this](MachineBasicBlock *MBB) {
413 return !invertPredicateBranchInstr(MBB&: *MBB);
414 });
415
416 if (Failed != BranchMBBs.end()) {
417 // Couldn't invert one of the branches. Roll back the prefix we
418 // already inverted and the compare-mode flip.
419 for (MachineBasicBlock *MBB : llvm::make_range(x: BranchMBBs.begin(), y: Failed))
420 invertPredicateBranchInstr(MBB&: *MBB);
421 invertScalarCompareInstr(MI);
422 return nullptr;
423 }
424
425 for (MachineInstr *SelectMI : SelectInstrs)
426 invertSelpInstr(MI&: *SelectMI, TII: *this);
427
428 return &MI;
429}
430