Skip to main content

dag_ml_data_core/
collation.rs

1use std::collections::BTreeSet;
2
3use serde::{Deserialize, Serialize};
4
5use crate::error::{DataError, Result};
6use crate::handle::CoordinatorFeatureBlock;
7use crate::ids::{ObservationId, RepresentationId, SampleId};
8
9#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum CollationPadding {
12    #[default]
13    None,
14    Right,
15    Left,
16    Center,
17}
18
19#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
20pub struct CollationPolicy {
21    #[serde(default)]
22    pub padding: CollationPadding,
23    #[serde(default)]
24    pub truncate: bool,
25    #[serde(default)]
26    pub batch_container: Option<String>,
27    #[serde(default = "default_true")]
28    pub emit_mask: bool,
29    #[serde(default)]
30    pub max_length: Option<usize>,
31    #[serde(default)]
32    pub pad_value: f64,
33}
34
35impl Default for CollationPolicy {
36    fn default() -> Self {
37        Self {
38            padding: CollationPadding::None,
39            truncate: false,
40            batch_container: None,
41            emit_mask: true,
42            max_length: None,
43            pad_value: 0.0,
44        }
45    }
46}
47
48fn default_true() -> bool {
49    true
50}
51
52#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
53pub struct NumericCollationInputBlock {
54    pub block_id: String,
55    pub representation_id: RepresentationId,
56    pub observation_ids: Vec<ObservationId>,
57    pub sample_ids: Vec<SampleId>,
58    pub rows: Vec<Vec<Option<f64>>>,
59    #[serde(default)]
60    pub feature_names: Option<Vec<String>>,
61}
62
63#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
64pub struct NumericTensorBlock {
65    pub block_id: String,
66    pub representation_id: RepresentationId,
67    pub batch_container: String,
68    pub observation_ids: Vec<ObservationId>,
69    pub sample_ids: Vec<SampleId>,
70    pub shape: Vec<usize>,
71    pub values: Vec<f64>,
72    #[serde(default)]
73    pub presence_mask: Option<Vec<bool>>,
74    #[serde(default)]
75    pub validity_mask: Option<Vec<bool>>,
76    #[serde(default)]
77    pub feature_names: Option<Vec<String>>,
78}
79
80pub fn numeric_input_from_feature_block(
81    block: &CoordinatorFeatureBlock,
82) -> Result<NumericCollationInputBlock> {
83    validate_feature_block_shape(block)?;
84    let rows = block
85        .values
86        .iter()
87        .enumerate()
88        .map(|(row_idx, row)| {
89            row.iter()
90                .enumerate()
91                .map(|(feature_idx, value)| match value {
92                    serde_json::Value::Null => Ok(None),
93                    serde_json::Value::Number(number) => number.as_f64().map(Some).ok_or_else(|| {
94                        DataError::Validation(format!(
95                            "feature block `{}` row `{}` feature `{}` contains a non-f64 numeric value",
96                            block.feature_set_id,
97                            block.observation_ids[row_idx],
98                            block.feature_names[feature_idx]
99                        ))
100                    }),
101                    _ => Err(DataError::Validation(format!(
102                        "feature block `{}` row `{}` feature `{}` must be numeric or null for collation",
103                        block.feature_set_id,
104                        block.observation_ids[row_idx],
105                        block.feature_names[feature_idx]
106                    ))),
107                })
108                .collect::<Result<Vec<_>>>()
109        })
110        .collect::<Result<Vec<_>>>()?;
111    Ok(NumericCollationInputBlock {
112        block_id: block.feature_set_id.clone(),
113        representation_id: block.representation_id.clone(),
114        observation_ids: block.observation_ids.clone(),
115        sample_ids: block.sample_ids.clone(),
116        rows,
117        feature_names: Some(block.feature_names.clone()),
118    })
119}
120
121pub fn collate_feature_block(
122    block: &CoordinatorFeatureBlock,
123    policy: &CollationPolicy,
124) -> Result<NumericTensorBlock> {
125    let input = numeric_input_from_feature_block(block)?;
126    collate_numeric_block(&input, policy)
127}
128
129pub fn collate_numeric_block(
130    block: &NumericCollationInputBlock,
131    policy: &CollationPolicy,
132) -> Result<NumericTensorBlock> {
133    validate_numeric_input_block(block)?;
134    validate_collation_policy(policy)?;
135    let target_len = target_length(block, policy)?;
136    validate_feature_names_for_collation(block, policy, target_len)?;
137
138    let batch = block.rows.len();
139    let cells = batch.checked_mul(target_len).ok_or_else(|| {
140        DataError::Validation("collation dimensions exceed addressable capacity".into())
141    })?;
142    let mut values = reserve_collation::<f64>(cells)?;
143    let mut presence = reserve_collation::<bool>(cells)?;
144    let mut validity = reserve_collation::<bool>(cells)?;
145    let mut has_invalid = false;
146    for row in &block.rows {
147        let projected = project_row(row, target_len, policy)?;
148        values.extend(
149            projected
150                .values
151                .iter()
152                .map(|value| value.unwrap_or(policy.pad_value)),
153        );
154        presence.extend(projected.presence.iter().copied());
155        for (value, present) in projected.values.iter().zip(projected.presence.iter()) {
156            let valid = *present && value.is_some();
157            has_invalid |= !valid;
158            validity.push(valid);
159        }
160    }
161
162    Ok(NumericTensorBlock {
163        block_id: block.block_id.clone(),
164        representation_id: block.representation_id.clone(),
165        batch_container: policy
166            .batch_container
167            .clone()
168            .unwrap_or_else(|| "ndarray".to_string()),
169        observation_ids: block.observation_ids.clone(),
170        sample_ids: block.sample_ids.clone(),
171        shape: vec![batch, target_len],
172        values,
173        presence_mask: policy.emit_mask.then_some(presence),
174        validity_mask: has_invalid.then_some(validity),
175        feature_names: projected_feature_names(block.feature_names.as_deref(), target_len, policy)?,
176    })
177}
178
179struct ProjectedRow {
180    values: Vec<Option<f64>>,
181    presence: Vec<bool>,
182}
183
184fn reserve_collation<T>(len: usize) -> Result<Vec<T>> {
185    let mut values = Vec::new();
186    values
187        .try_reserve_exact(len)
188        .map_err(|error| DataError::Validation(format!("collation allocation failed: {error}")))?;
189    Ok(values)
190}
191
192fn validate_collation_policy(policy: &CollationPolicy) -> Result<()> {
193    if policy.max_length == Some(0) {
194        return Err(DataError::Validation(
195            "collation max_length must be greater than zero".to_string(),
196        ));
197    }
198    if let Some(container) = &policy.batch_container {
199        if container.trim().is_empty() {
200            return Err(DataError::Validation(
201                "collation batch_container must not be empty".to_string(),
202            ));
203        }
204    }
205    if !policy.pad_value.is_finite() {
206        return Err(DataError::Validation(
207            "collation pad_value must be finite".to_string(),
208        ));
209    }
210    Ok(())
211}
212
213fn validate_numeric_input_block(block: &NumericCollationInputBlock) -> Result<()> {
214    if block.block_id.trim().is_empty() {
215        return Err(DataError::Validation(
216            "collation input block_id is empty".to_string(),
217        ));
218    }
219    if block.observation_ids.is_empty() {
220        return Err(DataError::Validation(format!(
221            "collation input block `{}` contains no rows",
222            block.block_id
223        )));
224    }
225    if block.observation_ids.len() != block.sample_ids.len()
226        || block.sample_ids.len() != block.rows.len()
227    {
228        return Err(DataError::Validation(format!(
229            "collation input block `{}` row identity/value lengths differ",
230            block.block_id
231        )));
232    }
233    let mut observations = BTreeSet::new();
234    for (idx, row) in block.rows.iter().enumerate() {
235        if !observations.insert(&block.observation_ids[idx]) {
236            return Err(DataError::Validation(format!(
237                "collation input block `{}` contains duplicate observation `{}`",
238                block.block_id, block.observation_ids[idx]
239            )));
240        }
241        if row.is_empty() {
242            return Err(DataError::Validation(format!(
243                "collation input block `{}` row `{}` is empty",
244                block.block_id, block.observation_ids[idx]
245            )));
246        }
247        for value in row.iter().flatten() {
248            if !value.is_finite() {
249                return Err(DataError::Validation(format!(
250                    "collation input block `{}` row `{}` contains a non-finite value",
251                    block.block_id, block.observation_ids[idx]
252                )));
253            }
254        }
255    }
256    if let Some(feature_names) = &block.feature_names {
257        if feature_names.is_empty() {
258            return Err(DataError::Validation(format!(
259                "collation input block `{}` has empty feature_names",
260                block.block_id
261            )));
262        }
263        if feature_names.iter().any(|name| name.trim().is_empty()) {
264            return Err(DataError::Validation(format!(
265                "collation input block `{}` has an empty feature name",
266                block.block_id
267            )));
268        }
269        let mut names = BTreeSet::new();
270        for name in feature_names {
271            if !names.insert(name) {
272                return Err(DataError::Validation(format!(
273                    "collation input block `{}` has duplicate feature `{name}`",
274                    block.block_id
275                )));
276            }
277        }
278    }
279    Ok(())
280}
281
282fn validate_feature_block_shape(block: &CoordinatorFeatureBlock) -> Result<()> {
283    if block.feature_set_id.trim().is_empty() {
284        return Err(DataError::Validation(
285            "feature block feature_set_id is empty".to_string(),
286        ));
287    }
288    if block.feature_names.is_empty() {
289        return Err(DataError::Validation(format!(
290            "feature block `{}` contains no features",
291            block.feature_set_id
292        )));
293    }
294    if block.observation_ids.len() != block.sample_ids.len()
295        || block.sample_ids.len() != block.values.len()
296    {
297        return Err(DataError::Validation(format!(
298            "feature block `{}` row identity/value lengths differ",
299            block.feature_set_id
300        )));
301    }
302    for (idx, row) in block.values.iter().enumerate() {
303        if row.len() != block.feature_names.len() {
304            return Err(DataError::Validation(format!(
305                "feature block `{}` row `{}` has {} values for {} features",
306                block.feature_set_id,
307                block.observation_ids[idx],
308                row.len(),
309                block.feature_names.len()
310            )));
311        }
312    }
313    Ok(())
314}
315
316fn validate_feature_names_for_collation(
317    block: &NumericCollationInputBlock,
318    policy: &CollationPolicy,
319    target_len: usize,
320) -> Result<()> {
321    let Some(feature_names) = &block.feature_names else {
322        return Ok(());
323    };
324    for row in &block.rows {
325        if row.len() != feature_names.len() {
326            return Err(DataError::Validation(format!(
327                "collation input block `{}` named features require rectangular rows",
328                block.block_id
329            )));
330        }
331    }
332    if target_len > feature_names.len() {
333        return Err(DataError::Validation(format!(
334            "collation input block `{}` cannot pad named feature rows",
335            block.block_id
336        )));
337    }
338    if target_len < feature_names.len() && !policy.truncate {
339        return Err(DataError::Validation(format!(
340            "collation input block `{}` feature names require truncation for max_length",
341            block.block_id
342        )));
343    }
344    Ok(())
345}
346
347fn target_length(block: &NumericCollationInputBlock, policy: &CollationPolicy) -> Result<usize> {
348    let max_observed = block.rows.iter().map(Vec::len).max().unwrap_or(0);
349    let target = policy.max_length.unwrap_or(max_observed);
350    if target == 0 {
351        return Err(DataError::Validation(format!(
352            "collation input block `{}` produced an empty target length",
353            block.block_id
354        )));
355    }
356    if !policy.truncate && block.rows.iter().any(|row| row.len() > target) {
357        return Err(DataError::Validation(format!(
358            "collation input block `{}` has rows longer than max_length without truncate",
359            block.block_id
360        )));
361    }
362    if policy.padding == CollationPadding::None && block.rows.iter().any(|row| row.len() < target) {
363        return Err(DataError::Validation(format!(
364            "collation input block `{}` has ragged rows but padding is none",
365            block.block_id
366        )));
367    }
368    Ok(target)
369}
370
371fn project_row(
372    row: &[Option<f64>],
373    target_len: usize,
374    policy: &CollationPolicy,
375) -> Result<ProjectedRow> {
376    let truncated = if row.len() > target_len {
377        if !policy.truncate {
378            return Err(DataError::Validation(
379                "collation row is longer than target length without truncate".to_string(),
380            ));
381        }
382        let start = truncate_start(row.len(), target_len, policy.padding);
383        row[start..start + target_len].to_vec()
384    } else {
385        row.to_vec()
386    };
387
388    if truncated.len() == target_len {
389        return Ok(ProjectedRow {
390            presence: vec![true; target_len],
391            values: truncated,
392        });
393    }
394    if policy.padding == CollationPadding::None {
395        return Err(DataError::Validation(
396            "collation row is shorter than target length but padding is none".to_string(),
397        ));
398    }
399
400    let missing = target_len - truncated.len();
401    let (left_pad, right_pad) = match policy.padding {
402        CollationPadding::None => unreachable!("padding none handled above"),
403        CollationPadding::Right => (0, missing),
404        CollationPadding::Left => (missing, 0),
405        CollationPadding::Center => (missing / 2, missing - (missing / 2)),
406    };
407    let mut values = reserve_collation::<Option<f64>>(target_len)?;
408    let mut presence = reserve_collation::<bool>(target_len)?;
409    values.extend(std::iter::repeat_n(None, left_pad));
410    presence.extend(std::iter::repeat_n(false, left_pad));
411    values.extend(truncated);
412    presence.extend(std::iter::repeat_n(true, target_len - left_pad - right_pad));
413    values.extend(std::iter::repeat_n(None, right_pad));
414    presence.extend(std::iter::repeat_n(false, right_pad));
415    Ok(ProjectedRow { values, presence })
416}
417
418fn truncate_start(row_len: usize, target_len: usize, padding: CollationPadding) -> usize {
419    match padding {
420        CollationPadding::Left => row_len - target_len,
421        CollationPadding::Center => (row_len - target_len) / 2,
422        CollationPadding::None | CollationPadding::Right => 0,
423    }
424}
425
426fn projected_feature_names(
427    feature_names: Option<&[String]>,
428    target_len: usize,
429    policy: &CollationPolicy,
430) -> Result<Option<Vec<String>>> {
431    let Some(feature_names) = feature_names else {
432        return Ok(None);
433    };
434    if target_len == feature_names.len() {
435        return Ok(Some(feature_names.to_vec()));
436    }
437    if target_len > feature_names.len() {
438        return Err(DataError::Validation(
439            "named feature collation cannot add padded feature names".to_string(),
440        ));
441    }
442    let start = truncate_start(feature_names.len(), target_len, policy.padding);
443    Ok(Some(feature_names[start..start + target_len].to_vec()))
444}
445
446#[cfg(test)]
447mod tests {
448    use super::*;
449    use crate::ids::RepresentationId;
450    use serde_json::json;
451
452    fn obs(value: &str) -> ObservationId {
453        ObservationId::new(value).unwrap()
454    }
455
456    fn sample(value: &str) -> SampleId {
457        SampleId::new(value).unwrap()
458    }
459
460    fn feature_block() -> CoordinatorFeatureBlock {
461        CoordinatorFeatureBlock {
462            feature_set_id: "x".to_string(),
463            representation_id: RepresentationId::new("tabular_numeric").unwrap(),
464            feature_names: vec!["f0".to_string(), "f1".to_string()],
465            observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
466            sample_ids: vec![sample("S001"), sample("S002")],
467            values: vec![vec![json!(1.0), json!(2.0)], vec![json!(3.0), json!(4.0)]],
468        }
469    }
470
471    #[test]
472    fn collates_rectangular_feature_block_to_row_major_tensor() {
473        let tensor = collate_feature_block(
474            &feature_block(),
475            &CollationPolicy {
476                emit_mask: false,
477                ..Default::default()
478            },
479        )
480        .unwrap();
481
482        assert_eq!(tensor.shape, vec![2, 2]);
483        assert_eq!(tensor.values, vec![1.0, 2.0, 3.0, 4.0]);
484        assert_eq!(tensor.presence_mask, None);
485        assert_eq!(tensor.validity_mask, None);
486        assert_eq!(
487            tensor.feature_names,
488            Some(vec!["f0".to_string(), "f1".to_string()])
489        );
490    }
491
492    #[test]
493    fn right_padding_emits_presence_and_validity_masks() {
494        let block = NumericCollationInputBlock {
495            block_id: "seq".to_string(),
496            representation_id: RepresentationId::new("sequence_tensor").unwrap(),
497            observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
498            sample_ids: vec![sample("S001"), sample("S002")],
499            rows: vec![vec![Some(1.0), Some(2.0)], vec![Some(3.0), None]],
500            feature_names: None,
501        };
502        let tensor = collate_numeric_block(
503            &block,
504            &CollationPolicy {
505                padding: CollationPadding::Right,
506                max_length: Some(3),
507                pad_value: -1.0,
508                ..Default::default()
509            },
510        )
511        .unwrap();
512
513        assert_eq!(tensor.shape, vec![2, 3]);
514        assert_eq!(tensor.values, vec![1.0, 2.0, -1.0, 3.0, -1.0, -1.0]);
515        assert_eq!(
516            tensor.presence_mask,
517            Some(vec![true, true, false, true, true, false])
518        );
519        assert_eq!(
520            tensor.validity_mask,
521            Some(vec![true, true, false, true, false, false])
522        );
523    }
524
525    #[test]
526    fn no_padding_refuses_ragged_rows() {
527        let block = NumericCollationInputBlock {
528            block_id: "seq".to_string(),
529            representation_id: RepresentationId::new("sequence_tensor").unwrap(),
530            observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
531            sample_ids: vec![sample("S001"), sample("S002")],
532            rows: vec![vec![Some(1.0), Some(2.0)], vec![Some(3.0)]],
533            feature_names: None,
534        };
535
536        let err = collate_numeric_block(&block, &CollationPolicy::default()).unwrap_err();
537
538        assert!(err.to_string().contains("ragged rows"));
539    }
540
541    #[test]
542    fn left_truncation_keeps_suffix_and_projects_feature_names() {
543        let tensor = collate_feature_block(
544            &feature_block(),
545            &CollationPolicy {
546                padding: CollationPadding::Left,
547                truncate: true,
548                max_length: Some(1),
549                ..Default::default()
550            },
551        )
552        .unwrap();
553
554        assert_eq!(tensor.shape, vec![2, 1]);
555        assert_eq!(tensor.values, vec![2.0, 4.0]);
556        assert_eq!(tensor.feature_names, Some(vec!["f1".to_string()]));
557    }
558
559    #[test]
560    fn collation_refuses_non_numeric_feature_values() {
561        let mut block = feature_block();
562        block.values[0][0] = json!("bad");
563
564        let err = collate_feature_block(&block, &CollationPolicy::default()).unwrap_err();
565
566        assert!(err.to_string().contains("must be numeric or null"));
567    }
568}