Skip to content

Fix #16: keep_all_cls_pred with heterogeneous head sizes - #52

Merged
jkobject merged 2 commits into
mainfrom
fix/keep-all-cls-pred-cantini
May 20, 2026
Merged

jkobject merged 2 commits into
mainfrom
fix/keep-all-cls-pred-cantini

Conversation

@jkobject

Copy link
Copy Markdown
Collaborator

Fixes #16. Three related bugs in the Embedder path when keep_all_cls_pred=True with multiple labels in pred_embedding (e.g. cell_type_ontology_term_id + disease_ontology_term_id).

Bug 1 — torch.stack with mismatched head sizes

scprint/model/model.py:_predict stacked per-class logits along a new dim, but heads have different n_classes (424 vs 62 in the report), so:

RuntimeError: stack expects each tensor to be equal size,
but got [32, 424] at entry 0 and [32, 62] at entry 1

Fix: when keep_all_cls_pred=True, store per-head logits in a dict[clsname -> tensor] and concatenate per head across batches. The argmax branch (keep_all_cls_pred=False) is unchanged.

on_validation_epoch_end is also updated to all_gather the dict shape.

Bug 2 — CUDA tensor in pd.DataFrame

scprint/tasks/cell_emb.py wrapped a CUDA tensor in pd.DataFrame →

TypeError: can't convert cuda:0 device type tensor to numpy.

Fix: move to CPU + numpy before the DataFrame call.

Bug 3 — pd.concat positional args

adata.obs = pd.concat(adata.obs, allclspred) is a positional bug — the second arg is interpreted as axis.
Fix: pd.concat([adata.obs, allclspred], axis=1).

Output shape note

With keep_all_cls_pred=True, all per-head probability columns are concatenated into adata.obs (one column per class, named via label_decoders). For 424 + 62 classes this is ~486 columns — that's intentional but worth flagging.

Test plan

  • AST-checked.
  • Functional smoke test requires a checkpoint + GPU; would appreciate @Prachi-Priyam re-running their reproducer with this branch installed:
    pip install git+https://github.com/cantinilab/scPRINT.git@fix/keep-all-cls-pred-cantini

CC @Prachi-Priyam — thanks for the very detailed report.

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>

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 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".

Comment thread scprint/model/model.py
Comment on lines +1581 to +1584
self.pred = {
clsname: output["cls_output_" + clsname]
for clsname in self.classes
}

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge 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 👍 / 👎.

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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; update on_validation_epoch_end to all_gather dict values.
  • In the Embedder task, convert tensors to CPU/NumPy before building a pd.DataFrame, and fix pd.concat usage 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.

Comment thread scprint/model/model.py
Comment on lines +1581 to +1584
self.pred = {
clsname: output["cls_output_" + clsname]
for clsname in self.classes
}
Comment thread scprint/tasks/cell_emb.py Outdated
Comment on lines +217 to +218
allclspred = pd.concat(dfs, axis=1)
adata.obs = pd.concat([adata.obs, allclspred], axis=1)
Comment thread scprint/tasks/cell_emb.py Outdated
Comment on lines +209 to +213
tensor = model.pred[cl]
if hasattr(tensor, "detach"):
tensor = tensor.detach().cpu().numpy()
dfs.append(
pd.DataFrame(
Comment thread scprint/tasks/cell_emb.py Outdated
Comment on lines +202 to +205
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 = []
Comment thread scprint/tasks/cell_emb.py Outdated
Comment on lines +208 to +215
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-commenter

codecov-commenter commented May 20, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 67.64706% with 11 lines in your changes missing coverage. Please review.
✅ Project coverage is 41.51%. Comparing base (491f640) to head (df997eb).

Files with missing lines Patch % Lines
scprint/model/model.py 58.82% 7 Missing ⚠️
scprint/tasks/cell_emb.py 76.47% 4 Missing ⚠️
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.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

- 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.
@jkobject

Copy link
Copy Markdown
Collaborator Author

Thanks @chatgpt-codex-connector & @copilot — all five points addressed in df997eb:

Codex P1 / Copilot #1 — make_adata crash on dict. Real bug, you're right. log_adata now passes None into utils.make_adata when self.pred is a dict; make_adata only ever receives the argmax tensor or None. The dict logits keep flowing to adata.obs via the Embedder caller, which is the only place that knows the per-head label decoders anyway. (Kept self.pred as the dict rather than splitting into a second attribute, because on_validation_epoch_end also needs to all_gather it and that pathway is now explicit.)

Copilot #2 — empty pd.concat. Guarded the whole block with if self.keep_all_cls_pred and model.classes and model.pred is not None and an inner if dfs before pd.concat. No-classes models won't crash.

Copilot #3 — logits vs probabilities. You're right, ClsDecoder returns raw logits. I left them raw (softmax loses information and we have no monotonic-preference reason here) but updated the inline comment, and prefixed the columns with <clsname>__ so they're clearly distinct from the standard pred_<cls> argmax columns and don't collide across heads.

Copilot #4 — missing regression test. Added one in tests/test_base.py: it runs the existing fixture model with keep_all_cls_pred=True over two heads of different n_classes (cell_type_ontology_term_id + disease_ontology_term_id) and asserts the per-class columns are added without raising. Reuses the same small CKPT + h5ad, so no extra fixtures or CI cost.

Copilot #5 — chunked max_size_in_mem reset. Good catch. For keep_all_cls_pred=True the right long-term answer is to persist per-chunk logits into the h5ad parts and re-merge, which is a bigger refactor (touches the per-part writer in log_adata). For now I added a smaller defensive path: when tensor.shape[0] < adata.n_obs (i.e. buffer was reset mid-run), I align by the trailing rows and reindex with NaN-fill, so we get a partial but correct result instead of a misleading silent corruption. I'll open a follow-up issue to track the proper chunk-aware persistence.

@chatgpt-codex-connector

Copy link
Copy Markdown

To use Codex here, create an environment for this repo.

@jkobject

Copy link
Copy Markdown
Collaborator Author

Re #5 (chunked persistence): follow-up tracked at #54.

@jkobject
jkobject merged commit 2e8a011 into main May 20, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Openning an issue for Jalil's questions

3 participants