Skip to main content

data_beans/aux/
frozen_features.rs

1//! Load a pre-trained per-gene embedding (and optional per-gene bias)
2//! from parquet and strictly intersect its row axis against a caller's
3//! target feature axis.
4//!
5//! Used by `senna gbe / topic / cell-embedded-topic` to freeze the
6//! gene-side parameter table (`E_feat` in gbe, ρ in the ETM topic models)
7//! so cells train on a shared, pre-fit gene-relation space.
8//!
9//! Source formats supported via [`FrozenLoadArgs`]:
10//! - **gbe**: `{prefix}.dictionary.parquet` + `{prefix}.feature_bias.parquet`
11//!   (gene × H plus gene × 1 bias).
12//! - **topic / cell-embedded-topic**: `{prefix}.feature_embedding.parquet`
13//!   alone — bias defaults to zeros, which is what the topic models use
14//!   internally (no per-gene additive bias on ρ).
15//!
16//! Name resolution goes through [`FeatureNameKind`] so `TGFB1` and
17//! `ENSG00000105329_TGFB1` resolve to the same row.
18
19use crate::aux::feature_names::FeatureNameKind;
20use legume_numeric::matrix::traits::IoOps;
21use nalgebra::DMatrix;
22use rustc_hash::{FxHashMap, FxHashSet};
23
24/// Loaded + aligned frozen feature side ready to hand off to a candle
25/// engine (`graph-embedding-util` or one of the topic-model encoders).
26///
27/// Rows are reordered to follow the *target* axis. `keep_target_indices`
28/// records which positions in the caller's target feature axis survived
29/// the intersection — the caller MUST restrict its data (triplets,
30/// encoder D, decoder β) to these indices, otherwise the row order
31/// disagrees with the embedding rows.
32pub struct FrozenFeatureHost {
33    /// `[|keep|, H]`, rows in the same order as `keep_target_indices`.
34    pub e_feat: DMatrix<f32>,
35    /// `[|keep|]`. Zeros when no `bias_path` was given.
36    pub b_feat: Vec<f32>,
37    /// Indices into the *target* feature axis that matched a source row.
38    /// Length equals `e_feat.nrows()`.
39    pub keep_target_indices: Vec<usize>,
40    /// The *source* (dictionary) row each kept target row came from, parallel to
41    /// `keep_target_indices` — what a caller needs to look up any other table
42    /// keyed on the dictionary's rows (a module membership, say).
43    pub keep_src_indices: Vec<usize>,
44    /// The whole source table `[n_src, H]` and its row names, as read — so a
45    /// caller that needs the unmatched rows too (to place a gene by its
46    /// neighbours' rows) does not decode the file a second time.
47    pub src_e_feat: DMatrix<f32>,
48    pub src_names: Vec<Box<str>>,
49    /// Rows in the dictionary file, i.e. how many features the MODEL has.
50    ///
51    /// The only field that survives the intersection unfiltered, and the reason
52    /// it exists: `e_feat` and `keep_target_indices` are both already restricted
53    /// to the matched features, so a coverage fraction built from them is
54    /// identically 1 and tells a caller nothing. This is the denominator.
55    pub n_src: usize,
56    pub h: usize,
57}
58
59/// A rename of source row names, see [`FrozenLoadArgs::source_name_map`].
60pub type SourceNameMap<'a> = &'a dyn Fn(&str) -> Box<str>;
61
62pub struct FrozenLoadArgs<'a> {
63    /// Path to the `[D_src, H]` parquet (gbe `dictionary.parquet` or
64    /// topic `feature_embedding.parquet`). Row column 0 is the gene name.
65    pub dictionary_path: &'a str,
66    /// Optional path to a `[D_src, 1]` per-gene bias parquet (gbe
67    /// `feature_bias.parquet`). `None` → bias filled with zeros, which
68    /// matches the topic models' implicit "no per-gene bias on ρ".
69    pub bias_path: Option<&'a str>,
70    /// Caller's feature axis (e.g. `unified.feature_names` for gbe;
71    /// the topic models' `gene_names`). Output rows follow this order
72    /// after dropping unmatched entries.
73    pub target_feature_names: &'a [Box<str>],
74    /// Per-name canonicalization rule applied to both source and target
75    /// names before intersection. [`FeatureNameKind::Exact`] for strict
76    /// matching; [`FeatureNameKind::Gene { delim: '_' }`] is the typical
77    /// choice for scRNA gene IDs.
78    pub name_kind: FeatureNameKind,
79    /// Applied to each SOURCE row name that may match (every row, unless
80    /// [`load_frozen_feature_host_matching`] marks fewer) before
81    /// canonicalization, and kept as
82    /// the host's `src_names`: how a caller whose axis carries a row grammar
83    /// (`{gene}/count/spliced`) reads a plain gene table, lifting each bare
84    /// name into the grammar first. `None` = the names as read.
85    pub source_name_map: Option<SourceNameMap<'a>>,
86}
87
88/// Load the dictionary and match its rows to the target axis by canonical
89/// name. Among source rows with one canonical name the first wins.
90pub 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
94/// [`load_frozen_feature_host`], matching only the source rows `matchable`
95/// marks. It is handed the dictionary's row names as read (before
96/// [`FrozenLoadArgs::source_name_map`]) and returns one flag per row, BY
97/// POSITION: two rows may share a name (a cell type `CD4` beside the gene
98/// `CD4`), so marking by name would mark both and the first would still win.
99/// [`crate::aux::feature_types::feature_rows`] marks a mixed-type table's
100/// gene and region rows from its types table. A row left unmarked is never
101/// matched or renamed, yet stays in `src_names` / `src_e_feat`.
102pub 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    // Only a row that may match is renamed: an unmarked one keeps its name.
155    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        // First occurrence wins (see `load_frozen_feature_host`); `insert`
172        // would keep the last.
173        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    // A dictionary is a plain gene table. Source rows carrying the channelized
220    // row grammar ({gene}/{modality}/... ) that matched nothing usually mean
221    // the caller fed a channelized or co-embedding artifact; a PARTIAL match
222    // would otherwise proceed silently on the plain-name subset.
223    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        // Source: 4 genes × H=3. Source row "ENSG_DROP" has no target match.
302        let src = DMatrix::<f32>::from_row_slice(
303            4,
304            3,
305            &[
306                1.0, 2.0, 3.0, // TGFB1
307                4.0, 5.0, 6.0, // MYC
308                7.0, 8.0, 9.0, // ENSG_DROP (unmatched)
309                10.0, 11.0, 12.0, // TP53
310            ],
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        // Target: 5 genes; "FOO" and "BAR" don't appear in source.
321        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        // Kept target indices = positions of TP53, TGFB1, MYC in target order.
336        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        // Row 0 of e_feat should be source row for TP53 (= source row 3).
342        assert_eq!(host.e_feat[(0, 0)], 10.0);
343        assert_eq!(host.e_feat[(0, 2)], 12.0);
344        // Row 1: TGFB1 → source row 0.
345        assert_eq!(host.e_feat[(1, 0)], 1.0);
346        // Row 2: MYC → source row 1.
347        assert_eq!(host.e_feat[(2, 1)], 5.0);
348    }
349
350    /// A source of bare gene names read onto an axis that carries the row
351    /// grammar: the map lifts each source name into the grammar before the
352    /// canonical match, and the host reports the lifted names.
353    #[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        // Source uses ENSG-prefixed; target uses bare gene symbols.
407        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        // Row 0 (target TP53) ← source row 1.
428        assert_eq!(host.e_feat[(0, 0)], 3.0);
429        // Row 1 (target TGFB1) ← source row 0.
430        assert_eq!(host.e_feat[(1, 0)], 1.0);
431    }
432
433    /// An unmarked row is never matched, though it comes first and shares
434    /// the gene's name, and stays in the source table.
435    #[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        // Every row unmarked, or the wrong count: refused.
466        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        // The marks are asked of the names as read, before any rename.
476        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        // ...and only a marked row is renamed.
491        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        // Row 0 of output = target "B" = source row 1.
547        assert_eq!(host.b_feat, vec![-0.3, 0.5]);
548    }
549}