1#![allow(dead_code, unused_imports)]
2
3#[cfg(feature = "tensor")]
4pub use legume_numeric::candle_core::Tensor;
5pub use nalgebra::DMatrix;
6pub use nalgebra_sparse::{csc::CscMatrix, csr::CsrMatrix};
7#[cfg(feature = "ndarray")]
8pub use ndarray::prelude::*;
9
10pub const MAX_ROW_NAME_IDX: usize = 3;
11pub const MAX_COLUMN_NAME_IDX: usize = 10;
12pub const COLUMN_SEP: &str = "@";
13pub const ROW_SEP: &str = "_";
14
15use super::helpers::*;
16use super::meta::Metadata;
17
18use crate::sparse_data_visitors::styled_progress_bar;
19use clap::ValueEnum;
20use indicatif::ParallelProgressIterator;
21use legume_numeric::matrix::mtx_io::*;
22use legume_numeric::matrix::traits::*;
23use log::info;
24use rayon::prelude::*;
25use rustc_hash::FxHashMap as HashMap;
26use std::ops::Range;
27use std::sync::{Arc, Mutex};
28
29#[cfg(test)]
30mod tests;
31
32#[derive(ValueEnum, Clone, Debug, PartialEq)]
33#[clap(rename_all = "lowercase")]
34pub enum SparseIoBackend {
35 Zarr,
36 HDF5,
37}
38
39#[derive(Clone, Copy, Debug, PartialEq, Eq)]
43pub enum CsKey {
44 CscData,
45 CscIndices,
46 CscIndptr,
47 CsrData,
48 CsrIndices,
49 CsrIndptr,
50}
51
52const SLAB_NNZ: usize = 1 << 20;
56
57fn slab_end(
64 triplets: &[(u64, u64, f32)],
65 start: usize,
66 slab_nnz: usize,
67 n_major: usize,
68 major: impl Fn(&(u64, u64, f32)) -> u64,
69) -> (usize, u64) {
70 debug_assert!(slab_nnz > 0);
71 let nnz = triplets.len();
72 let mut end = (start + slab_nnz).min(nnz);
73 while end < nnz && major(&triplets[end]) == major(&triplets[end - 1]) {
74 end += 1;
75 }
76 let band_end = if end == nnz {
77 n_major as u64
78 } else {
79 major(&triplets[end])
80 };
81 (end, band_end)
82}
83
84pub trait SparseIo: Sync + Send {
85 type IndexIter: IntoIterator<Item = usize> + FromIterator<usize>;
86
87 #[cfg(feature = "ndarray")]
92 fn read_columns_ndarray(&self, columns: Self::IndexIter) -> anyhow::Result<Array2<f32>> {
96 let (nrow, ncol, triplets) = self.read_triplets_by_columns(columns)?;
97 Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
98 }
99
100 #[cfg(feature = "tensor")]
101 fn read_columns_tensor(&self, columns: Self::IndexIter) -> anyhow::Result<Tensor> {
105 let (nrow, ncol, triplets) = self.read_triplets_by_columns(columns)?;
106 Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
107 }
108
109 fn read_columns_dmatrix(&self, columns: Self::IndexIter) -> anyhow::Result<DMatrix<f32>> {
113 let (nrow, ncol, triplets) = self.read_triplets_by_columns(columns)?;
114 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
115 }
116
117 fn read_columns_csr(&self, columns: Self::IndexIter) -> anyhow::Result<CsrMatrix<f32>> {
121 let (nrow, ncol, triplets) = self.read_triplets_by_columns(columns)?;
122 CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
123 }
124
125 fn read_columns_csc(&self, columns: Self::IndexIter) -> anyhow::Result<CscMatrix<f32>> {
129 let (nrow, ncol, triplets) = self.read_triplets_by_columns(columns)?;
130 CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
131 }
132
133 fn csc_column_arrays(&self) -> Option<(&[u64], &[u64], &[f32])> {
139 None
140 }
141
142 #[cfg(feature = "ndarray")]
143 fn read_rows_ndarray(&self, rows: Self::IndexIter) -> anyhow::Result<Array2<f32>> {
147 let (nrow, ncol, triplets) = self.read_triplets_by_rows(rows)?;
148 Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
149 }
150
151 #[cfg(feature = "tensor")]
152 fn read_rows_tensor(&self, rows: Self::IndexIter) -> anyhow::Result<Tensor> {
156 let (nrow, ncol, triplets) = self.read_triplets_by_rows(rows)?;
157 Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
158 }
159
160 fn read_rows_dmatrix(&self, rows: Self::IndexIter) -> anyhow::Result<DMatrix<f32>> {
164 let (nrow, ncol, triplets) = self.read_triplets_by_rows(rows)?;
165 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
166 }
167
168 fn read_rows_csr(&self, rows: Self::IndexIter) -> anyhow::Result<CsrMatrix<f32>> {
172 let (nrow, ncol, triplets) = self.read_triplets_by_rows(rows)?;
173 CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
174 }
175
176 fn read_rows_csc(&self, rows: Self::IndexIter) -> anyhow::Result<CscMatrix<f32>> {
180 let (nrow, ncol, triplets) = self.read_triplets_by_rows(rows)?;
181 CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
182 }
183
184 fn import_mtx_file(&mut self, mtx_file: &str, index_by_row: bool) -> anyhow::Result<()> {
194 let (mut mtx_triplets, mtx_shape) = read_mtx_triplets(mtx_file)?;
195 info!("read mtx file: {}", mtx_file);
196 if mtx_triplets.is_empty() {
197 return Err(anyhow::anyhow!("No data in mtx file"));
198 }
199 self.record_mtx_shape(Some(mtx_shape))?;
200 info!("recording the column index");
201 self.record_triplets_by_col(&mut mtx_triplets)?;
202 if index_by_row {
203 info!("recording the row index");
204 self.record_triplets_by_row(&mut mtx_triplets)?;
205 }
206 Ok(())
207 }
208
209 fn import_dmatrix_by_row(&mut self, matrix: &DMatrix<f32>) -> anyhow::Result<()> {
216 let (nrow, ncol) = matrix.shape();
217 let mut mtx_triplets = dmatrix_to_triplets(matrix);
218 let mtx_shape = (nrow, ncol, mtx_triplets.len());
219 self.record_mtx_shape(Some(mtx_shape))?;
220 self.record_triplets_by_row(&mut mtx_triplets)
221 }
222
223 fn import_dmatrix_by_col(&mut self, matrix: &DMatrix<f32>) -> anyhow::Result<()> {
226 let (nrow, ncol) = matrix.shape();
227 let mut mtx_triplets = dmatrix_to_triplets(matrix);
228 let mtx_shape = (nrow, ncol, mtx_triplets.len());
229 self.record_mtx_shape(Some(mtx_shape))?;
230 self.record_triplets_by_col(&mut mtx_triplets)
231 }
232
233 #[cfg(feature = "ndarray")]
238 fn import_ndarray_by_row(&mut self, array: &Array2<f32>) -> anyhow::Result<()> {
241 let nrow = array.shape()[0];
242 let ncol = array.shape()[1];
243
244 let mut mtx_triplets = ndarray_to_triplets(array);
246
247 let nnz = mtx_triplets.len();
248 let mtx_shape = (nrow, ncol, nnz);
249 self.record_mtx_shape(Some(mtx_shape))?;
250
251 self.record_triplets_by_row(&mut mtx_triplets)
254 }
255
256 #[cfg(feature = "ndarray")]
257 fn import_ndarray_by_col(&mut self, array: &Array2<f32>) -> anyhow::Result<()> {
260 let nrow = array.shape()[0];
261 let ncol = array.shape()[1];
262
263 let mut mtx_triplets = ndarray_to_triplets(array);
265
266 let nnz = mtx_triplets.len();
267 let mtx_shape = (nrow, ncol, nnz);
268 self.record_mtx_shape(Some(mtx_shape))?;
269
270 self.record_triplets_by_col(&mut mtx_triplets)
273 }
274
275 #[allow(clippy::type_complexity)]
283 fn read_triplets_by_rows(
284 &self,
285 rows: Self::IndexIter,
286 ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)>;
287
288 #[allow(clippy::type_complexity)]
292 fn read_triplets_by_columns(
293 &self,
294 columns: Self::IndexIter,
295 ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)>;
296
297 #[allow(clippy::type_complexity)]
301 fn read_triplets_by_single_column(
302 &self,
303 col: usize,
304 ) -> anyhow::Result<(usize, usize, Vec<(u64, u64, f32)>)>;
305
306 fn to_mtx_file(&self, mtx_file: &str) -> anyhow::Result<()>;
309
310 fn num_rows(&self) -> Option<usize>;
312
313 fn num_columns(&self) -> Option<usize>;
315
316 fn num_non_zeros(&self) -> Option<usize>;
318
319 fn reopen_backend(&mut self) -> anyhow::Result<()>;
324
325 #[doc(hidden)]
329 fn note_streamed_nnz(&mut self, n: u64);
337
338 #[doc(hidden)]
340 fn streamed_nnz(&self) -> u64;
341
342 #[doc(hidden)]
344 fn reset_streamed_nnz(&mut self);
345
346 fn column_indptr(&self) -> &[u64];
351
352 fn column_nnz(&self, col: usize) -> Option<u64> {
358 let indptr = self.column_indptr();
359 let hi = *indptr.get(col + 1)?;
360 let lo = *indptr.get(col)?;
361 hi.checked_sub(lo)
362 }
363
364 fn register_row_names_file(&mut self, row_name_file: &str);
367
368 fn register_column_names_file(&mut self, column_name_file: &str);
371
372 fn register_row_names_vec(&mut self, rows: &[Box<str>]);
375
376 fn register_column_names_vec(&mut self, columns: &[Box<str>]);
379
380 fn register_names_file(
386 &mut self,
387 key: &str,
388 name_file: &str,
389 name_columns: Range<usize>,
390 name_sep: &str,
391 ) -> anyhow::Result<()>;
392
393 fn register_names_vec(&mut self, key: &str, names: &[Box<str>]) -> anyhow::Result<()>;
397
398 fn row_names(&self) -> anyhow::Result<Vec<Box<str>>>;
399
400 fn column_names(&self) -> anyhow::Result<Vec<Box<str>>>;
401
402 fn retrieve_registered_names(&self, key: &str) -> anyhow::Result<Vec<Box<str>>>;
405
406 fn metadata(&self) -> Metadata;
414
415 fn set_metadata(&mut self, meta: &Metadata) -> anyhow::Result<()>;
420
421 fn meta(&self, key: &str) -> Option<String> {
423 self.metadata().remove(key)
424 }
425
426 fn set_meta(&mut self, key: &str, value: &str) -> anyhow::Result<()> {
428 let mut meta = self.metadata();
429 meta.insert(key.to_string(), value.to_string());
430 self.set_metadata(&meta)
431 }
432
433 fn subset_columns_rows(
441 &mut self,
442 columns: Option<&Vec<usize>>,
443 rows: Option<&Vec<usize>>,
444 ) -> anyhow::Result<()> {
445 let ncol_data = self
446 .num_columns()
447 .ok_or_else(|| anyhow::anyhow!("missing shape information"))?;
448 let nrow_data = self
449 .num_rows()
450 .ok_or_else(|| anyhow::anyhow!("missing shape information"))?;
451
452 let distinct = |sel: &[usize], what: &str| -> anyhow::Result<()> {
461 anyhow::ensure!(!sel.is_empty(), "subset: empty {what} selection");
462 let mut seen = sel.to_vec();
463 seen.sort_unstable();
464 seen.dedup();
465 anyhow::ensure!(
466 seen.len() == sel.len(),
467 "subset: the {what} selection repeats an index ({} of {} are distinct)",
468 seen.len(),
469 sel.len()
470 );
471 Ok(())
472 };
473 if let Some(cols) = columns {
474 distinct(cols, "column")?;
475 }
476 if let Some(rs) = rows {
477 distinct(rs, "row")?;
478 }
479
480 let (old2new_cols, new_col_names) =
485 take_subset_indices_names_if_needed(columns, Some(ncol_data), self.column_names()?);
486 let (old2new_rows, new_row_names) =
487 take_subset_indices_names_if_needed(rows, Some(nrow_data), self.row_names()?);
488 let (new_ncol, new_nrow) = (new_col_names.len(), new_row_names.len());
489 anyhow::ensure!(new_ncol > 0, "subset: no column survived the selection");
490 anyhow::ensure!(new_nrow > 0, "subset: no row survived the selection");
491
492 let mut cols_new_order: Vec<(u64, u64)> =
494 old2new_cols.iter().map(|(&o, &n)| (n, o)).collect();
495 cols_new_order.sort_unstable();
496
497 let mut row_map: Vec<Option<u64>> = vec![None; nrow_data];
502 for (&old, &new) in &old2new_rows {
503 row_map[old as usize] = Some(new);
504 }
505 let monotone_rows = row_map.iter().flatten().is_sorted_by(|a, b| a < b);
506
507 let full_rows = rows.is_none();
515 let per_col_nnz: Vec<u64> = if full_rows {
516 cols_new_order
517 .iter()
518 .map(|&(_, old)| {
519 self.column_nnz(old as usize)
520 .ok_or_else(|| anyhow::anyhow!("subset: no indptr for column {old}"))
521 })
522 .collect::<anyhow::Result<_>>()?
523 } else {
524 let mut counts = vec![0u64; cols_new_order.len()];
529 let coarse = legume_numeric::matrix::utils::generate_minibatch_intervals(
530 cols_new_order.len(),
531 0,
532 Some(8192),
533 );
534 for (lb, ub) in coarse {
535 let old_cols: Vec<usize> = cols_new_order[lb..ub]
536 .iter()
537 .map(|&(_, o)| o as usize)
538 .collect();
539 let (_, _, triplets) =
540 self.read_triplets_by_columns(old_cols.into_iter().collect())?;
541 for (i, c_local, _) in triplets {
542 if row_map[i as usize].is_some() {
543 counts[lb + c_local as usize] += 1;
544 }
545 }
546 }
547 counts
548 };
549 let new_nnz: u64 = per_col_nnz.iter().sum();
550
551 let final_path = self.get_backend_file_name().to_string();
569 anyhow::ensure!(
570 !final_path.ends_with(".zip"),
571 "subset: {final_path} is a zip archive; convert it to a directory \
572 backend first (data-beans convert)"
573 );
574 let temp_path = format!("{final_path}.subset_tmp");
575 if std::path::Path::new(&temp_path).exists() {
576 crate::sparse_io::remove_backend_path(&temp_path)?;
577 }
578
579 {
580 let backend_kind = self.backend_type();
581 let mut out = crate::sparse_io::create_sparse_streaming_empty(
582 Some(&temp_path),
583 Some(&backend_kind),
584 )?;
585 out.begin_streaming_csc((new_nrow, new_ncol, new_nnz as usize))?;
586
587 let blocks = legume_numeric::matrix::utils::byte_budget_intervals(
590 &per_col_nnz,
591 crate::sparse_io::SLAB_BUDGET_BYTES,
592 crate::sparse_io::TRIPLET_BYTES,
593 );
594
595 let mut nnz_offset = 0u64;
596 for (lb, ub) in blocks {
597 let old_cols: Vec<usize> = cols_new_order[lb..ub]
601 .iter()
602 .map(|&(_, o)| o as usize)
603 .collect();
604 let (_, _, triplets) =
605 self.read_triplets_by_columns(old_cols.into_iter().collect())?;
606
607 let n_block = ub - lb;
608 let mut per_col: Vec<Vec<(u64, f32)>> = vec![Vec::new(); n_block];
609 for (i, c_local, x) in triplets {
610 if let Some(new_row) = row_map[i as usize] {
611 per_col[c_local as usize].push((new_row, x));
612 }
613 }
614 let mut local_colptr = Vec::with_capacity(n_block);
615 let mut row_indices = Vec::new();
616 let mut values = Vec::new();
617 for entries in &mut per_col {
618 if !monotone_rows {
619 entries.sort_unstable_by_key(|&(r, _)| r);
622 }
623 local_colptr.push(row_indices.len() as u64);
624 for &(r, x) in entries.iter() {
625 row_indices.push(r);
626 values.push(x);
627 }
628 }
629 out.append_csc_slab(lb as u64, nnz_offset, &local_colptr, &row_indices, &values)?;
630 nnz_offset += values.len() as u64;
631 }
632
633 out.finalize_streaming_csc()?;
634 out.build_csr_from_csc_streaming()?;
635 out.register_row_names_vec(&new_row_names);
636 out.register_column_names_vec(&new_col_names);
637 out.set_metadata(&self.metadata())?;
638 }
639
640 self.remove_backend_file()?;
645 std::fs::rename(&temp_path, &final_path)?;
646 self.reopen_backend()?;
647 self.clean_preloaded_columns();
648 self.clean_preloaded_rows();
649 info!("registered new data to {}", self.get_backend_file_name());
650 Ok(())
651 }
652
653 fn reorder_rows(&mut self, row_names_order: &[Box<str>]) -> anyhow::Result<()> {
656 let meta = self.metadata();
658 let new_col_names = self.column_names()?.clone();
659 let name2new = build_name2index_map(row_names_order);
660
661 let block_size = 100;
662
663 let old2new: HashMap<u64, u64> = self
664 .row_names()?
665 .into_par_iter()
666 .enumerate()
667 .filter_map(|(idx_old, name)| {
668 name2new
669 .get(&name)
670 .map(|&idx_new| (idx_old as u64, idx_new as u64))
671 })
672 .collect();
673
674 if let Some(ncol) = self.num_columns() {
675 let arc_triplets = Arc::new(Mutex::new(vec![]));
680
681 let nblock = ncol.div_ceil(block_size);
682
683 info!("remapping triplets ...");
684
685 (0..nblock)
686 .into_par_iter()
687 .progress_with(styled_progress_bar(nblock as u64, "blocks"))
688 .map(|b| {
689 let lb = (b * block_size) as u64;
690 let ub = ((b + 1) * block_size).min(ncol) as u64;
691 (lb, ub)
692 })
693 .for_each(|(lb, ub)| {
694 let (_, _, _triplets_b) = self
695 .read_triplets_by_columns(((lb as usize)..(ub as usize)).collect())
696 .unwrap();
697
698 let _triplets_b = _triplets_b.into_iter().filter_map(|(i, j_loc, x)| {
699 let j_glob = j_loc + lb;
700 old2new.get(&i).map(|&i_new| (i_new, j_glob, x))
701 });
702
703 {
704 let mut triplets = arc_triplets.lock().unwrap();
705 triplets.extend(_triplets_b);
706 }
707 });
708
709 self.remove_backend_file()?;
713
714 self.initialize_backend()?;
718
719 {
721 let mut row_col_val_triplets =
722 arc_triplets.lock().expect("failed to lock triplets");
723
724 let nnz = row_col_val_triplets.len();
725 debug_assert!(row_col_val_triplets.len() <= nnz); let new_nrow = row_names_order.len();
727 let mtx_shape = (new_nrow, ncol, nnz);
728
729 info!("sorting triplets ...");
730
731 self.record_mtx_shape(Some(mtx_shape))?;
732 self.record_triplets_by_col(&mut row_col_val_triplets)?;
733 self.record_triplets_by_row(&mut row_col_val_triplets)?;
734 }
735 self.read_column_indptr()?;
736 self.read_row_indptr()?;
737
738 self.register_row_names_vec(row_names_order);
739 self.register_column_names_vec(&new_col_names);
740 self.set_metadata(&meta)?;
741 info!("registered new data to {}", self.get_backend_file_name());
742 }
743
744 self.clean_preloaded_columns();
745 self.clean_preloaded_rows();
746 Ok(())
747 }
748 fn remove_backend_file(&self) -> anyhow::Result<()>;
752
753 fn initialize_backend(&mut self) -> anyhow::Result<()>;
755
756 fn record_mtx_shape(&mut self, mtx_shape: Option<(usize, usize, usize)>) -> anyhow::Result<()>;
757
758 fn record_triplets_by_row(
761 &mut self,
762 row_col_val_triplets: &mut Vec<(u64, u64, f32)>,
763 ) -> anyhow::Result<()> {
764 let nrow = self.num_rows().expect("should have `nrow`");
765 let ncol = self.num_columns().expect("should have `ncol`");
766 let nnz = row_col_val_triplets.len();
767
768 if nnz == 0 {
769 let csr_rowptr = vec![0u64; nrow + 1];
770 return self.record_csr_dataset_backend(&[], &[], &csr_rowptr);
771 }
772
773 row_col_val_triplets.par_sort_unstable_by_key(|&(row, col, _)| (row, col));
777
778 self.begin_streaming_csr((nrow, ncol, nnz))?;
779
780 let mut local_rowptr: Vec<u64> = Vec::new();
781 let mut cols: Vec<u64> = Vec::with_capacity(SLAB_NNZ);
782 let mut vals: Vec<f32> = Vec::with_capacity(SLAB_NNZ);
783
784 let mut start = 0_usize;
785 let mut row_offset = 0_u64;
786 while (row_offset as usize) < nrow {
787 let (end, band_end_row) =
788 slab_end(row_col_val_triplets, start, SLAB_NNZ, nrow, |t| t.0);
789
790 local_rowptr.clear();
791 cols.clear();
792 vals.clear();
793 let mut i = start;
794 for row in row_offset..band_end_row {
795 local_rowptr.push((i - start) as u64);
796 while i < end && row_col_val_triplets[i].0 == row {
797 cols.push(row_col_val_triplets[i].1);
798 vals.push(row_col_val_triplets[i].2);
799 i += 1;
800 }
801 }
802 debug_assert_eq!(i, end, "every entry of the band belongs to one of its rows");
803
804 self.append_csr_slab(row_offset, start as u64, &local_rowptr, &cols, &vals)?;
805 start = end;
806 row_offset = band_end_row;
807 }
808
809 self.finalize_streaming_csr()
810 }
811
812 fn record_triplets_by_col(
819 &mut self,
820 row_col_val_triplets: &mut Vec<(u64, u64, f32)>,
821 ) -> anyhow::Result<()> {
822 let nrow = self.num_rows().expect("should have `nrow`");
823 let ncol = self.num_columns().expect("should have `ncol`");
824 let nnz = row_col_val_triplets.len();
825
826 if nnz == 0 {
827 let csc_colptr = vec![0u64; ncol + 1];
828 return self.record_csc_dataset_backend(&[], &[], &csc_colptr);
829 }
830
831 row_col_val_triplets.par_sort_unstable_by_key(|&(row, col, _)| (col, row));
833
834 self.begin_streaming_csc((nrow, ncol, nnz))?;
835
836 let mut local_colptr: Vec<u64> = Vec::new();
837 let mut rows: Vec<u64> = Vec::with_capacity(SLAB_NNZ);
838 let mut vals: Vec<f32> = Vec::with_capacity(SLAB_NNZ);
839
840 let mut start = 0_usize;
841 let mut col_offset = 0_u64;
842 while (col_offset as usize) < ncol {
843 let (end, band_end_col) =
844 slab_end(row_col_val_triplets, start, SLAB_NNZ, ncol, |t| t.1);
845
846 local_colptr.clear();
847 rows.clear();
848 vals.clear();
849 let mut i = start;
850 for col in col_offset..band_end_col {
851 local_colptr.push((i - start) as u64);
852 while i < end && row_col_val_triplets[i].1 == col {
853 rows.push(row_col_val_triplets[i].0);
854 vals.push(row_col_val_triplets[i].2);
855 i += 1;
856 }
857 }
858 debug_assert_eq!(
859 i, end,
860 "every entry of the band belongs to one of its columns"
861 );
862
863 self.append_csc_slab(col_offset, start as u64, &local_colptr, &rows, &vals)?;
864 start = end;
865 col_offset = band_end_col;
866 }
867
868 self.finalize_streaming_csc()
869 }
870
871 fn record_csr_dataset_backend(
880 &mut self,
881 csr_cols: &[u64],
882 csr_vals: &[f32],
883 csr_rowptr: &[u64],
884 ) -> anyhow::Result<()>;
885
886 fn record_csc_dataset_backend(
896 &mut self,
897 csc_rows: &[u64],
898 csc_vals: &[f32],
899 csc_colptr: &[u64],
900 ) -> anyhow::Result<()>;
901
902 fn cs_create(&mut self, key: CsKey, len: usize) -> anyhow::Result<()>;
905
906 fn cs_write_u64(&mut self, key: CsKey, offset: u64, data: &[u64]) -> anyhow::Result<()>;
909
910 fn cs_write_f32(&mut self, key: CsKey, offset: u64, data: &[f32]) -> anyhow::Result<()>;
913
914 fn begin_streaming_csc(&mut self, shape: (usize, usize, usize)) -> anyhow::Result<()> {
919 self.reset_streamed_nnz();
923 let (_, ncol, nnz) = shape;
924 self.record_mtx_shape(Some(shape))?;
925 self.cs_create(CsKey::CscData, nnz)?;
926 self.cs_create(CsKey::CscIndices, nnz)?;
927 self.cs_create(CsKey::CscIndptr, ncol + 1)?;
928 Ok(())
929 }
930
931 fn append_csc_slab(
940 &mut self,
941 col_offset: u64,
942 nnz_offset: u64,
943 local_colptr: &[u64],
944 row_indices: &[u64],
945 values: &[f32],
946 ) -> anyhow::Result<()> {
947 anyhow::ensure!(
953 row_indices.len() == values.len(),
954 "append_csc_slab: {} row indices vs {} values",
955 row_indices.len(),
956 values.len()
957 );
958 anyhow::ensure!(
959 local_colptr.first().copied() == Some(0) || local_colptr.is_empty(),
960 "append_csc_slab: local_colptr must start at 0"
961 );
962 anyhow::ensure!(
963 local_colptr.windows(2).all(|w| w[0] <= w[1]),
964 "append_csc_slab: local_colptr must be monotone non-decreasing"
965 );
966 if let Some(&last) = local_colptr.last() {
967 anyhow::ensure!(
968 last <= values.len() as u64,
969 "append_csc_slab: colptr claims {last} entries, slab holds {}",
970 values.len()
971 );
972 }
973 if let Some(nrow) = self.num_rows() {
974 if let Some(&bad) = row_indices.iter().find(|&&r| r >= nrow as u64) {
975 anyhow::bail!("append_csc_slab: row index {bad} outside the {nrow}-row matrix");
976 }
977 }
978 for (c, &start) in local_colptr.iter().enumerate() {
981 let end = local_colptr
982 .get(c + 1)
983 .copied()
984 .unwrap_or(values.len() as u64) as usize;
985 anyhow::ensure!(
986 row_indices[start as usize..end]
987 .windows(2)
988 .all(|w| w[0] < w[1]),
989 "append_csc_slab: rows within column {} of this band must be \
990 strictly ascending — repeated rows usually mean duplicate \
991 (row, col) coordinates in the source (an MTX with repeated \
992 entries, or a union remap folding rows together)",
993 col_offset as usize + c
994 );
995 }
996
997 let shifted: Vec<u64> = local_colptr.iter().map(|&p| p + nnz_offset).collect();
998 self.cs_write_u64(CsKey::CscIndptr, col_offset, &shifted)?;
999 self.cs_write_u64(CsKey::CscIndices, nnz_offset, row_indices)?;
1000 self.cs_write_f32(CsKey::CscData, nnz_offset, values)?;
1001 self.note_streamed_nnz(values.len() as u64);
1002 Ok(())
1003 }
1004
1005 fn finalize_streaming_csc(&mut self) -> anyhow::Result<()> {
1008 let ncol = self
1009 .num_columns()
1010 .ok_or_else(|| anyhow::anyhow!("ncol not set before finalize_streaming_csc"))?;
1011 let nnz = self
1012 .num_non_zeros()
1013 .ok_or_else(|| anyhow::anyhow!("nnz not set before finalize_streaming_csc"))?;
1014 self.cs_write_u64(CsKey::CscIndptr, ncol as u64, &[nnz as u64])?;
1015 self.read_column_indptr()?;
1016
1017 let indptr = self.column_indptr();
1026 anyhow::ensure!(
1027 indptr.len() == ncol + 1,
1028 "finalize_streaming_csc: indptr has {} entries, expected {}",
1029 indptr.len(),
1030 ncol + 1
1031 );
1032 anyhow::ensure!(
1033 indptr.first().copied() == Some(0),
1034 "finalize_streaming_csc: indptr[0] = {:?}, expected 0 — the first \
1035 slab was never appended",
1036 indptr.first()
1037 );
1038 if let Some(w) = indptr.windows(2).position(|w| w[0] > w[1]) {
1039 anyhow::bail!(
1040 "finalize_streaming_csc: indptr decreases at column {w} — slabs \
1041 were appended with a gap or overlap in their nnz offsets"
1042 );
1043 }
1044 let appended = self.streamed_nnz();
1050 anyhow::ensure!(
1051 appended == nnz as u64,
1052 "finalize_streaming_csc: {appended} entries appended but {nnz} \
1053 declared — the difference reads back as fill values wearing real \
1054 entries' positions"
1055 );
1056 Ok(())
1057 }
1058
1059 fn begin_streaming_csr(&mut self, shape: (usize, usize, usize)) -> anyhow::Result<()> {
1062 self.reset_streamed_nnz();
1063 let (nrow, _, nnz) = shape;
1064 self.record_mtx_shape(Some(shape))?;
1065 self.cs_create(CsKey::CsrData, nnz)?;
1066 self.cs_create(CsKey::CsrIndices, nnz)?;
1067 self.cs_create(CsKey::CsrIndptr, nrow + 1)?;
1068 Ok(())
1069 }
1070
1071 fn append_csr_slab(
1081 &mut self,
1082 row_offset: u64,
1083 nnz_offset: u64,
1084 local_rowptr: &[u64],
1085 col_indices: &[u64],
1086 values: &[f32],
1087 ) -> anyhow::Result<()> {
1088 anyhow::ensure!(
1089 col_indices.len() == values.len(),
1090 "append_csr_slab: {} column indices vs {} values",
1091 col_indices.len(),
1092 values.len()
1093 );
1094 anyhow::ensure!(
1095 local_rowptr.first().copied() == Some(0) || local_rowptr.is_empty(),
1096 "append_csr_slab: local_rowptr must start at 0"
1097 );
1098 anyhow::ensure!(
1099 local_rowptr.windows(2).all(|w| w[0] <= w[1]),
1100 "append_csr_slab: local_rowptr must be monotone non-decreasing"
1101 );
1102 if let Some(&last) = local_rowptr.last() {
1103 anyhow::ensure!(
1104 last <= values.len() as u64,
1105 "append_csr_slab: rowptr claims {last} entries, slab holds {}",
1106 values.len()
1107 );
1108 }
1109 if let Some(ncol) = self.num_columns() {
1110 if let Some(&bad) = col_indices.iter().find(|&&c| c >= ncol as u64) {
1111 anyhow::bail!(
1112 "append_csr_slab: column index {bad} outside the {ncol}-column matrix"
1113 );
1114 }
1115 }
1116 for (r, &start) in local_rowptr.iter().enumerate() {
1117 let end = local_rowptr
1118 .get(r + 1)
1119 .copied()
1120 .unwrap_or(values.len() as u64) as usize;
1121 anyhow::ensure!(
1122 col_indices[start as usize..end]
1123 .windows(2)
1124 .all(|w| w[0] < w[1]),
1125 "append_csr_slab: columns within row {} of this band must be \
1126 strictly ascending — repeated columns usually mean duplicate \
1127 (row, col) coordinates in the source",
1128 row_offset as usize + r
1129 );
1130 }
1131
1132 let shifted: Vec<u64> = local_rowptr.iter().map(|&p| p + nnz_offset).collect();
1133 self.cs_write_u64(CsKey::CsrIndptr, row_offset, &shifted)?;
1134 self.cs_write_u64(CsKey::CsrIndices, nnz_offset, col_indices)?;
1135 self.cs_write_f32(CsKey::CsrData, nnz_offset, values)?;
1136 self.note_streamed_nnz(values.len() as u64);
1137 Ok(())
1138 }
1139
1140 fn finalize_streaming_csr(&mut self) -> anyhow::Result<()> {
1144 let nrow = self
1145 .num_rows()
1146 .ok_or_else(|| anyhow::anyhow!("nrow not set before finalize_streaming_csr"))?;
1147 let nnz = self
1148 .num_non_zeros()
1149 .ok_or_else(|| anyhow::anyhow!("nnz not set before finalize_streaming_csr"))?;
1150 self.cs_write_u64(CsKey::CsrIndptr, nrow as u64, &[nnz as u64])?;
1151 self.read_row_indptr()?;
1152
1153 let appended = self.streamed_nnz();
1154 anyhow::ensure!(
1155 appended == nnz as u64,
1156 "finalize_streaming_csr: {appended} entries appended but {nnz} \
1157 declared — the slabs did not cover the matrix"
1158 );
1159 Ok(())
1160 }
1161
1162 fn build_csr_from_csc_streaming(&mut self) -> anyhow::Result<()> {
1166 let nrow = self
1167 .num_rows()
1168 .ok_or_else(|| anyhow::anyhow!("nrow not set before build_csr_from_csc_streaming"))?;
1169 let ncol = self
1170 .num_columns()
1171 .ok_or_else(|| anyhow::anyhow!("ncol not set before build_csr_from_csc_streaming"))?;
1172 let nnz = self
1173 .num_non_zeros()
1174 .ok_or_else(|| anyhow::anyhow!("nnz not set before build_csr_from_csc_streaming"))?;
1175
1176 if nnz == 0 {
1177 self.cs_create(CsKey::CsrData, 0)?;
1178 self.cs_create(CsKey::CsrIndices, 0)?;
1179 self.cs_create(CsKey::CsrIndptr, nrow + 1)?;
1180 let zeros = vec![0u64; nrow + 1];
1181 self.cs_write_u64(CsKey::CsrIndptr, 0, &zeros)?;
1182 self.read_row_indptr()?;
1183 return Ok(());
1184 }
1185
1186 const COL_BLOCK: usize = 1024;
1187 let n_col_blocks = ncol.div_ceil(COL_BLOCK);
1188 let bar1 = styled_progress_bar(n_col_blocks as u64, "transpose count");
1189 let mut row_counts = vec![0u64; nrow];
1190 let mut col_lo = 0usize;
1191 while col_lo < ncol {
1192 let col_hi = (col_lo + COL_BLOCK).min(ncol);
1193 let cols: Self::IndexIter = (col_lo..col_hi).collect();
1194 let (_, _, triplets) = self.read_triplets_by_columns(cols)?;
1195 for (row_i, _, _) in &triplets {
1196 row_counts[*row_i as usize] += 1;
1197 }
1198 col_lo = col_hi;
1199 bar1.inc(1);
1200 }
1201 bar1.finish_and_clear();
1202
1203 let mut rowptr = vec![0u64; nrow + 1];
1204 let mut acc = 0u64;
1205 for i in 0..nrow {
1206 rowptr[i] = acc;
1207 acc += row_counts[i];
1208 }
1209 rowptr[nrow] = acc;
1210 debug_assert_eq!(acc, nnz as u64);
1211
1212 self.cs_create(CsKey::CsrData, nnz)?;
1213 self.cs_create(CsKey::CsrIndices, nnz)?;
1214 self.cs_create(CsKey::CsrIndptr, nrow + 1)?;
1215 self.cs_write_u64(CsKey::CsrIndptr, 0, &rowptr)?;
1216
1217 const TRANSPOSE_BAND_BYTES: usize = 256 * 1024 * 1024;
1221 let avg_density = nnz.div_ceil(nrow.max(1));
1222 let band_rows = (TRANSPOSE_BAND_BYTES / (12 * avg_density.max(1)))
1223 .max(1)
1224 .min(nrow);
1225 let n_bands = nrow.div_ceil(band_rows);
1226
1227 let bar2 = styled_progress_bar(n_bands as u64, "transpose scatter");
1228 let mut band_lo = 0usize;
1229 while band_lo < nrow {
1230 let band_hi = (band_lo + band_rows).min(nrow);
1231 let band_nnz_start = rowptr[band_lo];
1232 let band_nnz_end = rowptr[band_hi];
1233 let band_nnz = (band_nnz_end - band_nnz_start) as usize;
1234
1235 if band_nnz == 0 {
1236 band_lo = band_hi;
1237 bar2.inc(1);
1238 continue;
1239 }
1240
1241 let mut out_indices = vec![0u64; band_nnz];
1242 let mut out_values = vec![0f32; band_nnz];
1243 let mut cursor = vec![0u64; band_hi - band_lo];
1244
1245 let mut col_lo = 0usize;
1246 while col_lo < ncol {
1247 let col_hi = (col_lo + COL_BLOCK).min(ncol);
1248 let cols: Self::IndexIter = (col_lo..col_hi).collect();
1249 let (_, _, triplets) = self.read_triplets_by_columns(cols)?;
1250 for &(row_i, col_j_local, x) in &triplets {
1251 let row_i_us = row_i as usize;
1252 if row_i_us >= band_lo && row_i_us < band_hi {
1253 let band_idx = row_i_us - band_lo;
1254 let col_j_global = col_j_local + col_lo as u64;
1259 let offset_in_band =
1260 (rowptr[band_lo + band_idx] - band_nnz_start) + cursor[band_idx];
1261 out_indices[offset_in_band as usize] = col_j_global;
1262 out_values[offset_in_band as usize] = x;
1263 cursor[band_idx] += 1;
1264 }
1265 }
1266 col_lo = col_hi;
1267 }
1268
1269 self.cs_write_u64(CsKey::CsrIndices, band_nnz_start, &out_indices)?;
1270 self.cs_write_f32(CsKey::CsrData, band_nnz_start, &out_values)?;
1271
1272 band_lo = band_hi;
1273 bar2.inc(1);
1274 }
1275 bar2.finish_and_clear();
1276
1277 self.read_row_indptr()?;
1278 Ok(())
1279 }
1280
1281 fn read_row_indptr(&mut self) -> anyhow::Result<()>;
1283
1284 fn read_column_indptr(&mut self) -> anyhow::Result<()>;
1286
1287 fn preload_columns(&mut self) -> anyhow::Result<()>;
1289
1290 fn clean_preloaded_columns(&mut self);
1292
1293 fn preload_rows(&mut self) -> anyhow::Result<()>;
1295
1296 fn clean_preloaded_rows(&mut self);
1298
1299 fn get_backend_file_name(&self) -> &str;
1301
1302 fn backend_type(&self) -> SparseIoBackend;
1304}