Skip to main content

dspy_rs/data/
dataloader.rs

1use anyhow::Result;
2use arrow::array::{Array, StringArray};
3use csv::{ReaderBuilder, WriterBuilder};
4use hf_hub::api::sync::Api;
5use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
6use rayon::prelude::*;
7use reqwest;
8use std::fs;
9use std::io::Cursor;
10use std::{collections::HashMap, path::Path};
11
12use crate::{Example, is_url, string_record_to_example};
13
14pub struct DataLoader;
15
16impl DataLoader {
17    pub fn load_json(
18        path: &str,
19        lines: bool,
20        input_keys: Vec<String>,
21        output_keys: Vec<String>,
22    ) -> Result<Vec<Example>> {
23        let data = if is_url(path) {
24            let response = reqwest::blocking::get(path)?;
25            response.text()?
26        } else {
27            fs::read_to_string(path)?
28        };
29
30        let examples: Vec<Example> = if lines {
31            let lines = data.lines().collect::<Vec<&str>>();
32
33            lines
34                .par_iter()
35                .map(|line| {
36                    Example::new(
37                        serde_json::from_str(line).unwrap(),
38                        input_keys.clone(),
39                        output_keys.clone(),
40                    )
41                })
42                .collect()
43        } else {
44            vec![Example::new(
45                serde_json::from_str(&data).unwrap(),
46                input_keys.clone(),
47                output_keys.clone(),
48            )]
49        };
50        Ok(examples)
51    }
52
53    pub fn save_json(path: &str, examples: Vec<Example>, lines: bool) -> Result<()> {
54        let data = if lines {
55            examples
56                .into_iter()
57                .map(|example| serde_json::to_string(&example).unwrap())
58                .collect::<Vec<String>>()
59                .join("\n")
60        } else {
61            serde_json::to_string(&examples).unwrap()
62        };
63        fs::write(path, data)?;
64        Ok(())
65    }
66
67    pub fn load_csv(
68        path: &str,
69        delimiter: char,
70        input_keys: Vec<String>,
71        output_keys: Vec<String>,
72        has_headers: bool,
73    ) -> Result<Vec<Example>> {
74        let records = if is_url(path) {
75            let response = reqwest::blocking::get(path)?.bytes()?.to_vec();
76            let cursor = Cursor::new(response);
77
78            let records: Vec<_> = ReaderBuilder::new()
79                .delimiter(delimiter as u8)
80                .has_headers(has_headers)
81                .from_reader(cursor)
82                .into_records()
83                .collect::<Result<Vec<_>, _>>()?;
84
85            records
86        } else {
87            let records: Vec<_> = ReaderBuilder::new()
88                .delimiter(delimiter as u8)
89                .has_headers(has_headers)
90                .from_path(path)?
91                .into_records()
92                .collect::<Result<Vec<_>, _>>()?;
93
94            records
95        };
96
97        let examples = records
98            .par_iter()
99            .map(|row| {
100                string_record_to_example(row.clone(), input_keys.clone(), output_keys.clone())
101            })
102            .collect();
103
104        Ok(examples)
105    }
106
107    pub fn save_csv(path: &str, examples: Vec<Example>, delimiter: char) -> Result<()> {
108        let mut writer = WriterBuilder::new()
109            .delimiter(delimiter as u8)
110            .from_path(path)?;
111        let headers = examples[0].data.keys().cloned().collect::<Vec<String>>();
112        writer.write_record(&headers)?;
113        for example in examples {
114            writer.write_record(
115                example
116                    .data
117                    .values()
118                    .cloned()
119                    .map(|value| value.to_string())
120                    .collect::<Vec<String>>(),
121            )?;
122        }
123        Ok(())
124    }
125
126    #[allow(clippy::while_let_on_iterator)]
127    pub fn load_parquet(
128        path: &str,
129        input_keys: Vec<String>,
130        output_keys: Vec<String>,
131    ) -> Result<Vec<Example>> {
132        let file_path = Path::new(path);
133
134        let file = fs::File::open(file_path)?;
135        let builder = ParquetRecordBatchReaderBuilder::try_new(file)?;
136        let mut record_batch_reader = builder.build()?;
137
138        let mut examples = Vec::new();
139        while let Some(record_batch_result) = record_batch_reader.next() {
140            let record_batch = record_batch_result?;
141            let schema = record_batch.schema();
142            let num_rows = record_batch.num_rows();
143
144            // Process each row
145            for row_idx in 0..num_rows {
146                let mut data = HashMap::new();
147
148                for col_idx in 0..record_batch.num_columns() {
149                    let column = record_batch.column(col_idx);
150                    let column_name = schema.field(col_idx).name();
151
152                    if let Some(string_array) = column.as_any().downcast_ref::<StringArray>()
153                        && !string_array.is_null(row_idx)
154                    {
155                        let value = string_array.value(row_idx);
156                        data.insert(column_name.to_string(), value.to_string().into());
157                    }
158                }
159
160                if !data.is_empty() {
161                    examples.push(Example::new(data, input_keys.clone(), output_keys.clone()));
162                }
163            }
164        }
165        Ok(examples)
166    }
167
168    pub fn load_hf(
169        dataset_id: &str,
170        input_keys: Vec<String>,
171        output_keys: Vec<String>,
172        subset: &str,
173        split: &str,
174        verbose: bool,
175    ) -> Result<Vec<Example>> {
176        let api = Api::new()?;
177        let repo = api.dataset(dataset_id.to_string());
178
179        // Get metadata and list of files using info()
180        let metadata = repo.info()?;
181        let files: Vec<&str> = metadata
182            .siblings
183            .iter()
184            .map(|sib| sib.rfilename.as_str())
185            .collect();
186
187        let examples: Vec<_> = files
188            .par_iter()
189            .filter_map(|file: &&str| {
190                let extension = file.split(".").last().unwrap();
191                if !file.ends_with(".parquet")
192                    && !extension.ends_with("json")
193                    && !extension.ends_with("jsonl")
194                    && !extension.ends_with("csv")
195                {
196                    if verbose {
197                        println!("Skipping file by extension: {file}");
198                    }
199                    return None;
200                }
201
202                if (!subset.is_empty() && !file.contains(subset))
203                    || (!split.is_empty() && !file.contains(split))
204                {
205                    if verbose {
206                        println!("Skipping file by subset or split: {file}");
207                    }
208                    return None;
209                }
210
211                let file_path = repo.get(file).unwrap();
212                let os_str = file_path.as_os_str().to_str().unwrap();
213
214                if verbose {
215                    println!("Loading file: {os_str}");
216                }
217
218                if os_str.ends_with(".parquet") {
219                    DataLoader::load_parquet(os_str, input_keys.clone(), output_keys.clone()).ok()
220                } else if os_str.ends_with(".json") || os_str.ends_with(".jsonl") {
221                    let is_jsonl = os_str.ends_with(".jsonl");
222                    DataLoader::load_json(os_str, is_jsonl, input_keys.clone(), output_keys.clone())
223                        .ok()
224                } else if os_str.ends_with(".csv") {
225                    DataLoader::load_csv(os_str, ',', input_keys.clone(), output_keys.clone(), true)
226                        .ok()
227                } else {
228                    None
229                }
230            })
231            .flatten()
232            .collect();
233
234        if verbose {
235            println!("Loaded {} examples", examples.len());
236        }
237        Ok(examples)
238    }
239}