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 reopen_backend(&mut self) -> anyhow::Result<()> {
558        // Path-addressed store: rebuild it, refresh the resident indptrs — and
559        // DROP the decoded-chunk LRU caches, which pin arrays of the store they
560        // were built on. Keeping them served pre-swap chunk contents against
561        // post-swap indptrs: a matrix that read back with row indices past its
562        // own nrow, from a file that was byte-for-byte correct on disk.
563        let store = Arc::new(FilesystemStore::new(&self.file_name)?);
564        self.read_store = store.clone();
565        self.write_store = Some(store);
566        self.by_column_data_cache = Arc::new(OnceLock::new());
567        self.by_column_indices_cache = Arc::new(OnceLock::new());
568        self.by_row_data_cache = Arc::new(OnceLock::new());
569        self.by_row_indices_cache = Arc::new(OnceLock::new());
570        self.streamed_nnz = 0;
571        self.read_column_indptr()?;
572        self.read_row_indptr()?;
573        Ok(())
574    }
575
576    fn note_streamed_nnz(&mut self, n: u64) {
577        self.streamed_nnz += n;
578    }
579
580    fn streamed_nnz(&self) -> u64 {
581        self.streamed_nnz
582    }
583
584    fn reset_streamed_nnz(&mut self) {
585        self.streamed_nnz = 0;
586    }
587
588    fn read_column_indptr(&mut self) -> anyhow::Result<()> {
589        use zarrs::array::Array as ZArray;
590        let key = "/by_column/indptr";
591        if let Ok(indptr) = ZArray::open(self.read_store.clone(), key) {
592            let indptr_vec = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
593            self.by_column_indptr.clear();
594            self.by_column_indptr.extend(indptr_vec);
595        }
596        Ok(())
597    }
598
599    fn clean_preloaded_columns(&mut self) {
600        self.by_column_data = None;
601        self.by_column_indices = None;
602    }
603
604    /// preload columns' values and indices
605    fn preload_columns(&mut self) -> anyhow::Result<()> {
606        if let Some(nnz) = self.num_non_zeros() {
607            if !crate::sparse_io::preload_within_budget(nnz, "column") {
608                return Ok(());
609            }
610        }
611        use zarrs::array::Array as ZArray;
612
613        let key = "/by_column/data";
614        let data = ZArray::open(self.read_store.clone(), key)?;
615        let key = "/by_column/indices";
616        let indices = ZArray::open(self.read_store.clone(), key)?;
617
618        let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
619        let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
620
621        self.by_column_indices = Some(indices);
622        self.by_column_data = Some(data);
623        Ok(())
624    }
625
626    fn clean_preloaded_rows(&mut self) {
627        self.by_row_data = None;
628        self.by_row_indices = None;
629    }
630
631    /// preload rows' values and indices
632    fn preload_rows(&mut self) -> anyhow::Result<()> {
633        if let Some(nnz) = self.num_non_zeros() {
634            if !crate::sparse_io::preload_within_budget(nnz, "row") {
635                return Ok(());
636            }
637        }
638        use zarrs::array::Array as ZArray;
639
640        let data = ZArray::open(self.read_store.clone(), KEY_BY_ROW_DATA)?;
641        let indices = ZArray::open(self.read_store.clone(), KEY_BY_ROW_INDICES)?;
642
643        let data = data.retrieve_array_subset::<Vec<f32>>(&data.subset_all())?;
644        let indices = indices.retrieve_array_subset::<Vec<u64>>(&indices.subset_all())?;
645
646        self.by_row_indices = Some(indices);
647        self.by_row_data = Some(data);
648        Ok(())
649    }
650
651    /// Helper function to keep the matrix shape
652    fn record_mtx_shape(&mut self, mtx_shape: Option<(usize, usize, usize)>) -> anyhow::Result<()> {
653        if let Some((nrow, ncol, nnz)) = mtx_shape {
654            let ws = self.write_store()?;
655            let read_store = self.read_store.clone();
656
657            let check_set_attr = |attr_name: &str, value: usize| -> anyhow::Result<()> {
658                let old_value = Self::_get_group_attr::<usize>(read_store.clone(), "/", attr_name);
659                let new_value = serde_json::to_value(value)?;
660
661                match old_value {
662                    Some(old_value) => {
663                        if old_value != new_value {
664                            return Err(anyhow!("{} mismatch", attr_name));
665                        }
666                    }
667                    _ => {
668                        Self::_set_group_attr(ws.clone(), "/", attr_name, &new_value)?;
669                    }
670                }
671                Ok(())
672            };
673
674            check_set_attr("nrow", nrow)?;
675            check_set_attr("ncol", ncol)?;
676            check_set_attr("nnz", nnz)?;
677        }
678        Ok(())
679    }
680
681    /// Helper function to create a new zarr backend file
682    fn initialize_backend(&mut self) -> anyhow::Result<()> {
683        use zarrs::group::GroupBuilder;
684
685        self.remove_backend_file()?;
686        let zarr_file = &self.file_name;
687        let store = Arc::new(FilesystemStore::new(zarr_file)?);
688        let root = GroupBuilder::new().build(store.clone(), "/")?;
689        root.store_metadata()?;
690
691        self.read_store = store.clone();
692        self.write_store = Some(store);
693        self.file_name = zarr_file.to_string();
694        self.max_column_name_idx = MAX_COLUMN_NAME_IDX;
695        self.max_row_name_idx = MAX_ROW_NAME_IDX;
696        self.by_column_indptr = vec![];
697        self.by_row_indptr = vec![];
698
699        Ok(())
700    }
701
702    /// Clean up the backend file
703    fn remove_backend_file(&self) -> anyhow::Result<()> {
704        let backend = std::path::Path::new(&self.file_name);
705        if backend.exists() {
706            if backend.is_file() {
707                std::fs::remove_file(backend)?;
708            } else {
709                std::fs::remove_dir_all(backend)?;
710            }
711        }
712        Ok(())
713    }
714
715    /// Access file name of the zarr backend
716    fn get_backend_file_name(&self) -> &str {
717        &self.file_name
718    }
719
720    fn backend_type(&self) -> SparseIoBackend {
721        SparseIoBackend::Zarr
722    }
723
724    /// Export the data to a mtx file. This will take time.
725    /// * `mtx_file`: mtx file to be written
726    fn to_mtx_file(&self, mtx_file: &str) -> anyhow::Result<()> {
727        if let (Some(ncol), Some(nrow), Some(nnz)) =
728            (self.num_columns(), self.num_rows(), self.num_non_zeros())
729        {
730            let (nrow, ncol, nnz) = (nrow, ncol, nnz);
731
732            let mut buf = open_buf_writer(mtx_file)?;
733            shared::write_mtx_header(&mut buf, nrow, ncol, nnz)?;
734
735            let (indptr, data, indices) = self.open_csc_triplets()?;
736            let indptr = indptr.retrieve_array_subset::<Vec<u64>>(&indptr.subset_all())?;
737            debug_assert!(indptr.len() == ncol + 1);
738
739            // Stream the CSC value/row-index arrays in large sequential blocks
740            // (instead of two retrieves per column) and walk `indptr` to map
741            // each stored nonzero back to its column. CSC values are laid out
742            // column-major in column order, so emitting them in storage order
743            // reproduces the same column-by-column output as a per-column scan;
744            // empty columns are skipped by advancing the column pointer.
745            let total_nnz = indptr[ncol];
746            let mut jj = 0usize; // column owning the running nnz position
747            let mut pos = 0u64;
748            while pos < total_nnz {
749                let end = (pos + MTX_STREAM_BLOCK).min(total_nnz);
750                let subset = Self::create_subset(pos..end);
751                let data_block = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
752                let indices_block = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
753
754                for (k, (&val, &ii)) in data_block.iter().zip(&indices_block).enumerate() {
755                    let global = pos + k as u64;
756                    // advance to the column owning this nonzero (skips empties)
757                    while jj + 1 < indptr.len() && indptr[jj + 1] <= global {
758                        jj += 1;
759                    }
760                    // 1-based indices
761                    writeln!(buf, "{}\t{}\t{}", ii as usize + 1, jj + 1, val)?;
762                }
763                pos = end;
764            }
765            buf.flush()?;
766            Ok(())
767        } else {
768            Err(anyhow!("Unable to figure out the size of the backend data"))
769        }
770    }
771
772    /// Set row names for the matrix
773    /// * `row_name_file`: a file each line contains row name words
774    fn register_row_names_file(&mut self, row_name_file: &str) {
775        let _ = self.register_names_file(
776            "/row_names",
777            row_name_file,
778            0..self.max_row_name_idx,
779            ROW_SEP,
780        );
781    }
782
783    /// Set row names for the matrix
784    /// * `rows`: a vector of row names
785    fn register_row_names_vec(&mut self, rows: &[Box<str>]) {
786        let _ = self.register_names_vec("/row_names", rows);
787    }
788
789    /// Set column names for the matrix
790    /// * `column_name_file`: a file each line contains column name words
791    fn register_column_names_file(&mut self, column_name_file: &str) {
792        let _ = self.register_names_file(
793            "/column_names",
794            column_name_file,
795            0..self.max_column_name_idx,
796            COLUMN_SEP,
797        );
798    }
799
800    /// Set column names for the matrix
801    /// * `columns`: a vector of column names
802    fn register_column_names_vec(&mut self, columns: &[Box<str>]) {
803        let _ = self.register_names_vec("/column_names", columns);
804    }
805
806    /// Number of rows in the matrix
807    fn num_rows(&self) -> Option<usize> {
808        Self::_num_rows(self.read_store.clone())
809    }
810
811    /// Number of columns in the matrix
812    fn num_columns(&self) -> Option<usize> {
813        Self::_num_columns(self.read_store.clone())
814    }
815
816    /// Number of non-zero elements in the matrix
817    fn num_non_zeros(&self) -> Option<usize> {
818        Self::_num_nnz(self.read_store.clone())
819    }
820
821    /// Add arbitrary names (a vector of strings)
822    /// * `group_name`: group name
823    /// * `name_file`: a file each line contains name words
824    /// * `name_columns`: range of columns to be used for name
825    /// * `name_sep`: separator for name columns
826    fn register_names_file(
827        &mut self,
828        key: &str,
829        name_file: &str,
830        name_columns: Range<usize>,
831        name_sep: &str,
832    ) -> anyhow::Result<()> {
833        let names = parse_name_file(name_file, name_columns, name_sep)?;
834        self.new_filled_vector(key, data_type::string(), &names)?;
835        Ok(())
836    }
837
838    /// Add arbitrary names (a vector of strings)
839    /// * `group_name`: group name
840    /// * `names`: a file each line contains name words
841    fn register_names_vec(&mut self, key: &str, names: &[Box<str>]) -> anyhow::Result<()> {
842        let names_vec: Vec<String> = names.iter().map(|x| x.to_string()).collect();
843        self.new_filled_vector(key, data_type::string(), &names_vec)?;
844        Ok(())
845    }
846
847    fn row_names(&self) -> anyhow::Result<Vec<Box<str>>> {
848        self.retrieve_registered_names("/row_names")
849    }
850
851    fn column_names(&self) -> anyhow::Result<Vec<Box<str>>> {
852        self.retrieve_registered_names("/column_names")
853    }
854
855    /// Get back the registered names
856    /// * `key`: key for the registered names
857    fn retrieve_registered_names(&self, key: &str) -> anyhow::Result<Vec<Box<str>>> {
858        Ok(self
859            ._retrieve_vector::<String>(key)?
860            .into_iter()
861            .map(|s| s.into_boxed_str())
862            .collect())
863    }
864
865    /// Read columns within the range and return a vector of triplets (row, col, value)
866    /// * `col` : usize
867    ///
868    fn read_triplets_by_single_column(
869        &self,
870        j_data: usize,
871    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
872        use zarrs::array::Array as ZArray;
873
874        debug_assert!(!self.by_column_indptr.is_empty()); // pre-loaded
875        debug_assert!(j_data < self.num_columns().unwrap_or(0)); //
876
877        let indptr = &self.by_column_indptr;
878
879        debug_assert!((j_data + 1) < indptr.len());
880        debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
881
882        let nrow = self
883            .num_rows()
884            .ok_or(anyhow!("can't figure out the number of rows"))?;
885
886        if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
887            let ncol_out = 1;
888            let jj = 0;
889
890            // [start, end)
891            let start = indptr[j_data] as usize;
892            let end = indptr[j_data + 1] as usize;
893            let ret: Vec<(u64, u64, f32)> = indices[start..end]
894                .iter()
895                .zip(data[start..end].iter())
896                .map(|(&ii, &x_ij)| (ii, jj, x_ij))
897                .collect();
898
899            Ok((nrow, ncol_out, ret))
900        } else {
901            let key = "/by_column/data";
902            let data = ZArray::open(self.read_store.clone(), key)?;
903            let key = "/by_column/indices";
904            let indices = ZArray::open(self.read_store.clone(), key)?;
905
906            let ncol_out = 1;
907            let jj = 0;
908
909            // [start, end)
910            let start = indptr[j_data];
911            let end = indptr[j_data + 1];
912
913            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((end - start) as usize);
914
915            if start < end {
916                let subset = Self::create_subset(start..end);
917                let data_slice = data.retrieve_array_subset::<Vec<f32>>(&subset)?;
918                let indices_slice = indices.retrieve_array_subset::<Vec<u64>>(&subset)?;
919
920                for k in 0..(end - start) {
921                    let x_ij = data_slice[k as usize];
922                    let ii = indices_slice[k as usize];
923                    debug_assert!((ii as usize) < nrow);
924                    ret.push((ii, jj, x_ij));
925                }
926            }
927
928            Ok((nrow, ncol_out, ret))
929        }
930    }
931
932    /// Read columns within the range and return dense `ndarray::Array2`
933    /// * `columns` : range e.g., 0..3 -> [0, 1, 2] or vec![0, 1, 2]
934    ///
935    fn read_triplets_by_columns(
936        &self,
937        columns: Self::IndexIter,
938    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
939        debug_assert!(!self.by_column_indptr.is_empty());
940        let indptr = &self.by_column_indptr;
941        let columns_vec = columns.into_iter().collect::<Vec<usize>>();
942
943        debug_assert!(indptr.len() > self.num_columns().unwrap_or(0));
944
945        let nrow = self
946            .num_rows()
947            .ok_or(anyhow!("can't figure out the number of rows"))?;
948
949        let ncol = self
950            .num_columns()
951            .ok_or(anyhow!("can't figure out the number of columns"))?;
952
953        let ncol_out = columns_vec.len();
954
955        if let (Some(data), Some(indices)) = (&self.by_column_data, &self.by_column_indices) {
956            let min_start = columns_vec
957                .iter()
958                .map(|&j_data| indptr[j_data])
959                .min()
960                .unwrap_or(0);
961
962            let max_end = columns_vec
963                .iter()
964                .map(|&j_data| indptr[j_data + 1])
965                .max()
966                .unwrap_or(0);
967
968            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity((max_end - min_start) as usize);
969
970            for (jj, &j_data) in columns_vec.iter().enumerate() {
971                let jj = jj as u64;
972                let start = indptr[j_data] as usize;
973                let end = indptr[j_data + 1] as usize;
974                for (&ii, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
975                    ret.push((ii, jj, x_ij));
976                }
977            }
978
979            Ok((nrow, ncol_out, ret))
980        } else {
981            // CSC: tag = output column, inner = row. Sort by indptr.start so
982            // abutting/overlapping ranges fuse into one retrieve; across-chunk
983            // redundancy for non-mergeable ranges is caught by the chunk cache.
984            let mut tagged: Vec<(u64, u64, u64)> = columns_vec
985                .iter()
986                .enumerate()
987                .filter_map(|(jj, &j_data)| {
988                    if j_data >= ncol {
989                        return None;
990                    }
991                    let start = indptr[j_data];
992                    let end = indptr[j_data + 1];
993                    (start < end).then_some((jj as u64, start, end))
994                })
995                .collect();
996            tagged.sort_by_key(|&(_, start, _)| start);
997
998            let data_cache = self.cache_for(&self.by_column_data_cache, KEY_BY_COLUMN_DATA)?;
999            let indices_cache =
1000                self.cache_for(&self.by_column_indices_cache, KEY_BY_COLUMN_INDICES)?;
1001
1002            let opts = zarrs::array::CodecOptions::default();
1003            let ret = shared::coalesce_and_emit(
1004                &tagged,
1005                nrow,
1006                |jj, ii, val| (ii, jj, val),
1007                |s, e| {
1008                    let subset = Self::create_subset(s..e);
1009                    let data_buf =
1010                        <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1011                            Vec<f32>,
1012                        >(data_cache, &subset, &opts)?;
1013                    let indices_buf =
1014                        <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1015                            Vec<u64>,
1016                        >(indices_cache, &subset, &opts)?;
1017                    Ok((data_buf, indices_buf))
1018                },
1019            )?;
1020            Ok((nrow, ncol_out, ret))
1021        }
1022    }
1023
1024    fn csc_column_arrays(&self) -> Option<(&[u64], &[u64], &[f32])> {
1025        match (
1026            self.by_column_data.as_ref(),
1027            self.by_column_indices.as_ref(),
1028        ) {
1029            (Some(data), Some(indices)) if !self.by_column_indptr.is_empty() => Some((
1030                self.by_column_indptr.as_slice(),
1031                indices.as_slice(),
1032                data.as_slice(),
1033            )),
1034            _ => None,
1035        }
1036    }
1037
1038    /// Read rows within the range and return a vector of triplets (row, col, value)
1039    /// * `rows` : range e.g., 0..3 -> [0, 1, 2] or vec![0, 1, 2]
1040    ///
1041    fn read_triplets_by_rows(
1042        &self,
1043        rows: Self::IndexIter,
1044    ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)> {
1045        debug_assert!(!self.by_row_indptr.is_empty());
1046        let indptr = &self.by_row_indptr;
1047        debug_assert!(indptr.len() > self.num_rows().unwrap_or(0));
1048
1049        let rows_vec = rows.into_iter().collect::<Vec<_>>();
1050
1051        let (nrow, ncol) = match (self.num_rows(), self.num_columns()) {
1052            (Some(nrow), Some(ncol)) => (nrow, ncol),
1053            _ => return Err(anyhow!("Unable to figure out the size of the backend data")),
1054        };
1055        let nrow_out = rows_vec.len();
1056
1057        if let (Some(data), Some(indices)) = (&self.by_row_data, &self.by_row_indices) {
1058            let mut nnz_total: usize = 0;
1059            let valid: Vec<(u64, usize)> = rows_vec
1060                .iter()
1061                .enumerate()
1062                .filter_map(|(ii, &i_data)| {
1063                    if i_data >= nrow {
1064                        return None;
1065                    }
1066                    nnz_total += (indptr[i_data + 1] - indptr[i_data]) as usize;
1067                    Some((ii as u64, i_data))
1068                })
1069                .collect();
1070
1071            let mut ret: Vec<(u64, u64, f32)> = Vec::with_capacity(nnz_total);
1072            for (ii, i_data) in valid {
1073                let start = indptr[i_data] as usize;
1074                let end = indptr[i_data + 1] as usize;
1075                for (&jj, &x_ij) in indices[start..end].iter().zip(data[start..end].iter()) {
1076                    ret.push((ii, jj, x_ij));
1077                }
1078            }
1079            return Ok((nrow_out, ncol, ret));
1080        }
1081
1082        // CSR: tag = output row, inner = column.
1083        let mut tagged: Vec<(u64, u64, u64)> = rows_vec
1084            .iter()
1085            .enumerate()
1086            .filter_map(|(ii, &i_data)| {
1087                if i_data >= nrow {
1088                    return None;
1089                }
1090                debug_assert!((i_data + 1) < indptr.len());
1091                let start = indptr[i_data];
1092                let end = indptr[i_data + 1];
1093                (start < end).then_some((ii as u64, start, end))
1094            })
1095            .collect();
1096        tagged.sort_by_key(|&(_, start, _)| start);
1097
1098        let data_cache = self.cache_for(&self.by_row_data_cache, KEY_BY_ROW_DATA)?;
1099        let indices_cache = self.cache_for(&self.by_row_indices_cache, KEY_BY_ROW_INDICES)?;
1100
1101        let opts = zarrs::array::CodecOptions::default();
1102        let ret = shared::coalesce_and_emit(
1103            &tagged,
1104            ncol,
1105            |ii, jj, val| (ii, jj, val),
1106            |s, e| {
1107                let subset = Self::create_subset(s..e);
1108                let data_buf = <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<
1109                    Vec<f32>,
1110                >(data_cache, &subset, &opts)?;
1111                let indices_buf =
1112                    <_ as zarrs::array::chunk_cache::ChunkCache>::retrieve_array_subset::<Vec<u64>>(
1113                        indices_cache,
1114                        &subset,
1115                        &opts,
1116                    )?;
1117                Ok((data_buf, indices_buf))
1118            },
1119        )?;
1120        Ok((nrow_out, ncol, ret))
1121    }
1122    /// CSR data structure in Zarr backend
1123    ///
1124    /// ```text
1125    ///     └── by_row
1126    ///         ├── data
1127    ///         ├── indices (column indices)
1128    ///         └── isndptr (row pointers)
1129    /// ```
1130    fn record_csr_dataset_backend(
1131        &mut self,
1132        csr_cols: &[u64],
1133        csr_vals: &[f32],
1134        csr_rowptr: &[u64],
1135    ) -> anyhow::Result<()> {
1136        // open or create the group "/by_row"
1137        let key = "/by_row";
1138        self._add_group(key)?;
1139
1140        let key = "/by_row/data";
1141        self.new_filled_vector(key, data_type::float32(), csr_vals)?;
1142        let key = "/by_row/indices";
1143        self.new_filled_vector(key, data_type::uint64(), csr_cols)?;
1144        let key = "/by_row/indptr";
1145        self.new_filled_vector(key, data_type::uint64(), csr_rowptr)?;
1146
1147        Ok(())
1148    }
1149
1150    /// CSC data structure in Zarr backend
1151    ///
1152    /// ```text
1153    /// Helper function to record the CSC dataset
1154    ///     ├── by_column
1155    ///     │   ├── data
1156    ///     │   ├── indices (row indices)
1157    ///     │   └── indptr (column pointers)
1158    /// ```
1159    fn record_csc_dataset_backend(
1160        &mut self,
1161        csc_rows: &[u64],
1162        csc_vals: &[f32],
1163        csc_colptr: &[u64],
1164    ) -> anyhow::Result<()> {
1165        // open or create the group "/by_column"
1166        let key = "/by_column";
1167        self._add_group(key)?;
1168
1169        let key = "/by_column/data";
1170        self.new_filled_vector(key, data_type::float32(), csc_vals)?;
1171        let key = "/by_column/indices";
1172        self.new_filled_vector(key, data_type::uint64(), csc_rows)?;
1173        let key = "/by_column/indptr";
1174        self.new_filled_vector(key, data_type::uint64(), csc_colptr)?;
1175
1176        Ok(())
1177    }
1178
1179    fn cs_create(&mut self, key: CsKey, len: usize) -> anyhow::Result<()> {
1180        let (group, path, dt, elem_bytes) = match key {
1181            CsKey::CscData => (
1182                "/by_column",
1183                "/by_column/data",
1184                data_type::float32(),
1185                std::mem::size_of::<f32>(),
1186            ),
1187            CsKey::CscIndices => (
1188                "/by_column",
1189                "/by_column/indices",
1190                data_type::uint64(),
1191                std::mem::size_of::<u64>(),
1192            ),
1193            CsKey::CscIndptr => (
1194                "/by_column",
1195                "/by_column/indptr",
1196                data_type::uint64(),
1197                std::mem::size_of::<u64>(),
1198            ),
1199            CsKey::CsrData => (
1200                "/by_row",
1201                "/by_row/data",
1202                data_type::float32(),
1203                std::mem::size_of::<f32>(),
1204            ),
1205            CsKey::CsrIndices => (
1206                "/by_row",
1207                "/by_row/indices",
1208                data_type::uint64(),
1209                std::mem::size_of::<u64>(),
1210            ),
1211            CsKey::CsrIndptr => (
1212                "/by_row",
1213                "/by_row/indptr",
1214                data_type::uint64(),
1215                std::mem::size_of::<u64>(),
1216            ),
1217        };
1218        self._add_group(group)?;
1219        self.create_shaped_vector(path, dt, elem_bytes, len)
1220    }
1221
1222    fn cs_write_u64(&mut self, key: CsKey, offset: u64, data: &[u64]) -> anyhow::Result<()> {
1223        let path = match key {
1224            CsKey::CscIndices => "/by_column/indices",
1225            CsKey::CscIndptr => "/by_column/indptr",
1226            CsKey::CsrIndices => "/by_row/indices",
1227            CsKey::CsrIndptr => "/by_row/indptr",
1228            CsKey::CscData | CsKey::CsrData => {
1229                return Err(anyhow!("cs_write_u64 called on f32 slot {:?}", key));
1230            }
1231        };
1232        self.write_slab_u64(path, offset, data)
1233    }
1234
1235    fn cs_write_f32(&mut self, key: CsKey, offset: u64, data: &[f32]) -> anyhow::Result<()> {
1236        let path = match key {
1237            CsKey::CscData => "/by_column/data",
1238            CsKey::CsrData => "/by_row/data",
1239            _ => {
1240                return Err(anyhow!("cs_write_f32 called on u64 slot {:?}", key));
1241            }
1242        };
1243        self.write_slab_f32(path, offset, data)
1244    }
1245}