Diffusion-Augmented Markov Decision Processes for Maximum Entropy Reinforcement Learning
Authors: Sebastian Sanokowski, Kaustubh Patil, Majid Khadiv
Organizations: Munich Institute of Robotics and Machine Intelligence (MIRMI), Technical University Munich · Practical Project Student Exchange Program, Technical University Munich · MIT World Peace University
Diffusion models provide an expressive framework for sampling from complex, unnormalized distributions. In this work, we extend Maximum Entropy Reinforcement Learning (ME-RL) to diffusion-based policies by introducing Diffusion-Augmented Markov Decision Processes (DA-MDPs). DA-MDPs interpret each reverse-diffusion transition as an individual reinforcement-learning decision, while only the final denoised action is executed in the environment. Our DA-MDPs follow from a principled derivation based on the variational-inference formulation of ME-RL. By augmenting policy and target trajectories with intermediate diffusion variables, we obtain a tractable reverse-KL upper bound via the data-processing inequality. This bound decomposes across denoising transitions, yielding diffusion-augmented variants of soft rewards, value functions, and local policy objectives. This provides a general framework for adapting ME-RL algorithms to diffusion policies while differentiating through only one diffusion transition at a time. We instantiate the framework with PPO, REPPO, and a maximum-entropy extension of WPO. Experiments demonstrate improved continuous-control performance, benefits from additional diffusion steps, and memory-efficient training. On the StackCube and PushT manipulation tasks, DA-MDP methods learn alternative successful strategies from the same initial state and achieve higher success rates and generally higher success-weighted mode entropy than the Gaussian ME-RL baseline. We also demonstrate successful training when using action chunking.
Figures & tables
Figure 1: From environment trajectories to local denoising decisions. (a) Each action produces an environment reward and an entropy contribution before the environment advances. (b) A fresh denoising chain generates an action at each environment state. Dashed purple guides expand the highlighted policy call at st into local DA-MDP decisions. Gray arrows between denoising decisions carry the sampled action into the next augmented state; only the final action at0 changes the environment state. Gold boxes show the local contributions and their infinite-horizon discounted returns.
Figure 2: Action multimodality and behavioral diversity on the Multimodal Agent benchmark. (a) REPPO and DPPO remain unimodal; DA: PPO, DA: REPPO, and DA: WPO assign probability mass to both reward modes. (b) REPPO produces directed motion, and DPPO produces cyclic trajectories; the DA-MDP methods exhibit more diverse trajectories. Figure 11 in the appendix provides the full comparison, including TruDi and state-visitation histograms.
Figure 3: Multimodal manipulation. (a) Illustrative motions. (b) IID StackCube success and success-weighted mode entropy pH : mean ± SD across three training seeds. (c) PushT success, successful CW/CCW fractions of all attempts, and pH : one seed, 95% sampling intervals.
Figure 4: (a) Equal-task-weighted IQM of normalized returns across twelve tasks with 95% stratified-bootstrap confidence intervals. (b) G1 returns at the reported checkpoints (mean ± SD over seeds). Headers give method-specific interaction counts; bold marks the best depth per method. (c) Memory–runtime trade-off between DA: REPPO and TruDi on CheetahRun at K=16 .
Appendix figures & tables15 assets
Supplementary material from the paper’s appendix.
Appendix
Symbol
Meaning
t∈{0,1,…}
Environment-time index in the stationary discounted setting; T is used only for the finite trajectory prefixes that introduce the DPI.
k∈{K,…,1}
Reverse-diffusion index; atK is prior noise and at0 is executed in the environment.
t~,s~t~,a~t~
Augmented index, state, and action: t~=tK+(K−k) , s~t~=(st,atk,k) , and a~t~=atk−1 .
qθ,πθ
Learned reverse policy/kernel and forward diffusion kernel, respectively; π denotes an unnormalized reward target and π its normalized version.
θ
Current actor parameters with respect to which gradients are taken.
θroll
Frozen behavior-policy snapshot that generated the stored rollout states and actions.
Appendix
Table 1: Notation used throughout the paper.
Figure 5: Continuous-control performance across twelve tasks. Curves show IQM returns; shaded regions indicate 95% bootstrap confidence intervals.
Figure 6: Product-Beta forward transitions in one dimension. (a) Varying m shifts the mean at fixed concentration. (b) Increasing κ narrows the distribution at fixed mean. (c) With uk−1=0.8 , ρ controls the influence of the current value on the next latent. At ρ=0 , the transition follows the Beta(2,2) reference distribution; increasing ρ moves the conditional mean toward uk−1 and reduces the variance.
Policy family
Executed action
Terminal log-Jacobian in actor/critic
Entropy used for temperature tuning
Target factors c
tanh , latent entropy
a=tanh(x0)
Omitted
HX
−4,−5,−6
tanh , transformed entropy
a=tanh(x0)
Included at k=1
HX+E[jtanh(x0)]
−5,−6,−7
Product-Beta / Jacobi
a=2u0−1
Constant dlog2
Bounded path-space entropy
−4,−5,−6
Appendix
Table 2: Policy families in the action-transform ablation. The Gaussian variants differ in the entropy terms used by the actor, critic, and temperature estimator. The transformed-action formulation includes the terminal tanh log-Jacobian in all three components, whereas the latent formulation omits it.
Figure 7: Action-transform and entropy-measure ablation. Curves show mean evaluation return over four training seeds, with 95% bootstrap confidence intervals. The aggregate weights the five tasks equally and bootstraps seeds within each task. Gaussian policies use ODE sampling; Product-Beta uses stochastic reverse transitions and a learned forward schedule initialized from the cosine coefficients and constrained by a forward chain-KL bound of 0.08 .
Figure 8: StackCube spatial evaluation: 540 positions, 512 rollouts per position. Rows show strict success p , blue-on-red fraction q among first strict successes, and success-weighted entropy pH . Bottom-row hue encodes q and brightness scales as pH ; black means pH=0 . Gray denotes no successes; cyan crosses flag fewer than 25 successes. Hatching excludes overlaps. Circles, stars, and squares mark initial end-effector, blue-cube, and red-cube positions. Budgets differ across methods.
Figure 9: Training success for all manipulation runs, grouped by task: StackCube above and PushT below. Lines show trailing 1 M-interaction means; faint traces show raw logs and markers identify final measurements. Original and resumed stages share a physical-interaction axis. Curves use randomized training starts and are not fixed-state multimodality evaluations. All three seeds are shown separately; panel budgets differ.
Figure 10: G1 diffusion-step ablation: deterministic evaluation learning curves for DA: REPPO, DA: PPO, and DA: WPO (left to right), using seven, eight, and six training seeds. Colors indicate K ; bands show ±1 sample standard deviation, without temporal smoothing. Panels share a return scale and use method-specific interaction ranges. Endpoint markers correspond to the table in Fig. 4(b) : 50.07 M, 200.02 M, and 69.07 M interactions. WPO is displayed through 70 M. Evaluation follows prior and reverse-kernel means and excludes entropy bonuses.
Figure 11: Extended Results on the Multimodal Agent benchmark, which show that DA: WPO and TruDi are also able to cover the multimodal action distribution. Panel (c) complements the action and trajectory visualizations with normalized state-visitation counts; it illustrates why high state entropy alone does not imply multimodal action coverage.
Method
Marginal action- sequence KL ↓
State gap ↓
Diffusion gap ↓
Method-specific trajectory KL ↓
REPPO
927.01125
2.07944
N/A
929.091
DA: REPPO
3.74249±0.00012
0.06882±0.00012
284.603±0.104
288.415±0.104
Appendix
Table 3: Numerical estimates of the DPI-gap decomposition on the Multimodal Agent benchmark. The trajectory KL in the final column is method-specific because only DA-MDP contains auxiliary diffusion variables. These stochastic checkpoints achieve mean per-step environment rewards of 0.9981 (REPPO) and 0.9572 (DA: REPPO).
Comparison
Parameter block
Cosine
Relative error
Local surrogate
Joint
1.0000
1.16×10−16
Local surrogate
Diffusion coefficients
1.0000
2.21×10−16
Forward gradient omitted
Diffusion coefficients
0.2065
1.054
Appendix
Table 4: Comparison of expected gradients at θ=θroll . Relative error is ∥g−gtraj∥2/∥gtraj∥2 .
Figure 12: Controlled NVML process-memory–runtime configuration sweep for the complete K=16 , 4×512 -policy sweep. Point labels denote (samples per minibatch, optimizer updates). Every measured configuration is shown with a filled marker, and lines connect configurations in order of increasing samples per minibatch.
Method
Minibatches
Samples/minibatch
NVML peak memory
Runtime
Relative result
DA: REPPO
32
32,768
3.047 GiB
3.362 s
5.74× less memory
TruDi
16
8,192
17.502 GiB
1.646 s
2.04× faster
Appendix
Table 5: Controlled NVML process-memory–runtime comparison of DA: REPPO and TruDi. DA-MDP minibatches contain diffusion-MDP samples, whereas TruDi minibatches contain environment samples. “Updates” denotes the number of minibatch optimizer steps required to process one epoch corresponding to 131,072 environment transitions.
Parameter
DA: PPO
DA: REPPO
DA: WPO
Rollout and training budget
Environment interactions
50 M
Parallel environments
2048
1024
1024
Environment steps per rollout, per env.
64
Local transitions per rollout, per env.
512
Minibatches per epoch
64
Appendix
Table 6: Common MuJoCo hyperparameters for the JAX training presets.
Parameter
Value
Parallel environments / policy steps per rollout
1024/128
Training budget / episode horizon
250.35 M interactions / 100 steps
Action chunk size / diffusion steps
1/8
Epochs / minibatches per epoch
8/64
Discount / trace parameter
0.9/0.98
Actor / critic hidden dimensions
(512,512,512)/(512,512)
Appendix
Table 7: Hyperparameters of the unchunked DA:REPPO PushT experiment.
Recent advances in reinforcement learning (RL) have achieved great successes by leveraging the multimodality and exploration capability of diffusion policies. Among these approaches, one representative branch focuses on the sampling-based policy optimization. This design enables better exploration capability of the diffusion model, particularly at the beginning of training, but suffer from low exploitation in Q-value information, resulting in a slow policy convergence. Another branch pays attention to gradient-based policy optimization, which sufficiently exploits the gradient of the Q function yet tends to collapse into a unimodal policy with low diversity. To address this issue, we propose CGPO, \textbf{C}ritic-\textbf{G}uided diffusion \textbf{P}olicy \textbf{O}ptimization, which effectively balances exploration and exploitation with the training-free guidance technique integrated into the denoising process of diffusion policy. Concretely, CGPO steers action generation toward high-value regions defined by the critic network and uses the guided actions as regression objectives. In this manner, CGPO reduces the time required to obtain high-quality actions and improves final performance with better balance between the exploration-exploitation tradeoff. We validate the effectiveness of CGPO on 5 MuJoCo locomotion tasks, and CGPO achieves state-of-the-art performance compared with existing diffusion-based RL methods. Notably, CGPO is the first success to incorporate diffusion policy into real-world RL, with its superior performance on Franka robot arm grasping tasks. Our official page is released at https://dingsht.tech/cgpo-webpage.
Diffusion models have recently emerged as expressive policy representations for online reinforcement learning (RL). However, their iterative generative processes introduce substantial training and inference overhead. To overcome this limitation, we propose to represent policies using MeanFlow models, a class of few-step flow-based generative models, to improve training and inference efficiency over diffusion-based RL approaches. To promote exploration, we optimize MeanFlow policies under the maximum entropy RL framework via soft policy iteration, and address two key challenges specific to MeanFlow policies: action likelihood evaluation and soft policy improvement. Experiments on MuJoCo, DeepMind Control Suite and HumanoidBench benchmarks demonstrate that our method, Mean Flow Policy Optimization (MFPO), achieves performance comparable to or exceeding current diffusion-based baselines while considerably reducing training and inference time. Our code is available at https://github.com/dongxiaoyi-xyz/MFPO.
Xiaoyi Dong, Xi Sheryl Zhang, Jian Cheng
C2DL, Institute of Automation, Chinese Academy of Sciences · School of Artificial Intelligence, University of Chinese Academy of Sciences · AiRiA +1
We formulate reinforcement learning (RL) in continuous time with discrete state spaces and possibly arbitrary action spaces via a stochastic control approach, where the state dynamics are modeled as a controlled continuous-time Markov chain (CTMC). We consider policy optimization problems and derive the corresponding policy gradient methods, leading to continuous-time variants of proximal policy optimization (PPO) and group relative policy optimization (GRPO). As a primary application, we develop a complete continuous-time RL framework for fine-tuning score-based discrete diffusion models. The proposed framework enables reward-driven optimization without requiring differentiability on the reward signals. In contrast to the existing GRPO-based approaches that only rely on terminal rewards, our formulation allows intermediate reward or advantage signals to be incorporated throughout the denoising trajectory. Importantly, when specialized to masked diffusion models (MDMs), our framework encompasses a rich class of policy parameterizations over the vocabulary simplex with analytically tractable probability ratios, providing a unified perspective on exploration and policy optimization in MDMs. For masked diffusion large language models (dLLMs), we further propose trajectory subsampling techniques to efficiently estimate computationally prohibitive trajectory likelihoods, reducing the computational cost of computing per-position probability ratios. We showcase the effectiveness of our methods on both low-dimensional entropy-regularized optimization problems and RL post-training of dLLMs on mathematical reasoning and coding tasks.
Zikun Zhang, Jiayuan Sheng, David D. Yao +1
Department of Industrial Engineering and Operations Research, Columbia University, New York, NY 10027