1#![allow(dead_code)]
2
3use super::*;
4use crate::sparse_io::index_audit;
5
6#[inline]
10fn remap_row(
11 l2g: &[usize],
12 g2c: &[Option<usize>],
13 row: u64,
14 didx: usize,
15 col: usize,
16) -> anyhow::Result<Option<usize>> {
17 let Some(&g) = l2g.get(row as usize) else {
18 anyhow::bail!(
19 "backend {didx}, column {col}: row index {row} (0x{row:016x}) is outside \
20 the backend's {} rows",
21 l2g.len()
22 );
23 };
24 g2c.get(g).copied().ok_or_else(|| {
25 anyhow::anyhow!(
26 "backend {didx}, column {col}: row index {row} maps to global row {g}, \
27 outside the global row table"
28 )
29 })
30}
31
32impl SparseIoVec {
33 pub(super) fn read_column_offset(
45 &self,
46 glob: usize,
47 col_offset: &mut usize,
48 triplets: &mut Vec<(u64, u64, f32)>,
49 ) -> anyhow::Result<()> {
50 let off = *col_offset as u64;
51 let g2c = self.global_to_compact_row.as_slice();
52 for source in self.col_to_data[glob].iter() {
53 let didx = source.backend as usize;
54 let loc = source.local_col as usize;
55 let (_, _, loc_triplets) = self.data_vec[didx].read_triplets_by_single_column(loc)?;
56 let l2g = self.data_local_to_global_row[didx].as_slice();
57 for (i, _j, v) in loc_triplets {
58 if let Some(c) = g2c[l2g[i as usize]] {
59 triplets.push((c as u64, off, v));
60 }
61 }
62 }
63 *col_offset += 1;
64 Ok(())
65 }
66
67 pub fn for_each_triplet<I, F>(
80 &self,
81 cells: I,
82 chunk_cols: usize,
83 mut f: F,
84 ) -> anyhow::Result<(usize, usize)>
85 where
86 I: Iterator<Item = usize>,
87 F: FnMut(u64, u64, f32),
88 {
89 let nrow = self.num_rows();
90 let cells: Vec<usize> = cells.collect();
91 let ncol = cells.len();
92 let chunk = chunk_cols.max(1);
93 let g2c = self.global_to_compact_row.as_slice();
94
95 for slab_start in (0..ncol).step_by(chunk) {
96 let slab_end = (slab_start + chunk).min(ncol);
97
98 let mut backend_groups: HashMap<usize, Vec<(usize, usize)>> = HashMap::default();
105 for (k, &glob) in cells[slab_start..slab_end].iter().enumerate() {
106 let out_col = slab_start + k;
107 for source in self.col_to_data[glob].iter() {
108 backend_groups
109 .entry(source.backend as usize)
110 .or_default()
111 .push((source.local_col as usize, out_col));
112 }
113 }
114
115 for (&didx, group) in &backend_groups {
116 let local_cols: Vec<usize> = group.iter().map(|&(loc, _)| loc).collect();
117 let (_, _, group_triplets) =
118 self.data_vec[didx].read_triplets_by_columns(local_cols)?;
119
120 let l2g = self.data_local_to_global_row[didx].as_slice();
121 if self.data_has_intra_row_merges[didx] {
122 let mut acc: HashMap<(u64, u64), f32> = HashMap::default();
128 for (i, j, v) in group_triplets {
129 if let Some(c) = g2c[l2g[i as usize]] {
130 let out_col = group[j as usize].1 as u64;
131 *acc.entry((c as u64, out_col)).or_insert(0.0) += v;
132 }
133 }
134 for ((r, c), v) in acc {
135 f(r, c, v);
136 }
137 } else {
138 for (i, j, v) in group_triplets {
139 if let Some(c) = g2c[l2g[i as usize]] {
140 let out_col = group[j as usize].1 as u64;
141 f(c as u64, out_col, v);
142 }
143 }
144 }
145 }
146 }
147
148 Ok((nrow, ncol))
149 }
150
151 #[allow(clippy::type_complexity)]
157 pub fn columns_triplets<I>(
158 &self,
159 cells: I,
160 ) -> anyhow::Result<((usize, usize), Vec<(u64, u64, f32)>)>
161 where
162 I: Iterator<Item = usize>,
163 {
164 let cells: Vec<usize> = cells.collect();
165 let one_slab = cells.len().max(1);
166 let mut triplets = Vec::new();
167 let dims = self.for_each_triplet(cells.into_iter(), one_slab, |r, c, v| {
168 triplets.push((r, c, v));
169 })?;
170 Ok((dims, triplets))
171 }
172
173 #[cfg(feature = "ndarray")]
174 pub fn read_columns_ndarray<I>(&self, cells: I) -> anyhow::Result<ndarray::Array2<f32>>
175 where
176 I: Iterator<Item = usize>,
177 {
178 let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
179 ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
180 }
181
182 pub fn read_columns_dmatrix<I>(&self, cells: I) -> anyhow::Result<nalgebra::DMatrix<f32>>
183 where
184 I: Iterator<Item = usize>,
185 {
186 let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
187 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
188 }
189
190 pub fn read_columns_csc<I>(&self, cells: I) -> anyhow::Result<CscMatrix<f32>>
201 where
202 I: Iterator<Item = usize>,
203 {
204 let cells: Vec<usize> = cells.collect();
205 let nrow = self.num_rows();
206 let ncol = cells.len();
207
208 let mut backend_groups: Vec<Vec<(usize, usize)>> =
215 (0..self.data_vec.len()).map(|_| Vec::new()).collect();
216 for (out_col, &glob) in cells.iter().enumerate() {
217 for source in self.col_to_data[glob].iter() {
218 backend_groups[source.backend as usize].push((source.local_col as usize, out_col));
219 }
220 }
221
222 let mut buckets: Vec<Vec<(u32, f32)>> = (0..ncol).map(|_| Vec::new()).collect();
224 let g2c = self.global_to_compact_row.as_slice();
225
226 for (didx, group) in backend_groups.iter().enumerate() {
227 if group.is_empty() {
228 continue;
229 }
230 let l2g = self.data_local_to_global_row[didx].as_slice();
231
232 if let Some((indptr, indices, values)) = self.data_vec[didx].csc_column_arrays() {
233 for &(loc, out_col) in group {
237 if loc + 1 >= indptr.len() {
238 continue;
239 }
240 index_audit::check_slot(
241 "preloaded column",
242 loc,
243 indptr[loc],
244 indptr[loc + 1],
245 indices.len().min(values.len()),
246 )?;
247 let s = indptr[loc] as usize;
248 let e = indptr[loc + 1] as usize;
249 let bucket = &mut buckets[out_col];
250 bucket.reserve(e - s);
251 for k in s..e {
252 if let Some(c) = remap_row(l2g, g2c, indices[k], didx, loc)? {
253 bucket.push((c as u32, values[k]));
254 }
255 }
256 }
257 } else {
258 let cols: Vec<usize> = group.iter().map(|&(loc, _)| loc).collect();
265 let (_, _, trip) = self.data_vec[didx].read_triplets_by_columns(cols)?;
266 for (i, jj, v) in trip {
267 let &(loc, out_col) = group.get(jj as usize).ok_or_else(|| {
268 anyhow::anyhow!(
269 "backend {didx}: read returned output column {jj} of {}",
270 group.len()
271 )
272 })?;
273 if let Some(c) = remap_row(l2g, g2c, i, didx, loc)? {
274 buckets[out_col].push((c as u32, v));
275 }
276 }
277 }
278 }
279
280 let total_nnz: usize = buckets.iter().map(|b| b.len()).sum();
282 let mut col_offsets: Vec<usize> = Vec::with_capacity(ncol + 1);
283 let mut row_indices: Vec<usize> = Vec::with_capacity(total_nnz);
284 let mut values: Vec<f32> = Vec::with_capacity(total_nnz);
285 col_offsets.push(0);
286
287 for bucket in &mut buckets {
288 let strictly_sorted_unique = bucket.windows(2).all(|w| w[0].0 < w[1].0);
300 if !strictly_sorted_unique {
301 bucket.sort_by_key(|&(r, _)| r);
302 let mut write = 0usize;
304 let mut read = 0usize;
305 while read < bucket.len() {
306 let (r, mut v) = bucket[read];
307 read += 1;
308 while read < bucket.len() && bucket[read].0 == r {
309 v += bucket[read].1;
310 read += 1;
311 }
312 bucket[write] = (r, v);
313 write += 1;
314 }
315 bucket.truncate(write);
316 }
317 for &(r, v) in bucket.iter() {
318 row_indices.push(r as usize);
319 values.push(v);
320 }
321 col_offsets.push(row_indices.len());
322 }
323
324 CscMatrix::try_from_csc_data(nrow, ncol, col_offsets, row_indices, values)
325 .map_err(|e| anyhow::anyhow!("CSC construction failed: {:?}", e))
326 }
327
328 pub fn read_columns_csr<I>(&self, cells: I) -> anyhow::Result<CsrMatrix<f32>>
329 where
330 I: Iterator<Item = usize>,
331 {
332 let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
333 nalgebra_sparse::CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
334 }
335
336 #[cfg(feature = "tensor")]
337 pub fn read_columns_tensor<I>(&self, cells: I) -> anyhow::Result<Tensor>
338 where
339 I: Iterator<Item = usize>,
340 {
341 let ((nrow, ncol), triplets) = self.columns_triplets(cells)?;
342 Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
343 }
344
345 #[allow(clippy::type_complexity)]
350 pub fn rows_triplets<I>(
351 &self,
352 rows: I,
353 ) -> anyhow::Result<((usize, usize), Vec<(u64, u64, f32)>)>
354 where
355 I: Iterator<Item = usize>,
356 {
357 let rows_compact: Vec<usize> = rows.collect();
358 let nrow_out = rows_compact.len();
359 let ncol_out = self.num_columns();
360
361 let n_compact = self.cached_num_rows;
362 let compact_to_global = self.compact_to_global_row.as_slice();
363
364 let mut triplets: Vec<(u64, u64, f32)> = Vec::new();
365 let mut local_to_out: Vec<usize> = Vec::with_capacity(rows_compact.len());
366 for didx in 0..self.data_vec.len() {
367 let g2l = &self.data_global_to_local_row[didx];
368
369 local_to_out.clear();
370 let mut local_rows: Vec<usize> = Vec::with_capacity(rows_compact.len());
371 for (out_row, &c) in rows_compact.iter().enumerate() {
372 if c >= n_compact {
373 continue;
374 }
375 let g = compact_to_global[c];
376 if let Some(&l) = g2l.get(&g) {
377 local_rows.push(l);
378 local_to_out.push(out_row);
379 }
380 }
381 if local_rows.is_empty() {
382 continue;
383 }
384
385 let (_, _, group_triplets) = self.data_vec[didx].read_triplets_by_rows(local_rows)?;
386 let cols_map = self
387 .data_to_cols
388 .get(&didx)
389 .ok_or_else(|| anyhow::anyhow!("missing data_to_cols entry for didx {}", didx))?;
390 triplets.reserve(group_triplets.len());
391 for (i, j, v) in group_triplets {
392 let mapped = cols_map[j as usize];
394 if mapped == usize::MAX {
395 continue;
396 }
397 let out_row = local_to_out[i as usize] as u64;
398 let out_col = mapped as u64;
399 triplets.push((out_row, out_col, v));
400 }
401 }
402
403 Ok(((nrow_out, ncol_out), triplets))
404 }
405
406 #[cfg(feature = "ndarray")]
407 pub fn read_rows_ndarray<I>(&self, rows: I) -> anyhow::Result<ndarray::Array2<f32>>
408 where
409 I: Iterator<Item = usize>,
410 {
411 let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
412 ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
413 }
414
415 pub fn read_rows_dmatrix<I>(&self, rows: I) -> anyhow::Result<nalgebra::DMatrix<f32>>
416 where
417 I: Iterator<Item = usize>,
418 {
419 let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
420 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
421 }
422
423 pub fn read_rows_csc<I>(&self, rows: I) -> anyhow::Result<CscMatrix<f32>>
424 where
425 I: Iterator<Item = usize>,
426 {
427 let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
428 nalgebra_sparse::CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
429 }
430
431 pub fn read_rows_csr<I>(&self, rows: I) -> anyhow::Result<CsrMatrix<f32>>
432 where
433 I: Iterator<Item = usize>,
434 {
435 let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
436 nalgebra_sparse::CsrMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)
437 }
438
439 #[cfg(feature = "tensor")]
440 pub fn read_rows_tensor<I>(&self, rows: I) -> anyhow::Result<Tensor>
441 where
442 I: Iterator<Item = usize>,
443 {
444 let ((nrow, ncol), triplets) = self.rows_triplets(rows)?;
445 Tensor::from_nonzero_triplets(nrow, ncol, &triplets)
446 }
447}