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 meshremat_mesh_gnn– wraps the mesh GNN that processes mesh-based featuresremat_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:
- Extracting only the
.datatensor fields from modality objects - Passing those tensors to
hk.remat - 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, andremat_mesh_to_grid_gnninMultimodalityForward - Memory-compute tradeoff: Enabling these flags reduces peak GPU memory by recomputing activations during backpropagation
- Safe wrapper:
modality_data_rematinupdate_blocks_utils.pypreserves coordinate metadata while rematerializing tensor data - Recommended practice: Start with
remat_mesh_gnn=Truefor 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →