use crate::math::digamma_tensor;
use crate::math::ln_gamma_tensor;
use crate::math::logsumexp;
use crate::math::logsumexp_mat;
use crate::math::logsumexp_pair;
use burn_tensor::backend::Backend;
use burn_tensor::Tensor;
type E = Box<dyn std::error::Error>;
const COLS: usize = 0;
const ROWS: usize = 1;
pub fn compute_norm<B: Backend>(
gamma_z: Tensor::<B, 2>,
dl_dphi: Tensor::<B, 2>,
) -> Tensor::<B, 1> {
let dl_dphi_pos = dl_dphi.clone().clamp_min(f32::MIN_POSITIVE);
let dl_dphi_neg = dl_dphi.clone().neg().clamp_min(f32::MIN_POSITIVE);
let temp1_pos = gamma_z.clone().add(dl_dphi_pos.log());
let temp1_neg = gamma_z.clone().add(dl_dphi_neg.log());
let log_colsums_pos = logsumexp(temp1_pos.clone(), COLS);
let log_colsums_neg = logsumexp(temp1_neg.clone(), COLS);
let colsums = log_colsums_pos.exp().sub(log_colsums_neg.exp());
let temp2 = dl_dphi.clone().sub(colsums);
let temp2_pos = temp2.clone().clamp_min(f32::MIN_POSITIVE);
let temp2_neg = temp2.clone().neg().clamp_min(f32::MIN_POSITIVE);
let temp3_pos_pos = temp2_pos.clone().log().add(temp1_pos.clone());
let temp3_pos_neg = temp2_pos.log().add(temp1_neg.clone());
let temp3_neg_neg = temp2_neg.clone().log().add(temp1_neg);
let temp3_neg_pos = temp2_neg.log().add(temp1_pos);
let sum_pos_pos = logsumexp_mat(temp3_pos_pos);
let sum_pos_neg = logsumexp_mat(temp3_pos_neg);
let sum_neg_neg = logsumexp_mat(temp3_neg_neg);
let sum_neg_pos = logsumexp_mat(temp3_neg_pos);
let norm_pos = logsumexp_pair(sum_pos_pos, sum_neg_neg);
let norm_neg = logsumexp_pair(sum_pos_neg, sum_neg_pos);
norm_pos.exp().sub(norm_neg.exp()).clamp_min(f32::MIN_POSITIVE)
}
pub fn mixt_negnatgrad<B: Backend>(
logl: Tensor::<B, 2>,
gamma_z: Tensor::<B, 2>,
log_n_k: Tensor::<B, 1>,
oldnorm_t: Tensor::<B, 1>,
oldstep: Tensor::<B, 2>,
) -> (Tensor::<B, 2>, Tensor::<B, 1>) {
const F32_EXP_OVERFLOW: f32 = 80.0;
let digamma_input = log_n_k.clone().mask_where(
log_n_k.clone().greater_equal_elem(F32_EXP_OVERFLOW),
log_n_k.zeros_like().add_scalar(1.0)
);
let digamma_input_exp = digamma_input.exp();
let digamma_vals = digamma_tensor(digamma_input_exp);
let digamma_n_k = log_n_k.clone().mask_where(log_n_k.clone().lower_elem(F32_EXP_OVERFLOW), digamma_vals).sub_scalar(1.0);
let gradient = logl.add(digamma_n_k.unsqueeze_dim(ROWS)).sub(gamma_z.clone());
let newnorm_t = compute_norm(gamma_z, gradient.clone());
let beta_fr_t = newnorm_t.clone().div(oldnorm_t);
let beta_fr_t = beta_fr_t.unsqueeze_dim(COLS);
let oldstep = oldstep.mul(beta_fr_t);
let step = gradient.add(oldstep.clone());
(step, newnorm_t)
}
pub fn update_n_k<B: Backend>(
gamma_z: Tensor::<B, 2>,
log_counts: Tensor::<B, 1>,
alpha0: Tensor::<B, 1>,
) -> Tensor::<B, 1> {
let tmp_counts = log_counts.unsqueeze_dim(COLS);
let temp = gamma_z.add(tmp_counts);
let log_n_k = logsumexp(temp, ROWS).squeeze_dim(ROWS);
logsumexp_pair(log_n_k, alpha0.log())
}
pub fn elbo_rcg_mat<B: Backend>(
logl: Tensor::<B, 2>,
gamma_z: Tensor::<B, 2>,
log_counts: Tensor::<B, 1>,
log_n_k: Tensor::<B, 1>,
) -> Tensor::<B, 1> {
let log_n_k_exp = log_n_k.exp();
let lgamma_n_k = ln_gamma_tensor(log_n_k_exp);
let logl_adj = logl.sub(gamma_z.clone());
let gamma_z_adj = gamma_z.add(log_counts.unsqueeze_dim(COLS));
let logl_adj_pos = logl_adj.clone().clamp_min(f32::MIN_POSITIVE);
let logl_adj_neg = logl_adj.neg().clamp_min(f32::MIN_POSITIVE);
let log_term_pos = gamma_z_adj.clone().add(logl_adj_pos.log());
let log_term_neg = gamma_z_adj.add(logl_adj_neg.log());
let colsums_pos_log = logsumexp(log_term_pos, ROWS);
let colsums_neg_log = logsumexp(log_term_neg, ROWS);
let colsums = colsums_pos_log.exp().sub(colsums_neg_log.exp());
let bounds = colsums.add(lgamma_n_k.unsqueeze_dim(ROWS));
bounds.sum()
}
pub fn rcg_optl_mat<B: Backend>(
logl: Tensor::<B, 2>,
log_counts: Tensor::<B, 1>,
alpha0: Tensor::<B, 1>,
tolerance: f64,
max_iters: usize,
) -> Result<Tensor::<B, 2>, E> {
let gamma_z_init = 1_f32.ln() - (logl.dims()[COLS] as f32).ln();
let mut gamma_z = Tensor::<B, 2>::full(logl.shape(), gamma_z_init, &logl.device());
let mut log_n_k = update_n_k(gamma_z.clone(), log_counts.clone(), alpha0.clone());
let mut oldstep = logl.zeros_like();
let mut step;
let mut oldbound = Tensor::<B, 1>::from_data([-f32::INFINITY], &logl.device());
let mut oldnorm = Tensor::<B, 1>::from_data([f32::INFINITY], &logl.device());
let mut iter = 0;
while iter < max_iters {
(step, oldnorm) = mixt_negnatgrad(logl.clone(), gamma_z.clone(), log_n_k.clone(), oldnorm.clone(), oldstep.clone());
let gamma_z_stepped = gamma_z.clone().add(step.clone());
let oldm = logsumexp(gamma_z_stepped.clone(), COLS);
let gamma_z_new = gamma_z_stepped.sub(oldm.clone());
log_n_k = update_n_k(gamma_z_new.clone(), log_counts.clone(), alpha0.clone());
let bound = elbo_rcg_mat(logl.clone(), gamma_z_new.clone(), log_counts.clone(), log_n_k.clone());
let bound_f: f32 = bound.clone().into_data().iter().next().unwrap();
let oldbound_f: f32 = oldbound.clone().into_data().iter().next().unwrap();
let bounds_converged = approx::relative_eq!(bound_f, oldbound_f, epsilon = tolerance as f32);
let bounds_finite = bound_f < f32::INFINITY && oldbound_f < f32::INFINITY;
if bounds_finite && bounds_converged {
let lse = logsumexp(gamma_z_new.clone(), COLS);
gamma_z = gamma_z_new.sub(lse);
break;
} else if bound_f < oldbound_f && (bound_f - oldbound_f).abs() > tolerance as f32 {
let gamma_z_reverted = gamma_z.clone();
log_n_k = update_n_k(gamma_z_reverted.clone(), log_counts.clone(), alpha0.clone());
oldbound = bound;
oldstep = logl.zeros_like();
oldnorm = oldnorm.zeros_like().add_scalar(f32::INFINITY);
} else {
oldstep = step;
oldbound = bound;
gamma_z = gamma_z_new;
}
iter += 1;
}
let lse = logsumexp(gamma_z.clone(), COLS);
let gamma_z_fin = gamma_z.sub(lse);
Ok(gamma_z_fin)
}
#[cfg(test)]
mod tests {
use assert_approx_eq::assert_approx_eq;
#[test]
fn mixt_negnatgrad() {
use burn::backend::ndarray::NdArray;
use burn_tensor::Tensor;
use super::mixt_negnatgrad;
let device = Default::default();
type Backend = NdArray<f32>;
let gamma_z = Tensor::<Backend, 2>::from_data(
[
[ -0.861124, -0.824187, -0.737067, -0.830991, -0.792902, -0.702885, -0.76075, -0.719832, -0.622649, -0.742541 ],
[ -1.01295, -0.976009, -0.888889, -0.982813, -0.944725, -0.854708, -0.912572, -0.871654, -0.774472, -1.26242 ],
[ -2.33926, -2.30233, -2.21521, -2.67719, -2.6391, -2.54908, -6.91527, -6.87435, -6.77717, -2.22068 ],
[ -2.13905, -2.47017, -6.69137, -2.10891, -2.43888, -6.65719, -2.03867, -2.36581, -6.57695, -2.02046 ],
],
&device,
);
let log_n_k = Tensor::<Backend, 1>::from_data(
[
4857.97, 3905.03, 701.053, 903.946,
],
&device,
).log();
let logl = Tensor::<Backend, 2>::from_data(
[
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.371713 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.371713, -0.371713, -0.371713, -4.60517, -4.60517, -4.60517, -0.0100503 ],
[ -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503 ],
],
&device,
);
let oldstep = Tensor::<Backend, 2>::from_data(
[
[20.177214, 20.177094, 20.176817, 20.177105, 20.176983, 20.176708, 20.176857, 20.176735, 20.17646, 20.17661],
[20.171112, 20.17099, 20.170715, 20.171001, 20.17088, 20.170609, 20.170753, 20.170631, 20.170357, 20.170507],
[20.177217, 20.177097, 20.176823, 20.17711, 20.176989, 20.176714, 20.176863, 20.17674, 20.176468, 20.176615],
[20.177156, 20.177038, 20.176762, 20.177048, 20.176928, 20.176651, 20.1768, 20.17668, 20.176403, 20.176554],
],
&device,
);
let oldnorm = Tensor::<Backend, 1>::from_data(
[
523.7,
],
&device,
);
let expected = Tensor::<Backend, 2>::from_data(
[
[8.636083, 8.599144, 8.51202, 8.605948, 8.567857, 8.477837, 8.535704, 8.494783, 8.397596, 8.517491],
[8.56944, 8.532496, 8.445373, 8.539301, 8.50121, 8.41119, 8.469056, 8.428136, 8.33095, 8.457237],
[8.177816, 8.140885, 8.053761, 8.154082, 8.115991, 8.025967, 8.158702, 8.117781, 8.0205965, 8.059228],
[8.231952, 8.201407, 8.189146, 8.20181, 8.170115, 8.154964, 8.131567, 8.097042, 8.074721, 8.113353],
],
&device,
);
let expected_norm = Tensor::<Backend, 1>::from_data([7.7018056], &device);
let (got, newnorm) = mixt_negnatgrad::<Backend>(logl, gamma_z, log_n_k, oldnorm, oldstep);
let got_data = got.into_data();
let newnorm_data = newnorm.into_data();
let expected_data = expected.into_data();
let expected_norm_data = expected_norm.into_data();
got_data.iter().zip(expected_data.iter()).for_each(|(x, y): (f32, f32)| { assert_approx_eq!(x, y, 1e-5) });
newnorm_data.iter().zip(expected_norm_data.iter()).for_each(|(x, y): (f32, f32)| { assert_approx_eq!(x, y, 1e-5) });
}
#[test]
fn compute_norm() {
use burn::backend::ndarray::NdArray;
use burn_tensor::Tensor;
use super::compute_norm;
let device = Default::default();
type Backend = NdArray<f32>;
let gamma_z = Tensor::<Backend, 2>::from_data(
[
[ -0.861124, -0.824187, -0.737067, -0.830991, -0.792902, -0.702885, -0.76075, -0.719832, -0.622649, -0.742541 ],
[ -1.01295, -0.976009, -0.888889, -0.982813, -0.944725, -0.854708, -0.912572, -0.871654, -0.774472, -1.26242 ],
[ -2.33926, -2.30233, -2.21521, -2.67719, -2.6391, -2.54908, -6.91527, -6.87435, -6.77717, -2.22068 ],
[ -2.13905, -2.47017, -6.69137, -2.10891, -2.43888, -6.65719, -2.03867, -2.36581, -6.57695, -2.02046 ],
],
&device,
);
let dl_dphi = Tensor::<Backend, 2>::from_data(
[
[ 8.33935, 8.30241, 8.21529, 8.30921, 8.27113, 8.18111, 8.23897, 8.19806, 8.10087, 8.22076 ],
[ 8.27279, 8.23585, 8.14873, 8.24266, 8.20457, 8.11455, 8.17241, 8.1315, 8.03431, 8.1606 ],
[ 7.88108, 7.84415, 7.75703, 7.85735, 7.81926, 7.72924, 7.86197, 7.82105, 7.72387, 7.7625 ],
[ 7.93521, 7.90467, 7.89242, 7.90508, 7.87339, 7.85823, 7.83484, 7.80032, 7.778, 7.81663 ],
],
&device,
);
let expected: f64 = 7.701816558837891;
let got: f64 = compute_norm::<Backend>(gamma_z, dl_dphi).into_data().iter().next().unwrap();
assert_approx_eq!(expected, got, 1e-4);
}
#[test]
fn update_n_k() {
use burn::backend::ndarray::NdArray;
use burn_tensor::Tensor;
use super::update_n_k;
let device = Default::default();
type Backend = NdArray<f32>;
let gamma_z = Tensor::<Backend, 2>::from_data(
[
[ -0.681538, -0.662494, -0.617806, -0.667704, -0.648392, -0.603055, -0.635526, -0.615577, -0.568692, -0.557316 ],
[ -0.951042, -0.931998, -0.887311, -0.937208, -0.917896, -0.872559, -0.905031, -0.885081, -0.838196, -1.18688 ],
[ -3.09143, -3.07238, -3.0277, -3.43766, -3.41835, -3.37301, -7.62022, -7.60027, -7.55338, -2.96721 ],
[ -2.77441, -3.11543, -7.28548, -2.76058, -3.10133, -7.27073, -2.7284, -3.06852, -7.23637, -2.65019 ],
],
&device,
);
let log_counts = Tensor::<Backend, 1>::from_data(
[
7.681099, 7.04316, 6.849066, 5.278115, 5.164786, 5.062595, 6.947937, 6.863803, 7.277248, 7.666222
],
&device,
);
let alpha0 = Tensor::<Backend, 1>::from_data(
[
1.0, 1.0, 1.0, 1.0
],
&device,
);
let expected = Tensor::<Backend, 1>::from_data(
[
5585.01, 3983.44, 327.192, 472.355
],
&device,
).log();
let got = update_n_k::<Backend>(gamma_z, log_counts, alpha0);
let got_data = got.into_data();
let expected_data = expected.into_data();
got_data.iter().zip(expected_data.iter()).for_each(|(x, y): (f32, f32)| { assert_approx_eq!(x, y, 1e-2) });
}
#[test]
fn elbo_rcg_mat() {
use burn::backend::ndarray::NdArray;
use burn_tensor::Tensor;
use super::elbo_rcg_mat;
let device = Default::default();
type Backend = NdArray<f32>;
let logl = Tensor::<Backend, 2>::from_data(
[
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.371713 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.371713, -0.371713, -0.371713, -4.60517, -4.60517, -4.60517, -0.0100503 ],
[ -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503 ],
],
&device,
);
let log_counts = Tensor::<Backend, 1>::from_data(
[
7.681099, 7.04316, 6.849066, 5.278115, 5.164786, 5.062595, 6.947937, 6.863803, 7.277248, 7.666222
],
&device,
);
let gamma_z = Tensor::<Backend, 2>::from_data(
[
[ -0.681538, -0.662494, -0.617806, -0.667704, -0.648392, -0.603055, -0.635526, -0.615577, -0.568692, -0.557316 ],
[ -0.951042, -0.931998, -0.887311, -0.937208, -0.917896, -0.872559, -0.905031, -0.885081, -0.838196, -1.18688 ],
[ -3.09143, -3.07238, -3.0277, -3.43766, -3.41835, -3.37301, -7.62022, -7.60027, -7.55338, -2.96721 ],
[ -2.77441, -3.11543, -7.28548, -2.76058, -3.10133, -7.27073, -2.7284, -3.06852, -7.23637, -2.65019 ],
],
&device,
);
let log_n_k = Tensor::<Backend, 1>::from_data(
[
5585.01, 3983.44, 327.192, 472.355
],
&device,
).log();
let expected = (-699.064_f64 + 85494_f64).ln();
let got: f64 = elbo_rcg_mat::<Backend>(logl, gamma_z, log_counts, log_n_k).into_data().iter().next().unwrap();
assert_approx_eq!(expected, got, 1e-1);
}
#[test]
fn rcg_optl_mat() {
use burn::backend::ndarray::NdArray;
use burn_tensor::Tensor;
use super::rcg_optl_mat;
let device = Default::default();
type Backend = NdArray<f32>;
let logl = Tensor::<Backend, 2>::from_data(
[
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.0100503, -0.371713 ],
[ -0.0100503, -0.0100503, -0.0100503, -0.371713, -0.371713, -0.371713, -4.60517, -4.60517, -4.60517, -0.0100503 ],
[ -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503, -0.371713, -4.60517, -0.0100503 ],
],
&device,
);
let log_counts = Tensor::<Backend, 1>::from_data(
[
7.681099, 7.04316, 6.849066, 5.278115, 5.164786, 5.062595, 6.947937, 6.863803, 7.277248, 7.666222
],
&device,
);
let alpha0 = Tensor::<Backend, 1>::from_data(
[
1.0, 1.0, 1.0, 1.0
],
&device,
);
let expected = Tensor::<Backend, 2>::from_data(
[
[ -0.0010899, -0.00104044, -0.000928571, -0.00104519, -0.000995734, -0.000883857, -0.000944069, -0.000894604, -0.000782716, -0.000853449 ],
[ -7.15745, -7.1574, -7.15729, -7.15741, -7.15736, -7.15725, -7.15731, -7.15726, -7.15715, -7.51888 ],
[ -8.82298, -8.82293, -8.82282, -9.1846, -9.18455, -9.18444, -13.418, -13.4179, -13.4178, -8.82274 ],
[ -8.72199, -9.0836, -13.3169, -8.72195, -9.08356, -13.3169, -8.72184, -9.08346, -13.3168, -8.72175 ],
],
&device,
);
let got = rcg_optl_mat::<Backend>(logl, log_counts, alpha0, 1e-7_f64, 100_usize).unwrap();
let got_data = got.into_data();
let expected_data = expected.into_data();
got_data.iter().zip(expected_data.iter()).for_each(|(x, y): (f32, f32)| { assert_approx_eq!(x, y, 1_f32) });
}
}