use crate::matrix::traits::RandomizedAlgs;
use candle_core::{DType, Device, Result, Tensor};
use crate::candle::sgvb::BlackBoxLikelihood;
pub struct RssSvd {
x_tilde: Tensor,
d_reg_inv: Tensor,
singular_values: Tensor,
singular_values_reg: Tensor,
v_mat: Tensor,
vt: Tensor,
lambda: f64,
effective_rank: usize,
}
impl RssSvd {
pub fn from_genotypes(
x: &Tensor,
max_rank: usize,
lambda: f64,
device: &Device,
) -> Result<Self> {
let (n, p) = x.dims2()?;
let dtype = x.dtype();
let x_data: Vec<f32> = x.to_dtype(DType::F32)?.flatten_all()?.to_vec1()?;
let x_nal = nalgebra::DMatrix::from_row_slice(n, p, &x_data);
let scale = 1.0 / (n as f32).sqrt();
let x_scaled = x_nal * scale;
let (_, d_nal, v_nal) = x_scaled
.rsvd(max_rank)
.map_err(|e| candle_core::Error::Msg(format!("rSVD failed: {}", e)))?;
Self::from_raw_svd(d_nal, v_nal, p, lambda, dtype, device)
}
pub fn project_zscores(&self, z: &Tensor) -> Result<Tensor> {
let vt_z = self.vt.matmul(z)?; vt_z.broadcast_mul(&self.d_reg_inv.unsqueeze(1)?) }
pub fn x_design(&self) -> &Tensor {
&self.x_tilde
}
pub fn singular_values(&self) -> &Tensor {
&self.singular_values
}
pub fn singular_values_reg(&self) -> &Tensor {
&self.singular_values_reg
}
pub fn v_mat(&self) -> &Tensor {
&self.v_mat
}
pub fn lambda(&self) -> f64 {
self.lambda
}
pub fn effective_rank(&self) -> usize {
self.effective_rank
}
pub fn estimate_ldsc_intercept(
d_sq: &[f32],
y_raw: &[Vec<f32>],
num_traits: usize,
) -> (Vec<f32>, Vec<f32>) {
let k = d_sq.len();
if k == 0 || num_traits == 0 {
return (vec![1.0; num_traits], vec![0.0; num_traits]);
}
let mean_x: f32 = d_sq.iter().sum::<f32>() / k as f32;
let var_x: f32 = d_sq.iter().map(|&x| (x - mean_x) * (x - mean_x)).sum();
let (intercepts, slopes): (Vec<f32>, Vec<f32>) = (0..num_traits)
.map(|tt| {
let y2: Vec<f32> = (0..k).map(|kk| y_raw[kk][tt] * y_raw[kk][tt]).collect();
let mean_y: f32 = y2.iter().sum::<f32>() / k as f32;
let cov: f32 = (0..k)
.map(|kk| (d_sq[kk] - mean_x) * (y2[kk] - mean_y))
.sum();
let slope = if var_x > 1e-12 { cov / var_x } else { 0.0 };
let intercept = (mean_y - slope * mean_x).max(1.0);
(intercept, slope)
})
.unzip();
(intercepts, slopes)
}
pub fn estimate_lambda(d_sq: &[f32], y_raw: &[Vec<f32>]) -> f64 {
let k = d_sq.len();
let t = if k > 0 { y_raw[0].len() } else { return 0.0 };
let grad = |lam: f64| -> f64 {
let mut g = 0.0f64;
for kk in 0..k {
let dk2 = d_sq[kk] as f64;
let sigma_sq = (1.0 - lam) * dk2 + lam;
if sigma_sq <= 1e-12 {
continue;
}
let coeff = 1.0 - dk2; let inv_s = 1.0 / sigma_sq;
let sum_y2: f64 = y_raw[kk][..t]
.iter()
.map(|&y| {
let y = y as f64;
y * y
})
.sum();
g += coeff * (t as f64 * inv_s - sum_y2 * inv_s * inv_s);
}
0.5 * g
};
let mut lo = 0.0f64;
let mut hi = 1.0f64;
if grad(lo) >= 0.0 {
return 0.0;
}
if grad(hi) <= 0.0 {
return 1.0;
}
for _ in 0..50 {
let mid = (lo + hi) * 0.5;
if grad(mid) < 0.0 {
lo = mid;
} else {
hi = mid;
}
}
(lo + hi) * 0.5
}
pub fn from_genotypes_estimate_lambda(
x: &Tensor,
z: &Tensor,
max_rank: usize,
device: &Device,
) -> Result<Self> {
let (n, p) = x.dims2()?;
let dtype = x.dtype();
let x_data: Vec<f32> = x.to_dtype(DType::F32)?.flatten_all()?.to_vec1()?;
let x_nal = nalgebra::DMatrix::from_row_slice(n, p, &x_data);
let scale = 1.0 / (n as f32).sqrt();
let x_scaled = x_nal * scale;
let (_, d_nal, v_nal) = x_scaled
.rsvd(max_rank)
.map_err(|e| candle_core::Error::Msg(format!("rSVD failed: {}", e)))?;
let k = d_nal.len();
let d_sq: Vec<f32> = d_nal.iter().map(|&di| di * di).collect();
let v_data: Vec<f32> = v_nal.iter().cloned().collect();
let z_data: Vec<f32> = z.to_dtype(DType::F32)?.flatten_all()?.to_vec1()?;
let (_, t) = z.dims2()?;
let y_raw: Vec<Vec<f32>> = (0..k)
.map(|kk| {
(0..t)
.map(|tt| {
(0..p)
.map(|j| v_data[kk * p + j] * z_data[j * t + tt])
.sum()
})
.collect()
})
.collect();
let lambda = Self::estimate_lambda(&d_sq, &y_raw);
Self::from_raw_svd(d_nal, v_nal, p, lambda, dtype, device)
}
fn from_raw_svd(
d_nal: nalgebra::DVector<f32>,
v_nal: nalgebra::DMatrix<f32>,
p: usize,
lambda: f64,
dtype: DType,
device: &Device,
) -> Result<Self> {
let k = d_nal.len();
let d_vec: Vec<f32> = d_nal
.iter()
.map(|&di| ((di as f64) * (di as f64) + lambda).sqrt() as f32)
.collect();
let d_orig: Vec<f32> = d_nal.iter().cloned().collect();
let d_inv_vec: Vec<f32> = d_vec.iter().map(|&di| 1.0 / di).collect();
let v_data: Vec<f32> = v_nal.iter().cloned().collect();
let v_row_major: Vec<f32> = (0..p)
.flat_map(|row| {
let v_ref = &v_data;
(0..k).map(move |col| v_ref[col * p + row])
})
.collect();
let v_mat = Tensor::from_vec(v_row_major, (p, k), device)?.to_dtype(dtype)?;
let vt = v_mat.t()?;
let d_t = Tensor::from_vec(d_orig, (k,), device)?.to_dtype(dtype)?;
let d_reg_t = Tensor::from_vec(d_vec, (k,), device)?.to_dtype(dtype)?;
let d_reg_inv_t = Tensor::from_vec(d_inv_vec, (k,), device)?.to_dtype(dtype)?;
let x_tilde = vt.broadcast_mul(&d_reg_t.unsqueeze(1)?)?;
Ok(Self {
x_tilde,
d_reg_inv: d_reg_inv_t,
singular_values: d_t,
singular_values_reg: d_reg_t,
v_mat,
vt,
lambda,
effective_rank: k,
})
}
}
pub struct RssLikelihood {
y_tilde: Tensor,
}
impl RssLikelihood {
pub fn new(svd: &RssSvd, z: &Tensor) -> Result<Self> {
let y_tilde = svd.project_zscores(z)?;
Ok(Self { y_tilde })
}
pub fn from_projected(y_tilde: Tensor) -> Self {
Self { y_tilde }
}
pub fn y_tilde(&self) -> &Tensor {
&self.y_tilde
}
}
impl BlackBoxLikelihood for RssLikelihood {
fn log_likelihood(&self, etas: &[&Tensor]) -> Result<Tensor> {
let mut eta = etas[0].clone();
for e in &etas[1..] {
eta = eta.broadcast_add(e)?;
}
let diff_sq = eta.broadcast_sub(&self.y_tilde)?.sqr()?; diff_sq.sum(2)?.sum(1)? * (-0.5) }
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn test_rss_construction() -> Result<()> {
let device = Device::Cpu;
let n = 100;
let p = 50;
let t = 3;
let k = 30;
let lambda = 0.1 / k as f64;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let z = Tensor::randn(0f32, 1f32, (p, t), &device)?;
let svd = RssSvd::from_genotypes(&x, k, lambda, &device)?;
let rss = RssLikelihood::new(&svd, &z)?;
assert_eq!(svd.effective_rank(), k);
assert_eq!(svd.x_design().dims(), &[k, p]);
assert_eq!(rss.y_tilde().dims(), &[k, t]);
assert_eq!(svd.v_mat().dims(), &[p, k]);
println!("n={}, p={}, K={}, T={}, λ={:.4e}", n, p, k, t, svd.lambda());
Ok(())
}
#[test]
fn test_svd_reuse() -> Result<()> {
let device = Device::Cpu;
let n = 100;
let p = 40;
let k = 20;
let lambda = 0.1 / k as f64;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let svd = RssSvd::from_genotypes(&x, k, lambda, &device)?;
let z1 = Tensor::randn(0f32, 1f32, (p, 3), &device)?;
let z2 = Tensor::randn(0f32, 1f32, (p, 5), &device)?;
let rss1 = RssLikelihood::new(&svd, &z1)?;
let rss2 = RssLikelihood::new(&svd, &z2)?;
assert_eq!(rss1.y_tilde().dims(), &[k, 3]);
assert_eq!(rss2.y_tilde().dims(), &[k, 5]);
assert_eq!(svd.x_design().dims(), &[k, p]);
Ok(())
}
#[test]
fn test_rss_evaluation() -> Result<()> {
let device = Device::Cpu;
let n = 100;
let p = 30;
let t = 2;
let s = 5;
let k = 20;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let z = Tensor::randn(0f32, 1f32, (p, t), &device)?;
let svd = RssSvd::from_genotypes(&x, k, 0.1 / k as f64, &device)?;
let rss = RssLikelihood::new(&svd, &z)?;
let beta = Tensor::randn(0f32, 0.1f32, (s, p, t), &device)?;
let eta = svd.x_design().unsqueeze(0)?.broadcast_matmul(&beta)?;
let llik = rss.log_likelihood(&[&eta])?;
assert_eq!(llik.dims(), &[s]);
let vals: Vec<f32> = llik.to_vec1()?;
for v in &vals {
assert!(v.is_finite());
}
Ok(())
}
#[test]
fn test_rss_susie_recovery() -> Result<()> {
use crate::candle::sgvb::{
local_reparam_loss, GaussianPrior, RegressionSGVB, SGVBConfig, SusieVar,
};
use candle_nn::{Optimizer, VarBuilder, VarMap};
let device = Device::Cpu;
let dtype = DType::F32;
let n = 200;
let p = 50;
let t = 1;
let l = 3;
let k = 40;
let lambda = 0.1 / k as f64;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let mut beta_data = vec![0.0f32; p];
beta_data[10] = 3.0;
let true_beta = Tensor::from_vec(beta_data, (p, t), &device)?;
let y = (x.matmul(&true_beta)? + Tensor::randn(0f32, 1f32, (n, t), &device)?)?;
let z = (x.t()?.matmul(&y)? / (n as f64).sqrt())?;
let svd = RssSvd::from_genotypes(&x, k, lambda, &device)?;
let rss = RssLikelihood::new(&svd, &z)?;
println!("K = {}, λ = {:.4e}", svd.effective_rank(), svd.lambda());
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie = SusieVar::new(vb.pp("susie"), l, p, t)?;
let prior = GaussianPrior::new(vb.pp("prior"), 1.0)?;
let config = SGVBConfig::new(50);
let model = RegressionSGVB::from_variational(susie, svd.x_design().clone(), prior, config);
let mut optimizer = candle_nn::AdamW::new_lr(varmap.all_vars(), 0.01)?;
for i in 0..300 {
let loss = local_reparam_loss(&model, &rss, 50, 1.0)?;
optimizer.backward_step(&loss)?;
if i % 50 == 0 {
let lv: f32 = loss.to_scalar()?;
let pip10: f32 = model.variational.pip()?.get(10)?.get(0)?.to_scalar()?;
println!("iter {}: loss={:.4}, PIP[10]={:.4}", i, lv, pip10);
}
}
let pip = model.variational.pip()?;
let pip_10: f32 = pip.get(10)?.get(0)?.to_scalar()?;
let mut other_sum = 0.0f32;
for j in 0..p {
if j != 10 {
other_sum += pip.get(j)?.get(0)?.to_scalar::<f32>()?;
}
}
let other_mean = other_sum / (p - 1) as f32;
println!("PIP[10]={:.4}, others mean={:.4}", pip_10, other_mean);
assert!(pip_10 > other_mean * 3.0);
Ok(())
}
#[test]
fn test_regularization_effect() -> Result<()> {
let device = Device::Cpu;
let n = 50;
let p = 100;
let k = 30;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let svd_small = RssSvd::from_genotypes(&x, k, 1e-6, &device)?;
let svd_default = RssSvd::from_genotypes(&x, k, 0.1 / k as f64, &device)?;
let d_small: Vec<f32> = svd_small.singular_values_reg().to_vec1()?;
let d_default: Vec<f32> = svd_default.singular_values_reg().to_vec1()?;
let last = d_small.len() - 1;
println!(
"Smallest D̃: small_λ={:.4e}, default_λ={:.4e}",
d_small[last], d_default[last]
);
assert!(d_default[last] >= d_small[last]);
Ok(())
}
#[test]
fn test_ldsc_intercept_recovery() {
let k = 200;
let num_traits = 3;
let true_a = 2.0f32;
let true_h = 0.5f32;
let d_sq: Vec<f32> = (0..k).map(|i| 1.0 + 10.0 * (i as f32 / k as f32)).collect();
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
let mut rng = StdRng::seed_from_u64(42);
let y_raw: Vec<Vec<f32>> = (0..k)
.map(|kk| {
let var = true_h * d_sq[kk] + true_a;
let std = var.sqrt();
(0..num_traits)
.map(|_| rng.random::<f32>() * 2.0 * std - std) .collect()
})
.collect();
let big_t = 50;
let y_raw_big: Vec<Vec<f32>> = (0..k)
.map(|kk| {
let var = true_h * d_sq[kk] + true_a;
let std = var.sqrt();
(0..big_t)
.map(|_| {
let u1: f32 = rng.random::<f32>().max(1e-10);
let u2: f32 = rng.random::<f32>();
std * (-2.0 * u1.ln()).sqrt() * (2.0 * std::f32::consts::PI * u2).cos()
})
.collect()
})
.collect();
let (intercepts, slopes) = RssSvd::estimate_ldsc_intercept(&d_sq, &y_raw_big, big_t);
let mean_a: f32 = intercepts.iter().sum::<f32>() / big_t as f32;
let mean_h: f32 = slopes.iter().sum::<f32>() / big_t as f32;
println!(
"LDSC intercept: mean_a={:.3} (true={}), mean_h={:.3} (true={})",
mean_a, true_a, mean_h, true_h
);
assert!(
(mean_a - true_a).abs() < 1.0,
"intercept too far: {}",
mean_a
);
assert!((mean_h - true_h).abs() < 0.5, "slope too far: {}", mean_h);
let (intercepts_3, _) = RssSvd::estimate_ldsc_intercept(&d_sq, &y_raw, num_traits);
for &a in &intercepts_3 {
assert!(a >= 1.0, "intercept should be clamped >= 1.0, got {}", a);
}
}
#[test]
fn test_ldsc_intercept_no_inflation() {
let k = 100;
let num_traits = 20;
let d_sq: Vec<f32> = (0..k).map(|i| 0.5 + 5.0 * (i as f32 / k as f32)).collect();
use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng};
let mut rng = StdRng::seed_from_u64(123);
let y_raw: Vec<Vec<f32>> = (0..k)
.map(|_| {
(0..num_traits)
.map(|_| {
let u1: f32 = rng.random::<f32>().max(1e-10);
let u2: f32 = rng.random::<f32>();
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f32::consts::PI * u2).cos()
})
.collect()
})
.collect();
let (intercepts, _slopes) = RssSvd::estimate_ldsc_intercept(&d_sq, &y_raw, num_traits);
let mean_a: f32 = intercepts.iter().sum::<f32>() / num_traits as f32;
println!("No-inflation LDSC: mean_a={:.3}", mean_a);
assert!(mean_a < 2.0, "intercept should be near 1.0, got {}", mean_a);
}
}