README.md
July 26, 2026 · View on GitHub
Quickstart | Installation | Reference docs | Paper | NeurIPS 2020
Molecular dynamics is a workhorse of modern computational condensed matter physics. It is frequently used to simulate materials to observe how small scale interactions can give rise to complex large-scale phenomenology. Most molecular dynamics packages (e.g. HOOMD Blue or LAMMPS) are complicated, specialized pieces of code that are many thousands of lines long. They typically involve significant code duplication to allow for running simulations on CPU and GPU. Additionally, large amounts of code is often devoted to taking derivatives of quantities to compute functions of interest (e.g. gradients of energies to compute forces).
However, recent work in machine learning has led to significant software developments that might make it possible to write more concise molecular dynamics simulations that offer a range of benefits. Here we target JAX, which allows us to write python code that gets compiled to XLA and allows us to run on CPU, GPU, or TPU. Moreover, JAX allows us to take derivatives of python code. Thus, not only is this molecular dynamics simulation automatically hardware accelerated, it is also end-to-end differentiable. This should allow for some interesting experiments that we're excited to explore.
JAX, MD is a research project that is currently under development. Expect sharp edges and possibly some API breaking changes as we continue to support a broader set of simulations. JAX MD is a functional and data driven library. Data is stored in arrays or tuples of arrays and functions transform data from one state to another.
Getting Started
For a video introducing JAX MD along with a demo, check out this talk from the Physics meets Machine Learning series:
To get started playing around with JAX MD check out the following colab notebooks on Google Cloud without needing to install anything. For a very simple introduction, I would recommend the Minimization example. For an example of a bunch of the features of JAX MD, check out the JAX MD cookbook.
- JAX MD Cookbook
- Custom Potentials
- Flocking
- Meta Optimization
- Swap Monte Carlo (Cargese Summer School)
- Implicit Differentiation
- Athermal Linear Elasticity
- Smash a Sand Castle
JAX MD also comes with self contained python scripts which you run locally if you have JAX MD installed:
- Fire minimization
- NVE Simulation
- NVT Simulation
- NPT Simulation
- NVE with Neighbor Lists
- Neural Network Potentials
See FEATURES.md for a tour of the main library components.
Installation
With uv
We recommend using uv for local development:
git clone https://github.com/jax-md/jax-md
cd jax-md
uv sync --no-default-groups
To include testing and documentation dependencies:
uv sync --no-default-groups --group testing --group docs
With pip
If you want to use pip instead, install the latest published release with:
python -m pip install jax-md --upgrade
Or install a local source checkout:
python -m pip install -e .
To include testing and documentation dependencies with pip, using the --group flag from pip 25.1 or newer:
python -m pip install -e . --group testing --group docs
Development
JAX MD is under active development. Please don't hesitate to open feature requests to help us guide development. We more than welcome contributions! See CONTRIBUTING.md for how to get set up.
Tests
Tests live in tests/ and run in double precision:
Run the full suite with uv:
JAX_ENABLE_X64=1 uv run --no-sync pytest
Run a specific test suite with uv:
JAX_ENABLE_X64=1 uv run --no-sync pytest tests/<suite>_test.py
Without uv, install the testing dependencies and run pytest directly:
python -m pip install -e . --group testing
JAX_ENABLE_X64=1 python -m pytest
JAX_ENABLE_X64=1 python -m pytest tests/<suite>_test.py
Run the suites affected by your change. See CONTRIBUTING.md for the full development setup.
Technical gotchas
GPU
You must follow JAX's GPU installation instructions to enable GPU support.
64-bit precision
To enable 64-bit precision, set the respective JAX flag before importing jax_md (see the JAX guide), for example:
import jax
jax.config.update("jax_enable_x64", True)
Publications
JAX MD has been used in the following publications. If you don't see your paper on the list, but you used JAX MD let us know and we'll add it to the list!
- Molecular Simulations with a Pretrained Neural Network and Universal Pairwise Force Fields. (J. Am. Chem. Soc. 2025)
A. Kabylda, J. T. Frank, S. Suárez-Dou, A. Khabibrakhmanov, L. Medrano Sandonas, O. T. Unke, S. Chmiela, K.-R. Müller, and A. Tkatchenko - Designing precise dynamical steady states in disordered networks. (Machine Learning: Science and Technology (2025))
M. Berneman and D. Hexner - Generalized design of sequence-ensemble-function relationships for intrinsically disordered proteins
R. K. Krueger, M. P. Brenner, and K. Shrinivas - Tuning colloidal reactions. (PRL 2024)
R. K. Krueger, E. M. King, and M. P. Brenner - Programming patchy particles for materials assembly design. (PNAS 2024)
E. M. King, CX. Du, QZ. Zhu, S. S. Schoenholz, and M. P. Brenner - LATTE: an atomic environment descriptor based on Cartesian tensor contractions. (arXiv 2024)
F. Pellegrini, S. Gironcoli, E. Küçükbenli - PySAGES: flexible, advanced sampling methods accelerated with GPUs. (npj Computational Materials 2024)
P. F. Zubieta Rico, et al. - Scaling deep learning for materials discovery (Nature 2023)
A. Merchant, et al. - LapTrack: linear assignment particle tracking with tunable metrics. (Bioinformatics 2023)
Yohsuke T Fukai and Kyogo Kawaguchi - A Differentiable Neural-Network Force Field for Ionic Liquids. (J. Chem. Inf. Model. 2022)
H. Montes-Campos, J. Carrete, S. Bichelmaier, L. M. Varela, and G. K. H. Madsen - Correlation Tracking: Using simulations to interpolate highly correlated particle tracks. (Phys. Rev. E. 2022)
E. M. King, Z. Wang, D. A. Weitz, F. Spaepen, and M. P. Brenner - Optimal Control of Nonequilibrium Systems Through Automatic Differentiation.
M. C. Engel, J. A. Smith, and M. P. Brenner - Graph Neural Networks Accelerated Molecular Dynamics. (J. Chem. Phys. 2022)
Z. Li, K. Meidani, P. Yadav, and A. B. Farimani - Gradients are Not All You Need.
L. Metz, C. D. Freeman, S. S. Schoenholz, and T. Kachman - Lagrangian Neural Network with Differential Symmetries and Relational Inductive Bias.
R. Bhattoo, S. Ranu, and N. M. A. Krishnan - Efficient and Modular Implicit Differentiation.
M. Blondel, Q. Berthet, M. Cuturi, R. Frostig, S. Hoyer, F. Llinares-López, F. Pedregosa, and J.-P. Vert - Learning neural network potentials from experimental data via Differentiable Trajectory Reweighting.
(Nature Communications 2021)
S. Thaler and J. Zavadlav - Learn2Hop: Learned Optimization on Rough Landscapes. (ICML 2021)
A. Merchant, L. Metz, S. S. Schoenholz, and E. D. Cubuk - Designing self-assembling kinetics with differentiable statistical physics models. (PNAS 2021)
C. P. Goodrich, E. M. King, S. S. Schoenholz, E. D. Cubuk, and M. P. Brenner
Citation
If you use the code in a publication, please cite the repo using the .bib,
@inproceedings{jaxmd2020,
author = {Schoenholz, Samuel S. and Cubuk, Ekin D.},
booktitle = {Advances in Neural Information Processing Systems},
publisher = {Curran Associates, Inc.},
title = {JAX M.D. A Framework for Differentiable Physics},
url = {https://papers.nips.cc/paper/2020/file/83d3d4b6c9579515e1679aca8cbc8033-Paper.pdf},
volume = {33},
year = {2020}
}
If you use functionalities related to RigidBody, please cite the following paper using the .bib,
@article{king2024programming,
title={Programming patchy particles for materials assembly design},
author={King, Ella M and Du, Chrisy Xiyu and Zhu, Qian-Ze and Schoenholz, Samuel S and Brenner, Michael P},
journal={Proceedings of the National Academy of Sciences},
volume={121},
number={27},
pages={e2311891121},
year={2024},
publisher={National Academy of Sciences}
}
