use rand::{Rng, RngExt};
use std::f32::consts::PI;
use super::elliptical_slice::{BRACKET_MIN_WIDTH, MAX_BRACKET_ITERS};
pub struct BatchStep {
pub value: Vec<f32>,
pub lnpdf: Vec<f32>,
pub fallbacks: usize,
pub rounds: usize,
}
pub fn elliptical_slice_batch<R: Rng>(
cur: &[f32],
prior_draw: &[f32],
cur_lnpdf: &[f32],
rngs: &mut [R],
lnpdf: &mut impl FnMut(&[f32], &[u32], &mut [f32]),
) -> BatchStep {
let n = cur.len();
assert_eq!(prior_draw.len(), n, "prior_draw must be one per item");
assert_eq!(cur_lnpdf.len(), n, "cur_lnpdf must be one per item");
assert_eq!(rngs.len(), n, "rngs must be one per item");
let mut value = cur.to_vec();
let mut out_lnpdf = cur_lnpdf.to_vec();
let mut hh = vec![0.0f32; n];
let mut angle = vec![0.0f32; n];
let mut lo = vec![0.0f32; n];
let mut hi = vec![0.0f32; n];
for i in 0..n {
let u: f32 = rngs[i].random();
hh[i] = u.ln() + cur_lnpdf[i];
let phi: f32 = rngs[i].random_range(0.0..2.0 * PI);
angle[i] = phi;
lo[i] = phi - 2.0 * PI;
hi[i] = phi;
}
let mut active: Vec<u32> = (0..n as u32).collect();
let mut x = vec![0.0f32; n];
let mut ll = vec![0.0f32; n];
let mut fallbacks = 0usize;
let mut rounds = 0usize;
for _ in 0..MAX_BRACKET_ITERS {
if active.is_empty() {
break;
}
x.clear();
for &i in &active {
let i = i as usize;
let a = angle[i];
x.push(cur[i] * a.cos() + prior_draw[i] * a.sin());
}
ll.resize(active.len(), 0.0);
lnpdf(&x, &active, &mut ll[..active.len()]);
rounds += 1;
let mut still = Vec::with_capacity(active.len());
for (slot, &i) in active.iter().enumerate() {
let idx = i as usize;
if ll[slot] > hh[idx] {
value[idx] = x[slot];
out_lnpdf[idx] = ll[slot];
continue;
}
if angle[idx] < 0.0 {
lo[idx] = angle[idx];
} else {
hi[idx] = angle[idx];
}
if hi[idx] - lo[idx] < BRACKET_MIN_WIDTH {
fallbacks += 1;
continue;
}
angle[idx] = rngs[idx].random_range(lo[idx]..hi[idx]);
still.push(i);
}
active = still;
}
fallbacks += active.len();
BatchStep {
value,
lnpdf: out_lnpdf,
fallbacks,
rounds,
}
}
#[cfg(test)]
#[path = "elliptical_slice_batch_tests.rs"]
mod elliptical_slice_batch_tests;