1use crate::aux::feature_names::FeatureNameKind;
20use legume_numeric::matrix::traits::IoOps;
21use nalgebra::DMatrix;
22use rustc_hash::{FxHashMap, FxHashSet};
23
24pub struct FrozenFeatureHost {
33 pub e_feat: DMatrix<f32>,
35 pub b_feat: Vec<f32>,
37 pub keep_target_indices: Vec<usize>,
40 pub keep_src_indices: Vec<usize>,
44 pub src_e_feat: DMatrix<f32>,
48 pub src_names: Vec<Box<str>>,
49 pub n_src: usize,
56 pub h: usize,
57}
58
59pub type SourceNameMap<'a> = &'a dyn Fn(&str) -> Box<str>;
61
62pub struct FrozenLoadArgs<'a> {
63 pub dictionary_path: &'a str,
66 pub bias_path: Option<&'a str>,
70 pub target_feature_names: &'a [Box<str>],
74 pub name_kind: FeatureNameKind,
79 pub source_name_map: Option<SourceNameMap<'a>>,
86}
87
88pub fn load_frozen_feature_host(args: FrozenLoadArgs) -> anyhow::Result<FrozenFeatureHost> {
91 load_frozen_feature_host_matching(args, |names| Ok(vec![true; names.len()]))
92}
93
94pub fn load_frozen_feature_host_matching(
103 args: FrozenLoadArgs,
104 matchable: impl FnOnce(&[Box<str>]) -> anyhow::Result<Vec<bool>>,
105) -> anyhow::Result<FrozenFeatureHost> {
106 let dict = <DMatrix<f32> as IoOps>::from_parquet(args.dictionary_path)?;
107 let n_src = dict.rows.len();
108 let h = dict.mat.ncols();
109 anyhow::ensure!(
110 h > 0 && dict.mat.nrows() == n_src,
111 "{}: malformed dictionary (rows={}, mat dims={}x{})",
112 args.dictionary_path,
113 n_src,
114 dict.mat.nrows(),
115 h
116 );
117
118 let src_bias: Vec<f32> = match args.bias_path {
119 None => vec![0.0; n_src],
120 Some(p) => {
121 let bias = <DMatrix<f32> as IoOps>::from_parquet(p)?;
122 anyhow::ensure!(
123 bias.rows == dict.rows,
124 "{} row names disagree with {} (both files must come from the same training run)",
125 p,
126 args.dictionary_path
127 );
128 anyhow::ensure!(
129 bias.mat.ncols() == 1,
130 "{}: expected 1 data column (bias), got {}",
131 p,
132 bias.mat.ncols()
133 );
134 (0..n_src).map(|i| bias.mat[(i, 0)]).collect()
135 }
136 };
137
138 let matchable = matchable(&dict.rows)
139 .map_err(|e| anyhow::anyhow!("{}: marking its rows: {e}", args.dictionary_path))?;
140 anyhow::ensure!(
141 matchable.len() == n_src,
142 "{}: {} row flags for {} rows",
143 args.dictionary_path,
144 matchable.len(),
145 n_src
146 );
147 let n_matchable = matchable.iter().filter(|&&m| m).count();
148 anyhow::ensure!(
149 n_src == 0 || n_matchable > 0,
150 "{}: none of its {} rows is marked as one that may match a feature",
151 args.dictionary_path,
152 n_src
153 );
154 let src_names: Vec<Box<str>> = match args.source_name_map {
156 Some(f) => dict
157 .rows
158 .iter()
159 .zip(&matchable)
160 .map(|(n, &m)| if m { f(n) } else { n.clone() })
161 .collect(),
162 None => dict.rows,
163 };
164 let mut src_by_canon: FxHashMap<Box<str>, usize> = FxHashMap::default();
165 let mut src_dupes = 0usize;
166 for (i, name) in src_names.iter().enumerate() {
167 if !matchable[i] {
168 continue;
169 }
170 let canon = args.name_kind.canonicalize(name);
171 if let std::collections::hash_map::Entry::Vacant(e) = src_by_canon.entry(canon) {
174 e.insert(i);
175 } else {
176 src_dupes += 1;
177 }
178 }
179 if src_dupes > 0 {
180 log::warn!(
181 "{}: {} source rows had duplicate canonical names — kept first occurrence",
182 args.dictionary_path,
183 src_dupes
184 );
185 }
186 let shadowing = src_names
187 .iter()
188 .zip(&matchable)
189 .filter(|(n, &m)| !m && src_by_canon.contains_key(&args.name_kind.canonicalize(n)))
190 .count();
191 if shadowing > 0 {
192 log::info!(
193 "{}: {} unmarked rows share a name with a matchable row and were passed over",
194 args.dictionary_path,
195 shadowing
196 );
197 }
198
199 let mut keep_target_indices = Vec::new();
200 let mut keep_src_indices = Vec::new();
201 for (target_i, name) in args.target_feature_names.iter().enumerate() {
202 let canon = args.name_kind.canonicalize(name);
203 if let Some(&src_i) = src_by_canon.get(&canon) {
204 keep_target_indices.push(target_i);
205 keep_src_indices.push(src_i);
206 }
207 }
208 anyhow::ensure!(
209 !keep_target_indices.is_empty(),
210 "No feature names matched between {} (n={}) and target axis (n={}) under {:?} \
211 — check the gene-name kind (Exact / Gene / Locus / Mixed) and source axis",
212 args.dictionary_path,
213 n_src,
214 args.target_feature_names.len(),
215 args.name_kind
216 );
217
218 let unique_src_used: FxHashSet<usize> = keep_src_indices.iter().copied().collect();
219 let channelized_unmatched = src_names
224 .iter()
225 .enumerate()
226 .filter(|(i, r)| {
227 matchable[*i]
228 && !unique_src_used.contains(i)
229 && crate::aux::feature_rows::parse_feature_row(r).is_some()
230 })
231 .count();
232 if channelized_unmatched > 0 {
233 log::warn!(
234 "{}: {} unmatched source rows carry the channelized row grammar — is this a raw gene dictionary, or a channelized/co-embedding output?",
235 args.dictionary_path,
236 channelized_unmatched
237 );
238 }
239 let matchable_note = if n_matchable < n_src {
240 format!("; {n_matchable} source rows may match")
241 } else {
242 String::new()
243 };
244 log::info!(
245 "Frozen feature side from {}: {}/{} target features matched (H={}, {} of {} source rows reused{}, kind={:?})",
246 args.dictionary_path,
247 keep_target_indices.len(),
248 args.target_feature_names.len(),
249 h,
250 unique_src_used.len(),
251 n_src,
252 matchable_note,
253 args.name_kind
254 );
255
256 let k = keep_target_indices.len();
257 let mut e_feat = DMatrix::<f32>::zeros(k, h);
258 let mut b_feat = Vec::with_capacity(k);
259 for (out_i, &src_i) in keep_src_indices.iter().enumerate() {
260 for j in 0..h {
261 e_feat[(out_i, j)] = dict.mat[(src_i, j)];
262 }
263 b_feat.push(src_bias[src_i]);
264 }
265
266 Ok(FrozenFeatureHost {
267 e_feat,
268 b_feat,
269 keep_target_indices,
270 keep_src_indices,
271 src_e_feat: dict.mat,
272 src_names,
273 n_src,
274 h,
275 })
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281 use legume_numeric::matrix::traits::IoOps;
282
283 fn write_test_parquet(
284 path: &str,
285 rows: &[&str],
286 row_axis: &str,
287 cols: &[&str],
288 data: &DMatrix<f32>,
289 ) {
290 let row_names: Vec<Box<str>> = rows.iter().map(|s| (*s).into()).collect();
291 let col_names: Vec<Box<str>> = cols.iter().map(|s| (*s).into()).collect();
292 data.to_parquet_with_names(path, (Some(&row_names), Some(row_axis)), Some(&col_names))
293 .unwrap();
294 }
295
296 #[test]
297 fn strict_intersection_drops_unmatched_and_preserves_target_order() {
298 let dir = tempfile::tempdir().unwrap();
299 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
300
301 let src = DMatrix::<f32>::from_row_slice(
303 4,
304 3,
305 &[
306 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ],
311 );
312 write_test_parquet(
313 &dict_path,
314 &["TGFB1", "MYC", "ENSG_DROP", "TP53"],
315 "gene",
316 &["h0", "h1", "h2"],
317 &src,
318 );
319
320 let target: Vec<Box<str>> = ["FOO", "TP53", "TGFB1", "BAR", "MYC"]
322 .iter()
323 .map(|s| (*s).into())
324 .collect();
325
326 let host = load_frozen_feature_host(FrozenLoadArgs {
327 dictionary_path: &dict_path,
328 bias_path: None,
329 target_feature_names: &target,
330 name_kind: FeatureNameKind::Exact,
331 source_name_map: None,
332 })
333 .unwrap();
334
335 assert_eq!(host.keep_target_indices, vec![1, 2, 4]);
337 assert_eq!(host.h, 3);
338 assert_eq!(host.e_feat.nrows(), 3);
339 assert_eq!(host.b_feat, vec![0.0, 0.0, 0.0]);
340
341 assert_eq!(host.e_feat[(0, 0)], 10.0);
343 assert_eq!(host.e_feat[(0, 2)], 12.0);
344 assert_eq!(host.e_feat[(1, 0)], 1.0);
346 assert_eq!(host.e_feat[(2, 1)], 5.0);
348 }
349
350 #[test]
354 fn a_source_name_map_is_applied_before_matching_and_kept_in_src_names() {
355 let dir = tempfile::tempdir().unwrap();
356 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
357 let src = DMatrix::<f32>::from_row_slice(2, 2, &[1.0, 2.0, 3.0, 4.0]);
358 write_test_parquet(
359 &dict_path,
360 &["TGFB1", "MYC/count/unspliced"],
361 "gene",
362 &["h0", "h1"],
363 &src,
364 );
365 let target: Vec<Box<str>> = [
366 "ENSG_TGFB1/count/spliced",
367 "ENSG_MYC/count/spliced",
368 "ENSG_MYC/count/unspliced",
369 ]
370 .iter()
371 .map(|s| (*s).into())
372 .collect();
373 let lift = |n: &str| -> Box<str> {
374 if n.contains('/') {
375 n.into()
376 } else {
377 format!("{n}/count/spliced").into()
378 }
379 };
380 let host = load_frozen_feature_host(FrozenLoadArgs {
381 dictionary_path: &dict_path,
382 bias_path: None,
383 target_feature_names: &target,
384 name_kind: FeatureNameKind::Gene { delim: '_' },
385 source_name_map: Some(&lift),
386 })
387 .unwrap();
388 assert_eq!(host.keep_target_indices, vec![0, 2]);
389 assert_eq!(host.keep_src_indices, vec![0, 1]);
390 assert_eq!(
391 host.src_names,
392 vec![
393 Box::<str>::from("TGFB1/count/spliced"),
394 Box::<str>::from("MYC/count/unspliced")
395 ]
396 );
397 assert_eq!(host.e_feat[(0, 0)], 1.0);
398 assert_eq!(host.e_feat[(1, 1)], 4.0);
399 }
400
401 #[test]
402 fn gene_canon_matches_across_delim_variants() {
403 let dir = tempfile::tempdir().unwrap();
404 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
405
406 let src = DMatrix::<f32>::from_row_slice(2, 2, &[1.0, 2.0, 3.0, 4.0]);
408 write_test_parquet(
409 &dict_path,
410 &["ENSG00000105329_TGFB1", "ENSG00000141510_TP53"],
411 "gene",
412 &["h0", "h1"],
413 &src,
414 );
415 let target: Vec<Box<str>> = ["TP53", "TGFB1"].iter().map(|s| (*s).into()).collect();
416
417 let host = load_frozen_feature_host(FrozenLoadArgs {
418 dictionary_path: &dict_path,
419 bias_path: None,
420 target_feature_names: &target,
421 name_kind: FeatureNameKind::Gene { delim: '_' },
422 source_name_map: None,
423 })
424 .unwrap();
425
426 assert_eq!(host.keep_target_indices, vec![0, 1]);
427 assert_eq!(host.e_feat[(0, 0)], 3.0);
429 assert_eq!(host.e_feat[(1, 0)], 1.0);
431 }
432
433 #[test]
436 fn only_the_marked_rows_match() {
437 let dir = tempfile::tempdir().unwrap();
438 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
439 let src = DMatrix::<f32>::from_row_slice(3, 2, &[9.0, 9.0, 1.0, 2.0, 3.0, 4.0]);
440 write_test_parquet(
441 &dict_path,
442 &["CD4", "CD4", "MYC"],
443 "feature",
444 &["h0", "h1"],
445 &src,
446 );
447 let target: Vec<Box<str>> = ["MYC", "CD4"].iter().map(|s| (*s).into()).collect();
448 let args = || FrozenLoadArgs {
449 dictionary_path: &dict_path,
450 bias_path: None,
451 target_feature_names: &target,
452 name_kind: FeatureNameKind::Exact,
453 source_name_map: None,
454 };
455 let host =
456 load_frozen_feature_host_matching(args(), |_| Ok(vec![false, true, true])).unwrap();
457 assert_eq!(host.keep_target_indices, vec![0, 1]);
458 assert_eq!(host.keep_src_indices, vec![2, 1]);
459 assert_eq!(
460 host.e_feat.row(1).iter().copied().collect::<Vec<_>>(),
461 [1.0, 2.0]
462 );
463 assert_eq!(host.src_names.len(), 3);
464
465 let err = |m: Vec<bool>| {
467 load_frozen_feature_host_matching(args(), |_| Ok(m))
468 .err()
469 .unwrap()
470 .to_string()
471 };
472 assert!(err(vec![false; 3]).contains("none of its 3 rows is marked"));
473 assert!(err(vec![true]).contains("1 row flags for 3 rows"));
474
475 let lift = |n: &str| -> Box<str> { format!("{n}/count/spliced").into() };
477 let mut seen: Vec<Box<str>> = Vec::new();
478 let renamed = FrozenLoadArgs {
479 source_name_map: Some(&lift),
480 ..args()
481 };
482 assert!(load_frozen_feature_host_matching(renamed, |names| {
483 seen = names.to_vec();
484 Ok(vec![false; names.len()])
485 })
486 .is_err());
487 let read: Vec<Box<str>> = vec!["CD4".into(), "CD4".into(), "MYC".into()];
488 assert_eq!(seen, read);
489
490 let target: Vec<Box<str>> = vec!["CD4/count/spliced".into()];
492 let host = load_frozen_feature_host_matching(
493 FrozenLoadArgs {
494 target_feature_names: &target,
495 source_name_map: Some(&lift),
496 ..args()
497 },
498 |_| Ok(vec![false, true, true]),
499 )
500 .unwrap();
501 assert_eq!(&*host.src_names[0], "CD4");
502 assert_eq!(&*host.src_names[1], "CD4/count/spliced");
503 assert_eq!(host.keep_src_indices, vec![1]);
504 }
505
506 #[test]
507 fn empty_intersection_errors() {
508 let dir = tempfile::tempdir().unwrap();
509 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
510 let src = DMatrix::<f32>::from_row_slice(2, 2, &[1.0, 2.0, 3.0, 4.0]);
511 write_test_parquet(&dict_path, &["A", "B"], "gene", &["h0", "h1"], &src);
512 let target: Vec<Box<str>> = ["C", "D"].iter().map(|s| (*s).into()).collect();
513 let result = load_frozen_feature_host(FrozenLoadArgs {
514 dictionary_path: &dict_path,
515 bias_path: None,
516 target_feature_names: &target,
517 name_kind: FeatureNameKind::Exact,
518 source_name_map: None,
519 });
520 let err = match result {
521 Ok(_) => panic!("expected empty-intersection error"),
522 Err(e) => e,
523 };
524 assert!(err.to_string().contains("No feature names matched"));
525 }
526
527 #[test]
528 fn bias_loaded_when_provided() {
529 let dir = tempfile::tempdir().unwrap();
530 let dict_path = dir.path().join("d.parquet").to_str().unwrap().to_string();
531 let bias_path = dir.path().join("b.parquet").to_str().unwrap().to_string();
532 let src = DMatrix::<f32>::from_row_slice(2, 2, &[1.0, 2.0, 3.0, 4.0]);
533 write_test_parquet(&dict_path, &["A", "B"], "gene", &["h0", "h1"], &src);
534 let bias = DMatrix::<f32>::from_row_slice(2, 1, &[0.5, -0.3]);
535 write_test_parquet(&bias_path, &["A", "B"], "gene", &["bias"], &bias);
536
537 let target: Vec<Box<str>> = ["B", "A"].iter().map(|s| (*s).into()).collect();
538 let host = load_frozen_feature_host(FrozenLoadArgs {
539 dictionary_path: &dict_path,
540 bias_path: Some(&bias_path),
541 target_feature_names: &target,
542 name_kind: FeatureNameKind::Exact,
543 source_name_map: None,
544 })
545 .unwrap();
546 assert_eq!(host.b_feat, vec![-0.3, 0.5]);
548 }
549}