datui_lib/formats/
hf_splits.rs1use 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;