How to Configure Sharding for Multi-TPU Training with JAX Device Meshes in WeatherNext
WeatherNext configures multi-TPU training through a global JAX mesh with named axes and PartitionSpec constraints, using sharding_utils.set_sharding() to declaratively partition tensors across TPU cores.
The WeatherNext codebase from Google DeepMind implements large-scale weather forecasting models that require distributed training across hundreds or thousands of TPU cores. Understanding how to configure JAX device mesh sharding is essential for running these models efficiently. This guide explains the three-part sharding workflow used throughout the repository, from mesh initialization to applying PartitionSpec constraints.
Understanding WeatherNext's Sharding Architecture
WeatherNext separates sharding configuration into two core modules:
weathernext/utils/sharding.py— defines mesh axis names and global mesh queriesweathernext/utils/sharding_utils.py— provides sharding application utilities
The design assumes your training script has already established a global mesh through jax.experimental.pjit or pmap. The WeatherNext utilities then query this mesh and apply sharding constraints to individual tensors.
Step 1: Verify the Global JAX Mesh Exists
Before applying any sharding, confirm that a global mesh is available. The is_global_mesh_defined() function in weathernext/utils/sharding.py performs this check:
from weathernext.utils import sharding
if not sharding.is_global_mesh_defined():
raise RuntimeError(
"Global TPU mesh not defined — launch with `pjit` or `pmap`"
)
# Retrieve the current mesh
mesh = sharding.get_global_mesh() # Lines 62-66 in sharding.py
The get_global_mesh() helper returns the active mesh from jax.experimental.maps.global_mesh. If you're using a standard JAX distributed setup, this returns the mesh created by your training launcher (e.g., --jax_distributed_init_method or XLA flags).
Step 2: Choose Named Mesh Axes for Your Data
WeatherNext establishes a convention for axis names that describes how different tensor dimensions map to hardware. These constants are defined in weathernext/utils/sharding.py lines 25-60:
| Axis Name | Purpose | Lines |
|---|---|---|
BATCH_AXIS = "batch" |
Standard data-parallel batch dimension | 25-27 |
SAMPLE_LOCAL_AXIS = "sample_local" |
Ensemble members local to a device | 28-32 |
SAMPLE_PROCESS_AXIS = "sample_process" |
Ensemble members sharded across processes | 33-40 |
LOCAL_SPATIAL_AXIS = "spatial_local" |
Spatial dimensions within one process | 50-56 |
SPATIAL_PROCESS_AXIS = "spatial_process" |
Spatial dimensions across processes | 50-56 |
OUTER_VMAP_AXIS = "outer_vmap" |
Optional leading dimension for vmap |
57-60 |
For convenience, WeatherNext provides composite axis groups (lines 44-56):
# From sharding.py lines 44-56
ENSEMBLE_AXES = (SAMPLE_LOCAL_AXIS, SAMPLE_PROCESS_AXIS)
BATCH_LIKE_AXES = (BATCH_AXIS, *ENSEMBLE_AXES)
SPATIAL_AXES = (SPATIAL_PROCESS_AXIS, LOCAL_SPATIAL_AXIS)
Use BATCH_LIKE_AXES when your tensor has combined batch and ensemble dimensions. Use SPATIAL_AXES for horizontal grid dimensions (latitude, longitude, or processed spatial features).
Step 3: Apply PartitionSpec with set_sharding()
The set_sharding() function in weathernext/utils/sharding_utils.py (lines 41-49) wraps jax.lax.with_sharding_constraint with a safety check for mesh availability:
from weathernext.utils import sharding, sharding_utils
import jax
def shard_tensor(tensor):
"""Apply WeatherNext's standard sharding pattern."""
spec = jax.sharding.PartitionSpec(
sharding.SPATIAL_AXES, # First dimension: spatial parallelism
sharding.BATCH_LIKE_AXES, # Second dimension: batch + ensemble parallelism
)
return sharding_utils.set_sharding(tensor, partition_spec=spec)
Key behavior of set_sharding():
- Returns the tensor unchanged if no global mesh exists (safe for CPU debugging)
- Uses
jax.lax.with_sharding_constraintwhen a mesh is active - Accepts any
jax.sharding.PartitionSpecorShardingobject
Complete Multi-TPU Configuration Example
Here's a production pattern from weathernext/weathernext2/architecture.py (lines 30-42), adapted for clarity:
import jax
from weathernext.utils import sharding, sharding_utils
import weathernext.utils.update_blocks as update_blocks
def encode_with_sharding(latent_grid_data):
"""
Encode grid data with proper TPU sharding for WeatherNext inference.
"""
# Ensure mesh is available
assert sharding.is_global_mesh_defined(), "TPU mesh required"
# Build the sharding spec: spatial dims × batch dims
spec = jax.sharding.PartitionSpec(
sharding.SPATIAL_AXES, # e.g., ("spatial_process", "spatial_local")
sharding.BATCH_LIKE_AXES, # e.g., ("batch", "sample_local", "sample_process")
)
# Apply sharding constraint to the tensor's main data
sharded_main = sharding_utils.set_sharding(
latent_grid_data.data.main,
partition_spec=spec,
)
# Update the data structure with sharded tensor
return update_blocks.update_main_data(
latent_grid_data,
sharded_main,
)
For nested structures, use jax.tree.map to apply sharding uniformly:
# Shard all leaves in a pytree
sharded_state = jax.tree.map(
lambda x: sharding_utils.set_sharding(x, my_spec),
model_state,
)
Debugging and Inspecting Sharding
WeatherNext includes optional inspection utilities. Set DISABLE_INSPECT_SHARDING=False (environment variable or module constant) to enable:
from weathernext.utils import sharding_utils
sharding_utils.inspect_sharding_if_available(
sharded_grid,
label="encoder_output" # Appears in logs
)
This prints the actual sharding layout of each tensor, helping verify that your PartitionSpec maps to hardware as intended.
Temporary Meshes for Device Placement
In some cases—such as rollout or checkpoint loading—you may need a temporary mesh for device_put operations. See weathernext/utils/rollout.py for the device_put_sharded pattern, which creates a local mesh scope for transferring data to devices without affecting the global mesh configuration.
Summary
- Verify mesh availability with
sharding.is_global_mesh_defined()before applying any sharding constraints - Use named axis constants from
sharding.py(SPATIAL_AXES,BATCH_LIKE_AXES) to ensure consistent dimension mapping across the codebase - Apply
PartitionSpecthroughsharding_utils.set_sharding()for safe, mesh-conditional sharding that works in both distributed and single-device contexts - Inspect with
inspect_sharding_if_available()to debug sharding layouts during development - Follow the architecture.py pattern for model code that shards latent grid data across spatial and batch dimensions
Frequently Asked Questions
What happens if I call set_sharding without a global mesh defined?
The function returns the tensor unchanged. This safety behavior in sharding_utils.py lines 41-49 lets you write sharding-agnostic code that runs on both TPU pods and single devices. For production training, always verify mesh availability with is_global_mesh_defined().
How do I choose between SPATIAL_AXES and BATCH_LIKE_AXES for my tensor dimensions?
Match the axis group to your tensor's semantic structure. Use SPATIAL_AXES for dimensions representing physical space (lat/lon grids, spatial features). Use BATCH_LIKE_AXES for independent samples (batch items, ensemble members). The order in your PartitionSpec matters—WeatherNext typically places spatial first, batch second, following the convention in architecture.py.
Can I use custom axis names beyond WeatherNext's defaults?
Yes. The constants in sharding.py are conventions, not requirements. Pass any tuple of strings to PartitionSpec that matches your mesh's axis names. However, using WeatherNext's predefined axes ensures compatibility with pretrained checkpoints and distributed training scripts in the repository.
Where is sharding actually applied in the WeatherNext forward pass?
The primary application sites are weathernext/weathernext2/architecture.py (encoder/decoder sharding) and weathernext/utils/xarray_dense.py (dense layer intermediates). Search for set_sharding calls to find all usage patterns in the codebase.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →