Skip to main content

conspire/io/npy/
mod.rs

1#[cfg(test)]
2mod test;
3
4use crate::io::Write;
5use std::{
6    fs::File,
7    io::{self, BufReader, BufWriter, Error, ErrorKind, Read, Result},
8    mem::size_of,
9    path::Path,
10    str::from_utf8,
11};
12
13pub trait NpyType: Copy {
14    const DESCR: &'static str;
15    const SIZE: usize;
16    fn write_le(self, buffer: &mut Vec<u8>);
17    fn read_le(bytes: &[u8]) -> Self;
18    #[cfg(target_endian = "little")]
19    fn read_from<R: Read>(file: &mut R, count: usize) -> Result<Vec<Self>> {
20        let mut data = Vec::<Self>::with_capacity(count);
21        unsafe {
22            let bytes =
23                std::slice::from_raw_parts_mut(data.as_mut_ptr() as *mut u8, count * Self::SIZE);
24            file.read_exact(bytes)
25                .map_err(|_| invalid("truncated .npy data".into()))?;
26            data.set_len(count);
27        }
28        Ok(data)
29    }
30    #[cfg(target_endian = "big")]
31    fn read_from<R: Read>(file: &mut R, count: usize) -> Result<Vec<Self>> {
32        let mut bytes = vec![0u8; count * Self::SIZE];
33        file.read_exact(&mut bytes)
34            .map_err(|_| invalid("truncated .npy data".into()))?;
35        Ok(Self::read_le_all(&bytes))
36    }
37    #[allow(clippy::chunks_exact_to_as_chunks)]
38    fn read_le_all(bytes: &[u8]) -> Vec<Self> {
39        bytes.chunks_exact(Self::SIZE).map(Self::read_le).collect()
40    }
41    #[cfg(target_endian = "little")]
42    fn write_le_all<W: io::Write>(data: &[Self], file: &mut W) -> Result<()> {
43        let bytes = unsafe {
44            std::slice::from_raw_parts(data.as_ptr() as *const u8, std::mem::size_of_val(data))
45        };
46        file.write_all(bytes)
47    }
48    #[cfg(target_endian = "big")]
49    fn write_le_all<W: io::Write>(data: &[Self], file: &mut W) -> Result<()> {
50        file.write_all(&Self::write_le_bytes(data))
51    }
52    fn write_le_bytes(data: &[Self]) -> Vec<u8> {
53        let mut buffer = Vec::with_capacity(std::mem::size_of_val(data));
54        for &value in data {
55            value.write_le(&mut buffer);
56        }
57        buffer
58    }
59}
60
61macro_rules! npy_type {
62    ($type:ty, $descr:literal) => {
63        impl NpyType for $type {
64            const DESCR: &'static str = $descr;
65            const SIZE: usize = size_of::<$type>();
66            fn write_le(self, buffer: &mut Vec<u8>) {
67                buffer.extend_from_slice(&self.to_le_bytes());
68            }
69            fn read_le(bytes: &[u8]) -> Self {
70                Self::from_le_bytes(bytes.try_into().unwrap())
71            }
72        }
73    };
74}
75npy_type!(u8, "|u1");
76npy_type!(i8, "|i1");
77npy_type!(u16, "<u2");
78npy_type!(i16, "<i2");
79npy_type!(u32, "<u4");
80npy_type!(i32, "<i4");
81npy_type!(u64, "<u8");
82npy_type!(i64, "<i8");
83npy_type!(f32, "<f4");
84npy_type!(f64, "<f8");
85
86/// An in-memory NumPy array.
87///
88/// The elements are flat in `data`, alongside their `shape` and whether they are
89/// laid out in `fortran_order` (column-major) rather than C order. Read with
90/// [`Npy::read`] / [`Npy::read_from`] and written through the
91/// [`Write`](crate::io::Write) trait or [`Npy::write_to`]. The element type `T`
92/// is any [`NpyType`]: the ten primitive integer and float types.
93pub struct Npy<T> {
94    pub data: Vec<T>,
95    pub shape: Vec<usize>,
96    pub fortran_order: bool,
97}
98
99impl<T, P> Write<P> for Npy<T>
100where
101    T: NpyType,
102    P: AsRef<Path>,
103{
104    type Error = Error;
105    fn write(&self, path: P) -> Result<()> {
106        self.write_to(&mut BufWriter::new(File::create(path)?))
107    }
108}
109
110impl<T: NpyType> Npy<T> {
111    pub fn write_to<W: io::Write>(&self, file: &mut W) -> Result<()> {
112        write_npy(&self.data, &self.shape, self.fortran_order, file)
113    }
114}
115
116fn write_npy<T: NpyType, W: io::Write>(
117    data: &[T],
118    shape: &[usize],
119    fortran_order: bool,
120    file: &mut W,
121) -> Result<()> {
122    let order = if fortran_order { "True" } else { "False" };
123    let dims: String = shape.iter().map(|d| format!("{d}, ")).collect();
124    let mut header = format!(
125        "{{'descr': '{}', 'fortran_order': {order}, 'shape': ({dims}), }}",
126        T::DESCR
127    );
128    let pad = (64 - (10 + header.len() + 1) % 64) % 64;
129    header.push_str(&" ".repeat(pad));
130    header.push('\n');
131    file.write_all(b"\x93NUMPY")?;
132    file.write_all(&[1, 0])?;
133    file.write_all(&(header.len() as u16).to_le_bytes())?;
134    file.write_all(header.as_bytes())?;
135    T::write_le_all(data, file)?;
136    file.flush()
137}
138
139impl<T: NpyType> Npy<T> {
140    pub fn read<P: AsRef<Path>>(path: P) -> Result<Self> {
141        let mut file = BufReader::new(File::open(path)?);
142        Self::read_from(&mut file)
143    }
144    pub fn read_from<R: Read>(file: &mut R) -> Result<Self> {
145        let mut prefix = [0u8; 10];
146        file.read_exact(&mut prefix)
147            .map_err(|_| invalid("not a .npy file".into()))?;
148        if &prefix[..6] != b"\x93NUMPY" {
149            return Err(invalid("not a .npy file".into()));
150        }
151        let header_length = match prefix[6] {
152            1 => u16::from_le_bytes([prefix[8], prefix[9]]) as usize,
153            2 => {
154                let mut rest = [0u8; 2];
155                file.read_exact(&mut rest)?;
156                u32::from_le_bytes([prefix[8], prefix[9], rest[0], rest[1]]) as usize
157            }
158            other => return Err(invalid(format!("unsupported .npy version {other}"))),
159        };
160        let mut header_bytes = vec![0u8; header_length];
161        file.read_exact(&mut header_bytes)?;
162        let header =
163            from_utf8(&header_bytes).map_err(|_| invalid("non-UTF-8 .npy header".into()))?;
164        let descr = quoted(header, "'descr':").ok_or_else(|| invalid("no descr".into()))?;
165        if descr.starts_with('>') || descr.get(1..) != T::DESCR.get(1..) {
166            return Err(invalid(format!(
167                "dtype {descr} does not match {}",
168                T::DESCR
169            )));
170        }
171        let fortran_order = header.contains("'fortran_order': True");
172        let shape = shape(header)?;
173        let count: usize = shape.iter().product();
174        let data = T::read_from(file, count)?;
175        Ok(Npy {
176            data,
177            shape,
178            fortran_order,
179        })
180    }
181}
182
183fn quoted<'a>(header: &'a str, key: &str) -> Option<&'a str> {
184    let at = header.find(key)? + key.len();
185    let open = header[at..].find('\'')? + at + 1;
186    let close = header[open..].find('\'')? + open;
187    Some(&header[open..close])
188}
189
190fn shape(header: &str) -> Result<Vec<usize>> {
191    let at = header
192        .find("'shape':")
193        .ok_or_else(|| invalid("no shape".into()))?;
194    let open = header[at..]
195        .find('(')
196        .ok_or_else(|| invalid("malformed shape".into()))?
197        + at
198        + 1;
199    let close = header[open..]
200        .find(')')
201        .ok_or_else(|| invalid("malformed shape".into()))?
202        + open;
203    header[open..close]
204        .split(',')
205        .map(str::trim)
206        .filter(|dim| !dim.is_empty())
207        .map(|dim| {
208            dim.parse()
209                .map_err(|_| invalid(format!("bad shape entry {dim}")))
210        })
211        .collect()
212}
213
214fn invalid(message: String) -> Error {
215    Error::new(ErrorKind::InvalidData, message)
216}