use std::{collections::HashMap, sync::Arc};
use carton_macros::for_each_numeric_carton_type;
use lunchbox::{
types::{MaybeSend, MaybeSync, ReadableFile},
ReadableFileSystem,
};
use serde::{Deserialize, Serialize};
use crate::{
info::PossiblyLoaded,
types::{GenericStorage, Tensor, TensorStorage, TypedStorage},
};
#[derive(Default, Serialize, Deserialize)]
struct IndexToml {
tensor: Vec<TensorInfo>,
}
#[derive(Default, Serialize, Deserialize)]
struct TensorInfo {
name: String,
dtype: String,
shape: Option<Vec<u64>>,
file: Option<String>,
inner: Vec<String>,
}
#[derive(Default, Serialize, Deserialize)]
struct StringsToml {
data: Vec<String>,
}
pub(crate) fn save_tensors<T>(
tensor_data_path: &std::path::Path,
tensors: HashMap<String, &Tensor<T>>,
) -> crate::error::Result<()>
where
T: TensorStorage,
{
let mut index_toml = IndexToml::default();
let (nested, mut unnested) = tensors
.into_iter()
.partition::<HashMap<_, _>, _>(|(_, v)| matches!(v, Tensor::NestedTensor(_)));
for (k, v) in nested {
if let Tensor::NestedTensor(items) = v {
let mut nt = TensorInfo {
name: k.strip_prefix("@tensor_data/").unwrap().to_owned(),
dtype: "nested".into(),
..Default::default()
};
for (idx, t) in items.into_iter().enumerate() {
let inner_name = format!("_carton_nested_inner_{k}_{idx}");
if matches!(t, Tensor::NestedTensor(_)) {
panic!("NestedTensors cannot contain NestedTensors");
}
nt.inner.push(inner_name.clone());
if unnested.insert(inner_name, t).is_some() {
panic!("Tensor names starting with `_carton_nested_inner_` are reserved.")
}
}
index_toml.tensor.push(nt);
} else {
unreachable!("This shouldn't happen because we partitioned above")
}
}
for (tensor_idx, (k, v)) in unnested.iter().enumerate() {
if let Tensor::String(t) = v {
let string_tensor = StringsToml {
data: t.view().as_standard_layout().into_iter().collect(),
};
let fname = format!("tensor_{tensor_idx}.toml");
index_toml.tensor.push(TensorInfo {
name: k.strip_prefix("@tensor_data/").unwrap().to_owned(),
dtype: "string".into(),
shape: Some(t.view().shape().into_iter().map(|v| *v as u64).collect()),
file: Some(fname.clone()),
..Default::default()
});
let serialized = toml::to_string_pretty(&string_tensor).unwrap();
std::fs::write(tensor_data_path.join(fname), serialized).unwrap();
} else {
for_each_numeric_carton_type! {
match v {
Tensor::NestedTensor(_) => {
unreachable!("This shouldn't happen because we partitioned above")
}
Tensor::String(_) => unreachable!(
"This shouldn't happen because we handled string tensors immediately above"
),
$(
Tensor::$CartonType(v) => {
let view = v.view();
let array = view.as_standard_layout();
#[cfg(not(target_endian = "little"))]
compile_error!("Writing tensor_data to disk is currently only supported on little-endian platforms");
let bytes_per_elem = bytes_per_elem(&view);
let total_bytes = array.len() * bytes_per_elem;
let data = unsafe { std::slice::from_raw_parts(array.as_ptr() as *const u8, total_bytes) };
let fname = format!("tensor_{tensor_idx}.bin");
index_toml.tensor.push(TensorInfo {
name: k.strip_prefix("@tensor_data/").unwrap().to_owned(),
dtype: $TypeStr.into(),
shape: Some(array.shape().into_iter().map(|v| *v as u64).collect()),
file: Some(fname.clone()),
..Default::default()
});
std::fs::write(tensor_data_path.join(fname), data).unwrap();
}
)*
};
}
}
}
let serialized = toml::to_string_pretty(&index_toml).unwrap();
std::fs::write(tensor_data_path.join("index.toml"), serialized).unwrap();
Ok(())
}
fn bytes_per_elem<T>(_array: &ndarray::ArrayViewD<T>) -> usize {
std::mem::size_of::<T>()
}
pub(crate) async fn load_tensors<T>(
fs: &Arc<T>,
tensor_data_path: &lunchbox::path::Path,
) -> crate::error::Result<HashMap<String, PossiblyLoaded<Tensor<GenericStorage>>>>
where
T: ReadableFileSystem + MaybeSend + MaybeSync + 'static,
T::FileType: ReadableFile + MaybeSend + MaybeSync + 'static,
{
let index_toml: IndexToml =
toml::from_slice(&fs.read(tensor_data_path.join("index.toml")).await.unwrap()).unwrap();
let mut unnested: HashMap<String, PossiblyLoaded<Tensor<GenericStorage>>> = HashMap::new();
for t in &index_toml.tensor {
for_each_numeric_carton_type! {
let loader = match t.dtype.as_str() {
"nested" => {
continue;
},
"string" => {
let shape: Vec<_> = t.shape.as_ref().unwrap().iter().map(|v| *v as usize).collect();
let fname = t.file.clone().unwrap();
let fs = fs.clone();
let path = tensor_data_path.join(fname);
PossiblyLoaded::from_loader(Box::pin(async move {
let data = fs.read(path).await.unwrap();
let strings: StringsToml = toml::from_slice(&data).unwrap();
Tensor::String(ndarray::ArrayD::<String>::from_shape_vec(shape, strings.data).unwrap())
}))
},
$(
$TypeStr => {
let shape: Vec<_> = t.shape.as_ref().unwrap().iter().map(|v| *v as usize).collect();
let fname = t.file.clone().unwrap();
let fs = fs.clone();
let path = tensor_data_path.join(fname);
PossiblyLoaded::from_loader(Box::pin(async move {
let data = fs.read(path).await.unwrap();
#[cfg(not(target_endian = "little"))]
compile_error!("Reading tensor_data from disk is currently only supported on little-endian platforms");
let bytes_per_elem = std::mem::size_of::<$RustType>();
let numel = data.len() / bytes_per_elem;
let typed_data = unsafe { std::slice::from_raw_parts(data.as_ptr() as *const $RustType, numel) }.to_vec();
Tensor::$CartonType(ndarray::ArrayD::<$RustType>::from_shape_vec(shape, typed_data).unwrap())
}))
},
)*
dtype => panic!("Found tensor with unknown type {dtype}. You may need to upgrade the version of Carton you're using.")
};
unnested.insert(t.name.clone(), loader);
}
}
let mut out: HashMap<_, _> = index_toml
.tensor
.into_iter()
.filter_map(|item| {
if item.dtype == "nested" {
let inner: Vec<_> = item
.inner
.into_iter()
.map(|name| unnested.remove(&name).unwrap())
.collect();
Some((
item.name,
PossiblyLoaded::from_loader(Box::pin(async move {
let mut tensors = Vec::new();
for item in inner {
tensors.push(item.into_get().await.unwrap());
}
Tensor::NestedTensor(tensors)
})),
))
} else {
None
}
})
.collect();
out.extend(unnested);
Ok(out)
}