Skip to main content

antecedent_data/
project.rs

1//! Column projection: narrow a table to the variables needed after identification.
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5use std::collections::{HashMap, HashSet};
6
7use antecedent_core::{CausalSchemaBuilder, VariableId};
8
9use crate::dataset::TabularData;
10use crate::error::DataError;
11use crate::storage::OwnedColumnarStorage;
12use crate::table::TableView;
13
14/// Dense-id remap produced by [`TabularData::project`].
15///
16/// Maps original (source) variable ids onto contiguous projected ids `0..k-1`.
17#[derive(Clone, Debug)]
18pub struct IdRemap {
19    old_to_new: HashMap<VariableId, VariableId>,
20}
21
22impl IdRemap {
23    /// Map an original id to its projected dense id.
24    ///
25    /// # Errors
26    ///
27    /// [`DataError::UnknownVariable`] when `old` was not included in the projection.
28    pub fn map(&self, old: VariableId) -> Result<VariableId, DataError> {
29        self.old_to_new.get(&old).copied().ok_or(DataError::UnknownVariable { id: old })
30    }
31
32    /// Number of projected columns.
33    #[must_use]
34    pub fn len(&self) -> usize {
35        self.old_to_new.len()
36    }
37
38    /// Whether empty.
39    #[must_use]
40    pub fn is_empty(&self) -> bool {
41        self.old_to_new.is_empty()
42    }
43}
44
45impl TabularData {
46    /// Project onto `ids` (order preserved, duplicates dropped).
47    ///
48    /// Column value buffers are Arc-shared when possible; schema and column ids
49    /// are rebuilt as contiguous `0..k-1`. Analysis mask and weights are retained.
50    ///
51    /// # Errors
52    ///
53    /// Unknown variable, empty selection, or schema construction failure.
54    pub fn project(&self, ids: &[VariableId]) -> Result<(Self, IdRemap), DataError> {
55        let mut seen = HashSet::new();
56        let mut ordered: Vec<VariableId> = Vec::with_capacity(ids.len());
57        for &id in ids {
58            if seen.insert(id) {
59                ordered.push(id);
60            }
61        }
62        if ordered.is_empty() {
63            return Err(DataError::EmptySelection {
64                context: "column projection: no variables requested",
65            });
66        }
67
68        let storage = self.storage();
69        let schema = storage.schema();
70        let mut builder = CausalSchemaBuilder::new();
71        let mut cols = Vec::with_capacity(ordered.len());
72        let mut old_to_new = HashMap::with_capacity(ordered.len());
73
74        for (new_idx, &old_id) in ordered.iter().enumerate() {
75            let meta = schema.get(old_id).map_err(|_| DataError::UnknownVariable { id: old_id })?;
76            builder
77                .add_variable(
78                    std::sync::Arc::clone(&meta.name),
79                    meta.value_type.clone(),
80                    meta.role_hints,
81                    meta.unit.clone(),
82                    meta.category_domain,
83                    meta.measurement.clone(),
84                )
85                .map_err(|e| DataError::Schema(e.to_string()))?;
86            let new_id = VariableId::from_raw(u32::try_from(new_idx).map_err(|_| {
87                DataError::InvalidArgument {
88                    message: "projected schema exceeds VariableId range".into(),
89                }
90            })?);
91            old_to_new.insert(old_id, new_id);
92            let col = storage
93                .columns()
94                .get(old_id.as_usize())
95                .ok_or(DataError::UnknownVariable { id: old_id })?;
96            cols.push(col.with_id(new_id));
97        }
98
99        let new_schema = builder.build().map_err(|e| DataError::Schema(e.to_string()))?;
100        let new_storage = OwnedColumnarStorage::try_new(
101            new_schema,
102            cols,
103            storage.analysis_mask().cloned(),
104            storage.weights().map(std::sync::Arc::from),
105        )?;
106        Ok((Self::new(new_storage), IdRemap { old_to_new }))
107    }
108}
109
110/// Deduplicate variable ids while preserving first-seen order.
111#[must_use]
112pub fn dedupe_variable_ids(ids: impl IntoIterator<Item = VariableId>) -> Vec<VariableId> {
113    let mut seen = HashSet::new();
114    let mut out = Vec::new();
115    for id in ids {
116        if seen.insert(id) {
117            out.push(id);
118        }
119    }
120    out
121}
122
123#[cfg(test)]
124mod tests {
125    #![allow(clippy::cast_precision_loss)]
126
127    use super::*;
128    use crate::column::{Float64Column, OwnedColumn, ValidityBitmap};
129    use antecedent_core::{MeasurementSpec, RoleHint, SmallRoleSet, ValueType};
130    use std::sync::Arc;
131
132    fn float_table(names: &[&str], rows: usize) -> TabularData {
133        let mut b = CausalSchemaBuilder::new();
134        for name in names {
135            b.add_variable(
136                *name,
137                ValueType::Continuous,
138                SmallRoleSet::from_hint(RoleHint::Context),
139                None,
140                None,
141                MeasurementSpec::default(),
142            )
143            .unwrap();
144        }
145        let schema = b.build().unwrap();
146        let cols: Vec<OwnedColumn> = names
147            .iter()
148            .enumerate()
149            .map(|(i, _)| {
150                let id = VariableId::from_raw(u32::try_from(i).unwrap());
151                let values: Arc<[f64]> =
152                    (0..rows).map(|r| (r + i * 100) as f64).collect::<Vec<_>>().into();
153                OwnedColumn::Float64(
154                    Float64Column::new(id, values, ValidityBitmap::all_valid(rows)).unwrap(),
155                )
156            })
157            .collect();
158        TabularData::new(OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap())
159    }
160
161    #[test]
162    fn project_preserves_values_and_remaps_ids() {
163        let data = float_table(&["noise", "t", "y", "z", "extra"], 4);
164        let t = VariableId::from_raw(1);
165        let y = VariableId::from_raw(2);
166        let z = VariableId::from_raw(3);
167        let (proj, remap) = data.project(&[t, y, z]).unwrap();
168        assert_eq!(proj.schema().len(), 3);
169        assert_eq!(proj.schema().get(VariableId::from_raw(0)).unwrap().name.as_ref(), "t");
170        assert_eq!(proj.schema().get(VariableId::from_raw(1)).unwrap().name.as_ref(), "y");
171        assert_eq!(proj.schema().get(VariableId::from_raw(2)).unwrap().name.as_ref(), "z");
172        assert_eq!(remap.map(t).unwrap(), VariableId::from_raw(0));
173        assert_eq!(remap.map(y).unwrap(), VariableId::from_raw(1));
174        assert_eq!(remap.map(z).unwrap(), VariableId::from_raw(2));
175        assert!(remap.map(VariableId::from_raw(0)).is_err());
176
177        let view = proj.column(VariableId::from_raw(0)).unwrap();
178        let crate::column::ColumnView::Float64(c) = view else {
179            panic!("expected float64");
180        };
181        assert_eq!(c.values.as_slice(), &[100.0, 101.0, 102.0, 103.0]);
182    }
183
184    #[test]
185    fn project_shares_float_buffers() {
186        let data = float_table(&["t", "y", "noise"], 8);
187        let t = VariableId::from_raw(0);
188        let y = VariableId::from_raw(1);
189        let before = match data.storage().columns()[0] {
190            OwnedColumn::Float64(ref c) => c.values.as_slice().as_ptr(),
191            _ => panic!("float"),
192        };
193        let (proj, _) = data.project(&[t, y]).unwrap();
194        let after = match proj.storage().columns()[0] {
195            OwnedColumn::Float64(ref c) => c.values.as_slice().as_ptr(),
196            _ => panic!("float"),
197        };
198        assert_eq!(before, after);
199    }
200}