Skip to main content

webdataset_core/
sample.rs

1//! The [`Sample`] type: one training example, assembled from the files in a
2//! tar archive that share a basename.
3
4use indexmap::map::Iter as MapIter;
5
6use crate::fields::Fields;
7
8use crate::error::{Error, Result};
9use crate::prelude::*;
10use crate::value::Value;
11
12/// The field holding the sample's basename within its shard.
13pub const KEY: &str = "__key__";
14/// The field holding the URL of the shard the sample came from.
15pub const URL: &str = "__url__";
16/// The field holding the local path of the shard, when it was read from disk.
17pub const LOCAL_PATH: &str = "__local_path__";
18/// A truthy value in this field marks the sample as unusable.
19pub const BAD: &str = "__bad__";
20
21/// The prefix and suffix that mark a field as metadata rather than data.
22pub const META_PREFIX: &str = "__";
23/// See [`META_PREFIX`].
24pub const META_SUFFIX: &str = "__";
25
26/// One training example: an insertion-ordered map from file extension to value.
27///
28/// Fields whose names start with `__` are metadata and are skipped by decoding,
29/// encoding, and key extraction.
30///
31/// ```
32/// use webdataset_core::{Sample, Value};
33///
34/// let mut sample = Sample::with_key("image0001");
35/// sample.insert("cls", Value::Int(7));
36/// sample.insert("txt", Value::Text("a caption".into()));
37///
38/// assert_eq!(sample.key(), Some("image0001"));
39/// assert_eq!(sample.get("cls").and_then(Value::as_i64), Some(7));
40/// assert_eq!(sample.field_names(), vec!["cls", "txt"]);
41/// ```
42#[derive(Debug, Clone, Default, PartialEq)]
43pub struct Sample {
44    fields: Fields,
45}
46
47impl Sample {
48    /// An empty sample with no fields at all.
49    pub fn new() -> Sample {
50        Sample::default()
51    }
52
53    /// An empty sample carrying only its `__key__`.
54    pub fn with_key(key: impl Into<String>) -> Sample {
55        let mut s = Sample::new();
56        s.insert(KEY, Value::Text(key.into()));
57        s
58    }
59
60    /// The sample's `__key__`, if it has one.
61    pub fn key(&self) -> Option<&str> {
62        self.get(KEY).and_then(Value::as_str)
63    }
64
65    /// The URL of the shard this sample came from.
66    pub fn url(&self) -> Option<&str> {
67        self.get(URL).and_then(Value::as_str)
68    }
69
70    /// The on-disk path of the shard, when it was read or cached locally.
71    pub fn local_path(&self) -> Option<&str> {
72        self.get(LOCAL_PATH).and_then(Value::as_str)
73    }
74
75    /// Set the `__key__` field.
76    pub fn set_key(&mut self, key: impl Into<String>) {
77        self.insert(KEY, Value::Text(key.into()));
78    }
79
80    /// Set the `__url__` field.
81    pub fn set_url(&mut self, url: impl Into<String>) {
82        self.insert(URL, Value::Text(url.into()));
83    }
84
85    /// Look up a field.
86    pub fn get(&self, name: &str) -> Option<&Value> {
87        self.fields.get(name)
88    }
89
90    /// Look up a field for modification.
91    pub fn get_mut(&mut self, name: &str) -> Option<&mut Value> {
92        self.fields.get_mut(name)
93    }
94
95    /// Look up the first field that exists among several alternatives.
96    ///
97    /// Alternatives are given either as a slice or, following the Python API,
98    /// as a single `;`-separated string (see [`Sample::get_first_spec`]).
99    pub fn get_first(&self, names: &[&str]) -> Option<&Value> {
100        names.iter().find_map(|n| self.get(n))
101    }
102
103    /// Look up the first field named by a `;`-separated alternation such as
104    /// `"png;jpg;jpeg"`.
105    pub fn get_first_spec(&self, spec: &str) -> Option<&Value> {
106        spec.split(';').find_map(|n| self.get(n))
107    }
108
109    /// Like [`Sample::get_first_spec`] but reports which alternatives were tried.
110    pub fn require_first_spec(&self, spec: &str) -> Result<&Value> {
111        self.get_first_spec(spec).ok_or_else(|| Error::MissingKey {
112            wanted: spec.split(';').map(str::to_string).collect(),
113            available: self.fields.keys().cloned().collect(),
114        })
115    }
116
117    /// Insert a field, returning the value it replaced.
118    pub fn insert(&mut self, name: impl Into<String>, value: impl Into<Value>) -> Option<Value> {
119        self.fields.insert(name.into(), value.into())
120    }
121
122    /// Remove a field, returning its value.
123    pub fn remove(&mut self, name: &str) -> Option<Value> {
124        self.fields.shift_remove(name)
125    }
126
127    /// Whether a field is present.
128    pub fn contains_key(&self, name: &str) -> bool {
129        self.fields.contains_key(name)
130    }
131
132    /// The number of fields, metadata included.
133    pub fn len(&self) -> usize {
134        self.fields.len()
135    }
136
137    /// Whether the sample has no fields at all.
138    pub fn is_empty(&self) -> bool {
139        self.fields.is_empty()
140    }
141
142    /// Iterate over all fields in insertion order.
143    pub fn iter(&self) -> MapIter<'_, String, Value> {
144        self.fields.iter()
145    }
146
147    /// All field names, metadata included.
148    pub fn keys(&self) -> impl Iterator<Item = &str> {
149        self.fields.keys().map(String::as_str)
150    }
151
152    /// The names of the data (non-`__`) fields.
153    pub fn field_names(&self) -> Vec<&str> {
154        self.keys().filter(|k| !is_meta(k)).collect()
155    }
156
157    /// Borrow the underlying map.
158    pub fn as_map(&self) -> &Fields {
159        &self.fields
160    }
161
162    /// Consume the sample and return the underlying map.
163    pub fn into_map(self) -> Fields {
164        self.fields
165    }
166
167    /// Whether this sample should be passed downstream.
168    ///
169    /// Mirrors `valid_sample` in the Python implementation: a sample is valid
170    /// when it has at least one field and is not marked `__bad__`.
171    pub fn is_valid(&self) -> bool {
172        !self.fields.is_empty() && !matches!(self.get(BAD), Some(Value::Bool(true)))
173    }
174
175    /// Rename a field, preserving its position in the field order.
176    pub fn rename(&mut self, from: &str, to: impl Into<String>) -> bool {
177        let Some(index) = self.fields.get_index_of(from) else {
178            return false;
179        };
180        let (_, value) = self.fields.shift_remove_index(index).expect("index was just looked up");
181        self.fields.shift_insert(index, to.into(), value);
182        true
183    }
184}
185
186impl From<Fields> for Sample {
187    fn from(fields: Fields) -> Sample {
188        Sample { fields }
189    }
190}
191
192impl FromIterator<(String, Value)> for Sample {
193    fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Sample {
194        Sample { fields: iter.into_iter().collect() }
195    }
196}
197
198impl<'a> IntoIterator for &'a Sample {
199    type Item = (&'a String, &'a Value);
200    type IntoIter = MapIter<'a, String, Value>;
201
202    fn into_iter(self) -> Self::IntoIter {
203        self.fields.iter()
204    }
205}
206
207impl IntoIterator for Sample {
208    type Item = (String, Value);
209    type IntoIter = indexmap::map::IntoIter<String, Value>;
210
211    fn into_iter(self) -> Self::IntoIter {
212        self.fields.into_iter()
213    }
214}
215
216impl core::ops::Index<&str> for Sample {
217    type Output = Value;
218
219    fn index(&self, name: &str) -> &Value {
220        self.get(name).unwrap_or_else(|| panic!("no field {name:?} in sample; have {:?}", self.field_names()))
221    }
222}
223
224/// Whether a field name denotes metadata (`__key__`, `__url__`, ...).
225pub fn is_meta(name: &str) -> bool {
226    name.starts_with(META_PREFIX)
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232
233    #[test]
234    fn separates_metadata_from_data_fields() {
235        let mut s = Sample::with_key("k");
236        s.set_url("pipe:cat shard.tar");
237        s.insert("png", Value::Bytes(vec![1, 2, 3].into()));
238        assert_eq!(s.field_names(), vec!["png"]);
239        assert_eq!(s.key(), Some("k"));
240        assert_eq!(s.url(), Some("pipe:cat shard.tar"));
241        assert_eq!(s.len(), 3);
242    }
243
244    #[test]
245    fn resolves_alternatives_in_order() {
246        let mut s = Sample::with_key("k");
247        s.insert("jpg", Value::Int(1));
248        s.insert("png", Value::Int(2));
249        assert_eq!(s.get_first_spec("png;jpg").and_then(Value::as_i64), Some(2));
250        assert_eq!(s.get_first_spec("cls;jpg").and_then(Value::as_i64), Some(1));
251        assert!(s.get_first_spec("cls;wnid").is_none());
252        assert!(s.require_first_spec("cls").is_err());
253    }
254
255    #[test]
256    fn treats_empty_and_bad_samples_as_invalid() {
257        assert!(!Sample::new().is_valid());
258        let mut s = Sample::with_key("k");
259        assert!(s.is_valid());
260        s.insert(BAD, Value::Bool(true));
261        assert!(!s.is_valid());
262    }
263
264    #[test]
265    fn renames_in_place() {
266        let mut s = Sample::with_key("k");
267        s.insert("a", Value::Int(1));
268        s.insert("b", Value::Int(2));
269        assert!(s.rename("a", "z"));
270        assert_eq!(s.field_names(), vec!["z", "b"]);
271        assert!(!s.rename("nope", "x"));
272    }
273}