Skip to main content

data_beans/sparse_io_vector/
batch.rs

1#![allow(dead_code)]
2
3use super::*;
4
5impl SparseIoVec {
6    /////////////////////////////////////
7    // data dictionary related methods //
8    /////////////////////////////////////
9
10    /// Register batch membership information along with the feature
11    /// matrix for quick look up operations.
12    ///
13    /// # Arguments
14    /// * `feature_matrix` - A feature matrix where each column corresponds to a cell.
15    /// * `batch_membership` - A vector of batch membership information for each cell.
16    #[cfg(feature = "ndarray")]
17    pub fn register_batches_ndarray<T>(
18        &mut self,
19        feature_matrix: &ndarray::Array2<f32>,
20        batch_membership: &[T],
21    ) -> anyhow::Result<()>
22    where
23        T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
24    {
25        {
26            debug_assert_eq!(batch_membership.len(), feature_matrix.ncols());
27        }
28        self._register_batches(
29            feature_matrix,
30            batch_membership,
31            |feature_matrix, batch_cells| {
32                let columns = batch_cells
33                    .iter()
34                    .map(|&c| feature_matrix.column(c))
35                    .collect::<Vec<_>>();
36                ColumnDict::<usize>::from_ndarray_views(columns, batch_cells.clone())
37            },
38        )
39    }
40
41    /// Register batch membership information along with the feature
42    /// matrix for quick look up operations.
43    ///
44    /// # Arguments
45    /// * `feature_matrix` - A feature matrix where each column corresponds to a cell.
46    /// * `batch_membership` - A vector of batch membership information for each cell.
47    pub fn register_batches_dmatrix<T>(
48        &mut self,
49        feature_matrix: &nalgebra::DMatrix<f32>,
50        batch_membership: &[T],
51    ) -> anyhow::Result<()>
52    where
53        T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
54    {
55        {
56            debug_assert_eq!(batch_membership.len(), feature_matrix.ncols());
57        }
58
59        self._register_batches(
60            feature_matrix,
61            batch_membership,
62            |feature_matrix, batch_cells| {
63                let columns = batch_cells
64                    .iter()
65                    .map(|&c| feature_matrix.column(c))
66                    .collect::<Vec<_>>();
67                ColumnDict::<usize>::from_dvector_views(columns, batch_cells.clone())
68            },
69        )
70    }
71
72    fn _register_batches<M, F, T>(
73        &mut self,
74        feature_matrix: &M,
75        batch_membership: &[T],
76        create_column_dict: F,
77    ) -> anyhow::Result<()>
78    where
79        M: Sync,
80        F: Fn(&M, &Vec<usize>) -> ColumnDict<usize> + Sync,
81        T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
82    {
83        let batches = partition_by_membership(batch_membership, None);
84
85        let ntot = self.num_columns();
86        let mut col_to_batch = vec![0; ntot];
87
88        // A ColumnDict build is internally parallel (instant-distance's rayon
89        // insert) only ABOVE EXACT_THRESHOLD points; smaller batches use the
90        // exact backend and build single-threaded. So when batches are many
91        // (>= n_threads) run the outer loop in parallel — that is the only
92        // parallelism the small batches get, and large ones just nest into the
93        // same work-stealing pool (safe, no true oversubscription). When batches
94        // are few, build sequentially so each large build owns the pool and only
95        // one HNSW index is under construction at a time (bounds peak memory).
96        let n_threads = rayon::current_num_threads();
97        let outer_parallel = batches.len() >= n_threads;
98
99        info!(
100            "building per-batch kNN indices ({} batches, {} cells, {}) ...",
101            batches.len(),
102            ntot,
103            if outer_parallel {
104                "parallel over batches"
105            } else {
106                "sequential over batches"
107            }
108        );
109
110        let prog_bar =
111            crate::sparse_data_visitors::styled_progress_bar(batches.len() as u64, "batches kNN");
112
113        // Canonical batch id = SORTED LABEL order (consistent with
114        // `register_batch_membership`), so per-batch stats / δ columns carry the
115        // same batch ids as the rest of the pipeline. Previously this sorted by
116        // descending size and used the processing position as the id, which
117        // silently permuted batch ids vs the label order — a latent bug for δ /
118        // `AdjMethod::Batch` consumers.
119        let mut batches_vec: Vec<_> = batches.into_iter().collect();
120        batches_vec.sort_by(|a, b| a.0.to_string().cmp(&b.0.to_string()));
121        // Schedule largest batches first (LPT) so the heavy tail overlaps with
122        // the small batches. The canonical id rides along; `sort_by_key(idx)`
123        // below restores canonical order.
124        let mut enumerated: Vec<_> = batches_vec.iter().enumerate().collect();
125        enumerated.sort_by_key(|(_, (_, cells))| std::cmp::Reverse(cells.len()));
126        let mut idx_name_glob_dict: Vec<_> = if outer_parallel {
127            enumerated
128                .into_par_iter()
129                .progress_with(prog_bar.clone())
130                .map(|(batch_index, (batch_name, batch_glob_indices))| {
131                    (
132                        batch_index,
133                        batch_name.to_string().into_boxed_str(),
134                        batch_glob_indices.clone(),
135                        create_column_dict(feature_matrix, batch_glob_indices),
136                    )
137                })
138                .collect()
139        } else {
140            enumerated
141                .into_iter()
142                .progress_with(prog_bar.clone())
143                .map(|(batch_index, (batch_name, batch_glob_indices))| {
144                    (
145                        batch_index,
146                        batch_name.to_string().into_boxed_str(),
147                        batch_glob_indices.clone(),
148                        create_column_dict(feature_matrix, batch_glob_indices),
149                    )
150                })
151                .collect()
152        };
153        prog_bar.finish_and_clear();
154
155        idx_name_glob_dict.sort_by_key(|&(idx, _, _, _)| idx);
156
157        let mut batch_names = vec![];
158        let mut batch_to_cols = vec![];
159        let mut dictionaries = vec![];
160
161        for (batch_idx, batch_name, glob_indices, dict) in idx_name_glob_dict.into_iter() {
162            dict.names()
163                .iter()
164                .for_each(|&cell| col_to_batch[cell] = batch_idx);
165
166            batch_names.push(batch_name);
167            batch_to_cols.push(glob_indices);
168            dictionaries.push(dict);
169        }
170
171        self.derived.batch_knn_lookup = Some(dictionaries);
172        self.derived.col_to_batch = Some(col_to_batch);
173        self.derived.batch_to_cols = Some(batch_to_cols);
174        self.derived.batch_idx_to_name = Some(batch_names);
175
176        if self.num_batches() > 2 {
177            self.sort_batch_proximity()?;
178        }
179
180        Ok(())
181    }
182
183    fn sort_batch_proximity(&mut self) -> anyhow::Result<()> {
184        let lookups = self
185            .derived
186            .batch_knn_lookup
187            .as_ref()
188            .ok_or(anyhow::anyhow!("no knn lookup"))?;
189
190        use nalgebra::DMatrix;
191
192        info!("retrieving batch-specific lookups");
193        let batch_data = lookups
194            .iter()
195            .flat_map(|dict| {
196                let data: Vec<f32> = dict.points().flatten().copied().collect();
197                let ncols = dict.num_points();
198                let nrows = data.len() / ncols;
199                DMatrix::from_vec(nrows, ncols, data)
200                    .column_mean()
201                    .data
202                    .as_vec()
203                    .clone()
204            })
205            .collect::<Vec<_>>();
206
207        let ncols = self.num_batches();
208        let nrows = batch_data.len() / ncols;
209        let batch_features = DMatrix::<f32>::from_vec(nrows, ncols, batch_data);
210
211        info!(
212            "built feature matrix across batches: {} x {}",
213            batch_features.nrows(),
214            batch_features.ncols()
215        );
216
217        let nbatches = self.num_batches();
218        let batches = (0..nbatches).collect();
219
220        let dict = ColumnDict::<usize>::from_dvector_views(
221            batch_features.column_iter().collect(),
222            batches,
223        );
224
225        let ret: Vec<Vec<usize>> = (0..nbatches)
226            .into_par_iter()
227            .map(|b| {
228                dict.search_by_query_name(&b, nbatches, false)
229                    .map(|(others, _)| others)
230            })
231            .collect::<anyhow::Result<Vec<Vec<usize>>>>()?;
232        self.derived.between_batch_proximity = Some(ret);
233
234        Ok(())
235    }
236
237    pub fn batch_name_map(&self) -> Option<HashMap<Box<str>, usize>> {
238        self.derived.batch_idx_to_name.as_ref().map(|names| {
239            names
240                .iter()
241                .enumerate()
242                .map(|(idx, name)| (name.clone(), idx))
243                .collect::<HashMap<Box<str>, usize>>()
244        })
245    }
246
247    pub fn num_batches(&self) -> usize {
248        if let Some(v) = &self.derived.batch_to_cols {
249            v.len()
250        } else if let Some(v) = &self.derived.batch_knn_lookup {
251            v.len()
252        } else {
253            0
254        }
255    }
256
257    /// Borrow the per-batch HNSW lookups populated by
258    /// `build_hnsw_per_batch` / `register_batches_dmatrix`. Returns `None`
259    /// before the indices have been built.
260    pub fn batch_knn_lookup(&self) -> Option<&Vec<ColumnDict<usize>>> {
261        self.derived.batch_knn_lookup.as_ref()
262    }
263
264    /// Register batch membership information without building HNSW
265    /// indices. This is a lightweight alternative to `register_batches_dmatrix`
266    /// for use with pb-sample based batch correction.
267    pub fn register_batch_membership<T>(&mut self, batch_membership: &[T])
268    where
269        T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
270    {
271        let batches = partition_by_membership(batch_membership, None);
272        let ntot = self.num_columns();
273        let mut col_to_batch = vec![0; ntot];
274
275        let mut sorted_batches: Vec<_> = batches.into_iter().collect();
276        sorted_batches.sort_by(|a, b| a.0.to_string().cmp(&b.0.to_string()));
277
278        let mut batch_names = Vec::with_capacity(sorted_batches.len());
279        let mut batch_to_cols = Vec::with_capacity(sorted_batches.len());
280
281        for (batch_idx, (batch_name, glob_indices)) in sorted_batches.into_iter().enumerate() {
282            for &cell in &glob_indices {
283                col_to_batch[cell] = batch_idx;
284            }
285            batch_names.push(batch_name.to_string().into_boxed_str());
286            batch_to_cols.push(glob_indices);
287        }
288
289        self.derived.col_to_batch = Some(col_to_batch);
290        self.derived.batch_to_cols = Some(batch_to_cols);
291        self.derived.batch_idx_to_name = Some(batch_names);
292    }
293
294    /// Declare that each column stands for more than one observation.
295    ///
296    /// A column is normally one cell, so every statistic that divides by a
297    /// count adds `1` per column. That breaks when a column is a *summary* of
298    /// many cells — a carried pseudobulk, or a bulk sample — because the
299    /// per-cell rate `μ = Σy / n` would divide a whole group's counts by one.
300    ///
301    /// With multiplicities registered, a column holding the **mean** profile of
302    /// `m` cells and a weight of `m` contributes exactly what those `m` cells
303    /// would have: `m·mean` to the sums and `m` to the count.
304    ///
305    /// Absent (the default) every column weighs `1`, and every accumulation is
306    /// bit-for-bit what it was before this existed.
307    ///
308    /// # Errors
309    /// If `multiplicity` is not one entry per column, or holds a non-finite or
310    /// non-positive weight — a zero would silently delete a column from the
311    /// denominator while leaving its counts in the numerator.
312    pub fn register_column_multiplicity(&mut self, multiplicity: &[f32]) -> anyhow::Result<()> {
313        let ntot = self.num_columns();
314        anyhow::ensure!(
315            multiplicity.len() == ntot,
316            "column multiplicity has {} entries but there are {ntot} columns",
317            multiplicity.len(),
318        );
319        if let Some((i, w)) = multiplicity
320            .iter()
321            .enumerate()
322            .find(|(_, w)| !w.is_finite() || **w <= 0.0)
323        {
324            anyhow::bail!("column {i} has multiplicity {w}; weights must be finite and positive");
325        }
326        self.derived.col_multiplicity = Some(multiplicity.to_vec());
327        Ok(())
328    }
329
330    /// Weight of a single column — `1.0` when no multiplicities are registered.
331    #[must_use]
332    pub fn column_multiplicity(&self, col: usize) -> f32 {
333        self.derived
334            .col_multiplicity
335            .as_ref()
336            .map_or(1.0, |m| m[col])
337    }
338
339    /// True when any column stands for more than one observation.
340    #[must_use]
341    pub fn has_column_multiplicity(&self) -> bool {
342        self.derived.col_multiplicity.is_some()
343    }
344
345    /// The whole multiplicity vector, one weight per column — `None` when no
346    /// multiplicities are registered (every column is one observation).
347    ///
348    /// Prefer this over gathering [`Self::column_multiplicity`] in a loop:
349    /// callers were rebuilding the vector element-by-element, re-encoding the
350    /// "absent means 1.0" default at every site.
351    #[must_use]
352    pub fn column_multiplicities(&self) -> Option<&[f32]> {
353        self.derived.col_multiplicity.as_deref()
354    }
355
356    pub fn batch_names(&self) -> Option<Vec<Box<str>>> {
357        self.derived.batch_idx_to_name.clone()
358    }
359
360    pub fn batch_to_columns(&self, batch: usize) -> Option<&Vec<usize>> {
361        if let Some(batch_to_cols) = &self.derived.batch_to_cols {
362            Some(&batch_to_cols[batch])
363        } else {
364            None
365        }
366    }
367
368    pub fn get_batch_membership<I>(&self, cells: I) -> Vec<usize>
369    where
370        I: Iterator<Item = usize>,
371    {
372        let cell_to_batch = self
373            .derived
374            .col_to_batch
375            .as_ref()
376            .expect("cell_to_batch not initialized");
377        cells.into_iter().map(|c| cell_to_batch[c]).collect()
378    }
379
380    pub fn column_names(&self) -> anyhow::Result<Vec<Box<str>>> {
381        debug_assert_eq!(self.num_columns(), self.column_names_with_data_tag.len());
382        Ok(self.column_names_with_data_tag.clone())
383    }
384}