#![allow(dead_code)]
use crate::sparse_io::*;
use legume_numeric::matrix::common_io::*;
use log::info;
use std::ops::Range;
use std::sync::{Arc, OnceLock};
use zarrs::array::chunk_cache::ChunkCacheDecodedLruChunkLimit;
use zarrs::array::{data_type, ArraySubset, DataType};
use zarrs::filesystem::FilesystemStore;
use zarrs::storage::ReadableListableStorageTraits as ZReadStorageTraits;
const DEFAULT_CACHE_CHUNK_CAP: u64 = 4;
fn cache_chunk_cap() -> u64 {
static CAP: OnceLock<u64> = OnceLock::new();
*CAP.get_or_init(|| {
std::env::var("LEGUME_ZARR_CACHE_CAP")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(DEFAULT_CACHE_CHUNK_CAP)
})
}
const KEY_BY_COLUMN_DATA: &str = "/by_column/data";
const KEY_BY_COLUMN_INDICES: &str = "/by_column/indices";
const KEY_BY_ROW_DATA: &str = "/by_row/data";
const KEY_BY_ROW_INDICES: &str = "/by_row/indices";
use anyhow::anyhow;
use crate::sparse_backend::shared;
use crate::utilities::io_helpers::{chunk_elems, parse_name_file};
const COMPRESSION_LEVEL: i32 = 5;
const MTX_STREAM_BLOCK: u64 = 1 << 20;
#[derive(Clone)]
pub struct SparseMtxData {
read_store: Arc<dyn ZReadStorageTraits>,
write_store: Option<Arc<FilesystemStore>>,
file_name: String,
max_row_name_idx: usize,
max_column_name_idx: usize,
by_column_indptr: Vec<u64>,
streamed_nnz: u64,
by_row_indptr: Vec<u64>,
by_column_indices: Option<Vec<u64>>,
by_column_data: Option<Vec<f32>>,
by_row_indices: Option<Vec<u64>>,
by_row_data: Option<Vec<f32>>,
by_column_data_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
by_column_indices_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
by_row_data_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
by_row_indices_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
}
impl SparseMtxData {
fn write_store(&self) -> anyhow::Result<&Arc<FilesystemStore>> {
self.write_store
.as_ref()
.ok_or_else(|| anyhow!("store is read-only (zip archive)"))
}
}
impl SparseMtxData {
pub fn new(zarr_file: Option<&str>) -> anyhow::Result<Self> {
Self::create_backend(zarr_file)
}
fn create_backend(zarr_file: Option<&str>) -> anyhow::Result<Self> {
match zarr_file {
Some(backend_file) => Self::register_backend_file(backend_file),
None => {
let backend_file = create_temp_dir_file(".zarr")?;
let backend_file = backend_file
.to_str()
.ok_or_else(|| anyhow::anyhow!("Failed to convert path to string"))?;
Self::register_backend_file(backend_file)
}
}
}
pub fn open(backend_file: &str) -> anyhow::Result<Self> {
let (read_store, write_store) = crate::zarr_io::open_zarr_store_rw(backend_file)?;
if (
Self::_num_rows(read_store.clone()),
Self::_num_columns(read_store.clone()),
Self::_num_nnz(read_store.clone()),
) == (None, None, None)
{
anyhow::bail!("Couldn't figure out the size of this sparse matrix data");
}
let mut ret = Self {
read_store,
write_store,
file_name: backend_file.to_string(),
max_row_name_idx: MAX_ROW_NAME_IDX,
max_column_name_idx: MAX_COLUMN_NAME_IDX,
by_column_indptr: vec![],
streamed_nnz: 0,
by_row_indptr: vec![],
by_column_indices: None,
by_column_data: None,
by_row_indices: None,
by_row_data: None,
by_column_data_cache: Arc::new(OnceLock::new()),
by_column_indices_cache: Arc::new(OnceLock::new()),
by_row_data_cache: Arc::new(OnceLock::new()),
by_row_indices_cache: Arc::new(OnceLock::new()),
};
ret.read_column_indptr()?;
ret.read_row_indptr()?;
Ok(ret)
}
pub fn from_mtx_file(
mtx_file: &str,
backend_file: Option<&str>,
index_by_row: Option<bool>,
) -> anyhow::Result<Self> {
let zarr_file = backend_file
.map(|s| s.to_string())
.unwrap_or_else(|| format!("{}.zarr", mtx_file));
info!("backend file: {}", zarr_file);
let mut ret = Self::register_backend_file(&zarr_file)?;
ret.import_mtx_file(mtx_file, index_by_row == Some(true))?;
info!("created sparse backend from {}", mtx_file);
Ok(ret)
}
pub fn from_ndarray(
array: &Array2<f32>,
zarr_file: Option<&str>,
index_by_row: Option<bool>,
) -> anyhow::Result<Self> {
let mut ret = Self::create_backend(zarr_file)?;
ret.import_ndarray_by_col(array)?;
ret.read_column_indptr()?;
if index_by_row == Some(true) {
ret.import_ndarray_by_row(array)?;
ret.read_row_indptr()?;
}
Ok(ret)
}
pub fn from_dmatrix(
matrix: &DMatrix<f32>,
zarr_file: Option<&str>,
index_by_row: Option<bool>,
) -> anyhow::Result<Self> {
let mut ret = Self::create_backend(zarr_file)?;
ret.import_dmatrix_by_col(matrix)?;
ret.read_column_indptr()?;
if index_by_row == Some(true) {
ret.import_dmatrix_by_row(matrix)?;
ret.read_row_indptr()?;
}
Ok(ret)
}
pub fn print_hierarchy(&self) -> anyhow::Result<()> {
use zarrs::config::MetadataRetrieveVersion;
let node =
zarrs::node::Node::open_opt(&self.read_store, "/", &MetadataRetrieveVersion::Default)?;
let tree = node.hierarchy_tree();
info!("hierarchy_tree:\n{}", tree);
Ok(())
}
fn register_backend_file(zarr_file: &str) -> anyhow::Result<Self> {
use zarrs::group::GroupBuilder;
let store = Arc::new(FilesystemStore::new(zarr_file)?);
let root = GroupBuilder::new().build(store.clone(), "/")?;
root.store_metadata()?;
Ok(Self {
read_store: store.clone(),
write_store: Some(store),
file_name: zarr_file.to_string(),
max_row_name_idx: MAX_ROW_NAME_IDX,
max_column_name_idx: MAX_COLUMN_NAME_IDX,
by_column_indptr: vec![],
streamed_nnz: 0,
by_row_indptr: vec![],
by_column_indices: None,
by_column_data: None,
by_row_indices: None,
by_row_data: None,
by_column_data_cache: Arc::new(OnceLock::new()),
by_column_indices_cache: Arc::new(OnceLock::new()),
by_row_data_cache: Arc::new(OnceLock::new()),
by_row_indices_cache: Arc::new(OnceLock::new()),
})
}
fn new_filled_vector<V>(&mut self, key: &str, dt: DataType, vec: &[V]) -> anyhow::Result<()>
where
V: zarrs::array::Element,
{
use zarrs::array::codec::ZstdCodec;
use zarrs::array::ArrayBuilder;
use zarrs::array::FillValue;
let ws = self.write_store()?;
let nelem = vec.len();
let chunk_size = chunk_elems(nelem, std::mem::size_of::<V>());
let fill = if dt == data_type::float32() {
FillValue::from(zarrs::array::ZARR_NAN_F32)
} else if dt == data_type::uint64() {
FillValue::from(0u64)
} else if dt == data_type::string() {
FillValue::from("")
} else {
FillValue::from(0)
};
let array = ArrayBuilder::new(
vec![vec.len() as u64], vec![chunk_size as u64], dt, fill, )
.bytes_to_bytes_codecs(vec![Arc::new(ZstdCodec::new(COMPRESSION_LEVEL, false))])
.build(ws.clone(), key)?;
array.store_metadata()?;
let subset = Self::create_subset(0..vec.len() as u64);
array.store_array_subset(&subset, vec)?;
Ok(())
}
fn _open_vector(
&self,
key: &str,
) -> anyhow::Result<zarrs::array::Array<dyn ZReadStorageTraits>> {
use zarrs::array::Array as ZArray;
let ret = ZArray::open(self.read_store.clone(), key)?;
Ok(ret)
}
fn create_shaped_vector(
&mut self,
key: &str,
dt: DataType,
elem_bytes: usize,
nelem: usize,
) -> anyhow::Result<()> {
use zarrs::array::codec::ZstdCodec;
use zarrs::array::ArrayBuilder;
use zarrs::array::FillValue;
let ws = self.write_store()?;
let chunk_size = chunk_elems(nelem, elem_bytes);
let fill = if dt == data_type::float32() {
FillValue::from(zarrs::array::ZARR_NAN_F32)
} else if dt == data_type::uint64() {
FillValue::from(0u64)
} else {
FillValue::from(0)
};
let array = ArrayBuilder::new(
vec![nelem.max(1) as u64],
vec![chunk_size.max(1) as u64],
dt,
fill,
)
.bytes_to_bytes_codecs(vec![Arc::new(ZstdCodec::new(COMPRESSION_LEVEL, false))])
.build(ws.clone(), key)?;
array.store_metadata()?;
Ok(())
}
fn _open_writable_vector(
&self,
key: &str,
) -> anyhow::Result<zarrs::array::Array<FilesystemStore>> {
use zarrs::array::Array as ZArray;
let ws = self.write_store()?.clone();
let ret = ZArray::open(ws, key)?;
Ok(ret)
}
fn write_slab_u64(&mut self, key: &str, offset: u64, data: &[u64]) -> anyhow::Result<()> {
if data.is_empty() {
return Ok(());
}
let array = self._open_writable_vector(key)?;
let subset = Self::create_subset(offset..offset + data.len() as u64);
array.store_array_subset(&subset, data)?;
Ok(())
}
fn write_slab_f32(&mut self, key: &str, offset: u64, data: &[f32]) -> anyhow::Result<()> {
if data.is_empty() {
return Ok(());
}
let array = self._open_writable_vector(key)?;
let subset = Self::create_subset(offset..offset + data.len() as u64);
array.store_array_subset(&subset, data)?;
Ok(())
}
#[allow(clippy::type_complexity)]
fn open_csc_triplets(
&self,
) -> anyhow::Result<(
zarrs::array::Array<dyn ZReadStorageTraits>,
zarrs::array::Array<dyn ZReadStorageTraits>,
zarrs::array::Array<dyn ZReadStorageTraits>,
)> {
Ok((
self._open_vector("/by_column/indptr")?,
self._open_vector("/by_column/data")?,
self._open_vector("/by_column/indices")?,
))
}
#[inline]
fn create_subset(range: Range<u64>) -> ArraySubset {
ArraySubset::new_with_ranges(&[range])
}
fn open_chunk_cache(
read_store: &Arc<dyn ZReadStorageTraits>,
key: &str,
) -> anyhow::Result<ChunkCacheDecodedLruChunkLimit> {
use zarrs::array::Array as ZArray;
use zarrs::storage::ReadableStorageTraits;
let storage_readable: Arc<dyn ReadableStorageTraits> = read_store.clone().readable();
let arr = ZArray::open(read_store.clone(), key)?;
let arr_arc = Arc::new(arr.with_storage(storage_readable));
Ok(ChunkCacheDecodedLruChunkLimit::new(
arr_arc,
cache_chunk_cap(),
))
}
fn cache_for<'a>(
&'a self,
cell: &'a OnceLock<ChunkCacheDecodedLruChunkLimit>,
key: &str,
) -> anyhow::Result<&'a ChunkCacheDecodedLruChunkLimit> {
if let Some(cache) = cell.get() {
return Ok(cache);
}
let _ = cell.set(Self::open_chunk_cache(&self.read_store, key)?);
Ok(cell.get().expect("OnceLock populated above"))
}
fn _retrieve_vector<V>(&self, key: &str) -> anyhow::Result<Vec<V>>
where
V: zarrs::array::ElementOwned,
{
let data = self._open_vector(key)?;
let ntot = data.shape()[0];
let subset = Self::create_subset(0..ntot);
Ok(data.retrieve_array_subset::<Vec<V>>(&subset)?)
}
fn _set_group_attr<V>(
store: Arc<FilesystemStore>,
group_name: &str,
attr_name: &str,
value: &V,
) -> anyhow::Result<()>
where
V: serde::Serialize,
{
use zarrs::group::Group;
let mut group = Group::open(store, group_name)?;
let new_value = serde_json::to_value(value)?;
group
.attributes_mut()
.insert((*attr_name).to_string(), new_value);
group.store_metadata()?;
Ok(())
}
fn _get_group_attr<V>(
store: Arc<dyn ZReadStorageTraits>,
group_name: &str,
attr_name: &str,
) -> Option<V>
where
V: serde::de::DeserializeOwned,
{
zarrs::group::Group::open(store, group_name)
.ok()
.and_then(|grp| grp.attributes().get(attr_name).cloned())
.and_then(|attr| serde_json::from_value(attr).ok())
}
fn _num_nnz(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
Self::_get_group_attr::<usize>(store, "/", "nnz")
}
fn _num_rows(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
Self::_get_group_attr::<usize>(store, "/", "nrow")
}
fn _num_columns(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
Self::_get_group_attr::<usize>(store, "/", "ncol")
}
fn _add_group(&mut self, group_name: &str) -> anyhow::Result<()> {
use zarrs::group::Group;
let ws = self.write_store()?;
if Group::open(ws.clone(), group_name).is_err() {
let new_group = zarrs::group::GroupBuilder::new().build(ws.clone(), group_name)?;
new_group.store_metadata()?;
}
Ok(())
}
}
impl SparseIo for SparseMtxData {
type IndexIter = Vec<usize>;
fn read_row_indptr(&mut self) -> anyhow::Result<()> {
use zarrs::array::Array as Zarray;
let key = "/by_row/indptr";
if let Ok(indptr) = Zarray::open(self.read_store.clone(), key) {
let indptr_vec = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
self.by_row_indptr.clear();
self.by_row_indptr.extend(indptr_vec);
}
Ok(())
}
fn column_indptr(&self) -> &[u64] {
&self.by_column_indptr
}
fn reopen_backend(&mut self) -> anyhow::Result<()> {
let store = Arc::new(FilesystemStore::new(&self.file_name)?);
self.read_store = store.clone();
self.write_store = Some(store);
self.by_column_data_cache = Arc::new(OnceLock::new());
self.by_column_indices_cache = Arc::new(OnceLock::new());
self.by_row_data_cache = Arc::new(OnceLock::new());
self.by_row_indices_cache = Arc::new(OnceLock::new());
self.streamed_nnz = 0;
self.read_column_indptr()?;
self.read_row_indptr()?;
Ok(())
}
fn note_streamed_nnz(&mut self, n: u64) {
self.streamed_nnz += n;
}
fn streamed_nnz(&self) -> u64 {
self.streamed_nnz
}
fn reset_streamed_nnz(&mut self) {
self.streamed_nnz = 0;
}
fn read_column_indptr(&mut self) -> anyhow::Result<()> {
use zarrs::array::Array as ZArray;
let key = "/by_column/indptr";
if let Ok(indptr) = ZArray::open(self.read_store.clone(), key) {
let indptr_vec = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
self.by_column_indptr.clear();
self.by_column_indptr.extend(indptr_vec);
}
Ok(())
}
fn clean_preloaded_columns(&mut self) {
self.by_column_data = None;
self.by_column_indices = None;
}
fn preload_columns(&mut self) -> anyhow::Result<()> {
if let Some(nnz) = self.num_non_zeros() {
if !crate::sparse_io::preload_within_budget(nnz, "column") {
return Ok(());
}
}
use zarrs::array::Array as ZArray;
let key = "/by_column/data";
let data = ZArray::open(self.read_store.clone(), key)?;
let key = "/by_column/indices";
let indices = ZArray::open(self.read_store.clone(), key)?;
let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
self.by_column_indices = Some(indices);
self.by_column_data = Some(data);
Ok(())
}
fn clean_preloaded_rows(&mut self) {
self.by_row_data = None;
self.by_row_indices = None;
}
fn preload_rows(&mut self) -> anyhow::Result<()> {
if let Some(nnz) = self.num_non_zeros() {
if !crate::sparse_io::preload_within_budget(nnz, "row") {
return Ok(());
}
}
use zarrs::array::Array as ZArray;
let data = ZArray::open(self.read_store.clone(), KEY_BY_ROW_DATA)?;
let indices = ZArray::open(self.read_store.clone(), KEY_BY_ROW_INDICES)?;
let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
self.by_row_indices = Some(indices);
self.by_row_data = Some(data);
Ok(())
}
fn record_mtx_shape(&mut self, mtx_shape: Option<(usize, usize, usize)>) -> anyhow::Result<()> {
if let Some((nrow, ncol, nnz)) = mtx_shape {
let ws = self.write_store()?;
let read_store = self.read_store.clone();
let check_set_attr = |attr_name: &str, value: usize| -> anyhow::Result<()> {
let old_value = Self::_get_group_attr::<usize>(read_store.clone(), "/", attr_name);
let new_value = serde_json::to_value(value)?;
match old_value {
Some(old_value) => {
if old_value != new_value {
return Err(anyhow!("{} mismatch", attr_name));
}
}
_ => {
Self::_set_group_attr(ws.clone(), "/", attr_name, &new_value)?;
}
}
Ok(())
};
check_set_attr("nrow", nrow)?;
check_set_attr("ncol", ncol)?;
check_set_attr("nnz", nnz)?;
}
Ok(())
}
fn initialize_backend(&mut self) -> anyhow::Result<()> {
use zarrs::group::GroupBuilder;
self.remove_backend_file()?;
let zarr_file = &self.file_name;
let store = Arc::new(FilesystemStore::new(zarr_file)?);
let root = GroupBuilder::new().build(store.clone(), "/")?;
root.store_metadata()?;
self.read_store = store.clone();
self.write_store = Some(store);
self.file_name = zarr_file.to_string();
self.max_column_name_idx = MAX_COLUMN_NAME_IDX;
self.max_row_name_idx = MAX_ROW_NAME_IDX;
self.by_column_indptr = vec![];
self.by_row_indptr = vec![];
Ok(())
}
fn remove_backend_file(&self) -> anyhow::Result<()> {
let backend = std::path::Path::new(&self.file_name);
if backend.exists() {
if backend.is_file() {
std::fs::remove_file(backend)?;
} else {
std::fs::remove_dir_all(backend)?;
}
}
Ok(())
}
fn get_backend_file_name(&self) -> &str {
&self.file_name
}
fn backend_type(&self) -> SparseIoBackend {
SparseIoBackend::Zarr
}
fn to_mtx_file(&self, mtx_file: &str) -> anyhow::Result<()> {
if let (Some(ncol), Some(nrow), Some(nnz)) =
(self.num_columns(), self.num_rows(), self.num_non_zeros())
{
let (nrow, ncol, nnz) = (nrow, ncol, nnz);
let mut buf = open_buf_writer(mtx_file)?;
shared::write_mtx_header(&mut buf, nrow, ncol, nnz)?;
let (indptr, data, indices) = self.open_csc_triplets()?;
let indptr = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
debug_assert!(indptr.len() == ncol + 1);
let total_nnz = indptr[ncol];
let mut jj = 0usize; let mut pos = 0u64;
while pos < total_nnz {
let end = (pos + MTX_STREAM_BLOCK).min(total_nnz);
let subset = Self::create_subset(pos..end);
let data_block = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
let indices_block = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
for (k, (&val, &ii)) in data_block.iter().zip(&indices_block).enumerate() {
let global = pos + k as u64;
while jj + 1 < indptr.len() && indptr[jj + 1] <= global {
jj += 1;
}
writeln!(buf, "{}\t{}\t{}", ii as usize + 1, jj + 1, val)?;
}
pos = end;
}
buf.flush()?;
Ok(())
} else {
Err(anyhow!("Unable to figure out the size of the backend data"))
}
}
fn register_row_names_file(&mut self, row_name_file: &str) {
let _ = self.register_names_file(
"/row_names",
row_name_file,
0..self.max_row_name_idx,
ROW_SEP,
);
}
fn register_row_names_vec(&mut self, rows: &[Box<str>]) {
let _ = self.register_names_vec("/row_names", rows);
}
fn register_column_names_file(&mut self, column_name_file: &str) {
let _ = self.register_names_file(
"/column_names",
column_name_file,
0..self.max_column_name_idx,
COLUMN_SEP,
);
}
fn register_column_names_vec(&mut self, columns: &[Box<str>]) {
let _ = self.register_names_vec("/column_names", columns);
}
fn num_rows(&self) -> Option<usize> {
Self::_num_rows(self.read_store.clone())
}
fn num_columns(&self) -> Option<usize> {
Self::_num_columns(self.read_store.clone())
}
fn num_non_zeros(&self) -> Option<usize> {
Self::_num_nnz(self.read_store.clone())
}
fn register_names_file(
&mut self,
key: &str,
name_file: &str,
name_columns: Range<usize>,
name_sep: &str,
) -> anyhow::Result<()> {
let names = parse_name_file(name_file, name_columns, name_sep)?;
self.new_filled_vector(key, data_type::string(), &names)?;
Ok(())
}
fn register_names_vec(&mut self, key: &str, names: &[Box<str>]) -> anyhow::Result<()> {
let names_vec: Vec<String> = names.iter().map(|x| x.to_string()).collect();
self.new_filled_vector(key, data_type::string(), &names_vec)?;
Ok(())
}
fn row_names(&self) -> anyhow::Result<Vec<Box<str>>> {
self.retrieve_registered_names("/row_names")
}
fn column_names(&self) -> anyhow::Result<Vec<Box<str>>> {
self.retrieve_registered_names("/column_names")
}
fn retrieve_registered_names(&self, key: &str) -> anyhow::Result<Vec<Box<str>>> {
Ok(self
._retrieve_vector::<String>(key)?
.into_iter()
.map(|s| s.into_boxed_str())
.collect())
}
fn read_triplets_by_single_column(
&self,
j_data: usize,
) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
use zarrs::array::Array as ZArray;
debug_assert!(!self.by_column_indptr.is_empty()); debug_assert!(j_data < self.num_columns().unwrap_or(0));
let indptr = &self.by_column_indptr;
debug_assert!((j_data + 1) < indptr.len());
debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
let nrow = self
.num_rows()
.ok_or(anyhow!("can't figure out the number of rows"))?;
if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
let ncol_out = 1;
let jj = 0;
let start = indptr[j_data] as usize;
let end = indptr[j_data + 1] as usize;
let ret: Vec<(u64, u64, f32)> = indices[start..end]
.iter()
.zip(data[start..end].iter())
.map(|(&ii, &x_ij)| (ii, jj, x_ij))
.collect();
Ok((nrow, ncol_out, ret))
} else {
let key = "/by_column/data";
let data = ZArray::open(self.read_store.clone(), key)?;
let key = "/by_column/indices";
let indices = ZArray::open(self.read_store.clone(), key)?;
let ncol_out = 1;
let jj = 0;
let start = indptr[j_data];
let end = indptr[j_data + 1];
let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((end - start) as usize);
if start < end {
let subset = Self::create_subset(start..end);
let data_slice = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
let indices_slice = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
for k in 0..(end - start) {
let x_ij = data_slice[k as usize];
let ii = indices_slice[k as usize];
debug_assert!((ii as usize) < nrow);
ret.push((ii, jj, x_ij));
}
}
Ok((nrow, ncol_out, ret))
}
}
fn read_triplets_by_columns(
&self,
columns: Self::IndexIter,
) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
debug_assert!(!self.by_column_indptr.is_empty());
let indptr = &self.by_column_indptr;
let columns_vec = columns.into_iter().collect::<Vec<usize>>();
debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
let nrow = self
.num_rows()
.ok_or(anyhow!("can't figure out the number of rows"))?;
let ncol = self
.num_columns()
.ok_or(anyhow!("can't figure out the number of columns"))?;
let ncol_out = columns_vec.len();
if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
let min_start = columns_vec
.iter()
.map(|&j_data| indptr[j_data])
.min()
.unwrap_or(0);
let max_end = columns_vec
.iter()
.map(|&j_data| indptr[j_data + 1])
.max()
.unwrap_or(0);
let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((max_end - min_start) as usize);
for (jj, &j_data) in columns_vec.iter().enumerate() {
let jj = jj as u64;
let start = indptr[j_data] as usize;
let end = indptr[j_data + 1] as usize;
for (&ii, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
ret.push((ii, jj, x_ij));
}
}
Ok((nrow, ncol_out, ret))
} else {
let mut tagged: Vec<(u64, u64, u64)> = columns_vec
.iter()
.enumerate()
.filter_map(|(jj, &j_data)| {
if j_data >= ncol {
return None;
}
let start = indptr[j_data];
let end = indptr[j_data + 1];
(start < end).then_some((jj as u64, start, end))
})
.collect();
tagged.sort_by_key(|&(_, start, _)| start);
let data_cache = self.cache_for(&self.by_column_data_cache, KEY_BY_COLUMN_DATA)?;
let indices_cache =
self.cache_for(&self.by_column_indices_cache, KEY_BY_COLUMN_INDICES)?;
let opts = zarrs::array::CodecOptions::default();
let ret = shared::coalesce_and_emit(
&tagged,
nrow,
|jj, ii, val| (ii, jj, val),
|s, e| {
let subset = Self::create_subset(s..e);
let data_buf =
<_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
Vec<f32>,
>(data_cache, &subset, &opts)?;
let indices_buf =
<_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
Vec<u64>,
>(indices_cache, &subset, &opts)?;
Ok((data_buf, indices_buf))
},
)?;
Ok((nrow, ncol_out, ret))
}
}
fn csc_column_arrays(&self) -> Option<(&[u64], &[u64], &[f32])> {
match (
self.by_column_data.as_ref(),
self.by_column_indices.as_ref(),
) {
(Some(data), Some(indices)) if !self.by_column_indptr.is_empty() => Some((
self.by_column_indptr.as_slice(),
indices.as_slice(),
data.as_slice(),
)),
_ => None,
}
}
fn read_triplets_by_rows(
&self,
rows: Self::IndexIter,
) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
debug_assert!(!self.by_row_indptr.is_empty());
let indptr = &self.by_row_indptr;
debug_assert!(indptr.len() > self.num_rows().unwrap_or(0));
let rows_vec = rows.into_iter().collect::<Vec<_>>();
let (nrow, ncol) = match (self.num_rows(), self.num_columns()) {
(Some(nrow), Some(ncol)) => (nrow, ncol),
_ => return Err(anyhow!("Unable to figure out the size of the backend data")),
};
let nrow_out = rows_vec.len();
if let (Some(data), Some(indices)) = (&self.by_row_data, &self.by_row_indices) {
let mut nnz_total: usize = 0;
let valid: Vec<(u64, usize)> = rows_vec
.iter()
.enumerate()
.filter_map(|(ii, &i_data)| {
if i_data >= nrow {
return None;
}
nnz_total += (indptr[i_data + 1] - indptr[i_data]) as usize;
Some((ii as u64, i_data))
})
.collect();
let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity(nnz_total);
for (ii, i_data) in valid {
let start = indptr[i_data] as usize;
let end = indptr[i_data + 1] as usize;
for (&jj, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
ret.push((ii, jj, x_ij));
}
}
return Ok((nrow_out, ncol, ret));
}
let mut tagged: Vec<(u64, u64, u64)> = rows_vec
.iter()
.enumerate()
.filter_map(|(ii, &i_data)| {
if i_data >= nrow {
return None;
}
debug_assert!((i_data + 1) < indptr.len());
let start = indptr[i_data];
let end = indptr[i_data + 1];
(start < end).then_some((ii as u64, start, end))
})
.collect();
tagged.sort_by_key(|&(_, start, _)| start);
let data_cache = self.cache_for(&self.by_row_data_cache, KEY_BY_ROW_DATA)?;
let indices_cache = self.cache_for(&self.by_row_indices_cache, KEY_BY_ROW_INDICES)?;
let opts = zarrs::array::CodecOptions::default();
let ret = shared::coalesce_and_emit(
&tagged,
ncol,
|ii, jj, val| (ii, jj, val),
|s, e| {
let subset = Self::create_subset(s..e);
let data_buf = <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
Vec<f32>,
>(data_cache, &subset, &opts)?;
let indices_buf =
<_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<Vec<u64>>(
indices_cache,
&subset,
&opts,
)?;
Ok((data_buf, indices_buf))
},
)?;
Ok((nrow_out, ncol, ret))
}
fn record_csr_dataset_backend(
&mut self,
csr_cols: &[u64],
csr_vals: &[f32],
csr_rowptr: &[u64],
) -> anyhow::Result<()> {
let key = "/by_row";
self._add_group(key)?;
let key = "/by_row/data";
self.new_filled_vector(key, data_type::float32(), csr_vals)?;
let key = "/by_row/indices";
self.new_filled_vector(key, data_type::uint64(), csr_cols)?;
let key = "/by_row/indptr";
self.new_filled_vector(key, data_type::uint64(), csr_rowptr)?;
Ok(())
}
fn record_csc_dataset_backend(
&mut self,
csc_rows: &[u64],
csc_vals: &[f32],
csc_colptr: &[u64],
) -> anyhow::Result<()> {
let key = "/by_column";
self._add_group(key)?;
let key = "/by_column/data";
self.new_filled_vector(key, data_type::float32(), csc_vals)?;
let key = "/by_column/indices";
self.new_filled_vector(key, data_type::uint64(), csc_rows)?;
let key = "/by_column/indptr";
self.new_filled_vector(key, data_type::uint64(), csc_colptr)?;
Ok(())
}
fn cs_create(&mut self, key: CsKey, len: usize) -> anyhow::Result<()> {
let (group, path, dt, elem_bytes) = match key {
CsKey::CscData => (
"/by_column",
"/by_column/data",
data_type::float32(),
std::mem::size_of::<f32>(),
),
CsKey::CscIndices => (
"/by_column",
"/by_column/indices",
data_type::uint64(),
std::mem::size_of::<u64>(),
),
CsKey::CscIndptr => (
"/by_column",
"/by_column/indptr",
data_type::uint64(),
std::mem::size_of::<u64>(),
),
CsKey::CsrData => (
"/by_row",
"/by_row/data",
data_type::float32(),
std::mem::size_of::<f32>(),
),
CsKey::CsrIndices => (
"/by_row",
"/by_row/indices",
data_type::uint64(),
std::mem::size_of::<u64>(),
),
CsKey::CsrIndptr => (
"/by_row",
"/by_row/indptr",
data_type::uint64(),
std::mem::size_of::<u64>(),
),
};
self._add_group(group)?;
self.create_shaped_vector(path, dt, elem_bytes, len)
}
fn cs_write_u64(&mut self, key: CsKey, offset: u64, data: &[u64]) -> anyhow::Result<()> {
let path = match key {
CsKey::CscIndices => "/by_column/indices",
CsKey::CscIndptr => "/by_column/indptr",
CsKey::CsrIndices => "/by_row/indices",
CsKey::CsrIndptr => "/by_row/indptr",
CsKey::CscData | CsKey::CsrData => {
return Err(anyhow!("cs_write_u64 called on f32 slot {:?}", key));
}
};
self.write_slab_u64(path, offset, data)
}
fn cs_write_f32(&mut self, key: CsKey, offset: u64, data: &[f32]) -> anyhow::Result<()> {
let path = match key {
CsKey::CscData => "/by_column/data",
CsKey::CsrData => "/by_row/data",
_ => {
return Err(anyhow!("cs_write_f32 called on u64 slot {:?}", key));
}
};
self.write_slab_f32(path, offset, data)
}
}