Skip to main content

data_beans/sparse_io_vector/
matched.rs

1#![allow(dead_code)]
2
3use super::*;
4
5impl SparseIoVec {
6    /////////////////////
7    // matched columns //
8    /////////////////////
9
10    /// Take columns matched with the given `cells` on a specific
11    /// target batch
12    ///
13    /// # Arguments
14    /// * `cells` - global column indices
15    /// * `target_batch` - a batch for targeted kNN search
16    /// * `knn` - k-nearest neighbours
17    /// * `skip_same_batch` - skip the same batch
18    ///
19    /// # Returns
20    /// * shape - (nrows, ncols)
21    /// * triplets
22    /// * distances
23    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]; // this cell's batch
56
57            if skip_same_batch && source_batch == target_batch {
58                continue; // skip cells in the same batch
59            }
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) {
67                    if glob == glob_matched {
68                        continue; // avoid identical cell pairs
69                    }
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    /// Take columns matched with the given `cells`
88    ///
89    /// # Arguments
90    /// * `cells` - global column indices
91    /// * `target_batches` - the batches for targeted kNN search
92    /// * `knn_columns` - k-nearest neighbours of columns
93    /// * `skip_same_batch` - skip the same batch
94    ///
95    /// # Returns
96    /// * shape - (nrows, ncols)
97    /// * triplets - a vector of triplets
98    /// * `source_columns` - a vector of the source columns
99    /// * distance - a vector of distances between the matched columns
100    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    /// Take columns with the neighbourhood of given `cells`
161    ///
162    /// # Arguments
163    /// * `cells` - global column indices
164    /// * `knn_batches` - k-nearest neighbour batches
165    /// * `knn_columns` - k-nearest neighbour columns
166    /// * `skip_same_batch` - skip the same batch
167    ///
168    /// # Returns
169    /// * shape - (nrows, ncols)
170    /// * triplets
171    /// * source positions
172    /// * distances
173    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        // Pre-compute neighbouring batches per source batch
209        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 matched.into_iter().zip(matched_distances) {
239                        if glob_index == glob_matched_index {
240                            continue;
241                        }
242                        self.read_column_offset(glob_matched_index, &mut ncol, &mut triplets)?;
243                        source_columns.push(glob_index);
244                        matched_columns.push(glob_matched_index);
245                        distances.push(dist);
246                    }
247                }
248            }
249        }
250
251        Ok(TripletsMatched {
252            shape: (nrow, ncol),
253            triplets,
254            source_columns,
255            matched_columns,
256            distances,
257        })
258    }
259
260    fn query_columns_by_data_triplets<T>(
261        &self,
262        query: T,
263        knn_per_batch: usize,
264    ) -> anyhow::Result<TripletsMatched>
265    where
266        T: MakeVecPoint,
267    {
268        let lookups = self
269            .derived
270            .batch_knn_lookup
271            .as_ref()
272            .ok_or(anyhow::anyhow!("no knn lookup"))?;
273
274        let nrow = self.num_rows();
275        let mut ncol = 0_usize;
276
277        let approx_knn = self.num_batches() * knn_per_batch;
278        let mut triplets = Vec::with_capacity(approx_knn);
279        let mut source_columns = Vec::with_capacity(approx_knn);
280        let mut matched_columns = Vec::with_capacity(approx_knn);
281        let mut distances = Vec::with_capacity(approx_knn);
282
283        let q = query.to_vp();
284        for lookup in lookups {
285            let (matched, matched_distances) =
286                lookup.search_by_query_data(q.as_slice(), knn_per_batch)?;
287
288            for (&glob_idx, &dist) in matched.iter().zip(matched_distances.iter()) {
289                self.read_column_offset(glob_idx, &mut ncol, &mut triplets)?;
290                source_columns.push(glob_idx);
291                matched_columns.push(glob_idx);
292                distances.push(dist);
293            }
294        }
295
296        Ok(TripletsMatched {
297            shape: (nrow, ncol),
298            triplets,
299            source_columns,
300            matched_columns,
301            distances,
302        })
303    }
304
305    /// Take columns within the neighbourhood of given `cells`
306    ///
307    /// # Arguments
308    /// * `cells` - global column indices
309    /// * `target_batches` - the batches for targeted kNN search
310    /// * `knn_batches` - k-nearest neighbour batches
311    /// * `knn_columns` - k-nearest neighbour columns
312    /// * `skip_same_batch` - skip the same batch
313    ///
314    /// # Returns
315    /// * the knn-matched matrix
316    /// * `source_columns` - a vector of the source columns
317    /// * a vector of distances between the matched columns
318    #[allow(clippy::type_complexity)]
319    pub fn read_neighbouring_columns_csc<I>(
320        &self,
321        cells: I,
322        knn_batches: usize,
323        knn_columns: usize,
324        skip_same_batch: bool,
325        skip_batches: Option<&[usize]>,
326    ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<usize>, Vec<f32>)>
327    where
328        I: Iterator<Item = usize>,
329    {
330        let TripletsMatched {
331            shape: (nrow, ncol),
332            triplets,
333            source_columns,
334            matched_columns,
335            distances,
336        } = self.neighbouring_columns_triplets(
337            cells,
338            knn_batches,
339            knn_columns,
340            skip_same_batch,
341            skip_batches,
342        )?;
343
344        Ok((
345            CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
346            source_columns,
347            matched_columns,
348            distances,
349        ))
350    }
351
352    #[cfg(feature = "ndarray")]
353    /// Take columns neighbouring with the given `cells`
354    ///
355    /// # Arguments
356    /// * `cells` - global column indices
357    /// * `target_batches` - the batches for targeted kNN search
358    /// * `knn` - k-nearest neighbours
359    /// * `skip_same_batch` - skip the same batch
360    ///
361    /// # Returns
362    /// * the knn-neighbouring matrix
363    /// * `source_columns` - a vector of the source columns
364    /// * distances - a vector of distances between the neighbouring columns
365    pub fn read_neighbouring_columns_ndarray<I>(
366        &self,
367        cells: I,
368        knn_batches: usize,
369        knn_columns: usize,
370        skip_same_batch: bool,
371        skip_batches: Option<&[usize]>,
372    ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
373    where
374        I: Iterator<Item = usize>,
375    {
376        let TripletsMatched {
377            shape: (nrow, ncol),
378            triplets,
379            source_columns,
380            distances,
381            ..
382        } = self.neighbouring_columns_triplets(
383            cells,
384            knn_batches,
385            knn_columns,
386            skip_same_batch,
387            skip_batches,
388        )?;
389        Ok((
390            ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
391            source_columns,
392            distances,
393        ))
394    }
395
396    /// Take columns neighbouring with the given `cells`
397    ///
398    /// # Arguments
399    /// * `cells` - global column indices
400    /// * `target_batches` - the batches for targeted kNN search
401    /// * `knn` - k-nearest neighbours
402    /// * `skip_same_batch` - skip the same batch
403    ///
404    /// # Returns
405    /// * the knn-neighbouring matrix
406    /// * `source_columns` - a vector of the source columns
407    /// * distances - a vector of distances between the neighbouring columns
408    pub fn read_neighbouring_columns_dmatrix<I>(
409        &self,
410        cells: I,
411        knn_batches: usize,
412        knn_columns: usize,
413        skip_same_batch: bool,
414        skip_batches: Option<&[usize]>,
415    ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
416    where
417        I: Iterator<Item = usize>,
418    {
419        let TripletsMatched {
420            shape: (nrow, ncol),
421            triplets,
422            source_columns,
423            distances,
424            ..
425        } = self.neighbouring_columns_triplets(
426            cells,
427            knn_batches,
428            knn_columns,
429            skip_same_batch,
430            skip_batches,
431        )?;
432        Ok((
433            DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
434            source_columns,
435            distances,
436        ))
437    }
438
439    /// Take columns matched with the given `cells`
440    ///
441    /// # Arguments
442    /// * `cells` - global column indices
443    /// * `target_batches` - the batches for targeted kNN search
444    /// * `knn` - k-nearest neighbours
445    /// * `skip_same_batch` - skip the same batch
446    ///
447    /// # Returns
448    /// * the knn-matched matrix
449    /// * `source_columns` - a vector of the source columns
450    /// * a vector of distances between the matched columns
451    pub fn read_matched_columns_csc<I>(
452        &self,
453        cells: I,
454        target_batches: &[usize],
455        knn: usize,
456        skip_same_batch: bool,
457    ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<f32>)>
458    where
459        I: Iterator<Item = usize>,
460    {
461        let TripletsMatched {
462            shape: (nrow, ncol),
463            triplets,
464            source_columns,
465            distances,
466            ..
467        } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
468        Ok((
469            CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
470            source_columns,
471            distances,
472        ))
473    }
474
475    #[cfg(feature = "ndarray")]
476    /// Take columns matched with the given `cells`
477    ///
478    /// # Arguments
479    /// * `cells` - global column indices
480    /// * `target_batches` - the batches for targeted kNN search
481    /// * `knn` - k-nearest neighbours
482    /// * `skip_same_batch` - skip the same batch
483    ///
484    /// # Returns
485    /// * the knn-matched matrix
486    /// * `source_columns` - a vector of the source columns
487    /// * distances - a vector of distances between the matched columns
488    pub fn read_matched_columns_ndarray<I>(
489        &self,
490        cells: I,
491        target_batches: &[usize],
492        knn: usize,
493        skip_same_batch: bool,
494    ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
495    where
496        I: Iterator<Item = usize>,
497    {
498        let TripletsMatched {
499            shape: (nrow, ncol),
500            triplets,
501            source_columns,
502            distances,
503            ..
504        } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
505
506        Ok((
507            ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
508            source_columns,
509            distances,
510        ))
511    }
512
513    /// Take columns matched with the given `cells`
514    ///
515    /// # Arguments
516    /// * `cells` - global column indices
517    /// * `target_batches` - the batches for targeted kNN search
518    /// * `knn` - k-nearest neighbours
519    /// * `skip_same_batch` - skip the same batch
520    ///
521    /// # Returns
522    /// * the knn-matched matrix
523    /// * `source_columns` - a vector of the source columns
524    /// * distances - a vector of distances between the matched columns
525    pub fn read_matched_columns_dmatrix<I>(
526        &self,
527        cells: I,
528        target_batches: &[usize],
529        knn: usize,
530        skip_same_batch: bool,
531    ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
532    where
533        I: Iterator<Item = usize>,
534    {
535        let TripletsMatched {
536            shape: (nrow, ncol),
537            triplets,
538            source_columns,
539            distances,
540            ..
541        } = self.matched_columns_triplets(cells, target_batches, knn, skip_same_batch)?;
542
543        Ok((
544            DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
545            source_columns,
546            distances,
547        ))
548    }
549
550    /// Query columns with projection data
551    ///
552    /// # Arguments
553    /// * `query` - a data vector
554    /// * `knn_per_batch` - k-nearest neighbour columns per batch
555    ///
556    /// # Returns
557    /// * the knn-matched matrix
558    /// * `source_columns` - a vector of the source columns
559    /// * a vector of distances between the matched columns
560    pub fn query_columns_by_data_csc<T>(
561        &self,
562        query: T,
563        knn_per_batch: usize,
564    ) -> anyhow::Result<(CscMatrix<f32>, Vec<usize>, Vec<f32>)>
565    where
566        T: MakeVecPoint,
567    {
568        let TripletsMatched {
569            shape: (nrow, ncol),
570            triplets,
571            source_columns,
572            distances,
573            ..
574        } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
575
576        Ok((
577            CscMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
578            source_columns,
579            distances,
580        ))
581    }
582
583    #[cfg(feature = "ndarray")]
584    /// Query columns with projection data
585    ///
586    /// # Arguments
587    /// * `query` - a data vector
588    /// * `knn_per_batch` - k-nearest neighbour columns per batch
589    ///
590    /// # Returns
591    /// * the knn-matched matrix
592    /// * `source_columns` - a vector of the source columns
593    /// * a vector of distances between the matched columns
594    pub fn query_columns_by_data_ndarray<T>(
595        &self,
596        query: T,
597        knn_per_batch: usize,
598    ) -> anyhow::Result<(ndarray::Array2<f32>, Vec<usize>, Vec<f32>)>
599    where
600        T: MakeVecPoint,
601    {
602        let TripletsMatched {
603            shape: (nrow, ncol),
604            triplets,
605            source_columns,
606            distances,
607            ..
608        } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
609
610        Ok((
611            ndarray::Array2::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
612            source_columns,
613            distances,
614        ))
615    }
616
617    /// Query columns with projection data
618    ///
619    /// # Arguments
620    /// * `query` - a data vector
621    /// * `knn_per_batch` - k-nearest neighbour columns per batch
622    ///
623    /// # Returns
624    /// * the knn-matched matrix
625    /// * `source_columns` - a vector of the source columns
626    /// * a vector of distances between the matched columns
627    pub fn query_columns_by_data_dmatrix<T>(
628        &self,
629        query: T,
630        knn_per_batch: usize,
631    ) -> anyhow::Result<(nalgebra::DMatrix<f32>, Vec<usize>, Vec<f32>)>
632    where
633        T: MakeVecPoint,
634    {
635        let TripletsMatched {
636            shape: (nrow, ncol),
637            triplets,
638            source_columns,
639            distances,
640            ..
641        } = self.query_columns_by_data_triplets(query, knn_per_batch)?;
642
643        Ok((
644            DMatrix::<f32>::from_nonzero_triplets(nrow, ncol, &triplets)?,
645            source_columns,
646            distances,
647        ))
648    }
649}