use crate::matrix::rand_util::name_seed;
use candle_core::{DType, 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(())
}
pub fn seed_declared_vars(varmap: &VarMap, seed: u64, skip: impl Fn(&str) -> bool) -> Result<()> {
use rand_distr::{Distribution, Normal};
let at_end = |name: &str, suffix: &str| name == suffix || name.ends_with(&format!(".{suffix}"));
let tbl = varmap.data().lock().unwrap();
for (name, var) in tbl.iter() {
if skip(name) {
continue;
}
let dims = var.dims().to_vec();
let n = var.elem_count();
let sibling_in = name
.strip_suffix(".bias")
.and_then(|stem| tbl.get(&format!("{stem}.weight")))
.and_then(|w| w.dims().get(1).copied());
if sibling_in.is_none() {
let values: Vec<f64> = var.flatten_all()?.to_dtype(DType::F64)?.to_vec1()?;
if values.windows(2).all(|w| w[0] == w[1]) {
continue;
}
}
let mut rng = StdRng::seed_from_u64(name_seed(seed, name));
let normal = |std: f64, rng: &mut StdRng| -> Vec<f64> {
let dist = Normal::new(0.0, std).expect("a finite positive deviation");
(0..n).map(|_| dist.sample(rng)).collect()
};
let draw: Vec<f64> = if let Some(in_dim) = sibling_in {
let bound = 1.0 / (in_dim as f64).sqrt();
(0..n)
.map(|_| (rng.random::<f64>() * 2.0 - 1.0) * bound)
.collect()
} else if at_end(name, "modules.logits") {
normal(
crate::candle::feature_embedding::INIT_LOGIT_JITTER,
&mut rng,
)
} else if at_end(name, crate::candle::lora::U_VAR_NAME) {
normal((1.0 / *dims.last().unwrap_or(&1) as f64).sqrt(), &mut rng)
} else {
let fan_in = if dims.len() < 2 {
1
} else {
dims[1] * dims.iter().skip(2).product::<usize>()
};
normal(2f64.sqrt() / (fan_in as f64).sqrt(), &mut rng)
};
let t = Tensor::from_vec(draw, dims.as_slice(), var.device())?.to_dtype(var.dtype())?;
var.set(&t)?;
}
Ok(())
}