PM++: Multi-GPU Particle-Mesh Cosmology
PM++ is a JAX-based, differentiable particle-mesh cosmology code built on PMWD
ideas and extended for multi-GPU simulations. The active implementation is
imported as pmpp and lives in src/pmpp/; tests/pmwd/ retains the PMWD
reference implementation used exclusively for validation.
The documented baseline uses multiple GPUs so every example exercises the distributed ownership, mesh-halo, and FFT paths.
Installation
PM++ supports JAX 0.9.1 through 0.10 (declared as jax>=0.9.1,<0.11). Install
the JAX build for your accelerator and driver first by following the official
JAX installation guide.
On a computer with two visible GPUs:
python -m venv ~/.venvs/pmpp
source ~/.venvs/pmpp/bin/activate
python -m pip install --upgrade pip
python -m pip install pmpp jupyter
# Optional: compile accelerated routing for this machine. This requires nvcc.
pmpp-build-cuda-routing
# The checkout supplies the notebooks; PM++ itself remains pip-installed.
git clone https://github.com/rouzib/PMpp.git ~/PMpp
cd ~/PMpp
jupyter lab docs/source/notebooks
Current Scope
- Multi-GPU PM N-body simulation with JAX.
- Preferred
mesh_halomulti-GPU mode. - PMWD comparison tests for forward and gradient correctness.
- Distributed FFT support for sharded meshes.
- LPT, Boltzmann/growth utilities, scatter/gather, and power-spectrum tools.
- Potential-correction models under
src/pmpp/corrections/.
Repository Layout
PMpp/
|-- src/pmpp/ # Active importable PM++ package
| |-- core/ # Configuration and shared utilities
| |-- cosmology/ # Cosmology, transfer, and growth
| |-- initial_conditions/ # White noise, modes, and LPT
| |-- numerics/ # Local FFT and ODE primitives
| |-- distributed/ # Multi-GPU FFT, halos, and routing
| |-- cic/ # Scatter, gather, and Pallas CIC
| |-- nbody/ # Particles, gravity, integrator, observers
| |-- corrections/ # Optional correction models
| |-- analysis/ # Power spectra and plotting
| `-- extras/ # CAMELS and QUIJOTE adapters
|-- tests/ # Regression and gradient tests
| `-- pmwd/ # Test-only PMWD reference implementation
|-- docs/source/notebooks/ # Pre-executed documentation notebooks
`-- docs/ # Project documentation
Import through the feature packages shown above. The former flat module paths were removed as part of this architecture change.
Minimal Multi-GPU Setup
New code should use the nested MultiGPUConfiguration object. The older
top-level compute_mesh= compatibility path still exists, but is not preferred.
import jax
import jax.numpy as jnp
from pmpp import Configuration, MultiGPUConfiguration
from pmpp.distributed import create_compute_mesh
res = 256
box_size = 1000.0 # Mpc/h
ptcl_grid_shape = (res, res, res)
ptcl_spacing = box_size / res
gpu_devices = [device for device in jax.devices() if device.platform == "gpu"]
if len(gpu_devices) < 2:
raise RuntimeError("This multi-GPU example requires at least 2 GPUs.")
selected_devices = gpu_devices
compute_mesh = create_compute_mesh(selected_devices)
num_devices = len(selected_devices)
conf = Configuration(
ptcl_spacing,
ptcl_grid_shape,
mesh_shape=1,
multigpu=MultiGPUConfiguration(
compute_mesh=compute_mesh,
mode="mesh_halo",
),
max_ptcl_per_slice=int((res**3 / num_devices) * 1.8),
max_share_ptcl=50_000,
max_halo_share_ptcl=50_000,
max_share_gather_ptcl=200_000,
float_dtype=jnp.float32,
)
Capacity overflows are correctness failures. If a run reports overflow in particle migration, halo rebuild, or gather exchange buffers, increase the corresponding capacity and rerun.
Minimal Multi-GPU Forward Run
import jax
import jax.numpy as jnp
from pmpp import Configuration, MultiGPUConfiguration
from pmpp.cic import scatter
from pmpp.cosmology import SimpleLCDM, boltzmann
from pmpp.distributed import create_compute_mesh
from pmpp.initial_conditions import linear_modes, lpt, white_noise
from pmpp.nbody import nbody
res = 32
box_size = 100.0
gpu_devices = [device for device in jax.devices() if device.platform == "gpu"]
if len(gpu_devices) < 2:
raise RuntimeError("This PM++ simulation requires at least two GPUs")
selected_devices = gpu_devices
conf = Configuration(
box_size / res,
(res, res, res),
mesh_shape=1,
multigpu=MultiGPUConfiguration(
compute_mesh=create_compute_mesh(selected_devices),
mode="mesh_halo",
),
float_dtype=jnp.float32,
)
@jax.jit
def simulate(seed):
cosmo = boltzmann(SimpleLCDM(conf), conf)
noise = white_noise(seed, conf)
modes = linear_modes(noise, cosmo, conf)
particles = lpt(modes, cosmo, conf)
particles = nbody(particles, cosmo, conf)
return particles, scatter(particles, conf)
ptcl_final, density = simulate(0)
density.block_until_ready()
print(density.shape)
print(float(density.mean()))
Expected sanity checks:
- density shape matches the mesh;
- density mean is close to
1.0; - no capacity warnings appear.
Multi-GPU Modes
Prefer mesh_halo for current multi-GPU work:
- particles are stored authoritatively on their owning slab;
- particles migrate between slabs when needed;
- mesh halos are exchanged for local stencil operations;
- it is generally faster than the older particle-halo path for both smaller and larger simulation boxes.
particle_halo remains useful for comparison and legacy validation.
Performance defaults
mesh_halo always uses canonical sparse routing and packed migration
collectives. pallas_cic=True uses paired Pallas gather/scatter on qualified
float32 GPU setups; unsupported platforms warn and fall back to reference JAX.
CUDA routing is selected automatically when its optional FFI is qualified.
See the optimization guide for the
measured forward and AD recommendations.
Development
Install the development and documentation tools from an editable checkout:
python -m pip install -e ".[dev,docs]"
PM++ uses YAPF 0.43.0 with the project style defined in pyproject.toml.
Format the active package and maintained tests, then verify that no formatting
changes remain:
python -m yapf --in-place --recursive src tests
python -m yapf --diff --recursive src tests
See the contributor guide for the complete implementation, validation, and documentation workflow.
Testing
Focused gravity checks:
/home/rouzib/.virtualenvs/PMPP/bin/python -m pytest \
tests/test_grad_gravity.py \
tests/test_gravity_particle_nyquist_filter.py \
-q
Mesh-halo scatter/gather:
/home/rouzib/.virtualenvs/PMPP/bin/python -m pytest tests/test_mesh_halo_scatter_gather.py -q
End-to-end gradient:
/home/rouzib/.virtualenvs/PMPP/bin/python -m pytest tests/test_grad_nbody.py -q
Notebooks
The documentation gallery contains six pre-executed notebooks:
- first simulation and configuration
- resolution-consistent initial conditions evolved from $32^3$ through $256^3$
- a multi-GPU
mesh_halorun - observers and analysis
- differentiation with finite-difference checks.
Read the Docs renders committed outputs and does not execute the notebooks. Restart kernels after code changes. Re-run every notebook with all visible GPUs in a clean temporary copy before committing its outputs.
License
PM++ is distributed under the BSD-3-Clause license. See LICENSE.
PM++ is based on PMWD and retains the original PMWD BSD 3-Clause notice in
THIRD_PARTY_NOTICES.md. The test-only tests/pmwd/
package is kept as a reference implementation for validation.
Documentation build
Install the documentation extra and build the Sphinx site locally:
python -m pip install -e ".[docs]"
sphinx-build -W --keep-going -b html docs/source docs/build/html
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file pmpp-0.2.2.tar.gz.
File metadata
- Download URL: pmpp-0.2.2.tar.gz
- Upload date:
- Size: 182.8 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/7.0.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
9f0fe3b5f33c54eecddf76aa0427e5544cd147055de7cdbb4304c98dbfaa9595
|
|
| MD5 |
f641ae39ac15ad0fb047541cc270bf5a
|
|
| BLAKE2b-256 |
a6859b4069a87779becccbf2c721a5303d1d1c7e880e06869b2f802f6630a0ca
|
Provenance
The following attestation bundles were made for pmpp-0.2.2.tar.gz:
Publisher:
publish-to-pypi.yml on rouzib/PMpp
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
pmpp-0.2.2.tar.gz -
Subject digest:
9f0fe3b5f33c54eecddf76aa0427e5544cd147055de7cdbb4304c98dbfaa9595 - Sigstore transparency entry: 2402341013
- Sigstore integration time:
-
Permalink:
rouzib/PMpp@63be2bec3d13fafcd5bceb735f17814a3104f880 -
Branch / Tag:
refs/heads/master - Owner: https://github.com/rouzib
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish-to-pypi.yml@63be2bec3d13fafcd5bceb735f17814a3104f880 -
Trigger Event:
workflow_dispatch
-
Statement type:
File details
Details for the file pmpp-0.2.2-py3-none-any.whl.
File metadata
- Download URL: pmpp-0.2.2-py3-none-any.whl
- Upload date:
- Size: 213.2 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/7.0.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
db584aab8d39ac019e1c1e9e61d3fb2abdcc00d5ac2d89dbe2994e1cb3fcf925
|
|
| MD5 |
31ecb99d96aa08e298250a19b3ba133d
|
|
| BLAKE2b-256 |
35583bae6945adda1db456caea745bc711685e153260679ac560f19bc691e425
|
Provenance
The following attestation bundles were made for pmpp-0.2.2-py3-none-any.whl:
Publisher:
publish-to-pypi.yml on rouzib/PMpp
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
pmpp-0.2.2-py3-none-any.whl -
Subject digest:
db584aab8d39ac019e1c1e9e61d3fb2abdcc00d5ac2d89dbe2994e1cb3fcf925 - Sigstore transparency entry: 2402341149
- Sigstore integration time:
-
Permalink:
rouzib/PMpp@63be2bec3d13fafcd5bceb735f17814a3104f880 -
Branch / Tag:
refs/heads/master - Owner: https://github.com/rouzib
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish-to-pypi.yml@63be2bec3d13fafcd5bceb735f17814a3104f880 -
Trigger Event:
workflow_dispatch
-
Statement type: