use std::collections::VecDeque;
use crate::common::InstSeq;
use crate::uarch::bpred::btb::{BranchKind, Btb, BtbHit};
use crate::uarch::bpred::direction::{BranchClass, DirectionPredictor, Jump, Retired};
use crate::uarch::bpred::ras::{Ras, RasHistory};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ControlInst {
Branch {
target: u64,
},
Jump {
target: u64,
link: Option<u64>,
},
IndirectJump {
returns: bool,
link: Option<u64>,
},
}
impl ControlInst {
const fn class(self) -> BranchClass {
match self {
Self::Branch { .. } => BranchClass::Conditional,
Self::Jump { .. } | Self::IndirectJump { .. } => BranchClass::Unconditional,
}
}
const fn predicts_indirect_target(self) -> bool {
matches!(self, Self::IndirectJump { returns: false, .. })
}
#[must_use]
pub fn from_encoding(pc: u64, size: u64, inst: u32) -> Option<Self> {
use crate::isa::encoding::rv64i::opcodes;
use crate::isa::instruction::{InstructionBits, decode_b_type_imm, decode_j_type_imm};
use crate::isa::reg;
use crate::isa::reg::RegIdx;
let rd = inst.rd();
let rs1 = inst.rs1();
let is_link = |reg: RegIdx| reg == reg::REG_RA || reg == reg::REG_T0;
let link = is_link(rd).then(|| pc.wrapping_add(size));
match inst.opcode() {
opcodes::OP_BRANCH => {
Some(Self::Branch { target: pc.wrapping_add(decode_b_type_imm(inst) as u64) })
}
opcodes::OP_JAL => {
Some(Self::Jump { target: pc.wrapping_add(decode_j_type_imm(inst) as u64), link })
}
opcodes::OP_JALR => {
let returns = is_link(rs1) && (!is_link(rd) || rd != rs1);
Some(Self::IndirectJump { returns, link })
}
_ => None,
}
}
pub const fn kind(self) -> BranchKind {
match self {
Self::Branch { .. } => BranchKind::Conditional,
Self::Jump { link, .. } => BranchKind::Jump { call: link.is_some() },
Self::IndirectJump { returns, link } => {
BranchKind::Indirect { returns, call: link.is_some() }
}
}
}
pub fn from_btb(hit: BtbHit, next_pc: u64) -> Self {
let link = |call: bool| call.then_some(next_pc);
match hit.kind {
BranchKind::Conditional => Self::Branch { target: hit.target },
BranchKind::Jump { call } => Self::Jump { target: hit.target, link: link(call) },
BranchKind::Indirect { returns, call } => {
Self::IndirectJump { returns, link: link(call) }
}
}
}
}
#[derive(Debug)]
struct PredictorHistory<H> {
seq: InstSeq,
pc: u64,
inst: ControlInst,
taken: bool,
target: Option<u64>,
ras: RasHistory,
direction: H,
}
#[derive(Debug)]
pub struct BranchPredUnit<P: DirectionPredictor> {
direction: P,
btb: Btb,
ras: Ras,
in_flight: VecDeque<PredictorHistory<P::History>>,
}
impl<P: DirectionPredictor> BranchPredUnit<P> {
pub fn new(direction: P, btb_size: usize, btb_ways: usize, ras_size: usize) -> Self {
Self {
direction,
btb: Btb::new(btb_size, btb_ways),
ras: Ras::new(ras_size),
in_flight: VecDeque::new(),
}
}
#[cfg(test)]
pub const fn direction(&self) -> &P {
&self.direction
}
pub fn predict(&mut self, seq: InstSeq, pc: u64, inst: ControlInst) -> Option<u64> {
let mut ras = RasHistory::default();
let (target, direction) = match inst {
ControlInst::Branch { target } => {
let (taken, direction) = self.direction.lookup(pc, target);
(taken.then_some(target), direction)
}
ControlInst::Jump { target, link } => {
if let Some(link) = link {
self.ras.push(link, &mut ras);
}
(Some(target), self.direction.unconditional(pc, Jump::Direct))
}
ControlInst::IndirectJump { returns, link } => {
let target = if returns {
self.ras.pop(&mut ras)
} else {
self.direction
.indirect_target(pc)
.or_else(|| self.btb.lookup(pc).map(|hit| hit.target))
};
if let Some(link) = link {
self.ras.push(link, &mut ras);
}
(target, self.direction.unconditional(pc, Jump::Indirect))
}
};
let taken = target.is_some();
self.direction.update_histories(pc, taken, &direction);
self.in_flight.push_back(PredictorHistory { seq, pc, inst, taken, target, ras, direction });
target
}
pub fn squash_after(&mut self, keep: InstSeq) {
if self.squash_younger_than(Some(keep)) {
self.direction.squash_done();
}
}
pub fn squash_all(&mut self) {
if self.squash_younger_than(None) {
self.direction.squash_done();
}
}
pub fn btb_lookup(&self, pc: u64) -> Option<BtbHit> {
self.btb.lookup(pc)
}
pub fn is_predicted(&self, seq: InstSeq) -> bool {
self.in_flight.iter().any(|record| record.seq == seq)
}
pub fn discover(&mut self, seq: InstSeq, pc: u64, inst: ControlInst) -> (Option<u64>, bool) {
let squashed = self.squash_younger_than(Some(seq));
if squashed {
self.direction.squash_done();
}
let target = self.predict(seq, pc, inst);
if let Some(target) = target {
self.btb.update(pc, target, inst.kind());
}
(target, squashed)
}
pub fn correct_target(&mut self, seq: InstSeq, target: u64) {
self.mispredict(seq, true, target);
}
pub fn forget(&mut self, seq: InstSeq, pc: u64) {
let mut squashed = false;
while let Some(record) = self.in_flight.pop_back_if(|record| record.seq >= seq) {
self.ras.squash(record.ras);
self.direction.squash(&record.direction);
squashed = true;
}
if squashed {
self.direction.squash_done();
}
self.btb.invalidate(pc);
}
pub fn mispredict(&mut self, seq: InstSeq, taken: bool, target: u64) {
let squashed = self.squash_younger_than(Some(seq));
let Some(record) = self.in_flight.back_mut().filter(|record| record.seq == seq) else {
if squashed {
self.direction.squash_done();
}
return;
};
record.taken = taken;
record.target = taken.then_some(target);
self.direction.correct(record.pc, taken, &record.direction);
let (pc, inst) = (record.pc, record.inst);
if taken {
self.btb.update(pc, target, inst.kind());
}
}
pub fn commit(&mut self, done: InstSeq) {
while let Some(record) = self.in_flight.pop_front_if(|record| record.seq <= done) {
self.retire(&record);
}
}
fn retire(&mut self, record: &PredictorHistory<P::History>) {
let class = record.inst.class();
let indirect_target = record.target.filter(|_| record.inst.predicts_indirect_target());
let retired = Retired { class, taken: record.taken, indirect_target };
self.direction.commit(record.pc, retired, &record.direction);
}
fn squash_younger_than(&mut self, keep: Option<InstSeq>) -> bool {
let mut squashed = false;
while let Some(record) =
self.in_flight.pop_back_if(|record| keep.is_none_or(|keep| record.seq > keep))
{
self.ras.squash(record.ras);
self.direction.squash(&record.direction);
squashed = true;
}
squashed
}
}