use std::{
collections::HashMap,
fs::File,
io::{Read, Write},
path::Path,
};
use half::{bf16, f16};
use crate::{
array::Array,
dtype::{Complex64, Dtype, Element},
error::{
AllocFailurePayload, Error, FileIoPayload, FileOp, ParsePayload, Result,
UnsupportedDtypePayload,
},
ops::shape::{contiguous, transpose},
};
const MAGIC: [u8; 6] = [0x93, b'N', b'U', b'M', b'P', b'Y'];
#[derive(Debug)]
struct NpyParseError {
detail: String,
}
impl std::fmt::Display for NpyParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.detail)
}
}
impl std::error::Error for NpyParseError {}
fn npy_err(detail: impl Into<String>) -> Error {
Error::Parse(ParsePayload::new(
"io::npy",
"npy",
NpyParseError {
detail: detail.into(),
},
))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ByteOrder {
Big,
LittleOrNative,
}
impl ByteOrder {
const fn swaps_on_host(self) -> bool {
matches!(self, ByteOrder::Big) != cfg!(target_endian = "big")
}
}
fn parse_descr_dtype(descr: &str) -> Result<(Dtype, ByteOrder)> {
let descr_bytes = descr.as_bytes();
let (order, flag, bytes) = match descr_bytes.len() {
2 => (ByteOrder::LittleOrNative, None, descr_bytes),
3 => {
let flag = descr_bytes[0];
let order = match flag {
b'>' => ByteOrder::Big,
b'<' | b'|' | b'=' => ByteOrder::LittleOrNative,
_ => {
return Err(npy_err(format!(
"unsupported array-protocol typestring {descr:?} (invalid byte-order flag)"
)));
}
};
(order, Some(flag), &descr_bytes[1..3])
}
_ => {
return Err(npy_err(format!(
"unsupported array-protocol typestring {descr:?} (expected 2 or 3 chars)"
)));
}
};
let reject_bar_on_multibyte = |dtype: Dtype| -> Result<()> {
if flag == Some(b'|') && dtype_itemsize(dtype) > 1 {
return Err(npy_err(format!(
"unsupported array-protocol typestring {descr:?} (`|` byte-order flag is valid only for single-byte types)"
)));
}
Ok(())
};
if bytes == b"V2" {
reject_bar_on_multibyte(Dtype::BF16)?;
return Ok((Dtype::BF16, order));
}
let kind = bytes[0];
let size = bytes[1]
.checked_sub(b'0')
.filter(|s| *s <= 9)
.ok_or_else(|| {
npy_err(format!(
"unsupported array-protocol typestring {descr:?} (non-digit itemsize)"
))
})?;
let dtype = match (kind, size) {
(b'b', 1) => Dtype::Bool,
(b'i', 1) => Dtype::I8,
(b'i', 2) => Dtype::I16,
(b'i', 4) => Dtype::I32,
(b'i', 8) => Dtype::I64,
(b'u', 1) => Dtype::U8,
(b'u', 2) => Dtype::U16,
(b'u', 4) => Dtype::U32,
(b'u', 8) => Dtype::U64,
(b'f', 2) => Dtype::F16,
(b'f', 4) => Dtype::F32,
(b'f', 8) => Dtype::F64,
(b'c', 8) => Dtype::Complex64,
_ => {
return Err(npy_err(format!(
"unsupported array-protocol typestring {descr:?}"
)));
}
};
reject_bar_on_multibyte(dtype)?;
Ok((dtype, order))
}
const fn dtype_kind(dtype: Dtype) -> u8 {
match dtype {
Dtype::Bool => b'b',
Dtype::U8 | Dtype::U16 | Dtype::U32 | Dtype::U64 => b'u',
Dtype::I8 | Dtype::I16 | Dtype::I32 | Dtype::I64 => b'i',
Dtype::F16 | Dtype::F32 | Dtype::F64 => b'f',
Dtype::BF16 => b'V',
Dtype::Complex64 => b'c',
}
}
const fn dtype_itemsize(dtype: Dtype) -> usize {
match dtype {
Dtype::Bool | Dtype::U8 | Dtype::I8 => 1,
Dtype::U16 | Dtype::I16 | Dtype::F16 | Dtype::BF16 => 2,
Dtype::U32 | Dtype::I32 | Dtype::F32 => 4,
Dtype::U64 | Dtype::I64 | Dtype::F64 | Dtype::Complex64 => 8,
}
}
fn dtype_to_descr(dtype: Dtype) -> String {
let size = dtype_itemsize(dtype);
let order = if size > 1 {
if cfg!(target_endian = "big") {
'>'
} else {
'<'
}
} else {
'|'
};
format!("{order}{}{size}", dtype_kind(dtype) as char)
}
struct NpyHeader {
dtype: Dtype,
swap_endianness: bool,
fortran_order: bool,
shape: Vec<usize>,
data_offset: usize,
}
fn parse_header(bytes: &[u8]) -> Result<NpyHeader> {
let prefix = bytes
.get(0..8)
.ok_or_else(|| npy_err("file shorter than the 8-byte magic + version prefix"))?;
if prefix[0..6] != MAGIC {
return Err(npy_err("invalid magic (not a numpy .npy stream)"));
}
let (major, minor) = (prefix[6], prefix[7]);
let header_len_size = match (major, minor) {
(1, 0) => 2usize,
(2, 0) | (3, 0) => 4usize,
_ => return Err(npy_err("unsupported .npy format version")),
};
let len_start = 8;
let len_end = len_start + header_len_size;
let len_bytes = bytes
.get(len_start..len_end)
.ok_or_else(|| npy_err("truncated header-length field"))?;
let header_len = match header_len_size {
2 => u16::from_le_bytes([len_bytes[0], len_bytes[1]]) as usize,
_ => u32::from_le_bytes([len_bytes[0], len_bytes[1], len_bytes[2], len_bytes[3]]) as usize,
};
let dict_start = len_end;
let dict_end = dict_start
.checked_add(header_len)
.ok_or_else(|| npy_err("header-length overflow"))?;
let dict_bytes = bytes
.get(dict_start..dict_end)
.ok_or_else(|| npy_err("header dict extends past end of file"))?;
let header =
std::str::from_utf8(dict_bytes).map_err(|_| npy_err("header dict is not valid UTF-8"))?;
let HeaderFields {
descr: dtype_descr,
fortran_order,
shape,
} = parse_header_fields(header)?;
let (dtype, byte_order) = parse_descr_dtype(&dtype_descr)?;
let swap_endianness = byte_order.swaps_on_host();
Ok(NpyHeader {
dtype,
swap_endianness,
fortran_order,
shape,
data_offset: dict_end,
})
}
struct HeaderFields {
descr: String,
fortran_order: bool,
shape: Vec<usize>,
}
fn skip_ws(b: &[u8], mut i: usize) -> usize {
while matches!(b.get(i), Some(b' ' | b'\t')) {
i += 1;
}
i
}
fn read_single_quoted(b: &[u8], i: usize) -> Result<(String, usize)> {
if b.get(i) != Some(&b'\'') {
return Err(npy_err("expected a single-quoted token"));
}
let start = i + 1;
let mut j = start;
loop {
match b.get(j) {
Some(b'\'') => break,
Some(_) => j += 1,
None => return Err(npy_err("unterminated single-quoted token")),
}
}
let content = std::str::from_utf8(&b[start..j])
.map_err(|_| npy_err("single-quoted token is not valid UTF-8"))?
.to_string();
Ok((content, j + 1))
}
fn read_bool_at(b: &[u8], i: usize) -> Result<(bool, usize)> {
if b.get(i..i + 4) == Some(b"True") {
Ok((true, i + 4))
} else if b.get(i..i + 5) == Some(b"False") {
Ok((false, i + 5))
} else {
Err(npy_err(
"malformed 'fortran_order' value (expected True/False)",
))
}
}
fn read_shape_tuple_at(b: &[u8], i: usize) -> Result<(Vec<usize>, usize)> {
if b.get(i) != Some(&b'(') {
return Err(npy_err("malformed 'shape' value (no opening '(')"));
}
let start = i + 1;
let mut j = start;
loop {
match b.get(j) {
Some(b')') => break,
Some(_) => j += 1,
None => return Err(npy_err("malformed 'shape' value (no closing ')')")),
}
}
let inner = std::str::from_utf8(&b[start..j])
.map_err(|_| npy_err("malformed 'shape' value (non-UTF-8 interior)"))?;
let tokens: Vec<&str> = inner.split(',').collect();
if tokens.len() == 1 {
return if tokens[0].trim_matches([' ', '\t']).is_empty() {
Ok((Vec::new(), j + 1))
} else {
Err(npy_err(
"malformed 'shape' value (parenthesized integer is not a tuple; expected a trailing comma)",
))
};
}
let mut shape = Vec::with_capacity(tokens.len());
for (k, tok) in tokens.iter().enumerate() {
let trimmed = tok.trim_matches([' ', '\t']);
if trimmed.is_empty() {
if k + 1 == tokens.len() {
continue;
}
return Err(npy_err("empty dimension in shape tuple"));
}
let dim: usize = trimmed
.parse()
.map_err(|_| npy_err("non-integer dimension in shape tuple"))?;
shape.push(dim);
}
Ok((shape, j + 1))
}
fn parse_header_fields(header: &str) -> Result<HeaderFields> {
let b = header.as_bytes();
let mut i = skip_ws(b, 0);
if b.get(i) != Some(&b'{') {
return Err(npy_err("header dict missing opening '{'"));
}
i += 1;
let mut descr: Option<String> = None;
let mut fortran_order: Option<bool> = None;
let mut shape: Option<Vec<usize>> = None;
loop {
i = skip_ws(b, i);
match b.get(i) {
Some(b'}') => {
i += 1;
break;
}
Some(b'\'') => {}
_ => return Err(npy_err("expected a quoted key or '}' in header dict")),
}
let (key, after_key) = read_single_quoted(b, i)?;
i = skip_ws(b, after_key);
if b.get(i) != Some(&b':') {
return Err(npy_err(format!(
"header dict '{key}' key not followed by ':'"
)));
}
i = skip_ws(b, i + 1);
match key.as_str() {
"descr" => {
if descr.is_some() {
return Err(npy_err("header dict has duplicate 'descr' key"));
}
let (v, after) = read_single_quoted(b, i)?;
descr = Some(v);
i = after;
}
"fortran_order" => {
if fortran_order.is_some() {
return Err(npy_err("header dict has duplicate 'fortran_order' key"));
}
let (v, after) = read_bool_at(b, i)?;
fortran_order = Some(v);
i = after;
}
"shape" => {
if shape.is_some() {
return Err(npy_err("header dict has duplicate 'shape' key"));
}
let (v, after) = read_shape_tuple_at(b, i)?;
shape = Some(v);
i = after;
}
_ => return Err(npy_err(format!("header dict has unknown key '{key}'"))),
}
i = skip_ws(b, i);
match b.get(i) {
Some(b',') => {
i += 1;
continue;
}
Some(b'}') => {
i += 1;
break;
}
_ => return Err(npy_err(format!("trailing junk after '{key}' value"))),
}
}
i = skip_ws(b, i);
while matches!(b.get(i), Some(b'\n' | b'\r')) {
i += 1;
i = skip_ws(b, i);
}
if i != b.len() {
return Err(npy_err("trailing junk after header dict"));
}
let descr = descr.ok_or_else(|| npy_err("header dict missing 'descr' key"))?;
let fortran_order =
fortran_order.ok_or_else(|| npy_err("header dict missing 'fortran_order' key"))?;
let shape = shape.ok_or_else(|| npy_err("header dict missing 'shape' key"))?;
Ok(HeaderFields {
descr,
fortran_order,
shape,
})
}
fn shape_numel(shape: &[usize]) -> Result<usize> {
shape.iter().try_fold(1usize, |acc, &d| {
acc.checked_mul(d).ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::new(
"io::npy: shape product",
"usize",
))
})
})
}
fn build_typed_le_view<T>(data: &[u8], numel: usize, shape: &[usize]) -> Result<Option<Array>>
where
T: Element,
{
let size = std::mem::size_of::<T>();
let need = numel.checked_mul(size).ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::new(
"io::npy: numel * itemsize",
"usize",
))
})?;
if data.len() != need {
return Err(npy_err(format!(
"data length mismatch: need {need} bytes, have {}",
data.len()
)));
}
let ptr = data.as_ptr();
if !(ptr as usize).is_multiple_of(std::mem::align_of::<T>()) {
return Ok(None);
}
let view: &[T] = unsafe { std::slice::from_raw_parts(ptr.cast::<T>(), numel) };
Ok(Some(Array::from_slice(view, &shape.to_vec())?))
}
fn build_typed<T, const N: usize>(
data: &[u8],
numel: usize,
shape: &[usize],
swap: bool,
from_le: impl Fn([u8; N]) -> T,
) -> Result<Array>
where
T: Element,
{
debug_assert_eq!(N, std::mem::size_of::<T>());
let need = numel.checked_mul(N).ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::new(
"io::npy: numel * itemsize",
"usize",
))
})?;
if data.len() != need {
return Err(npy_err(format!(
"data length mismatch: need {need} bytes, have {}",
data.len()
)));
}
let mut out: Vec<T> = Vec::new();
out.try_reserve_exact(numel).map_err(|e| {
Error::AllocFailure(AllocFailurePayload::new(
"io::npy: element buffer",
"elements",
numel as u64,
e,
))
})?;
for chunk in data[..need].as_chunks::<N>().0 {
let mut elem = [0u8; N];
elem.copy_from_slice(chunk);
if swap {
elem.reverse();
}
out.push(from_le(elem));
}
Array::from_slice(&out, &shape.to_vec())
}
fn build_complex64_swapped(
data: &[u8],
numel: usize,
need: usize,
shape: &[usize],
swap: bool,
) -> Result<Array> {
let mut out: Vec<Complex64> = Vec::new();
out.try_reserve_exact(numel).map_err(|e| {
Error::AllocFailure(AllocFailurePayload::new(
"io::npy: complex buffer",
"elements",
numel as u64,
e,
))
})?;
for chunk in data[..need].as_chunks::<8>().0 {
let mut re = [chunk[0], chunk[1], chunk[2], chunk[3]];
let mut im = [chunk[4], chunk[5], chunk[6], chunk[7]];
if swap {
re.reverse();
im.reverse();
}
out.push(Complex64::new(
f32::from_le_bytes(re),
f32::from_le_bytes(im),
));
}
Array::from_slice(&out, &shape.to_vec())
}
fn build_array(header: &NpyHeader, data: &[u8]) -> Result<Array> {
let numel = shape_numel(&header.shape)?;
let swap = header.swap_endianness;
let need = numel
.checked_mul(dtype_itemsize(header.dtype))
.ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::new(
"io::npy: numel * itemsize",
"usize",
))
})?;
if data.len() != need {
return Err(npy_err(format!(
"data length mismatch: header declares {need} element bytes, file has {}",
data.len()
)));
}
let read_shape: Vec<usize> = if header.fortran_order {
header.shape.iter().rev().copied().collect()
} else {
header.shape.clone()
};
let can_view = !swap && cfg!(target_endian = "little");
macro_rules! build_numeric {
($T:ty, $N:literal, $from_le:expr) => {{
let viewed = if can_view {
build_typed_le_view::<$T>(data, numel, &read_shape)?
} else {
None
};
match viewed {
Some(a) => a,
None => build_typed::<$T, $N>(data, numel, &read_shape, swap, $from_le)?,
}
}};
}
let arr = match header.dtype {
Dtype::Bool => {
let need = numel;
if data.len() != need {
return Err(npy_err(format!(
"data length mismatch: need {need} bytes, have {}",
data.len()
)));
}
let mut out: Vec<bool> = Vec::new();
out.try_reserve_exact(numel).map_err(|e| {
Error::AllocFailure(AllocFailurePayload::new(
"io::npy: bool buffer",
"elements",
numel as u64,
e,
))
})?;
out.extend(data[..need].iter().map(|&b| b != 0));
Array::from_slice(&out, &read_shape)?
}
Dtype::U8 => build_numeric!(u8, 1, |b| b[0]),
Dtype::I8 => build_numeric!(i8, 1, |b| b[0] as i8),
Dtype::U16 => build_numeric!(u16, 2, u16::from_le_bytes),
Dtype::I16 => build_numeric!(i16, 2, i16::from_le_bytes),
Dtype::U32 => build_numeric!(u32, 4, u32::from_le_bytes),
Dtype::I32 => build_numeric!(i32, 4, i32::from_le_bytes),
Dtype::U64 => build_numeric!(u64, 8, u64::from_le_bytes),
Dtype::I64 => build_numeric!(i64, 8, i64::from_le_bytes),
Dtype::F16 => build_numeric!(f16, 2, |b| f16::from_bits(u16::from_le_bytes(b))),
Dtype::BF16 => build_numeric!(bf16, 2, |b| bf16::from_bits(u16::from_le_bytes(b))),
Dtype::F32 => build_numeric!(f32, 4, f32::from_le_bytes),
Dtype::F64 => build_numeric!(f64, 8, f64::from_le_bytes),
Dtype::Complex64 => {
let need = numel.checked_mul(8).ok_or_else(|| {
Error::ArithmeticOverflow(crate::error::ArithmeticOverflowPayload::new(
"io::npy: complex numel * 8",
"usize",
))
})?;
if data.len() != need {
return Err(npy_err(format!(
"data length mismatch: need {need} bytes, have {}",
data.len()
)));
}
let viewed = if can_view {
build_typed_le_view::<Complex64>(data, numel, &read_shape)?
} else {
None
};
match viewed {
Some(arr) => arr,
None => build_complex64_swapped(data, numel, need, &read_shape, swap)?,
}
}
};
if header.fortran_order {
contiguous(&transpose(&arr)?, false)
} else {
Ok(arr)
}
}
fn load_npy_bytes(bytes: &[u8]) -> Result<Array> {
let header = parse_header(bytes)?;
let data = bytes
.get(header.data_offset..)
.ok_or_else(|| npy_err("data offset past end of file"))?;
build_array(&header, data)
}
fn mmap_file(path: &Path, op: &'static str) -> Result<memmapix::Mmap> {
let file = File::open(path)
.map_err(|e| Error::FileIo(FileIoPayload::new(op, FileOp::Open, path.to_path_buf(), e)))?;
unsafe { memmapix::Mmap::map(&file) }
.map_err(|e| Error::FileIo(FileIoPayload::new(op, FileOp::Read, path.to_path_buf(), e)))
}
pub fn load_npy(path: &Path) -> Result<Array> {
let map = mmap_file(path, "io::load_npy: open")?;
load_npy_bytes(&map)
}
pub fn load_npz(path: &Path) -> Result<HashMap<String, Array>> {
let map = mmap_file(path, "io::load_npz: open")?;
let mapped: &[u8] = ↦
let mut archive = zip::ZipArchive::new(std::io::Cursor::new(mapped))
.map_err(|e| Error::Parse(ParsePayload::new("io::load_npz", "npz (zip archive)", e)))?;
let mut out = HashMap::new();
for i in 0..archive.len() {
let member = archive
.by_index(i)
.map_err(|e| Error::Parse(ParsePayload::new("io::load_npz: member", "npz member", e)))?;
let name = member.name().to_string();
let is_stored = member.compression() == zip::CompressionMethod::Stored;
let data_start = member.data_start();
let uncompressed = member.size();
drop(member);
let arr = if is_stored {
let data_start =
data_start.ok_or_else(|| npy_err("npz member has no resolved data offset"))?;
let start =
usize::try_from(data_start).map_err(|_| npy_err("npz member data offset exceeds usize"))?;
let len =
usize::try_from(uncompressed).map_err(|_| npy_err("npz member length exceeds usize"))?;
let end = start
.checked_add(len)
.ok_or_else(|| npy_err("npz member data range overflows"))?;
let member_bytes = mapped
.get(start..end)
.ok_or_else(|| npy_err("npz member data range extends past end of archive"))?;
load_npy_bytes(member_bytes)?
} else {
let mut member = archive
.by_index(i)
.map_err(|e| Error::Parse(ParsePayload::new("io::load_npz: member", "npz member", e)))?;
let mut bytes = Vec::new();
member.read_to_end(&mut bytes).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"io::load_npz: member read",
FileOp::Read,
path.to_path_buf(),
e,
))
})?;
load_npy_bytes(&bytes)?
};
let key = name
.strip_suffix(".npy")
.map(str::to_string)
.unwrap_or(name);
match out.entry(key) {
std::collections::hash_map::Entry::Occupied(e) => {
return Err(Error::Parse(ParsePayload::new(
"io::load_npz",
"npz",
NpyParseError {
detail: format!("duplicate array key {:?} in npz archive", e.key()),
},
)));
}
std::collections::hash_map::Entry::Vacant(e) => {
e.insert(arr);
}
}
}
Ok(out)
}
fn build_header_bytes(dtype: Dtype, shape: &[usize]) -> Result<Vec<u8>> {
let descr = dtype_to_descr(dtype);
let mut dict = String::new();
dict.push_str("{'descr': '");
dict.push_str(&descr);
dict.push_str("', 'fortran_order': False, 'shape': (");
for d in shape {
dict.push_str(&d.to_string());
dict.push_str(", ");
}
dict.push_str(")}");
let header_len = dict.len();
let is_v1 = header_len + 15 < u16::MAX as usize;
let len_field = if is_v1 { 2usize } else { 4usize };
let padding = (6 + 2 + len_field + header_len + 1) % 16;
dict.push_str(&" ".repeat(padding));
dict.push('\n');
let total_header_len = dict.len();
let mut out = Vec::new();
out.extend_from_slice(&MAGIC);
if is_v1 {
out.push(0x01);
out.push(0x00);
let len = u16::try_from(total_header_len)
.map_err(|_| npy_err("header length does not fit in the v1 u16 length field"))?;
out.extend_from_slice(&len.to_le_bytes());
} else {
out.push(0x02);
out.push(0x00);
let len = u32::try_from(total_header_len)
.map_err(|_| npy_err("header length does not fit in the v2 u32 length field"))?;
out.extend_from_slice(&len.to_le_bytes());
}
out.extend_from_slice(dict.as_bytes());
Ok(out)
}
fn typed_bytes<T: Element>(arr: &mut Array) -> Result<Vec<u8>> {
let slice = arr.as_slice::<T>()?;
let bytes = unsafe {
std::slice::from_raw_parts(slice.as_ptr().cast::<u8>(), std::mem::size_of_val(slice))
};
Ok(bytes.to_vec())
}
fn append_array_bytes(out: &mut Vec<u8>, arr: &mut Array) -> Result<()> {
let dtype = arr.dtype()?;
let bytes = match dtype {
Dtype::Bool => typed_bytes::<bool>(arr)?,
Dtype::U8 => typed_bytes::<u8>(arr)?,
Dtype::I8 => typed_bytes::<i8>(arr)?,
Dtype::U16 => typed_bytes::<u16>(arr)?,
Dtype::I16 => typed_bytes::<i16>(arr)?,
Dtype::U32 => typed_bytes::<u32>(arr)?,
Dtype::I32 => typed_bytes::<i32>(arr)?,
Dtype::U64 => typed_bytes::<u64>(arr)?,
Dtype::I64 => typed_bytes::<i64>(arr)?,
Dtype::F16 => typed_bytes::<f16>(arr)?,
Dtype::BF16 => typed_bytes::<bf16>(arr)?,
Dtype::F32 => typed_bytes::<f32>(arr)?,
Dtype::F64 => typed_bytes::<f64>(arr)?,
Dtype::Complex64 => typed_bytes::<Complex64>(arr)?,
};
out.extend_from_slice(&bytes);
Ok(())
}
fn save_npy_bytes(arr: &mut Array) -> Result<Vec<u8>> {
let dtype = arr.dtype()?;
if arr.size() == 0 {
return Err(Error::UnsupportedDtype(UnsupportedDtypePayload::new(
"io::save_npy: empty array",
dtype,
&[],
)));
}
let shape = arr.shape();
let mut out = build_header_bytes(dtype, &shape)?;
append_array_bytes(&mut out, arr)?;
Ok(out)
}
pub fn save_npy(path: &Path, array: &mut Array) -> Result<()> {
let bytes = save_npy_bytes(array)?;
let mut file = File::create(path).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"io::save_npy: create",
FileOp::Create,
path.to_path_buf(),
e,
))
})?;
file.write_all(&bytes).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"io::save_npy: write",
FileOp::Write,
path.to_path_buf(),
e,
))
})
}
pub fn save_npz(path: &Path, arrays: &mut HashMap<String, Array>) -> Result<()> {
save_npz_impl(path, arrays, zip::CompressionMethod::Stored)
}
pub fn save_npz_compressed(path: &Path, arrays: &mut HashMap<String, Array>) -> Result<()> {
save_npz_impl(path, arrays, zip::CompressionMethod::Deflated)
}
fn save_npz_impl(
path: &Path,
arrays: &mut HashMap<String, Array>,
method: zip::CompressionMethod,
) -> Result<()> {
let file = File::create(path).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"io::save_npz: create",
FileOp::Create,
path.to_path_buf(),
e,
))
})?;
let mut writer = zip::ZipWriter::new(file);
let options: zip::write::FileOptions<'_, ()> =
zip::write::FileOptions::default().compression_method(method);
for (name, arr) in arrays.iter_mut() {
let bytes = save_npy_bytes(arr)?;
let member = format!("{name}.npy");
writer.start_file(member, options).map_err(|e| {
Error::Parse(ParsePayload::new(
"io::save_npz: start_file",
"npz member",
e,
))
})?;
writer.write_all(&bytes).map_err(|e| {
Error::FileIo(FileIoPayload::new(
"io::save_npz: member write",
FileOp::Write,
path.to_path_buf(),
e,
))
})?;
}
writer
.finish()
.map_err(|e| Error::Parse(ParsePayload::new("io::save_npz: finish", "npz archive", e)))?;
Ok(())
}
#[cfg(test)]
mod tests;