use std::path::Path;
const FIRST: [&str; 3] = ["train", "validation", "test"];
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Splits {
pub split: Option<String>,
pub others: Vec<String>,
pub caches: usize,
}
pub fn is_cache(name: &str) -> bool {
name.starts_with("cache-")
}
pub fn split_of(name: &str) -> Option<&str> {
let stem = name.rsplit_once('.').map_or(name, |(stem, _)| stem);
let stem = without_shard(stem);
let (builder, split) = stem.rsplit_once('-')?;
let word = |c: char| c.is_ascii_alphanumeric() || c == '_' || c == '.';
(!builder.is_empty() && !split.is_empty() && split.chars().all(word)).then_some(split)
}
fn without_shard(stem: &str) -> &str {
let digits = |s: &str| !s.is_empty() && s.bytes().all(|b| b.is_ascii_digit());
stem.rsplit_once("-of-")
.filter(|(_, total)| digits(total))
.and_then(|(head, _)| head.rsplit_once('-'))
.filter(|(_, index)| digits(index))
.map_or(stem, |(head, _)| head)
}
pub fn choose(names: &[&str], table: Option<&str>) -> Result<(Vec<usize>, Splits), String> {
let data: Vec<usize> = (0..names.len()).filter(|&i| !is_cache(names[i])).collect();
let caches = names.len() - data.len();
if data.is_empty() {
return Err(format!(
"this directory holds only the cache files map() writes ({caches} cache-*.arrow), and no split"
));
}
let by_split: Option<Vec<(usize, &str)>> = data
.iter()
.map(|&i| split_of(names[i]).map(|split| (i, split)))
.collect();
let Some(by_split) = by_split else {
if let Some(table) = table {
return Err(format!(
"--table {table}: this directory's files name no splits to pick from"
));
}
let splits = Splits {
caches,
..Splits::default()
};
return Ok((data, splits));
};
let mut splits: Vec<&str> = by_split.iter().map(|(_, split)| *split).collect();
splits.sort_unstable();
splits.dedup();
let mut picked = pick(&splits, table)?;
let files = by_split
.iter()
.filter(|(_, split)| Some(*split) == picked.split.as_deref())
.map(|(i, _)| *i)
.collect();
picked.caches = caches;
Ok((files, picked))
}
const CACHE_MARKERS: [&str; 2] = ["dataset_info.json", "state.json"];
pub fn is_cache_dir(dir: &Path) -> bool {
CACHE_MARKERS.iter().any(|name| dir.join(name).is_file())
}
pub fn cache_splits(dir: &Path) -> Vec<String> {
if !is_cache_dir(dir) {
return Vec::new();
}
let Ok(read) = std::fs::read_dir(dir) else {
return Vec::new();
};
let names: Vec<String> = read
.flatten()
.filter_map(|e| e.file_name().into_string().ok())
.filter(|n| n.to_ascii_lowercase().ends_with(".arrow") && !is_cache(n))
.collect();
let mut splits: Vec<&str> = Vec::new();
for name in &names {
match split_of(name) {
Some(split) if !splits.contains(&split) => splits.push(split),
Some(_) => {}
None => return Vec::new(),
}
}
splits.sort();
pick(&splits, None).map_or_else(
|_| Vec::new(),
|picked| picked.split.into_iter().chain(picked.others).collect(),
)
}
pub fn split_place(path: &Path) -> Option<(std::path::PathBuf, String)> {
if path.exists() {
return None;
}
let dir = path.parent().filter(|d| !d.as_os_str().is_empty())?;
let name = path.file_name()?.to_str()?;
cache_splits(dir)
.into_iter()
.find(|split| split == name)
.map(|split| (dir.to_path_buf(), split))
}
pub fn pick(listed: &[&str], table: Option<&str>) -> Result<Splits, String> {
let mut offered: Vec<&str> = listed.to_vec();
offered.sort_by_key(|split| {
FIRST
.iter()
.position(|first| first == split)
.unwrap_or(FIRST.len())
});
let chosen = match table {
Some(table) => *offered.iter().find(|s| **s == table).ok_or_else(|| {
format!(
"No split named {table}; this directory holds {}",
offered.join(", ")
)
})?,
None => *offered
.first()
.ok_or_else(|| "this directory names no splits".to_string())?,
};
Ok(Splits {
split: Some(chosen.to_string()),
others: offered
.iter()
.filter(|s| **s != chosen)
.map(|s| s.to_string())
.collect(),
caches: 0,
})
}
pub fn dataset_dict(dir: &Path) -> Option<Vec<String>> {
let text = std::fs::read_to_string(dir.join(DATASET_DICT)).ok()?;
let splits = dict_splits(&text)?;
splits
.iter()
.all(|split| dir.join(split).is_dir())
.then_some(splits)
}
pub const DATASET_DICT: &str = "dataset_dict.json";
pub fn dict_splits(text: &str) -> Option<Vec<String>> {
let value: serde_json::Value = serde_json::from_str(text).ok()?;
let splits: Vec<String> = value
.get("splits")?
.as_array()?
.iter()
.map(|split| split.as_str().map(str::to_string))
.collect::<Option<_>>()?;
let plain = |split: &String| {
!split.is_empty()
&& split
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '.')
&& split != "."
&& split != ".."
};
(!splits.is_empty() && splits.iter().all(plain)).then_some(splits)
}
#[cfg(test)]
mod tests;