use crate::matrix::rand_util::name_seed;
use candle_core::{Result, Tensor};
use candle_nn::VarMap;
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
pub fn seed_uniform_vars(varmap: &VarMap, seed: u64, skip: impl Fn(&str) -> bool) -> Result<()> {
let tbl = varmap.data().lock().unwrap();
for (name, var) in tbl.iter() {
if skip(name) {
continue;
}
let dims = var.dims().to_vec();
let n: usize = dims.iter().product();
let draw: Vec<f32> = if name.ends_with(".bias") {
vec![0f32; n]
} else {
let bound = (1.0 / *dims.last().unwrap_or(&1) as f64).sqrt();
let mut rng = StdRng::seed_from_u64(name_seed(seed, name));
(0..n)
.map(|_| ((rng.random::<f64>() * 2.0 - 1.0) * bound) as f32)
.collect()
};
var.set(&Tensor::from_vec(draw, dims.as_slice(), var.device())?)?;
}
Ok(())
}