data_beans/sparse_io_vector/
batch.rs1#![allow(dead_code)]
2
3use super::*;
4
5impl SparseIoVec {
6 #[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 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 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 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 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 pub fn batch_knn_lookup(&self) -> Option<&Vec<ColumnDict<usize>>> {
261 self.derived.batch_knn_lookup.as_ref()
262 }
263
264 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 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 #[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 #[must_use]
341 pub fn has_column_multiplicity(&self) -> bool {
342 self.derived.col_multiplicity.is_some()
343 }
344
345 #[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}