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.