use candle_core::{DType, Device, Result, Tensor};
use candle_nn::VarBuilder;
use super::susie_util::pip_from_alpha;
use super::traits::{ComponentVariational, VariationalDistribution};
use super::variant_tree::VariantTree;
pub struct MultiLevelSusieVar {
tree: VariantTree,
logits_per_level: Vec<Tensor>,
masks_per_level: Vec<Tensor>,
path_indices: Vec<Tensor>,
beta_mean: Tensor,
beta_ln_std: Tensor,
num_components: usize,
temperature: f64,
}
impl MultiLevelSusieVar {
pub fn new(
vb: VarBuilder,
tree: VariantTree,
num_components: usize,
k: usize,
temperature: f64,
) -> Result<Self> {
let device = vb.device().clone();
let dtype = vb.dtype();
let mut logits_per_level = Vec::with_capacity(tree.depth);
let mut masks_per_level = Vec::with_capacity(tree.depth);
let mut path_indices = Vec::with_capacity(tree.depth);
for (d, level) in tree.levels.iter().enumerate() {
let logits = vb.get_with_hints(
(num_components, level.num_groups, level.max_children, k),
&format!("logits_level_{}", d),
candle_nn::Init::Const(0.0),
)?;
logits_per_level.push(logits);
let mask_data: Vec<f32> = level
.mask
.iter()
.flat_map(|row| row.iter().map(|&b| if b { 1.0 } else { 0.0 }))
.collect();
let mask =
Tensor::from_vec(mask_data, (level.num_groups, level.max_children), &device)?
.to_dtype(dtype)?;
masks_per_level.push(mask);
let idx_data: Vec<u32> = level.flat_path_indices.iter().map(|&i| i as u32).collect();
let idx = Tensor::from_vec(idx_data, (tree.num_variants,), &device)?;
path_indices.push(idx);
}
let beta_mean = vb.get_with_hints(
(num_components, tree.num_variants, k),
"beta_mean",
candle_nn::Init::Randn {
mean: 0.0,
stdev: 0.01,
},
)?;
let beta_ln_std = vb.get_with_hints(
(num_components, tree.num_variants, k),
"beta_ln_std",
candle_nn::Init::Const(0.0),
)?;
Ok(Self {
tree,
logits_per_level,
masks_per_level,
path_indices,
beta_mean,
beta_ln_std,
num_components,
temperature,
})
}
fn level_log_softmax(&self, d: usize) -> Result<Tensor> {
let logits = &self.logits_per_level[d];
let mask = &self.masks_per_level[d];
let scaled = if (self.temperature - 1.0).abs() > 1e-10 {
(logits / self.temperature)?
} else {
logits.clone()
};
let neg_inf_mask = ((1.0 - mask)? * (-1e30))?.unsqueeze(0)?.unsqueeze(3)?;
candle_nn::ops::log_softmax(&scaled.broadcast_add(&neg_inf_mask)?, 2)
}
pub fn log_alpha(&self) -> Result<Tensor> {
let p = self.tree.num_variants;
let (l_dim, _, _, k) = self.logits_per_level[0].dims4()?;
let dtype = self.logits_per_level[0].dtype();
let device = self.logits_per_level[0].device();
let mut total = Tensor::zeros((l_dim, p, k), dtype, device)?;
for d in 0..self.tree.depth {
let log_sm = self.level_log_softmax(d)?;
let indices = &self.path_indices[d];
let (_, g_d, c_d, _) = log_sm.dims4()?;
let flat = log_sm.reshape((l_dim, g_d * c_d, k))?;
let gathered = flat.index_select(indices, 1)?;
total = (total + gathered)?;
}
Ok(total)
}
pub fn alpha(&self) -> Result<Tensor> {
self.log_alpha()?.exp()
}
pub fn pip(&self) -> Result<Tensor> {
pip_from_alpha(&self.alpha()?)
}
pub fn theta_mean(&self) -> Result<Tensor> {
let alpha = self.alpha()?; let weighted = alpha.broadcast_mul(&self.beta_mean)?; weighted.sum(0) }
pub fn set_temperature(&mut self, temperature: f64) {
self.temperature = temperature;
}
pub fn temperature(&self) -> f64 {
self.temperature
}
pub fn beta_mean(&self) -> &Tensor {
&self.beta_mean
}
pub fn beta_std(&self) -> Result<Tensor> {
self.beta_ln_std.exp()
}
pub fn num_components(&self) -> usize {
self.num_components
}
pub fn device(&self) -> &Device {
self.logits_per_level[0].device()
}
pub fn dtype(&self) -> DType {
self.logits_per_level[0].dtype()
}
pub fn kl_categorical(&self, _prior_alpha: f64) -> Result<Tensor> {
let device = self.logits_per_level[0].device();
let dtype = self.logits_per_level[0].dtype();
let mut total_kl = Tensor::new(0f32, device)?.to_dtype(dtype)?;
for d in 0..self.tree.depth {
let mask = &self.masks_per_level[d];
let log_sm = self.level_log_softmax(d)?;
let alpha_d = log_sm.exp()?;
let children_count = mask.sum(1)?; let log_prior = children_count.log()?.unsqueeze(0)?.unsqueeze(2)?; let log_prior = log_prior.unsqueeze(3)?;
let mask_4d = mask.unsqueeze(0)?.unsqueeze(3)?.to_dtype(dtype)?;
let kl_elements = alpha_d
.broadcast_mul(&log_sm.broadcast_add(&log_prior)?)?
.broadcast_mul(&mask_4d)?;
let level_kl = kl_elements.sum_all()?;
total_kl = (total_kl + level_kl)?;
}
Ok(total_kl)
}
}
impl ComponentVariational for MultiLevelSusieVar {
fn alpha(&self) -> Result<Tensor> {
self.alpha()
}
fn beta_mean(&self) -> Result<Tensor> {
Ok(self.beta_mean().clone())
}
fn beta_std(&self) -> Result<Tensor> {
self.beta_std()
}
fn num_components(&self) -> usize {
self.num_components()
}
}
impl VariationalDistribution for MultiLevelSusieVar {
fn mean(&self) -> Result<Tensor> {
self.theta_mean()
}
fn var(&self) -> Result<Tensor> {
let alpha = self.alpha()?; let mu = &self.beta_mean; let sigma_sq = (&self.beta_ln_std * 2.0)?.exp()?;
let mu_sq = mu.sqr()?;
let second_moment_l = alpha.broadcast_mul(&(&sigma_sq + &mu_sq)?)?; let second_moment = second_moment_l.sum(0)?;
let first_moment_l = alpha.broadcast_mul(mu)?; let first_moment = first_moment_l.sum(0)?; let first_moment_sq = first_moment.sqr()?;
(second_moment - first_moment_sq)?.clamp(1e-8, f64::INFINITY)
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
use candle_nn::{VarBuilder, VarMap};
#[test]
fn test_multilevel_susie_shapes() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F32;
let l = 3;
let p = 100;
let k = 2;
let tree = VariantTree::regular(p, 10);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie = MultiLevelSusieVar::new(vb, tree, l, k, 1.0)?;
let alpha = susie.alpha()?;
assert_eq!(alpha.dims(), &[l, p, k]);
let pip = susie.pip()?;
assert_eq!(pip.dims(), &[p, k]);
let theta_mean = susie.theta_mean()?;
assert_eq!(theta_mean.dims(), &[p, k]);
let var = susie.var()?;
assert_eq!(var.dims(), &[p, k]);
Ok(())
}
#[test]
fn test_alpha_sums_to_one() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F64;
let l = 3;
let p = 50;
let k = 2;
let tree = VariantTree::regular(p, 10);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie = MultiLevelSusieVar::new(vb, tree, l, k, 1.0)?;
let alpha = susie.alpha()?; let alpha_sum = alpha.sum(1)?;
for i in 0..l {
for j in 0..k {
let sum: f64 = alpha_sum.get(i)?.get(j)?.to_scalar()?;
assert!(
(sum - 1.0).abs() < 1e-5,
"Alpha should sum to 1 for l={}, k={}, got {}",
i,
j,
sum
);
}
}
Ok(())
}
#[test]
fn test_alpha_sums_with_uneven_groups() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F64;
let l = 2;
let p = 23;
let k = 1;
let tree = VariantTree::regular(p, 10);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie = MultiLevelSusieVar::new(vb, tree, l, k, 1.0)?;
let alpha = susie.alpha()?;
let alpha_sum = alpha.sum(1)?;
for i in 0..l {
let sum: f64 = alpha_sum.get(i)?.get(0)?.to_scalar()?;
assert!(
(sum - 1.0).abs() < 1e-5,
"Alpha should sum to 1 for l={}, got {}",
i,
sum
);
}
Ok(())
}
#[test]
fn test_temperature_effect() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F64;
let p = 20;
let k = 1;
let tree = VariantTree::regular(p, 5);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie_sharp = MultiLevelSusieVar::new(vb, tree.clone(), 1, k, 0.1)?;
let alpha_sharp = susie_sharp.alpha()?;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie_smooth = MultiLevelSusieVar::new(vb, tree, 1, k, 10.0)?;
let alpha_smooth = susie_smooth.alpha()?;
let sum_sharp: f64 = alpha_sharp.sum(1)?.get(0)?.get(0)?.to_scalar()?;
let sum_smooth: f64 = alpha_smooth.sum(1)?.get(0)?.get(0)?.to_scalar()?;
assert!((sum_sharp - 1.0).abs() < 1e-5);
assert!((sum_smooth - 1.0).abs() < 1e-5);
let val: f64 = alpha_sharp.get(0)?.get(0)?.get(0)?.to_scalar()?;
assert!(
(val - 1.0 / p as f64).abs() < 1e-5,
"Expected ~{}, got {}",
1.0 / p as f64,
val
);
Ok(())
}
#[test]
fn test_with_linear_model() -> Result<()> {
use crate::candle::sgvb::traits::BlackBoxLikelihood;
use crate::candle::sgvb::{local_reparam_loss, GaussianPrior, RegressionSGVB, SGVBConfig};
let device = Device::Cpu;
let dtype = DType::F32;
let n = 30;
let p = 20;
let k = 1;
let l = 2;
let x = Tensor::randn(0f32, 1f32, (n, p), &device)?;
let y = Tensor::randn(0f32, 1f32, (n, k), &device)?;
let tree = VariantTree::regular(p, 5);
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let susie = MultiLevelSusieVar::new(vb.pp("susie"), tree, l, k, 1.0)?;
let prior = GaussianPrior::new(vb.pp("prior"), 1.0)?;
let config = SGVBConfig::default();
let model = RegressionSGVB::from_variational(susie, x, prior, config);
struct GaussianLik {
y: Tensor,
}
impl BlackBoxLikelihood for GaussianLik {
fn log_likelihood(&self, etas: &[&Tensor]) -> Result<Tensor> {
let eta = etas[0];
let diff_sq = eta.broadcast_sub(&self.y)?.sqr()?;
let log_prob = (diff_sq * (-0.5))?;
log_prob.sum(2)?.sum(1)
}
}
let likelihood = GaussianLik { y };
let loss = local_reparam_loss(&model, &likelihood, 10, 1.0)?;
assert!(loss.dims().is_empty());
let eta_mean = model.eta_mean()?;
assert_eq!(eta_mean.dims(), &[n, k]);
let coef_mean = model.coef_mean()?;
assert_eq!(coef_mean.dims(), &[p, k]);
Ok(())
}
}