Predicting immunotherapy response from gut metagenomes with protein embeddings, vector search and random forests
A walkthrough of how more than 100 million microbial protein sequences became a searchable embedding space, and how species-level and functional features from gut metagenomes feed response models compared across cancer types and cohorts — with simplified code for each step.
Immune checkpoint inhibitors do not work equally for all patients, and response rates are limited. Predicting response before treatment, and finding the biomarkers related to it, can support clinical decisions and drug development, and the gut microbiome is one of the most promising places to look. Two things make it hard. Microbiome signatures found in one cohort often fail in the next, because populations carry different species. And the underlying data — protein sequences, metagenomes, microbiome profiles — is large enough that any representation of what the microbes can do has to be built and searched at scale.
In this post we walk through the approach developed with CJ Bioscience researchers for their biomarker platform. More than 100 million protein sequences were clustered into about 6 million families with MMseqs2 and Foldseek, represented with protein language-model embeddings (ESM2 and gLM were evaluated), and indexed in a Chroma DB vector store that searches about 6 million sequences in 1–2 seconds. Each sample's metagenome is described by species-level and functional features, and random forests, a microbiome taxonomic language model and ensembles predict responders, compared across cancer types and cohorts. The response model improved prediction accuracy by more than 15% over the prior machine-learning baseline, and the related study was presented as a SITC 2024 poster.
Solution overview
The work has two halves that share one structure. A representation path runs offline and at scale: protein sequences are clustered, cluster representatives are embedded, and the embeddings are indexed for search. A modelling path turns each sample's metagenome into species-level and functional features and trains response models that are scored on cohorts left out of training. The clusters serve both halves — as the unit of the search index and as functional units for the models.
100M+ microbial protein sequences
Collapse into ~6M families
Protein language-model vectors
Search ~6M in 1–2 s
Species and functional profiles
Response models across cohorts
The numbered steps in the diagram:
- More than 100 million protein sequences from gut metagenomes are the raw material for describing what the microbes can do.
- MMseqs2 clusters by sequence similarity and Foldseek by structure predicted with AlphaFold or ESMFold; the sequences collapse into about 6 million clusters.
- Cluster representatives are embedded with protein language models — ESM2 and gLM were evaluated — and embedding processing ran about 60× faster than the baseline.
- A Chroma DB vector store searches about 6 million sequences in 1–2 seconds; UMAP projections support exploration.
- Each sample is described by its species-level profile and by functional contig embeddings, so a function carried by different species can look the same to a model.
- Random forests, a microbiome taxonomic language model, BMCARP and ensembles predict responders and are compared across cancer types and cohorts.
Technology stack
| Layer | Technology | What it does here |
|---|---|---|
| Clustering | MMseqs2 (sequence) · Foldseek (structure) | 100M+ sequences → about 6M clusters |
| Structure | AlphaFold · ESMFold | Predicted structures for structural clustering and function inference |
| Embedding | Protein language models (ESM2 and gLM evaluated) | Fixed-length vectors for cluster representatives |
| Search | Chroma DB vector store · UMAP | About 6M sequences searchable in 1–2 s; 2-D maps for exploration |
| Features | Species-level profiles · functional contig embeddings | Per-sample inputs to the response models |
| Prediction | Random forest · Microbiome Taxonomic Language Model · BMCARP · ensembles | Responder vs non-responder |
| Validation | Cross-cohort comparison | Does a signature transfer to a population it has not seen? |
Step 1: Collapse 100 million sequences into protein families
Gut metagenomes contain the same proteins many times over: near-identical copies from related strains, and from sample after sample. Embedding and indexing every copy would spend most of the compute on redundancy. Clustering first turned more than 100 million sequences into about 6 million clusters — protein families that act as functional units — and only their representatives need to be embedded and searched.
Two clustering approaches were evaluated with the researchers. MMseqs2 groups sequences by similarity and has a linear-time mode for very large catalogues. Foldseek groups proteins by 3-D structure, here predicted with AlphaFold or ESMFold. Structure is conserved longer than sequence, so structural clustering can join distant relatives that sequence identity keeps apart, and predicted structures also support function inference.
# Sequence clustering with MMseqs2 (linear-time mode); thresholds are placeholders
mmseqs createdb proteins.faa catalogue_db
mmseqs linclust catalogue_db clusters_db tmp --min-seq-id "$MIN_ID" -c "$MIN_COV"
mmseqs createsubdb clusters_db catalogue_db reps_db # one representative per cluster
mmseqs convert2fasta reps_db representatives.faa
mmseqs createtsv catalogue_db catalogue_db clusters_db members.tsv
# Structure clustering with Foldseek over predicted structures (AlphaFold / ESMFold)
foldseek createdb predicted_structures/ struct_db
foldseek cluster struct_db struct_clusters tmp_fs -c "$MIN_COV"
foldseek createtsv struct_db struct_db struct_clusters struct_members.tsvSimplified. One way to run each tool; thresholds, resources and the production workflow are not shown.
Step 2: Embed representatives with a protein language model
Each representative becomes a fixed-length vector from a protein language model. ESM2 embeds a protein from its own sequence. gLM, a genomic language model, contextualizes each protein's embedding with its neighbours on the same contig, which carries functional information a single sequence lacks. The embeddings were validated for three uses: sequence classification, function inference and search.
h_1 … h_L = PLM(sequence) per-residue hidden states, h_l ∈ ℝ^d e = (1/L) · Σ_l h_l mean-pooled protein embedding sim(q, c) = cos(e_q, e_c) used for search and for grouping in embedding space
import numpy as np
import torch
from transformers import AutoTokenizer, EsmModel
tok = AutoTokenizer.from_pretrained(ESM2_CHECKPOINT) # one of the public ESM2 sizes
model = EsmModel.from_pretrained(ESM2_CHECKPOINT, torch_dtype=torch.float16).cuda().eval()
@torch.inference_mode()
def embed(seqs: list[str], batch_size: int = 64, max_len: int = 1024) -> np.ndarray:
order = sorted(range(len(seqs)), key=lambda i: len(seqs[i])) # less padding per batch
out = np.zeros((len(seqs), model.config.hidden_size), dtype=np.float32)
for k in range(0, len(order), batch_size):
idx = order[k:k + batch_size]
enc = tok([seqs[i] for i in idx], return_tensors="pt", padding=True,
truncation=True, max_length=max_len).to("cuda")
h = model(**enc).last_hidden_state # (batch, length, d)
mask = enc["attention_mask"].unsqueeze(-1).to(h.dtype)
out[idx] = ((h * mask).sum(1) / mask.sum(1)).float().cpu().numpy() # mean pool
return outSimplified. Batched, mean-pooled ESM2 embeddings with Hugging Face transformers; the checkpoint and the production pipeline's optimizations are not shown.
Embedding is the most expensive stage at this scale. Processing time came down by about 60× against the baseline, and because only cluster representatives are embedded, the job stays tractable as the catalogue grows.
Step 3: Index the embeddings for second-scale search
The representative embeddings are stored in a Chroma DB vector store, which answers nearest-neighbour queries with an approximate HNSW index. A query protein is embedded the same way and compared with about 6 million vectors in 1 to 2 seconds, fast enough for interactive lookups. Clustering is what keeps the index at about 6 million vectors rather than more than 100 million. UMAP projections of the embedding space were used for exploration.
import chromadb
import numpy as np
client = chromadb.PersistentClient(path=INDEX_DIR)
reps = client.get_or_create_collection("cluster_representatives",
metadata={"hnsw:space": "cosine"})
def add_representatives(ids: list[str], vectors: np.ndarray, sizes: list[int]):
step = client.get_max_batch_size()
for i in range(0, len(ids), step):
reps.add(ids=ids[i:i + step],
embeddings=vectors[i:i + step].tolist(),
metadatas=[{"members": int(n)} for n in sizes[i:i + step]])
def search(query_vec: np.ndarray, k: int = 10):
"""Nearest cluster representatives to one embedded query protein."""
res = reps.query(query_embeddings=[query_vec.tolist()], n_results=k)
return list(zip(res["ids"][0], res["distances"][0], res["metadatas"][0]))Simplified. Collection names and metadata are illustrative.
The live model below shows the same idea with an inverted-file index: the query is compared with k-means centroids first, then only with the members of the nearest few clusters.
compare with all cost(q) = N distances per query
cluster index cost(q) = K + Σ_{c ∈ P(q)} |members(c)| P(q) = the n_probe centroids nearest to q
live model N = 12,000, K = 200, n_probe = 4 → a few hundred distances per query
Step 4: Describe each sample by species and by function
A sample's metagenome gives two views. The species-level profile says who is there. The functional view says what those organisms can do — which protein families or contig-level functions the sample carries, and how abundantly. The SITC 2024 study combined species-level microbiome profiles with functional contig embeddings.
The reason for the second view is transfer. Different populations carry different species, but those species can carry the same functional genes. A responder-associated function carried by one species in one cohort and by another species elsewhere is invisible to a species-only model trained on the first cohort; a functional feature sees it in both.
import numpy as np
def clr(rel_abundance: np.ndarray, eps: float = 1e-5) -> np.ndarray:
"""Centred log-ratio of species relative abundances (samples x species)."""
logx = np.log(rel_abundance + eps)
return logx - logx.mean(axis=1, keepdims=True)
def functional_profile(rel_abundance: np.ndarray, copies: np.ndarray) -> np.ndarray:
"""Samples x clusters: abundance-weighted copies of each protein cluster
carried by the species present. `copies` is species x clusters."""
return np.log(rel_abundance @ copies + 1e-4)
def features(rel_abundance, copies, use_function: bool = True) -> np.ndarray:
X = clr(rel_abundance)
if use_function:
X = np.hstack([X, functional_profile(rel_abundance, copies)])
return XSimplified. This is the live model's construction — centred log-ratios for species and abundance-weighted cluster counts for function; the study's exact featurization is not shown.
Step 5: Train response models and compare them across cohorts
The SITC 2024 study defined responder versus non-responder prediction tasks on 942 samples from nine public metagenomic cohorts across cancer types, and compared random forests, a Microbiome Taxonomic Language Model, BMCARP and ensembles. Random forests are a strong baseline for this kind of data: they cope with many correlated, zero-heavy abundance features, need little tuning, and expose feature importances for candidate discovery.
How the models are validated matters more than which one wins. Cross-validation that mixes samples from every cohort lets a model learn each cohort's quirks — sequencing, population, response definition — and report them as signal. Holding out a whole cohort at a time is the stricter test, the one on which species-level signatures most often fail, and the one the live model below uses.
data (x_i, y_i, k_i), y_i = 1 for a responder, k_i ∈ {1, …, 9} the cohort
fold k f_(−k) = fit( { (x_i, y_i) : k_i ≠ k } )
score AUC_k = AUC( { f_(−k)(x_i), y_i : k_i = k } )
report AUC_1 … AUC_9 and their mean, never only one pooled score
live model f = ridge logistic regression: min_β −Σ_i log p(y_i | x_i, β) + (λ/2)·‖β‖²
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import roc_auc_score
from sklearn.model_selection import LeaveOneGroupOut
def loco(X, y, cohort, **rf_params):
"""Train on all cohorts but one, score the held-out cohort, repeat."""
per_cohort, models = {}, []
for train, test in LeaveOneGroupOut().split(X, y, groups=cohort):
rf = RandomForestClassifier(class_weight="balanced", n_jobs=-1,
random_state=0, **rf_params)
rf.fit(X[train], y[train])
held_out = cohort[test][0]
per_cohort[held_out] = roc_auc_score(y[test], rf.predict_proba(X[test])[:, 1])
models.append(rf)
return per_cohort, models
# Same folds, two feature sets: species only vs species + function
auc_species, _ = loco(features(P, copies, use_function=False), y, cohort)
auc_both, models = loco(features(P, copies), y, cohort)Simplified. Leave-one-cohort-out evaluation of a random forest with scikit-learn; hyperparameters are left to the caller.
Why report every held-out cohort? A mean AUC can hide a cohort where the model is no better than chance. The per-cohort scores show whether a signature transfers, which is the question a biomarker has to answer.
Step 6: Keep the candidates that hold up across cohorts
The structure was designed for biomarker discovery as well as prediction, and the study explored microbial and contig candidates associated with response. A feature that ranks among the most important in one fold may be a quirk of the cohorts it was trained on; one that ranks highly fold after fold is a candidate worth taking back to the laboratory.
from collections import Counter
import numpy as np
def stable_candidates(models, names, top: int = 5):
"""Count the folds in which each feature is among the most important."""
hits = Counter()
for rf in models: # one model per held-out cohort
best = np.argsort(rf.feature_importances_)[::-1][:top]
hits.update(names[i] for i in best)
n_folds = len(models)
return [(name, k, n_folds) for name, k in hits.most_common()]
candidates = [c for c in stable_candidates(models, feature_names) if c[1] >= MIN_FOLDS]Simplified. Fold-stability ranking over the models from Step 5; the cut-off MIN_FOLDS is a placeholder.
Importance is not mechanism. A stable feature says where to look; whether that protein family affects response needs experimental work.
Try the live model
The live model below rebuilds both halves on generated data: protein embeddings clustered with k-means and used as a search index, and response models trained on eight cohorts and tested on the ninth.
Live model, computed entirely in your browser on generated data the size of the public data in the SITC 2024 poster: 942 samples in nine cohorts, each with its own species mix, and a response that depends on a function carried by different species in different cohorts. Top: a species-only model and a model with protein-family clusters, each trained on eight cohorts and tested on the ninth. Bottom: 12,000 generated 16-dimensional protein embeddings grouped into 200 clusters by k-means and searched by probing the four nearest centroids. Ridge logistic regression and k-means stand in for the study's models and the production clustering and index; nothing is pre-computed. Open the live model on its own page ↗
Results
The microbiome-based response model improved prediction accuracy by more than 15% over the prior machine-learning baseline, and the project produced a structure for surfacing biomarker candidates rather than only scores. On the representation side, more than 100 million protein sequences were clustered into about 6 million clusters, embedding processing ran about 60× faster than the baseline, and the vector store searches about 6 million sequences in 1 to 2 seconds.
The related study — 942 samples from nine public metagenomic cohorts, comparing random forests, a microbiome taxonomic language model, BMCARP and ensembles — was presented as a poster at SITC 2024. The work also led to patent preparation, press materials and follow-up collaboration discussions with external research organizations, and helped CJ Bioscience and the AI Center register a formal collaborative KPI project.
Lessons learned
- Build the representation before the model. Clustering, embedding and indexing at scale came first; without them there are no functional features to model and no way to inspect candidates.
- Validate the way the model will be used. A biomarker has to work in the next hospital, so evaluation holds out whole cohorts instead of mixing them.
- Describe function as well as taxonomy. What microbes can do transfers across populations better than which species happen to be present.
- Design for discovery, not only a score. Features that stay strong across folds give researchers a short list to test.
- One structure, two jobs. The clusters that make 100 million sequences tractable also keep the search index small.
Conclusion
Predicting immunotherapy response from the gut microbiome is limited less by the classifier than by representation and validation. Clustering more than 100 million protein sequences into families, embedding them with protein language models, indexing them for second-scale search, and combining species and functional features in models tested on held-out cohorts produced a response model above the prior baseline and a path to biomarker candidates.
The same recipe — compress a huge biological catalogue into searchable families, describe samples by function, and validate across populations — applies to other microbiome questions where species differ between cohorts but function is shared.
Limitations
- Association is not mechanism: biomarker candidates from observational cohorts need experimental validation before any clinical use.
- Public cohorts differ in sequencing, processing and response definitions; part of any cross-cohort gap is technical rather than biological.
- The live model's species, proteins, embeddings and responses are generated with a planted functional signal, and ridge logistic regression and k-means stand in for the study's models and the production clustering and index — it shows the mechanism of the approach, not the size of the real effect.
About the demo and confidentiality
Everything in the embedded model is generated. No patient-level data, cohort, sequence, model or result from CJ Bioscience or the AI Center appears in this post beyond the figures stated in the author's CV and the public SITC 2024 poster; code is simplified and written for illustration.