causal_hub/datasets/table/categorical/
complete.rs1use std::{
2 borrow::Cow,
3 fmt::Display,
4 io::{Read, Write},
5 sync::Arc,
6};
7
8use csv::{ReaderBuilder, WriterBuilder};
9use itertools::Itertools;
10use log::debug;
11use ndarray::prelude::*;
12use serde::{Deserialize, Serialize};
13
14use crate::{
15 datasets::{CatEv, CatEvT, Dataset},
16 io::CsvIO,
17 models::{CatSupport, HasLabels},
18 types::{Error, Labels, Result, Set},
19};
20
21pub type CatType = u8;
23pub type CatSample = Array1<CatType>;
25
26#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
28pub struct CatTable {
29 labels: Labels,
30 support: CatSupport,
31 shape: Array1<usize>,
32 values: Array2<CatType>,
33}
34
35pub struct CatTableEvidenceIter<'a> {
37 rows: ndarray::iter::LanesIter<'a, CatType, Ix1>,
38 support: &'a CatSupport,
39}
40
41impl<'a> Iterator for CatTableEvidenceIter<'a> {
42 type Item = Result<CatEv>;
43
44 fn next(&mut self) -> Option<Self::Item> {
45 let row = self.rows.next()?;
46
47 let evidences = row
48 .iter()
49 .enumerate()
50 .map(|(event, &state)| CatEvT::CertainPositive {
51 event,
52 state: state as usize,
53 });
54
55 Some(CatEv::new(self.support.clone(), evidences))
56 }
57}
58
59impl HasLabels for CatTable {
60 #[inline]
61 fn labels(&self) -> &Labels {
62 &self.labels
63 }
64}
65
66impl CatTable {
67 pub fn new(mut support: CatSupport, mut values: Array2<CatType>) -> Result<Self> {
94 debug!(
96 "Creating a new categorical dataset with {} variables and {} samples.",
97 support.len(),
98 values.nrows()
99 );
100
101 support.iter().try_for_each(|(label, state)| {
103 if state.len() > CatType::MAX as usize {
104 return Err(Error::InvalidParameter(
105 &format!("support[{label}]"),
106 &format!("should have less than 256 support, found {}", state.len()),
107 ));
108 }
109 Ok(())
110 })?;
111 if support.len() != values.ncols() {
113 return Err(Error::IncompatibleShape(
114 &format!("|support| = {}", support.len()),
115 &format!("|cols| = {}", values.ncols()),
116 ));
117 }
118 values
120 .fold_axis(Axis(0), 0, |&a, &b| if a > b { a } else { b })
121 .into_iter()
122 .enumerate()
123 .try_for_each(|(i, x)| {
124 let (label, support) = support
125 .get_index(i)
126 .ok_or_else(|| Error::IndexOutOfBounds(i))?;
127
128 if x >= support.len() as CatType {
129 return Err(Error::InvalidParameter(
130 &format!("values[.., '{label}']"),
131 &format!(
132 "must be less than the number of support ({}), found {x}",
133 support.len()
134 ),
135 ));
136 }
137 Ok(())
138 })?;
139
140 if !support.keys().is_sorted() {
142 let mut indices: Vec<usize> = (0..support.len()).collect();
144 let keys: Vec<_> = support.keys().collect();
146 indices.sort_by_key(|&i| keys[i]);
147 support.sort_keys();
149 let mut new_values = values.clone();
151 indices.into_iter().enumerate().for_each(|(i, j)| {
153 new_values.column_mut(i).assign(&values.column(j));
154 });
155 values = new_values;
157 }
158
159 values
161 .columns_mut()
162 .into_iter()
163 .zip(support.values_mut())
164 .try_for_each(|(mut col, support)| -> Result<_> {
165 if !support.is_sorted() {
167 let mut new_states = support.clone();
169 new_states.sort();
171 col.iter_mut().try_for_each(|value| -> Result<_> {
173 let state = &support[*value as usize];
175 *value = new_states
177 .get_index_of(state)
178 .ok_or_else(|| Error::MissingState(state))?
179 as CatType;
180 Ok(())
181 })?;
182 *support = new_states;
184 }
185 Ok(())
186 })?;
187
188 let labels = support.keys().cloned().collect();
190 let shape = support.values().map(Set::len).collect();
192
193 Ok(Self {
194 labels,
195 support,
196 shape,
197 values,
198 })
199 }
200
201 #[inline]
208 pub const fn support(&self) -> &CatSupport {
209 &self.support
210 }
211
212 #[inline]
219 pub const fn shape(&self) -> &Array1<usize> {
220 &self.shape
221 }
222}
223
224impl Display for CatTable {
225 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
226 let n = self
228 .labels()
229 .iter()
230 .chain(self.support().values().flatten())
231 .map(|x| x.len())
232 .max()
233 .unwrap_or(0);
234
235 let hline = std::iter::repeat_n("-", (n + 3) * self.labels().len() + 1).join("");
237 writeln!(f, "{hline}")?;
238 let header = self.labels().iter().map(|x| format!("{x:n$}")).join(" | ");
240 writeln!(f, "| {header} |")?;
241 let separator = (0..self.labels().len()).map(|_| "-".repeat(n)).join(" | ");
243 writeln!(f, "| {separator} |")?;
244 for row in self.values.rows() {
246 let row = row
248 .iter()
249 .enumerate()
250 .map(|(i, &x)| &self.support()[i][x as usize])
251 .map(|x| format!("{x:n$}"))
252 .join(" | ");
253 writeln!(f, "| {row} |")?;
254 }
255 writeln!(f, "{hline}")
257 }
258}
259
260impl Dataset for CatTable {
261 type Values = Array2<CatType>;
262 type Support = CatSupport;
263 type Evidence = CatEv;
264 type EvidenceIter<'a> = CatTableEvidenceIter<'a>;
265
266 #[inline]
267 fn values(&self) -> &Self::Values {
268 &self.values
269 }
270
271 #[inline]
272 fn support(&self) -> Cow<'_, Self::Support> {
273 Cow::Borrowed(&self.support)
274 }
275
276 fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
277 CatTableEvidenceIter {
278 rows: self.values.rows().into_iter(),
279 support: &self.support,
280 }
281 }
282
283 #[inline]
284 fn sample_size(&self) -> f64 {
285 self.values.nrows() as f64
286 }
287
288 fn select(&self, x: &Set<usize>) -> Result<Self> {
289 x.iter().try_for_each(|&i| {
291 if i >= self.values.ncols() {
292 return Err(Error::IndexOutOfBounds(i));
293 }
294 Ok(())
295 })?;
296
297 let support: CatSupport = x
299 .iter()
300 .map(|&i| {
301 self.support
302 .get_index(i)
303 .map(|(label, support)| (label.clone(), support.clone()))
304 .ok_or_else(|| Error::IndexOutOfBounds(i))
305 })
306 .collect::<Result<_>>()?;
307
308 let mut new_values = Array2::zeros((self.values.nrows(), x.len()));
310 x.iter().enumerate().for_each(|(j, &i)| {
312 new_values.column_mut(j).assign(&self.values.column(i));
313 });
314 let values = new_values;
316
317 Self::new(support, values)
319 }
320}
321
322impl CsvIO for CatTable {
323 fn from_csv_reader<R: Read>(reader: R) -> Result<Self> {
324 let mut reader = ReaderBuilder::new().has_headers(true).from_reader(reader);
326
327 if !reader.has_headers() {
329 return Err(Error::MissingHeader());
330 }
331
332 let labels: Labels = reader
334 .headers()?
335 .into_iter()
336 .map(|x| x.to_owned())
337 .collect();
338
339 let mut support: CatSupport = labels
341 .iter()
342 .map(|x| (x.clone(), Default::default()))
343 .collect();
344
345 let values: Vec<CatType> = reader.into_records().enumerate().try_fold(
347 Vec::new(),
348 |mut values, (i, row)| -> Result<_> {
349 let row = row.map_err(|evidence| Error::Csv(Arc::new(evidence)))?;
351 for (j, (x, support)) in row.into_iter().zip(support.values_mut()).enumerate() {
353 if x.is_empty() {
355 return Err(Error::MissingValue(i + 1, j + 1));
356 }
357 let (idx, _) = support.insert_full(x.to_owned());
359 values.push(idx as CatType);
361 }
362
363 Ok(values)
364 },
365 )?;
366
367 let values = Array1::from_vec(values);
369
370 let ncols = labels.len();
372 let nrows = values.len() / ncols;
373 let values = values.into_shape_with_order((nrows, ncols))?;
375
376 Self::new(support, values)
378 }
379
380 fn to_csv_writer<W: Write>(&self, writer: W) -> Result<()> {
381 let mut writer = WriterBuilder::new().has_headers(true).from_writer(writer);
383
384 writer.write_record(self.labels.iter())?;
386
387 for row in self.values.rows() {
389 let record = row
391 .iter()
392 .zip(self.support().values())
393 .map(|(&x, support)| &support[x as usize]);
394 writer.write_record(record)?;
396 }
397
398 Ok(())
399 }
400}