mgxlens-v2 / README.md
vedantM's picture
Upload folder using huggingface_hub
78b2584 verified
|
Raw
History Blame Contribute Delete
10.6 kB
# mgxLens
A retrieval-based metagenomic taxonomic profiler. Given short DNA reads (~150 bp) from a
sequencing run, it reports the estimated **relative abundance** of microbial taxa at the
**genus** and **species** level.
**Reference coverage.** The index covers **13,399 distinct species-level SGBs across 3,297 genera**
— the full MetaPhlAn vJan25 SGB reference. Every indexed SGB resolves to a named Linnaean species
(no unnamed placeholder bins among the *indexed* organisms; the naming rate over indexed SGBs is
100%). This means "closed-world" is about organisms *outside MetaPhlAn's reference entirely* — not
gaps within it. An organism that is in MetaPhlAn's reference is representable; one that is not will
be misassigned to its nearest reference genus (see Caveats → closed-world).
---
## 1. What's in the bundle
| File | What it is |
|---|---|
| `encoder.pt` | The neural encoder (seqLens-89M backbone + attention pooling → 256-d embedding). Turns a DNA string into a vector. |
| `index.faiss` | A FAISS vector index of ~2.32M reference marker embeddings. Cosine search via inner product. |
| `index.clades.npy` | Row-aligned to the FAISS index: the SGB (clade) id for each indexed vector. |
| `index.markers.npy` | Row-aligned marker source ids (provenance; not needed for inference). |
| `index.config.json` | Index build parameters (window=150, nwin=5, emb_dim=256, max_length=128). |
| `lineage.tsv` | `clade_id → genus, species, display_name`. Turns a retrieved SGB into a taxon name (display_name is MetaPhlAn-style). |
| `mgx_encoder.py` | The encoder class + `load_encoder()` / `embed()` helpers. Import this; don't reimplement. |
| `MANIFEST.json` | Machine-readable summary + file checksums + caveats. |
**Encoder and index are from the same training run** (`index.config.json.source_ckpt` ==
`encoder.pt`). Do not mix this encoder with any other index or vice-versa — the vectors would be
meaningless.
---
## 2. Environment
```bash
pip install torch transformers faiss-cpu numpy
# GPU optional: faiss-gpu + a CUDA torch build speed up large batches, but faiss-cpu is fine for
# moderate throughput. The encoder runs on CPU or GPU; GPU is ~10-50x faster for embedding.
```
The encoder downloads its base model (`omicseye/seqLens_4096_512_89M-at-base-multi`) from
HuggingFace on first load via `transformers`, so the serving host needs network access on first
run (or pre-cache the HF model). `trust_remote_code=True` is required for that base model.
---
## 3. How it works (the inference recipe)
For each input read:
1. **Window** the read to 150 bp (if longer, take a centered 150 bp window; if ≤150, use as-is).
2. **Embed** it with the encoder → a 256-d L2-normalized vector.
3. **Search** the FAISS index for the top-k nearest reference vectors (k=25). Because vectors are
L2-normalized and the index is inner-product, the score **is** cosine similarity ∈ [-1, 1].
4. **Gate**: if the top-1 cosine is below a threshold Ï„ (start with **0.66**), the read is
*abstained* (dropped, counted as "unclassified"). Otherwise it's assigned.
5. **Assign**: the read's taxon is the clade (`index.clades.npy`) of its top-1 hit. Map that clade
to genus/species via `lineage.tsv`.
6. **Aggregate** across all assigned reads: count reads per genus (and per species), normalize to
fractions → the abundance profile. Report the unclassified fraction separately.
That's the whole pipeline. It's nearest-neighbor retrieval, not a trained classifier.
---
## 4. Reference implementation (copy-paste starting point)
```python
import numpy as np, faiss, torch
from collections import defaultdict
from mgx_encoder import load_encoder, embed # ships in this bundle
BUNDLE = "." # path to the bundle dir
K = 25 # neighbors per read
GATE = 0.66 # top-1 cosine threshold; below -> abstain
WINDOW = 150
# ---- load once at startup ----
device = "cuda" if torch.cuda.is_available() else "cpu"
model, tok, cfg = load_encoder(f"{BUNDLE}/encoder.pt", device)
index = faiss.read_index(f"{BUNDLE}/index.faiss")
clades = np.load(f"{BUNDLE}/index.clades.npy", allow_pickle=True) # (N,) SGB id per vector
lineage = {} # clade_id -> (genus, species, display_name)
with open(f"{BUNDLE}/lineage.tsv") as fh:
next(fh)
for line in fh:
cid, g, s, disp = line.rstrip("\n").split("\t")
lineage[cid] = (g, s, disp)
def center_window(seq, w=WINDOW):
if len(seq) <= w:
return seq
s = (len(seq) - w) // 2
return seq[s:s+w]
def profile(reads, batch_size=4096):
"""reads: list of DNA strings. Returns genus/species abundance + unclassified fraction.
Species keys are MetaPhlAn-style display names ('Escherichia coli', or
'Escherichia sp. (SGB123)' for unnamed genome bins)."""
genus_counts, species_counts = defaultdict(float), defaultdict(float)
n_total = len(reads); n_assigned = 0
for i in range(0, n_total, batch_size):
batch = [center_window(r) for r in reads[i:i+batch_size]]
ml = cfg["max_length"] if "max_length" in cfg else 128
Z = embed(model, tok, batch, device, max_length=ml).astype("float32")
D, I = index.search(Z, K) # D=cosine sims, I=index rows
for r in range(len(batch)):
if D[r, 0] < GATE: # gate: abstain
continue
clade = str(clades[I[r, 0]]) # top-1 hit's clade
g, s, disp = lineage.get(clade, ("", "", ""))
if not g:
continue
n_assigned += 1
genus_counts[g] += 1
if disp:
species_counts[disp] += 1 # display name, not raw s__ / SGB id
def norm(d):
tot = sum(d.values()) or 1.0
return {k: v/tot for k, v in sorted(d.items(), key=lambda x: -x[1])}
return {
"genus": norm(genus_counts),
"species": norm(species_counts),
"n_reads": n_total,
"n_assigned": n_assigned,
"unclassified_fraction": 1.0 - (n_assigned / n_total if n_total else 0.0),
}
# ---- example ----
if __name__ == "__main__":
reads = ["ACGT..."] # your reads (from a FASTQ)
print(profile(reads))
```
**Reading FASTQ**: reads come from `.fastq`/`.fastq.gz`. Use `pysam`, `Bio.SeqIO`, or a
2nd/4th-line parser to get the sequence strings, then pass the list to `profile()`. Paired-end
R1/R2 can both be fed as independent reads.
---
## 5. Serving it (FastAPI sketch)
```python
from fastapi import FastAPI, UploadFile
app = FastAPI()
# model/index/lineage loaded once at module import (see section 4)
@app.post("/profile")
async def profile_endpoint(file: UploadFile):
reads = parse_fastq(await file.read()) # your FASTQ parser -> list[str]
return profile(reads) # JSON: {genus:{...}, species:{...}, ...}
```
Load the model, index, and lineage **once at startup** (they're large — encoder 342 MB, index
2.3 GB). Never reload per request. The index holds ~2.3 GB in RAM; size the host accordingly
(≥8 GB RAM recommended, ≥16 GB comfortable).
---
## 6. Input / output contract
**Input**: DNA reads as strings (from FASTQ), ~150 bp each. Non-ACGT characters are tolerated by
the tokenizer but degrade the embedding; upstream QC/trimming (e.g. fastp) is assumed.
**Output** (JSON):
```json
{
"genus": {"escherichia": 0.29, "pseudomonas": 0.24, "...": 0.0},
"species": {"Escherichia coli": 0.21, "Pseudomonas aeruginosa": 0.18, "...": 0.0},
"n_reads": 200000,
"n_assigned": 94459,
"unclassified_fraction": 0.528
}
```
Abundances are fractions summing to ~1.0 within each level (over assigned reads). All 13,399
indexed SGBs have real species names, so species keys are Linnaean names (the `Genus sp. (SGBxxxx)`
fallback only appears if the index is rebuilt to include unnamed bins).
---
## 7. Tunable knobs
| Knob | Default | Effect |
|---|---|---|
| `GATE` (top-1 cosine) | 0.66 | Higher = stricter, more reads abstained, fewer false assignments (but see Caveats — it does **not** reliably reject novel organisms). |
| `K` (neighbors) | 25 | Only top-1 is used for assignment here; k>1 matters if you switch to k-NN vote aggregation. |
| aggregation | top-1 vote | Simplest. Alternatives (k-NN vote, similarity-weighted) exist in the research code but top-1 is the documented default. |
---
## 8. Caveats — read before building product on this
- **Proof of concept.** Not accuracy-tuned. The goal of this handoff is to enable infrastructure
and integration work, not to ship a validated diagnostic.
- **Genus is the reliable level; species is approximate.** At 150 bp, reads from different species
in the same genus are near-indistinguishable to the encoder — the model reliably gets the
*genus* right but frequently assigns the wrong *species within that genus*. Report species
output, but do not treat it as trustworthy.
- **Species naming.** All 13,399 indexed SGBs have real Linnaean species names, so normal output
shows names like "Escherichia coli", not bin ids. (The `display_name` logic still falls back to
`Genus sp. (SGBxxxx)` for unnamed bins — this only matters if the index is ever rebuilt to
include the unnamed SGBs that exist elsewhere in the MetaPhlAn taxonomy; the shipped index has
none.)
- **Closed-world — this is the big one.** The model can only assign a read to a taxon that exists
in its reference index. An organism that is **not** in the reference will still be confidently
assigned to its nearest reference genus and reported as present. There is effectively **no
novelty rejection** (abstain rate ~0% in evaluation), and this behavior was **never quantified**.
**Do not deploy this against arbitrary/open real-world samples** and present the output as
complete or trustworthy without adding a novelty-rejection layer. It is validated only in the
closed-world setting where the sample's organisms are known to be in the reference.
- **Evaluated on one mock community (Zymo D6300).** Broader validation was not done.
---
## 9. Provenance
- Encoder: `enc_attention_bg1.best_genus.pt` — seqLens-89M backbone, attention pooling, 256-d
projection, trained with contrastive InfoNCE + background negatives, checkpoint selected on
real-Zymo genus accuracy.
- Index: `index_bakeoff_mw_notest` — leakage-fixed (held-out test markers excluded).
- Base model pulled at runtime: `omicseye/seqLens_4096_512_89M-at-base-multi`.
See `MANIFEST.json` for file checksums and the machine-readable summary.