#![allow(clippy::cast_precision_loss, clippy::cast_possible_truncation, clippy::cast_sign_loss)]
use rand::{Rng, RngExt, seq::IndexedMutRandom};
pub struct SkipSampler<'a, I, R> {
iterator: I,
rng: &'a mut R,
remaining_samples: usize,
remaining_population: usize,
use_method_a: bool,
vprime: f32,
max_start_index: usize,
rem_samples_inv: f32,
threshold: usize,
}
impl<'a, I: Iterator, R: Rng> SkipSampler<'a, I, R> {
const ALPHAINV: usize = 13;
pub fn new(iterator: I, target: usize, total_items: usize, rng: &'a mut R) -> std::io::Result<Self> {
if target > total_items {
return Err(std::io::Error::other(format!(
"The sample target size {target} should not be larger than the population size {total_items}!"
)));
}
let remaining_samples = target;
let remaining_population = total_items;
let max_start_index = remaining_population - remaining_samples + 1;
let rem_samples_inv = (remaining_samples as f32).recip();
let vprime = rng.random::<f32>().powf(rem_samples_inv);
Ok(SkipSampler {
iterator,
remaining_samples,
remaining_population,
max_start_index,
rem_samples_inv,
vprime,
threshold: Self::ALPHAINV * target,
use_method_a: false,
rng,
})
}
fn next_skip(&mut self) -> Option<usize> {
if self.remaining_samples == 0 {
return None;
}
if self.use_method_a {
return Some(self.next_skip_method_a());
}
if self.threshold >= self.remaining_population || self.remaining_samples == 1 {
self.use_method_a = true;
return Some(self.next_skip_method_a());
}
let rem_samples_min1_inv = ((self.remaining_samples - 1) as f32).recip();
loop {
let mut x = self.remaining_population as f32 * (1.0 - self.vprime);
let mut skip = x as usize;
while skip >= self.max_start_index {
self.vprime = self.rng.random::<f32>().powf(self.rem_samples_inv);
x = self.remaining_population as f32 * (1.0 - self.vprime);
skip = x as usize;
}
let y1 = (self.rng.random::<f32>() * self.remaining_population as f32 / self.max_start_index as f32)
.powf(rem_samples_min1_inv);
self.vprime = y1
* (1.0 - x / self.remaining_population as f32)
* (self.max_start_index as f32 / (self.max_start_index - skip) as f32);
if self.vprime <= 1.0 {
return Some(self.yield_skip(skip, rem_samples_min1_inv));
}
let mut y2 = 1.0;
let (mut bottom, limit) = if self.remaining_samples - 1 > skip {
(
self.remaining_population - self.remaining_samples,
self.remaining_population - skip,
)
} else {
(self.remaining_population - skip - 1, self.max_start_index)
};
for top in (limit..self.remaining_population).rev() {
y2 *= top as f32 / bottom as f32;
bottom -= 1;
}
if self.remaining_population as f32 / (self.remaining_population as f32 - x)
>= y1 * y2.powf(rem_samples_min1_inv)
{
self.vprime = self.rng.random::<f32>().powf(rem_samples_min1_inv);
return Some(self.yield_skip(skip, rem_samples_min1_inv));
}
self.vprime = self.rng.random::<f32>().powf(self.rem_samples_inv);
}
}
fn yield_skip(&mut self, skip: usize, rem_samples_min1_inv: f32) -> usize {
self.remaining_population -= skip + 1;
self.remaining_samples -= 1;
self.rem_samples_inv = rem_samples_min1_inv;
self.max_start_index -= skip;
self.threshold -= Self::ALPHAINV;
skip
}
fn next_skip_method_a(&mut self) -> usize {
if self.remaining_samples > 1 {
let mut top = self.remaining_population - self.remaining_samples;
let mut skip = 0;
let variate = self.rng.random::<f32>();
let mut quot = top as f32 / self.remaining_population as f32;
while quot > variate {
skip += 1;
top -= 1;
self.remaining_population -= 1;
quot *= top as f32 / self.remaining_population as f32;
}
self.remaining_population -= 1;
self.remaining_samples -= 1;
skip
} else {
let skip = self.rng.random_range(0..self.remaining_population);
self.remaining_population -= skip + 1;
self.remaining_samples = 0;
skip
}
}
}
impl<I: Iterator, R: Rng> Iterator for SkipSampler<'_, I, R> {
type Item = I::Item;
fn next(&mut self) -> Option<Self::Item> {
let skip = self.next_skip()?;
self.iterator.nth(skip)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
(0, Some(self.remaining_samples))
}
fn nth(&mut self, n: usize) -> Option<Self::Item> {
let mut inner_n = self.next_skip()?;
for _ in 0..n {
inner_n += 1 + self.next_skip()?;
}
self.iterator.nth(inner_n)
}
}
pub struct BernoulliSampler<'a, I, R> {
iterator: I,
prob: f32,
rng: &'a mut R,
}
impl<I, R> Iterator for BernoulliSampler<'_, I, R>
where
I: Iterator,
R: Rng,
{
type Item = I::Item;
fn next(&mut self) -> Option<Self::Item> {
loop {
let item = self.iterator.next()?;
if self.rng.random::<f32>() < self.prob {
return Some(item);
}
}
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
(0, self.iterator.size_hint().1)
}
}
pub fn downsample_reservoir<I, T, R>(iter: I, rng: &mut R, target: usize) -> Vec<T>
where
I: Iterator<Item = T>,
R: Rng, {
let mut reservoir = Vec::with_capacity(target);
let n_sample = target;
let mut r = rng.random::<f32>();
let mut w = (r.ln() / n_sample as f32).exp();
r = rng.random::<f32>();
let mut s = (r.ln() / (1.0 - w).ln()).floor() as usize;
for (i, sample) in iter.enumerate() {
if i < n_sample {
reservoir.push(sample);
}
else if s == 0 {
if let Some(slot) = reservoir.choose_mut(rng) {
*slot = sample;
}
r = rng.random::<f32>();
w *= (r.ln() / n_sample as f32).exp();
r = rng.random::<f32>();
s = (r.ln() / (1.0 - w).ln()).floor() as usize;
}
else {
s -= 1;
}
}
reservoir
}
pub trait DownsampleBernoulli: Iterator + Sized {
#[inline]
fn downsample_bernoulli<R>(self, prob: f32, rng: &mut R) -> BernoulliSampler<'_, Self, R> {
BernoulliSampler {
iterator: self,
prob,
rng,
}
}
}
impl<I: Iterator> DownsampleBernoulli for I {}
pub trait DownsampleKnownSize: ExactSizeIterator + Sized {
fn downsample_known_size<R>(self, rng: &mut R, target: usize) -> std::io::Result<SkipSampler<'_, Self, R>>
where
R: Rng;
}
impl<I> DownsampleKnownSize for I
where
I: ExactSizeIterator,
{
#[inline]
fn downsample_known_size<R>(self, rng: &mut R, target: usize) -> std::io::Result<SkipSampler<'_, I, R>>
where
R: Rng, {
let total_items = self.len();
SkipSampler::new(self, target, total_items, rng)
}
}