Skip to content

keep_all_cls_pred=True: persist per-chunk logits across max_size_in_mem boundary #54

Description

@jkobject

Follow-up from review of #52.

scPrint._predict resets self.pred to None after every max_size_in_mem-sized chunk (so the GPU buffer doesn't grow unbounded), and log_adata() writes the current chunk to disk as a numbered predict_part_<n>.h5ad.

For the standard keep_all_cls_pred=False argmax path, that's fine: each part is self-contained and the argmaxes live in adata.obs.

For keep_all_cls_pred=True, the new dict[clsname → tensor] only ever contains the last buffer at the moment Embedder.__call__ writes columns into adata.obs. #52 added a defensive trailing-rows alignment so we don't silently corrupt, but the proper fix is to persist the per-head logits into each predict_part_<n>.h5ad (or accumulate them on CPU) and re-merge in the Embedder caller.

Probably wants to:

  • pass model.pred (when dict) through log_adata and write it into adata.obsm["cls_logits_<clsname>"] or similar in each part,
  • have Embedder.__call__ concat across all written parts at the end instead of relying on the live buffer.

Not urgent (the common case is one chunk), but worth tracking.

Refs #16, #52.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions