pmcore/estimation/nonparametric/
ipm.rs1use 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
8pub 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}