How to Configure Mesh Padding for Spatial Axis Operations in WeatherNext
WeatherNext provides three utilities in 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:
get_num_padded_edges()
Computes the target padded size by rounding up to the nearest multiple.
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.
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 orderinglinearly_distributed— Spreads padding edges uniformly; improves load balancing across mesh partitions
pad_edges()
Applies padding to edge features, senders, and receivers tensors.
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:
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. The MeshTransformer class expects:
sendersandreceiversas 1-D integer arrays of lengthpadded_sizeedge_featuresas a chex-compatible pytree with leading dimensionpadded_sizefeature_padding_valuetreated as a sentinel that downstream attention/gather ops ignore
For gather-scatter primitives in 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 |
Core utilities: get_num_padded_edges(), pad_edges(), get_indices_padding_locations() |
weathernext/utils/mesh_transformer.py |
Consumes padded tensors in MeshTransformer spatial-axis blocks |
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.pyhandle all configuration:get_num_padded_edges()for size calculation,pad_edges()for tensor application, andget_indices_padding_locations()for index mapping. -
linearly_distributedmode spreads padding edges uniformly across the list, optimizing load balance for mesh-parallel kernels on TPU/GPU hardware. -
Set
feature_padding_valueto a recognizable sentinel (e.g.,0.0or-1.0) that downstream operations can identify and ignore. -
Integrate with
MeshTransformerandgather_scatter_ops.pyby 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. 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:
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:
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().
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 →