# How to Configure Sharding for Multi-TPU Training with JAX Device Meshes in WeatherNext

> Learn to configure JAX device meshes for multi-TPU training in WeatherNext. Discover how to declaratively partition tensors across TPU cores using PartitionSpec for efficient distributed computation.

- Repository: [Google DeepMind/weathernext](https://github.com/google-deepmind/weathernext)
- Tags: how-to-guide
- Published: 2026-08-12

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py)** — defines mesh axis names and global mesh queries
- **[`weathernext/utils/sharding_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding.py) performs this check:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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):

```python

# 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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/sharding_utils.py) (lines 41-49) wraps `jax.lax.with_sharding_constraint` with a safety check for mesh availability:

```python
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_constraint` when a mesh is active
- Accepts any `jax.sharding.PartitionSpec` or `Sharding` object

## Complete Multi-TPU Configuration Example

Here's a production pattern from [`weathernext/weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) (lines 30-42), adapted for clarity:

```python
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:

```python

# 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:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/sharding.py) (`SPATIAL_AXES`, `BATCH_LIKE_AXES`) to ensure consistent dimension mapping across the codebase
- **Apply `PartitionSpec`** through `sharding_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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/architecture.py).

### Can I use custom axis names beyond WeatherNext's defaults?

Yes. The constants in [`sharding.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) (encoder/decoder sharding) and [`weathernext/utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/xarray_dense.py) (dense layer intermediates). Search for `set_sharding` calls to find all usage patterns in the codebase.