1use std::collections::HashSet;
14
15use prospicio_core::{Error, Result};
16
17use crate::copula::target_ranks;
18use crate::predictive::{ComponentKey, KeyValue, PredictiveDistribution};
19use crate::provenance::Provenance;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum Pairing {
24 Independent,
27 SameSimulations,
31}
32
33impl PredictiveDistribution {
34 pub fn join(
74 parts: &[(&str, &PredictiveDistribution)],
75 dim: &str,
76 pairing: Pairing,
77 ) -> Result<Self> {
78 let (_, first) = parts
79 .first()
80 .ok_or_else(|| Error::Data("join needs at least one part".into()))?;
81 let n = first.n_sims();
82 let mut labels = HashSet::new();
83 let mut seeds = HashSet::new();
84 let mut dims: Vec<String> = vec![dim.to_string()];
85 for (label, pd) in parts {
86 if pd.n_sims() != n {
87 return Err(Error::Data(format!(
88 "part {label:?} has {} simulations, the first has {n}",
89 pd.n_sims()
90 )));
91 }
92 if !labels.insert(*label) {
93 return Err(Error::Data(format!("part label {label:?} is repeated")));
94 }
95 if pd.dims().iter().any(|d| d == dim) {
96 return Err(Error::Data(format!(
97 "part {label:?} already has a dimension {dim:?}"
98 )));
99 }
100 let prov = pd.provenance();
101 if let (Pairing::Independent, Some(seed), Some(scheme)) =
102 (pairing, prov.seed, prov.stream_scheme.as_ref())
103 {
104 if !seeds.insert((seed, scheme.clone())) {
105 return Err(Error::Data(format!(
106 "part {label:?} reuses seed {seed} with the same stream scheme as an \
107 earlier part, so their simulations would share random numbers; \
108 simulate it with another seed"
109 )));
110 }
111 }
112 for d in pd.dims() {
113 if !dims.contains(d) {
114 dims.push(d.clone());
115 }
116 }
117 }
118 let mut components: Vec<ComponentKey> = Vec::new();
119 for (label, pd) in parts {
120 for key in pd.components() {
121 let mut k = vec![KeyValue::from(*label)];
122 for d in &dims[1..] {
123 match pd.dims().iter().position(|x| x == d) {
124 Some(j) => k.push(key[j].clone()),
125 None => k.push(KeyValue::from("")),
126 }
127 }
128 components.push(k);
129 }
130 }
131 let width = components.len();
132 let mut draws = Vec::with_capacity(n * width);
133 for i in 0..n {
134 for (_, pd) in parts {
135 draws.extend_from_slice(pd.row(i).expect("in range"));
136 }
137 }
138 let mut provenance = Provenance::new("join").param("pairing", format!("{pairing:?}"));
139 for (label, pd) in parts {
140 let p = pd.provenance();
141 let seed = p.seed.map_or_else(|| "none".to_string(), |s| s.to_string());
142 provenance = provenance.param(
143 format!("part:{label}"),
144 format!("model {}, seed {seed}", p.model),
145 );
146 }
147 Self::from_draws(dims, components, draws, provenance)
148 }
149
150 pub fn reorder_groups(&self, dim: &str, correlation: &[f64], seed: u64) -> Result<Self> {
179 let d = self
180 .dims()
181 .iter()
182 .position(|x| x == dim)
183 .ok_or_else(|| Error::Data(format!("no dimension {dim:?}")))?;
184 let mut groups: Vec<KeyValue> = Vec::new();
185 let group_of: Vec<usize> = self
186 .components()
187 .iter()
188 .map(|k| match groups.iter().position(|g| *g == k[d]) {
189 Some(g) => g,
190 None => {
191 groups.push(k[d].clone());
192 groups.len() - 1
193 }
194 })
195 .collect();
196 let (n, m, g) = (self.n_sims(), self.n_components(), groups.len());
197 let ranks = target_ranks(n, g, correlation, seed)?;
198 let order: Vec<Vec<usize>> = (0..g)
200 .map(|gi| {
201 let totals: Vec<f64> = (0..n)
202 .map(|i| {
203 let row = self.row(i).expect("in range");
204 (0..m).filter(|&j| group_of[j] == gi).map(|j| row[j]).sum()
205 })
206 .collect();
207 let mut idx: Vec<usize> = (0..n).collect();
208 idx.sort_by(|&a, &b| totals[a].total_cmp(&totals[b]));
209 idx
210 })
211 .collect();
212 let mut draws = vec![0.0; n * m];
213 for i in 0..n {
214 for (j, &gi) in group_of.iter().enumerate() {
215 let source = order[gi][ranks[gi][i]];
216 draws[i * m + j] = self.row(source).expect("in range")[j];
217 }
218 }
219 let provenance = self.provenance().clone().param(
220 "reorder_groups",
221 format!("{dim}: {correlation:?}, seed {seed}"),
222 );
223 Self::from_draws(
224 self.dims().to_vec(),
225 self.components().to_vec(),
226 draws,
227 provenance,
228 )
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use super::*;
235 use crate::Empirical;
236
237 fn part(
238 dim: &str,
239 keys: &[i64],
240 n: usize,
241 seed: u64,
242 f: impl Fn(usize, usize) -> f64,
243 ) -> PredictiveDistribution {
244 let draws = (0..n)
245 .flat_map(|i| (0..keys.len()).map(move |j| (i, j)))
246 .map(|(i, j)| f(i, j))
247 .collect();
248 PredictiveDistribution::from_draws(
249 vec![dim.into()],
250 keys.iter().map(|k| vec![KeyValue::Int(*k)]).collect(),
251 draws,
252 Provenance::new("test").seed(seed, "chacha20/sim-index/v1"),
253 )
254 .unwrap()
255 }
256
257 #[test]
258 fn join_checks_shapes_labels_and_seeds() {
259 let a = part("origin", &[1, 2], 10, 1, |i, j| (i + j) as f64);
260 let b = part("lob", &[7], 10, 2, |i, _| i as f64);
261 let same_seed = part("lob", &[7], 10, 1, |i, _| i as f64);
262 let short = part("lob", &[7], 9, 3, |i, _| i as f64);
263 assert!(
264 PredictiveDistribution::join(
265 &[("a", &a), ("b", &same_seed)],
266 "risk",
267 Pairing::Independent
268 )
269 .is_err()
270 );
271 assert!(
272 PredictiveDistribution::join(
273 &[("a", &a), ("b", &same_seed)],
274 "risk",
275 Pairing::SameSimulations
276 )
277 .is_ok()
278 );
279 assert!(
280 PredictiveDistribution::join(&[("a", &a), ("b", &short)], "risk", Pairing::Independent)
281 .is_err()
282 );
283 assert!(
284 PredictiveDistribution::join(&[("a", &a), ("a", &b)], "risk", Pairing::Independent)
285 .is_err()
286 );
287 assert!(
288 PredictiveDistribution::join(&[("a", &a)], "origin", Pairing::Independent).is_err()
289 );
290 let j = PredictiveDistribution::join(&[("a", &a), ("b", &b)], "risk", Pairing::Independent)
291 .unwrap();
292 assert_eq!(
293 j.components()[2],
294 vec![KeyValue::from("b"), KeyValue::from(""), KeyValue::Int(7)]
295 );
296 assert_eq!(
297 j.marginal(&j.components()[1].clone()).unwrap().draws(),
298 a.marginal(&vec![KeyValue::Int(2)]).unwrap().draws()
299 );
300 }
301
302 #[test]
303 fn reorder_groups_keeps_marginals_and_sets_rank_correlation() {
304 let n = 4000;
305 let a = part("origin", &[1, 2], n, 1, |i, j| {
306 ((i * 7919 + j * 31) % n) as f64 + j as f64
307 });
308 let b = part("lob", &[7], n, 2, |i, _| ((i * 104_729) % n) as f64);
309 let j = PredictiveDistribution::join(&[("a", &a), ("b", &b)], "risk", Pairing::Independent)
310 .unwrap();
311 let r = j.reorder_groups("risk", &[1.0, 0.8, 0.8, 1.0], 3).unwrap();
312 for c in j.components() {
313 let mut x = j.marginal(c).unwrap().draws().to_vec();
314 let mut y = r.marginal(c).unwrap().draws().to_vec();
315 x.sort_by(f64::total_cmp);
316 y.sort_by(f64::total_cmp);
317 assert_eq!(x, y);
318 }
319 let rows_a: HashSet<(u64, u64)> = (0..n)
321 .map(|i| {
322 let row = j.row(i).unwrap();
323 (row[0].to_bits(), row[1].to_bits())
324 })
325 .collect();
326 assert!((0..n).all(|i| {
327 let row = r.row(i).unwrap();
328 rows_a.contains(&(row[0].to_bits(), row[1].to_bits()))
329 }));
330 let by = r.aggregate(&["risk"]).unwrap();
332 let (ta, tb): (Vec<f64>, Vec<f64>) = (0..n)
333 .map(|i| (by.row(i).unwrap()[0], by.row(i).unwrap()[1]))
334 .unzip();
335 let rank = |v: &[f64]| {
336 let mut idx: Vec<usize> = (0..v.len()).collect();
337 idx.sort_by(|&p, &q| v[p].total_cmp(&v[q]));
338 let mut r = vec![0.0; v.len()];
339 for (k, &i) in idx.iter().enumerate() {
340 r[i] = k as f64;
341 }
342 r
343 };
344 let (ra, rb) = (rank(&ta), rank(&tb));
345 let mean = (n as f64 - 1.0) / 2.0;
346 let cov: f64 = ra
347 .iter()
348 .zip(&rb)
349 .map(|(x, y)| (x - mean) * (y - mean))
350 .sum();
351 let var: f64 = ra.iter().map(|x| (x - mean).powi(2)).sum();
352 let rho = cov / var;
353 let want = 6.0 / std::f64::consts::PI * (0.4f64).asin();
354 assert!((rho - want).abs() < 0.03, "{rho} vs {want}");
355 assert!(r.reorder_groups("nope", &[1.0], 1).is_err());
356 }
357}