1//===-- RISCVInsertReadWriteCSR.cpp - Insert Read/Write of RISC-V CSR -----===//
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// This file implements the machine function pass to insert read/write of CSR-s
9// of the RISC-V instructions.
10//
11// Currently the pass implements:
12// -Writing and saving frm before an RVV floating-point instruction with a
13// static rounding mode and restores the value after.
14//
15//===----------------------------------------------------------------------===//
16
17#include "MCTargetDesc/RISCVBaseInfo.h"
18#include "RISCV.h"
19#include "RISCVSubtarget.h"
20#include "llvm/CodeGen/MachineFunctionPass.h"
21using namespace llvm;
22
23#define DEBUG_TYPE "riscv-insert-read-write-csr"
24#define RISCV_INSERT_READ_WRITE_CSR_NAME "RISC-V Insert Read/Write CSR Pass"
25
26namespace {
27
28class RISCVInsertReadWriteCSR : public MachineFunctionPass {
29 const TargetInstrInfo *TII;
30
31public:
32 static char ID;
33
34 RISCVInsertReadWriteCSR() : MachineFunctionPass(ID) {}
35
36 bool runOnMachineFunction(MachineFunction &MF) override;
37
38 void getAnalysisUsage(AnalysisUsage &AU) const override {
39 AU.setPreservesCFG();
40 MachineFunctionPass::getAnalysisUsage(AU);
41 }
42
43 StringRef getPassName() const override {
44 return RISCV_INSERT_READ_WRITE_CSR_NAME;
45 }
46
47private:
48 bool emitWriteRoundingMode(MachineBasicBlock &MBB);
49 bool emitWriteRoundingModeOpt(MachineBasicBlock &MBB);
50};
51
52} // end anonymous namespace
53
54char RISCVInsertReadWriteCSR::ID = 0;
55
56INITIALIZE_PASS(RISCVInsertReadWriteCSR, DEBUG_TYPE,
57 RISCV_INSERT_READ_WRITE_CSR_NAME, false, false)
58
59// TODO: Use more accurate rounding mode at the start of MBB.
60bool RISCVInsertReadWriteCSR::emitWriteRoundingModeOpt(MachineBasicBlock &MBB) {
61 bool Changed = false;
62 MachineInstr *LastFRMChanger = nullptr;
63 unsigned CurrentRM = RISCVFPRndMode::DYN;
64 Register SavedFRM;
65
66 for (MachineInstr &MI : MBB) {
67 if (MI.getOpcode() == RISCV::SwapFRMImm ||
68 MI.getOpcode() == RISCV::WriteFRMImm) {
69 CurrentRM = MI.getOperand(i: 0).getImm();
70 SavedFRM = Register();
71 continue;
72 }
73
74 if (MI.getOpcode() == RISCV::WriteFRM) {
75 CurrentRM = RISCVFPRndMode::DYN;
76 SavedFRM = Register();
77 continue;
78 }
79
80 if (MI.isCall() || MI.isInlineAsm() ||
81 MI.readsRegister(Reg: RISCV::FRM, /*TRI=*/nullptr)) {
82 // Restore FRM before unknown operations.
83 if (SavedFRM.isValid())
84 BuildMI(BB&: MBB, I&: MI, MIMD: MI.getDebugLoc(), MCID: TII->get(Opcode: RISCV::WriteFRM))
85 .addReg(RegNo: SavedFRM);
86 CurrentRM = RISCVFPRndMode::DYN;
87 SavedFRM = Register();
88 continue;
89 }
90
91 assert(!MI.modifiesRegister(RISCV::FRM, /*TRI=*/nullptr) &&
92 "Expected that MI could not modify FRM.");
93
94 int FRMIdx = RISCVII::getFRMOpNum(Desc: MI.getDesc());
95 if (FRMIdx < 0)
96 continue;
97 unsigned InstrRM = MI.getOperand(i: FRMIdx).getImm();
98
99 LastFRMChanger = &MI;
100
101 // Make MI implicit use FRM.
102 MI.addOperand(Op: MachineOperand::CreateReg(Reg: RISCV::FRM, /*IsDef*/ isDef: false,
103 /*IsImp*/ isImp: true));
104 Changed = true;
105
106 // Skip if MI uses same rounding mode as FRM.
107 if (InstrRM == CurrentRM)
108 continue;
109
110 if (!SavedFRM.isValid()) {
111 // Save current FRM value to SavedFRM.
112 MachineRegisterInfo *MRI = &MBB.getParent()->getRegInfo();
113 SavedFRM = MRI->createVirtualRegister(RegClass: &RISCV::GPRRegClass);
114 BuildMI(BB&: MBB, I&: MI, MIMD: MI.getDebugLoc(), MCID: TII->get(Opcode: RISCV::SwapFRMImm), DestReg: SavedFRM)
115 .addImm(Val: InstrRM);
116 } else {
117 // Don't need to save current FRM when SavedFRM having value.
118 BuildMI(BB&: MBB, I&: MI, MIMD: MI.getDebugLoc(), MCID: TII->get(Opcode: RISCV::WriteFRMImm))
119 .addImm(Val: InstrRM);
120 }
121 CurrentRM = InstrRM;
122 }
123
124 // Restore FRM if needed.
125 if (SavedFRM.isValid()) {
126 assert(LastFRMChanger && "Expected valid pointer.");
127 MachineInstrBuilder MIB =
128 BuildMI(MF&: *MBB.getParent(), MIMD: {}, MCID: TII->get(Opcode: RISCV::WriteFRM))
129 .addReg(RegNo: SavedFRM);
130 MBB.insertAfter(I: LastFRMChanger, MI: MIB);
131 }
132
133 return Changed;
134}
135
136// This function also swaps frm and restores it when encountering an RVV
137// floating point instruction with a static rounding mode.
138bool RISCVInsertReadWriteCSR::emitWriteRoundingMode(MachineBasicBlock &MBB) {
139 bool Changed = false;
140 for (MachineInstr &MI : MBB) {
141 int FRMIdx = RISCVII::getFRMOpNum(Desc: MI.getDesc());
142 if (FRMIdx < 0)
143 continue;
144
145 unsigned FRMImm = MI.getOperand(i: FRMIdx).getImm();
146
147 // The value is a hint to this pass to not alter the frm value.
148 if (FRMImm == RISCVFPRndMode::DYN)
149 continue;
150
151 Changed = true;
152
153 // Save
154 MachineRegisterInfo *MRI = &MBB.getParent()->getRegInfo();
155 Register SavedFRM = MRI->createVirtualRegister(RegClass: &RISCV::GPRRegClass);
156 BuildMI(BB&: MBB, I&: MI, MIMD: MI.getDebugLoc(), MCID: TII->get(Opcode: RISCV::SwapFRMImm),
157 DestReg: SavedFRM)
158 .addImm(Val: FRMImm);
159 MI.addOperand(Op: MachineOperand::CreateReg(Reg: RISCV::FRM, /*IsDef*/ isDef: false,
160 /*IsImp*/ isImp: true));
161 // Restore
162 MachineInstrBuilder MIB =
163 BuildMI(MF&: *MBB.getParent(), MIMD: {}, MCID: TII->get(Opcode: RISCV::WriteFRM))
164 .addReg(RegNo: SavedFRM);
165 MBB.insertAfter(I: MI, MI: MIB);
166 }
167 return Changed;
168}
169
170bool RISCVInsertReadWriteCSR::runOnMachineFunction(MachineFunction &MF) {
171 // Skip if the vector extension is not enabled.
172 const RISCVSubtarget &ST = MF.getSubtarget<RISCVSubtarget>();
173 if (!ST.hasVInstructions())
174 return false;
175
176 TII = ST.getInstrInfo();
177
178 bool Changed = false;
179
180 for (MachineBasicBlock &MBB : MF) {
181 if (!ST.getCLOpts().frm_insert_opt)
182 Changed |= emitWriteRoundingMode(MBB);
183 else
184 Changed |= emitWriteRoundingModeOpt(MBB);
185 }
186
187 return Changed;
188}
189
190FunctionPass *llvm::createRISCVInsertReadWriteCSRPass() {
191 return new RISCVInsertReadWriteCSR();
192}
193