dspy_rs/data/
dataloader.rs1use 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 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 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}