use candle_core::{Result, Tensor};
use super::regression_linear::RegressionSGVB;
use super::traits::{
AnalyticalKL, BlackBoxLikelihood, LocalReparamSample, Prior, VariationalDistribution,
};
pub struct CompositeModel<V, P> {
pub modules: Vec<RegressionSGVB<V, P>>,
}
impl<V: VariationalDistribution, P: Prior> CompositeModel<V, P> {
pub fn new(modules: Vec<RegressionSGVB<V, P>>) -> Self {
assert!(
!modules.is_empty(),
"CompositeModel requires at least one module"
);
Self { modules }
}
pub fn num_modules(&self) -> usize {
self.modules.len()
}
}
pub fn composite_local_reparam_loss<V, P, L>(
model: &CompositeModel<V, P>,
likelihood: &L,
num_samples: usize,
kl_weight: f64,
) -> Result<Tensor>
where
V: VariationalDistribution,
P: Prior + AnalyticalKL,
L: BlackBoxLikelihood,
{
let samples: Vec<LocalReparamSample> = model
.modules
.iter()
.map(|m| m.forward(num_samples))
.collect::<Result<_>>()?;
samples_local_reparam_loss(&samples, likelihood, kl_weight)
}
pub fn samples_local_reparam_loss<L>(
samples: &[LocalReparamSample],
likelihood: &L,
kl_weight: f64,
) -> Result<Tensor>
where
L: BlackBoxLikelihood,
{
let etas: Vec<&Tensor> = samples.iter().map(|s| &s.eta).collect();
let llik = likelihood.log_likelihood(&etas)?;
let llik = if llik.rank() > 1 { llik.sum(1)? } else { llik };
let mut total_kl = samples[0].kl.clone();
for s in &samples[1..] {
total_kl = (&total_kl + &s.kl)?;
}
let elbo = (llik.mean(0)? - (total_kl * kl_weight)?)?;
elbo.neg()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::candle::sgvb::{GaussianPrior, GaussianRegressionSGVB, SGVBConfig};
use candle_core::{DType, Device, Tensor};
use candle_nn::{Optimizer, VarBuilder, VarMap};
struct TestGaussianLikelihood {
y: Tensor,
}
impl TestGaussianLikelihood {
fn new(y: Tensor) -> Self {
Self { y }
}
}
impl BlackBoxLikelihood for TestGaussianLikelihood {
fn log_likelihood(&self, etas: &[&Tensor]) -> Result<Tensor> {
assert!(etas.len() >= 2, "TestGaussianLikelihood requires 2 etas");
let mu = etas[0]; let log_var_raw = etas[1];
let log_var = log_var_raw.clamp(-10.0, 10.0)?;
let ln_2pi = (2.0 * std::f64::consts::PI).ln();
let diff = mu.broadcast_sub(&self.y)?;
let diff_sq = diff.sqr()?;
let var = log_var.exp()?;
let scaled_diff_sq = (diff_sq / &var)?;
let log_prob = ((scaled_diff_sq + &log_var)? + ln_2pi)? * (-0.5);
log_prob?.sum(2)?.sum(1)
}
}
#[test]
fn test_composite_model_construction() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F32;
let n = 50;
let p_mean = 10;
let p_var = 5;
let k = 2;
let x_mean = Tensor::randn(0f32, 1f32, (n, p_mean), &device)?;
let x_var = Tensor::randn(0f32, 1f32, (n, p_var), &device)?;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let prior_mean = GaussianPrior::new(vb.pp("prior_mean"), 1.0)?;
let prior_var = GaussianPrior::new(vb.pp("prior_var"), 1.0)?;
let config = SGVBConfig::default();
let model_mean =
GaussianRegressionSGVB::new(vb.pp("mean"), x_mean, k, prior_mean, config.clone())?;
let model_var =
GaussianRegressionSGVB::new(vb.pp("var"), x_var, k, prior_var, config.clone())?;
let composite = CompositeModel::new(vec![model_mean, model_var]);
assert_eq!(composite.num_modules(), 2);
let samples: Vec<LocalReparamSample> = composite
.modules
.iter()
.map(|m| m.forward(10))
.collect::<Result<_>>()?;
assert_eq!(samples.len(), 2);
assert_eq!(samples[0].eta.dims(), &[10, n, k]);
assert_eq!(samples[1].eta.dims(), &[10, n, k]);
Ok(())
}
#[test]
fn test_heteroscedastic_regression() -> Result<()> {
let device = Device::Cpu;
let dtype = DType::F32;
let n = 100;
let p_mean = 5;
let p_var = 3;
let k = 1;
let x_mean = Tensor::randn(0f32, 1f32, (n, p_mean), &device)?;
let x_var = Tensor::randn(0f32, 1f32, (n, p_var), &device)?;
let true_mean = (x_mean.narrow(1, 0, 1)? * 2.0)?;
let true_log_var = (x_var.narrow(1, 0, 1)? * 0.5)?;
let true_std = (true_log_var.clone() / 2.0)?.exp()?;
let noise = Tensor::randn(0f32, 1f32, (n, k), &device)?;
let y = (true_mean + (noise * true_std)?)?;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, dtype, &device);
let prior_mean = GaussianPrior::new(vb.pp("prior_mean"), 1.0)?;
let prior_var = GaussianPrior::new(vb.pp("prior_var"), 1.0)?;
let config = SGVBConfig::new(30);
let model_mean =
GaussianRegressionSGVB::new(vb.pp("mean"), x_mean, k, prior_mean, config.clone())?;
let model_var =
GaussianRegressionSGVB::new(vb.pp("var"), x_var, k, prior_var, config.clone())?;
let composite = CompositeModel::new(vec![model_mean, model_var]);
let likelihood = TestGaussianLikelihood::new(y);
let mut optimizer = candle_nn::AdamW::new_lr(varmap.all_vars(), 0.01)?;
for i in 0..300 {
let loss = composite_local_reparam_loss(&composite, &likelihood, 30, 1.0)?;
optimizer.backward_step(&loss)?;
if i % 100 == 0 {
let loss_val: f32 = loss.to_scalar()?;
let elbo_val = -loss_val;
println!(
"iter {:4}: loss = {:10.4}, ELBO = {:10.4}",
i, loss_val, elbo_val
);
}
}
let mean_coef = composite.modules[0].coef_mean()?;
let mean_coef_0: f32 = mean_coef.get(0)?.get(0)?.to_scalar()?;
println!("\nMean coef[0] (true=2.0): {:.4}", mean_coef_0);
assert!(
mean_coef_0 > 0.5,
"Mean coef[0] should be > 0.5, got {}",
mean_coef_0
);
Ok(())
}
}