# How `extra_dims_to_split` in `xarray_dense` Controls Weight Sharding in WeatherNeXt

> Understand how extra_dims_to_split in xarray_dense controls weight sharding in WeatherNeXt. Learn to shard parameters across dimensions like time and pressure levels for efficient model training.

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

---

**`extra_dims_to_split` determines which non-preserved dimensions of an `xarray.DataArray` are split into separate weight shards, enabling fine-grained parameter sharding across coordinates like time or pressure level.**

The `xarray_dense` module in the [google-deepmind/weathernext](https://github.com/google-deepmind/weathernext) repository provides dense encoding and decoding for geospatial weather data. The `extra_dims_to_split` parameter (exposed as `dims_to_split` in the code) is central to how model weights are distributed across the feature dimensions of multi-dimensional weather variables. This article explains the mechanism using actual source code from [`weathernext/utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/xarray_dense.py).

## What `extra_dims_to_split` Does

When you initialize a `DataArrayDictDenseEncoder` or `DataArrayDictDenseDecoder`, the `dims_to_split` argument specifies which dimensions beyond the preserved ones should trigger weight sharding. According to the implementation in [`utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/xarray_dense.py) (lines 00136-00150), the encoder processes each `xarray.DataArray` through these stages:

1. **Preserve leading dimensions** — Dimensions listed in `preserved_dims` (typically `"batch"`, `"lat"`, `"lon"`) remain as axes of the underlying array.

2. **Stack split dimensions** — Dimensions in `dims_to_split` are collapsed via `xr.DataArray.stack` into a single **"channels"** dimension.

3. **Concatenate and split weights** — The flattened arrays are concatenated, then sliced back into separate parameters per unique coordinate combination.

The actual sharding occurs in `_SplitInputMatMul._initialize_and_get_params` (lines 00178-00195), where parameter names are constructed as `w_{variable}_{coordinate_strings}`.

## How Weight Sharding Granularity Works

The granularity of sharding depends directly on which dimensions you include in `extra_dims_to_split`:

| Configuration | Result |
|-------------|--------|
| `dims_to_split=()` | Single weight matrix per variable (no sharding) |
| `dims_to_split=("time",)` | Separate weight slice per time coordinate |
| `dims_to_split=("time", "level")` | Sharding over the Cartesian product of time and level |

For example, with geopotential data at multiple pressure levels and forecast steps, setting `dims_to_split=("time",)` creates independent weights for each time coordinate—`w_geopotential_time=-10`, `w_geopotential_time=0`, and so on. This allows the model to learn level-specific or time-specific transformations without sharing parameters across those coordinates.

## Code Example: Configuring `dims_to_split`

Below is a practical example demonstrating how `extra_dims_to_split` affects parameter initialization. The code creates two weather variables with extra dimensions and configures an encoder that shards weights over the `"time"` dimension:

```python
import xarray as xr
import numpy as np
import weathernext.utils.xarray_dense as xd
import haiku as hk
import jax

# Create dummy weather data with extra dimensions.

geopotential = xr.DataArray(
    np.random.randn(2, 3, 4, 2, 2),   # batch, lat, lon, time, level

    dims=["batch", "lat", "lon", "time", "level"],
    coords={"time": [-10, 0], "level": [1, 2]},
)

temperature = xr.DataArray(
    np.random.randn(2, 3, 4, 2),      # batch, lat, lon, time

    dims=["batch", "lat", "lon", "time"],
    coords={"time": [10, 20]},
)

data_mapping = {
    "geopotential": geopotential,
    "2m_temperature": temperature,
}

# Encoder with time-based weight sharding.

encoder = xd.DataArrayDictDenseEncoder(
    name="enc",
    preserved_dims=("batch", "lat", "lon"),
    dims_to_split=("time",),      # <-- extra_dims_to_split controls sharding

    hidden_size=64,
    output_size=32,
    num_hidden_layers=2,
)

def forward(mapping):
    return encoder(data_array_mapping=mapping)

# Initialize and inspect sharded parameters.

init_fn = hk.transform(forward).init
params = init_fn(jax.random.PRNGKey(0), data_mapping)

# Parameters now contain time-specific shards:

# "w_geopotential_time=-10_level=1,2"

# "w_geopotential_time=0_level=1,2"

# "w_2m_temperature_time=10"

# "w_2m_temperature_time=20"

```

The resulting `params` dictionary contains separate entries for each combination of split dimension coordinates, confirming that `extra_dims_to_split` has triggered the expected sharding behavior.

## Implementation Details in [`xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/xarray_dense.py)

Three key functions in [`utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/xarray_dense.py) implement the sharding logic:

- **`_flatten_data_array_mapping`** (lines 00848-00973): Stacks dimensions listed in `dims_to_split` into a channel axis using `xr.DataArray.stack`.

- **`_SplitInputMatMul._initialize_and_get_params`** (lines 00178-00195): Creates per-channel weight shards by slicing the concatenated weight matrix and assigning unique parameter names based on coordinate values.

- **`_SplitOutputLinear._initialize_and_get_params`** (lines 00686-00746): Applies the same sharding pattern to the decoder output layer, ensuring symmetry between encoding and decoding.

The sharding is **logical rather than physical**—parameters remain stored as a single concatenated array at runtime, but the initializer provides fine-grained access to coordinate-specific slices. This design enables efficient parallelism and controlled model capacity expansion without fragmenting the underlying storage.

## When to Use Different `dims_to_split` Configurations

Choosing the right `extra_dims_to_split` value involves trade-offs between model capacity and parameter efficiency:

- **Empty tuple `()`**: Use when you want shared transformations across all coordinates, minimizing parameter count.

- **Single dimension `"time"`**: Appropriate when temporal evolution requires distinct processing (e.g., different physics at different forecast leads).

- **Multiple dimensions `("time", "level")`**: Enable when interactions between coordinates are complex and warrant independent parameter sets, at the cost of increased memory usage.

## Summary

- **`extra_dims_to_split`** (the `dims_to_split` argument) controls which `xarray.DataArray` dimensions become separate weight shards in `DataArrayDictDenseEncoder` and `DataArrayDictDenseDecoder`.

- The parameter is processed in [`utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/xarray_dense.py) by stacking split dimensions into channels, then slicing the weight matrix in `_initialize_and_get_params`.

- Sharding granularity ranges from none (empty tuple) to full Cartesian products of multiple dimensions, enabling flexible trade-offs between model expressiveness and efficiency.

- The resulting parameter names encode coordinate values, making the sharding explicit in the parameter tree.

## Frequently Asked Questions

### What happens if I don't specify `dims_to_split`?

If you pass an empty tuple or omit the argument, the encoder creates a single weight matrix per variable with no sharding. All coordinate values share the same transformation parameters, reducing model capacity but also reducing memory requirements.

### Can I shard over pressure level instead of time?

Yes. Any dimension present in your `xarray.DataArray` can be listed in `dims_to_split`. For atmospheric models, sharding over `"level"` allows distinct transformations at different pressure heights, which may improve representation of vertically varying phenomena.

### Is the sharding physical or logical?

The sharding is **logical**. At runtime, parameters remain concatenated in a single array, but the initializer in `_SplitInputMatMul` provides named views into coordinate-specific slices. This preserves efficient vectorized operations while allowing coordinate-specific initialization and optimization.

### Does sharding affect inference speed?

Sharding primarily affects memory layout and gradient computation during training. The runtime overhead during inference is minimal because the same concatenated weight matrix is used regardless of sharding configuration.