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.

10 min read • View on GitHub • More from google-research

A towering storm column rises above a tiled accelerator mesh, with a source file unfurling on one side and partitioned tensor blocks on the other. The scene explains how the project turns distributed execution into a compiler problem instead of a manual message-passing problem.
Swirl-Jatmos makes the scale story look like a layout problem, not a cluster-management problem.
Key Takeaways

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.

DimensionSwirl-JatmosWRF / ICONswirl-lmjax-cfd
LanguagePython + JAXFortran and C++TensorFlowJAX
Parallelism modelNamedSharding over a device meshExplicit MPI domain decompositionTPU-aware distributed executionJAX transformations and sharding
Hardware biasTPU and GPUCPU clusters, with some accelerator pathsTPUGPU and TPU
Physics scope3D atmospheric LES with radiation, microphysics, and pressure solvesOperational weather and climate modelingCFD for variable-density flowsCore CFD building blocks
DifferentiabilityDesigned to stay close to end-to-end gradientsLimited and often indirectResearch-oriented, not the main focusA core advantage
Best fitReadable accelerator-native research codeBattle-tested forecasting systemsAccelerator-scaled fluid simulationResearch 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.

One timestep, from mesh assignment to checkpoint, rendered as a compiler-shaped pipeline.

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.

A close-up of a structured grid under pressure, with one crack running through the lattice and a tensor-shaped clamp pulling the field back into alignment. The image explains why pressure coupling is the numerical bottleneck and why a fast diagonalization solver is a performance story, not a detail.
The pressure solve is where local updates become a global constraint.

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.