Skip to main content

data_beans/sparse_io_vector/
read.rs

1#![allow(dead_code)]
2
3use super::*;
4
5impl SparseIoVec {
6    ////////////////////
7    // access columns //
8    ////////////////////
9
10    /// Read a single column by global index and append offset triplets.
11    /// Under [`ColumnAlignment::Union`] one global column can be backed
12    /// by multiple `(didx, loc)` pairs — their triplets all land in the
13    /// same output column and `col_offset` advances by 1.
14    ///
15    /// `pub(super)` so the matched/neighbour read paths in `matched.rs` (a
16    /// sibling module) can reuse it.
17    pub(super) fn read_column_offset(
18        &self,
19        glob: usize,
20        col_offset: &mut usize,
21        triplets: &mut Vec<(u64, u64, f32)>,
22    ) -> anyhow::Result<()> {
23        let off = *col_offset as u64;
24        let g2c = self.global_to_compact_row.as_slice();
25        for source in self.col_to_data[glob].iter() {
26            let didx = source.backend as usize;
27            let loc = source.local_col as usize;
28            let (_, _, loc_triplets) = self.data_vec[didx].read_triplets_by_single_column(loc)?;
29            let l2g = self.data_local_to_global_row[didx].as_slice();
30            for (i, _j, v) in loc_triplets {
31                if let Some(c) = g2c[l2g[i as usize]] {
32                    triplets.push((c as u64, off, v));
33                }
34            }
35        }
36        *col_offset += 1;
37        Ok(())
38    }
39
40    /// Stream every nonzero of the selected `cells` (global column
41    /// indices) through `f(row, col, val)` **without** materializing the
42    /// full triplet vector. `row` is a compact-row index and `col` runs
43    /// over `0..ncol` in the iteration order of `cells` — exactly the
44    /// indices [`Self::columns_triplets`] would emit.
45    ///
46    /// Columns are processed in slabs of at most `chunk_cols`, so the
47    /// only transient allocation is one backend slab's worth of
48    /// `(u64, u64, f32)` triplets: peak memory is bounded by `chunk_cols`,
49    /// not by the total nnz. Callers can therefore build a compact edge
50    /// list (e.g. 12-byte triplets) directly and never pay for the wide
51    /// intermediate. Returns the `(nrow, ncol)` dimensions.
52    pub fn for_each_triplet<I, F>(
53        &self,
54        cells: I,
55        chunk_cols: usize,
56        mut f: F,
57    ) -> anyhow::Result<(usize, usize)>
58    where
59        I: Iterator<Item = usize>,
60        F: FnMut(u64, u64, f32),
61    {
62        let nrow = self.num_rows();
63        let cells: Vec<usize> = cells.collect();
64        let ncol = cells.len();
65        let chunk = chunk_cols.max(1);
66        let g2c = self.global_to_compact_row.as_slice();
67
68        for slab_start in (0..ncol).step_by(chunk) {
69            let slab_end = (slab_start + chunk).min(ncol);
70
71            // Group this slab's cells by backend, tracking (local_col,
72            // out_col) where out_col is the index into the *full* `cells`
73            // sequence. Under `ColumnAlignment::Union` one global cell can
74            // contribute entries to multiple backend groups (one per
75            // backend that observed it); all those reads target the same
76            // out_col.
77            let mut backend_groups: HashMap<usize, Vec<(usize, usize)>> = HashMap::default();
78            for (k, &glob) in cells[slab_start..slab_end].iter().enumerate() {
79                let out_col = slab_start + k;
80                for source in self.col_to_data[glob].iter() {
81                    backend_groups
82                        .entry(source.backend as usize)
83                        .or_default()
84                        .push((source.local_col as usize, out_col));
85                }
86            }
87
88            for (&didx, group) in &backend_groups {
89                let local_cols: Vec<usize> = group.iter().map(|&(loc, _)| loc).collect();
90                let (_, _, group_triplets) =
91                    self.data_vec[didx].read_triplets_by_columns(local_cols)?;
92
93                let l2g = self.data_local_to_global_row[didx].as_slice();
94                if self.data_has_intra_row_merges[didx] {
95                    // Canonicalizer collapsed >=2 local rows in this dataset
96                    // to the same global. Sum into a HashMap so downstream
97                    // consumers don't see duplicate (row, col) entries. Every
98                    // entry for a given out_col lives in this slab, so
99                    // per-slab accumulation is exact.
100                    let mut acc: HashMap<(u64, u64), f32> = HashMap::default();
101                    for (i, j, v) in group_triplets {
102                        if let Some(c) = g2c[l2g[i as usize]] {
103                            let out_col = group[j as usize].1 as u64;
104                            *acc.entry((c as u64, out_col)).or_insert(0.0) += v;
105                        }
106                    }
107                    for ((r, c), v) in acc {
108                        f(r, c, v);
109                    }
110                } else {
111                    for (i, j, v) in group_triplets {
112                        if let Some(c) = g2c[l2g[i as usize]] {
113                            let out_col = group[j as usize].1 as u64;
114                            f(c as u64, out_col, v);
115                        }
116                    }
117                }
118            }
119        }
120
121        Ok((nrow, ncol))
122    }
123
124    /// Collect all nonzeros of the selected `cells` into one triplet
125    /// vector. Thin wrapper over [`Self::for_each_triplet`] with a single
126    /// slab spanning every column (identical one-pass behavior); prefer
127    /// `for_each_triplet` when the result is consumed once, to avoid the
128    /// full-width intermediate.
129    #[allow(clippy::type_complexity)]
130    pub fn columns_triplets<I>(
131        &self,
132        cells: I,
133    ) -> anyhow::Result<((usize, usize), Vec<(u64, u64, f32)>)>
134    where
135        I: Iterator<Item = usize>,
136    {
137        let cells: Vec<usize> = cells.collect();
138        let one_slab = cells.len().max(1);
139        let mut triplets = Vec::new();
140        let dims = self.for_each_triplet(cells.into_iter(), one_slab, |r, c, v| {
141            triplets.push((r, c, v));
142        })?;
143        Ok((dims, triplets))
144    }
145
146    #[cfg(feature = "ndarray")]
147    pub fn read_columns_ndarray<I>(&self, cells: I) -> anyhow::Result<ndarray::Array2<f32>>
148    where
149        I: Iterator<Item = usize>,
150    {
151        let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
152        ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
153    }
154
155    pub fn read_columns_dmatrix<I>(&self, cells: I) -> anyhow::Result<nalgebra::DMatrix<f32>>
156    where
157        I: Iterator<Item = usize>,
158    {
159        let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
160        DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
161    }
162
163    /// Direct-slice CSC read: bypasses the `triplet → COO → CSC` roundtrip
164    /// when the underlying backends have preloaded column arrays.
165    ///
166    /// For each cell, we slice `(indices, values)` straight out of the
167    /// backend's preloaded `by_column_indices` / `by_column_data`, remap
168    /// row indices through `l2g` then `g2c` once, drop entries that fall
169    /// outside the row intersection, and assemble final CSC arrays in one
170    /// pass per column. Backends that aren't preloaded fall back to
171    /// per-column triplet reads, which still avoids the global triplet
172    /// vec and the column-major sort inside `CscMatrix::from(&coo)`.
173    pub fn read_columns_csc<I>(&self, cells: I) -> anyhow::Result<CscMatrix<f32>>
174    where
175        I: Iterator<Item = usize>,
176    {
177        let cells: Vec<usize> = cells.collect();
178        let nrow = self.num_rows();
179        let ncol = cells.len();
180
181        // Group cells by backend, carrying the output column index so we
182        // can scatter directly into the final per-column buckets. Under
183        // `ColumnAlignment::Union` one global cell can be observed by
184        // multiple backends — the loop visits every `(didx, loc)` for
185        // the same `out_col`, so the bucket accumulates contributions
186        // from all observing backends.
187        let mut backend_groups: Vec<Vec<(usize, usize)>> =
188            (0..self.data_vec.len()).map(|_| Vec::new()).collect();
189        for (out_col, &glob) in cells.iter().enumerate() {
190            for source in self.col_to_data[glob].iter() {
191                backend_groups[source.backend as usize].push((source.local_col as usize, out_col));
192            }
193        }
194
195        // Per-output-column buckets of (compact_row, value).
196        let mut buckets: Vec<Vec<(u32, f32)>> = (0..ncol).map(|_| Vec::new()).collect();
197        let g2c = self.global_to_compact_row.as_slice();
198
199        for (didx, group) in backend_groups.iter().enumerate() {
200            if group.is_empty() {
201                continue;
202            }
203            let l2g = self.data_local_to_global_row[didx].as_slice();
204
205            if let Some((indptr, indices, values)) = self.data_vec[didx].csc_column_arrays() {
206                // Fast path: zero-copy slicing into preloaded arrays.
207                for &(loc, out_col) in group {
208                    if loc + 1 >= indptr.len() {
209                        continue;
210                    }
211                    let s = indptr[loc] as usize;
212                    let e = indptr[loc + 1] as usize;
213                    let bucket = &mut buckets[out_col];
214                    bucket.reserve(e - s);
215                    for k in s..e {
216                        let local_row = indices[k] as usize;
217                        if let Some(c) = g2c[l2g[local_row]] {
218                            bucket.push((c as u32, values[k]));
219                        }
220                    }
221                }
222            } else {
223                // Cold path: route through `read_triplets_by_columns` so the
224                // backend can coalesce abutting indptr ranges into a single
225                // zarr/hdf5 retrieval (zarr uses `coalesce_and_emit` + chunk
226                // LRU cache). For a contiguous block of N cells this becomes
227                // ONE retrieval instead of N — the dominant win on cold reads
228                // from `.zarr.zip` over a slow disk.
229                let cols: Vec<usize> = group.iter().map(|&(loc, _)| loc).collect();
230                let (_, _, trip) = self.data_vec[didx].read_triplets_by_columns(cols)?;
231                for (i, jj, v) in trip {
232                    let (_, out_col) = group[jj as usize];
233                    if let Some(c) = g2c[l2g[i as usize]] {
234                        buckets[out_col].push((c as u32, v));
235                    }
236                }
237            }
238        }
239
240        // Assemble final CSC arrays in one pass.
241        let total_nnz: usize = buckets.iter().map(|b| b.len()).sum();
242        let mut col_offsets: Vec<usize> = Vec::with_capacity(ncol + 1);
243        let mut row_indices: Vec<usize> = Vec::with_capacity(total_nnz);
244        let mut values: Vec<f32> = Vec::with_capacity(total_nnz);
245        col_offsets.push(0);
246
247        for bucket in &mut buckets {
248            // Canonical CSC requires within-column row indices sorted
249            // ascending AND unique. On-disk indices are sorted by local
250            // row, but two sources can land on the same compact row:
251            //   (a) a row canonicalizer that maps two local rows in the
252            //       same backend to one global row;
253            //   (b) `ColumnAlignment::Union` where multiple backends
254            //       contribute to one output column.
255            // The composition `l2g[g2c[..]]` is monotonic in the
256            // single-source, no-canonicalizer case (the historical
257            // fast path) — detect it cheaply and skip the rebuild;
258            // otherwise sort and fold duplicates by summing.
259            let strictly_sorted_unique = bucket.windows(2).all(|w| w[0].0 < w[1].0);
260            if !strictly_sorted_unique {
261                bucket.sort_by_key(|&(r, _)| r);
262                // Compact duplicates in-place: sum values for equal rows.
263                let mut write = 0usize;
264                let mut read = 0usize;
265                while read < bucket.len() {
266                    let (r, mut v) = bucket[read];
267                    read += 1;
268                    while read < bucket.len() && bucket[read].0 == r {
269                        v += bucket[read].1;
270                        read += 1;
271                    }
272                    bucket[write] = (r, v);
273                    write += 1;
274                }
275                bucket.truncate(write);
276            }
277            for &(r, v) in bucket.iter() {
278                row_indices.push(r as usize);
279                values.push(v);
280            }
281            col_offsets.push(row_indices.len());
282        }
283
284        CscMatrix::try_from_csc_data(nrow, ncol, col_offsets, row_indices, values)
285            .map_err(|e| anyhow::anyhow!("CSC construction failed: {:?}", e))
286    }
287
288    pub fn read_columns_csr<I>(&self, cells: I) -> anyhow::Result<CsrMatrix<f32>>
289    where
290        I: Iterator<Item = usize>,
291    {
292        let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
293        nalgebra_sparse::CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
294    }
295
296    #[cfg(feature = "tensor")]
297    pub fn read_columns_tensor<I>(&self, cells: I) -> anyhow::Result<Tensor>
298    where
299        I: Iterator<Item = usize>,
300    {
301        let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
302        Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
303    }
304
305    /// Build (shape, triplets) for the requested compact rows across
306    /// all backends. Output column index is the SparseIoVec-global
307    /// column (concatenation of backends in push order); output row
308    /// index is the position in `rows`.
309    #[allow(clippy::type_complexity)]
310    pub fn rows_triplets<I>(
311        &self,
312        rows: I,
313    ) -> anyhow::Result<((usize, usize), Vec<(u64, u64, f32)>)>
314    where
315        I: Iterator<Item = usize>,
316    {
317        let rows_compact: Vec<usize> = rows.collect();
318        let nrow_out = rows_compact.len();
319        let ncol_out = self.num_columns();
320
321        let n_compact = self.cached_num_rows;
322        let compact_to_global = self.compact_to_global_row.as_slice();
323
324        let mut triplets: Vec<(u64, u64, f32)> = Vec::new();
325        let mut local_to_out: Vec<usize> = Vec::with_capacity(rows_compact.len());
326        for didx in 0..self.data_vec.len() {
327            let g2l = &self.data_global_to_local_row[didx];
328
329            local_to_out.clear();
330            let mut local_rows: Vec<usize> = Vec::with_capacity(rows_compact.len());
331            for (out_row, &c) in rows_compact.iter().enumerate() {
332                if c >= n_compact {
333                    continue;
334                }
335                let g = compact_to_global[c];
336                if let Some(&l) = g2l.get(&g) {
337                    local_rows.push(l);
338                    local_to_out.push(out_row);
339                }
340            }
341            if local_rows.is_empty() {
342                continue;
343            }
344
345            let (_, _, group_triplets) = self.data_vec[didx].read_triplets_by_rows(local_rows)?;
346            let cols_map = self
347                .data_to_cols
348                .get(&didx)
349                .ok_or_else(|| anyhow::anyhow!("missing data_to_cols entry for didx {}", didx))?;
350            triplets.reserve(group_triplets.len());
351            for (i, j, v) in group_triplets {
352                // `usize::MAX` marks a cell dropped by `mask_columns`.
353                let mapped = cols_map[j as usize];
354                if mapped == usize::MAX {
355                    continue;
356                }
357                let out_row = local_to_out[i as usize] as u64;
358                let out_col = mapped as u64;
359                triplets.push((out_row, out_col, v));
360            }
361        }
362
363        Ok(((nrow_out, ncol_out), triplets))
364    }
365
366    #[cfg(feature = "ndarray")]
367    pub fn read_rows_ndarray<I>(&self, rows: I) -> anyhow::Result<ndarray::Array2<f32>>
368    where
369        I: Iterator<Item = usize>,
370    {
371        let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
372        ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
373    }
374
375    pub fn read_rows_dmatrix<I>(&self, rows: I) -> anyhow::Result<nalgebra::DMatrix<f32>>
376    where
377        I: Iterator<Item = usize>,
378    {
379        let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
380        DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
381    }
382
383    pub fn read_rows_csc<I>(&self, rows: I) -> anyhow::Result<CscMatrix<f32>>
384    where
385        I: Iterator<Item = usize>,
386    {
387        let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
388        nalgebra_sparse::CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
389    }
390
391    pub fn read_rows_csr<I>(&self, rows: I) -> anyhow::Result<CsrMatrix<f32>>
392    where
393        I: Iterator<Item = usize>,
394    {
395        let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
396        nalgebra_sparse::CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
397    }
398
399    #[cfg(feature = "tensor")]
400    pub fn read_rows_tensor<I>(&self, rows: I) -> anyhow::Result<Tensor>
401    where
402        I: Iterator<Item = usize>,
403    {
404        let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
405        Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
406    }
407}