Discussion page, 28 September 2026

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.

per column, top first: 32 input channels x 30 levels t-2 block t-1 block t0 block: T, q, divT, divq, SHF, LHF, ps, SOLIN memory: own previous dT and dq Conv1d, k = 3 32 to 128 ch 10 pre-activation residual blocks each: BN, ReLU, Conv1d k=3, BN, ReLU, Conv1d k=3, + identity 128 channels, 20 conv layers receptive field: whole column ReLU, Conv1d 128 to 4 ch super net outputs dT (heating) dq (moistening) Qc, Qi (condensate) cloud net, 4 to 6 ch Conv1d 4 to 64, 2 residual blocks, Conv1d 64 to 6 number, fraction, snow Between the nets: procedural rectifiers in the deployed order (CFL clips on heating and moistening, negative-moisture floor, supersaturation relaxation, negative-precipitation fixer). Precipitation is the column integral of the drying. Qc and Qi go to the cloud net and to radiation. Every column of the globe (about 20,000 at T63) is one batch element; the same weights serve all columns. Parameters: 1,004,676 in the super net, 51,910 in the cloud net. Cost in CAM5 at 2 degrees: about 1.7 s per model day.
Four panels: input profiles on 30 levels, one convolution with weights shared over levels, receptive field versus depth, output tendency profiles
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

What we know it does not do

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

ItemValueNote
Gradient step, 12-step window, compiled11 sone A100
Compile of the 12-step backward30 minonce per forcing year, cached on disk
GPU memory, 12-step window29 GB40 GB with radiation live at 32 column chunks
Host memory during compile58 GBforces 2 GPUs per shared-queue job
Episode generation2 to 3 min3-day spin-up, one GPU
5-day free rollout for a demo1 minforward only
Compile of an alternative nested form of the 12-step backward42 minversus 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.

EpochPrecip termWind term x 0.05Window mean P (mm/day)Finite steps
12.2179.242.71121 / 200
22.0269.232.62124 / 200
31.8719.232.65110 / 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 day12345
RMSE vs IMERG, original (mm/day)8.728.869.209.299.56
RMSE vs IMERG, trained7.958.428.839.029.31
Tropical pattern correlation, original0.280.230.250.210.16
Tropical pattern correlation, trained0.290.270.280.220.18
Five-day mean precipitation maps for the original and trained nets against IMERG, their differences, and zonal means
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.
RMSE and correlation versus lead day and daily precipitation maps at two leads
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 mean precipitation for original and trained nets against IMERG
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:

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

What more GPU would buy

Questions for discussion

Where a computing-science view would change what we do.

  1. 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?
  2. 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?
  3. 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?
  4. 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?
  5. 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?
  6. 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.