Training a convection emulator inside a differentiable GCM
Why the ResCu convection net is a one-dimensional CNN, what happened in the first online training pass in the JAX climate model, and where the open problems are. Written to be argued with.
Context in one page
ResCu is a neural-network replacement for the convection and cloud parameterisation of a climate model. It was trained offline on a super-parameterised model (SPCAM), where every grid column contains a small cloud-resolving model. The net learns to map the column state the host provides to the tendencies the embedded cloud model produced. It runs inside CAM5 (the hybrid model is called HyCAM) and now inside a JAX port of a spectral GCM (jcm, built on the dinosaur dynamical core), where the whole model, dynamics and physics, is differentiable.
The offline-trained net is good at short lead but has a persistent climatological error in the JAX host: precipitation forms a double band around the equator, with too much rain south of it and too little in the northern intertropical convergence zone. A ladder of nudging experiments showed that the host winds, not the net's column response, build that band: with the winds relaxed to reanalysis the bias mostly disappears. Offline training cannot see this feedback. Online training, where the loss is evaluated after the net and the host have interacted for some time, can in principle reach it.
Host model
jcm on a T63 spectral grid (about 1.9 degrees), 38 levels, 1800 s time step, ERA5 initial states, prescribed ERA5 surface heat and moisture fluxes, interactive surface drag, RRTMGP radiation.
Trainable part
The ResCu super net only, 1,062,218 parameters together with the cloud net. Everything else is fixed physics, differentiated but not trained.
Targets
IMERG daily precipitation and ERA5 winds on model levels, 2001 to 2008 for training, 2009 to 2014 held out.
Hardware
Perlmutter shared GPU queue, one A100 40 GB per job plus a second card claimed only for host memory (compiles need about 58 GB of host RAM). Login-node GPUs kill any process above about 30 GB of host memory.
Why the emulator is a one-dimensional CNN
The question from the computing side: of all the architectures, why a convolution along the vertical?
In a global model each atmospheric column is treated independently by the physics. The input to a convection scheme is a set of vertical profiles (temperature, humidity, the large-scale forcing on each) plus a few surface scalars, and the output is a set of vertical profiles of tendencies. So the natural data object is a multi-channel one-dimensional signal along the level axis: 30 levels, 32 input channels, 4 output channels. That is exactly what a 1-D CNN consumes.
How the net treats one column. (a) The inputs are profiles on the 30 CAM5 levels, top first, with scalars broadcast along the column. (b) One convolution reads 3 adjacent levels across all channels and writes one value per filter at that level; the same weights slide over all 30 levels, and 128 filters give 128 output channels. (c) Each layer widens the vertical view by two levels, so from the 15th layer every output level sees the whole column. (d) The outputs are again profiles: heating, drying and condensate. Profiles are idealised.
The reasons, in order of weight
Locality in the vertical. Convective transport is a vertical flux divergence, cloud formation depends on the layers just above and below, and stability is a local vertical gradient. A kernel of width 3 expresses those operators directly; a stack of blocks extends the receptive field to the full column, which deep convection needs.
Weight sharing across levels. A plume or an inversion behaves similarly whether it sits at level 12 or level 20. Sharing the kernel over levels gives approximate translation equivariance in height and keeps the net small: a million parameters, versus a dense network over the flattened 960-element column that would spend most of its parameters on level-to-level coupling it must learn from scratch.
Residual blocks. The net predicts tendencies, small corrections to a state; pre-activation residual blocks with identity shortcuts train stably at depth and start close to the identity, which matches the physics of a small tendency.
Memory instead of recurrence. A cloud-resolving model has state that persists between host time steps. ResCu approximates that with a two-step history of inputs and its own previous outputs fed back as channels, which is cheaper and simpler to deploy than a recurrent cell and proved sufficient offline.
Deployability. Column independence means the net batches over all columns of the globe in one call, in Fortran (through FTorch in CAM5) or in JAX. Nothing in the architecture couples neighbouring columns, so the host's parallel decomposition is untouched.
What we know it does not do
It is deterministic. The cloud-resolving truth is stochastic; the net predicts a conditional mean. Cloud condensate in particular is under-dispersed, which is one reason the cloud net exists as a separate, smaller stage. A diffusion-model variant was explored for this reason.
It has no horizontal context. Organised convection (the MJO, squall lines) is emergent from the host coupling, not represented inside the net.
A vertical convolution-plus-attention variant was built and coupled. It did not remove the coupled biases either, which is part of why the emphasis moved from architecture to training procedure.
Online training setup
Each iteration has two steps. Generate: run the model with the current weights from many start dates and save the state after a spin-up. Optimise: from each saved state, roll the model forward inside the gradient for a window, compare with observations, and back-propagate through dynamics and physics into the net.
Samples
200 episodes: random 6-hourly start instants in 2001 to 2008, each an ERA5 initial state spun up 3 days with the current net. Each episode holds the window start state (472 MB) and 2 days of targets. Six held-out episodes from 2009 to 2014 serve the demos.
Window in pass 1
12 time steps of 1800 s, so 6 hours of model evolution inside the gradient. The activations of every step are checkpointed and recomputed in the backward pass; the 12-step gradient fits in 29 GB.
Loss
Square-root-space mean squared error of the window's precipitation against the IMERG day, plus 0.05 times a mass-weighted wind error against ERA5 at the end of the window, plus a penalty on the global-mean precipitation error (weight 100) that stops the optimiser from removing light rain everywhere.
Optimiser
Adam, gradient clipped at norm 1, learning rate 3e-6 for the first 206 steps and 1e-5 after. One episode per gradient step, no batching (one episode saturates the GPU). Non-finite gradients are skipped.
What is differentiated and what is not
The forward model is complete: dynamics, the net, the rectifiers, vertical diffusion, surface fluxes and radiation all run every step. In the backward pass the radiation scheme is detached: its heating is treated as a fixed forcing along the trajectory rather than differentiated with respect to the state. The reason is practical. The exact radiation backward needs about 300 GB of temporaries unless chunked, and the chunked version produces NaN on its own. Operational 4D-Var systems make the same kind of simplification in their adjoint physics.
Measured, not assumed. A finite-difference check of the full forward loss along the detached-gradient direction gives a ratio of 0.98 to 1.02 to the detached gradient on a 1-hour window (step sizes 1e-6 to 1e-4). The 12-step and 2-day versions of the check are queued, and the ratio as a function of window length is the number to report.
Consequence. The cloud net receives no gradient at all: its outputs reach the loss only through radiation. Its Adam moments are exactly zero after 600 steps. Pass 1 trained the super net alone. For precipitation targets this is acceptable and even reasonable; a 6-hour window is only sensitive to precipitation anyway. Training the cloud net needs either radiation in the backward pass or a cloud or outgoing-longwave target.
Cost figures
Item
Value
Note
Gradient step, 12-step window, compiled
11 s
one A100
Compile of the 12-step backward
30 min
once per forcing year, cached on disk
GPU memory, 12-step window
29 GB
40 GB with radiation live at 32 column chunks
Host memory during compile
58 GB
forces 2 GPUs per shared-queue job
Episode generation
2 to 3 min
3-day spin-up, one GPU
5-day free rollout for a demo
1 min
forward only
Compile of an alternative nested form of the 12-step backward
42 min
versus 3.5 min to load the cached straight-line form
Pass 1 results
Three epochs over the 200 episodes: 603 gradient evaluations, 355 finite updates, 248 lost to NaN. Because each epoch visits the same 200 windows, the epoch means are directly comparable.
Epoch
Precip term
Wind term x 0.05
Window mean P (mm/day)
Finite steps
1
2.217
9.24
2.71
121 / 200
2
2.026
9.23
2.62
124 / 200
3
1.871
9.23
2.65
110 / 200
The precipitation term fell 16 percent, improved in every one of the 200 episodes, and was still falling linearly at the end: the fit is far from saturated. The wind term, four times larger in the loss, did not move at all. Within 6 hours the winds are fixed by the initial condition, so that term cannot be learned by a convection net, yet it dominates the gradient.
Held-out free rollouts with the final weights
Lead day
1
2
3
4
5
RMSE vs IMERG, original (mm/day)
8.72
8.86
9.20
9.29
9.56
RMSE vs IMERG, trained
7.95
8.42
8.83
9.02
9.31
Tropical pattern correlation, original
0.28
0.23
0.25
0.21
0.16
Tropical pattern correlation, trained
0.29
0.27
0.28
0.22
0.18
Five-day means over the six held-out starts, final pass-1 weights. The change removes rain along the southern band and over the warm pool and keeps the northern ITCZ; the error against IMERG drops from 2.92 to 2.65 mm/day.Scores by lead day (shading is the standard error over the six starts) and daily maps at lead days 1 and 3 for one start.Thirty-day free rollouts, intermediate weights (panel titles say 5-day; they are 30-day means). The gain is confined to the first five days. The 30-day climatology, including the double band, is essentially unchanged: southern zonal peak 5.8 to 5.7 mm/day against IMERG's 4.4, northern 5.2 to 5.1 against 6.2. A 6-hour window cannot see the feedback that builds the band.
Problems, ranked
1. Half the gradient steps come back NaN
The forward pass is always finite. On a NaN step every leaf of the super net's gradient is NaN. What we established today:
The same episode with the same weights gives a finite gradient in a fresh process, and NaN on every retry within the process that first produced NaN.
Two differently compiled executables of the same loss (a straight-line unrolled form and a nested-group form) fail on independent subsets of the same 25 episodes: 9 and 7 failures, 2 shared, where independence predicts 2.5.
Ruled out: the episode, the weights, the time step, jax 0.11.2 versus 0.11.1, CUDA command buffers, XLA autotuning.
Our reading is a memory-state-dependent kernel in the compiled backward, most likely an uninitialised read, whose effect is deterministic within a process because the allocation sequence repeats. The mitigation now installed is a second executable used only when the first gives NaN, which should cut the loss from a third of the steps to under a tenth. It is a workaround, not an understanding.
2. Compile time and memory scale with the window
The 12-step backward compiles in 30 to 40 minutes per forcing year, and the compile grows with the number of unrolled steps: a 5-day window (240 steps) unrolled would compile for hours. A scan over checkpointed 12-step groups, which compiles one group and loops it, is written and being tested. That form was blamed for NaN early on; given problem 1, that blame is probably misattributed.
3. The loss is badly proportioned
The wind term is unlearnable at 6 hours and dominates the gradient. A precipitation-only pass (wind weight zero, global-mean penalty kept) is running now. Winds become a meaningful target only at 2-day or 5-day windows, where the net's heating has had time to spin the circulation up.
4. Short windows cannot reach the climatology
This is the central scientific problem and the reason for everything above. The double band is a host-circulation feedback with a time scale of days. Training must integrate over that time scale, which means longer windows and the cost that comes with them.
Plan and success criteria
The goal for the next weeks is a demonstration that online training produces a significant precipitation improvement in free runs of days, weeks and a month, convincing enough to justify a larger GPU allocation.
running
Precipitation-only training on 6-hour windows, three slots of 100 steps, from the original weights. Tests whether the loss drops faster without the wind term and whether a 6-hour window can move the climatology at all.
queued
2-day windows (48 steps): unrolled nested form and scan form side by side, for compile time, step time, memory and NaN rate. The episodes already hold 2 days of targets.
next
5-day windows: extend the episode targets to five IMERG days and the matching ERA5 wind instants (no new spin-up), scan form, wind term back in with a weight chosen from the 2-day results.
each pass
Five-day and 30-day free rollouts from the six held-out starts, scored against IMERG, same figure format every time.
Criteria we propose to hold ourselves to
Five-day rollouts: RMSE and tropical pattern correlation better than the original net at every lead day. Pass 1 meets this, by 3 to 9 percent.
Thirty-day rollouts: the 30-day-mean pattern error against IMERG clearly down (pass 1: 2.39 to 2.30, not enough), the southern zonal peak at least halfway from 5.8 toward IMERG's 4.4 mm/day, the northern peak not weakened.
No drift and a stable global mean over 30 days.
What more GPU would buy
80 GB cards: radiation live in the backward pass at manageable chunking, so the cloud net can be trained and the adjoint is closed.
Several cards per job: one episode per card with gradient accumulation across cards, so a batch of episodes per update instead of one.
Dedicated nodes: compiles of long windows once, without the 2-hour shared-queue slot as the unit of work.
Questions for discussion
Where a computing-science view would change what we do.
Localising the NaN. A gradient that is NaN under one XLA executable and finite under another, deterministic within a process, for the same inputs. What is the right tool: compute-sanitizer initcheck on the compiled kernels, XLA's HLO dumps with buffer assignment, deterministic-ops flags (which failed to compile here), or bisecting the backward by stopping gradients at intermediate points? Is the two-executable fallback an acceptable engineering answer for now?
Memory versus compile time for long windows. Straight-line unrolling with per-step checkpointing fits 12 steps in 29 GB; nesting groups extends it at the price of compile time. Is there a better remat policy or a scan structure that keeps the compile at one group while preserving the checkpoint structure? Is pipeline parallelism over the window across GPUs realistic in JAX?
Batching. One episode per update saturates a 40 GB card. Gradient accumulation across cards is the obvious route; is there anything smarter than data parallelism over episodes for a model where each sample is a global state?
Loss design. Square-root transform of daily precipitation, a global-mean penalty, a wind term that only becomes learnable at longer windows. Is a curriculum over window length (6 hours, then 2 days, then 5 days, warm-starting each) the right structure, and how would one detect that a longer window has started to matter?
Architecture. Given a differentiable host, is the 1-D CNN still the right emulator, or does online training favour something with horizontal context, a learned memory state, or stochastic outputs? What would it take to make the cloud net trainable without radiation gradients?
Precision. The dynamical core defaulted to bfloat16 matrix multiplies and had to be forced to float32; the finite-difference closure check is at the edge of float32 noise. Where would mixed precision be safe in this pipeline and where not?
Code: the trainer runs/train/train_episodes.py (episode loop, loss, two-executable fallback), the generator runs/train/generate_episodes.py, the evaluator runs/train/rollout_eval.py, all in the jcm_rescu project on Perlmutter. The design note with the full derivation is analysis/online_training_design_20260927.pdf.