Fix #16: keep_all_cls_pred with heterogeneous head sizes - #52
Conversation
Three related bugs reported in #16: 1. scPrint._predict: with keep_all_cls_pred=True, the code stacked the per-class logits with torch.stack, but classification heads have different n_classes (e.g. 424 for cell_type vs 62 for disease), so stacking raises 'stack expects each tensor to be equal size'. Store per-class logits in a dict {clsname: tensor} instead, and concatenate per head across batches. 2. on_validation_epoch_end: handle the dict shape when all_gather'ing self.pred. 3. Embedder.__call__ (cell_emb.py): - move logits to CPU/numpy before wrapping in pd.DataFrame (CUDA tensors cannot be converted directly). - fix pd.concat(adata.obs, allclspred) -> pd.concat([adata.obs, allclspred], axis=1) (positional misuse, second arg was being interpreted as 'axis'). The keep_all_cls_pred=False (argmax) path is unchanged. Refs #16. Co-authored-by: Prachi-Priyam <noreply@github.com>
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: b84af7addb
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
| self.pred = { | ||
| clsname: output["cls_output_" + clsname] | ||
| for clsname in self.classes | ||
| } |
There was a problem hiding this comment.
Keep
self.pred tensor-compatible before logging adata
When keep_all_cls_pred=True, this change stores self.pred as a dict of per-head logits, but model.log_adata() still passes pred into utils.make_adata, which unconditionally calls pred.to(...) (in scprint/model/utils.py). That causes an AttributeError as soon as prediction logging runs (e.g., Embedder.__call__ always calls model.log_adata), so the keep-all workflow now crashes before returning results.
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Pull request overview
Fixes failures in the embedding/prediction path when keep_all_cls_pred=True and multiple classification heads have different n_classes, and corrects pandas usage when appending per-class outputs into adata.obs.
Changes:
- Store per-head classification logits in a
dict[clsname -> tensor]instead of stacking mismatched head outputs; updateon_validation_epoch_endtoall_gatherdict values. - In the Embedder task, convert tensors to CPU/NumPy before building a
pd.DataFrame, and fixpd.concatusage to concatenate as columns.
Reviewed changes
Copilot reviewed 2 out of 2 changed files in this pull request and generated 5 comments.
| File | Description |
|---|---|
scprint/tasks/cell_emb.py |
Builds per-head DataFrames from model.pred and concatenates them into adata.obs, with CPU/NumPy conversion and corrected pd.concat call. |
scprint/model/model.py |
Changes _predict to store per-head logits in a dict when keeping all predictions; updates validation epoch-end gathering accordingly. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| self.pred = { | ||
| clsname: output["cls_output_" + clsname] | ||
| for clsname in self.classes | ||
| } |
| allclspred = pd.concat(dfs, axis=1) | ||
| adata.obs = pd.concat([adata.obs, allclspred], axis=1) |
| tensor = model.pred[cl] | ||
| if hasattr(tensor, "detach"): | ||
| tensor = tensor.detach().cpu().numpy() | ||
| dfs.append( | ||
| pd.DataFrame( |
| if self.keep_all_cls_pred: | ||
| allclspred = model.pred | ||
| columns = [] | ||
| # model.pred is a dict[clsname -> tensor[n_cells, n_classes_cl]] | ||
| # (heads have different n_classes), so concatenate per head. | ||
| dfs = [] |
| columns = [model.label_decoders[cl][i] for i in range(n)] | ||
| tensor = model.pred[cl] | ||
| if hasattr(tensor, "detach"): | ||
| tensor = tensor.detach().cpu().numpy() | ||
| dfs.append( | ||
| pd.DataFrame( | ||
| tensor, columns=columns, index=adata.obs.index | ||
| ) |
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #52 +/- ##
==========================================
+ Coverage 41.11% 41.51% +0.40%
==========================================
Files 27 27
Lines 2875 2900 +25
==========================================
+ Hits 1182 1204 +22
- Misses 1693 1696 +3 ☔ View full report in Codecov by Sentry. 🚀 New features to boost your workflow:
|
- log_adata: pass None into utils.make_adata when self.pred is a dict (keep_all_cls_pred=True). make_adata calls pred.to(...) and would crash on a dict. The dict logits are still written into adata.obs by Embedder.__call__. - Embedder: guard the keep_all_cls_pred block against empty model.classes / missing keys / pd.concat([]) and against chunked prediction where model.pred only contains the last buffer. - Embedder: prefix per-class columns with '<clsname>__' so 424 + 62 multi-head probabilities can't accidentally collide and don't shadow the standard pred_<cls> argmax columns. Renamed conceptually to 'raw logits' (ClsDecoder output is unnormalised) -- adjusted the inline comment to reflect that. - tests/test_base.py: add a regression test that runs Embedder with keep_all_cls_pred=True over two heads of different n_classes (cell_type + disease) and asserts the per-head columns are added without raising.
|
Thanks @chatgpt-codex-connector & @copilot — all five points addressed in df997eb: Codex P1 / Copilot #1 — Copilot #2 — empty Copilot #3 — logits vs probabilities. You're right, Copilot #4 — missing regression test. Added one in Copilot #5 — chunked |
|
To use Codex here, create an environment for this repo. |
Fixes #16. Three related bugs in the Embedder path when
keep_all_cls_pred=Truewith multiple labels inpred_embedding(e.g.cell_type_ontology_term_id+disease_ontology_term_id).Bug 1 —
torch.stackwith mismatched head sizesscprint/model/model.py:_predictstacked per-class logits along a new dim, but heads have differentn_classes(424 vs 62 in the report), so:Fix: when
keep_all_cls_pred=True, store per-head logits in adict[clsname -> tensor]and concatenate per head across batches. The argmax branch (keep_all_cls_pred=False) is unchanged.on_validation_epoch_endis also updated toall_gatherthe dict shape.Bug 2 — CUDA tensor in
pd.DataFramescprint/tasks/cell_emb.pywrapped a CUDA tensor inpd.DataFrame→Fix: move to CPU + numpy before the
DataFramecall.Bug 3 —
pd.concatpositional argsadata.obs = pd.concat(adata.obs, allclspred)is a positional bug — the second arg is interpreted asaxis.Fix:
pd.concat([adata.obs, allclspred], axis=1).Output shape note
With
keep_all_cls_pred=True, all per-head probability columns are concatenated intoadata.obs(one column per class, named vialabel_decoders). For 424 + 62 classes this is ~486 columns — that's intentional but worth flagging.Test plan
pip install git+https://github.com/cantinilab/scPRINT.git@fix/keep-all-cls-pred-cantiniCC @Prachi-Priyam — thanks for the very detailed report.