use std::hash::Hasher;
use crate::sig::{U32Hasher, U64Hasher};
const EMPTY: u8 = 0xFF;
const INITIAL_CAP: usize = 64;
const GROUP: usize = 8;
const SWAR_LO: u64 = 0x0101_0101_0101_0101;
const SWAR_HI: u64 = 0x8080_8080_8080_8080;
pub(crate) trait CountKey: Copy + Eq + Default {
fn ghash(self) -> u64;
}
impl CountKey for u64 {
#[inline(always)]
fn ghash(self) -> u64 {
let mut h = U64Hasher::default();
h.write_u64(self);
h.finish()
}
}
impl CountKey for u32 {
#[inline(always)]
fn ghash(self) -> u64 {
let mut h = U32Hasher::default();
h.write_u32(self);
h.finish()
}
}
#[inline(always)]
fn load_group(ctrl: &[u8], pos: usize) -> u64 {
u64::from_le_bytes(ctrl[pos..pos + GROUP].try_into().unwrap())
}
#[inline(always)]
fn match_lanes(group: u64, byte: u8) -> u64 {
let x = group ^ (SWAR_LO.wrapping_mul(byte as u64));
x.wrapping_sub(SWAR_LO) & !x & SWAR_HI
}
#[inline(always)]
fn lowest_lane(mask: u64) -> usize {
(mask.trailing_zeros() >> 3) as usize
}
pub(crate) struct CountSet<K> {
ctrl: Vec<u8>,
keys: Vec<K>,
counts: Vec<u16>,
cap: usize,
mask: usize,
len: usize,
}
impl<K: CountKey> CountSet<K> {
pub(crate) fn new() -> Self {
Self { ctrl: Vec::new(), keys: Vec::new(), counts: Vec::new(), cap: 0, mask: 0, len: 0 }
}
pub(crate) fn counts(&self) -> impl Iterator<Item = u16> + '_ {
(0..self.cap).filter(move |&i| self.ctrl[i] & 0x80 == 0).map(move |i| self.counts[i])
}
#[inline]
pub(crate) fn bump(&mut self, key: K) {
if (self.len + 1) * 8 > self.cap * 7 {
self.grow();
}
let h = key.ghash();
let h2 = ((h >> 57) as u8) & 0x7F;
let mut pos = (h as usize) & self.mask;
let mut stride = 0usize;
loop {
let group = load_group(&self.ctrl, pos);
let mut m = match_lanes(group, h2);
while m != 0 {
let i = (pos + lowest_lane(m)) & self.mask;
if self.keys[i] == key {
self.counts[i] = self.counts[i].saturating_add(1);
return;
}
m &= m - 1;
}
let e = match_lanes(group, EMPTY);
if e != 0 {
let i = (pos + lowest_lane(e)) & self.mask;
self.set_ctrl(i, h2);
self.keys[i] = key;
self.len += 1;
return;
}
stride += GROUP;
pos = (pos + stride) & self.mask;
}
}
#[inline(always)]
fn set_ctrl(&mut self, i: usize, v: u8) {
self.ctrl[i] = v;
if i < GROUP - 1 {
self.ctrl[i + self.cap] = v;
}
}
fn grow(&mut self) {
let new_cap = if self.cap == 0 { INITIAL_CAP } else { self.cap * 2 };
let new_mask = new_cap - 1;
let mut nctrl = vec![EMPTY; new_cap + GROUP - 1];
let mut nkeys = vec![K::default(); new_cap];
let mut ncounts = vec![2u16; new_cap];
for i in 0..self.cap {
if self.ctrl[i] & 0x80 != 0 {
continue; }
let (key, cnt) = (self.keys[i], self.counts[i]);
let h = key.ghash();
let h2 = ((h >> 57) as u8) & 0x7F;
let mut pos = (h as usize) & new_mask;
let mut stride = 0usize;
loop {
let e = match_lanes(load_group(&nctrl, pos), EMPTY);
if e != 0 {
let j = (pos + lowest_lane(e)) & new_mask;
nctrl[j] = h2;
if j < GROUP - 1 {
nctrl[j + new_cap] = h2;
}
nkeys[j] = key;
ncounts[j] = cnt;
break;
}
stride += GROUP;
pos = (pos + stride) & new_mask;
}
}
self.ctrl = nctrl;
self.keys = nkeys;
self.counts = ncounts;
self.cap = new_cap;
self.mask = new_mask;
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
fn splitmix(state: &mut u64) -> u64 {
*state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = *state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
#[test]
fn count_is_observations_not_bumps() {
let mut cs: CountSet<u64> = CountSet::new();
for _ in 0..4 {
cs.bump(42);
}
assert_eq!(cs.counts().count(), 1);
assert_eq!(cs.counts().next().unwrap(), 5);
}
#[test]
fn bump_matches_a_shadow_map_over_many_colliding_keys() {
let mut cs: CountSet<u64> = CountSet::new();
let mut shadow: HashMap<u64, u32> = HashMap::new(); let mut st = 0xDEAD_BEEF_u64;
for _ in 0..300_000 {
let key = splitmix(&mut st) & 0x1FFF; cs.bump(key);
*shadow.entry(key).or_insert(0) += 1;
}
assert_eq!(cs.counts().count(), shadow.len(), "distinct-key count");
let mut got: Vec<u16> = cs.counts().collect();
let mut want: Vec<u16> =
shadow.values().map(|&b| (b + 1).min(u16::MAX as u32) as u16).collect();
got.sort_unstable();
want.sort_unstable();
assert_eq!(got, want, "per-key counts");
}
#[test]
fn counts_saturate_at_u16_max() {
let mut cs: CountSet<u32> = CountSet::new();
for _ in 0..70_000 {
cs.bump(7u32);
}
assert_eq!(cs.counts().count(), 1);
assert_eq!(cs.counts().next().unwrap(), u16::MAX);
}
#[test]
fn u32_keys_stay_distinct() {
let mut cs: CountSet<u32> = CountSet::new();
cs.bump(1);
cs.bump(1);
cs.bump(2);
let mut got: Vec<u16> = cs.counts().collect();
got.sort_unstable();
assert_eq!(cs.counts().count(), 2);
assert_eq!(got, vec![2, 3]); }
}