Swirl-Jatmos: The Atmospheric Simulator That Lets JAX Do the MPI
A TPU-first Python framework for 3D weather physics, where sharding, pressure solves, and checkpoints are handled as compiler-aware infrastructure, not hand-built cluster code.
- Swirl-Jatmos turns distributed atmospheric simulation into a JAX sharding problem, which is the repo's sharpest surprise.
- The project changes the execution model, not the physics, so the solver still carries the weight of serious atmospheric LES.
- Its fast diagonalization pressure solve matters because the pressure step is where structured fluid codes usually pay the highest global cost.
- Because the whole stack lives in JAX, the model is closer to differentiable Earth modeling than legacy MPI-era weather code.
Weather codes are usually judged on physics first and ergonomics second. Swirl-Jatmos flips that order. It uses JAX sharding, Mesh, and NamedSharding so the parallelism layer looks like normal Python, while the solver still carries the machinery of atmospheric large-eddy simulation.
Why atmospheric simulation is such a brutal test case
Atmospheric simulation is a hard benchmark because the equations are stiff, the domain is three-dimensional, and the pressure step couples the whole grid. You do not get to hide weak numerical choices behind pretty syntax. If the time step, grid layout, or solve strategy is off, the model pays for it immediately.
| Dimension | Swirl-Jatmos | WRF / ICON | swirl-lm | jax-cfd |
|---|---|---|---|---|
| Language | Python + JAX | Fortran and C++ | TensorFlow | JAX |
| Parallelism model | NamedSharding over a device mesh | Explicit MPI domain decomposition | TPU-aware distributed execution | JAX transformations and sharding |
| Hardware bias | TPU and GPU | CPU clusters, with some accelerator paths | TPU | GPU and TPU |
| Physics scope | 3D atmospheric LES with radiation, microphysics, and pressure solves | Operational weather and climate modeling | CFD for variable-density flows | Core CFD building blocks |
| Differentiability | Designed to stay close to end-to-end gradients | Limited and often indirect | Research-oriented, not the main focus | A core advantage |
| Best fit | Readable accelerator-native research code | Battle-tested forecasting systems | Accelerator-scaled fluid simulation | Research and data generation |
The trick: sharding is part of the API
The important move is that the code describes placement, not message passing. In sim_initializer.py, the mesh layout is declared up front, the arrays are constrained to that layout, and XLA handles the distributed mechanics. The developer shapes the grid. The compiler figures out how to keep it coherent across devices.
mesh = jax.sharding.Mesh(devices, ('x', 'y', 'z'))
state_sharding = jax.sharding.NamedSharding(
mesh,
jax.sharding.PartitionSpec('x', 'y', None),
)
state = jax.lax.with_sharding_constraint(state, state_sharding)
for _ in range(num_steps):
state = rk3_step(state, config)
state = solve_pressure(state)
state = checkpoint_manager.save(state)
Read that as a control-flow sketch, not as the whole program. The point is that the sharding choice sits beside the solver logic, which is exactly why the repo feels different from older MPI-first codes. The programmer is still in charge of the math. The compiler is in charge of the logistics.
Inside the solver: RK3, C-grid, and the pressure step
The solver uses a third-order Runge-Kutta loop on an Arakawa C-grid. Velocities live on cell faces, scalars live in cell centers, and the driver keeps prognostic, diagnostic, and auxiliary state separate so each piece can move through the pipeline cleanly. That is not a toy model. It is the real shape of atmospheric numerics.
The pressure solve is the choke point. Swirl-Jatmos uses a fast diagonalization approach in linalg/fast_diagonalization_solver.py, which matters because pressure is global and global steps are where many fluid codes lose their speed. A faster Poisson solve does not just shave cycles. It changes what scale the whole model can reach.
How it stacks up against nearby models
This is where the repo separates itself from both legacy atmospheric codes and nearby accelerator-native projects. Compared with WRF or ICON, it removes most of the manual parallelism surface. Compared with swirl-lm, it swaps TensorFlow for JAX and leans harder into sharding as an API. Compared with jax-cfd, it is much more opinionated about full atmospheric physics rather than core CFD primitives.
What differentiable weather unlocks
The real upside is not just speed. Because the whole stack lives in JAX, the model is closer to differentiable Earth modeling, which makes it useful for data assimilation, parameter tuning, and hybrid physics-ML workflows. That is the deeper bet here. If the simulator stays readable while the compiler handles the distribution, the research loop gets a lot shorter.