How to Configure FGN Architecture with `mesh_num_splits` in WeatherNext: A Complete Guide
Set the mesh_num_splits parameter in the ForwardPass Haiku module to control how many times the icosahedral mesh is recursively subdivided, which directly determines the spatial resolution of the Fully-Connected Graph Neural (FGN) architecture.
The mesh_num_splits parameter is a core configuration lever in Google DeepMind's WeatherNext FGN architecture. Located in weathernext/weathernext2/architecture.py, this integer value governs the granularity of the spherical mesh that underpins the model's graph neural network operations. Understanding how to properly configure this parameter enables you to balance prediction accuracy against computational resource requirements.
What mesh_num_splits Controls in the FGN Architecture
The FGN architecture builds its geometric foundation on a recursively subdivided icosahedral mesh. The mesh_num_splits constructor argument in ForwardPass specifies how many refinement iterations are applied to this base mesh.
In weathernext/weathernext2/architecture.py lines 59-84, the ForwardPass.__init__ method defines:
def __init__(
self,
latent_dense_kwargs: dense.DenseLayerKwargs,
output_dense_kwargs: dense.DenseLayerKwargsExceptOutputSize,
spatial_features_kwargs: ...,
mesh_num_splits: int, # <-- the key parameter
points_to_mesh_model_ctor: ...,
mesh_model_ctor: ...,
mesh_to_grid_model_ctor: ...,
mesh_padding_kwargs: ...,
) -> None:
This value is immediately stored on the instance at line 131:
self._mesh_num_splits = mesh_num_splits
How mesh_num_splits Affects Mesh Resolution
During the forward pass (lines 211-215), the mesh is constructed using:
input_mesh_data = TriangularMeshData.with_icosahedral_mesh(
data=None, splits_list=[2] * self._mesh_num_splits)
The expression [2] * self._mesh_num_splits creates a list where each element represents splitting each edge twice per iteration. This recursive subdivision follows the pattern:
mesh_num_splits |
Faces | Approximate Nodes |
|---|---|---|
| 1 | 80 | ~42 |
| 2 | 320 | ~162 |
| 3 | 1,280 | ~642 |
| 4 | 5,120 | ~2,562 |
Each increment of mesh_num_splits quadruples the number of faces while approximately quadrupling the node count. This exponential growth directly impacts memory consumption and computational throughput.
Configuring mesh_num_splits: Two Approaches
Direct Instantiation with Haiku
For explicit module construction, pass mesh_num_splits as a keyword argument to ForwardPass:
import haiku as hk
from weathernext.weathernext2 import architecture
from weathernext.utils import dense, sharding
forward_cfg = dict(
latent_dense_kwargs=dense.DenseLayerKwargs(...),
output_dense_kwargs=dense.DenseLayerKwargsExceptOutputSize(...),
spatial_features_kwargs=dict(...),
mesh_num_splits=3, # 3 splits → 1,280 faces, finer resolution
points_to_mesh_model_ctor=architecture.PointsMeshUpdateConstructor(...),
mesh_model_ctor=architecture.MeshUpdateConstructor(...),
mesh_to_grid_model_ctor=architecture.PointsMeshUpdateConstructor(...),
mesh_padding_kwargs=None,
)
def build_forward():
return architecture.ForwardPass(**forward_cfg)
forward = hk.transform(build_forward)
params = forward.init(rng_key, inputs=..., targets_template=..., forcings=..., is_training=True)
Fiddle Configuration System
WeatherNext uses Fiddle for structured configuration. Set mesh_num_splits within a fdl.Config:
import fiddle as fdl
from weathernext.weathernext2 import architecture
cfg = fdl.Config(
architecture.ForwardPass,
latent_dense_kwargs=...,
output_dense_kwargs=...,
spatial_features_kwargs=...,
mesh_num_splits=4, # set desired refinement level
points_to_mesh_model_ctor=...,
mesh_model_ctor=...,
mesh_to_grid_model_ctor=...,
)
This integrates with WeatherNext's configuration pipeline for reproducible experiment tracking.
Selecting the Right mesh_num_splits Value
Consider three factors when configuring this parameter:
- Computational budget — Higher values increase both memory usage and per-step latency due to larger mesh operations
- Target variable resolution — Finer atmospheric features (e.g., localized precipitation) benefit from
mesh_num_splits >= 3 - Downstream task requirements — Global forecasting at 0.25° resolution typically uses
mesh_num_splits=2or3according to the WeatherNext source code
The default value of 2 (320 faces) provides a baseline suitable for many research experiments, while production deployments may tune this based on hardware constraints.
Key Source Files for mesh_num_splits Configuration
| File | Purpose |
|---|---|
weathernext/weathernext2/architecture.py |
Defines ForwardPass and mesh_num_splits parameter handling |
weathernext/utils/data_modalities.py |
Contains TriangularMeshData class for mesh instantiation |
weathernext/weathernext2/architecture_utils.py |
Helper utilities for mesh dtype and structural operations |
weathernext/utils/sharding.py |
Sharding specifications that interact with mesh padding |
Summary
mesh_num_splitsis an integer constructor argument toForwardPassinweathernext/weathernext2/architecture.pythat controls icosahedral mesh refinement- Each increment quadruples face count, exponentially increasing spatial resolution and computational cost
- Configure via direct instantiation with Haiku or Fiddle configs for integration with WeatherNext's experiment system
- Typical values range from 2-4, with 2 as a common default for development and 3-4 for higher-resolution applications
- The parameter is stored as
self._mesh_num_splitsand consumed at forward-pass time viaTriangularMeshData.with_icosahedral_mesh()
Frequently Asked Questions
What happens if I set mesh_num_splits too high?
Memory usage grows quadratically with mesh node count, and the GNN message-passing operations become prohibitively expensive. Values above 5 typically require specialized distributed training setups or reduced batch sizes to fit within accelerator memory.
Can I change mesh_num_splits after model initialization?
No. The mesh topology is baked into the model structure at init time through TriangularMeshData.with_icosahedral_mesh(). To use a different resolution, you must reinitialize the ForwardPass module with a new configuration.
How does mesh_num_splits interact with model sharding?
The mesh_padding_kwargs parameter in ForwardPass works alongside mesh_num_splits to ensure mesh partitions align with device boundaries. Consult weathernext/utils/sharding.py for sharding specifications that account for your chosen mesh size.
What's the difference between mesh_num_splits and the splits_list argument?
splits_list in TriangularMeshData.with_icosahedral_mesh() allows per-iteration split counts, but ForwardPass internally fixes this to [2] * mesh_num_splits for uniform subdivision. The mesh_num_splits parameter provides a simpler integer API for this common case.
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 →