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.into_iter()) {
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
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    /// Take columns within the neighbourhood of given `cells`
308    ///
309    /// # Arguments
310    /// * `cells` - global column indices
311    /// * `target_batches` - the batches for targeted kNN search
312    /// * `knn_batches` - k-nearest neighbour batches
313    /// * `knn_columns` - k-nearest neighbour columns
314    /// * `skip_same_batch` - skip the same batch
315    ///
316    /// # Returns
317    /// * the knn-matched matrix
318    /// * `source_columns` - a vector of the source columns
319    /// * a vector of distances between the matched columns
320    #[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    /// Take columns neighbouring with the given `cells`
356    ///
357    /// # Arguments
358    /// * `cells` - global column indices
359    /// * `target_batches` - the batches for targeted kNN search
360    /// * `knn` - k-nearest neighbours
361    /// * `skip_same_batch` - skip the same batch
362    ///
363    /// # Returns
364    /// * the knn-neighbouring matrix
365    /// * `source_columns` - a vector of the source columns
366    /// * distances - a vector of distances between the neighbouring columns
367    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    /// Take columns neighbouring with the given `cells`
399    ///
400    /// # Arguments
401    /// * `cells` - global column indices
402    /// * `target_batches` - the batches for targeted kNN search
403    /// * `knn` - k-nearest neighbours
404    /// * `skip_same_batch` - skip the same batch
405    ///
406    /// # Returns
407    /// * the knn-neighbouring matrix
408    /// * `source_columns` - a vector of the source columns
409    /// * distances - a vector of distances between the neighbouring columns
410    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    /// Take columns matched with the given `cells`
442    ///
443    /// # Arguments
444    /// * `cells` - global column indices
445    /// * `target_batches` - the batches for targeted kNN search
446    /// * `knn` - k-nearest neighbours
447    /// * `skip_same_batch` - skip the same batch
448    ///
449    /// # Returns
450    /// * the knn-matched matrix
451    /// * `source_columns` - a vector of the source columns
452    /// * a vector of distances between the matched columns
453    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    /// Take columns matched with the given `cells`
479    ///
480    /// # Arguments
481    /// * `cells` - global column indices
482    /// * `target_batches` - the batches for targeted kNN search
483    /// * `knn` - k-nearest neighbours
484    /// * `skip_same_batch` - skip the same batch
485    ///
486    /// # Returns
487    /// * the knn-matched matrix
488    /// * `source_columns` - a vector of the source columns
489    /// * distances - a vector of distances between the matched columns
490    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    /// Take columns matched with the given `cells`
516    ///
517    /// # Arguments
518    /// * `cells` - global column indices
519    /// * `target_batches` - the batches for targeted kNN search
520    /// * `knn` - k-nearest neighbours
521    /// * `skip_same_batch` - skip the same batch
522    ///
523    /// # Returns
524    /// * the knn-matched matrix
525    /// * `source_columns` - a vector of the source columns
526    /// * distances - a vector of distances between the matched columns
527    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    /// Query columns with projection data
553    ///
554    /// # Arguments
555    /// * `query` - a data vector
556    /// * `knn_per_batch` - k-nearest neighbour columns per batch
557    ///
558    /// # Returns
559    /// * the knn-matched matrix
560    /// * `source_columns` - a vector of the source columns
561    /// * a vector of distances between the matched columns
562    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    /// Query columns with projection data
587    ///
588    /// # Arguments
589    /// * `query` - a data vector
590    /// * `knn_per_batch` - k-nearest neighbour columns per batch
591    ///
592    /// # Returns
593    /// * the knn-matched matrix
594    /// * `source_columns` - a vector of the source columns
595    /// * a vector of distances between the matched columns
596    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    /// Query columns with projection data
620    ///
621    /// # Arguments
622    /// * `query` - a data vector
623    /// * `knn_per_batch` - k-nearest neighbour columns per batch
624    ///
625    /// # Returns
626    /// * the knn-matched matrix
627    /// * `source_columns` - a vector of the source columns
628    /// * a vector of distances between the matched columns
629    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}