use std::collections::HashMap;
use crate::config::{Config, MemDepPredictorKind};
use crate::uarch::pipeline::rob::{Rob, RobTag};
use super::predictor::{MdpStats, MemDepPredictor, MemPrediction};
use super::store_set::StoreSetPredictor;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub enum MemDepState {
#[default]
None,
Bypass,
WaitAll,
WaitFor(RobTag),
Resolved(RobTag),
}
#[derive(Debug)]
pub struct MemDepUnit {
predictor: PredictorKind,
deps: HashMap<u32, DepRecord>,
stats: MdpStats,
}
#[derive(Debug)]
enum PredictorKind {
Blind,
StoreSet(Box<StoreSetPredictor>),
}
#[derive(Debug)]
struct DepRecord {
barrier: RobTag,
resolved: bool,
}
impl MemDepUnit {
pub fn new(config: &Config) -> Self {
let predictor = match config.pipeline.mem_dep_predictor {
MemDepPredictorKind::Blind => PredictorKind::Blind,
MemDepPredictorKind::StoreSet => PredictorKind::StoreSet(Box::new(
StoreSetPredictor::new(&config.pipeline.store_set),
)),
};
Self { predictor, deps: HashMap::new(), stats: MdpStats::default() }
}
pub fn dispatch(
&mut self,
pc: u64,
rob_tag: RobTag,
is_load: bool,
is_store: bool,
is_atomic: bool,
) -> MemDepState {
match &mut self.predictor {
PredictorKind::Blind => {
if is_load {
self.stats.predictions_wait_all += 1;
MemDepState::WaitAll
} else {
MemDepState::None
}
}
PredictorKind::StoreSet(predictor) => {
if !is_load && !is_store {
return MemDepState::None;
}
predictor.note_mem_op();
if is_atomic {
if is_load {
self.stats.predictions_wait_all += 1;
return MemDepState::WaitAll;
}
return MemDepState::None;
}
let prediction = predictor.predict(pc, rob_tag, is_store);
if is_store {
predictor.register_store(pc, rob_tag);
}
match prediction {
MemPrediction::NoDep => {
if is_load {
self.stats.predictions_bypass += 1;
MemDepState::Bypass
} else {
MemDepState::None
}
}
MemPrediction::DepOn(barrier) => {
if barrier.is_older_than(rob_tag) {
let _ =
self.deps.insert(rob_tag.0, DepRecord { barrier, resolved: false });
self.stats.predictions_wait_for += 1;
MemDepState::WaitFor(barrier)
} else {
if is_load {
self.stats.predictions_bypass += 1;
MemDepState::Bypass
} else {
MemDepState::None
}
}
}
}
}
}
}
pub fn store_resolved(&mut self, store_rob_tag: RobTag) -> Option<RobTag> {
let mut any_woken = false;
for dep in self.deps.values_mut() {
if dep.barrier == store_rob_tag && !dep.resolved {
dep.resolved = true;
any_woken = true;
}
}
if any_woken { Some(store_rob_tag) } else { None }
}
pub fn issued(&mut self, rob_tag: RobTag) {
let _ = self.deps.remove(&rob_tag.0);
}
pub fn violation(&mut self, load_pc: u64, store_pc: u64) {
self.stats.violations += 1;
if let PredictorKind::StoreSet(predictor) = &mut self.predictor {
predictor.train(load_pc, store_pc);
}
}
pub fn flush(&mut self) {
self.deps.clear();
if let PredictorKind::StoreSet(predictor) = &mut self.predictor {
predictor.flush();
}
}
pub fn flush_after(&mut self, keep_tag: RobTag, rob: &Rob) {
self.deps.retain(|&tag, _| RobTag(tag).is_older_or_eq(keep_tag));
if let PredictorKind::StoreSet(predictor) = &mut self.predictor {
predictor.flush_after(keep_tag);
for entry in rob.iter_in_order() {
if entry.ctrl.mem_write {
predictor.rebuild_lfst_entry(entry.pc, entry.tag);
}
}
}
}
#[cfg(test)]
pub fn stats(&self) -> MdpStats {
self.stats.clone()
}
pub fn take_stats(&mut self) -> MdpStats {
std::mem::take(&mut self.stats)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{Config, MemDepPredictorKind, StoreSetConfig};
fn blind_config() -> Config {
let mut c = Config::default();
c.pipeline.mem_dep_predictor = MemDepPredictorKind::Blind;
c
}
fn store_set_config() -> Config {
let mut c = Config::default();
c.pipeline.mem_dep_predictor = MemDepPredictorKind::StoreSet;
c.pipeline.store_set = StoreSetConfig { ssit_size: 64, lfst_size: 16, clear_period: 0 };
c
}
#[test]
fn test_blind_dispatch_loads_wait_all() {
let config = blind_config();
let mut mdu = MemDepUnit::new(&config);
assert_eq!(mdu.dispatch(0x1000, RobTag(1), true, false, false), MemDepState::WaitAll);
assert_eq!(mdu.stats().predictions_wait_all, 1);
}
#[test]
fn test_blind_dispatch_stores_none() {
let config = blind_config();
let mut mdu = MemDepUnit::new(&config);
assert_eq!(mdu.dispatch(0x2000, RobTag(2), false, true, false), MemDepState::None);
}
#[test]
fn test_blind_dispatch_non_memory_none() {
let config = blind_config();
let mut mdu = MemDepUnit::new(&config);
assert_eq!(mdu.dispatch(0x3000, RobTag(3), false, false, false), MemDepState::None);
}
#[test]
fn test_store_set_unknown_pc_bypass() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
assert_eq!(mdu.dispatch(0x1000, RobTag(1), true, false, false), MemDepState::Bypass);
assert_eq!(mdu.stats().predictions_bypass, 1);
}
#[test]
fn test_store_set_trained_dep() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
let load_pc = 0x1000;
let store_pc = 0x2000;
mdu.violation(load_pc, store_pc);
let s1 = RobTag(5);
assert_eq!(mdu.dispatch(store_pc, s1, false, true, false), MemDepState::None);
let l1 = RobTag(10);
assert_eq!(mdu.dispatch(load_pc, l1, true, false, false), MemDepState::WaitFor(s1));
assert_eq!(mdu.stats().predictions_wait_for, 1);
}
#[test]
fn test_store_resolved_wakeup() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
let load_pc = 0x1000;
let store_pc = 0x2000;
mdu.violation(load_pc, store_pc);
let s1 = RobTag(5);
let _ = mdu.dispatch(store_pc, s1, false, true, false);
let l1 = RobTag(10);
let _ = mdu.dispatch(load_pc, l1, true, false, false);
let woken = mdu.store_resolved(s1);
assert_eq!(woken, Some(s1));
}
#[test]
fn test_issued_cleans_up() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
let load_pc = 0x1000;
let store_pc = 0x2000;
mdu.violation(load_pc, store_pc);
let s1 = RobTag(5);
let _ = mdu.dispatch(store_pc, s1, false, true, false);
let l1 = RobTag(10);
let _ = mdu.dispatch(load_pc, l1, true, false, false);
mdu.issued(l1);
assert!(mdu.deps.is_empty());
}
#[test]
fn test_flush_clears_deps() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
let load_pc = 0x1000;
let store_pc = 0x2000;
mdu.violation(load_pc, store_pc);
let s1 = RobTag(5);
let _ = mdu.dispatch(store_pc, s1, false, true, false);
let l1 = RobTag(10);
let _ = mdu.dispatch(load_pc, l1, true, false, false);
mdu.flush();
assert!(mdu.deps.is_empty());
}
#[test]
fn test_younger_barrier_ignored() {
let config = store_set_config();
let mut mdu = MemDepUnit::new(&config);
let load_pc = 0x1000;
let store_pc = 0x2000;
mdu.violation(load_pc, store_pc);
let l1 = RobTag(1);
assert_eq!(mdu.dispatch(load_pc, l1, true, false, false), MemDepState::Bypass);
let s1 = RobTag(5);
let _ = mdu.dispatch(store_pc, s1, false, true, false);
let l2 = RobTag(10);
assert_eq!(mdu.dispatch(load_pc, l2, true, false, false), MemDepState::WaitFor(s1));
}
}