Skip to main content

webdataset_core/
tensor.rs

1//! A minimal dense n-dimensional array.
2//!
3//! WebDataset stores tensors in two formats: NumPy's `.npy` and the 8-byte
4//! aligned `.ten` ("tenbin") format. Both are described by an element type, a
5//! shape, and a blob of C-ordered, native-endian element data, which is exactly
6//! what [`Tensor`] holds. Keeping the representation this simple means the
7//! crate does not force a particular array library on its users while still
8//! allowing zero-copy round trips through the archive formats.
9
10use bytes::Bytes;
11
12use crate::error::{Error, Result};
13use crate::prelude::*;
14
15/// The element type of a [`Tensor`].
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
17pub enum DType {
18    /// IEEE-754 binary16. Read and widened to `f64`; never produced by this
19    /// crate, since Rust has no stable 16-bit float.
20    F16,
21    /// IEEE-754 binary32, NumPy's `float32`.
22    F32,
23    /// IEEE-754 binary64, NumPy's `float64`.
24    F64,
25    /// Signed 8-bit.
26    I8,
27    /// Signed 16-bit.
28    I16,
29    /// Signed 32-bit.
30    I32,
31    /// Signed 64-bit.
32    I64,
33    /// Unsigned 8-bit, which is what a decoded image is made of.
34    U8,
35    /// Unsigned 16-bit.
36    U16,
37    /// Unsigned 32-bit.
38    U32,
39    /// Unsigned 64-bit.
40    U64,
41}
42
43impl DType {
44    /// Size of one element in bytes.
45    pub const fn size(self) -> usize {
46        match self {
47            DType::I8 | DType::U8 => 1,
48            DType::F16 | DType::I16 | DType::U16 => 2,
49            DType::F32 | DType::I32 | DType::U32 => 4,
50            DType::F64 | DType::I64 | DType::U64 => 8,
51        }
52    }
53
54    /// The NumPy long name, e.g. `"float32"`.
55    pub const fn long_name(self) -> &'static str {
56        match self {
57            DType::F16 => "float16",
58            DType::F32 => "float32",
59            DType::F64 => "float64",
60            DType::I8 => "int8",
61            DType::I16 => "int16",
62            DType::I32 => "int32",
63            DType::I64 => "int64",
64            DType::U8 => "uint8",
65            DType::U16 => "uint16",
66            DType::U32 => "uint32",
67            DType::U64 => "uint64",
68        }
69    }
70
71    /// The two character NumPy short name, e.g. `"f4"`.
72    pub const fn short_name(self) -> &'static str {
73        match self {
74            DType::F16 => "f2",
75            DType::F32 => "f4",
76            DType::F64 => "f8",
77            DType::I8 => "i1",
78            DType::I16 => "i2",
79            DType::I32 => "i4",
80            DType::I64 => "i8",
81            DType::U8 => "u1",
82            DType::U16 => "u2",
83            DType::U32 => "u4",
84            DType::U64 => "u8",
85        }
86    }
87
88    /// Parse either the long (`"float32"`) or short (`"f4"`) NumPy name.
89    pub fn parse(name: &str) -> Result<DType> {
90        let dt = match name {
91            "float16" | "f2" => DType::F16,
92            "float32" | "f4" => DType::F32,
93            "float64" | "f8" => DType::F64,
94            "int8" | "i1" => DType::I8,
95            "int16" | "i2" => DType::I16,
96            "int32" | "i4" => DType::I32,
97            "int64" | "i8" => DType::I64,
98            "uint8" | "u1" => DType::U8,
99            "uint16" | "u2" => DType::U16,
100            "uint32" | "u4" => DType::U32,
101            "uint64" | "u8" => DType::U64,
102            other => return Err(Error::unsupported(format!("dtype {other}"))),
103        };
104        Ok(dt)
105    }
106
107    /// Whether this is a floating point type.
108    pub const fn is_float(self) -> bool {
109        matches!(self, DType::F16 | DType::F32 | DType::F64)
110    }
111}
112
113/// A dense, C-ordered array of numbers.
114///
115/// Element bytes are stored in native endianness, matching what NumPy and the
116/// tenbin format write on the same machine.
117#[derive(Debug, Clone, PartialEq, Eq)]
118pub struct Tensor {
119    dtype: DType,
120    shape: Vec<usize>,
121    data: Bytes,
122}
123
124impl Tensor {
125    /// Build a tensor from raw native-endian element bytes.
126    ///
127    /// Fails when `data` does not hold exactly `shape.product() * dtype.size()`
128    /// bytes.
129    pub fn new(dtype: DType, shape: Vec<usize>, data: impl Into<Bytes>) -> Result<Tensor> {
130        let data = data.into();
131        let expected = shape.iter().product::<usize>() * dtype.size();
132        if data.len() != expected {
133            return Err(Error::format(format!(
134                "tensor data is {} bytes but shape {shape:?} of {} needs {expected}",
135                data.len(),
136                dtype.long_name()
137            )));
138        }
139        Ok(Tensor { dtype, shape, data })
140    }
141
142    /// Build a one-dimensional `f32` tensor.
143    pub fn from_f32(values: &[f32]) -> Tensor {
144        Self::from_f32_shaped(values, vec![values.len()])
145    }
146
147    /// Build an `f32` tensor with an explicit shape.
148    ///
149    /// # Panics
150    /// Panics if `shape` does not describe exactly `values.len()` elements.
151    pub fn from_f32_shaped(values: &[f32], shape: Vec<usize>) -> Tensor {
152        assert_eq!(shape.iter().product::<usize>(), values.len(), "shape does not match value count");
153        let mut data = Vec::with_capacity(values.len() * 4);
154        for v in values {
155            data.extend_from_slice(&v.to_ne_bytes());
156        }
157        Tensor { dtype: DType::F32, shape, data: data.into() }
158    }
159
160    /// Build a `u8` tensor with an explicit shape.
161    ///
162    /// # Panics
163    /// Panics if `shape` does not describe exactly `values.len()` elements.
164    pub fn from_u8_shaped(values: Vec<u8>, shape: Vec<usize>) -> Tensor {
165        assert_eq!(shape.iter().product::<usize>(), values.len(), "shape does not match value count");
166        Tensor { dtype: DType::U8, shape, data: values.into() }
167    }
168
169    /// The element type.
170    pub fn dtype(&self) -> DType {
171        self.dtype
172    }
173
174    /// The shape, outermost dimension first.
175    pub fn shape(&self) -> &[usize] {
176        &self.shape
177    }
178
179    /// The number of dimensions.
180    pub fn ndim(&self) -> usize {
181        self.shape.len()
182    }
183
184    /// The total number of elements.
185    pub fn numel(&self) -> usize {
186        self.shape.iter().product()
187    }
188
189    /// The raw native-endian element bytes.
190    pub fn data(&self) -> &Bytes {
191        &self.data
192    }
193
194    /// Consume the tensor and return its raw element bytes.
195    pub fn into_data(self) -> Bytes {
196        self.data
197    }
198
199    /// Change the shape, keeping the element data.
200    ///
201    /// Fails when the new shape describes a different number of elements.
202    pub fn reshape(&self, shape: Vec<usize>) -> Result<Tensor> {
203        Tensor::new(self.dtype, shape, self.data.clone())
204    }
205
206    /// Copy the elements out as `f64`, whatever the stored type.
207    pub fn to_f64_vec(&self) -> Vec<f64> {
208        let n = self.numel();
209        let mut out = Vec::with_capacity(n);
210        let d = &self.data[..];
211        for i in 0..n {
212            out.push(self.element_at(d, i));
213        }
214        out
215    }
216
217    /// Copy the elements out as `f32`, whatever the stored type.
218    pub fn to_f32_vec(&self) -> Vec<f32> {
219        self.to_f64_vec().into_iter().map(|v| v as f32).collect()
220    }
221
222    /// Read a single element by flat index, widened to `f64`.
223    ///
224    /// # Panics
225    /// Panics if `index` is out of bounds.
226    pub fn get(&self, index: usize) -> f64 {
227        assert!(index < self.numel(), "index {index} out of bounds");
228        self.element_at(&self.data[..], index)
229    }
230
231    fn element_at(&self, d: &[u8], i: usize) -> f64 {
232        let s = self.dtype.size();
233        let b = &d[i * s..(i + 1) * s];
234        match self.dtype {
235            DType::U8 => b[0] as f64,
236            DType::I8 => b[0] as i8 as f64,
237            DType::U16 => u16::from_ne_bytes([b[0], b[1]]) as f64,
238            DType::I16 => i16::from_ne_bytes([b[0], b[1]]) as f64,
239            DType::F16 => f16_to_f64(u16::from_ne_bytes([b[0], b[1]])),
240            DType::U32 => u32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
241            DType::I32 => i32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
242            DType::F32 => f32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
243            DType::U64 => u64::from_ne_bytes(b.try_into().expect("8 bytes")) as f64,
244            DType::I64 => i64::from_ne_bytes(b.try_into().expect("8 bytes")) as f64,
245            DType::F64 => f64::from_ne_bytes(b.try_into().expect("8 bytes")),
246        }
247    }
248}
249
250/// Widen an IEEE-754 binary16 bit pattern to `f64`.
251fn f16_to_f64(bits: u16) -> f64 {
252    let sign = if bits & 0x8000 != 0 { -1.0f64 } else { 1.0 };
253    let exponent = ((bits >> 10) & 0x1f) as i32;
254    let mantissa = (bits & 0x3ff) as f64;
255    match exponent {
256        0 => sign * mantissa * exp2(-24),
257        0x1f if mantissa == 0.0 => sign * f64::INFINITY,
258        0x1f => f64::NAN,
259        _ => sign * (1.0 + mantissa / 1024.0) * exp2(exponent - 15),
260    }
261}
262
263/// Two raised to `n`, computed without `std`'s float intrinsics.
264///
265/// Powers of two are exactly representable, and binary16 only ever needs
266/// exponents in `-24..=16`, so repeated halving or doubling is exact.
267fn exp2(n: i32) -> f64 {
268    let mut out = 1.0f64;
269    for _ in 0..n.abs() {
270        if n >= 0 {
271            out *= 2.0;
272        } else {
273            out /= 2.0;
274        }
275    }
276    out
277}
278
279#[cfg(test)]
280mod tests {
281    use super::*;
282
283    #[test]
284    fn round_trips_f32_values() {
285        let t = Tensor::from_f32_shaped(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
286        assert_eq!(t.shape(), &[2, 3]);
287        assert_eq!(t.numel(), 6);
288        assert_eq!(t.dtype(), DType::F32);
289        assert_eq!(t.to_f64_vec(), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
290        assert_eq!(t.get(4), 5.0);
291    }
292
293    #[test]
294    fn rejects_mismatched_shapes() {
295        let err = Tensor::new(DType::F32, vec![3], vec![0u8; 8]);
296        assert!(err.is_err());
297    }
298
299    #[test]
300    fn reshapes_without_copying_semantics() {
301        let t = Tensor::from_f32(&[1.0, 2.0, 3.0, 4.0]);
302        let r = t.reshape(vec![2, 2]).unwrap();
303        assert_eq!(r.shape(), &[2, 2]);
304        assert!(t.reshape(vec![3, 3]).is_err());
305    }
306
307    #[test]
308    fn parses_dtype_names() {
309        assert_eq!(DType::parse("float32").unwrap(), DType::F32);
310        assert_eq!(DType::parse("f4").unwrap(), DType::F32);
311        assert_eq!(DType::parse("u1").unwrap(), DType::U8);
312        assert!(DType::parse("complex64").is_err());
313    }
314}