causal_hub/datasets/table/categorical/
incomplete.rs1use std::{
2 borrow::Cow,
3 io::{Read, Write},
4 sync::Arc,
5};
6
7use csv::{ReaderBuilder, WriterBuilder};
8use ndarray::prelude::*;
9use serde::{Deserialize, Serialize};
10
11use crate::{
12 datasets::{
13 CatEv, CatEvT, CatTable, CatType, CatWtdTable, Dataset, IncDataset, MissingMechanism,
14 MissingTable,
15 },
16 estimators::{BE, CPDEstimator},
17 io::CsvIO,
18 models::{CPD, CatSupport, HasLabels},
19 set, support,
20 types::{Error, Labels, Result, Set},
21};
22
23#[derive(Clone, Debug, Serialize, Deserialize)]
25pub struct CatIncTable {
26 labels: Labels,
27 support: CatSupport,
28 shape: Array1<usize>,
29 values: Array2<CatType>,
30 missing: MissingTable,
31}
32
33pub struct CatIncTableEvidenceIter<'a> {
35 rows: ndarray::iter::LanesIter<'a, CatType, Ix1>,
36 support: &'a CatSupport,
37 missing: CatType,
38}
39
40impl<'a> Iterator for CatIncTableEvidenceIter<'a> {
41 type Item = Result<CatEv>;
42
43 fn next(&mut self) -> Option<Self::Item> {
44 let row = self.rows.next()?;
45
46 let evidences = row.iter().enumerate().filter_map(|(event, &state)| {
47 (state != self.missing).then_some(CatEvT::CertainPositive {
48 event,
49 state: state as usize,
50 })
51 });
52
53 Some(CatEv::new(self.support.clone(), evidences))
54 }
55}
56
57impl HasLabels for CatIncTable {
58 #[inline]
59 fn labels(&self) -> &Labels {
60 &self.labels
61 }
62}
63
64impl CatIncTable {
65 pub fn new(mut support: CatSupport, mut values: Array2<CatType>) -> Result<Self> {
92 support.iter().try_for_each(|(label, state)| {
94 if state.len() > CatType::MAX as usize {
95 return Err(Error::InvalidParameter(
96 label,
97 &format!("should have less than 256 support, found {}", state.len()),
98 ));
99 }
100 Ok(())
101 })?;
102 if support.len() != values.ncols() {
104 return Err(Error::IncompatibleShape(
105 &support.len().to_string(),
106 &values.ncols().to_string(),
107 ));
108 }
109 let max_values = values.fold_axis(
111 Axis(0),
112 0,
113 |&a, &b| if a > b || b == Self::MISSING { a } else { b },
115 );
116 max_values.into_iter().enumerate().try_for_each(|(i, x)| {
117 if x >= support[i].len() as CatType {
118 return Err(Error::IndexOutOfBounds(x as usize));
119 }
120 Ok(())
121 })?;
122
123 if !support.keys().is_sorted() {
125 let mut indices: Vec<usize> = (0..support.len()).collect();
127 indices.sort_by(|&i, &j| {
129 support
130 .get_index(i)
131 .map(|(l, _)| l)
132 .cmp(&support.get_index(j).map(|(l, _)| l))
133 });
134 support.sort_keys();
136 let mut new_values = values.clone();
138 indices.into_iter().enumerate().for_each(|(i, j)| {
140 new_values.column_mut(i).assign(&values.column(j));
141 });
142 values = new_values;
144 }
145
146 values
148 .columns_mut()
149 .into_iter()
150 .zip(support.values_mut())
151 .try_for_each(|(mut col, support)| -> Result<_> {
152 if !support.is_sorted() {
154 let mut new_states = support.clone();
156 new_states.sort();
158 col.iter_mut().try_for_each(|value| -> Result<_> {
160 if *value != Self::MISSING {
162 *value = new_states
164 .get_index_of(&support[*value as usize])
165 .ok_or_else(|| Error::MissingState(&support[*value as usize]))?
166 as CatType;
167 }
168 Ok(())
169 })?;
170 *support = new_states;
172 }
173 Ok(())
174 })?;
175
176 let labels: Labels = support.keys().cloned().collect();
178 let shape = support.values().map(Set::len).collect();
180
181 let missing_mask = values.mapv(|x| x == Self::MISSING);
183 let missing = MissingTable::new(labels.clone(), missing_mask)?;
185
186 Ok(Self {
187 labels,
188 support,
189 shape,
190 values,
191 missing,
192 })
193 }
194
195 #[inline]
202 pub const fn support(&self) -> &CatSupport {
203 &self.support
204 }
205
206 #[inline]
213 pub const fn shape(&self) -> &Array1<usize> {
214 &self.shape
215 }
216}
217
218impl Dataset for CatIncTable {
219 type Values = Array2<CatType>;
220 type Support = CatSupport;
221 type Evidence = CatEv;
222 type EvidenceIter<'a> = CatIncTableEvidenceIter<'a>;
223
224 #[inline]
225 fn values(&self) -> &Self::Values {
226 &self.values
227 }
228
229 #[inline]
230 fn support(&self) -> Cow<'_, Self::Support> {
231 Cow::Borrowed(&self.support)
232 }
233
234 fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
235 CatIncTableEvidenceIter {
236 rows: self.values.rows().into_iter(),
237 support: &self.support,
238 missing: Self::MISSING,
239 }
240 }
241
242 #[inline]
243 fn sample_size(&self) -> f64 {
244 self.values.nrows() as f64
245 }
246
247 fn select(&self, x: &Set<usize>) -> Result<Self> {
248 x.iter().try_for_each(|&i| {
250 if i >= self.values.ncols() {
251 return Err(Error::IndexOutOfBounds(i));
252 }
253 Ok(())
254 })?;
255
256 let support: CatSupport = x
258 .iter()
259 .map(|&i| {
260 self.support
261 .get_index(i)
262 .map(|(label, support)| (label.clone(), support.clone()))
263 .ok_or_else(|| Error::IndexOutOfBounds(i))
264 })
265 .collect::<Result<_>>()?;
266
267 let mut new_values = Array2::zeros((self.values.nrows(), x.len()));
269 x.iter().enumerate().for_each(|(j, &i)| {
271 new_values.column_mut(j).assign(&self.values.column(i));
272 });
273 let values = new_values;
275
276 Self::new(support, values)
278 }
279}
280
281impl IncDataset for CatIncTable {
282 type Missing = CatType;
283 const MISSING: Self::Missing = CatType::MAX;
284
285 type Complete = CatTable;
286 type Weighted = CatWtdTable;
287
288 #[inline]
289 fn missing(&self) -> &MissingTable {
290 &self.missing
291 }
292
293 fn ipw_weights(
294 &self,
295 d_u: &Self::Complete,
296 u: &Set<usize>,
297 pr: &MissingMechanism,
298 ) -> Result<Array1<f64>> {
299 let pr_iter = u.iter().filter_map(|&ri| pr.get(&ri).map(|pri| (ri, pri)));
301 let pr_iter = pr_iter.filter(|(_, pri)| !pri.is_empty());
303
304 let beta_i = |d_u: &Self::Complete, ri: usize, pri: &Set<usize>| -> Result<Array1<f64>> {
306 let d_pri_rpri = self.pw_deletion(pri)?;
310 let d_pri_ri_rpri = self.pw_deletion(&(&set![ri] | pri))?;
311 let x_pri_rpri = d_pri_rpri.indices_from(pri, self.labels())?;
313 let x_pri_ri_rpri = d_pri_ri_rpri.indices_from(pri, self.labels())?;
314 let p_pri_rpri = BE::new(&d_pri_rpri).fit(&x_pri_rpri, &set![])?;
316 let p_pri_ri_rpri = BE::new(&d_pri_ri_rpri).fit(&x_pri_ri_rpri, &set![])?;
317
318 let x_pri_u = d_u.indices_from(pri, self.labels())?;
320
321 let mut b_pri_rpri = Array::zeros(d_u.values().nrows());
323 let mut b_pri_ri_rpri = b_pri_rpri.clone();
324 for (d_u_j, (b_pri_rpri_j, b_pri_ri_rpri_j)) in d_u
326 .values()
327 .rows()
328 .into_iter()
329 .zip(b_pri_rpri.iter_mut().zip(b_pri_ri_rpri.iter_mut()))
330 {
331 let pri_j = x_pri_u.iter().map(|&j| d_u_j[j]).collect();
333 *b_pri_rpri_j = p_pri_rpri.pf(&pri_j, &array![])?;
335 *b_pri_ri_rpri_j = p_pri_ri_rpri.pf(&pri_j, &array![])?;
336 }
337 Ok(b_pri_rpri / b_pri_ri_rpri)
339 };
340
341 let mut beta = Array::ones(d_u.values().nrows());
343 for (ri, pri) in pr_iter {
344 let beta_i = beta_i(d_u, ri, pri)?;
345 beta *= &beta_i;
346 }
347
348 if beta.sum() > 0. {
350 beta *= (beta.len() as f64) / beta.sum();
351 }
352
353 Ok(beta)
354 }
355
356 fn lw_deletion(&self) -> Result<Self::Complete> {
357 let mut new_values = Array::zeros((
359 self.missing.complete_rows_count(), self.values.ncols(),
361 ));
362
363 let rows = self
365 .values
366 .rows()
367 .into_iter()
368 .zip(self.missing.missing_mask_by_rows())
369 .filter_map(|(row, &is_complete)| if !is_complete { Some(row) } else { None });
371
372 rows.zip(new_values.rows_mut())
374 .for_each(|(row, mut new_row)| new_row.assign(&row));
375
376 Self::Complete::new(self.support.clone(), new_values)
378 }
379
380 fn pw_deletion(&self, x: &Set<usize>) -> Result<Self::Complete> {
381 if x.is_empty() {
383 let stats = support![];
384 let v = Array::default((0, 0));
385 return Self::Complete::new(stats, v);
386 }
387
388 x.iter().try_for_each(|&i| {
390 if i >= self.values.ncols() {
391 return Err(Error::IndexOutOfBounds(i));
392 }
393 Ok(())
394 })?;
395
396 let mut cols = x.clone();
398 cols.sort();
400
401 let rows: Vec<_> = self
403 .missing
404 .missing_mask()
405 .rows()
406 .into_iter()
407 .enumerate()
408 .filter_map(|(i, row)| {
409 if !cols.iter().any(|&j| row[j]) {
411 Some(i)
412 } else {
413 None
414 }
415 })
416 .collect();
417
418 let new_values = Array::from_shape_fn(
420 (rows.len(), cols.len()), |(i, j)| self.values[[rows[i], cols[j]]],
422 );
423
424 let new_states = cols
426 .iter()
427 .map(|&j| {
428 self.support
429 .get_index(j)
430 .map(|(label, state)| (label.clone(), state.clone()))
431 .ok_or_else(|| Error::IndexOutOfBounds(j))
432 })
433 .collect::<Result<_>>()?;
434
435 Self::Complete::new(new_states, new_values)
437 }
438
439 fn ipw_deletion(&self, x: &Set<usize>, pr: &MissingMechanism) -> Result<Self::Weighted> {
440 if x.is_empty() {
442 let stats = support![];
443 let v = Array::default((0, 0));
444 let w = Array::default(0);
445 return Self::Weighted::new(Self::Complete::new(stats, v)?, w);
446 }
447
448 x.iter().try_for_each(|&i| {
450 if i >= self.values.ncols() {
451 return Err(Error::IndexOutOfBounds(i));
452 }
453 Ok(())
454 })?;
455 pr.keys().try_for_each(|&i| {
457 if i >= self.values.ncols() {
458 return Err(Error::IndexOutOfBounds(i));
459 }
460 Ok(())
461 })?;
462 if !pr.keys().is_sorted() {
464 return Err(Error::InvalidParameter(
465 "missing_mechanism",
466 "keys must be sorted.",
467 ));
468 }
469 if !pr.values().all(|pri| pri.iter().is_sorted()) {
470 return Err(Error::InvalidParameter(
471 "missing_mechanism",
472 "values must be sorted.",
473 ));
474 }
475
476 let mut u = x.clone();
478 let mut pru: Set<_> = x
479 .iter()
480 .flat_map(|&x| pr.get(&x).cloned())
481 .flatten()
482 .collect();
483 while !pru.is_subset(&u) {
485 u.extend(pru.drain(..));
486 pru.extend(u.iter().flat_map(|&u| pr.get(&u).cloned()).flatten());
487 }
488 u.sort();
490
491 let d_u = self.pw_deletion(&u)?;
493 let b_u = self.ipw_weights(&d_u, &u, pr)?;
495
496 let x = d_u.indices_from(x, self.labels())?;
498 let d_x = d_u.select(&x)?;
500
501 Self::Weighted::new(d_x, b_u)
503 }
504
505 fn aipw_deletion(&self, x: &Set<usize>, pr: &MissingMechanism) -> Result<Self::Weighted> {
506 if x.is_empty() {
508 let stats = support![];
509 let v = Array::default((0, 0));
510 let w = Array::default(0);
511 return Self::Weighted::new(Self::Complete::new(stats, v)?, w);
512 }
513
514 x.iter().try_for_each(|&i| {
516 if i >= self.values.ncols() {
517 return Err(Error::IndexOutOfBounds(i));
518 }
519 Ok(())
520 })?;
521 pr.keys().try_for_each(|&i| {
523 if i >= self.values.ncols() {
524 return Err(Error::IndexOutOfBounds(i));
525 }
526 Ok(())
527 })?;
528 if !pr.keys().is_sorted() {
530 return Err(Error::InvalidParameter(
531 "missing_mechanism",
532 "keys must be sorted.",
533 ));
534 }
535 if !pr.values().all(|pri| pri.iter().is_sorted()) {
536 return Err(Error::InvalidParameter(
537 "missing_mechanism",
538 "values must be sorted.",
539 ));
540 }
541
542 let mut w = x.clone();
544 let prw: Set<_> = x
545 .iter()
546 .flat_map(|x| pr.get(x).cloned())
547 .flatten()
548 .collect();
549 w.sort();
551
552 let v_m = self.missing().partially_observed();
554 if (&(&prw - &w) & v_m).is_empty() {
556 return self.ipw_deletion(x, pr); };
558
559 let d_x = self.pw_deletion(x)?;
561 let b_x = Array::ones(d_x.values().nrows()); Self::Weighted::new(d_x, b_x)
564 }
565}
566
567impl CsvIO for CatIncTable {
568 fn from_csv_reader<R: Read>(reader: R) -> Result<Self> {
569 let mut reader = ReaderBuilder::new().has_headers(true).from_reader(reader);
571
572 if !reader.has_headers() {
574 return Err(Error::MissingHeader());
575 }
576
577 let labels: Labels = reader
579 .headers()?
580 .into_iter()
581 .map(|x| x.to_owned())
582 .collect();
583
584 let mut support: CatSupport = labels
586 .iter()
587 .map(|x| (x.clone(), Default::default()))
588 .collect();
589
590 let values: Vec<CatType> =
592 reader
593 .into_records()
594 .try_fold(Vec::new(), |mut values, row| -> Result<_> {
595 let row = row.map_err(|evidence| Error::Csv(Arc::new(evidence)))?;
597 values.extend(
599 row.into_iter()
600 .zip(support.values_mut())
601 .map(|(x, support)| {
602 if x.is_empty() {
604 Self::MISSING
605 } else {
606 let (x, _) = support.insert_full(x.to_owned());
608 x as CatType
610 }
611 }),
612 );
613
614 Ok(values)
615 })?;
616
617 let ncols = labels.len();
619 let nrows = values.len() / ncols;
620 let values = Array1::from_vec(values).into_shape_with_order((nrows, ncols))?;
622
623 Self::new(support, values)
625 }
626
627 fn to_csv_writer<W: Write>(&self, writer: W) -> Result<()> {
628 let mut writer = WriterBuilder::new().has_headers(true).from_writer(writer);
630
631 writer.write_record(self.labels.iter())?;
633
634 let missing = String::new();
636
637 for row in self.values.rows() {
639 let record = row.iter().zip(self.support().values());
641 let record = record.map(|(&x, support)| {
643 if x == Self::MISSING {
645 return &missing;
646 }
647 &support[x as usize]
649 });
650 writer.write_record(record)?;
652 }
653
654 Ok(())
655 }
656}