entrenar/config/train/batches/
parquet.rs1use 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
13struct ColumnPair<'a> {
15 input_name: &'a str,
16 target_name: &'a str,
17}
18
19fn 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
27fn 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
35fn 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
42fn 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
54fn 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
77fn 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
95pub 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 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 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}