#![allow(dead_code)]
use serde::Deserialize;
use std::path::Path;
pub fn lcg(seed: u64) -> impl FnMut() -> f32 {
let mut s = seed;
move || {
s = s
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((s >> 33) as f32) / (1u32 << 31) as f32
}
}
pub fn fill_random(mut rng: impl FnMut() -> f32, out: &mut [f32]) {
for v in out {
*v = rng();
}
}
pub fn normal(mut rng: impl FnMut() -> f32) -> f32 {
let mut sum = 0.0;
for _ in 0..12 {
sum += rng();
}
sum - 6.0
}
pub fn coverage(intervals: impl Iterator<Item = (f64, f64)>, labels: &[f32]) -> (f64, f64) {
let mut covered = 0;
let mut width = 0.0;
for ((lower, upper), &label) in intervals.zip(labels) {
let label = f64::from(label);
covered += usize::from(lower <= label && label <= upper);
width += upper - lower;
}
let n = labels.len() as f64;
(covered as f64 / n, width / n)
}
pub fn accuracy(classes: &[u32], labels: &[f32]) -> f32 {
classes
.iter()
.zip(labels)
.filter(|(c, l)| **c as f32 == **l)
.count() as f32
/ labels.len() as f32
}
#[derive(Deserialize)]
pub struct Dataset {
pub n_rows: usize,
pub n_test: usize,
pub n_cols: usize,
pub num_round: usize,
pub objective: String,
pub num_class: usize,
pub metric: String,
pub max_depth: usize,
pub eta: f64,
pub lambda: f64,
pub max_bin: usize,
pub base_score: f64,
pub seed: u64,
}
pub fn read_f32(path: &Path) -> std::io::Result<Vec<f32>> {
let bytes = std::fs::read(path)?;
if bytes.len() % 4 != 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"dataset byte count must be divisible by four",
));
}
Ok(bytes
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect())
}
pub fn load_meta(dir: &Path) -> Result<Dataset, Box<dyn std::error::Error>> {
Ok(serde_json::from_slice(&std::fs::read(
dir.join("meta.json"),
)?)?)
}