#[derive(Debug, Clone, Copy, PartialEq)]
struct Entry {
rmin: f32,
rmax: f32,
wmin: f32,
value: f32,
}
impl Entry {
#[inline]
fn new(rmin: f32, rmax: f32, wmin: f32, value: f32) -> Self {
Entry {
rmin,
rmax,
wmin,
value,
}
}
#[inline]
fn rmin_next(self) -> f32 {
self.rmin + self.wmin
}
#[inline]
fn rmax_prev(self) -> f32 {
self.rmax - self.wmin
}
}
fn sketch_epsilon(max_bin: usize, n_samples: usize) -> f64 {
let n = n_samples.max(1);
1.0 / max_bin.min(n) as f64
}
fn limit_size_level(maxn: usize, eps: f64) -> usize {
if maxn == 0 {
return 1;
}
let internal_eps = eps / 2.0;
let mut nlevel = 1u32;
loop {
let limit = ((f64::from(nlevel) / internal_eps).ceil() as usize + 1).min(maxn);
if (1usize << nlevel) * limit >= maxn {
return limit;
}
nlevel += 1;
}
}
fn summary_budget(max_bin: usize, n: usize) -> usize {
limit_size_level(n, sketch_epsilon(max_bin, n))
}
fn set_from_sorted(queue: &[(f32, f32)], out: &mut Vec<Entry>) {
out.clear();
let mut wsum = 0f32;
let mut i = 0;
while i < queue.len() {
let value = queue[i].0;
let mut w = queue[i].1;
let mut j = i + 1;
while j < queue.len() && queue[j].0 == value {
w += queue[j].1;
j += 1;
}
out.push(Entry::new(wsum, wsum + w, w, value));
wsum += w;
i = j;
}
}
fn set_prune(data: &mut Vec<Entry>, maxsize: usize) {
if maxsize == 0 {
data.clear();
return;
}
let src_size = data.len();
if src_size <= maxsize {
return;
}
if maxsize == 1 {
data.truncate(1);
return;
}
let d = data.as_mut_slice();
let begin = d[0].rmax;
let range = d[src_size - 1].rmin - d[0].rmax;
let n = maxsize - 1;
let tail = d[src_size - 1];
let mut left = d[1];
let mut right = d[2];
let mut len = 1usize;
let (mut i, mut lastidx) = (1usize, 0usize);
for k in 1..n {
let dx2 = 2.0f32 * ((k as f32 * range) / n as f32 + begin);
while i < src_size - 1 && dx2 >= right.rmax + right.rmin {
i += 1;
left = right;
if i < src_size - 1 {
right = d[i + 1];
}
}
if i == src_size - 1 {
break;
}
if dx2 < left.rmin_next() + right.rmax_prev() {
if i != lastidx {
d[len] = left;
len += 1;
lastidx = i;
}
} else if i + 1 != lastidx {
d[len] = right;
len += 1;
lastidx = i + 1;
}
}
if lastidx != src_size - 1 {
d[len] = tail;
len += 1;
}
data.truncate(len);
}
fn fix_error(data: &mut [Entry]) {
let (mut prev_rmin, mut prev_rmax) = (0f32, 0f32);
for e in data.iter_mut() {
if e.rmin < prev_rmin {
e.rmin = prev_rmin;
} else {
prev_rmin = e.rmin;
}
if e.rmax < prev_rmax {
e.rmax = prev_rmax;
}
let rmin_next = e.rmin_next();
if e.rmax < rmin_next {
e.rmax = rmin_next;
}
prev_rmax = e.rmax;
}
}
fn set_combine(this: &mut Vec<Entry>, other: &[Entry], workspace: &mut Vec<Entry>) {
if other.is_empty() {
return;
}
if this.is_empty() {
this.extend_from_slice(other);
return;
}
workspace.clear();
let (a, b) = (this.as_slice(), other);
let (mut ia, mut ib) = (0usize, 0usize);
let (mut aprev_rmin, mut bprev_rmin) = (0f32, 0f32);
while ia < a.len() && ib < b.len() {
let (ea, eb) = (a[ia], b[ib]);
if ea.value == eb.value {
workspace.push(Entry::new(
ea.rmin + eb.rmin,
ea.rmax + eb.rmax,
ea.wmin + eb.wmin,
ea.value,
));
aprev_rmin = ea.rmin_next();
bprev_rmin = eb.rmin_next();
ia += 1;
ib += 1;
} else if ea.value < eb.value {
workspace.push(Entry::new(
ea.rmin + bprev_rmin,
ea.rmax + eb.rmax_prev(),
ea.wmin,
ea.value,
));
aprev_rmin = ea.rmin_next();
ia += 1;
} else {
workspace.push(Entry::new(
eb.rmin + aprev_rmin,
eb.rmax + ea.rmax_prev(),
eb.wmin,
eb.value,
));
bprev_rmin = eb.rmin_next();
ib += 1;
}
}
if ia < a.len() {
let brmax = b[b.len() - 1].rmax;
for ea in &a[ia..] {
workspace.push(Entry::new(
ea.rmin + bprev_rmin,
ea.rmax + brmax,
ea.wmin,
ea.value,
));
}
}
if ib < b.len() {
let armax = a[a.len() - 1].rmax;
for eb in &b[ib..] {
workspace.push(Entry::new(
eb.rmin + aprev_rmin,
eb.rmax + armax,
eb.wmin,
eb.value,
));
}
}
fix_error(workspace);
this.clear();
this.extend_from_slice(workspace);
}
fn set_prune_sorted(sorted: &[(f32, f32)], max_size: usize, out: &mut Vec<Entry>) {
out.clear();
if sorted.is_empty() {
return;
}
let mut sum_total = 0f64;
let mut unique_values = 0usize;
for (i, &(v, w)) in sorted.iter().enumerate() {
if i == 0 || sorted[i - 1].0 != v {
unique_values += 1;
}
sum_total += f64::from(w);
}
let (mut rmin, mut wmin) = (0f64, 0f64);
let mut last_value = 0f32;
if unique_values <= max_size {
for (i, &(v, w)) in sorted.iter().enumerate() {
if i == 0 {
last_value = v;
wmin = f64::from(w);
continue;
}
if last_value == v {
wmin += f64::from(w);
continue;
}
let rmax = rmin + wmin;
out.push(Entry::new(
rmin as f32,
rmax as f32,
wmin as f32,
last_value,
));
rmin = rmax;
last_value = v;
wmin = f64::from(w);
}
let rmax = rmin + wmin;
out.push(Entry::new(
rmin as f32,
rmax as f32,
wmin as f32,
last_value,
));
return;
}
let mut next_goal = -1f64;
for &(v, w) in sorted {
if next_goal == -1.0 {
next_goal = 0.0;
last_value = v;
wmin = f64::from(w);
continue;
}
if last_value == v {
wmin += f64::from(w);
} else {
let rmax = rmin + wmin;
let mut size = out.len();
if rmax >= next_goal && size != max_size {
if size == 0 || last_value > out[size - 1].value {
out.push(Entry::new(
rmin as f32,
rmax as f32,
wmin as f32,
last_value,
));
size += 1;
}
next_goal = if size == max_size {
sum_total * 2.0 + f64::from(1e-5f32)
} else {
f64::from((size as f64 * sum_total / max_size as f64) as f32)
};
}
rmin = rmax;
wmin = f64::from(w);
last_value = v;
}
}
let size = out.len();
if size == 0 || last_value > out[size - 1].value {
let rmax = rmin + wmin;
out.push(Entry::new(
rmin as f32,
rmax as f32,
wmin as f32,
last_value,
));
}
}
fn query_cut_values(data: &[Entry], max_bin: usize, out: &mut Vec<f32>) {
if data.is_empty() {
out.push(1e-5);
return;
}
let n = data.len();
let advance = |mut cursor: usize, value: f32| {
while cursor < n && data[cursor].value <= value {
cursor += 1;
}
cursor
};
let mut last_cut = data[0].value;
let mut next_value = advance(1, last_cut);
if n <= max_bin {
while next_value < n {
let cpt = data[next_value].value;
out.push(cpt);
last_cut = cpt;
next_value = advance(next_value + 1, last_cut);
}
} else {
let total = f64::from(data[n - 1].rmax);
let mut q = 0usize;
for i in 1..max_bin {
let rank2 = 2.0 * (i as f64 * total / max_bin as f64);
while q < n - 2 && rank2 >= f64::from(data[q + 1].rmin + data[q + 1].rmax) {
q += 1;
}
let queried = if rank2 < f64::from(data[q].rmin_next() + data[q + 1].rmax_prev()) {
data[q]
} else {
data[q + 1]
};
let mut cpt = queried.value;
if cpt <= last_cut {
next_value = advance(next_value, last_cut);
if next_value == n {
break;
}
cpt = data[next_value].value;
} else if next_value < n && data[next_value].value <= cpt {
next_value = advance(next_value + 1, cpt);
}
out.push(cpt);
last_cut = cpt;
}
}
let cpt = data[n - 1].value;
out.push(cpt + (cpt.abs() + 1e-5));
}
pub(crate) struct WQSketch {
max_bin: usize,
limit_size: usize,
queue: Vec<(f32, f32)>,
levels: Vec<Vec<Entry>>,
temp: Vec<Entry>,
workspace: Vec<Entry>,
num_elements: usize,
}
impl WQSketch {
pub(crate) fn new(n_values: usize, max_bin: usize) -> Self {
let limit_size = limit_size_level(n_values, sketch_epsilon(max_bin, n_values));
WQSketch {
max_bin,
limit_size,
queue: Vec::new(),
levels: Vec::new(),
temp: Vec::new(),
workspace: Vec::new(),
num_elements: 0,
}
}
#[inline]
pub(crate) fn push(&mut self, value: f32, weight: f32) {
if weight == 0.0 {
return;
}
self.num_elements += 1;
if !self.queue_push(value, weight) {
self.flush_queue();
self.queue_push(value, weight);
}
}
#[inline]
fn queue_push(&mut self, value: f32, weight: f32) -> bool {
if let Some(last) = self.queue.last_mut()
&& last.0 == value
{
last.1 += weight;
return true;
}
if self.queue.len() == 2 * self.limit_size {
return false;
}
self.queue.push((value, weight));
true
}
fn flush_queue(&mut self) {
self.queue.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
set_from_sorted(&self.queue, &mut self.temp);
self.queue.clear();
self.push_summary();
}
fn push_summary(&mut self) {
let mut l = 0usize;
loop {
if self.levels.len() <= l {
self.levels.push(Vec::new());
}
set_prune(&mut self.temp, self.limit_size);
set_combine(&mut self.temp, &self.levels[l], &mut self.workspace);
self.levels[l].clear();
if self.temp.len() <= self.limit_size {
break;
}
l += 1;
}
self.levels[l].extend_from_slice(&self.temp);
}
pub(crate) fn push_sorted(&mut self, sorted: &[(f32, f32)]) {
self.num_elements += sorted.iter().filter(|&&(_, w)| w != 0.0).count();
set_prune_sorted(sorted, self.max_bin, &mut self.temp);
if !sorted.is_empty() {
self.push_summary();
}
}
fn summary(mut self, max_size: usize) -> Vec<Entry> {
self.flush_queue();
let prune_size = max_size.max(self.limit_size);
let mut out = Vec::new();
for level in &self.levels {
set_combine(&mut out, level, &mut self.workspace);
set_prune(&mut out, prune_size);
}
set_prune(&mut out, max_size);
out
}
pub(crate) fn cut_values(self, out: &mut Vec<f32>) {
let max_bin = self.max_bin;
let budget = summary_budget(max_bin, self.num_elements);
let summary = self.summary(budget);
query_cut_values(&summary, summary.len().min(max_bin), out);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cuts_of(values: &[f32], max_bin: usize) -> Vec<f32> {
let mut sketch = WQSketch::new(values.len(), max_bin);
for &v in values {
sketch.push(v, 1.0);
}
let mut out = Vec::new();
sketch.cut_values(&mut out);
out
}
#[test]
fn few_distinct_values_skip_the_minimum_and_add_sentinel() {
assert_eq!(
cuts_of(&[2.0, 0.0, 1.0, 1.0], 256),
vec![1.0, 2.0, 2.0 + 2.0 + 1e-5]
);
}
#[test]
fn empty_feature_gets_the_upstream_placeholder() {
assert_eq!(cuts_of(&[], 256), vec![1e-5]);
}
#[test]
fn quantile_cuts_are_strictly_increasing_and_bounded() {
let values: Vec<f32> = (0..5000)
.map(|i| ((i * 7919) % 5000) as f32 * 0.001)
.collect();
let cuts = cuts_of(&values, 64);
assert!(cuts.len() <= 64);
assert!(cuts.windows(2).all(|w| w[0] < w[1]));
assert!(*cuts.last().unwrap() > 4.999);
}
#[test]
fn zero_weight_values_do_not_contribute() {
let mut sketch = WQSketch::new(4, 256);
for (v, w) in [(0.0, 1.0), (5.0, 0.0), (1.0, 1.0), (7.0, 0.0)] {
sketch.push(v, w);
}
let mut out = Vec::new();
sketch.cut_values(&mut out);
assert_eq!(out, vec![1.0, 1.0 + 1.0 + 1e-5]);
}
#[test]
fn sorted_path_matches_streaming_path_when_exact() {
let values: Vec<f32> = (0..100).map(|i| (i % 37) as f32).collect();
let streaming = cuts_of(&values, 256);
let mut sorted: Vec<(f32, f32)> = values.iter().map(|&v| (v, 1.0)).collect();
sorted.sort_by(|a, b| a.0.total_cmp(&b.0));
let mut sketch = WQSketch::new(values.len(), 256);
sketch.push_sorted(&sorted);
let mut out = Vec::new();
sketch.cut_values(&mut out);
assert_eq!(out, streaming);
}
}