1use std::path::Path;
11
12const FIRST: [&str; 3] = ["train", "validation", "test"];
15
16#[derive(Debug, Clone, Default, PartialEq, Eq)]
18pub struct Splits {
19 pub split: Option<String>,
22 pub others: Vec<String>,
24 pub caches: usize,
26}
27
28pub fn is_cache(name: &str) -> bool {
30 name.starts_with("cache-")
31}
32
33pub fn split_of(name: &str) -> Option<&str> {
36 let stem = name.rsplit_once('.').map_or(name, |(stem, _)| stem);
37 let stem = without_shard(stem);
38 let (builder, split) = stem.rsplit_once('-')?;
39 let word = |c: char| c.is_ascii_alphanumeric() || c == '_' || c == '.';
40 (!builder.is_empty() && !split.is_empty() && split.chars().all(word)).then_some(split)
41}
42
43fn without_shard(stem: &str) -> &str {
45 let digits = |s: &str| !s.is_empty() && s.bytes().all(|b| b.is_ascii_digit());
46 stem.rsplit_once("-of-")
47 .filter(|(_, total)| digits(total))
48 .and_then(|(head, _)| head.rsplit_once('-'))
49 .filter(|(_, index)| digits(index))
50 .map_or(stem, |(head, _)| head)
51}
52
53pub fn choose(names: &[&str], table: Option<&str>) -> Result<(Vec<usize>, Splits), String> {
60 let data: Vec<usize> = (0..names.len()).filter(|&i| !is_cache(names[i])).collect();
61 let caches = names.len() - data.len();
62 if data.is_empty() {
63 return Err(format!(
64 "this directory holds only the cache files map() writes ({caches} cache-*.arrow), and no split"
65 ));
66 }
67 let by_split: Option<Vec<(usize, &str)>> = data
68 .iter()
69 .map(|&i| split_of(names[i]).map(|split| (i, split)))
70 .collect();
71 let Some(by_split) = by_split else {
72 if let Some(table) = table {
73 return Err(format!(
74 "--table {table}: this directory's files name no splits to pick from"
75 ));
76 }
77 let splits = Splits {
78 caches,
79 ..Splits::default()
80 };
81 return Ok((data, splits));
82 };
83 let mut splits: Vec<&str> = by_split.iter().map(|(_, split)| *split).collect();
84 splits.sort_unstable();
85 splits.dedup();
86 let mut picked = pick(&splits, table)?;
87 let files = by_split
88 .iter()
89 .filter(|(_, split)| Some(*split) == picked.split.as_deref())
90 .map(|(i, _)| *i)
91 .collect();
92 picked.caches = caches;
93 Ok((files, picked))
94}
95
96const CACHE_MARKERS: [&str; 2] = ["dataset_info.json", "state.json"];
98
99pub fn is_cache_dir(dir: &Path) -> bool {
101 CACHE_MARKERS.iter().any(|name| dir.join(name).is_file())
102}
103
104pub fn cache_splits(dir: &Path) -> Vec<String> {
108 if !is_cache_dir(dir) {
109 return Vec::new();
110 }
111 let Ok(read) = std::fs::read_dir(dir) else {
112 return Vec::new();
113 };
114 let names: Vec<String> = read
115 .flatten()
116 .filter_map(|e| e.file_name().into_string().ok())
117 .filter(|n| n.to_ascii_lowercase().ends_with(".arrow") && !is_cache(n))
118 .collect();
119 let mut splits: Vec<&str> = Vec::new();
120 for name in &names {
121 match split_of(name) {
122 Some(split) if !splits.contains(&split) => splits.push(split),
123 Some(_) => {}
124 None => return Vec::new(),
126 }
127 }
128 splits.sort();
129 pick(&splits, None).map_or_else(
130 |_| Vec::new(),
131 |picked| picked.split.into_iter().chain(picked.others).collect(),
132 )
133}
134
135pub fn split_place(path: &Path) -> Option<(std::path::PathBuf, String)> {
138 if path.exists() {
139 return None;
140 }
141 let dir = path.parent().filter(|d| !d.as_os_str().is_empty())?;
142 let name = path.file_name()?.to_str()?;
143 cache_splits(dir)
144 .into_iter()
145 .find(|split| split == name)
146 .map(|split| (dir.to_path_buf(), split))
147}
148
149pub fn pick(listed: &[&str], table: Option<&str>) -> Result<Splits, String> {
153 let mut offered: Vec<&str> = listed.to_vec();
154 offered.sort_by_key(|split| {
155 FIRST
156 .iter()
157 .position(|first| first == split)
158 .unwrap_or(FIRST.len())
159 });
160 let chosen = match table {
161 Some(table) => *offered.iter().find(|s| **s == table).ok_or_else(|| {
162 format!(
163 "No split named {table}; this directory holds {}",
164 offered.join(", ")
165 )
166 })?,
167 None => *offered
168 .first()
169 .ok_or_else(|| "this directory names no splits".to_string())?,
170 };
171 Ok(Splits {
172 split: Some(chosen.to_string()),
173 others: offered
174 .iter()
175 .filter(|s| **s != chosen)
176 .map(|s| s.to_string())
177 .collect(),
178 caches: 0,
179 })
180}
181
182pub fn dataset_dict(dir: &Path) -> Option<Vec<String>> {
186 let text = std::fs::read_to_string(dir.join(DATASET_DICT)).ok()?;
187 let splits = dict_splits(&text)?;
188 splits
189 .iter()
190 .all(|split| dir.join(split).is_dir())
191 .then_some(splits)
192}
193
194pub const DATASET_DICT: &str = "dataset_dict.json";
196
197pub fn dict_splits(text: &str) -> Option<Vec<String>> {
200 let value: serde_json::Value = serde_json::from_str(text).ok()?;
201 let splits: Vec<String> = value
202 .get("splits")?
203 .as_array()?
204 .iter()
205 .map(|split| split.as_str().map(str::to_string))
206 .collect::<Option<_>>()?;
207 let plain = |split: &String| {
208 !split.is_empty()
209 && split
210 .chars()
211 .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '.')
212 && split != "."
213 && split != ".."
214 };
215 (!splits.is_empty() && splits.iter().all(plain)).then_some(splits)
216}
217
218#[cfg(test)]
219mod tests {
220 use super::*;
221
222 #[test]
223 fn a_name_says_its_split_with_or_without_a_shard() {
224 assert_eq!(split_of("imdb-train.arrow"), Some("train"));
225 assert_eq!(split_of("imdb-test-00001-of-00004.arrow"), Some("test"));
226 assert_eq!(split_of("squad_v2-validation.arrow"), Some("validation"));
227 assert_eq!(split_of("wiki-40b-train_sft.arrow"), Some("train_sft"));
228 assert_eq!(split_of("data-00000-of-00003.arrow"), None, "save_to_disk");
229 assert_eq!(split_of("people.arrow"), None);
230 assert_eq!(split_of("-train.arrow"), None);
231 }
232
233 fn picked(names: &[&str], table: Option<&str>) -> (Vec<String>, Splits) {
234 let (files, splits) = choose(names, table).unwrap();
235 (
236 files.iter().map(|&i| names[i].to_string()).collect(),
237 splits,
238 )
239 }
240
241 #[test]
242 fn train_opens_and_the_other_splits_are_named() {
243 let names = [
244 "p-test.arrow",
245 "p-train-00000-of-00002.arrow",
246 "cache-0f3c.arrow",
247 "p-train-00001-of-00002.arrow",
248 "p-validation.arrow",
249 ];
250 let (files, splits) = picked(&names, None);
251 assert_eq!(
252 files,
253 [
254 "p-train-00000-of-00002.arrow",
255 "p-train-00001-of-00002.arrow"
256 ]
257 );
258 assert_eq!(
259 splits,
260 Splits {
261 split: Some("train".into()),
262 others: vec!["validation".into(), "test".into()],
263 caches: 1,
264 }
265 );
266 let (files, splits) = picked(&names, Some("validation"));
267 assert_eq!(files, ["p-validation.arrow"]);
268 assert_eq!(splits.others, ["train", "test"]);
269 let error = choose(&names, Some("dev")).unwrap_err();
270 assert!(error.contains("train, validation, test"), "{error}");
271 }
272
273 #[test]
274 fn without_train_the_first_split_opens() {
275 let (files, splits) = picked(&["x-zeta.arrow", "x-alpha.arrow"], None);
276 assert_eq!(files, ["x-alpha.arrow"]);
277 assert_eq!(splits.others, ["zeta"]);
278 }
279
280 #[test]
283 fn the_splits_datasets_names_come_first() {
284 let names = [
285 "p-extra.arrow",
286 "p-test.arrow",
287 "p-validation.arrow",
288 "p-a_more.arrow",
289 ];
290 let (files, splits) = picked(&names, None);
291 assert_eq!(files, ["p-validation.arrow"]);
292 assert_eq!(splits.others, ["test", "a_more", "extra"]);
293 let splits = pick(&["zz", "test", "aa", "train"], None).unwrap();
294 assert_eq!(splits.split.as_deref(), Some("train"));
295 assert_eq!(splits.others, ["test", "zz", "aa"]);
296 }
297
298 #[test]
301 fn a_dataset_dict_names_its_split_directories() {
302 assert_eq!(
303 dict_splits(r#"{"splits": ["train", "test"]}"#),
304 Some(vec!["train".to_string(), "test".to_string()])
305 );
306 for bad in [
307 r#"{"splits": []}"#,
308 r#"{"splits": ["../x"]}"#,
309 r#"{"splits": [".."]}"#,
310 r#"{"splits": ["a/b"]}"#,
311 r#"{"splits": [1]}"#,
312 r#"{"other": ["train"]}"#,
313 "not json",
314 ] {
315 assert_eq!(dict_splits(bad), None, "{bad}");
316 }
317 let dir = tempfile::tempdir().unwrap();
318 std::fs::write(
319 dir.path().join(DATASET_DICT),
320 r#"{"splits": ["train", "test"]}"#,
321 )
322 .unwrap();
323 std::fs::create_dir(dir.path().join("train")).unwrap();
324 assert_eq!(dataset_dict(dir.path()), None, "test/ is missing");
325 std::fs::create_dir(dir.path().join("test")).unwrap();
326 assert_eq!(
327 dataset_dict(dir.path()),
328 Some(vec!["train".to_string(), "test".to_string()])
329 );
330 }
331
332 #[test]
333 fn shards_that_name_no_split_are_one_table() {
334 let names = [
335 "data-00000-of-00002.arrow",
336 "data-00001-of-00002.arrow",
337 "cache-1.arrow",
338 ];
339 let (files, splits) = picked(&names, None);
340 assert_eq!(files.len(), 2);
341 assert_eq!((splits.split, splits.caches), (None, 1));
342 assert!(choose(&names, Some("train")).is_err());
343 assert!(choose(&["cache-1.arrow"], None).is_err());
344 }
345}