MDP & Timestep RL for Discrete Diffusion LLMs
Semi-AR (Block AR); Agentic multi turn; Only terminal/outcome reward, \(\gamma=1\) as the convention for LLM RL.
Setup
Index
Some canvas with idx \(c\) (length \(C\)) includes denoising timestep index (not token index) \([n_c, n_c+T_c)\), \(T_c\) is not constant over \(c\) by early-stopping. Global denoising step \(n\).
From \(n\), a tuple of indices can be specified
\[ n\mapsto (j(n),c(n),t(n)) \]- \(j(n)\): (global) turn index
- \(c(n)\): global canvas index
- \(t(n)\): timestep index in the canvas
thus
\[ n=\sum_{c< c(n)}T_c+t(n). \]State
For some timestep \(n\), define its realized state \(s_n\) (and r.v. \(S\))
\[ s_n=(h_n,x_n,z_n,t(n)) \]- \(h_n\): current fixed context (=prefill), including
prompt,observation(tool output, etc.), previous answers (denoised tokens in previous turns & committed canvases)- Only these and all of these are in KV cache. Equivalent to the content in KV cache.
- \(x_n\in \mathcal {V}^C\): the current canvas, i.e. positions including noises (mask/uniform/…) and partially denoised tokens.
- \(z_n\): Self-conditioning, e.g. in DiffusionGemma the self-conditioning is a logit-based representation.
- \(t(n)\): the local timestep, for scheduling the temperature annealing & stopping criteria?
Model
Token vocabulary \(\mathcal V\). Diffusion language model \(p_\theta\) to denoise over timesteps, or its equivalent policy \(\pi_\theta\).
\[ p_\theta:S\mapsto \Delta^{C}_{\mathcal V} \]which is a multi-multi distribution creating \(C\) prob vectors (for a new canvas \(x_{n+1}\)) over token vocab \(\mathcal V\).
- \(p_\theta^{(i)}\) or \(p_\theta^i\) means the probability of \(i\)-th position by \(p_\theta\), where the superscription is applied over other prob symbols.
Conditional Independence
Assume \(x_{n+1}^i\) is independent to \(x_{n+1}^{j\ne i}\) given \(s_n\). This assumption is vital to make \(x_{n+1}^i\) solely sampled from \(p_\theta^i\).
Action
-
From \(s_n\) to \(s_{n+1}\) (or \(x_n\) to \(x_{n+1}\)), the changes of tokens are three parts: \(U_n\), \(D_n\), \(E_n\). For a pos idx \(i\) in canvas \(x_{n+1}\),
\[ x^i_{n+1}\sim\begin{cases}p^{i}_\theta(\cdot\mid s_n), & i\in U_n,\\=f^i_\theta(s_n), & i\in D_n,\\\kappa(\cdot\mid x^i_n), & i\in E_n.\end{cases} \]- \(U_n\): positions where new tokens are sampled from \(p_\theta\) (regardless whatever old tokens are)
- vanilla covered by policy gradient with a standard score term
- \(D_n\): positions where a \(\theta\)-related non-policy (deterministic) function \(f_\theta\) (e.g. argmax of \(p_\theta\)) from \(s_n\) produces the token. Argmax appears in inference where max time steps are reached or in early stopping; while nondeterministic created tokens are basically \(U_n\) or \(E_n\)
- Nondifferentiable, in training, sampling is applied so token changes are in \(U_n\) and \(D_n=\empty\), but in inference \(D_n\) is nonempty since early stopping/max-time is reached.
- Is self-conditioning this category?
- \(E_n\): where a \(\theta\)-unrelated kernel distribution \(\kappa\) produces the token, which can be discrete uniform \(\kappa=\mathcal U(\mathcal V)\) / setting as
[MASK](\(\kappa=\mathbf 1[x=\text{MASK}]\)) / copying \(x_n^i\) (\(\kappa=\mathbf 1[x=x_n^i]\))
- \(U_n\): positions where new tokens are sampled from \(p_\theta\) (regardless whatever old tokens are)
-
An action is to denoise from \(x_n\) to \(x_{n+1}\), with possible re-masking/noising of tokens, written as a map:
\[ A:\mathcal V^C\to \mathcal V^C \]and its realization
\[ a_n:x_n\mapsto x_{n+1} \]from which we can write probability distribution (vector) as (a realized form):
\[ \begin{aligned} \pi_\theta(a_n|s_n)&=\text{Pr}(x_{n+1}|s_n) \\&=\text{Pr}(U_n,D_n,E_n|s_n)\prod_{i\in U_n}\text{Pr}(x_{n+1}^i|s_n)\prod_{i\in D_n}\text{Pr}(x_{n+1}^i|s_n)\prod_{i\in E_n}\text{Pr}(x_{n+1}^i|s_n) \\&=P_\theta(U_n,D_n,E_n|s_n)\prod_{i\in U_n}p_\theta^i(x_{n+1}|s_n)\prod_{i\in D_n}\mathbf{1}[x_{n+1}=f_\theta^i(s_n)]\prod_{i\in E_n}\kappa(x_{n+1}^i|x_n^i) \end{aligned} \]Choosing \(U_n,D_n,E_n\) also follow \(\theta\) e.g. EB sampler or confidence-based greedy denoising by \(P_\theta\), but it’s piecewise constant so \(\nabla_\theta P_\theta\) is almost zero everywhere; Also it’s true for \(\nabla_\theta f_\theta\). And also their terms are 1 when in realization, so
\[ \pi_\theta(a_n|s_n)=\prod_{i\in U_n}\pi_\theta(x_{n+1}^i|s_n;i)\prod_{i\in E_n}\kappa(x_{n+1}^i|x_n^i) \] -
Thus we have the log policy:
\[ \log \pi_\theta(a_n|s_n)=\sum_{i\in U_n}\log\pi_\theta(x_{n+1}^i|s_n;i)+\underbrace{\sum_{i\in E_n}\log \kappa(x_{n+1}^i|x_n^i)}_{\text{policy-unrelated constant}} \]
Trajectory & Rewarding
Thus a (long-horizon or whatever, agentic or general) trajectory is \(\Tau\), realized to \(\tau\):
\[ \tau=(s_0,a_0,r_1,s_1,a_1,\dots,a_{T-1},r_T,s_T) \]In a group that will be \(\tau_j\), \(s_{j,n}\) and \(x_{j,n}^i\) as examples. Specifically in an outcome rewarding scheme, \(r_T=R(\tau)\) and \(r_{n\ne T}=0\).
\[ \text{Pr}(\tau;\theta)=\text{Pr}(s_0)\prod_{n=0}^{T-1}\pi_\theta(a_n|s_n)\text{Pr}(s_{n+1},r_{n+1}|s_n,a_n) \]Assume the environment dynamics don’t depend on \(\theta\), then
\[ \nabla_\theta\log\text{Pr}(\tau;\theta)=\sum_{n=0}^{T-1}\nabla_\theta\log\pi_\theta(a_n|s_n) \]BTW we define \(R(\tau)=\sum_{n=1}^T \gamma^{n-1} r_n\).
Objective & Policy Gradient
Our objective is, classically,
\[ \max_\theta E_{\Tau\sim\pi_\theta}\left[\sum_{n=1}^TR_n\right]\equiv \max_{\theta}E_{\Tau\sim\pi_\theta}[R(\Tau)] \]and after we define \(G_n=\sum_{t=n+1}^T \gamma^{t-n-1}R_t\), the objective is
\[ \max_\theta E_{\Tau\sim \pi_\theta}[G_0] \]defining \(J(\theta)=E_{\Tau\sim\pi_\theta}[G_0]\), and then by log-derivative trick (expanded below),
\[ \begin{aligned} \nabla_\theta J(\theta)&= \sum_{\tau}R(\tau)\nabla_\theta\text{Pr}(\tau;\theta) \\&= \sum_{\tau}\text{Pr}(\tau;\theta) R(\tau)\nabla_\theta\log \text{Pr}(\tau;\theta) \\&=E_{\Tau\sim\pi_\theta}[R(\Tau)\nabla_\theta\log \text{Pr}(\Tau;\theta)] \\&=E_{\Tau\sim\pi_\theta}[G_0\nabla_\theta\log \text{Pr}(\Tau;\theta)] \\&=E_{\Tau\sim\pi_\theta}[G_0\sum_{n=0}^{T-1} \nabla_\theta\log \pi_\theta(A_n|S_n)] \\&=\sum_{n=0}^{T-1}E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_0}[G_0\nabla_\theta\log \pi_\theta(A_n|S_n)] \end{aligned} \]Since expanding \(G_0\),
\[ G_0=\sum_{t=0}^{n-1}\gamma^t R_{t+1}+\gamma^n G_n, \]we define the reward-to-date \(H_n\) by
\[ H_n=\sum_{t=0}^{n-1}\gamma^tR_{t+1}=G_0-\gamma^nG_n \]where \(E_{A_n\sim\pi_\theta(\cdot|S_n)}[H_n\nabla_\theta\log \pi_\theta(A_n|S_n)]=0\); since \(H_n\) is determined by historical variables before \(A_n\) and \(S_n\),
\[ \begin{aligned} E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_0}[H_n\nabla_\theta\log \pi_\theta(A_n|S_n)]&= E_{A_{0:n-1},S_{0:n},R_{0:n}}[H_nE_{A_n}[\nabla_\theta\log \pi_\theta(A_n|S_n)|A_{0:n-1},S_{0:n},R_{0:n}]\\&=E[H_n\nabla_\theta 1]\\&=0. \end{aligned} \]Thus the policy gradient is
\[ \nabla_\theta J(\theta)=\sum_{n=0}^{T-1}E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_n}[\gamma^nG_n\nabla_\theta\log \pi_\theta(A_n|S_n)] \]with \(\gamma=1\) that’s
\[ \nabla_\theta J(\theta)=\sum_{n=0}^{T-1}E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_n}[G_n\nabla_\theta\log \pi_\theta(A_n|S_n)]. \]Adding Baseline
Consider an expectation where the state is realized:
\[ E_{A_n\sim\pi_\theta(\cdot|s_n)}[b(s_n)\nabla_\theta\log\pi_\theta(A_n|s_n)] \]whose expected value is equal to \(b(s_n)E_{A_n\sim \pi_\theta(\cdot|s_n)} [\nabla_\theta \log \pi_\theta(A_n|s_n)]=0\), for any realization of state. Hence
\[ \nabla_\theta J(\theta)=\sum_{n=0}^{T-1}E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_n}[\gamma^n(G_n-b(S_n))\nabla_\theta\log \pi_\theta(A_n|S_n)] \]as the REINFORCE gradient formula.
Q and V
Define \(Q_n=Q_n(S_n,A_n)=E[G_n|S_n,A_n]\), realized as \(Q_n(s,a)=E[G_n|S_n=s,A_n=a]\).
\[ \begin{aligned} \nabla_\theta J(\theta)&=\sum_{n=0}^{T-1}E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_n}[\gamma^n G_n\nabla_\theta\log \pi_\theta(A_n|S_n)] \\&=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[E[G_n\nabla_\theta\log\pi_\theta(A_n|S_n)|S_n,A_n]] \\&=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[E[G_n|S_n,A_n]~\nabla_\theta\log\pi_\theta(A_n|S_n)] \\&=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[Q_n(S_n,A_n)\nabla_\theta \log\pi_\theta(A_n|S_n)] \\&=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[(Q_n(S_n,A_n)-b(S_n))\nabla_\theta \log\pi_\theta(A_n|S_n)] \end{aligned} \]Also there’s a thing called state-value function, \(V_n(S_n)\) can be the baseline:
\[ \nabla_\theta J(\theta)=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[(Q_n(S_n,A_n)-V_n(S_n))\nabla_\theta \log\pi_\theta(A_n|S_n)] \]Write \(\text{Adv}_n(S_n,A_n):=Q_n(S_n,A_n)-V_n(S_n)\),
\[ \nabla_\theta J(\theta)=\sum_{n=0}^{T-1}\gamma^n E_{A_n\sim\dots,S_n}[\text{Adv}_n(S_n,A_n)\nabla_\theta \log\pi_\theta(A_n|S_n)] \]Critic
Here we also introduce a value estimator network, i.e. critic whose objective is to estimate \(V_n(S_n)\):
\[ \hat V_\phi(S_n)\approx V_n(S_n)=\mathbb E_{\pi_\theta}[G_n\mid S_n] \]For every trajectory, it’s objective is to minimize
\[ L(\phi)=\sum_{n=0}^{T-1}(\hat V_\phi(S_n)-G_n)^2. \]Monte Carlo Estimator
A Monte-Carlo estimator from above is
\[ \sum_{n=0}^{T-1}\gamma^n \text{Adv}_n(S_n,A_n)\nabla_\theta \log\pi_\theta(A_n|S_n) \]such advantage can be estimated by an unbiased MC estimator. Since
\[ E_{A_n\sim\pi_\theta(\cdot|S_n),S_n,G_n}[(G_n-b(S_n))\nabla_\theta\log \pi_\theta(A_n|S_n)]=E_{A_n\sim\dots,S_n}[(Q_n(S_n,A_n)-V_n(S_n))\nabla_\theta \log\pi_\theta(A_n|S_n)] \]the estimator can be also this:
\[ \sum_{n=0}^{T-1}\gamma^n (G_n-b(S_n))\nabla_\theta \log\pi_\theta(A_n|S_n) \]Setting \(b(S_n)\) as \(V(S_n)=E[G_n|S_n]\) and estimating \(V(S_n)\) by a network \(\phi\), i.e. \(\hat V_{n,\phi}(S_n)\) with terminal
\[\hat V_{T,\phi}(S_T):=0\], we create a MC estimator with baseline and \(\phi\):
\[ {\text{MC}}_n (S_n,A_n)=G_n-\hat V_{n,\phi}(S_n) \]For a general advantage estimator \(\widehat{\text{Adv}}_n\),
\[ \hat g=\sum_{n=0}^{T-1}\gamma^n \widehat {\text{Adv}}_n (S_n,A_n)\nabla_\theta \log\pi_\theta(A_n|S_n) \]Temporal Difference
Define the temporal difference (TD) estimator:
\[ \hat\delta_n = R_{n+1}+\gamma \hat V_{n+1,\phi}(S_{n+1})-\hat V_{n,\phi}(S_n) \]Since \(Q_n(s,a)=E_{R_{n+1},S_{n+1}}[R_{n+1}+\gamma V_{n+1}(S_{n+1})|S_n=s,A_n=a]\), the true advantage is
\[ \mathrm{Adv}_n(s,a)=Q_n(s,a)-V_n(s)=\mathbb E\big[\underbrace{R_{n+1}+\gamma V_{n+1}(S_{n+1})-V_n(s)}_{\delta_n}\,\big|\,S_n=s,A_n=a\big] \]where \(\delta_n\) is the true temporal difference as an unbiased estimate of advantage, able to be estimated by \(\hat \delta_n\).
This is considered having smaller variance than MC estimator \(G_n-b(S_n)\) since it relies only on \(R_{n+1},S_{n+1}\) (\(S_n\) is realized).
For outcome reward schemes, \(R_{n+1\ne T}=0\) so \(\delta_{n<T-1}=\gamma V_{n+1}(S_{n+1})-V_n(S_n)\) while \(\delta_{T-1}=R_T-V_{T-1}(S_{T-1})\).
k-step TD
\[ \hat\delta^{(k)}_n=\sum_{j=0}^{k-1}\gamma^j\hat\delta_{n+j}=\sum_{j=0}^{k-1}\gamma^{j}R_{n+j+1}+\gamma^{k}\,\hat V_{n+k,\phi}(S_{n+k})-\hat V_{n,\phi}(S_n) \]If \(n+k>T\) then \(k\leftarrow T-n\); note \(\hat V_{T,\phi}(S_T)=0\). It has smaller bias (since 1-step TD relies on \(\phi\) causing bias) but larger variance.
Generalized Advantage Estimation
\[ \hat\delta^{\mathrm{GAE}(\gamma,\lambda)}_n=\sum_{j=0}^{T-n-1}(\gamma\lambda)^{j}\hat\delta_{n+j} \]when \(\lambda=0\) it’s 1-step TD \(\hat\delta_n\); when \(\lambda=1\) it’s MC estimator \(G_n\) w/o baseline.
This is what SAO applies to estimate, where in a token-wise (not our timestep-wise settings here) setting, \(l\) is the number of all actions (effective tokens? \(T\)?), \(\lambda_{\text{policy}}=1-{1\over \alpha l}\) by VAPO to make GAE estimates as above and \(\lambda_{\text{critic}}=1\) to train the critic.
By the way,