Skip to content
Open
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
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,10 @@

## NEW FUNCTIONALITY

* Add new method: MMContext, which is a pre-trained multi-modal embedding model, that maps free text annotations and count data into a shared latent space and is trained with contrastive learning. For now included only as a zeroshot method.

## NEW FUNCTIONALITY

* Add new method: CellMapper, which is a k-NN based approach to map cells across representations and can be used for label projection. Two versions are included here, one based on PCA or CCA embeddings (`linear`) and one based on an scvi embedding (`scvi`) (PR #22)

## MAJOR CHANGES
Expand Down
76 changes: 76 additions & 0 deletions src/methods/mmcontext_zeroshot/config.vsh.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
__merge__: ../../api/base_method.yaml

name: mmcontext_zeroshot
label: mmcontext (zero shot)
summary: "Zero-shot reference mapping using a pretrained mmcontext model."
description: |
mmcontext is a pretrained model for single-cell context modeling. In this zero-shot
variant, a pretrained mmcontext SentenceTransformer model is used to embed cells and
assign cell-type labels in a zero-shot fashion, given a set of label names from the
test data. The training data is not used in this method.

references:
doi: 10.64898/2025.12.08.692934

links:
documentation: https://github.com/mengerj/mmcontext/blob/main/README.md
repository: https://github.com/mengerj/mmcontext

info:
preferred_normalization: counts

arguments:
- name: --model_id
type: string
description: >
Hugging Face model ID for the pretrained mmcontext SentenceTransformer model.
If not provided, a default model ID will be used.
required: false
- name: --batch_size
type: integer
default: 64
description: Batch size used for mmcontext inference.
- name: --cache_dir
type: string
description: >
Directory used by mmcontext and adata_hf_datasets to cache intermediate
files and model artifacts.
required: true
- name: --initial_embedding
type: string
default: gs10k
description: >
Name of the initial embedding method used by InitialEmbedder (e.g. gs10k).
- name: --n_top
type: integer
default: 5
description: >
Number of top candidate labels for mmcontext to consider when annotating
cells; the best scoring label will be used as the prediction.
- name: --force_model
type: boolean
default: false
description: >
Skip the sanity-check that `--initial_embedding` is contained within `--model_id`.
Only use this if you are sure the initial embedder matches the one used to train
the mmcontext model checkpoint.

resources:
- type: python_script
path: script.py
- path: /src/utils/exit_codes.py

engines:
- type: docker
image: openproblems/base_python:1
setup:
- type: docker
run: |
pip install "mmcontext @ git+https://github.com/mengerj/mmcontext.git@main"

runners:
- type: executable
- type: nextflow
directives:
label: [midtime, midmem, midcpu]

175 changes: 175 additions & 0 deletions src/methods/mmcontext_zeroshot/script.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
import sys
import pathlib
import time

import anndata as ad
import numpy as np
import scanpy as sc
import torch
from adata_hf_datasets import InitialEmbedder
from adata_hf_datasets.dataset import AnnDataSetConstructor
from mmcontext.embed.model_utils import prepare_model_and_embed
from mmcontext.eval import OmicsQueryAnnotator
from sentence_transformers import SentenceTransformer

## VIASH START
par = {
"input_train": "resources_test/task_label_projection/cxg_immune_cell_atlas/train.h5ad",
"input_test": "resources_test/task_label_projection/cxg_immune_cell_atlas/test.h5ad",
"output": "output.h5ad",
# mmcontext-specific parameters
"model_id": None,
"batch_size": 64,
"cache_dir": "mmcontext_cache",
"initial_embedding": "gs10k",
"n_top": 5,
"force_model": False,
}
meta = {
"name": "mmcontext_zeroshot",
# "resources_dir" will be injected by Viash when built; default here for local testing
"resources_dir": "target/executable/methods/mmcontext_zeroshot",
}
## VIASH END

sys.path.append(meta["resources_dir"])



def compute_mmcontext_embeddings(
adata_path: str,
adata: ad.AnnData,
cache_dir: str,
model_id: str,
initial_embedding: str,
) -> np.ndarray:
"""
Compute mmcontext embeddings for all cells in `adata` using the same
pipeline as your external script, adapted for normalized-then-log data.
"""
cache_dir_path = pathlib.Path(cache_dir)
cache_dir_path.mkdir(parents=True, exist_ok=True)

# Initial embeddings (e.g. gs10k)
emb_key = f"X_{initial_embedding}"
print("=== Computing initial embeddings ===", flush=True)
embedder = InitialEmbedder(method=initial_embedding)
embedder.prepare(adata)
embeddings = embedder.embed(adata)
adata.obsm[emb_key] = embeddings

# Persist processed AnnData so AnnDataSetConstructor can link to it
processed_path = adata_path.replace(".h5ad", "_pp.h5ad")
print(f"=== Writing processed AnnData to {processed_path} ===", flush=True)
adata.write_h5ad(processed_path)

print("=== Building AnnData dataset for mmcontext ===", flush=True)
adata.obs["sample_idx"] = adata.obs.index
constructor = AnnDataSetConstructor(dataset_format="single")
constructor.add_anndata(
adata,
sentence_keys=["sample_idx"],
adata_link=processed_path,
)
ds = constructor.get_dataset()

print("=== Loading MMContext model ===", flush=True)
model = SentenceTransformer(model_id, trust_remote_code=True)

print("=== Computing MMContext embeddings ===", flush=True)
emb_df, _ = prepare_model_and_embed(
model,
ds,
index_col="sample_idx",
adata_download_dir=str(cache_dir_path),
main_col="cell_sentence_1",
layer_key=emb_key,
)

# Ensure embeddings align with adata.obs index order
emb_df_sorted = emb_df.set_index("sample_idx").reindex(
list(adata.obs.index)
).reset_index()
embedding_matrix = np.vstack(emb_df_sorted["embedding"].to_numpy())
adata.obsm["mmcontext_emb"] = embedding_matrix
return embedding_matrix


print("Reading input files", flush=True)
input_train = ad.read_h5ad(par["input_train"])
input_test = ad.read_h5ad(par["input_test"])

# Example organism check; adjust or remove depending on mmcontext’s capabilities.
#if input_train.uns.get("dataset_organism") not in (None, "homo_sapiens"):
# exit_non_applicable(
# "mmcontext_zeroshot currently only supports human data "
# f'(dataset_organism == "{input_train.uns["dataset_organism"]}")'
# )

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Device: '{device}'", flush=True)

print("Preparing test data for mmcontext", flush=True)
if "counts" in input_test.layers:
input_test.X = input_test.layers["counts"]

# Updated preprocessing: normalize first, then log-transform
print("=== Preprocessing test data (normalize -> log1p) ===", flush=True)
sc.pp.normalize_total(input_test)
sc.pp.log1p(input_test)

# Determine model ID
default_model_id = "jo-mengr/mmcontext-pubmedbert-gs10k"
model_id = par.get("model_id") or default_model_id
print(f"Using mmcontext model ID: {model_id}", flush=True)

# Quick safety check: initial embedding should match the checkpoint naming.
if (not par.get("force_model", False)) and (par["initial_embedding"] not in model_id):
raise ValueError(
"The configured initial embedding method does not appear to match the mmcontext model.\n"
f"- initial_embedding: '{par['initial_embedding']}'\n"
f"- model_id: '{model_id}'\n\n"
"The initial embedding method needs to match the one that was used to train the model.\n"
"If you are absolutely sure they match, re-run with `--force_model true` to skip this check."
)

start_time = time.time()

embeddings = compute_mmcontext_embeddings(
adata_path=par["input_test"],
adata=input_test,
cache_dir=par["cache_dir"],
model_id=model_id,
initial_embedding=par["initial_embedding"],
)

print("=== Creating annotator and predicting labels ===", flush=True)
annotator_model = SentenceTransformer(model_id, trust_remote_code=True)
annotator = OmicsQueryAnnotator(annotator_model)

# Use unique labels from the test data as the zero-shot label set
label_space = input_test.obs["label"].unique()
annotator.annotate_omics_data(
input_test,
labels=label_space,
emb_key="mmcontext_emb",
n_top=par["n_top"],
)

# OmicsQueryAnnotator writes the best label to obs["best_label"]
input_test.obs["label_pred"] = input_test.obs["best_label"]
input_test.obs["label_pred"] = input_test.obs["label_pred"].astype("category")

elapsed_time = time.time() - start_time
print(f"mmcontext zero-shot prediction completed in {elapsed_time:.2f} seconds", flush=True)

print("Write output AnnData to file", flush=True)
output = ad.AnnData(
obs=input_test.obs[["label_pred"]],
uns={
"method_id": meta["name"],
"dataset_id": input_test.uns["dataset_id"],
"normalization_id": input_test.uns["normalization_id"],
},
)
output.write_h5ad(par["output"], compression="gzip")