Stage 2: residual diffusion¶
ERA5ResidualDiffusion keeps diffusion while changing the generated variable from absolute wind speed to a signed correction around a dense baseline.
In the main two-stage system, that baseline is the frozen output of the deterministic Stage 1 model:
Read the complete Stage 1 → Stage 2 handoff first.
The implementation can also use ERA5 directly as a portable ablation, but the deterministic checkpoint is the intended stacked workflow.
Residual transform¶
Residuals concentrate near zero but have a rare intense positive tail. A linear mapping would compress useful small corrections. The model uses an odd, invertible asinh transform:
The checked-in preset uses s = 5 m/s and c = 80 m/s. Zero remains exactly zero, small corrections receive useful resolution, and the full tail maps into diffusion space [-1,1].
Inputs, masks, and output¶
On observed pixels, the target is the transformed physical SAR-minus-baseline residual. Outside the SAR swath, the target is zero residual with a weak configured weight. Invalid baseline pixels receive zero loss.
The denoiser receives:
1 noisy residual
+ 24 prepared GEO / ERA5 / geometry / solar / mask channels
+ 1 exact frozen baseline
+ 1 baseline-valid mask
= 27 U-Net input channels
Sampling inverts the residual transform, adds the result to the same baseline in m/s, applies physical wind bounds, and only then maps to normalized wind for image and storm-structure metrics.
Validation reports baseline_mae_ms and mae_skill_vs_baseline, measuring the sampled refinement directly against the frozen field it is supposed to improve.
Probabilistic refinement¶
The deterministic-baseline preset keeps epsilon diffusion as the generative objective and adds:
- Min-SNR weighting across noise levels;
- separately normalized SAR and off-swath losses;
- extra emphasis for inner-core and high-wind pixels;
- weak gradient, spectrum, low-frequency, and total-variation losses at low-to-medium noise; and
- classifier-free guidance through condition dropout.
The spectrum term compares amplitude rather than phase. The low-frequency term keeps members tied to the broad baseline, while total variation suppresses pixel-scale ringing in the correction without smoothing Stage 1 itself. The selected structured-asinh default also adds robust peak, radial-profile, soft exceedance-area, multi-scale, and target-relative annular terms.
Ten percent condition dropout trains an unconditional branch without changing U-Net shape. The checked-in preset samples at guidance_scale: 1.2, the compromise selected from the K=10 sweep; higher values reduce raw maximum-wind bias but narrow ensemble coverage, while lower values preserve more diversity.
Validation ensemble¶
Validation uses four stable latent members on its first reconstruction batch and reports:
- CRPS and ensemble spread;
- pairwise diversity;
- ensemble-mean and best-member MAE;
- gradient sharpness ratio;
- log-spectrum error; and
probabilistic_refinement_score.
Checkpoints use the composite score. Inspect individual members when judging sharpness because averaging plausible alternatives is expected to blur them.
Train Stage 2¶
GEO2WF_BASELINE_CKPT=/path/to/deterministic.ckpt \
uv run geo2wf-train \
data=geo_sar_common10_era5 \
model=residual_diffusion_deterministic_baseline
The deterministic module is loaded as a frozen child, kept in evaluation mode, excluded from the optimizer, and saved with the residual-diffusion checkpoint.
For the ERA5-only ablation:
Continue to Sampling or Evaluation.