From b84af7addbf0a9464d75190217abec3073eb1b5b Mon Sep 17 00:00:00 2001 From: jkobject Date: Wed, 20 May 2026 08:24:53 +0000 Subject: [PATCH 1/2] fix(embedder): keep_all_cls_pred with heterogeneous head sizes Three related bugs reported in cantinilab/scPRINT#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 cantinilab/scPRINT#16. Co-authored-by: Prachi-Priyam --- scprint/model/model.py | 67 ++++++++++++++++++++++----------------- scprint/tasks/cell_emb.py | 21 ++++++++---- 2 files changed, 52 insertions(+), 36 deletions(-) diff --git a/scprint/model/model.py b/scprint/model/model.py index f02df39..82ac532 100644 --- a/scprint/model/model.py +++ b/scprint/model/model.py @@ -1374,11 +1374,15 @@ def on_validation_epoch_end(self): """@see pl.LightningModule""" self.embs = self.all_gather(self.embs).view(-1, self.embs.shape[-1]) self.info = self.all_gather(self.info).view(-1, self.info.shape[-1]) - self.pred = ( - self.all_gather(self.pred).view(-1, self.pred.shape[-1]) - if self.pred is not None - else None - ) + if self.pred is None: + pass + elif isinstance(self.pred, dict): + self.pred = { + k: self.all_gather(v).view(-1, v.shape[-1]) + for k, v in self.pred.items() + } + else: + self.pred = self.all_gather(self.pred).view(-1, self.pred.shape[-1]) self.pos = self.all_gather(self.pos).view(-1, self.pos.shape[-1]) if self.trainer.state.stage != "sanity_check": if self.trainer.is_global_zero: @@ -1568,20 +1572,23 @@ def _predict( if self.embs is None: self.embs = torch.mean(cell_embs[:, ind, :], dim=1) # self.embs = output["cls_output_" + "cell_type_ontology_term_id"] - self.pred = ( - torch.stack( + if len(self.classes) == 0: + self.pred = None + elif self.keep_all_cls_pred: + # Heads have different n_classes (e.g. 424 vs 62), so a + # stacked tensor is not well-defined. Store per-class logits + # in a dict so downstream code can concat per head. + self.pred = { + clsname: output["cls_output_" + clsname] + for clsname in self.classes + } + else: + self.pred = torch.stack( [ - ( - torch.argmax(output["cls_output_" + clsname], dim=1) - if not self.keep_all_cls_pred - else output["cls_output_" + clsname] - ) + torch.argmax(output["cls_output_" + clsname], dim=1) for clsname in self.classes ] ).transpose(0, 1) - if len(self.classes) > 0 - else None - ) self.pos = gene_pos self.expr_pred = ( [output["mean"], output["disp"], output["zero_logits"]] @@ -1593,25 +1600,27 @@ def _predict( # [self.embs, output["cls_output_" + "cell_type_ontology_term_id"]] [self.embs, torch.mean(cell_embs[:, ind, :], dim=1)] ) - self.pred = torch.cat( - [ - self.pred, - ( + if len(self.classes) == 0: + pass # keep self.pred = None + elif self.keep_all_cls_pred: + for clsname in self.classes: + self.pred[clsname] = torch.cat( + [self.pred[clsname], output["cls_output_" + clsname]] + ) + else: + self.pred = torch.cat( + [ + self.pred, torch.stack( [ - ( - torch.argmax(output["cls_output_" + clsname], dim=1) - if not self.keep_all_cls_pred - else output["cls_output_" + clsname] + torch.argmax( + output["cls_output_" + clsname], dim=1 ) for clsname in self.classes ] - ).transpose(0, 1) - if len(self.classes) > 0 - else None - ), - ], - ) + ).transpose(0, 1), + ], + ) self.pos = torch.cat([self.pos, gene_pos]) self.expr_pred = ( [ diff --git a/scprint/tasks/cell_emb.py b/scprint/tasks/cell_emb.py index a3e12fb..0dbffee 100644 --- a/scprint/tasks/cell_emb.py +++ b/scprint/tasks/cell_emb.py @@ -200,15 +200,22 @@ def __call__(self, model: torch.nn.Module, adata: AnnData, cache=False): pred_adata.obs.index = adata.obs.index adata.obs = pd.concat([adata.obs, pred_adata.obs], axis=1) 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 = [] for cl in model.classes: n = model.label_counts[cl] - columns += [model.label_decoders[cl][i] for i in range(n)] - allclspred = pd.DataFrame( - allclspred, columns=columns, index=adata.obs.index - ) - adata.obs = pd.concat(adata.obs, allclspred) + 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 + ) + ) + allclspred = pd.concat(dfs, axis=1) + adata.obs = pd.concat([adata.obs, allclspred], axis=1) metrics = {} if self.doclass and not self.keep_all_cls_pred: From df997eb14d236ab3a02c4d9f2841f406595884a0 Mon Sep 17 00:00:00 2001 From: jkobject Date: Wed, 20 May 2026 09:35:42 +0000 Subject: [PATCH 2/2] fix: address Codex/Copilot review on PR #52 - 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 '__' so 424 + 62 multi-head probabilities can't accidentally collide and don't shadow the standard pred_ 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. --- scprint/model/model.py | 8 +++++++- scprint/tasks/cell_emb.py | 30 +++++++++++++++++++++++------- tests/test_base.py | 33 +++++++++++++++++++++++++++++++++ 3 files changed, 63 insertions(+), 8 deletions(-) diff --git a/scprint/model/model.py b/scprint/model/model.py index 82ac532..feb7a40 100644 --- a/scprint/model/model.py +++ b/scprint/model/model.py @@ -1712,13 +1712,19 @@ def log_adata(self, gtclass=None, name=""): mdir = "data/" if not os.path.exists(mdir): os.makedirs(mdir) + # When keep_all_cls_pred=True, self.pred is a dict of per-head logits + # of different sizes. utils.make_adata expects an int-tensor of + # argmaxes (one column per class) and would crash on a dict, so we + # pass None: the full per-head logits are written into adata.obs by + # the Embedder caller (see scprint.tasks.cell_emb.Embedder.__call__). + pred_for_make = None if isinstance(self.pred, dict) else self.pred adata, fig = utils.make_adata( pos=self.pos, expr_pred=self.expr_pred, genes=self.genes, embs=self.embs, classes=self.classes, - pred=self.pred, + pred=pred_for_make, attention=self.attn.get(), label_decoders=self.label_decoders, labels_hierarchy=self.labels_hierarchy, diff --git a/scprint/tasks/cell_emb.py b/scprint/tasks/cell_emb.py index 0dbffee..c3f8cd5 100644 --- a/scprint/tasks/cell_emb.py +++ b/scprint/tasks/cell_emb.py @@ -199,23 +199,39 @@ def __call__(self, model: torch.nn.Module, adata: AnnData, cache=False): pred_adata.obs.index = adata.obs.index adata.obs = pd.concat([adata.obs, pred_adata.obs], axis=1) - if self.keep_all_cls_pred: + if self.keep_all_cls_pred and model.classes and model.pred is not None: # model.pred is a dict[clsname -> tensor[n_cells, n_classes_cl]] # (heads have different n_classes), so concatenate per head. + # Note: these are raw logits straight from the classification + # heads (ClsDecoder), not softmax probabilities. Column names + # are prefixed with the class id so they don't collide with + # the standard pred_ argmax columns. dfs = [] for cl in model.classes: + if cl not in model.pred: + continue n = model.label_counts[cl] - columns = [model.label_decoders[cl][i] for i in range(n)] + columns = [ + f"{cl}__{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( + if tensor.shape[0] != len(adata.obs.index): + # _predict resets self.pred to None when max_size_in_mem + # is exceeded (chunked logging). In that case the buffer + # only has the last chunk -- align by the trailing rows. + sliced_index = adata.obs.index[-tensor.shape[0]:] + df = pd.DataFrame(tensor, columns=columns, index=sliced_index) + df = df.reindex(adata.obs.index) + else: + df = pd.DataFrame( tensor, columns=columns, index=adata.obs.index ) - ) - allclspred = pd.concat(dfs, axis=1) - adata.obs = pd.concat([adata.obs, allclspred], axis=1) + dfs.append(df) + if dfs: + allclspred = pd.concat(dfs, axis=1) + adata.obs = pd.concat([adata.obs, allclspred], axis=1) metrics = {} if self.doclass and not self.keep_all_cls_pred: diff --git a/tests/test_base.py b/tests/test_base.py index ed8647a..e492fc4 100644 --- a/tests/test_base.py +++ b/tests/test_base.py @@ -98,6 +98,39 @@ def test_base(): col.startswith("pred_") for col in adata_emb.obs.columns ), "Classification failed" + # Regression test for #16: keep_all_cls_pred=True with heads of + # different n_classes used to raise + # RuntimeError: stack expects each tensor to be equal size + # because the per-head logits were stacked along a new dim. The fix + # stores them in a dict[clsname -> tensor] and the Embedder writes + # one column per (class, label) into adata.obs. + cell_embedder_all = Embedder( + batch_size=2, + num_workers=1, + how="random expr", + max_len=300, + doclass=True, + pred_embedding=[ + "cell_type_ontology_term_id", + "disease_ontology_term_id", + ], + doplot=False, + keep_all_cls_pred=True, + dtype=torch.float32, + ) + adata_emb_all, _ = cell_embedder_all(model, adata[:6, :]) + # Per-head probability/logit columns are emitted with the + # __