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=2 or 3 according 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_splits is an integer constructor argument to ForwardPass in weathernext/weathernext2/architecture.py that 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_splits and consumed at forward-pass time via TriangularMeshData.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:

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 →