1#![allow(dead_code)]
2
3use super::*;
4
5impl SparseIoVec {
6 fn matched_columns_triplets_on_one_target<I>(
24 &self,
25 cells: I,
26 target_batch: usize,
27 knn: usize,
28 skip_same_batch: bool,
29 ) -> anyhow::Result<TripletsMatched>
30 where
31 I: Iterator<Item = usize> + Clone,
32 {
33 let lookups = self
34 .derived
35 .batch_knn_lookup
36 .as_ref()
37 .ok_or(anyhow::anyhow!("no knn lookup"))?;
38
39 let cell_to_batch = self
40 .derived
41 .col_to_batch
42 .as_ref()
43 .ok_or(anyhow::anyhow!("no cell to batch"))?;
44
45 debug_assert!(target_batch < self.num_batches());
46
47 let nrow = self.num_rows();
48 let mut ncol = 0;
49 let mut triplets = Vec::new();
50 let mut distances = Vec::new();
51 let mut source_columns = Vec::new();
52 let mut matched_columns = Vec::new();
53
54 for glob in cells {
55 let source_batch = cell_to_batch[glob]; if skip_same_batch && source_batch == target_batch {
58 continue; }
60
61 if let (Some(source_lookup), Some(target_lookup)) =
62 (lookups.get(source_batch), lookups.get(target_batch))
63 {
64 let (matched, matched_distances) =
65 source_lookup.match_by_query_name_against(&glob, knn, target_lookup)?;
66 for (glob_matched, dist) in matched.into_iter().zip(matched_distances.into_iter()) {
67 if glob == glob_matched {
68 continue; }
70 self.read_column_offset(glob_matched, &mut ncol, &mut triplets)?;
71 source_columns.push(glob);
72 matched_columns.push(glob_matched);
73 distances.push(dist);
74 }
75 }
76 }
77
78 Ok(TripletsMatched {
79 shape: (nrow, ncol),
80 triplets,
81 source_columns,
82 matched_columns,
83 distances,
84 })
85 }
86
87 fn matched_columns_triplets<I>(
101 &self,
102 cells: I,
103 target_batches: &[usize],
104 knn_columns: usize,
105 skip_same_batch: bool,
106 ) -> anyhow::Result<TripletsMatched>
107 where
108 I: Iterator<Item = usize>,
109 {
110 let cells: Vec<usize> = cells.collect();
111
112 let nrows = self.num_rows();
113 let nbatches = self.num_batches();
114 let ncols = cells.len();
115 let approx_ncols = ncols * knn_columns * nbatches;
116
117 let mut tot_triplets: Vec<(u64, u64, f32)> = Vec::with_capacity(approx_ncols);
118 let mut tot_distances: Vec<f32> = Vec::with_capacity(approx_ncols);
119 let mut tot_sources: Vec<usize> = Vec::with_capacity(approx_ncols);
120 let mut tot_matched: Vec<usize> = Vec::with_capacity(approx_ncols);
121 let mut tot_ncells_matched: usize = 0;
122
123 for &target_b in target_batches.iter() {
124 let TripletsMatched {
125 shape,
126 triplets,
127 source_columns,
128 matched_columns,
129 distances,
130 } = self.matched_columns_triplets_on_one_target(
131 cells.iter().cloned(),
132 target_b,
133 knn_columns,
134 skip_same_batch,
135 )?;
136
137 tot_triplets.extend(
138 triplets
139 .into_iter()
140 .map(|(i, j, z_ij)| (i, j + (tot_ncells_matched as u64), z_ij)),
141 );
142
143 tot_distances.extend(distances);
144 tot_ncells_matched += shape.1;
145 tot_sources.extend(source_columns);
146 tot_matched.extend(matched_columns);
147 }
148
149 let shape = (nrows, tot_ncells_matched);
150
151 Ok(TripletsMatched {
152 shape,
153 triplets: tot_triplets,
154 source_columns: tot_sources,
155 matched_columns: tot_matched,
156 distances: tot_distances,
157 })
158 }
159
160 fn neighbouring_columns_triplets<I>(
174 &self,
175 cells: I,
176 knn_batches: usize,
177 knn_columns: usize,
178 skip_same_batch: bool,
179 skip_batches: Option<&[usize]>,
180 ) -> anyhow::Result<TripletsMatched>
181 where
182 I: Iterator<Item = usize>,
183 {
184 let lookups = self
185 .derived
186 .batch_knn_lookup
187 .as_ref()
188 .ok_or(anyhow::anyhow!("no knn lookup"))?;
189
190 let cell_to_batch = self
191 .derived
192 .col_to_batch
193 .as_ref()
194 .ok_or(anyhow::anyhow!("no cell to batch"))?;
195
196 let approx_ncol = knn_columns * knn_batches;
197
198 let nrow = self.num_rows();
199 let mut ncol = 0_usize;
200 let mut triplets = Vec::with_capacity(approx_ncol * nrow);
201
202 let mut distances = Vec::with_capacity(approx_ncol);
203 let mut source_columns = Vec::with_capacity(approx_ncol);
204 let mut matched_columns = Vec::with_capacity(approx_ncol);
205
206 let nbatches = self.num_batches();
207
208 let neighbouring_batches_by_source: Vec<Vec<usize>> = (0..nbatches)
210 .map(
211 |source_batch| match self.derived.between_batch_proximity.as_ref() {
212 Some(prox) => prox[source_batch]
213 .iter()
214 .copied()
215 .filter(|&b| skip_batches.is_none_or(|skip| !skip.contains(&b)))
216 .filter(|&b| !skip_same_batch || b != source_batch)
217 .collect(),
218 _ => (0..nbatches)
219 .filter(|&b| skip_batches.is_none_or(|skip| !skip.contains(&b)))
220 .filter(|&b| !skip_same_batch || b != source_batch)
221 .collect(),
222 },
223 )
224 .collect();
225
226 for glob_index in cells {
227 let source_batch = cell_to_batch[glob_index];
228
229 for &target_batch in &neighbouring_batches_by_source[source_batch] {
230 if let (Some(source_lookup), Some(target_lookup)) =
231 (lookups.get(source_batch), lookups.get(target_batch))
232 {
233 let (matched, matched_distances) = source_lookup.match_by_query_name_against(
234 &glob_index,
235 knn_columns,
236 target_lookup,
237 )?;
238 for (glob_matched_index, dist) in
239 matched.into_iter().zip(matched_distances.into_iter())
240 {
241 if glob_index == glob_matched_index {
242 continue;
243 }
244 self.read_column_offset(glob_matched_index, &mut ncol, &mut triplets)?;
245 source_columns.push(glob_index);
246 matched_columns.push(glob_matched_index);
247 distances.push(dist);
248 }
249 }
250 }
251 }
252
253 Ok(TripletsMatched {
254 shape: (nrow, ncol),
255 triplets,
256 source_columns,
257 matched_columns,
258 distances,
259 })
260 }
261
262 fn query_columns_by_data_triplets<T>(
263 &self,
264 query: T,
265 knn_per_batch: usize,
266 ) -> anyhow::Result<TripletsMatched>
267 where
268 T: MakeVecPoint,
269 {
270 let lookups = self
271 .derived
272 .batch_knn_lookup
273 .as_ref()
274 .ok_or(anyhow::anyhow!("no knn lookup"))?;
275
276 let nrow = self.num_rows();
277 let mut ncol = 0_usize;
278
279 let approx_knn = self.num_batches() * knn_per_batch;
280 let mut triplets = Vec::with_capacity(approx_knn);
281 let mut source_columns = Vec::with_capacity(approx_knn);
282 let mut matched_columns = Vec::with_capacity(approx_knn);
283 let mut distances = Vec::with_capacity(approx_knn);
284
285 let q = query.to_vp();
286 for lookup in lookups {
287 let (matched, matched_distances) =
288 lookup.search_by_query_data(q.as_slice(), knn_per_batch)?;
289
290 for (&glob_idx, &dist) in matched.iter().zip(matched_distances.iter()) {
291 self.read_column_offset(glob_idx, &mut ncol, &mut triplets)?;
292 source_columns.push(glob_idx);
293 matched_columns.push(glob_idx);
294 distances.push(dist);
295 }
296 }
297
298 Ok(TripletsMatched {
299 shape: (nrow, ncol),
300 triplets,
301 source_columns,
302 matched_columns,
303 distances,
304 })
305 }
306
307 #[allow(clippy::type_complexity)]
321 pub fn read_neighbouring_columns_csc<I>(
322 &self,
323 cells: I,
324 knn_batches: usize,
325 knn_columns: usize,
326 skip_same_batch: bool,
327 skip_batches: Option<&[usize]>,
328 ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<usize>, Vec<f32>)>
329 where
330 I: Iterator<Item = usize>,
331 {
332 let TripletsMatched {
333 shape: (nrow, ncol),
334 triplets,
335 source_columns,
336 matched_columns,
337 distances,
338 } = self.neighbouring_columns_triplets(
339 cells,
340 knn_batches,
341 knn_columns,
342 skip_same_batch,
343 skip_batches,
344 )?;
345
346 Ok((
347 CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
348 source_columns,
349 matched_columns,
350 distances,
351 ))
352 }
353
354 #[cfg(feature = "ndarray")]
355 pub fn read_neighbouring_columns_ndarray<I>(
368 &self,
369 cells: I,
370 knn_batches: usize,
371 knn_columns: usize,
372 skip_same_batch: bool,
373 skip_batches: Option<&[usize]>,
374 ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
375 where
376 I: Iterator<Item = usize>,
377 {
378 let TripletsMatched {
379 shape: (nrow, ncol),
380 triplets,
381 source_columns,
382 distances,
383 ..
384 } = self.neighbouring_columns_triplets(
385 cells,
386 knn_batches,
387 knn_columns,
388 skip_same_batch,
389 skip_batches,
390 )?;
391 Ok((
392 ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
393 source_columns,
394 distances,
395 ))
396 }
397
398 pub fn read_neighbouring_columns_dmatrix<I>(
411 &self,
412 cells: I,
413 knn_batches: usize,
414 knn_columns: usize,
415 skip_same_batch: bool,
416 skip_batches: Option<&[usize]>,
417 ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
418 where
419 I: Iterator<Item = usize>,
420 {
421 let TripletsMatched {
422 shape: (nrow, ncol),
423 triplets,
424 source_columns,
425 distances,
426 ..
427 } = self.neighbouring_columns_triplets(
428 cells,
429 knn_batches,
430 knn_columns,
431 skip_same_batch,
432 skip_batches,
433 )?;
434 Ok((
435 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
436 source_columns,
437 distances,
438 ))
439 }
440
441 pub fn read_matched_columns_csc<I>(
454 &self,
455 cells: I,
456 target_batches: &[usize],
457 knn: usize,
458 skip_same_batch: bool,
459 ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<f32>)>
460 where
461 I: Iterator<Item = usize>,
462 {
463 let TripletsMatched {
464 shape: (nrow, ncol),
465 triplets,
466 source_columns,
467 distances,
468 ..
469 } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
470 Ok((
471 CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
472 source_columns,
473 distances,
474 ))
475 }
476
477 #[cfg(feature = "ndarray")]
478 pub fn read_matched_columns_ndarray<I>(
491 &self,
492 cells: I,
493 target_batches: &[usize],
494 knn: usize,
495 skip_same_batch: bool,
496 ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
497 where
498 I: Iterator<Item = usize>,
499 {
500 let TripletsMatched {
501 shape: (nrow, ncol),
502 triplets,
503 source_columns,
504 distances,
505 ..
506 } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
507
508 Ok((
509 ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
510 source_columns,
511 distances,
512 ))
513 }
514
515 pub fn read_matched_columns_dmatrix<I>(
528 &self,
529 cells: I,
530 target_batches: &[usize],
531 knn: usize,
532 skip_same_batch: bool,
533 ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
534 where
535 I: Iterator<Item = usize>,
536 {
537 let TripletsMatched {
538 shape: (nrow, ncol),
539 triplets,
540 source_columns,
541 distances,
542 ..
543 } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
544
545 Ok((
546 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
547 source_columns,
548 distances,
549 ))
550 }
551
552 pub fn query_columns_by_data_csc<T>(
563 &self,
564 query: T,
565 knn_per_batch: usize,
566 ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<f32>)>
567 where
568 T: MakeVecPoint,
569 {
570 let TripletsMatched {
571 shape: (nrow, ncol),
572 triplets,
573 source_columns,
574 distances,
575 ..
576 } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
577
578 Ok((
579 CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
580 source_columns,
581 distances,
582 ))
583 }
584
585 #[cfg(feature = "ndarray")]
586 pub fn query_columns_by_data_ndarray<T>(
597 &self,
598 query: T,
599 knn_per_batch: usize,
600 ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
601 where
602 T: MakeVecPoint,
603 {
604 let TripletsMatched {
605 shape: (nrow, ncol),
606 triplets,
607 source_columns,
608 distances,
609 ..
610 } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
611
612 Ok((
613 ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
614 source_columns,
615 distances,
616 ))
617 }
618
619 pub fn query_columns_by_data_dmatrix<T>(
630 &self,
631 query: T,
632 knn_per_batch: usize,
633 ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
634 where
635 T: MakeVecPoint,
636 {
637 let TripletsMatched {
638 shape: (nrow, ncol),
639 triplets,
640 source_columns,
641 distances,
642 ..
643 } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
644
645 Ok((
646 DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
647 source_columns,
648 distances,
649 ))
650 }
651}