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 ordering
  • linearly_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:

  • 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, 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.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 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. 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:

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 →