#![allow(dead_code)]
use ndarray::{ArrayBase, Data, Dimension};
use ndarray_rand::rand::SeedableRng;
use ndarray_rand::rand::rngs::StdRng;
pub fn seeded_rng(seed: u64) -> StdRng {
StdRng::seed_from_u64(seed)
}
#[must_use = "bind the guard to a variable; an unbound guard clears the seed immediately"]
pub struct GlobalSeedGuard;
impl GlobalSeedGuard {
pub fn set(seed: u64) -> Self {
rustyml::set_global_seed(seed);
GlobalSeedGuard
}
}
impl Drop for GlobalSeedGuard {
fn drop(&mut self) {
rustyml::clear_global_seed();
}
}
pub fn assert_allclose<A, S1, S2, D>(actual: &ArrayBase<S1, D>, expected: &ArrayBase<S2, D>, eps: A)
where
A: approx::AbsDiffEq<Epsilon = A> + Copy + std::fmt::Debug,
S1: Data<Elem = A>,
S2: Data<Elem = A>,
D: Dimension,
{
assert_eq!(
actual.shape(),
expected.shape(),
"shape mismatch: actual {:?} vs expected {:?}",
actual.shape(),
expected.shape()
);
for (a, e) in actual.iter().zip(expected.iter()) {
assert!(
a.abs_diff_eq(e, eps),
"element mismatch: actual {a:?} vs expected {e:?} (eps {eps:?})"
);
}
}