About this project
jaxphys is a GPU-accelerated differentiable physics engine built on JAX. It targets a gap between heavyweight research codes (FEniCS, OpenFOAM, COMSOL) that are hard to install and not differentiable, and educational libraries that are CPU-only and toy-level. Everything is pure JAX, so simulations run under jax.jit, jax.vmap and jax.grad.
Covered domains and solvers:
- Classical mechanics: Lagrangian and Hamiltonian systems, symplectic Euler, leapfrog (Störmer-Verlet), Yoshida-4, RK4, Euler, velocity-Verlet N-body, rigid bodies, adaptive RK45 stepper.
- Quantum: split-operator Schrödinger solvers in 1D and 2D, finite-difference eigenstates, tight-binding band structures, Heisenberg spin chains via exact diagonalization, Lindblad master equation.
- Electromagnetism: FDTD (2D TM and 3D) with split-field PML, FDFD (2D TM) with stretched-coordinate PML, Boris-pushed charges, rectangular waveguide modes.
- Fluids: weakly compressible SPH, compressible Euler (1D MUSCL-HLLC), lattice Boltzmann D2Q9, vorticity-streamfunction Navier-Stokes.
- Statistical mechanics: Ising Metropolis (checkerboard) and Wolff cluster Monte Carlo, Boltzmann statistics (Monte Carlo sampling itself is not differentiable).
- Optics: ABCD ray tracing, Fraunhofer diffraction.
A key workflow is defining a Lagrangian and letting JAX autodiff produce equations of motion, then optimizing through entire trajectories. The optimize module provides gradient-based optimization, sensitivity analysis, parameter grid sweeps and refinement of the best sweep points. Integrators are documented with order and symplectic properties; symplectic integrators assume separable Hamilton equations, so LagrangianSystem.simulate accepts only euler and rk4, while HamiltonianSystem is used for symplectic integration.
The README reports CPU-only benchmarks against vectorized NumPy baselines (for example N-body 16.3x, FDTD 3D 8.4x, Schrödinger 2D 1.7x, SPH 0.6x where the SciPy cKDTree baseline wins on CPU), with the caveat that measurements come from a shared 4-core cloud container and no GPU numbers were measured. A validation suite checks solvers against analytic or independent results: Kepler energy error bounded for leapfrog and Yoshida-4 while RK4 drifts, measured convergence orders, Ising critical temperature from Binder-cumulant crossings, FDTD cavity modes against the Yee dispersion relation, Schrödinger norm conservation and free-packet accuracy, LBM channel decay and Poiseuille profile, Sod shock tube L1 error, FDFD versus the Hankel Green's function, and jax.grad gradients compared with central finite differences.
Installation is via pip install jaxphys; examples cover a double pendulum, projectile optimization, quantum tunneling, FDTD diffraction and 3D dipole radiation, Ising phase transition, Kármán vortex street, spacecraft trajectory targeting and a benchmark script. Development uses pytest, ruff and mypy, with an offline demo runnable through uv. The project is MIT licensed and cites standard physics and numerical-integration references.
Comments
0 Rating appears after 10 ratings
Sign in to join the discussion.