ConceptioArchivearXiv CS
arXiv CSopen access

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

Unknown · 2026 · arxiv_cs
arXiv CS · Papers · License: Open Access · 2026
Open Source ↗Direct PDF ↓
machine learning, deep learning, neural networks

2026-6-30

Diffusion Fine-tuning with Rewarded Moment Matching Distillation Alexis Jacq*,1 , Guillaume Couairon*,1 , Valentin De Bortoli1 , Quentin Berthet1 , Arnaud Doucet1 and Romuald Elie1

arXiv:2606.30414v1 [cs.LG] 29 Jun 2026

* Equal contributions, 1 Google DeepMind

Distillation and Reinforcement Learning (RL) fine-tuning are the primary pillars of diffusion post-training. While traditionally studied in isolation, the interaction between these phases remains poorly understood, and in particular how fine-tuning impacts the generative quality of distilled models. We introduce Rewarded Moment Matching Distillation (RMMD), a novel framework that simultaneously distills diffusion models and maximizes a reward function. RMMD preserves the high-fidelity “naturalness” characteristic of advanced distillation (such as 8-step Moment Matching) by adapting the sampling loop for on-policy training and repurposing the distillation loss as a proxy for integral KL regularization. By evaluating the FID-Reward Pareto fronts on ImageNet, we demonstrate that RMMD achieves superior trade-offs compared to single-step baselines (DI++) and multi-step competitors (DRaFT, HyperNoise). Finally, we apply RMMD to GenCast, a state-of-the-art weather forecasting model, to distill it while optimizing the Continuous Ranked Probability Score (CRPS) metric. The resulting distilled model achieves a 7.5× speedup while outperforming the teacher model on 93% of target weather variables, and being better calibrated. This proves that RMMD scales to complex, high-dimensional scientific domains.

1. Introduction Diffusion models (Ho et al., 2020; Song et al., 2021b) have become the backbone of high-fidelity image synthesis (Esser et al., 2024; Hoogeboom et al., 2025), largely because of their remarkable empirical performance and the fact that their training objective reduces to a stable regression problem over corrupted data. At inference time, however, generating a single sample requires iterating a learned denoiser dozens to hundreds of times, severely limiting throughput. Distillation compresses this into one or a few steps (Boffi et al., 2024; Geng et al., 2025c; Salimans et al., 2024; Song et al., 2023), while reward fine-tuning steers the model toward downstream objectives such as human preferences or safety criteria (Black et al., 2024; Clark et al., 2024; Fan et al., 2023; Uehara et al., 2024; Xu et al., 2023). Both are desirable: a practical model must be fast and aligned. Combining the two is non-trivial. Merging distillation and reward fine-tuning into a single training phase (Li et al., 2024; Luo, 2024; Ren et al., 2024) is appealing but fragile: reward maximization shifts generated samples out of the teacher’s input distribution, progressively invalidating the distillation signal. Fine-tuning an already distilled model avoids this drift, but optimizing a reward over a multistep chain is memory-intensive, and truncating backpropagation introduces bias. DRaFT-1 (Clark et al., 2024) is a method that differentiates only through the final step. ReFL (Xu et al., 2023) differentiates through a random step, but evaluates the reward on a still-noisy latent, giving an inaccurate gradient. HyperNoise (Eyring et al., 2025) sidesteps chain backpropagation entirely by perturbing the initial noise, but is structurally confined to low-frequency image changes. We propose Rewarded Moment Matching Distillation (RMMD), a two-phase procedure that connects distillation and fine-tuning in a principled way. In the first phase, we distill the base model via Moment-Matching Distillation (MMD) (Salimans et al., 2024), which matches intermediate denoising distributions along the sampling trajectory, yielding a multi-step student that closely tracks © 2026 Google DeepMind. All rights reserved

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

the teacher’s marginals. This model is then frozen as a stable distributional reference. In the second phase, we fine-tune the student with a single-step gradient on corrupted on-policy samples: we noise a current student sample to an intermediate timestep, take one denoising step, and evaluate the reward, capturing the full frequency range without multi-step unrolling. The central contribution of RMMD is recycling the moment-matching loss as a regularizer during fine-tuning. Rather than a heuristic penalty, this directly penalizes the discrepancy between the student’s intermediate denoising distributions and those of the frozen reference, giving a principled, interpretable knob (the regularization weight) to trade off reward against distributional fidelity. Finally, we apply Rewarded Moment Matching to distill a diffusion-based weather model, GenCast (Price et al., 2023). GenCast is sampled with 59 backbone evaluations (NFEs) for a single 12h forecast, a high cost which motivates a distillation approach. We apply RMMD to GenCast by considering the CRPS scoring rule as a reward function, which is a distance between the marginals of the generative distribution and the ground truth distribution. We show that optimizing CRPS at a 12 hours lead time provides consistent improvements of the model up to 7-day forecasts, and a much better dispersion of forecasts compared to MMD distillation alone. Our main contributions are: 1. The RMMD procedure, a novel method for diffusion distillation while maximizing a reward function; 2. empirical results demonstrating that RMMD consistently outperforms DRaFT, HyperNoise, and DI++ on FID–Reward Pareto fronts across diverse reward functions; and 3. using RMMD to improve GenCast, a state-of-the-art diffusion-based weather forecasting model. With RMMD, the distilled model is 7.5 times faster than the teacher model while being more accurate on 93% of the forecast weather variables.

2. Background and notations 2.1. Forward process and inference A forward diffusion process progressively corrupts clean data 𝑥0 ∈ ℝ𝑑 ∼ 𝑝0 from a target data distribution 𝑝0 , obtaining purely Gaussian noise 𝑥1 ∼ N (0, 𝐼𝑑 ) according to a schedule such that, conditional upon 𝑥0 , 𝑥𝑡 = 𝛼𝑡 𝑥0 + 𝜎𝑡 𝜀 ∼ 𝑝noise ( 𝑥𝑡 | 𝑥0 ) for 𝑡 ∈ (0, 1] with 𝜀 ∼ N (0, 𝐼𝑑 ) and 𝛼𝑡 and 𝜎𝑡 defining the signal-to-noise ratio SNR𝑡 = 𝛼2𝑡 /𝜎𝑡2 , where 𝛼0 = 1, 𝜎0 = 0 and 𝛼1 = 0, 𝜎1 = 1. We denote by 𝑝𝑡 the marginal distribution of 𝑥𝑡 . In diffusion modeling, a neural network Ψ is trained to predict the denoised data 𝔼[ 𝑥0 | 𝑥𝑡 ], given a timestep 𝑡 ∈ (0, 1] and input 𝑥𝑡 . This is equivalent to learning the score 𝑠 ( 𝑡, 𝑥 ) = ∇𝑥 log 𝑝𝑡 ( 𝑥 ) = ( 𝛼𝑡 𝔼[ 𝑥0 | 𝑥𝑡 = 𝑥 ] − 𝑥 )/𝜎𝑡2 . In this work, a network trained in this fashion plays the role of a teacher network. Diffusion modeling leverages this denoising network via the generative process, which consists in sampling 𝑥1 ∼ N (0, 𝐼𝑑 ), and iteratively for 𝐾 steps 𝑡 = 𝐾𝛿, ..., 𝛿 for 𝛿 = 1/ 𝐾 , denoising (𝑡 ) (𝑡 ) (𝑡 ) (𝑡 ) 𝑥ˆ0 = Ψ ( 𝑥𝑡 , 𝑡 ) and sampling 𝑥𝑡 − 𝛿 ∼ 𝑝cond ( 𝑥𝑡 − 𝛿 | 𝑥𝑡 , 𝑥ˆ0 ) = N ( 𝑥𝑡 − 𝛿 | 𝜇 𝑡 ( 𝑥𝑡 , 𝑥ˆ0 ) , Σ𝑡 ) where Σ𝑡 and 𝜇 𝑡 ( 𝑥𝑡 , 𝑥ˆ0 ) are functions of 𝛼𝑡 , 𝛼𝑡 −𝛿 , 𝜎𝑡 and 𝜎𝑡 −𝛿 (Song et al., 2021a). Considering the model Ψ as a teacher model, this mechanism induces a distribution 𝑝teacher for synthetic data, that we can use as a target in distillation. 2.2. Distributional distillation Backward sampling in diffusion models requires 𝐾 ≫ 1 steps, motivating distillation. A student model Φ𝜃 , defined as a deterministic function of 𝑥𝑡 , a timestep 𝑡 , and auxiliary noise 𝜉 ∼ 𝑞 ( 𝜉) (e.g. for 2

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

dropout), is trained so that its output 𝑥ˆ0 = Φ𝜃 ( 𝑥𝑡 , 𝑡, 𝜉) matches the teacher posterior sampled with the original many-step diffusion schedule obtained with Ψ: 𝑥ˆ0 | 𝑥𝑡 ∼ 𝑝teacher (·| 𝑥𝑡 ) .

(1)

The quality of the distilled model depends on how well 𝑥ˆ0 | 𝑥𝑡 matches the target distribution 𝑝teacher (·| 𝑥𝑡 ). If there is an exact equality of distributions for 𝑡 = 1, then the student can perfectly match the teacher in one step. Otherwise, multi-step sampling for a number of steps 𝐾𝑠𝑡𝑢𝑑𝑒𝑛𝑡 ≪ 𝐾 is often beneficial, as 𝑥ˆ0 in Moment Matching Distillation (Salimans et al., 2024) presented below. We write 𝜕𝜕𝜃 as shorthand 𝜕Φ𝜃 ( 𝑥𝑡 ,𝑡,𝜉 ) for , and use sg[·] for the stop-gradient operator. 𝜕𝜃 Diff-Instruct and DI++. Diff-Instruct (Luo et al., 2023) fixes 𝑡 = 1 and trains a single-step student by minimizing an integral KL divergence. Given 𝑥1 ∼ N (0, 𝐼𝑑 ), the student predicts 𝑥ˆ0 ∼ 𝑝𝜃 (ˆ 𝑥0 | 𝑥1 ); a re-noised sample 𝑥 𝑠′ ∼ 𝑝noise ( 𝑥 𝑠′ |sg[ˆ 𝑥0 ]) is then drawn at a uniformly sampled 𝑠 ∈ [0, 1]. Marginally, we have 𝑥 𝑠′ ∼ 𝑝𝜃 ( 𝑥 𝑠′ ). The loss and its gradient are defined as   LDI ( 𝜃) = 𝔼𝑠 𝑤 ( 𝑠) KL 𝑝𝜃 ( 𝑥 𝑠′ ) ∥ 𝑝teacher ( 𝑥 𝑠′ ) ,   ⊤  𝜕 𝑥ˆ0 ′ ′  ′ 𝑠student ( 𝑠, 𝑥 𝑠 ) − 𝑠teacher ( 𝑠, 𝑥 𝑠 ) , ∇𝜃 LDI ( 𝜃) = 𝔼𝑠,𝑥𝑠 𝑤 ( 𝑠) 𝜕𝜃

where 𝑠student and 𝑠teacher are the score functions of the student and teacher respectively. DI++ (Luo, 2024) augments this with a differentiable reward term:    ⊤  𝜕 𝑥ˆ0 ′ ′ ′ 𝑥0 ) , ∇𝜃 LDI++ ( 𝜃) = 𝔼𝑠,𝑥1 ,𝑥ˆ0 ,𝑥𝑠 𝑤 ( 𝑠) 𝑠student ( 𝑠, 𝑥 𝑠 ) − 𝑠teacher ( 𝑠, 𝑥 𝑠 ) − 𝜆 ∇𝑥ˆ0 𝑅 (ˆ 𝜕𝜃

where 𝜆 balances reward and KL regularization. The student may optionally be pre-distilled with Diff-Instruct, after which the KL term acts as a classical regularizer keeping the generator close to its initial distribution. Moment-Matching Distillation (MMD). MMD extends distillation to arbitrary timesteps 𝑡 ∼ Uniform[0, 1] by matching generalized moments of student and teacher posteriors (Salimans et al., 2024):   LMMD ( 𝜃) = 𝔼𝑡,𝑠,𝑥𝑡 ,𝑥ˆ0 ,𝑥𝑠′ 𝑥ˆ0⊤sg 𝑚student ( 𝑠, 𝑥 𝑠′ ) − 𝑚teacher ( 𝑠, 𝑥 𝑠′ ) , where the expectation is w.r.t. 𝑡 ∼ Uniform[0, 1], 𝑥𝑡 ∼ 𝑝𝑡 , a student prediction 𝑥ˆ0 ∼ 𝑝𝜃 (ˆ 𝑥0 | 𝑥𝑡 ), and a re-noised sample 𝑥 𝑠′ ∼ 𝑝cond ( 𝑥 𝑠′ | 𝑥𝑡 , sg[ˆ 𝑥0 ]) at 𝑠 ∈ [ 𝑡 −𝛿student , 𝑡 ). Here 𝑚teacher ( 𝑠, 𝑥 𝑠′ ) = 𝔼 𝑝 ( 𝑥0 | 𝑥𝑠′ ) [ 𝑥0 | 𝑥 𝑠′ ] is predicted by the teacher and 𝑚student ( 𝑠, 𝑥 𝑠′ ) = 𝔼 𝑝𝜃 (ˆ𝑥0 | 𝑥𝑠′ ) [ˆ 𝑥0 | 𝑥 𝑠′ ] by an auxiliary network trained ′ alongside the student to predict 𝑥ˆ0 from 𝑥 𝑠 . We detail the auxiliary’s loss in AppendixA.1. Since 𝑥 𝑠′ uses a stop-gradient on 𝜃, the loss gradient is  ⊤  𝜕 𝑥ˆ0 ′ ′  ′ ∇𝜃 LMMD ( 𝜃) = 𝔼𝑡,𝑠,𝑥𝑡 ,𝑥ˆ0 ,𝑥𝑠 𝑚student ( 𝑠, 𝑥 𝑠 ) − 𝑚teacher ( 𝑠, 𝑥 𝑠 ) . (2) 𝜕𝜃

When 𝑚student = 𝑚teacher everywhere, student and teacher marginals coincide (Salimans et al., 2024). Training specializes to step size 𝛿student , requiring a number of sampling steps of 𝐾student = 1/𝛿student ≪ 𝐾 for the student, using DDPM as for the teacher. We refer the reader to Salimans et al. (2024) for a complete presentation of MMD.

3

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

L2

sampling using

Reward Pure noise

MMD diffusion distribution trained student parameters frozen parameters stop-gradient, detached vector Loss function

pol

pol

Figure 1 | RMMD: an on-policy student sample 𝑥0 is re-noised to 𝑥𝑡 , from which the student predicts 𝑥ˆ0 . The loss combines a reward term 𝑅 (ˆ 𝑥0 ) with the moment-matching regularizer LMMD and the L2 regularizer LL2 .

3. Rewarded Moment Matching 3.1. From off-policy to on-policy reward optimization. A natural baseline is to add a reward term directly to the MMD gradient (eq. (2)):  ⊤   𝜕 𝑥ˆ0 ′ ′ ∇𝜃 Lnaive ( 𝜃) = 𝔼𝑡,𝑠,𝑥𝑡 ,𝑥ˆ0 ,𝑥𝑠′ 𝑚student ( 𝑠, 𝑥 𝑠 ) − 𝑚teacher ( 𝑠, 𝑥 𝑠 ) − 𝜆 ∇𝑥ˆ0 𝑅 (ˆ 𝑥0 ) . 𝜕𝜃

This is analogous to DI++, but introduces two problem: First, 𝑥ˆ0 is predicted from a re-noised off-policy data point 𝑥 𝑠′ , so an extra dataset is needed and fine-tuning is capped by the best reward achievable on the data. Second, as the student distribution shifts toward high-reward outputs, the data-marginal 𝑝𝑡 becomes an increasingly poor approximation of the student’s own intermediate distribution. pol

𝑥𝑡

We address both issues by sampling 𝑥𝑡 not from the noisy data 𝑝𝑡 but from noisy on-policy samples ∼ 𝑝noise ( 𝑥𝑡 |sg[˜ 𝑥0 ]) where 𝑥˜0 ∼ 𝑝𝜃 is obtained by sampling the student with 𝐾student steps. With this marginal, the reward objective now directly approximates the student’s expected reward:       pol LReward ( 𝜃) = 𝔼𝑡, 𝑥 pol , 𝜉 𝑅 Φ𝜃 ( 𝑥𝑡 , 𝑡, 𝜉) = 𝔼𝑥ˆ0 𝑅 𝑥ˆ0 ≃ 𝔼𝑥˜0 ∼ 𝑝𝜃 𝑅 (˜ 𝑥0 ) 𝑡

The quality of the approximation depends on the quality of the distilled model in Equation (1). With the stop-gradient on 𝑥˜0 , the gradient is given by  ⊤  𝜕 𝑥ˆ0 ∇𝜃 LReward ( 𝜃) = 𝔼𝑥ˆ0 ∇𝑥ˆ0 𝑅 (ˆ 𝑥0 ) . (3) 𝜕𝜃

3.2. RMMD objective We combine the on-policy reward gradient (Equation (3)) with the on-policy version of MMD regularization (Equation (2)), yielding the on-policy gradient  ⊤   𝜕 𝑥ˆ0 ∇𝜃 LRMMD ( 𝜃) = 𝔼𝑡,𝑠, 𝑥 pol , 𝑥ˆ ,𝑥 ′ 𝑚student ( 𝑠, 𝑥 𝑠′ ) − 𝑚teacher ( 𝑠, 𝑥 𝑠′ ) − 𝜆 ∇𝑥ˆ0 𝑅 (ˆ 𝑥0 ) . (4) 𝑡

0

𝑠

𝜕𝜃

4

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

The reference moment 𝑚teacher can be replaced by the frozen auxiliary model learned during the initial distillation phase. The process is illustrated in Figure 1. In contrast to DI++, this regularization is not equivalent to an integral KL divergence unless 𝑥 𝑠′ is drawn from the forward diffusion (rather than the conditional), in which case LMMD recovers the score-matching objective of DI as shown in Salimans et al. (2024). On top of the MMD distillation objective and reward objective, we also add a regularization loss that penalizes the student if its predictions are significantly different from the MMD-distilled model Φ𝜃0 , with the loss   pol pol LL2 ( 𝜃) = 𝔼𝑡,𝑥 pol ,𝜉 || Φ𝜃 ( 𝑥𝑡 , 𝑡, 𝜉) − Φ𝜃0 ( 𝑥𝑡 , 𝑡 )|| 2 , 𝑡

where Φ𝜃0 is used without noise 𝜉. Finally, we extend the RMMD objective (4) to reward functions that depend on multiple samples, (𝑖) e.g. for measuring diversity. We can sample (ˆ 𝑥0 ) 𝑖 and combine the associated MMD regularization (1) (2) losses with 𝑅 (ˆ 𝑥0 , 𝑥ˆ0 , ..).

4. Related Work Diffusion model distillation. A large body of work has focused on reducing the number of sampling steps required by diffusion models. Consistency Models (Song and Dhariwal, 2024; Song et al., 2023) and extensions like Consistency Trajectory Models (Kim et al., 2024) and Shortcut Models (Frans et al., 2025) enforce self-consistency along the probability-flow ODE, enabling one- or few-step generation. MeanFlow (Geng et al., 2025b) learns an average velocity field for one-step generation, and the improved iMF variant addresses training stability and recovers inference flexibility (Geng et al., 2025a). Multistep Consistency Models (Heek et al., 2024) extend Consistency Models by matching intermediate denoising marginals, yielding high-quality multi-step students that closely track the teacher. Score-distillation methods such as Diff-Instruct (Luo et al., 2023) minimize an Integral KL divergence between teacher and student score functions. In this work, we use Moment Matching Distillation (Salimans et al., 2024), which is a multi-step distillation method based on stochastic sampling that provides samples with image quality similar to the teacher when using 8 steps. Reward fine-tuning of diffusion models. Several methods directly optimize a differentiable reward along the sampling trajectory. DI++ (Luo, 2024) builds on top of the Diff-Instruct framework and also supports reward fine-tuning from a pre-distilled initialization. DRaFT (Clark et al., 2024) backpropagates through 𝐾 denoising steps, while DPOK (Fan et al., 2023) and DDPO (Black et al., 2024) cast the chain as a Markov decision process and apply policy-gradient methods. ReFL (Xu et al., 2023) backpropagates through a single randomly chosen step, which can be suboptimal as the reward is evaluated on a noisy intermediate prediction. Implicit Diffusion (Marion et al., 2025) frames fine-tuning as stochastic optimal control, enabling gradient computation through stochastic samplers. All of these methods operate on undistilled models and therefore incur the full multi-step inference cost during training. Joint distillation and reward fine-tuning. RG-LCD (Li et al., 2024), Hyper-SD (Ren et al., 2024), and DI++ (Luo, 2024) augment distillation training with a reward term, using the distillation objective as an implicit regularizer. A shared limitation is that online fine-tuning causes generated samples to drift outside the teacher’s distribution, gradually degrading the distillation signal. RewardInstruct (Luo et al., 2025) sidesteps distillation altogether by directly fine-tuning a few-step teacher to maximize a reward, a strategy sensitive to the expressiveness of the reward function. Fine-tuning distilled models. The closest line of work to ours fine-tunes an already distilled model in a separate phase. HyperNoise (Eyring et al., 2025) trains a lightweight network to perturb 5

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

noise inputs so that they follow a reward-shifted distribution, avoiding the memory cost of multi-step backpropagation. Because the perturbation acts only at the noise level, HyperNoise is constrained to low-frequency modifications, which can be insufficient for reward functions affecting high frequency structure. Our work shares the two-phase spirit of HyperNoise but differs in two key respects: we backpropagate through single-step predictions on corrupted on-policy samples rather than perturbing noise inputs, and we explicitly regularize fine-tuning with the moment-matching loss, providing a principled connection between the distillation and fine-tuning phases.

5. Experiments In this section, we first evaluate RMMD on ImageNet (Deng et al., 2009) with the U-Vit backbone from Simple Diffusion (Hoogeboom et al., 2023, 2025) and some selected reward functions. We justify the use of a multi-step regime for fine-tuning by comparing RMMD with DI++ (Luo, 2024), and then compare our method to other existing multi-step such as DRaFT (Clark et al., 2024) and HyperNoise (Eyring et al., 2025) on selected reward functions. Finally, we use RMMD to distill and improve the state-of-the-art diffusion-based weather model GenCast (Price et al., 2023). 5.1. Fine-tuning and evaluation details Distillation. The first stage of RMMD is MMD distillation (Salimans et al., 2024) without any reward optimization. We use MMD with 8 sampling steps, which we found to be a very strong baseline, and tuned some hyper-parameters to start from the strongest possible distilled model. For the fine-tuning phase, we keep the same hyper-parameters as for distillation (same batch sizes, optimizers, learning rates) except that we use 10,000 steps and decay the learning rate to zero. We optimize the combined online − 𝜆𝑅 + 𝜆 loss L = LMMD reg LL2 , for different values of 𝜆 and fixing 𝜆 reg = 𝜆 /2. Evaluation. Since there can be a trade-off between the two objectives of optimizing a reward function and preserving image quality, all our evaluations are based on FID-vs-Reward Pareto fronts. We fine-tune models with 𝐽 reward factors { 𝜆 𝑗 } 𝐽𝑗=1 and evaluate them every 2500 steps during training, evaluating the FID (Heusel et al., 2017) and Reward on 50, 000 generated samples. The FID-Reward Pareto front is obtained as the set of Pareto-optimal (FID, Reward) evaluation results. Reward functions. We evaluate RMMD with simple reward functions to validate the method, and then focus on more challenging real-case scenarios like weather forecasting. The black-andwhite reward is given by the pixel-wise distance between an image and its black-and-white version, averaging over the color channels. The Laplacian smoothing reward is given by the average distances between a pixel and its four neighbours (up, left, down, right). IS reward directly uses the Inception Score (Barratt and Sharma, 2018) as a reward, and CLIP-red is the CLIP (Radford et al., 2021) alignment score with the word “red”. 5.2. Strong distilled student with 8-steps sampling In the first stage, we employ models distilled using 8-step sampling, which provides an optimal balance between inference speed and image quality, evidenced by an FID score of 1.26 on ImageNet 64 × 64, closely matching the teacher’s score of 1.19. Furthermore, multi-step distillation significantly benefits RMMD; we validate this in Figure 2 by comparing our approach against the 1-step DI++ method. For a fair comparison, we trained the strongest possible 1-step baseline, achieved by distilling a model into 2 steps via MMD before further distilling it into 1 step with DI++ (FID 2.65). As shown in Figure 2, our 8-step RMMD consistently outperforms the 1-step DI++ alternative.

6

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

smooth

7.5

5.0

2.5

2.5 0.175 0.200 Reward

5.0

FID

5.0

5.0

is

7.5

FID

FID

7.5

bw

FID

clip_red

RMMD (8 steps) DI++ (1 step)

2.5

2.5 0.05 Reward

0.00

0.10 0.05 Reward

100 150 Reward

Figure 2 | The advantage of multi-step sampling: FID-Reward Pareto obtained with RMMD (8 steps) and DI++ (1 step). 5.3. Comparing to other multi-step fine-tuning methods We compare the reward fine-tuning stage of RMMD with HyperNoise (Eyring et al., 2025) and DRaFT (Clark et al., 2024), at both small (64 x 64) and large (512 x 512) image resolution, and both 2-step and 8-step regimes, in Figure 3. For DRaFT, we found that adding the L2 regularization leads to much better generations than just using LoRA in few-step distilled regimes (this regularization is called “KL regularization” in DRaFT (Clark et al., 2024), since it can be viewed as optimizing an integral KL divergence in non-distilled networks). In all settings, our fine-tuning method led to better Pareto fronts, especially in neural-network based rewards (CLIP alignment and Inception score). 8

7.5

0.05 Reward

clip_red

20

20

bw

30

FID

FID

0.20 Reward

0.25

FID

4

0.10 Reward

smooth

20

10 5

2

0.00

15

10

6

4

2.5

0.00

2

0.05

100

150 Reward

is

5.5 5.0

10 0.05 Reward

is

8

6

5.0

0.175 0.200 Reward

smooth

FID

2.5

bw

FID

5.0

10.0 FID

FID

7.5

RMMD DRaFT-1 DRaFT-1 + L2 Hypernoise

FID

clip_red

0.02 0.01 Reward

4.5

200

300 Reward

Figure 3 | FID-Reward Pareto for different multi-step fine-tuning methods. Top: ImageNet-64 images generated in 8 steps. Bottom: ImageNet-512 images generated in 2 steps. 5.4. Qualitative comparison Figure 4 shows samples obtained by fine-tuning over the same pre-distilled 2-step ImageNet512 model using the different multi-step methods. In each case, we start from the same initial noise 𝑥1 , and denoise using the same random seed, conditioning on the bullfinch bird class. DRaFT-1 adaptations only affect high frequency changes without modifying the content. DRaFT-2 (back-propagating the gradient on the “whole" 2-step sampling) is prone to reward hacking (blindly coloring everything in red for the CLIP alignment with ’red’). HyperNoise tends to move further away from the teacher distribution for similar reward. RMMD reaches the best trade-off between image quality and reward.

7

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

(a) Teacher

(b) DRaFT-1

(c) DRaFT-2

(d) HyperNoise

(e) RMMD (Ours)

Figure 4 | Visualization of fine-tuning behaviors using a 2-step ImageNet512 MMD teacher. Across all examples, we start from the same initial noise 𝑥1 and denoise using identical random seeds and class conditioning. The target rewards are CLIP alignments for the concepts ’red’, ’Picasso’, and ’watercolor’. While DRaFT-1 preserves the original semantic content, it introduces adversarial artifacts that exploit CLIP features (e.g., red shadows, cubic shapes, and patchy textures) and reduces overall sharpness. DRaFT-2 causes drastic distribution shifts and significantly deteriorates image quality. HyperNoise struggles to optimize complex style rewards like Picasso and watercolor, instead defaulting to broad color shifts (blue and white, respectively). In contrast, RMMD (Ours) successfully integrates subtle modifications to increase the reward without significantly deviating from the original data distribution. Further examples are provided in the Appendix.

8

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

5.5. Application of RMMD to Weather forecasting with GenCast distillation In this section, we use RMMD to improve the GenCast weather model, based on diffusion (Price et al., 2023). Weather forecasting consists in predicting the evolution of some physical variables (temperature, geopotential, humidity, wind) over time, starting from an initial condition. A gridded historical dataset of these fields, called ERA5, is publicly available (Hersbach et al., 2020) and is used to train machine learning-based weather models such as GenCast. Given the value of weather variables 𝑥 𝑡 at a given time 𝑡 , GenCast is trained to learn the distribution of these variables 12 hours later 𝑝 ( 𝑥 𝑡+𝛿 | 𝑥 𝑡 , 𝑥 𝑡 − 𝛿 ). This transition distribution is modeled with a conditional diffusion model, requiring 59 function evaluations (NFE) for sampling. To produce forecasts at a longer horizon, the model is rolled out auto-regressively. More details about GenCast are in Appendix C.2. To improve GenCast with RMMD, we use the Continuous Ranked Probability Score (CRPS) scoring rule (Ferro, 2014) as a reward function to optimize. CRPS measures the compatibility of a probabilistic forecast with an observation, and is minimum when the forecast follows the exact same distribution as the observation. It is computed separately for each dimension of the data and then averaged. Given an input state ( 𝑥 𝑡 , 𝑥 𝑡 −𝛿 ), we use sample 𝑀 predictions (ˆ 𝑥 𝑡+𝛿,𝑖 ) 𝑖𝑀=1 , called members, from the generative 𝑡 +𝛿 𝑡 𝑡 −𝛿 𝑡 +𝛿 model 𝑝𝜃 ( 𝑥 | 𝑥 , 𝑥 ). We also denote by 𝑥 the single sample from the ground truth distribution available in the dataset. The sample-based unbiased estimator for CRPS is then defined, for dimension 𝑟 and 𝜏 = 𝑡 + 𝛿, as 1 ∑︁ 𝑀

CRPS𝑟 ({ˆ 𝑥 𝜏,𝑖 } 𝑖𝑀=1 , 𝑥 𝜏 ) =

𝑀

𝑖=1

𝜏,𝑖 |ˆ 𝑥 ( 𝑟 ) − 𝑥 𝜏( 𝑟 ) | −

1

𝑀 ∑︁ 𝑀 ∑︁

2 𝑀 ( 𝑀 − 1) 𝑖=1 𝑗=1

𝜏, 𝑗 𝜏,𝑖 |ˆ 𝑥 ( 𝑟 ) − 𝑥ˆ( 𝑟 ) |

In theory, explicit CRPS optimization would be unnecessary if the generative model was perfectly recovering the true transition distribution 𝑝 ( 𝑥 𝑡+𝛿 | 𝑥 𝑡 , 𝑥 𝑡 −𝛿 ). However, in practice, diffusion models often suffer from under-dispersion, failing to capture the full variance of possible outcomes, and one of the strengths of RMMD is allowing us to address this problem. Experimental details. For our experiments, we use a GenCast teacher that was trained for 500k with a total batch size of 128, on 1 degree resolution maps (resulting in 180×360 latitude-longitude maps) available on GitHub (GoogleDeepMind, 2025). We then run the MMD algorithm for 300k steps with a total batch size of 16, which corresponds to 7.5% of the initial compute budget. We then continue distillation with RMMD (with CRPS), with a second phase of 300k optimization steps at a batch size of 16. CRPS being a multi-sample reward function, we use the 2-sample variant of RMMD presented in section 3.2 and use a global reward weighting of 𝜆 = 0.3. The final model can be sampled with 8 steps of DDPM sampling instead of 59 for the teacher, resulting in a 7.5× speed-up. For evaluation, we take each initialization date of the evaluation year (2018), roll out our model for 15 steps (maximum lead time of 7.5 days) with 𝑀 = 8 members, and compute CRPS separately at each lead time, following standard procedure. Evaluation metrics. Given a reference ground truth trajectory 𝑥 𝑡+𝛿:𝑡+ 𝐾𝛿 , we can measure how well the distribution of a set of samples (ˆ 𝑥 𝑡+𝛿:𝑡+ 𝐾𝛿,𝑖 ) 𝑖𝑀=1 matches the true distribution over trajectories with CRPS. The CRPS is computed separately for each physical variable 𝑣 and lead time 𝜏 = 𝑡 + 𝑘𝛿. In the remainder of the paper, we report relative CRPS improvement of a model over the default GenCast teacher model. These CRPS improvements are averaged over physical variables and lead times 𝜏 = 𝑡 + 𝑘𝛿 for 𝑘 ∈ {1, . . . , 𝐾 }, to give a representative summary of a model’s performance called “Rel. CRPS score”, e.g. in Table 1. We also define the “Win Rate” as the fraction of physical variables and lead times for which the model is better than the teacher model. 9

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

We also use Spread-skill ratio (Fortin et al., 2014) to measure whether the forecasts are correctly calibrated, a value < 1 indicating an under-dispersed forecast and a value > 1 indicating an overdispersed forecast; see Appendix C.1 for a definition. Evaluation of RMMD. We present an overview of our results in Table 1, where we compare the following models to the GenCast teacher (sampled with default parameters and 59 NFE): (i) A model distilled with MMD without any change to the original algorithm; (ii) A second model distilled with MMD by optimizing some hyperparameters (using 𝜂 = 0.5 in DDIM sampling and changing the diffusion noise schedule to 𝜌 = 100). It has an average CRPS improvement of 0.8% and is better than the teacher on 75% of variables.

Name Teacher Plain MMD MMD (best) RMMD w/ CRPS (offline) RMMD w/ CRPS (online)

CRPS improv.(↑) 0% -1.32% 0.82% 1.11% 1.51%

Win rate(↑) N/A 4.9% 75.0% 89.2% 93.0%

Table 1 | Summary table of ablations. Our best MMD model is 7.5× faster than the GenCast model, while being better on 75% of variables. The RMMD with CRPS (online) model is also 7.5× faster and better on 93% of variables.

Two models are distilled with RMMD and CRPS, (i) a first version where 𝑥0𝑡 is sampled offline pol

from the dataset; (ii) A second version where 𝑥0 is sampled from the current policy. This is more costly (about 2×) but is also more effective (+0.4% CRPS improvements, +3.8% Win Rate.). There are two main advantages of the on-policy variant: (i) the network is trained on noisy versions of its own generations, which is closer to the distribution of its inputs during sampling; (ii) For a given ( 𝑥 𝑡 −𝛿 , 𝑥 𝑡 ), there is only one transition 𝑥 𝑡 → 𝑥 𝑡+𝛿 in the training set; on-policy allows us to sample more transitions with the current policy, which can reduce overfitting. RMMD improves model calibration. In Figure 5, we observe that CRPS optimization also improves dispersion (measured by spread-skill ratio) compared to the reference MMD model used for initialization. However surprisingly, the on-policy version, which has better CRPS scores, has slightly worse dispersion compared to the off-policy model. Discussion on the reward objective. CRPS is the metric on which weather forecasts are evaluated at different lead times ( 𝑡 + 𝑘𝛿) 𝑘 ∈ [1..𝐾 ] along trajectories, and it is also the metric that is being optimized with RMMD. Therefore it could be seen as obvious that optimizing CRPS will improve evaluation metrics. However, we only optimize the next-state CRPS at a 12 hours time difference and we observe that the relative gain of the distilled model over the teacher actually increases for larger lead times in auto-regressive rollouts (see Figure 11 in Appendix), and the CRPS at larger lead times is not being optimized. We believe this shows that RMMD improves modeling of the distribution, because reward hacking for the next-state predictions would not yield improvements in auto-regressive rollouts. 5.6. Limitations The quality of generated samples is inherently limited by the performance of the distilled model prior to reward optimization. While RMMD successfully maintains the generation quality of the initial MMD-distilled model, it does not improve FID during the reward optimization phase. This limitation motivated our choice of a highly robust starting point: an 8-step MMD distillation. Because RMMD is an on-policy method (consistent with DRaFT-1 or HyperNoise), it necessitates sampling from the current policy at every training step. By utilizing 8 sampling steps with 10

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

geopotential at 500 hPa

spread-skill ratio

1.00

1.00

0.95

0.95

0.90

0.90

0.85

0.85

geopotential at 850 hPa

temperature at 300 hPa 1.00

Days

Days

u comp. of wind at 500 hPa v comp. of wind at 850 hPa

2m temperature

0.95

0.95

0.95

0.95

0.90

0.90

0.90

0.90

spread-skill ratio

1.00

0 1 2 3 4 5 6 7 8 Days

0.85

0 1 2 3 4 5 6 7 8

GenCast (teacher)

Days

0.85

0.975 0.950 0.925 Days

1.00

0 1 2 3 4 5 6 7 8

MMD (best)

Days

1.00

specific humidity at 700 hPa 1.000

0.90

0.90

1.00

0.85

temperature at 850 hPa

0.95

0.95

Days

1.00

10m u comp. of wind

0.900 1.00

Days

mean sea level pressure

0.95 0.90 0.85 0 1 2 3 4 5 6 7 8 Days

RMMD (on-policy)

0 1 2 3 4 5 6 7 8 Days

RMMD (off-policy)

Figure 5 | Spread-skill ratio of our RMMD-finetuned models. Optimizing CRPS greatly improves dispersion of generated weather states compared to using MMD alone, generally matching or improving GenCast’s dispersion, except for humidity. a stop-gradient on on-policy samples, we observed only a 2× computational slowdown compared to an off-policy variant. This trade-off is justified by a significant gain in accuracy, as measured on GenCast distillation. Finally, RMMD requires a differentiable reward function; it cannot optimize non-differentiable or black-box rewards without the use of a differentiable surrogate.

6. Conclusion We develop a fine-tuning method that leverages a model initially distilled via MMD, and optimizes differentiable rewards while using an on-policy version of the moment-matching loss for regularization. Our empirical evaluations on ImageNet demonstrate that RMMD shows superior trade-offs compared to single-step methods like DI++, ensuring higher generation quality through retained multi-step capabilities. Furthermore, it leads to better Pareto fronts than other multi-step fine-tuning approaches such as DRaFT and HyperNoise across the considered reward functions, ranging from simple pixel-level metrics to complex semantic signals like CLIP alignment. We have applied RMMD to weather forecasting with the distillation of the diffusion-based model GenCast, and demonstrated a 7.5× speed improvement while improving the model’s probabilistic predictions on 93% of variables. RMMD also solved the under-dispersion issue of MMD distillation. This paves the way for better diffusion-based models in science where there is a differentiable objective that can be optimized.

11

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

References S. Barratt and R. Sharma. A note on the inception score. arXiv preprint arXiv:1801.01973, 2018. K. Black, M. Janner, Y. Du, I. Kostrikov, and S. Levine. Training diffusion models with reinforcement learning. International Conference on Learning Representations, 2024. N. M. Boffi, M. S. Albergo, and E. Vanden-Eijnden. Flow map matching with stochastic interpolants: A mathematical framework for consistency models. arXiv preprint arXiv:2406.07507, 2024. K. Clark, P. Vicol, K. Swersky, and D. J. Fleet. Directly fine-tuning diffusion models on differentiable rewards. International Conference on Learning Representations, 2024. V. De Bortoli, A. Galashov, J. S. Guntupalli, G. Zhou, K. Murphy, A. Gretton, and A. Doucet. Distributional diffusion models with scoring rules. International Conference on Machine Learning, 2025. J. Deng, W. Dong, R. Socher, L.-J. Li, K. Li, and L. Fei-Fei. Imagenet: A large-scale hierarchical image database. In 2009 IEEE Conference on Computer Vision and Pattern Recognition. Ieee, 2009. P. Esser, S. Kulal, A. Blattmann, R. Entezari, J. Müller, H. Saini, Y. Levi, D. Lorenz, A. Sauer, F. Boesel, et al. Scaling rectified flow transformers for high-resolution image synthesis. In International Conference on Machine Learning, 2024. L. Eyring, S. Karthik, A. Dosovitskiy, N. Ruiz, and Z. Akata. Noise hypernetworks: Amortizing test-time compute in diffusion models. Advances in Neural Information Processing Systems, 2025. Y. Fan, O. Watkins, Y. Du, H. Liu, M. Ryu, C. Boutilier, P. Abbeel, M. Ghavamzadeh, K. Lee, and K. Lee. Dpok: Reinforcement learning for fine-tuning text-to-image diffusion models. Advances in Neural Information Processing Systems, 2023. C. Ferro. Fair scores for ensemble forecasts. Quarterly Journal of the Royal Meteorological Society, 140 (683):1917–1923, 2014. V. Fortin, M. Abaza, F. Anctil, and R. Turcotte. Why should ensemble spread match the rmse of the ensemble mean? Journal of Hydrometeorology, 15(4):1708–1713, 2014. K. Frans, D. Hafner, S. Levine, and P. Abbeel. One step diffusion via shortcut models. In The Thirteenth International Conference on Learning Representations, 2025. URL https://openreview.net/ forum?id=OlzB6LnXcS. Z. Geng, M. Deng, X. Bai, J. Z. Kolter, and K. He. Improved mean flows: On the challenges of fastforward generative models, 2025a. URL https://arxiv.org/abs/2512.02012. Z. Geng, M. Deng, X. Bai, J. Z. Kolter, and K. He. Mean flows for one-step generative modeling. arXiv preprint arXiv:2505.13447, 2025b. Z. Geng, M. Deng, X. Bai, J. Z. Kolter, and K. He. Mean flows for one-step generative modeling. In Advances in Neural Information Processing Systems, 2025c. GoogleDeepMind.

Google deepmind graphcast and gencast. google-deepmind/graphcast, 2025.

J. Heek, E. Hoogeboom, and T. Salimans. arXiv:2403.06807, 2024.

https://github.com/

Multistep consistency models.

arXiv preprint

12

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

H. Hersbach, B. Bell, P. Berrisford, S. Hirahara, A. Horányi, J. Muñoz-Sabater, J. Nicolas, C. Peubey, R. Radu, D. Schepers, et al. The era5 global reanalysis. Quarterly Journal of the Royal Meteorological Society, 146(730):1999–2049, 2020. M. Heusel, H. Ramsauer, T. Unterthiner, B. Nessler, and S. Hochreiter. Gans trained by a two time-scale update rule converge to a local nash equilibrium. Advances in Neural Information Processing Systems, 30, 2017. J. Ho, A. Jain, and P. Abbeel. Denoising diffusion probabilistic models. Advances in Neural Information Processing Systems, 2020. E. Hoogeboom, J. Heek, and T. Salimans. Simple diffusion: End-to-end diffusion for high resolution images. International Conference on Machine Learning, 2023. E. Hoogeboom, T. Mensink, J. Heek, K. Lamerigts, R. Gao, and T. Salimans. Simpler diffusion (sid2): 1.5 fid on imagenet512 with pixel-space diffusion. Conference on Computer Vision and Pattern Recognition, 2025. T. Karras, M. Aittala, T. Aila, and S. Laine. Elucidating the design space of diffusion-based generative models. Advances in Neural Information Processing Systems, 2022. D. Kim, C.-H. Lai, W.-H. Liao, N. Murata, Y. Takida, T. Uesaka, Y. He, Y. Mitsufuji, and S. Ermon. Consistency trajectory models: Learning probability flow ODE trajectory of diffusion. In The Twelfth International Conference on Learning Representations, 2024. URL https://openreview.net/ forum?id=ymjI8feDTD. J. Li, W. Feng, W. Chen, and W. Y. Wang. Reward guided latent consistency distillation. arXiv preprint arXiv:2403.11027, 2024. W. Luo. Diff-instruct++: Training one-step text-to-image generator model to align with human preferences. Transactions on Machine Learning Research, 2024. W. Luo, T. Hu, S. Zhang, J. Sun, Z. Li, and Z. Zhang. Diff-instruct: A universal approach for transferring knowledge from pre-trained diffusion models. Advances in Neural Information Processing Systems, 36:76525–76546, 2023. Y. Luo, T. Hu, W. Luo, K. Kawaguchi, and J. Tang. Reward-instruct: A reward-centric approach to fast photo-realistic image generation. Advances in Neural Information Processing Systems, 2025. P. Marion, A. Korba, P. Bartlett, M. Blondel, V. De Bortoli, A. Doucet, F. Llinares-López, C. Paquette, and Q. Berthet. Implicit diffusion: Efficient optimization through stochastic sampling. Artificial Intelligence and Statistics, 2025. I. Price, A. Sanchez-Gonzalez, F. Alet, T. R. Andersson, A. El-Kadi, D. Masters, T. Ewalds, J. Stott, S. Mohamed, P. Battaglia, et al. Gencast: Diffusion-based ensemble forecasting for medium-range weather. arXiv preprint arXiv:2312.15796, 2023. A. Radford, J. W. Kim, C. Hallacy, A. Ramesh, G. Goh, S. Agarwal, G. Sastry, A. Askell, P. Mishkin, J. Clark, et al. Learning transferable visual models from natural language supervision. International Conference on Machine Learning, 2021. Y. Ren, X. Xia, Y. Lu, J. Zhang, J. Wu, P. Xie, X. Wang, and X. Xiao. Hyper-sd: Trajectory segmented consistency model for efficient image synthesis. Advances in Neural Information Processing Systems, 2024. 13

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

T. Salimans, T. Mensink, J. Heek, and E. Hoogeboom. Multistep distillation of diffusion models via moment matching. Advances in Neural Information Processing Systems, 2024. J. Song, C. Meng, and S. Ermon. Denoising diffusion implicit models. International Conference on Learning Representations, 2021a. Y. Song and P. Dhariwal. Improved techniques for training consistency models. In The Twelfth International Conference on Learning Representations, 2024. URL https://openreview.net/ forum?id=WNzy9bRDvG. Y. Song, J. Sohl-Dickstein, D. P. Kingma, A. Kumar, S. Ermon, and B. Poole. Score-based generative modeling through stochastic differential equations. International Conference on Learning Representations, 2021b. Y. Song, P. Dhariwal, M. Chen, and I. Sutskever. Consistency models. International Conference on Machine Learning, 2023. M. Uehara, Y. Zhao, T. Biancalani, and S. Levine. Understanding reinforcement learning-based fine-tuning of diffusion models: A tutorial and review. arXiv preprint arXiv:2407.13734, 2024. J. Xu, X. Liu, Y. Wu, Y. Tong, Q. Li, M. Ding, J. Tang, and Y. Dong. Imagereward: Learning and evaluating human preferences for text-to-image generation. Advances in Neural Information Processing Systems, 2023.

A. Additional details For all experiments on ImageNet, we use the U-Vit backbone from Simple Diffusion (Hoogeboom et al., 2023, 2025). Our only modifications are to allow for a dropout rate of 0.1 in all transformer blocks. We use a pixel space based diffusion process, and, for 64 x 64 images, we use a shifted cosine schedule for distillation with a logSNR shift 𝑏 = log(2) (Hoogeboom et al., 2025), which we found to bring notable improvement even with a teacher trained with a symmetric cosine schedule, as we can see in Table 2. With this schedule, the SNR for timestep 𝑡 = 0.5 is equal to 2 instead of 1 for the default cosine schedule. A.1. Training the auxiliary model to predict the student’s moment We follow (Salimans et al., 2024) for training the auxiliary model alongside the student to predict the student’s moment 𝑚student ( 𝑠, 𝑥 𝑠′ ). For an auxiliary model 𝑚student = Φ𝜃aux (initialized with 𝜃0 ), the loss is written:   Lauxiliary ( 𝜃aux ) = 𝔼𝑡,𝑥ˆ0 ,𝑥𝑠′ || Φ𝜃aux ( 𝑥 𝑠′ ) − 𝑥ˆ0 || 2 + || Φ𝜃aux ( 𝑥 𝑠′ ) − Φ𝜃0 ( 𝑥 𝑠′ )|| 2 , where the first term is a regression to predict the student’s generation 𝑥ˆ0 (at noise level 𝑡 ), given the next denoising step 𝑥 𝑠′ . The second term is a L2 regularization to ensure the auxiliary weights stay close to the initial ones. A.2. Hyperparameters Here is the list of hyper-parameters that we use for RMMD, on the ImageNet64, ImageNet512 and ERA5 datasets: 14

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

Architecture Teacher Model Dropout (MMD) Dropout (RMMD) MMD steps RMMD steps (phase 2) Batch size Training hardware fine-tuning samples Data augmentation Optimizer gradient accumulation Learning Rate Reward weight 𝜆 Reg. weight 𝜆 𝑟𝑒𝑔 Noise schedule DDPM epsilon

ImageNet 64x64 U-ViT Simpler Diffusion (SiD2) 0.1 0.1 50k 10k 2048 16 TPU-v5 120M Random hflip Adam( 𝛽1 = 0.9, 𝛽2 = 0.99, 𝜖 = 1𝑒 − 12) 1 1e-5 variable 𝜆 /2 Cosine 1

ImageNet 512x512 U-ViT Simpler Diffusion (SiD2) 0.1 0.0 50k 10k 2048 16 TPU-v5 120M Random hflip Adam( 𝛽1 = 0.9, 𝛽2 = 0.99, 𝜖 = 1𝑒 − 12) 1 1e-5 variable 𝜆 /2 Cosine 1

ERA5 1º Graph Transformer GenCast 0.0 0.0 300k 300k 16 16 TPU-v6 9.6M None Adam( 𝛽1 = 0.9, 𝛽2 = 0.99, 𝜖 = 1𝑒 − 12) 8 1e-7 0.3 1 EDM w/ 𝜌 = 100 0.5

A.3. Comparison against competing methods We evaluate FID-Reward Pareto fronts by sweeping over reward scaling factor 𝜆 . For DRaFT (Clark et al., 2024) with LoRA, we sweep over LoRA weighting coefficients instead of reward factors; for DRaFT with LoRA mixed with L2 regularization 𝜆 𝑟𝑒𝑔 , we fix the best LoRA coefficient and sweep over 𝜆 𝑟𝑒𝑔 ).

B. Additional experiments B.1. Alternative sampling pol

Another strategy to sample pathwise 𝑥𝑡 , inspired by ReFL (Xu et al., 2023) consists in early stopping the denoising process at a random step 𝑡𝑠𝑡𝑜𝑝 ∈ {0, 𝛿student , . . . , (1 − 𝐾student ) 𝛿student }. The advantage is that it only trains at the time steps that matter, i.e. the 𝑥ˇ𝑡 seen at inference are the same as during training. This however limits the generalization of the moments matched by MMD, which performs better on continuous time steps 𝑡 (Salimans et al., 2024). We implemented and compared both sampling methods in Figure 6, and observed slightly better results with early stopping. Unfortunately, since our implementation operates on batches, we had to either use the same stopping time for a whole batch, or do the whole sampling and keep all time steps in memory to randomly draw a batch of different times (the latter solution led to the reported results, but is much more memory expensive). Figure 6 compares the Pareto fronts obtained with the two sampling methods (in both cases we apply both RMMD and L2 regularization). The discrete sampling performs slightly better on the CLIP-based reward and the smoothness, while the continuous sampling works slightly better with the Inception Score reward.

15

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

0.175 0.200 Reward

5.0

FID 0.05 Reward

0.00

is

7.5

2.5

2.5

2.5

2.5

smooth

5.0

5.0

5.0

7.5

7.5 FID

FID

7.5

bw

FID

clip_red

RMMD sampling ReFL sampling

0.10 0.05 Reward

100 150 Reward

pol

Figure 6 | FID-Reward Pareto obtained with methods for sampling 𝑥𝑡 : ReFL (discrete) sampling and continuous sampling. The lower and righter the better. B.2. Evaluating L2 regularization In this section, we evaluate a baseline that uses reward maximization along with L2 regularization, but without the MMD distillation objective. One risk with L2 regularization alone is that it can reduce diversity, since 𝔼[||ˆ 𝑥0,𝜃 − 𝑥ˆ0,𝜃0 || 2 |˜ 𝑥𝑡 ] = Tr(cov[ˆ 𝑥0,𝜃 |˜ 𝑥𝑡 ]) + Tr(cov[ˆ 𝑥0,𝜃0 |˜ 𝑥𝑡 ]) + ||𝔼[ˆ 𝑥0,𝜃 |˜ 𝑥𝑡 ] − 𝔼[ˆ 𝑥0,𝜃0 |˜ 𝑥𝑡 ] || 2 , where 𝑥ˆ0,𝜃 = Φ𝜃 (˜ 𝑥𝑡 , 𝑡, 𝜉) is learned and 𝑥ˆ0,𝜃0 = Φ𝜃0 (˜ 𝑥𝑡 , 𝑡 ) is the target. In our implementation, only the mean of 𝑝𝜃 (ˆ 𝑥0 |˜ 𝑥𝑡 ) is modeled and the variability is obtained by a dropout in the weights, so we avoid this side effect. In a setting where the variance is learned, a regularization based on a scoring rule (De Bortoli et al., 2025) instead could lead to better results. The discrete ReFL sampling combined with only this L2 regularization would be equivalent to the ReFL method introduced in (Xu et al., 2023), fine-tuning over a MMD-distilled model. Figure 7 shows FID-Reward Pareto obtained with the two regularization approaches on 4 different reward functions. These experiments are conducted at image resolution 64 x 64 and at the 8-step regime. In pixel-wise rewards (black-and-white and smoothness), MMD regularization is better than L2 regularization, but the opposite is observed in the CLIP alignment reward. The combination MMD regularization and L2 regularization stays closer to the best Pareto for these three rewards, and outperforms both isolated methods for the Inception score.

2.5

5.0

2.5 0.05 Reward

0.00

is 7.5

5.0

5.0 2.5

0.175 0.200 Reward

smooth

FID

5.0

7.5

7.5 FID

FID

7.5

bw

FID

clip_red

RMMD RMMD w/o L2 L2 only

2.5 0.10 0.05 Reward

100

150 Reward

200

Figure 7 | FID-Reward Pareto for different regularization techniques: Rewarded Moment Matching without L2 regularization, L2 regularization (L2) and RMMD (mix of both moment matching and L2 regularization). The lower and righter the better. Using both forms of regularization improve the FID-Reward Pareto fronts.

16

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

(a) Base

(b) CLIP ’red’

(c) Black and white

(d) Laplacian smooth (e) Inception Score

Figure 8 | Changes, for a similar seed and class, when maximizing the different reward functions, while staying at a decent FID level (FID <= 8) on ImageNet64. B.3. Reward functions Figure 8 displays the resulting changes when maximizing the different reward functions, while staying at an FID below 8 on ImageNet 64. For a given class, the denoising processes start from the same initial noise and sample each step using the same random seed. B.4. Effect of dropout on MMD We report in Table 2 the FID of generations obtained by adding stochasticity in the student’s network prediction via a dropout on our implementation of MMD (only during first distillation phase), as well as shifting the schedule for the 64x64 resolution. MMD (paper) MMD (ours) MMD + dropout MMD + dropout + shift

I64 - 8 steps 1.24 1.35 1.37 1.26

I64 - 2 steps 3.86 2.0 1.66 1.4

I512 - 2 steps 9.7 5.4 -

Table 2 | Effect of dropout and shifted schedule on our implementation of MMD on the FID.

B.5. Empirical comparison of DRaFT, HyperNoise and RMMD We provide more samples obtained with the different fine-tuning methods on Fig. 9 (CLIP alignment with the word ’watercolor’) and Fig. 10 (CLIP alignment with the word ’Picasso’).

C. Additional information on the distillation of GenCast C.1. Definition of metrics In this section, we provide the metrics that we use for evaluating our fine-tuned weather models.

17

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

teacher

DRaFT-1

DRaFT-2

HyperNoise

RMMD (ours)

Figure 9 | More samples of i512 image generated in 2 diffusion steps, using the teacher model, RMMD, DRaFT-1 or DRaFT-2, optimizing for the CLIP alignment with ’Watercolor’

18

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

teacher

DRaFT-1

DRaFT-2

HyperNoise

RMMD (ours)

Figure 10 | More samples of i512 image generated in 2 diffusion steps, using the teacher model, RMMD, DRaFT-1 or DRaFT-2, optimizing for the CLIP alignment with ’Picasso’

19

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

CRPS. Given an input state 𝑥 𝑡 , the generative models predict a distribution 𝑝 ( 𝑥 𝑡+𝛿 | 𝑥 𝑡 ) as a set of 𝑀 samples (ˆ 𝑥 𝑡+𝛿,𝑖 ) 𝑖 and a single sample 𝑥 𝑡+𝛿 from the ground truth distribution is available in the dataset. The CRPS (Ferro, 2014) for dimension 𝑟 and verification time 𝜏 = 𝑡 + 𝑘𝛿 is then defined as 1 ∑︁

1

𝑀

CRPS({ˆ 𝑥 𝜏,𝑖 } 𝑖 , 𝑥 𝜏 ) =

𝑀

|ˆ 𝑥 𝜏,𝑖 − 𝑥 𝜏 | −

𝑖=1

𝑀 ∑︁ 𝑀 ∑︁

2 𝑀 ( 𝑀 − 1) 𝑖=1 𝑗=1

|ˆ 𝑥 𝜏,𝑖 − 𝑥ˆ𝜏, 𝑗 |

The CRPS is then averaged for each initialization time 𝑡 . Ensemble Mean RMSE. The Ensemble Mean RMSE is the error between the ground truth and the average over members. It is defined for each physical variable 𝑣 and lead time 𝑘 as

√︄ EnsMeanRMSE𝑘,𝑣 =

1 ∑︁

∥ 𝑦 𝑡+𝑘𝛿 −

𝑇

1 ∑︁ 𝑡+𝑘𝛿,𝑖 2 𝑥 ∥ 𝑀

𝑡

𝑖

Spread-Skill Ratio. The Spread Skill ratio (Fortin et al., 2014) is another metric to assess whether the set of samples has the correct dispersion or spread. The (bias-corrected) spread-skill ratio is defined as

Spread𝑘 =

√︄

1 ∑︁ 1 ∑︁ 𝑇

𝑀 𝑡

√︂ SpreadSkillRatio𝑘 =

∥ 𝑥 𝑡+𝑘𝛿,𝑖 −

1 ∑︁ 𝑡+𝑘𝛿, 𝑗 2 𝑥 ∥ 𝑀

𝑖

𝑗

Spread𝑘 𝑀+1 𝑀 EnsMeanRMSE𝑘

The reasoning is as follows: under the assumption of a perfect ensemble forecast, the ensemble mean should have the same distance on average to the ground truth or to a random ensemble member, hence the spread-skill ratio should be equal to 1. If it is below one, the spread is too small (relative to the model’s error) and the model is said to be "underdispersive"; if it is above one, the model is "overdispersive". C.2. The GenCast model GenCast is a model trained to predict the transition distribution 𝑝 ( 𝑥 𝑡+𝛿 | 𝑥 𝑡 , 𝑥 𝑡 −𝛿 ) with 𝛿 = 12ℎ. 𝑥 𝑡 contains six variables sampled at different altitude levels in the atmosphere (parameterized by pressure): temperature (T), humidity (Q), Geopotential (Z), and the three components of the wind vector (U, V, W). These variables are sampled at pressure levels ranging from 50 hPa in the upper atmosphere to 1000 hPa close to the Earth’s surface. In addition, 𝑥𝑡 also contains surface variables, temperature at 2-meters (t2m), mean sea-level pressure (msl), u and v component of wind at 10meters (10u and 10v). We use a version of GenCast trained at a 1º resolution, resulting in equirectangular maps of size 180 × 360. With all the variables included, A state 𝑥 𝑡 is of dimension 5,313,600.

20

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

D. Additional experiments of RMMD for GenCast D.1. Score card of RMMD model We analyse the performance of the best RMMD-distilled model separately for each weather variable. In Figure 11, we show the per-variable comparison of the GenCast distilled model compared to the teacher. The distilled model with CRPS is better than the teacher on almost all weather variables and lead times, except humidity at small lead times (<2days) var

z

t

q

2t msl tp12hr

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000

-.1

var -.1

-.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1 -.1

u -.1 -.1 -.1

v

w

-.1

0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

10u 10v

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

Figure 11 | Per-variable CRPS relative improvements of the best RMMD-distilled model, relative to the GenCast teacher. Colorbars are scaled so to a maximum improvement/degradation of +5/-5%.

D.2. Making the strongest MMD baseline Without any change to MMD, the distilled model is worse than the teacher, as we can see in Figure 12. In DDPM, the posterior distribution 𝑝𝑐𝑜𝑛𝑑 ( 𝑥 𝑠 | 𝑥𝑡 , 𝑥ˆ0 ( 𝑡 ) ) is written: √︃ 2 2 𝜖 ˆ, 𝛾𝑠,𝑡 𝐼) 𝑝cond ( 𝑥 𝑠 | 𝑥𝑡 , 𝑥ˆ0 ( 𝑡 ) ) = N ( 𝛼𝑠 𝑥ˆ0 ( 𝑡 ) + 1 − 𝛼2𝑠 − 𝛾𝑠,𝑡 √︄ with

𝜖ˆ = ( 𝑥𝑡 − 𝛼𝑡 𝑥ˆ0 ( 𝑡 ) )/𝜎𝑡

and

𝜎𝑠 𝛾𝑠,𝑡 = 𝜂 𝜎𝑡

1−

𝛼2𝑡 𝛼2𝑠

With 𝜂 = 1 corresponding to the classical DDPM posterior distribution and 𝜂 = 0 being deterministic sampling, which is not compatible with Moment Matching Distillation, since there is no distribution. We found that using 𝜂 = 0.5 brings large improvements, as presented in Table 3. This parameter is used during sampling as well as during computing 𝑥 𝑠′ from 𝑥𝑡 and 𝑥ˆ0 in the training loss of MMD. Increasing this parameter decreases the effect of stochasticity during the sampling loop. We found 21

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

var

z

t

q

2t msl tp12hr

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000

var

.1 .1 .1 .1 .1 .1

u

v

-.1 -.1

w

0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

10u 10v

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

Figure 12 | Per-variable CRPS improvement of an MMD-distilled model, without any change to the MMD parameters, compared to the GenCast teacher. that can see that reducing this parameter reduces dispersion of the forecasts measured by spread-skill ratio. However, decreasing this parameter further did not bring additional gains, which we hypothesize is because MMD relies on 𝑥 𝑠′ being a non-deterministic function of 𝑥𝑡 and 𝑥ˆ0 . We also observed that changing the sampling schedule parameter 𝜌 from 𝜌 = 7 to 𝜌 = 100 (corresponding roughly to a uniform schedule in log-SNR parametrization) during distillation also improves CRPS scores and reduces under-dispersion.

Teacher Plain MMD MMD w/ 𝜂 = 0.5 MMD (best)

CRPS improv.(↑) 0% -1.32% 0.22% 0.82%

Win rate(↑) N/A 4.9% 50.7% 75.0%

Table 3 | Impact of parameter 𝜂 on CRPS improvement. Using 𝜂 = 0.5 improves performance a lot. In addition, using 𝜌 = 100 results in the best score for MMD. D.3. Churn analysis GenCast uses EDM sampling algorithm (Karras et al., 2022), which is a deterministic sampling algorithm with stochastic churn on top. In theory, churn is not needed to sample from the correct distribution. In practice, GenCast uses 𝑆𝑛𝑜𝑖𝑠𝑒 = 1.05, which empirically helps improve the quality of samples. We also found that GenCast without stochastic churn provides under-dispersed forecasts 22

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

(without enough diversity), and that stochastic churn helps to increase the spread of forecasts and therefore their calibration. The effect is quantified in Table 4 and separated per lead time and physical variable in Figure 13. Although RMMD is based on DDPM sampling, we can implement a variant of stochastic churn with the same characteristics as for EDM sampling. We evaluate this variant for MMD only, which motivated us to not use stochastic churn for both MMD and RMMD. Name GenCast Teacher GenCast Teacher w/o churn MMD MMD w/ churn

CRPS improv.(↑) 0% -0.57% 0.82% 0.11%

Win rate(↑) N/A 14.2% 75.0% 50.2%

Table 4 | Churn is helpful for the GenCast teacher but not for the distilled model. geopotential at 500 hPa

spread-skill ratio

1.00

1.00

0.95

0.95

0.90

0.90

0.85

0.85 Days

geopotential at 850 hPa

temperature at 300 hPa 1.00 0.95 0.90

Days

2m temperature

0.925

0.95

0.95

0.90

0.90

0.90

0.90

0.85

0.85

0.85

spread-skill ratio

0.95

GenCast (teacher)

Days

0 1 2 3 4 5 6 7 8 Days

GenCast (teacher) w/o churn

0.950

Days

0.95

0 1 2 3 4 5 6 7 8

1.000

0.90

1.00

Days

specific humidity at 700 hPa 0.975

1.00

0 1 2 3 4 5 6 7 8

temperature at 850 hPa

0.95

Days

u comp. of wind at 500 hPa v comp. of wind at 850 hPa

1.00

1.00

1.00

10m u comp. of wind

0.900 1.00

Days

mean sea level pressure

0.95 0.90 0.85 0 1 2 3 4 5 6 7 8

0 1 2 3 4 5 6 7 8

MMD (best)

MMD (best) w/ churn

Days

Days

Figure 13 | Spread-skill ratio of GenCast with and without inflated churn, and MMD-distilled model with and without inflated churn. We observe that the churn technique used in EDM is ineffective with MMD: It does not improve spread skill ratio and deteriorates CRPS scores. We make the following observations: • First, if we compare MMD to the GenCast teacher without churn, we see that the MMD model is more under-dispersive at short lead times (0-3 days) and better calibrated at longer lead times. • Second, we see that adding churn to the MMD model does not help to improve dispersion. We can also see that it degrades CRPS, in Table 4. • Our method does help to significantly improve spread-skill ratio over MMD (see Figure 5), and we believe it is a better way to fix the under-dispersion issue of GenCast than stochastic churn. D.4. RMMD improvement over MMD In this section, we present the per-variable CRPS improvement of the RMMD-distilled models compared to the MMD checkpoint trained during the first stage. Interestingly, the on-policy version of RMMD improves upon the MMD initialization on every variable, which is not the case of the off-policy version. 23

Diffusion Fine-tuning with Rewarded Moment Matching Distillation

var

z

t

q

2t msl tp12hr

var

z

t

q

2t msl tp12hr

var

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000

u

v

w

0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000

10u 10v

var

u

v

-.1 -.1

w

0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

10u 10v

PL

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000

PL

0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 50 100 150 200 250 300 400 500 600 700 850 925 1000 0.5 1 1.5 2 2.5 3 3.5 4 4.5 5 5.5 6 6.5 7 7.5 Lead time (days)

Figure 14 | CRPS improvements of RMMD-distilled models compared to the MMD-distilled models in the first phase. Top: Off-policy version of RMMD. Bottom: RMMD (on-policy).

24

Record · ID 321848 · SHA-256 43a90af91595133f
Retrieved via Conceptio — every document is proof-bundled with source, license, and retrieval metadata.