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
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.
| 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.
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).
AMR-LM integrates a mathematically rigorous interpretability module using Attention Rollout:
-
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$ ). -
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}$ ). - Heatmap Visualization: Highlights point mutations and specific motifs in the DNA string that caused the classification (e.g., conservative active sites of beta-lactamases).
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
- Python 3.8+ (Python 3.10 recommended)
- CUDA GPU (Highly recommended for training; runs CPU-only for inference)
- Clone the repository:
git clone https://github.com/amithgowda-m/AMR-LM.git cd AMR-LM - Install Python dependencies:
pip install -r requirements.txt
The project includes a Flask web interface that displays real-time predictions and the single-nucleotide XAI attention heatmap:
python app.pyOpen http://127.0.0.1:5000 in your browser.
To evaluate the model's throughput, specificity, and OOD performance:
python scripts/benchmark_validation.pyThis script updates the reports and outputs the validation plots:
results/figures/fig6_real_sample_abundance.pngresults/figures/fig7_ood_false_positive_rates.png
Windows:
run_all.batLinux / WSL:
bash run_all.shTo build and package the web application in a lightweight container:
- Build the Docker image:
docker build -t amr-lm-app . - Run the container (mounting the
modelsfolder to load checkpoints):docker run -p 5000:5000 -v $(pwd)/models:/app/models amr-lm-app
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
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}
}This project is for academic and research purposes. See individual database licenses (CARD, MEGARes, SARG) for data usage terms.