use crate::uarch::bpred::Ghr;
#[derive(Clone, Copy, Debug)]
pub struct FoldedHistory {
pub val: u64,
fold_width: usize,
hist_length: usize,
}
impl FoldedHistory {
pub const fn new(fold_width: usize, hist_length: usize) -> Self {
Self { val: 0, fold_width, hist_length }
}
#[inline]
pub const fn update(&mut self, new_bit: bool, old_bit: bool) {
let w = self.fold_width;
if w == 0 {
return;
}
let mask = (1u64 << w) - 1;
let msb = (self.val >> (w - 1)) & 1;
self.val = ((self.val << 1) | msb) & mask;
self.val ^= new_bit as u64;
self.val ^= (old_bit as u64) << (self.hist_length % w);
self.val &= mask;
}
pub fn recompute(&mut self, ghr: &Ghr) {
let w = self.fold_width;
if w == 0 {
self.val = 0;
return;
}
let mask = (1u64 << w) - 1;
let mut result = 0u64;
let num_words = self.hist_length.div_ceil(64);
for word_idx in 0..num_words {
let mut word = ghr.word(word_idx);
let bits_in_word = (self.hist_length - word_idx * 64).min(64);
if bits_in_word < 64 {
word &= (1u64 << bits_in_word) - 1;
}
if word == 0 {
continue;
}
let mut folded = 0u64;
let mut v = word;
while v != 0 {
folded ^= v & mask;
v >>= w;
}
let rot = (word_idx * 64) % w;
if rot > 0 {
folded = ((folded << rot) | (folded >> (w - rot))) & mask;
}
result ^= folded;
}
self.val = result & mask;
}
#[cfg(test)]
pub fn recompute_reference(&mut self, ghr: &Ghr) {
let w = self.fold_width;
if w == 0 {
self.val = 0;
return;
}
let mask = (1u64 << w) - 1;
self.val = 0;
for i in (0..self.hist_length).rev() {
let msb = (self.val >> (w - 1)) & 1;
self.val = ((self.val << 1) | msb) & mask;
self.val ^= ghr.bit(i) as u64;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_folded_history_update_matches_recompute() {
let hist_length = 44;
let fold_width = 11;
let mut ghr = Ghr::with_len(hist_length);
let mut csr = FoldedHistory::new(fold_width, hist_length);
for i in 0..200 {
let bit = (i * 7 + 3) % 2 == 0;
let old_bit = if hist_length > 0 { ghr.bit(hist_length - 1) } else { false };
csr.update(bit, old_bit);
ghr.push(bit);
}
let mut recomputed = FoldedHistory::new(fold_width, hist_length);
recomputed.recompute(&ghr);
assert_eq!(
csr.val, recomputed.val,
"Incremental CSR ({:#x}) != recomputed ({:#x})",
csr.val, recomputed.val
);
}
#[test]
fn test_folded_history_various_widths() {
for &(hist_len, fold_w) in &[(5, 3), (15, 9), (44, 11), (130, 10), (712, 11)] {
let mut ghr = Ghr::with_len(hist_len);
let mut csr = FoldedHistory::new(fold_w, hist_len);
for i in 0..300 {
let bit = i % 3 != 0;
let old_bit = if hist_len > 0 { ghr.bit(hist_len - 1) } else { false };
csr.update(bit, old_bit);
ghr.push(bit);
}
let mut recomputed = FoldedHistory::new(fold_w, hist_len);
recomputed.recompute(&ghr);
assert_eq!(
csr.val, recomputed.val,
"Mismatch for hist_len={hist_len}, fold_w={fold_w}: {:#x} != {:#x}",
csr.val, recomputed.val
);
}
}
#[test]
fn test_fast_recompute_matches_reference() {
let cases = [
(5, 11),
(5, 10),
(5, 9),
(5, 8),
(15, 11),
(15, 10),
(15, 9),
(44, 11),
(44, 10),
(130, 11),
(130, 10),
(247, 11),
(247, 10),
(375, 11),
(375, 10),
(512, 11),
(512, 10),
(712, 11),
(712, 10),
(712, 9),
(4, 9),
(8, 9),
(16, 10),
(32, 10),
(64, 11),
(128, 11),
(256, 11),
(512, 11),
(1, 1),
(2, 1),
(63, 7),
(64, 8),
(65, 8),
(127, 10),
(128, 10),
];
for &(hist_len, fold_w) in &cases {
let mut ghr = Ghr::with_len(hist_len);
for i in 0..500u64 {
ghr.push((i.wrapping_mul(7) ^ i.wrapping_mul(13)) & 1 != 0);
}
let mut fast = FoldedHistory::new(fold_w, hist_len);
fast.recompute(&ghr);
let mut reference = FoldedHistory::new(fold_w, hist_len);
reference.recompute_reference(&ghr);
assert_eq!(
fast.val, reference.val,
"Fast vs reference mismatch for hist_len={hist_len}, fold_w={fold_w}: \
fast={:#x} ref={:#x}",
fast.val, reference.val
);
}
}
}