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

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 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.

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

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

Three key functions in 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →