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

> Optimize WeatherNext memory with remat options. Learn to set remat_grid_to_mesh_gnn, remat_mesh_gnn, and remat_mesh_to_grid_gnn to True for efficient training.

- Repository: [Google DeepMind/weathernext](https://github.com/google-deepmind/weathernext)
- Tags: how-to-guide
- Published: 2026-08-12

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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:

```python
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:

```python
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:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) around line 266, the wrapping occurs conditionally:

```python
if remat_mesh_gnn:
    mesh_model = utils.modality_data_remat(mesh_model)

```

Similar patterns appear in [`weathernext/utils/xarray_dense.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/xarray_dense.py) (lines 147-155) where `DataArrayDictDenseEncoder` also exposes a `remat` argument, and in [`weathernext/utils/dense.py`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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.