How to Use `remat` Options for Memory Optimization in WeatherNext: A Complete Guide

Enable memory-efficient training in WeatherNext by setting remat_grid_to_mesh_gnn, remat_mesh_gnn, and remat_mesh_to_grid_gnn to True when constructing MultimodalityForward, which wraps JAX rematerialization around specific GNN sub-modules to trade compute for reduced activation storage.

WeatherNext is a graph-neural-network (GNN) based weather forecasting system developed by Google DeepMind. The model processes three distinct data modalities—grid-to-mesh, mesh, and mesh-to-grid—and training these large GNNs can quickly exhaust GPU memory due to cached intermediate activations. This article explains how to use remat options for memory optimization in WeatherNext by leveraging JAX's rematerialization capabilities through three boolean flags in the model constructor.

Understanding JAX Rematerialization in WeatherNext

JAX rematerialization (hk.remat) is a memory optimization technique that recomputes forward-pass activations during the backward pass instead of storing them. This reduces peak GPU memory usage at the cost of additional computation time.

In weathernext/weathernext2/architecture.py, the MultimodalityForward class exposes three constructor arguments that control rematerialization for each GNN sub-module:

  • remat_grid_to_mesh_gnn – wraps the grid-to-mesh GNN that converts latitude-longitude grids to an icosahedral mesh
  • remat_mesh_gnn – wraps the mesh GNN that processes mesh-based features
  • remat_mesh_gnn – wraps the mesh-to-grid GNN that projects mesh updates back to the regular grid

Each flag defaults to False. When set to True, the corresponding sub-module is wrapped with hk.remat via the helper function utils.update_blocks_utils.modality_data_remat.

The modality_data_remat Implementation

The core rematerialization logic lives in weathernext/utils/update_blocks_utils.py (lines 39-84). This helper safely handles WeatherNext's modality objects by:

  1. Extracting only the .data tensor fields from modality objects
  2. Passing those tensors to hk.remat
  3. Re-injecting non-tensor metadata (coordinates, lat/lon, etc.) after computation

This design ensures that rich coordinate metadata isn't lost during the rematerialization process, which would break downstream xarray operations.

Enabling Remat Flags: Code Examples

Rematerialize All Three GNN Blocks

For maximum memory reduction, enable all three remat options when constructing your model:

import weathernext.weathernext2.architecture as wn2

model = wn2.MultimodalityForward(
    latent_dense_kwargs=...,               # required dense-layer kwargs

    output_dense_kwargs=...,               # required dense-layer kwargs

    spatial_features_kwargs=...,           # spatial feature config

    mesh_num_splits=3,
    points_to_mesh_model_ctor=...,         # mapping of modality → constructor

    mesh_model_ctor=...,                   # mesh GNN constructor

    mesh_to_grid_model_ctor=...,           # mesh-to-grid GNN constructor

    # ---------------------------------------------------------

    remat_grid_to_mesh_gnn=True,   # wrap grid-to-mesh GNN

    remat_mesh_gnn=True,           # wrap mesh GNN

    remat_mesh_to_grid_gnn=True,   # wrap mesh-to-grid GNN

)

Selective Rematerialization: Target the Mesh GNN Only

The mesh GNN is typically the most memory-intensive block. If you want to minimize compute overhead while still achieving substantial memory savings, rematerialize only this component:

model = wn2.MultimodalityForward(
    ...,
    remat_grid_to_mesh_gnn=False,
    remat_mesh_gnn=True,          # only the mesh GNN is rematerialized

    remat_mesh_to_grid_gnn=False,
)

This configuration roughly halves peak GPU memory for the mesh processing stage while avoiding recomputation overhead in the lighter grid-to-mesh and mesh-to-grid projections.

Training with Rematerialization Enabled

Once constructed, the model runs transparently—the rematerialization happens automatically during the backward pass:

output = model(
    inputs=xr_inputs,
    targets_template=xr_targets,
    forcings=xr_forcings,
    is_training=True,  # rematerialization only affects training (backward pass)

)

Performance Tradeoffs of Remat Options

When using remat options for memory optimization in WeatherNext, consider these characteristics:

Configuration Peak Memory Reduction Compute Overhead Best Use Case
remat_mesh_gnn=True only ~30-40% Low Balanced training on mid-size GPUs
All three flags True ~50%+ Moderate Large models on memory-constrained hardware
All flags False (default) None None Inference or when memory is abundant

The exact savings depend on batch size, mesh resolution, and model depth. The mesh GNN block typically dominates activation memory due to its iterative message-passing structure.

Where Rematerialization Is Applied in the Source Code

In weathernext/weathernext2/architecture.py around line 266, the wrapping occurs conditionally:

if remat_mesh_gnn:
    mesh_model = utils.modality_data_remat(mesh_model)

Similar patterns appear in weathernext/utils/xarray_dense.py (lines 147-155) where DataArrayDictDenseEncoder also exposes a remat argument, and in weathernext/utils/dense.py (lines 42-50) for dense layer rematerialization. These implementations demonstrate the consistent application of modality_data_remat across the codebase.

Summary

  • Three boolean flags control rematerialization: remat_grid_to_mesh_gnn, remat_mesh_gnn, and remat_mesh_to_grid_gnn in MultimodalityForward
  • Memory-compute tradeoff: Enabling these flags reduces peak GPU memory by recomputing activations during backpropagation
  • Safe wrapper: modality_data_remat in update_blocks_utils.py preserves coordinate metadata while rematerializing tensor data
  • Recommended practice: Start with remat_mesh_gnn=True for the best memory-to-overhead ratio, then expand if needed

Frequently Asked Questions

What happens if I enable all three remat flags?

All three GNN sub-modules—grid-to-mesh, mesh, and mesh-to-grid—will be wrapped with hk.remat. This maximizes memory reduction but adds the most computational overhead during training, as each wrapped module's forward pass is recomputed during backpropagation.

Does rematerialization affect inference performance?

No. JAX rematerialization in hk.remat only affects the backward pass. Since is_training=False during inference, no recomputation occurs and the model runs at full speed with standard memory usage.

Why does WeatherNext need a special modality_data_remat helper?

WeatherNext uses rich modality objects that combine tensor data with xarray coordinates and metadata. The standard hk.remat would treat the entire object as a JAX traceable, potentially losing non-tensor attributes. The helper explicitly extracts .data tensors, applies rematerialization, then restores metadata, preserving the object structure required by downstream processing.

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 →