Skip to main content

pmcore/estimation/nonparametric/
ipm.rs

1use crate::estimation::nonparametric::{Psi, Weights};
2use anyhow::bail;
3use faer::linalg::triangular_solve::solve_lower_triangular_in_place;
4use faer::linalg::triangular_solve::solve_upper_triangular_in_place;
5use faer::{Col, Mat, Row};
6use rayon::prelude::*;
7
8/// Applies Burke's Interior Point Method (IPM) to solve a convex optimization problem.
9pub fn burke(psi: &Psi) -> anyhow::Result<(Weights, f64)> {
10    let log_scale = psi.log_scale();
11    let mut psi = psi.matrix().to_owned();
12
13    psi.row_iter_mut().try_for_each(|row| {
14        row.iter_mut().try_for_each(|x| {
15            if !x.is_finite() {
16                bail!("Input matrix must have finite entries")
17            } else {
18                *x = x.abs();
19                Ok(())
20            }
21        })
22    })?;
23
24    let (n_sub, n_point) = psi.shape();
25    let ecol: Col<f64> = Col::from_fn(n_point, |_| 1.0);
26    let erow: Row<f64> = Row::from_fn(n_sub, |_| 1.0);
27    let mut plam: Col<f64> = &psi * &ecol;
28    let eps: f64 = 1e-8;
29    let mut sig: f64 = 0.0;
30    let mut lam = ecol.clone();
31    let mut w: Col<f64> = Col::from_fn(plam.nrows(), |i| 1.0 / plam.get(i));
32    let mut ptw: Col<f64> = psi.transpose() * &w;
33
34    let ptw_max = ptw.iter().fold(f64::NEG_INFINITY, |acc, &x| x.max(acc));
35    let shrink = 2.0 * ptw_max;
36    lam *= shrink;
37    plam *= shrink;
38    w /= shrink;
39    ptw /= shrink;
40
41    let mut y: Col<f64> = &ecol - &ptw;
42    let mut r: Col<f64> = Col::from_fn(n_sub, |i| erow.get(i) - w.get(i) * plam.get(i));
43    let mut norm_r: f64 = r.iter().fold(0.0, |max, &val| max.max(val.abs()));
44    let sum_log_plam: f64 = plam.iter().map(|x| x.ln()).sum();
45    let sum_log_w: f64 = w.iter().map(|x| x.ln()).sum();
46    let mut gap: f64 = (sum_log_w + sum_log_plam).abs() / (1.0 + sum_log_plam);
47    let mut mu = lam.transpose() * &y / n_point as f64;
48
49    let mut psi_inner: Mat<f64> = Mat::zeros(psi.nrows(), psi.ncols());
50    let n_threads = faer::get_global_parallelism().degree();
51    let rows = psi.nrows();
52    let mut output: Vec<Mat<f64>> = (0..n_threads).map(|_| Mat::zeros(rows, rows)).collect();
53    let mut h: Mat<f64> = Mat::zeros(rows, rows);
54
55    while mu > eps || norm_r > eps || gap > eps {
56        let smu = sig * mu;
57        let inner = Col::from_fn(lam.nrows(), |i| lam.get(i) / y.get(i));
58        let w_plam = Col::from_fn(plam.nrows(), |i| plam.get(i) / w.get(i));
59
60        if psi.ncols() > n_threads * 128 {
61            psi_inner
62                .par_col_partition_mut(n_threads)
63                .zip(psi.par_col_partition(n_threads))
64                .zip(inner.par_partition(n_threads))
65                .zip(output.par_iter_mut())
66                .for_each(|(((mut psi_inner, psi), inner), output)| {
67                    psi_inner
68                        .as_mut()
69                        .col_iter_mut()
70                        .zip(psi.col_iter())
71                        .zip(inner.iter())
72                        .for_each(|((col, psi_col), inner_val)| {
73                            col.iter_mut().zip(psi_col.iter()).for_each(|(x, psi_val)| {
74                                *x = psi_val * inner_val;
75                            });
76                        });
77                    faer::linalg::matmul::triangular::matmul(
78                        output.as_mut(),
79                        faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
80                        faer::Accum::Replace,
81                        &psi_inner,
82                        faer::linalg::matmul::triangular::BlockStructure::Rectangular,
83                        psi.transpose(),
84                        faer::linalg::matmul::triangular::BlockStructure::Rectangular,
85                        1.0,
86                        faer::Par::Seq,
87                    );
88                });
89
90            let mut first_iter = true;
91            for output in &output {
92                if first_iter {
93                    h.copy_from(output);
94                    first_iter = false;
95                } else {
96                    h += output;
97                }
98            }
99        } else {
100            psi_inner
101                .as_mut()
102                .col_iter_mut()
103                .zip(psi.col_iter())
104                .zip(inner.iter())
105                .for_each(|((col, psi_col), inner_val)| {
106                    col.iter_mut().zip(psi_col.iter()).for_each(|(x, psi_val)| {
107                        *x = psi_val * inner_val;
108                    });
109                });
110            faer::linalg::matmul::triangular::matmul(
111                h.as_mut(),
112                faer::linalg::matmul::triangular::BlockStructure::TriangularLower,
113                faer::Accum::Replace,
114                &psi_inner,
115                faer::linalg::matmul::triangular::BlockStructure::Rectangular,
116                psi.transpose(),
117                faer::linalg::matmul::triangular::BlockStructure::Rectangular,
118                1.0,
119                faer::Par::Seq,
120            );
121        }
122
123        for i in 0..h.nrows() {
124            h[(i, i)] += w_plam[i];
125        }
126
127        let uph = match h.llt(faer::Side::Lower) {
128            Ok(llt) => llt,
129            Err(_) => {
130                bail!("Error during Cholesky decomposition. The matrix might not be positive definite. This is usually due to model misspecification or numerical issues.")
131            }
132        };
133        let uph = uph.L().transpose().to_owned();
134
135        let smuyinv: Col<f64> = Col::from_fn(ecol.nrows(), |i| smu * (ecol[i] / y[i]));
136        let psi_dot_muyinv: Col<f64> = &psi * &smuyinv;
137        let rhsdw: Row<f64> = Row::from_fn(erow.ncols(), |i| erow[i] / w[i] - psi_dot_muyinv[i]);
138        let mut dw = Mat::from_fn(rhsdw.ncols(), 1, |i, _j| *rhsdw.get(i));
139
140        solve_lower_triangular_in_place(uph.transpose().as_ref(), dw.as_mut(), faer::Par::rayon(0));
141        solve_upper_triangular_in_place(uph.as_ref(), dw.as_mut(), faer::Par::rayon(0));
142
143        let dw = dw.col(0);
144        let dy = -(psi.transpose() * dw);
145        let inner_times_dy = Col::from_fn(ecol.nrows(), |i| inner[i] * dy[i]);
146        let dlam: Row<f64> =
147            Row::from_fn(ecol.nrows(), |i| smuyinv[i] - lam[i] - inner_times_dy[i]);
148
149        let ratio_dlam_lam = Row::from_fn(lam.nrows(), |i| dlam[i] / lam[i]);
150        let min_ratio_dlam = ratio_dlam_lam.iter().cloned().fold(f64::INFINITY, f64::min);
151        let mut alfpri: f64 = -1.0 / min_ratio_dlam.min(-0.5);
152        alfpri = (0.99995 * alfpri).min(1.0);
153
154        let ratio_dy_y = Row::from_fn(y.nrows(), |i| dy[i] / y[i]);
155        let min_ratio_dy = ratio_dy_y.iter().cloned().fold(f64::INFINITY, f64::min);
156        let ratio_dw_w = Row::from_fn(dw.nrows(), |i| dw[i] / w[i]);
157        let min_ratio_dw = ratio_dw_w.iter().cloned().fold(f64::INFINITY, f64::min);
158        let mut alfdual = -1.0 / min_ratio_dy.min(-0.5);
159        alfdual = alfdual.min(-1.0 / min_ratio_dw.min(-0.5));
160        alfdual = (0.99995 * alfdual).min(1.0);
161
162        lam += alfpri * dlam.transpose();
163        w += alfdual * dw;
164        y += alfdual * &dy;
165
166        mu = lam.transpose() * &y / n_point as f64;
167        plam = &psi * &lam;
168        r = Col::from_fn(n_sub, |i| erow.get(i) - w.get(i) * plam.get(i));
169        ptw -= alfdual * dy;
170
171        norm_r = r.norm_max();
172        let sum_log_plam: f64 = plam.iter().map(|x| x.ln()).sum();
173        let sum_log_w: f64 = w.iter().map(|x| x.ln()).sum();
174        gap = (sum_log_w + sum_log_plam).abs() / (1.0 + sum_log_plam);
175
176        if mu < eps && norm_r > eps {
177            sig = 1.0;
178        } else {
179            let candidate1 = (1.0 - alfpri).powi(2);
180            let candidate2 = (1.0 - alfdual).powi(2);
181            let candidate3 = (norm_r - mu) / (norm_r + 100.0 * mu);
182            sig = candidate1.max(candidate2).max(candidate3).min(0.3);
183        }
184    }
185
186    lam /= n_sub as f64;
187    let obj = (psi * &lam).iter().map(|x| x.ln()).sum::<f64>() + log_scale;
188    let lam_sum: f64 = lam.iter().sum();
189    lam = &lam / lam_sum;
190
191    Ok((lam.into(), obj))
192}
193
194#[cfg(test)]
195mod tests {
196    use super::*;
197    use approx::assert_relative_eq;
198    use faer::Mat;
199
200    #[test]
201    fn test_burke_identity() {
202        let n = 100;
203        let mat = Mat::identity(n, n);
204        let psi = Psi::from(mat);
205        let (lam, _) = burke(&psi).unwrap();
206
207        let expected = 1.0 / n as f64;
208        for i in 0..n {
209            assert_relative_eq!(lam[i], expected, epsilon = 1e-10);
210        }
211        assert_relative_eq!(lam.iter().sum::<f64>(), 1.0, epsilon = 1e-10);
212    }
213
214    #[test]
215    fn test_burke_uniform_square() {
216        let n_sub = 10;
217        let n_point = 10;
218        let mat = Mat::from_fn(n_sub, n_point, |_, _| 1.0);
219        let psi = Psi::from(mat);
220        let (lam, _) = burke(&psi).unwrap();
221
222        assert_relative_eq!(lam.iter().sum::<f64>(), 1.0, epsilon = 1e-10);
223        let expected = 1.0 / n_point as f64;
224        for i in 0..n_point {
225            assert_relative_eq!(lam[i], expected, epsilon = 1e-10);
226        }
227    }
228
229    #[test]
230    fn test_burke_restores_log_likelihood_scale() -> anyhow::Result<()> {
231        let log_likelihoods =
232            ndarray::Array2::from_shape_vec((2, 2), vec![-1000.0, -1001.0, -2.0, -4.0])?;
233        let psi = Psi::from_log_likelihoods(log_likelihoods)?;
234
235        let (weights, objective) = burke(&psi)?;
236        let expected_objective = -1000.0 + (weights[0] + (-1.0_f64).exp() * weights[1]).ln() - 2.0
237            + (weights[0] + (-2.0_f64).exp() * weights[1]).ln();
238
239        assert_relative_eq!(objective, expected_objective, epsilon = 1e-8);
240        Ok(())
241    }
242
243    #[test]
244    fn test_burke_with_non_finite_values() {
245        let n_sub = 10;
246        let n_point = 10;
247        let mat = Mat::from_fn(n_sub, n_point, |i, j| {
248            if i == 0 && j == 0 {
249                f64::NAN
250            } else {
251                1.0
252            }
253        });
254        let psi = Psi::from(mat);
255        assert!(burke(&psi).is_err());
256    }
257}