Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 45 additions & 30 deletions scprint/model/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
}
Comment on lines +1581 to +1584

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

Comment on lines +1581 to +1584
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"]]
Expand All @@ -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 = (
[
Expand Down Expand Up @@ -1703,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,
Expand Down
39 changes: 31 additions & 8 deletions scprint/tasks/cell_emb.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,16 +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:
allclspred = model.pred
columns = []
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_<cls> 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)]
allclspred = pd.DataFrame(
allclspred, columns=columns, index=adata.obs.index
)
adata.obs = pd.concat(adata.obs, allclspred)
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()
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
)
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:
Expand Down
33 changes: 33 additions & 0 deletions tests/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# <class>__<label> prefix so they don't clash with the pred_* argmax.
all_cols = list(adata_emb_all.obs.columns)
assert any(c.startswith("cell_type_ontology_term_id__") for c in all_cols), (
"keep_all_cls_pred=True: missing per-class columns for cell_type"
)
assert any(c.startswith("disease_ontology_term_id__") for c in all_cols), (
"keep_all_cls_pred=True: missing per-class columns for disease"
)
# And no row was dropped.
assert adata_emb_all.n_obs == 6

# GRN inference
grn_inferer = GNInfer(
layer=[0, 1],
Expand Down
Loading