# How to Configure Mesh Padding for Spatial Axis Operations in WeatherNext

> Learn to configure mesh padding for spatial axis operations in WeatherNext. Utilize padding utilities for efficient graph edge management and balanced execution.

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

---

**WeatherNext provides three utilities in [`weathernext/utils/padding_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/padding_utils.py)—`get_num_padded_edges()`, `pad_edges()`, and `get_indices_padding_locations()`—to pad graph edges so their total count becomes a multiple of a hardware-aligned divisor, enabling static-shape compilation and balanced mesh-parallel execution.**

Spatial-axis operations in **WeatherNext**, such as graph convolutions and mesh attention, require fixed-size tensor shapes for efficient XLA compilation on TPU/GPU hardware. Because real-world atmospheric graphs have variable edge counts, the framework implements a deterministic padding system that rounds edge counts up to hardware-friendly multiples. This guide explains how to configure that padding using the official utilities from `google-deepmind/weathernext`.

## Understanding Mesh Padding Requirements

WeatherNext's **graph-mesh transformer** processes edges as dense tensors. Without padding, dynamic edge counts force XLA to generate suboptimal kernels with dynamic shapes. Padding solves this by:

- **Guaranteeing static shapes** for compiler-friendly code generation
- **Enabling mesh parallelism** by aligning tensor dimensions with TPU/GPU tile sizes
- **Preserving graph connectivity** through carefully placed dummy edges

The padding system supports two distribution strategies: **`ends`** (contiguous blocks) and **`linearly_distributed`** (interleaved throughout), with the latter preferred for load-balanced execution.

## Core Padding Utilities in padding_utils.py

All configuration happens through three public functions in [`weathernext/utils/padding_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/padding_utils.py):

### get_num_padded_edges()

Computes the target padded size by rounding up to the nearest multiple.

```python
from weathernext.utils.padding_utils import get_num_padded_edges

num_edges = 250
divisor = 16  # Match your hardware tile size

padded_size = get_num_padded_edges(num_edges, pad_edges_to_multiple_of=divisor)

# Result: 256 (250 rounded up to next multiple of 16)

```

The function signature accepts:

| Parameter | Description |
|-----------|-------------|
| `num_edges` | Integer count of original edges in the graph |
| `pad_edges_to_multiple_of` | Target divisor (typically 8, 16, or 32 for TPU) |

### get_indices_padding_locations()

Returns the indices where original edges should be placed within the padded array.

```python
from weathernext.utils.padding_utils import get_indices_padding_locations

original_locs = get_indices_padding_locations(
    len_before_padding=num_edges,
    len_after_padding=padded_size,
    mode="linearly_distributed"  # or "ends"

)

```

**Mode options:**

- **`ends`** — Places all padding at the start and end of the edge list; preserves original edge ordering
- **`linearly_distributed`** — Spreads padding edges uniformly; improves load balancing across mesh partitions

### pad_edges()

Applies padding to edge features, senders, and receivers tensors.

```python
from weathernext.utils.padding_utils import pad_edges

padded_features, padded_senders, padded_receivers = pad_edges(
    edge_features=edge_features,
    senders=senders,
    receivers=receivers,
    padded_size=padded_size,
    padding_mode="linearly_distributed",
    feature_padding_value=-1.0  # Sentinel value for "no data"

)

```

The function validates that original senders and receivers remain unmodified after insertion of padding edges.

## Complete Configuration Example

Here's a production-ready pattern for configuring mesh padding:

```python
import numpy as np
from weathernext.utils.padding_utils import (
    get_num_padded_edges,
    pad_edges,
)
import jax

# --- Example: Atmospheric graph with variable edges ----------------

num_nodes = 10_000
senders = np.random.randint(0, num_nodes, size=12_345)
receivers = np.random.randint(0, num_nodes, size=12_345)
edge_features = {
    "static_weights": np.random.randn(12_345, 32),
    "dynamic_state": np.random.randn(12_345, 64),
}

# --- Step 1: Choose divisor matching hardware ----------------------

# TPU v4: 128; TPU v3: 8; GPU: 8 or 16 depending on model

hardware_divisor = 128
num_edges = len(senders)

# --- Step 2: Compute padded edge count ---------------------------

padded_size = get_num_padded_edges(num_edges, hardware_divisor)
print(f"Original: {num_edges} edges → Padded: {padded_size} edges")

# Output: "Original: 12345 edges → Padded: 12544 edges"

# --- Step 3: Apply padding with linear distribution ---------------

padded_feat, padded_s, padded_r = pad_edges(
    edge_features,
    senders,
    receivers,
    padded_size,
    padding_mode="linearly_distributed",
    feature_padding_value=0.0,
)

# --- Step 4: Pass to mesh transformer ----------------------------

# from weathernext.utils.mesh_transformer import MeshTransformer

# transformer = MeshTransformer(...)

# output = transformer(padded_feat, padded_s, padded_r, node_features)

```

## Integrating with Spatial-Axis Operations

Padded tensors feed directly into [`weathernext/utils/mesh_transformer.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/mesh_transformer.py). The `MeshTransformer` class expects:

- `senders` and `receivers` as 1-D integer arrays of length `padded_size`
- `edge_features` as a **chex-compatible pytree** with leading dimension `padded_size`
- `feature_padding_value` treated as a sentinel that downstream attention/gather ops ignore

For gather-scatter primitives in [`weathernext/utils/gather_scatter_ops.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/gather_scatter_ops.py), the linear distribution mode ensures padding edges are spread evenly across mesh partitions, preventing straggling tiles at the list boundaries.

## Choosing Padding Parameters

| Scenario | Recommended Configuration |
|----------|---------------------------|
| **Maximum throughput** | `divisor=128`, `mode="linearly_distributed"` (TPU v4 optimized) |
| **Memory-constrained** | Smallest divisor ≥8 that eliminates dynamic shapes |
| **Debugging/validation** | `mode="ends"` to preserve contiguous original edges |
| **Multi-platform code** | `divisor=8` as portable baseline |

## Key Files Reference

| File Path | Purpose |
|-----------|---------|
| [`weathernext/utils/padding_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/padding_utils.py) | Core utilities: `get_num_padded_edges()`, `pad_edges()`, `get_indices_padding_locations()` |
| [`weathernext/utils/mesh_transformer.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/mesh_transformer.py) | Consumes padded tensors in `MeshTransformer` spatial-axis blocks |
| [`weathernext/utils/gather_scatter_ops.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/gather_scatter_ops.py) | JAX primitives requiring padded edge lists for mesh-parallel gather/scatter |

## Summary

- **Mesh padding** in WeatherNext ensures edge counts are multiples of hardware-aligned divisors for static-shape XLA compilation.

- Three functions in [`padding_utils.py`](https://github.com/google-deepmind/weathernext/blob/main/padding_utils.py) handle all configuration: `get_num_padded_edges()` for size calculation, `pad_edges()` for tensor application, and `get_indices_padding_locations()` for index mapping.

- **`linearly_distributed`** mode spreads padding edges uniformly across the list, optimizing load balance for mesh-parallel kernels on TPU/GPU hardware.

- Set `feature_padding_value` to a recognizable sentinel (e.g., `0.0` or `-1.0`) that downstream operations can identify and ignore.

- Integrate with `MeshTransformer` and [`gather_scatter_ops.py`](https://github.com/google-deepmind/weathernext/blob/main/gather_scatter_ops.py) by passing the padded tensors directly; no additional preprocessing required.

## Frequently Asked Questions

### What divisor value should I use for TPU v4?

Use **128** for optimal TPU v4 performance. This matches the hardware's matrix unit dimensions and enables maximum parallelism in [`weathernext/utils/gather_scatter_ops.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/gather_scatter_ops.py). For TPU v3 or GPU hardware, **8 or 16** are typical values. The divisor must be a power of two for most XLA backends.

### Why does WeatherNext use "linearly_distributed" padding instead of simple end-padding?

**Linear distribution spreads padding edges uniformly** throughout the edge list rather than clustering them at boundaries. This prevents "hot" tiles in mesh-parallel execution where one processing unit receives disproportionately many real edges while others idle on padding. According to the `google-deepmind/weathernext` source code, this mode is the default in `pad_edges()` for production deployments.

### How do I verify that padding hasn't corrupted my graph connectivity?

The `pad_edges()` function validates internally that original senders and receivers remain unchanged. For additional verification, compare outputs:

```python
orig_locs = get_indices_padding_locations(num_edges, padded_size, "linearly_distributed")
np.testing.assert_array_equal(padded_senders[orig_locs], senders)
np.testing.assert_array_equal(padded_receivers[orig_locs], receivers)

```

These assertions confirm that real edges occupy their expected positions in the padded arrays.

### Can I use different padding values for different edge feature components?

Yes. The `edge_features` parameter is a **chex-compatible pytree**, so you can pad each component with its own sentinel by preprocessing:

```python
edge_features = {
    "weights": pad_with_zero(weights),
    "mask": pad_with_neg_one(mask),  # different sentinel

}

```

However, `pad_edges()` applies a single `feature_padding_value` to all leaves. For component-specific values, call `jax.tree_map()` with custom padding logic after obtaining indices from `get_indices_padding_locations()`.