Skip to main content

entrenar/config/train/batches/
parquet.rs

1//! Parquet batch loading using alimentar
2
3use super::super::arrow::arrow_array_to_f32;
4use super::rebatch::rebatch;
5use crate::error::{Error, Result};
6use crate::train::Batch;
7use crate::Tensor;
8use alimentar::{ArrowDataset, Dataset};
9use arrow::datatypes::Schema;
10use arrow::record_batch::RecordBatch;
11use std::path::Path;
12
13/// Column detection result
14struct ColumnPair<'a> {
15    input_name: &'a str,
16    target_name: &'a str,
17}
18
19/// Detect input column from schema
20fn detect_input_column<'a>(column_names: &[&'a str]) -> Option<&'a str> {
21    column_names
22        .iter()
23        .find(|&&n| n == "input" || n == "input_ids" || n == "x" || n == "features")
24        .copied()
25}
26
27/// Detect target column from schema
28fn detect_target_column<'a>(column_names: &[&'a str]) -> Option<&'a str> {
29    column_names
30        .iter()
31        .find(|&&n| n == "target" || n == "output" || n == "labels" || n == "y")
32        .copied()
33}
34
35/// Detect input/target column pair from schema
36fn detect_columns<'a>(column_names: &[&'a str]) -> Option<ColumnPair<'a>> {
37    let input_name = detect_input_column(column_names)?;
38    let target_name = detect_target_column(column_names)?;
39    Some(ColumnPair { input_name, target_name })
40}
41
42/// Reject a parquet dataset whose input/target columns could not be identified.
43///
44/// Never substitutes demo data: training on fabricated examples while reporting
45/// success is worse than refusing to start.
46fn missing_columns_error(path: &Path, column_names: &[&str]) -> Error {
47    Error::ConfigError(format!(
48        "Could not find input/target columns in parquet '{}' (found: {column_names:?}). \
49         Expected a pair like input/target, x/y or features/labels.",
50        path.display()
51    ))
52}
53
54/// Convert a single record batch to a training batch
55fn record_batch_to_training_batch(
56    record_batch: &RecordBatch,
57    schema: &Schema,
58    input_name: &str,
59    target_name: &str,
60) -> Result<Batch> {
61    let input_idx = schema
62        .index_of(input_name)
63        .map_err(|e| Error::ConfigError(format!("Column not found: {e}")))?;
64    let target_idx = schema
65        .index_of(target_name)
66        .map_err(|e| Error::ConfigError(format!("Column not found: {e}")))?;
67
68    let input_array = record_batch.column(input_idx);
69    let target_array = record_batch.column(target_idx);
70
71    let input_data = arrow_array_to_f32(input_array)?;
72    let target_data = arrow_array_to_f32(target_array)?;
73
74    Ok(Batch::new(Tensor::from_vec(input_data, false), Tensor::from_vec(target_data, false)))
75}
76
77/// Process all record batches from dataset
78fn process_record_batches(dataset: &ArrowDataset, columns: &ColumnPair<'_>) -> Result<Vec<Batch>> {
79    let schema = dataset.schema();
80    let mut batches = Vec::new();
81
82    for record_batch in dataset.iter() {
83        let batch = record_batch_to_training_batch(
84            &record_batch,
85            &schema,
86            columns.input_name,
87            columns.target_name,
88        )?;
89        batches.push(batch);
90    }
91
92    Ok(batches)
93}
94
95/// Load batches from parquet file using alimentar
96pub fn load_parquet_batches(path: &Path, batch_size: usize) -> Result<Vec<Batch>> {
97    println!("  Loading parquet: {}", path.display());
98
99    let dataset = ArrowDataset::from_parquet(path).map_err(|e| {
100        Error::ConfigError(format!("Failed to load parquet {}: {}", path.display(), e))
101    })?;
102
103    println!("  Loaded {} rows from parquet", dataset.len());
104
105    let schema = dataset.schema();
106    let column_names: Vec<&str> = schema.fields().iter().map(|f| f.name().as_str()).collect();
107
108    let columns = match detect_columns(&column_names) {
109        Some(cols) => cols,
110        None => return Err(missing_columns_error(path, &column_names)),
111    };
112
113    println!("  Using columns: input='{}', target='{}'", columns.input_name, columns.target_name);
114
115    let mut batches = process_record_batches(&dataset, &columns)?;
116
117    // Re-batch to desired batch size if needed
118    if batches.len() > 1 && batch_size > 0 {
119        batches = rebatch(batches, batch_size);
120    }
121
122    Ok(batches)
123}
124
125#[cfg(test)]
126mod tests {
127    use super::*;
128    use arrow::array::{Float32Array, Float64Array, Int32Array};
129    use arrow::datatypes::{DataType, Field};
130    use std::sync::Arc;
131
132    fn make_test_schema() -> Schema {
133        Schema::new(vec![
134            Field::new("input", DataType::Float32, false),
135            Field::new("target", DataType::Float32, false),
136        ])
137    }
138
139    fn make_test_record_batch() -> RecordBatch {
140        let schema = Arc::new(make_test_schema());
141        let input = Float32Array::from(vec![1.0, 2.0, 3.0, 4.0]);
142        let target = Float32Array::from(vec![0.0, 1.0, 0.0, 1.0]);
143        RecordBatch::try_new(schema, vec![Arc::new(input), Arc::new(target)])
144            .expect("conversion should succeed")
145    }
146
147    #[test]
148    fn test_detect_input_column_input() {
149        let cols = vec!["input", "target"];
150        assert_eq!(detect_input_column(&cols), Some("input"));
151    }
152
153    #[test]
154    fn test_detect_input_column_input_ids() {
155        let cols = vec!["input_ids", "labels"];
156        assert_eq!(detect_input_column(&cols), Some("input_ids"));
157    }
158
159    #[test]
160    fn test_detect_input_column_x() {
161        let cols = vec!["x", "y"];
162        assert_eq!(detect_input_column(&cols), Some("x"));
163    }
164
165    #[test]
166    fn test_detect_input_column_features() {
167        let cols = vec!["features", "labels"];
168        assert_eq!(detect_input_column(&cols), Some("features"));
169    }
170
171    #[test]
172    fn test_detect_input_column_none() {
173        let cols = vec!["foo", "bar"];
174        assert_eq!(detect_input_column(&cols), None);
175    }
176
177    #[test]
178    fn test_detect_target_column_target() {
179        let cols = vec!["input", "target"];
180        assert_eq!(detect_target_column(&cols), Some("target"));
181    }
182
183    #[test]
184    fn test_detect_target_column_output() {
185        let cols = vec!["input", "output"];
186        assert_eq!(detect_target_column(&cols), Some("output"));
187    }
188
189    #[test]
190    fn test_detect_target_column_labels() {
191        let cols = vec!["features", "labels"];
192        assert_eq!(detect_target_column(&cols), Some("labels"));
193    }
194
195    #[test]
196    fn test_detect_target_column_y() {
197        let cols = vec!["x", "y"];
198        assert_eq!(detect_target_column(&cols), Some("y"));
199    }
200
201    #[test]
202    fn test_detect_target_column_none() {
203        let cols = vec!["foo", "bar"];
204        assert_eq!(detect_target_column(&cols), None);
205    }
206
207    #[test]
208    fn test_detect_columns_success() {
209        let cols = vec!["input", "target"];
210        let result = detect_columns(&cols);
211        assert!(result.is_some());
212        let pair = result.expect("operation should succeed");
213        assert_eq!(pair.input_name, "input");
214        assert_eq!(pair.target_name, "target");
215    }
216
217    #[test]
218    fn test_detect_columns_missing_input() {
219        let cols = vec!["foo", "target"];
220        assert!(detect_columns(&cols).is_none());
221    }
222
223    #[test]
224    fn test_detect_columns_missing_target() {
225        let cols = vec!["input", "bar"];
226        assert!(detect_columns(&cols).is_none());
227    }
228
229    #[test]
230    fn test_missing_columns_is_an_error_naming_the_columns() {
231        // Was `test_handle_missing_columns_returns_demo_batches`, which asserted the
232        // defect: an unreadable parquet silently became synthetic training data.
233        let cols = vec!["foo", "bar"];
234        let err = missing_columns_error(Path::new("/tmp/x.parquet"), &cols);
235        let msg = err.to_string();
236        assert!(msg.contains("/tmp/x.parquet"), "error must name the dataset: {msg}");
237        assert!(msg.contains("foo"), "error must name the columns it found: {msg}");
238    }
239
240    #[test]
241    fn test_record_batch_to_training_batch_success() {
242        let record_batch = make_test_record_batch();
243        let schema = make_test_schema();
244        let result = record_batch_to_training_batch(&record_batch, &schema, "input", "target");
245        assert!(result.is_ok());
246        let batch = result.expect("operation should succeed");
247        assert_eq!(batch.inputs.data().len(), 4);
248        assert_eq!(batch.targets.data().len(), 4);
249    }
250
251    #[test]
252    fn test_record_batch_to_training_batch_invalid_input_column() {
253        let record_batch = make_test_record_batch();
254        let schema = make_test_schema();
255        let result =
256            record_batch_to_training_batch(&record_batch, &schema, "nonexistent", "target");
257        assert!(result.is_err());
258    }
259
260    #[test]
261    fn test_record_batch_to_training_batch_invalid_target_column() {
262        let record_batch = make_test_record_batch();
263        let schema = make_test_schema();
264        let result = record_batch_to_training_batch(&record_batch, &schema, "input", "nonexistent");
265        assert!(result.is_err());
266    }
267
268    #[test]
269    fn test_record_batch_with_float64() {
270        let schema = Arc::new(Schema::new(vec![
271            Field::new("x", DataType::Float64, false),
272            Field::new("y", DataType::Float64, false),
273        ]));
274        let input = Float64Array::from(vec![1.0, 2.0, 3.0]);
275        let target = Float64Array::from(vec![0.0, 1.0, 2.0]);
276        let record_batch =
277            RecordBatch::try_new(schema.clone(), vec![Arc::new(input), Arc::new(target)])
278                .expect("conversion should succeed");
279
280        let result = record_batch_to_training_batch(&record_batch, &schema, "x", "y");
281        assert!(result.is_ok());
282    }
283
284    #[test]
285    fn test_record_batch_with_int32() {
286        let schema = Arc::new(Schema::new(vec![
287            Field::new("features", DataType::Int32, false),
288            Field::new("labels", DataType::Int32, false),
289        ]));
290        let input = Int32Array::from(vec![1, 2, 3]);
291        let target = Int32Array::from(vec![0, 1, 0]);
292        let record_batch =
293            RecordBatch::try_new(schema.clone(), vec![Arc::new(input), Arc::new(target)])
294                .expect("conversion should succeed");
295
296        let result = record_batch_to_training_batch(&record_batch, &schema, "features", "labels");
297        assert!(result.is_ok());
298    }
299
300    #[test]
301    fn test_column_pair_fields() {
302        let pair = ColumnPair { input_name: "input", target_name: "target" };
303        assert_eq!(pair.input_name, "input");
304        assert_eq!(pair.target_name, "target");
305    }
306}