Skip to main content

data_beans/sparse_backend/
zarr.rs

1#![allow(dead_code)]
2
3use crate::sparse_io::*;
4use legume_numeric::matrix::common_io::*;
5use log::info;
6use std::ops::Range;
7use std::sync::{Arc, OnceLock};
8use zarrs::array::chunk_cache::ChunkCacheDecodedLruChunkLimit;
9use zarrs::array::{data_type, ArraySubset, DataType};
10use zarrs::filesystem::FilesystemStore;
11use zarrs::storage::ReadableListableStorageTraits as ZReadStorageTraits;
12
13/// Decoded-chunk LRU capacity, per zarr array. Keeps cache memory low
14/// per backend so it scales when many backends are loaded simultaneously
15/// (e.g. multi-sample workloads). With ~1 MiB/chunk, default 4 chunks per
16/// array × 4 arrays = ~16 MiB upper bound per backend. Override with the
17/// `LEGUME_ZARR_CACHE_CAP` env var when total backend count × default
18/// would exceed the memory budget.
19const DEFAULT_CACHE_CHUNK_CAP: u64 = 4;
20
21fn cache_chunk_cap() -> u64 {
22    static CAP: OnceLock<u64> = OnceLock::new();
23    *CAP.get_or_init(|| {
24        std::env::var("LEGUME_ZARR_CACHE_CAP")
25            .ok()
26            .and_then(|s| s.parse().ok())
27            .unwrap_or(DEFAULT_CACHE_CHUNK_CAP)
28    })
29}
30
31const KEY_BY_COLUMN_DATA: &str = "/by_column/data";
32const KEY_BY_COLUMN_INDICES: &str = "/by_column/indices";
33const KEY_BY_ROW_DATA: &str = "/by_row/data";
34const KEY_BY_ROW_INDICES: &str = "/by_row/indices";
35
36use anyhow::anyhow;
37
38use crate::sparse_backend::shared;
39use crate::utilities::io_helpers::{chunk_elems, parse_name_file};
40
41const COMPRESSION_LEVEL: i32 = 5;
42
43/// Block size (in elements) for streaming the CSC value/row-index arrays when
44/// exporting to Matrix Market. ~1M elements ≈ 4 MiB f32 + 8 MiB u64 transient
45/// per block — large enough to amortize per-retrieve overhead, small enough to
46/// bound memory regardless of matrix size.
47const MTX_STREAM_BLOCK: u64 = 1 << 20;
48
49/// 10x-like cell-feature matrix with `zarr` backend (feature x cell)
50///
51/// ```text
52/// (root)
53///     ├── nrow
54///     ├── ncell
55///     ├── by_column
56///     │   ├── data
57///     │   ├── indices (row indices)
58///     │   └── indptr (column pointers)
59///     └── by_row
60///         ├── data
61///         ├── indices (column indices)
62///         └── indptr (row pointers)
63/// ```
64///
65#[derive(Clone)]
66pub struct SparseMtxData {
67    read_store: Arc<dyn ZReadStorageTraits>,
68    write_store: Option<Arc<FilesystemStore>>,
69    file_name: String,
70    max_row_name_idx: usize,
71    max_column_name_idx: usize,
72    by_column_indptr: Vec<u64>,
73    /// Streaming-write cursor: entries appended so far (see `note_streamed_nnz`).
74    streamed_nnz: u64,
75    by_row_indptr: Vec<u64>,
76    by_column_indices: Option<Vec<u64>>,
77    by_column_data: Option<Vec<f32>>,
78    by_row_indices: Option<Vec<u64>>,
79    by_row_data: Option<Vec<f32>>,
80    /// Persistent decoded-chunk LRU caches. `Arc<OnceLock<_>>` so that
81    /// `Clone`s share state, and the underlying `moka::sync::Cache` inside
82    /// `ChunkCacheDecodedLruChunkLimit` is internally thread-safe — no
83    /// external locking needed.
84    by_column_data_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
85    by_column_indices_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
86    by_row_data_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
87    by_row_indices_cache: Arc<OnceLock<ChunkCacheDecodedLruChunkLimit>>,
88}
89
90impl SparseMtxData {
91    /// Get the writable store, or error if this is a read-only backend (e.g. zip archive).
92    fn write_store(&self) -> anyhow::Result<&Arc<FilesystemStore>> {
93        self.write_store
94            .as_ref()
95            .ok_or_else(|| anyhow!("store is read-only (zip archive)"))
96    }
97}
98
99impl SparseMtxData {
100    /// Create an empty new `SparseMtxData` instance with a zarr
101    /// backend file If no `backend_file` is provided, a temporary
102    /// file will be created.
103    ///
104    /// * `backend_file` - Optional zarr backend file
105    pub fn new(zarr_file: Option<&str>) -> anyhow::Result<Self> {
106        Self::create_backend(zarr_file)
107    }
108
109    /// Helper to create a backend file (from provided path or temp file)
110    fn create_backend(zarr_file: Option<&str>) -> anyhow::Result<Self> {
111        match zarr_file {
112            Some(backend_file) => Self::register_backend_file(backend_file),
113            None => {
114                let backend_file = create_temp_dir_file(".zarr")?;
115                let backend_file = backend_file
116                    .to_str()
117                    .ok_or_else(|| anyhow::anyhow!("Failed to convert path to string"))?;
118                Self::register_backend_file(backend_file)
119            }
120        }
121    }
122
123    /// Create `SparseMtxData` instance from an existing zarr backend file
124    /// * `zarr_file` - zarr backend file (directory or `.zarr.zip`)
125    pub fn open(backend_file: &str) -> anyhow::Result<Self> {
126        let (read_store, write_store) = crate::zarr_io::open_zarr_store_rw(backend_file)?;
127
128        if (
129            Self::_num_rows(read_store.clone()),
130            Self::_num_columns(read_store.clone()),
131            Self::_num_nnz(read_store.clone()),
132        ) == (None, None, None)
133        {
134            anyhow::bail!("Couldn't figure out the size of this sparse matrix data");
135        }
136
137        let mut ret = Self {
138            read_store,
139            write_store,
140            file_name: backend_file.to_string(),
141            max_row_name_idx: MAX_ROW_NAME_IDX,
142            max_column_name_idx: MAX_COLUMN_NAME_IDX,
143            by_column_indptr: vec![],
144            streamed_nnz: 0,
145            by_row_indptr: vec![],
146            by_column_indices: None,
147            by_column_data: None,
148            by_row_indices: None,
149            by_row_data: None,
150            by_column_data_cache: Arc::new(OnceLock::new()),
151            by_column_indices_cache: Arc::new(OnceLock::new()),
152            by_row_data_cache: Arc::new(OnceLock::new()),
153            by_row_indices_cache: Arc::new(OnceLock::new()),
154        };
155
156        ret.read_column_indptr()?;
157        ret.read_row_indptr()?;
158
159        Ok(ret)
160    }
161
162    /// Create `SparseMtxData` from mtx file with `backend_file` as
163    /// the backend file.  If no `backend_file` is provided, it will
164    /// be the same as `mtx_file` with `.zarr` extension.
165    /// * `mtx_file`: mtx file to be read into zarr backend
166    /// * `backend_file`: zarr file to be associated with
167    /// * `index_by_row`: if true, the matrix will be indexed by row
168    pub fn from_mtx_file(
169        mtx_file: &str,
170        backend_file: Option<&str>,
171        index_by_row: Option<bool>,
172    ) -> anyhow::Result<Self> {
173        let zarr_file = backend_file
174            .map(|s| s.to_string())
175            .unwrap_or_else(|| format!("{}.zarr", mtx_file));
176
177        info!("backend file: {}", zarr_file);
178        let mut ret = Self::register_backend_file(&zarr_file)?;
179
180        ret.import_mtx_file(mtx_file, index_by_row == Some(true))?;
181
182        info!("created sparse backend from {}", mtx_file);
183        Ok(ret)
184    }
185
186    #[cfg(feature = "ndarray")]
187    /// Create a new `SparseMtxData` instance from an `ndarray` array
188    /// * `array` - 2D array to be added to the backend
189    /// * `backend_file` - Optional zarr backend file
190    /// * `index_by_row` - Optional flag to index by row (CSR format)
191    pub fn from_ndarray(
192        array: &Array2<f32>,
193        zarr_file: Option<&str>,
194        index_by_row: Option<bool>,
195    ) -> anyhow::Result<Self> {
196        let mut ret = Self::create_backend(zarr_file)?;
197
198        ret.import_ndarray_by_col(array)?;
199        ret.read_column_indptr()?;
200
201        if index_by_row == Some(true) {
202            ret.import_ndarray_by_row(array)?;
203            ret.read_row_indptr()?;
204        }
205        Ok(ret)
206    }
207
208    /// Create a new `SparseMtxData` instance from an `DMatrix` array
209    /// * `array` - 2D array to be added to the backend
210    /// * `backend_file` - Optional zarr backend file
211    /// * `index_by_row` - Optional flag to index by row (CSR format)
212    pub fn from_dmatrix(
213        matrix: &DMatrix<f32>,
214        zarr_file: Option<&str>,
215        index_by_row: Option<bool>,
216    ) -> anyhow::Result<Self> {
217        let mut ret = Self::create_backend(zarr_file)?;
218
219        ret.import_dmatrix_by_col(matrix)?;
220        ret.read_column_indptr()?;
221
222        if index_by_row == Some(true) {
223            ret.import_dmatrix_by_row(matrix)?;
224            ret.read_row_indptr()?;
225        }
226        Ok(ret)
227    }
228
229    /// Show the hierarchy of the zarr store
230    pub fn print_hierarchy(&self) -> anyhow::Result<()> {
231        use zarrs::config::MetadataRetrieveVersion;
232        let node =
233            zarrs::node::Node::open_opt(&self.read_store, "/", &MetadataRetrieveVersion::Default)?;
234        let tree = node.hierarchy_tree();
235        info!("hierarchy_tree:\n{}", tree);
236        Ok(())
237    }
238
239    /// Helper function to create a new zarr backend file
240    fn register_backend_file(zarr_file: &str) -> anyhow::Result<Self> {
241        use zarrs::group::GroupBuilder;
242        let store = Arc::new(FilesystemStore::new(zarr_file)?);
243        let root = GroupBuilder::new().build(store.clone(), "/")?;
244        root.store_metadata()?;
245
246        Ok(Self {
247            read_store: store.clone(),
248            write_store: Some(store),
249            file_name: zarr_file.to_string(),
250            max_row_name_idx: MAX_ROW_NAME_IDX,
251            max_column_name_idx: MAX_COLUMN_NAME_IDX,
252            by_column_indptr: vec![],
253            streamed_nnz: 0,
254            by_row_indptr: vec![],
255            by_column_indices: None,
256            by_column_data: None,
257            by_row_indices: None,
258            by_row_data: None,
259            by_column_data_cache: Arc::new(OnceLock::new()),
260            by_column_indices_cache: Arc::new(OnceLock::new()),
261            by_row_data_cache: Arc::new(OnceLock::new()),
262            by_row_indices_cache: Arc::new(OnceLock::new()),
263        })
264    }
265
266    /////////////////////
267    // backend related //
268    /////////////////////
269
270    /// Helper function to create a filled 1D array with the given
271    /// data type and fill value. This is the most useful function to
272    /// create a vector like data.
273    ///
274    /// * `key` - the key name
275    /// * `dt` - the data type among `DataType`
276    /// * `vec` - the vector to be stored
277    ///
278    fn new_filled_vector<V>(&mut self, key: &str, dt: DataType, vec: &[V]) -> anyhow::Result<()>
279    where
280        V: zarrs::array::Element,
281    {
282        use zarrs::array::codec::ZstdCodec;
283        use zarrs::array::ArrayBuilder;
284        use zarrs::array::FillValue;
285
286        let ws = self.write_store()?;
287
288        let nelem = vec.len();
289        let chunk_size = chunk_elems(nelem, std::mem::size_of::<V>());
290
291        let fill = if dt == data_type::float32() {
292            FillValue::from(zarrs::array::ZARR_NAN_F32)
293        } else if dt == data_type::uint64() {
294            FillValue::from(0u64)
295        } else if dt == data_type::string() {
296            FillValue::from("")
297        } else {
298            FillValue::from(0)
299        };
300
301        let array = ArrayBuilder::new(
302            vec![vec.len() as u64],  // array shape
303            vec![chunk_size as u64], // chunk shape
304            dt,                      // data type
305            fill,                    //
306        )
307        .bytes_to_bytes_codecs(vec![Arc::new(ZstdCodec::new(COMPRESSION_LEVEL, false))])
308        .build(ws.clone(), key)?;
309
310        array.store_metadata()?;
311
312        let subset = Self::create_subset(0..vec.len() as u64);
313        array.store_array_subset(&subset, vec)?;
314
315        Ok(())
316    }
317
318    fn _open_vector(
319        &self,
320        key: &str,
321    ) -> anyhow::Result<zarrs::array::Array<dyn ZReadStorageTraits>> {
322        use zarrs::array::Array as ZArray;
323        let ret = ZArray::open(self.read_store.clone(), key)?;
324        Ok(ret)
325    }
326
327    /// Create an empty fixed-shape 1-D array with the given data type,
328    /// chunk layout, and fill value. No data is written.
329    ///
330    /// Used by the streaming write path to pre-create `/by_column/*` and
331    /// `/by_row/*` arrays at their final size before any triplets land.
332    fn create_shaped_vector(
333        &mut self,
334        key: &str,
335        dt: DataType,
336        elem_bytes: usize,
337        nelem: usize,
338    ) -> anyhow::Result<()> {
339        use zarrs::array::codec::ZstdCodec;
340        use zarrs::array::ArrayBuilder;
341        use zarrs::array::FillValue;
342
343        let ws = self.write_store()?;
344
345        let chunk_size = chunk_elems(nelem, elem_bytes);
346
347        let fill = if dt == data_type::float32() {
348            FillValue::from(zarrs::array::ZARR_NAN_F32)
349        } else if dt == data_type::uint64() {
350            FillValue::from(0u64)
351        } else {
352            FillValue::from(0)
353        };
354
355        let array = ArrayBuilder::new(
356            vec![nelem.max(1) as u64],
357            vec![chunk_size.max(1) as u64],
358            dt,
359            fill,
360        )
361        .bytes_to_bytes_codecs(vec![Arc::new(ZstdCodec::new(COMPRESSION_LEVEL, false))])
362        .build(ws.clone(), key)?;
363
364        array.store_metadata()?;
365        Ok(())
366    }
367
368    /// Open an existing array through the writable filesystem store.
369    /// `_open_vector` uses the read-only handle, so it can't be used
370    /// to stage streaming writes.
371    fn _open_writable_vector(
372        &self,
373        key: &str,
374    ) -> anyhow::Result<zarrs::array::Array<FilesystemStore>> {
375        use zarrs::array::Array as ZArray;
376        let ws = self.write_store()?.clone();
377        let ret = ZArray::open(ws, key)?;
378        Ok(ret)
379    }
380
381    /// Write a `u64` slab at the given offset. The target array must
382    /// already exist (see [`create_shaped_vector`]).
383    fn write_slab_u64(&mut self, key: &str, offset: u64, data: &[u64]) -> anyhow::Result<()> {
384        if data.is_empty() {
385            return Ok(());
386        }
387        let array = self._open_writable_vector(key)?;
388        let subset = Self::create_subset(offset..offset + data.len() as u64);
389        array.store_array_subset(&subset, data)?;
390        Ok(())
391    }
392
393    /// Write an `f32` slab at the given offset.
394    fn write_slab_f32(&mut self, key: &str, offset: u64, data: &[f32]) -> anyhow::Result<()> {
395        if data.is_empty() {
396            return Ok(());
397        }
398        let array = self._open_writable_vector(key)?;
399        let subset = Self::create_subset(offset..offset + data.len() as u64);
400        array.store_array_subset(&subset, data)?;
401        Ok(())
402    }
403
404    #[allow(clippy::type_complexity)]
405    fn open_csc_triplets(
406        &self,
407    ) -> anyhow::Result<(
408        zarrs::array::Array<dyn ZReadStorageTraits>,
409        zarrs::array::Array<dyn ZReadStorageTraits>,
410        zarrs::array::Array<dyn ZReadStorageTraits>,
411    )> {
412        Ok((
413            self._open_vector("/by_column/indptr")?,
414            self._open_vector("/by_column/data")?,
415            self._open_vector("/by_column/indices")?,
416        ))
417    }
418
419    /// Helper to create an ArraySubset from a range
420    #[inline]
421    fn create_subset(range: Range<u64>) -> ArraySubset {
422        ArraySubset::new_with_ranges(&[range])
423    }
424
425    /// Build a decoded-chunk LRU cache for one zarr array. The cache
426    /// short-circuits redundant zstd decompression when the same chunk
427    /// is touched by multiple non-mergeable subset reads, both within
428    /// a single batch and across repeated calls (e.g. minibatch loops
429    /// over the same backend).
430    fn open_chunk_cache(
431        read_store: &Arc<dyn ZReadStorageTraits>,
432        key: &str,
433    ) -> anyhow::Result<ChunkCacheDecodedLruChunkLimit> {
434        use zarrs::array::Array as ZArray;
435        use zarrs::storage::ReadableStorageTraits;
436
437        let storage_readable: Arc<dyn ReadableStorageTraits> = read_store.clone().readable();
438        let arr = ZArray::open(read_store.clone(), key)?;
439        let arr_arc = Arc::new(arr.with_storage(storage_readable));
440        Ok(ChunkCacheDecodedLruChunkLimit::new(
441            arr_arc,
442            cache_chunk_cap(),
443        ))
444    }
445
446    /// Lazily build the persistent decoded-chunk cache. Concurrent
447    /// first-callers may both open the array, but only one `set` wins;
448    /// the loser's cache is dropped and `get` returns the winner.
449    fn cache_for<'a>(
450        &'a self,
451        cell: &'a OnceLock<ChunkCacheDecodedLruChunkLimit>,
452        key: &str,
453    ) -> anyhow::Result<&'a ChunkCacheDecodedLruChunkLimit> {
454        if let Some(cache) = cell.get() {
455            return Ok(cache);
456        }
457        let _ = cell.set(Self::open_chunk_cache(&self.read_store, key)?);
458        Ok(cell.get().expect("OnceLock populated above"))
459    }
460
461    fn _retrieve_vector<V>(&self, key: &str) -> anyhow::Result<Vec<V>>
462    where
463        V: zarrs::array::ElementOwned,
464    {
465        let data = self._open_vector(key)?;
466        let ntot = data.shape()[0];
467        let subset = Self::create_subset(0..ntot);
468        Ok(data.retrieve_array_subset::<Vec<V>>(&subset)?)
469    }
470
471    /////////////////////////////
472    // purely helper functions //
473    /////////////////////////////
474
475    /// Helper function to set an attribute from a group named `group_name`
476    fn _set_group_attr<V>(
477        store: Arc<FilesystemStore>,
478        group_name: &str,
479        attr_name: &str,
480        value: &V,
481    ) -> anyhow::Result<()>
482    where
483        V: serde::Serialize,
484    {
485        use zarrs::group::Group;
486        let mut group = Group::open(store, group_name)?;
487
488        let new_value = serde_json::to_value(value)?;
489        group
490            .attributes_mut()
491            .insert((*attr_name).to_string(), new_value);
492        group.store_metadata()?;
493        Ok(())
494    }
495
496    /// Helper function to get an attribute from a group named `group_name`
497    fn _get_group_attr<V>(
498        store: Arc<dyn ZReadStorageTraits>,
499        group_name: &str,
500        attr_name: &str,
501    ) -> Option<V>
502    where
503        V: serde::de::DeserializeOwned,
504    {
505        zarrs::group::Group::open(store, group_name)
506            .ok()
507            .and_then(|grp| grp.attributes().get(attr_name).cloned())
508            .and_then(|attr| serde_json::from_value(attr).ok())
509    }
510
511    fn _num_nnz(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
512        Self::_get_group_attr::<usize>(store, "/", "nnz")
513    }
514
515    fn _num_rows(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
516        Self::_get_group_attr::<usize>(store, "/", "nrow")
517    }
518
519    fn _num_columns(store: Arc<dyn ZReadStorageTraits>) -> Option<usize> {
520        Self::_get_group_attr::<usize>(store, "/", "ncol")
521    }
522
523    /// Helper function to add a group in the writable store
524    fn _add_group(&mut self, group_name: &str) -> anyhow::Result<()> {
525        use zarrs::group::Group;
526        let ws = self.write_store()?;
527
528        if Group::open(ws.clone(), group_name).is_err() {
529            let new_group = zarrs::group::GroupBuilder::new().build(ws.clone(), group_name)?;
530            new_group.store_metadata()?;
531        }
532
533        Ok(())
534    }
535}
536
537impl SparseIo for SparseMtxData {
538    type IndexIter = Vec<usize>;
539
540    /// Read row index pointers
541    fn read_row_indptr(&mut self) -> anyhow::Result<()> {
542        use zarrs::array::Array as Zarray;
543        let key = "/by_row/indptr";
544        if let Ok(indptr) = Zarray::open(self.read_store.clone(), key) {
545            let indptr_vec = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
546            self.by_row_indptr.clear();
547            self.by_row_indptr.extend(indptr_vec);
548        }
549        Ok(())
550    }
551
552    /// Read column index pointers
553    fn column_indptr(&self) -> &[u64] {
554        &self.by_column_indptr
555    }
556
557    fn metadata(&self) -> Metadata {
558        Self::_get_group_attr::<Metadata>(self.read_store.clone(), "/", meta::ATTR)
559            .unwrap_or_default()
560    }
561
562    fn set_metadata(&mut self, values: &Metadata) -> anyhow::Result<()> {
563        let mut root = zarrs::group::Group::open(self.write_store()?.clone(), "/")?;
564        if values.is_empty() {
565            root.attributes_mut().remove(meta::ATTR);
566        } else {
567            root.attributes_mut()
568                .insert(meta::ATTR.to_string(), serde_json::to_value(values)?);
569        }
570        root.store_metadata()?;
571        Ok(())
572    }
573
574    fn reopen_backend(&mut self) -> anyhow::Result<()> {
575        // Path-addressed store: rebuild it, refresh the resident indptrs — and
576        // DROP the decoded-chunk LRU caches, which pin arrays of the store they
577        // were built on. Keeping them served pre-swap chunk contents against
578        // post-swap indptrs: a matrix that read back with row indices past its
579        // own nrow, from a file that was byte-for-byte correct on disk.
580        let store = Arc::new(FilesystemStore::new(&self.file_name)?);
581        self.read_store = store.clone();
582        self.write_store = Some(store);
583        self.by_column_data_cache = Arc::new(OnceLock::new());
584        self.by_column_indices_cache = Arc::new(OnceLock::new());
585        self.by_row_data_cache = Arc::new(OnceLock::new());
586        self.by_row_indices_cache = Arc::new(OnceLock::new());
587        self.streamed_nnz = 0;
588        self.read_column_indptr()?;
589        self.read_row_indptr()?;
590        Ok(())
591    }
592
593    fn note_streamed_nnz(&mut self, n: u64) {
594        self.streamed_nnz += n;
595    }
596
597    fn streamed_nnz(&self) -> u64 {
598        self.streamed_nnz
599    }
600
601    fn reset_streamed_nnz(&mut self) {
602        self.streamed_nnz = 0;
603    }
604
605    fn read_column_indptr(&mut self) -> anyhow::Result<()> {
606        use zarrs::array::Array as ZArray;
607        let key = "/by_column/indptr";
608        if let Ok(indptr) = ZArray::open(self.read_store.clone(), key) {
609            let indptr_vec = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
610            self.by_column_indptr.clear();
611            self.by_column_indptr.extend(indptr_vec);
612        }
613        Ok(())
614    }
615
616    fn clean_preloaded_columns(&mut self) {
617        self.by_column_data = None;
618        self.by_column_indices = None;
619    }
620
621    /// preload columns' values and indices
622    fn preload_columns(&mut self) -> anyhow::Result<()> {
623        if let Some(nnz) = self.num_non_zeros() {
624            if !crate::sparse_io::preload_within_budget(nnz, "column") {
625                return Ok(());
626            }
627        }
628        use zarrs::array::Array as ZArray;
629
630        let key = "/by_column/data";
631        let data = ZArray::open(self.read_store.clone(), key)?;
632        let key = "/by_column/indices";
633        let indices = ZArray::open(self.read_store.clone(), key)?;
634
635        let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
636        let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
637
638        self.by_column_indices = Some(indices);
639        self.by_column_data = Some(data);
640        Ok(())
641    }
642
643    fn clean_preloaded_rows(&mut self) {
644        self.by_row_data = None;
645        self.by_row_indices = None;
646    }
647
648    /// preload rows' values and indices
649    fn preload_rows(&mut self) -> anyhow::Result<()> {
650        if let Some(nnz) = self.num_non_zeros() {
651            if !crate::sparse_io::preload_within_budget(nnz, "row") {
652                return Ok(());
653            }
654        }
655        use zarrs::array::Array as ZArray;
656
657        let data = ZArray::open(self.read_store.clone(), KEY_BY_ROW_DATA)?;
658        let indices = ZArray::open(self.read_store.clone(), KEY_BY_ROW_INDICES)?;
659
660        let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
661        let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
662
663        self.by_row_indices = Some(indices);
664        self.by_row_data = Some(data);
665        Ok(())
666    }
667
668    /// Helper function to keep the matrix shape
669    fn record_mtx_shape(&mut self, mtx_shape: Option<(usize, usize, usize)>) -> anyhow::Result<()> {
670        if let Some((nrow, ncol, nnz)) = mtx_shape {
671            let ws = self.write_store()?;
672            let read_store = self.read_store.clone();
673
674            let check_set_attr = |attr_name: &str, value: usize| -> anyhow::Result<()> {
675                let old_value = Self::_get_group_attr::<usize>(read_store.clone(), "/", attr_name);
676                let new_value = serde_json::to_value(value)?;
677
678                match old_value {
679                    Some(old_value) => {
680                        if old_value != new_value {
681                            return Err(anyhow!("{} mismatch", attr_name));
682                        }
683                    }
684                    _ => {
685                        Self::_set_group_attr(ws.clone(), "/", attr_name, &new_value)?;
686                    }
687                }
688                Ok(())
689            };
690
691            check_set_attr("nrow", nrow)?;
692            check_set_attr("ncol", ncol)?;
693            check_set_attr("nnz", nnz)?;
694        }
695        Ok(())
696    }
697
698    /// Helper function to create a new zarr backend file
699    fn initialize_backend(&mut self) -> anyhow::Result<()> {
700        use zarrs::group::GroupBuilder;
701
702        self.remove_backend_file()?;
703        let zarr_file = &self.file_name;
704        let store = Arc::new(FilesystemStore::new(zarr_file)?);
705        let root = GroupBuilder::new().build(store.clone(), "/")?;
706        root.store_metadata()?;
707
708        self.read_store = store.clone();
709        self.write_store = Some(store);
710        self.file_name = zarr_file.to_string();
711        self.max_column_name_idx = MAX_COLUMN_NAME_IDX;
712        self.max_row_name_idx = MAX_ROW_NAME_IDX;
713        self.by_column_indptr = vec![];
714        self.by_row_indptr = vec![];
715
716        Ok(())
717    }
718
719    /// Clean up the backend file
720    fn remove_backend_file(&self) -> anyhow::Result<()> {
721        let backend = std::path::Path::new(&self.file_name);
722        if backend.exists() {
723            if backend.is_file() {
724                std::fs::remove_file(backend)?;
725            } else {
726                std::fs::remove_dir_all(backend)?;
727            }
728        }
729        Ok(())
730    }
731
732    /// Access file name of the zarr backend
733    fn get_backend_file_name(&self) -> &str {
734        &self.file_name
735    }
736
737    fn backend_type(&self) -> SparseIoBackend {
738        SparseIoBackend::Zarr
739    }
740
741    /// Export the data to a mtx file. This will take time.
742    /// * `mtx_file`: mtx file to be written
743    fn to_mtx_file(&self, mtx_file: &str) -> anyhow::Result<()> {
744        if let (Some(ncol), Some(nrow), Some(nnz)) =
745            (self.num_columns(), self.num_rows(), self.num_non_zeros())
746        {
747            let (nrow, ncol, nnz) = (nrow, ncol, nnz);
748
749            let mut buf = open_buf_writer(mtx_file)?;
750            shared::write_mtx_header(&mut buf, nrow, ncol, nnz)?;
751
752            let (indptr, data, indices) = self.open_csc_triplets()?;
753            let indptr = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
754            debug_assert!(indptr.len() == ncol + 1);
755
756            // Stream the CSC value/row-index arrays in large sequential blocks
757            // (instead of two retrieves per column) and walk `indptr` to map
758            // each stored nonzero back to its column. CSC values are laid out
759            // column-major in column order, so emitting them in storage order
760            // reproduces the same column-by-column output as a per-column scan;
761            // empty columns are skipped by advancing the column pointer.
762            let total_nnz = indptr[ncol];
763            let mut jj = 0usize; // column owning the running nnz position
764            let mut pos = 0u64;
765            while pos < total_nnz {
766                let end = (pos + MTX_STREAM_BLOCK).min(total_nnz);
767                let subset = Self::create_subset(pos..end);
768                let data_block = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
769                let indices_block = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
770
771                for (k, (&val, &ii)) in data_block.iter().zip(&indices_block).enumerate() {
772                    let global = pos + k as u64;
773                    // advance to the column owning this nonzero (skips empties)
774                    while jj + 1 < indptr.len() && indptr[jj + 1] <= global {
775                        jj += 1;
776                    }
777                    // 1-based indices
778                    writeln!(buf, "{}\t{}\t{}", ii as usize + 1, jj + 1, val)?;
779                }
780                pos = end;
781            }
782            buf.flush()?;
783            Ok(())
784        } else {
785            Err(anyhow!("Unable to figure out the size of the backend data"))
786        }
787    }
788
789    /// Set row names for the matrix
790    /// * `row_name_file`: a file each line contains row name words
791    fn register_row_names_file(&mut self, row_name_file: &str) {
792        let _ = self.register_names_file(
793            "/row_names",
794            row_name_file,
795            0..self.max_row_name_idx,
796            ROW_SEP,
797        );
798    }
799
800    /// Set row names for the matrix
801    /// * `rows`: a vector of row names
802    fn register_row_names_vec(&mut self, rows: &[Box<str>]) {
803        let _ = self.register_names_vec("/row_names", rows);
804    }
805
806    /// Set column names for the matrix
807    /// * `column_name_file`: a file each line contains column name words
808    fn register_column_names_file(&mut self, column_name_file: &str) {
809        let _ = self.register_names_file(
810            "/column_names",
811            column_name_file,
812            0..self.max_column_name_idx,
813            COLUMN_SEP,
814        );
815    }
816
817    /// Set column names for the matrix
818    /// * `columns`: a vector of column names
819    fn register_column_names_vec(&mut self, columns: &[Box<str>]) {
820        let _ = self.register_names_vec("/column_names", columns);
821    }
822
823    /// Number of rows in the matrix
824    fn num_rows(&self) -> Option<usize> {
825        Self::_num_rows(self.read_store.clone())
826    }
827
828    /// Number of columns in the matrix
829    fn num_columns(&self) -> Option<usize> {
830        Self::_num_columns(self.read_store.clone())
831    }
832
833    /// Number of non-zero elements in the matrix
834    fn num_non_zeros(&self) -> Option<usize> {
835        Self::_num_nnz(self.read_store.clone())
836    }
837
838    /// Add arbitrary names (a vector of strings)
839    /// * `group_name`: group name
840    /// * `name_file`: a file each line contains name words
841    /// * `name_columns`: range of columns to be used for name
842    /// * `name_sep`: separator for name columns
843    fn register_names_file(
844        &mut self,
845        key: &str,
846        name_file: &str,
847        name_columns: Range<usize>,
848        name_sep: &str,
849    ) -> anyhow::Result<()> {
850        let names = parse_name_file(name_file, name_columns, name_sep)?;
851        self.new_filled_vector(key, data_type::string(), &names)?;
852        Ok(())
853    }
854
855    /// Add arbitrary names (a vector of strings)
856    /// * `group_name`: group name
857    /// * `names`: a file each line contains name words
858    fn register_names_vec(&mut self, key: &str, names: &[Box<str>]) -> anyhow::Result<()> {
859        let names_vec: Vec<String> = names.iter().map(|x| x.to_string()).collect();
860        self.new_filled_vector(key, data_type::string(), &names_vec)?;
861        Ok(())
862    }
863
864    fn row_names(&self) -> anyhow::Result<Vec<Box<str>>> {
865        self.retrieve_registered_names("/row_names")
866    }
867
868    fn column_names(&self) -> anyhow::Result<Vec<Box<str>>> {
869        self.retrieve_registered_names("/column_names")
870    }
871
872    /// Get back the registered names
873    /// * `key`: key for the registered names
874    fn retrieve_registered_names(&self, key: &str) -> anyhow::Result<Vec<Box<str>>> {
875        Ok(self
876            ._retrieve_vector::<String>(key)?
877            .into_iter()
878            .map(|s| s.into_boxed_str())
879            .collect())
880    }
881
882    /// Read columns within the range and return a vector of triplets (row, col, value)
883    /// * `col` : usize
884    ///
885    fn read_triplets_by_single_column(
886        &self,
887        j_data: usize,
888    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
889        use zarrs::array::Array as ZArray;
890
891        debug_assert!(!self.by_column_indptr.is_empty()); // pre-loaded
892        debug_assert!(j_data < self.num_columns().unwrap_or(0)); //
893
894        let indptr = &self.by_column_indptr;
895
896        debug_assert!((j_data + 1) < indptr.len());
897        debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
898
899        let nrow = self
900            .num_rows()
901            .ok_or(anyhow!("can't figure out the number of rows"))?;
902
903        if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
904            let ncol_out = 1;
905            let jj = 0;
906
907            // [start, end)
908            let start = indptr[j_data] as usize;
909            let end = indptr[j_data + 1] as usize;
910            let ret: Vec<(u64, u64, f32)> = indices[start..end]
911                .iter()
912                .zip(data[start..end].iter())
913                .map(|(&ii, &x_ij)| (ii, jj, x_ij))
914                .collect();
915
916            Ok((nrow, ncol_out, ret))
917        } else {
918            let key = "/by_column/data";
919            let data = ZArray::open(self.read_store.clone(), key)?;
920            let key = "/by_column/indices";
921            let indices = ZArray::open(self.read_store.clone(), key)?;
922
923            let ncol_out = 1;
924            let jj = 0;
925
926            // [start, end)
927            let start = indptr[j_data];
928            let end = indptr[j_data + 1];
929
930            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((end - start) as usize);
931
932            if start < end {
933                let subset = Self::create_subset(start..end);
934                let data_slice = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
935                let indices_slice = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
936
937                for k in 0..(end - start) {
938                    let x_ij = data_slice[k as usize];
939                    let ii = indices_slice[k as usize];
940                    debug_assert!((ii as usize) < nrow);
941                    ret.push((ii, jj, x_ij));
942                }
943            }
944
945            Ok((nrow, ncol_out, ret))
946        }
947    }
948
949    /// Read columns within the range and return dense `ndarray::Array2`
950    /// * `columns` : range e.g., 0..3 -> [0, 1, 2] or vec![0, 1, 2]
951    ///
952    fn read_triplets_by_columns(
953        &self,
954        columns: Self::IndexIter,
955    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
956        debug_assert!(!self.by_column_indptr.is_empty());
957        let indptr = &self.by_column_indptr;
958        let columns_vec = columns.into_iter().collect::<Vec<usize>>();
959
960        debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
961
962        let nrow = self
963            .num_rows()
964            .ok_or(anyhow!("can't figure out the number of rows"))?;
965
966        let ncol = self
967            .num_columns()
968            .ok_or(anyhow!("can't figure out the number of columns"))?;
969
970        let ncol_out = columns_vec.len();
971
972        if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
973            let min_start = columns_vec
974                .iter()
975                .map(|&j_data| indptr[j_data])
976                .min()
977                .unwrap_or(0);
978
979            let max_end = columns_vec
980                .iter()
981                .map(|&j_data| indptr[j_data + 1])
982                .max()
983                .unwrap_or(0);
984
985            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((max_end - min_start) as usize);
986
987            for (jj, &j_data) in columns_vec.iter().enumerate() {
988                let jj = jj as u64;
989                let start = indptr[j_data] as usize;
990                let end = indptr[j_data + 1] as usize;
991                for (&ii, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
992                    ret.push((ii, jj, x_ij));
993                }
994            }
995
996            Ok((nrow, ncol_out, ret))
997        } else {
998            // CSC: tag = output column, inner = row. Sort by indptr.start so
999            // abutting/overlapping ranges fuse into one retrieve; across-chunk
1000            // redundancy for non-mergeable ranges is caught by the chunk cache.
1001            let mut tagged: Vec<(u64, u64, u64)> = columns_vec
1002                .iter()
1003                .enumerate()
1004                .filter_map(|(jj, &j_data)| {
1005                    if j_data >= ncol {
1006                        return None;
1007                    }
1008                    let start = indptr[j_data];
1009                    let end = indptr[j_data + 1];
1010                    (start < end).then_some((jj as u64, start, end))
1011                })
1012                .collect();
1013            tagged.sort_by_key(|&(_, start, _)| start);
1014
1015            let data_cache = self.cache_for(&self.by_column_data_cache, KEY_BY_COLUMN_DATA)?;
1016            let indices_cache =
1017                self.cache_for(&self.by_column_indices_cache, KEY_BY_COLUMN_INDICES)?;
1018
1019            let opts = zarrs::array::CodecOptions::default();
1020            let ret = shared::coalesce_and_emit(
1021                &tagged,
1022                nrow,
1023                |jj, ii, val| (ii, jj, val),
1024                |s, e| {
1025                    let subset = Self::create_subset(s..e);
1026                    let data_buf =
1027                        <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1028                            Vec<f32>,
1029                        >(data_cache, &subset, &opts)?;
1030                    let indices_buf =
1031                        <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1032                            Vec<u64>,
1033                        >(indices_cache, &subset, &opts)?;
1034                    Ok((data_buf, indices_buf))
1035                },
1036            )?;
1037            Ok((nrow, ncol_out, ret))
1038        }
1039    }
1040
1041    fn csc_column_arrays(&self) -> Option<(&[u64], &[u64], &[f32])> {
1042        match (
1043            self.by_column_data.as_ref(),
1044            self.by_column_indices.as_ref(),
1045        ) {
1046            (Some(data), Some(indices)) if !self.by_column_indptr.is_empty() => Some((
1047                self.by_column_indptr.as_slice(),
1048                indices.as_slice(),
1049                data.as_slice(),
1050            )),
1051            _ => None,
1052        }
1053    }
1054
1055    /// Read rows within the range and return a vector of triplets (row, col, value)
1056    /// * `rows` : range e.g., 0..3 -> [0, 1, 2] or vec![0, 1, 2]
1057    ///
1058    fn read_triplets_by_rows(
1059        &self,
1060        rows: Self::IndexIter,
1061    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
1062        debug_assert!(!self.by_row_indptr.is_empty());
1063        let indptr = &self.by_row_indptr;
1064        debug_assert!(indptr.len() > self.num_rows().unwrap_or(0));
1065
1066        let rows_vec = rows.into_iter().collect::<Vec<_>>();
1067
1068        let (nrow, ncol) = match (self.num_rows(), self.num_columns()) {
1069            (Some(nrow), Some(ncol)) => (nrow, ncol),
1070            _ => return Err(anyhow!("Unable to figure out the size of the backend data")),
1071        };
1072        let nrow_out = rows_vec.len();
1073
1074        if let (Some(data), Some(indices)) = (&self.by_row_data, &self.by_row_indices) {
1075            let mut nnz_total: usize = 0;
1076            let valid: Vec<(u64, usize)> = rows_vec
1077                .iter()
1078                .enumerate()
1079                .filter_map(|(ii, &i_data)| {
1080                    if i_data >= nrow {
1081                        return None;
1082                    }
1083                    nnz_total += (indptr[i_data + 1] - indptr[i_data]) as usize;
1084                    Some((ii as u64, i_data))
1085                })
1086                .collect();
1087
1088            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity(nnz_total);
1089            for (ii, i_data) in valid {
1090                let start = indptr[i_data] as usize;
1091                let end = indptr[i_data + 1] as usize;
1092                for (&jj, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
1093                    ret.push((ii, jj, x_ij));
1094                }
1095            }
1096            return Ok((nrow_out, ncol, ret));
1097        }
1098
1099        // CSR: tag = output row, inner = column.
1100        let mut tagged: Vec<(u64, u64, u64)> = rows_vec
1101            .iter()
1102            .enumerate()
1103            .filter_map(|(ii, &i_data)| {
1104                if i_data >= nrow {
1105                    return None;
1106                }
1107                debug_assert!((i_data + 1) < indptr.len());
1108                let start = indptr[i_data];
1109                let end = indptr[i_data + 1];
1110                (start < end).then_some((ii as u64, start, end))
1111            })
1112            .collect();
1113        tagged.sort_by_key(|&(_, start, _)| start);
1114
1115        let data_cache = self.cache_for(&self.by_row_data_cache, KEY_BY_ROW_DATA)?;
1116        let indices_cache = self.cache_for(&self.by_row_indices_cache, KEY_BY_ROW_INDICES)?;
1117
1118        let opts = zarrs::array::CodecOptions::default();
1119        let ret = shared::coalesce_and_emit(
1120            &tagged,
1121            ncol,
1122            |ii, jj, val| (ii, jj, val),
1123            |s, e| {
1124                let subset = Self::create_subset(s..e);
1125                let data_buf = <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1126                    Vec<f32>,
1127                >(data_cache, &subset, &opts)?;
1128                let indices_buf =
1129                    <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<Vec<u64>>(
1130                        indices_cache,
1131                        &subset,
1132                        &opts,
1133                    )?;
1134                Ok((data_buf, indices_buf))
1135            },
1136        )?;
1137        Ok((nrow_out, ncol, ret))
1138    }
1139    /// CSR data structure in Zarr backend
1140    ///
1141    /// ```text
1142    ///     └── by_row
1143    ///         ├── data
1144    ///         ├── indices (column indices)
1145    ///         └── isndptr (row pointers)
1146    /// ```
1147    fn record_csr_dataset_backend(
1148        &mut self,
1149        csr_cols: &[u64],
1150        csr_vals: &[f32],
1151        csr_rowptr: &[u64],
1152    ) -> anyhow::Result<()> {
1153        // open or create the group "/by_row"
1154        let key = "/by_row";
1155        self._add_group(key)?;
1156
1157        let key = "/by_row/data";
1158        self.new_filled_vector(key, data_type::float32(), csr_vals)?;
1159        let key = "/by_row/indices";
1160        self.new_filled_vector(key, data_type::uint64(), csr_cols)?;
1161        let key = "/by_row/indptr";
1162        self.new_filled_vector(key, data_type::uint64(), csr_rowptr)?;
1163
1164        Ok(())
1165    }
1166
1167    /// CSC data structure in Zarr backend
1168    ///
1169    /// ```text
1170    /// Helper function to record the CSC dataset
1171    ///     ├── by_column
1172    ///     │   ├── data
1173    ///     │   ├── indices (row indices)
1174    ///     │   └── indptr (column pointers)
1175    /// ```
1176    fn record_csc_dataset_backend(
1177        &mut self,
1178        csc_rows: &[u64],
1179        csc_vals: &[f32],
1180        csc_colptr: &[u64],
1181    ) -> anyhow::Result<()> {
1182        // open or create the group "/by_column"
1183        let key = "/by_column";
1184        self._add_group(key)?;
1185
1186        let key = "/by_column/data";
1187        self.new_filled_vector(key, data_type::float32(), csc_vals)?;
1188        let key = "/by_column/indices";
1189        self.new_filled_vector(key, data_type::uint64(), csc_rows)?;
1190        let key = "/by_column/indptr";
1191        self.new_filled_vector(key, data_type::uint64(), csc_colptr)?;
1192
1193        Ok(())
1194    }
1195
1196    fn cs_create(&mut self, key: CsKey, len: usize) -> anyhow::Result<()> {
1197        let (group, path, dt, elem_bytes) = match key {
1198            CsKey::CscData => (
1199                "/by_column",
1200                "/by_column/data",
1201                data_type::float32(),
1202                std::mem::size_of::<f32>(),
1203            ),
1204            CsKey::CscIndices => (
1205                "/by_column",
1206                "/by_column/indices",
1207                data_type::uint64(),
1208                std::mem::size_of::<u64>(),
1209            ),
1210            CsKey::CscIndptr => (
1211                "/by_column",
1212                "/by_column/indptr",
1213                data_type::uint64(),
1214                std::mem::size_of::<u64>(),
1215            ),
1216            CsKey::CsrData => (
1217                "/by_row",
1218                "/by_row/data",
1219                data_type::float32(),
1220                std::mem::size_of::<f32>(),
1221            ),
1222            CsKey::CsrIndices => (
1223                "/by_row",
1224                "/by_row/indices",
1225                data_type::uint64(),
1226                std::mem::size_of::<u64>(),
1227            ),
1228            CsKey::CsrIndptr => (
1229                "/by_row",
1230                "/by_row/indptr",
1231                data_type::uint64(),
1232                std::mem::size_of::<u64>(),
1233            ),
1234        };
1235        self._add_group(group)?;
1236        self.create_shaped_vector(path, dt, elem_bytes, len)
1237    }
1238
1239    fn cs_write_u64(&mut self, key: CsKey, offset: u64, data: &[u64]) -> anyhow::Result<()> {
1240        let path = match key {
1241            CsKey::CscIndices => "/by_column/indices",
1242            CsKey::CscIndptr => "/by_column/indptr",
1243            CsKey::CsrIndices => "/by_row/indices",
1244            CsKey::CsrIndptr => "/by_row/indptr",
1245            CsKey::CscData | CsKey::CsrData => {
1246                return Err(anyhow!("cs_write_u64 called on f32 slot {:?}", key));
1247            }
1248        };
1249        self.write_slab_u64(path, offset, data)
1250    }
1251
1252    fn cs_write_f32(&mut self, key: CsKey, offset: u64, data: &[f32]) -> anyhow::Result<()> {
1253        let path = match key {
1254            CsKey::CscData => "/by_column/data",
1255            CsKey::CsrData => "/by_row/data",
1256            _ => {
1257                return Err(anyhow!("cs_write_f32 called on u64 slot {:?}", key));
1258            }
1259        };
1260        self.write_slab_f32(path, offset, data)
1261    }
1262}