Skip to content

Repository files navigation

🧬 AMR-LM

Cross-Database Generalization and Interpretability of DNA Language Models for Antimicrobial Resistance Gene Detection in Short Metagenomic Reads

Results • Interpretability (XAI) • Pipeline • Quick Start • Docker • Citation


Overview

AMR-LM fine-tunes DNABERT-2 (117M parameters) as a binary classifier to detect antimicrobial resistance (AMR) genes directly from DNA sequences. The model is trained on the CARD database and rigorously evaluated for:

  • Cross-database generalization → Validated on external MEGARes & SARG databases.
  • Short-read robustness → 150bp, 300bp, and 500bp simulated metagenomic fragment analysis.
  • Data leakage prevention → Rigorous k-mer overlap deduplication ($k=31$, threshold = 90%) between training and test sets.
  • Interpretability (XAI) → Single-nucleotide resolution attention map visualization using recursive Attention Rollout.
  • Metagenomic & OOD Validation → Profiling against host background (human DNA) and viral sequences to test specificity and out-of-distribution (OOD) safety.

The pipeline also benchmarks against gold-standard alignment-based tools (RGI) and deep learning models (DeepARG) to generate publication-ready tables and figures.


Key Results

1. Classification Benchmarks (150bp Fragments)

Model F1 (CARD) MCC (CARD) F1 (MEGARes) F1 (SARG) Inference Throughput
AMR-LM (Ours) 0.954 0.919 0.990 0.880 ~1.4–3.3 seqs/sec (CPU)
RGI (Alignment) 0.582 0.532 0.454 0.381 ~20.9 seqs/sec (CPU)
DeepARG (CNN) 0.534 0.484 0.363 0.349 ~24.8 seqs/sec (CPU)

Note: On short metagenomic reads (150bp), alignment-based tools fail due to incomplete alignments. AMR-LM leverages semantic representations to retain high F1 scores even on highly fragmented inputs.

2. Out-of-Distribution (OOD) Specificity

When evaluated against 80 out-of-distribution background sequences (consisting of human protein-coding exons and viral genomes):

  • False Positive Rate: 0.0% (0/80 false positives).
  • Mean AMR Probability: 0.264 (well below the 0.5 classification threshold, showing no confident hallucinations on host genomic backgrounds).

Interpretability (XAI)

AMR-LM integrates a mathematically rigorous interpretability module using Attention Rollout:

  1. Rollout Calculation: Instead of using naive last-layer attention, we recursively multiply the raw attention matrices across all 12 transformer layers ($R_l = A_l^{raw} \times R_{l-1}$) while accounting for residual connections ($A_l^{raw} = 0.5 \cdot I + 0.5 \cdot A_l$).
  2. BPE-to-Nucleotide Distribution: DNABERT-2 tokenizes sequences using Byte Pair Encoding (BPE). To resolve attention at single-nucleotide resolution, each subword token's rollout attention score is divided uniformly across its constituent bases ($\text{Score} / \text{Length}$).
  3. Heatmap Visualization: Highlights point mutations and specific motifs in the DNA string that caused the classification (e.g., conservative active sites of beta-lactamases).

Pipeline Overview

The pipeline consists of 9 sequential steps:

1. setup.py                          → Setup directory structure & download CARD
2. download_model.py                 → Locally cache pre-trained DNABERT-2 snapshots
3. download_negatives_and_test_dbs.py → Fetch MEGARes, SARG, and background sequences
4. preprocess.py                     → Deduplicate, perform k-mer filtering, and split data
5. run_baselines.py                  → Run RGI & DeepARG benchmarks
6. train_dnabert2.py                 → Fine-tune DNABERT-2 (supports LoRA adapters)
7. evaluate.py                       → Evaluate model on test sets across fragment sizes
8. benchmark_validation.py           → Run OOD checks, measure throughput, & save validation plots
9. generate_paper_figures/tables.py  → Compile LaTeX tables and 300 DPI publication curves

Quick Start

Prerequisites

  • Python 3.8+ (Python 3.10 recommended)
  • CUDA GPU (Highly recommended for training; runs CPU-only for inference)

Installation

  1. Clone the repository:
    git clone https://github.com/amithgowda-m/AMR-LM.git
    cd AMR-LM
  2. Install Python dependencies:
    pip install -r requirements.txt

Running the Web UI

The project includes a Flask web interface that displays real-time predictions and the single-nucleotide XAI attention heatmap:

python app.py

Open http://127.0.0.1:5000 in your browser.

Running Metagenomics & OOD Validation

To evaluate the model's throughput, specificity, and OOD performance:

python scripts/benchmark_validation.py

This script updates the reports and outputs the validation plots:

  • results/figures/fig6_real_sample_abundance.png
  • results/figures/fig7_ood_false_positive_rates.png

Running the Full Pipeline

Windows:

run_all.bat

Linux / WSL:

bash run_all.sh

Docker Deployment

To build and package the web application in a lightweight container:

  1. Build the Docker image:
    docker build -t amr-lm-app .
  2. Run the container (mounting the models folder to load checkpoints):
    docker run -p 5000:5000 -v $(pwd)/models:/app/models amr-lm-app

Project Structure

AMR-LM/
├── app.py                           # Flask web application server
├── database.py                      # SQLite database for prediction history
├── predictor.py                     # AMRPredictor class (Rollout attention & BPE mapping)
├── Dockerfile                       # Multi-stage optimized Docker file
├── requirements.txt                 # Python dependencies
├── run_all.bat                      # Windows pipeline runner
├── run_all.sh                       # Linux pipeline runner
├── scripts/
│   ├── setup.py                     # Directory structure & raw CARD setup
│   ├── download_model.py            # Local model cacher
│   ├── train_dnabert2.py            # Fine-tuning loop (supports FP16 and LoRA)
│   ├── dnabert2_loader.py           # PEFT-safe loader
│   ├── benchmark_validation.py      # Metagenomics & OOD validation benchmark
│   └── evaluate.py                  # Evaluation suite
├── models/                          # Model checkpoints (git-ignored)
├── data/                            # Processed data splits (git-ignored)
└── results/                         # Generated report logs, LaTeX tables, and plots
    ├── figures/                     # 300 DPI publication-ready PNG curves
    └── tables/                      # LaTeX output tables

Citation

If you use this pipeline or tool in your research, please cite:

@article{zhou2023dnabert2,
  title={DNABERT-2: Efficient Foundation Model and Benchmark For Multi-Species Genome},
  author={Zhou, Zhihan and Ji, Yanrong and Li, Weijian and Dutta, Pratik and Davuluri, Ramana and Liu, Han},
  journal={arXiv preprint arXiv:2306.15006},
  year={2023}
}

@article{alcock2023card,
  title={CARD 2023: expanded curation, support for machine learning, and resistome prediction at the Comprehensive Antibiotic Resistance Database},
  author={Alcock, Brian P and others},
  journal={Nucleic Acids Research},
  volume={51},
  number={D1},
  pages={D419--D428},
  year={2023}
}

License

This project is for academic and research purposes. See individual database licenses (CARD, MEGARes, SARG) for data usage terms.

About

Cross-database generalization for Antimicrobial Resistance (ARG) detection in short metagenomic reads using DNA Language Models

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages