| name | manifold-aware-rl-video |
| title | SAGE-GRPO: Manifold-Aware Exploration for Video Generation Reinforcement Learning |
| version | 0.0.3 |
| engine | skillxiv-v0.0.3-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2603.21872 |
| keywords | ["Reinforcement Learning","Video Generation","Diffusion","Manifold Learning"] |
| description | Constrain video GRPO policy updates to stay within pre-trained model's data manifold using dual-control exploration. Implement precise manifold-aware SDE with logarithmic noise variance correction (captures geometric signal decay standard methods miss). Apply gradient norm equalizer to balance learning across diffusion timesteps (mitigate vanishing/exploding gradients). Use dual trust region combining position control (anchored exploration) and velocity control (KL constraints) for stability-plasticity balance. |
Core Mechanism
SAGE-GRPO treats video generation as constrained exploration within the manifold defined by a pre-trained diffusion model.
Manifold Definition:
The pre-trained video model defines a valid video data manifold—the set of videos the model considers plausible. Policy optimization must keep updates within this manifold's vicinity, preventing drift to out-of-distribution proposals.
Exploration Constraints:
Rather than unconstrained policy gradient descent, apply multi-scale constraints:
- Micro-level: Per-timestep noise variance control
- Macro-level: Multi-step trust region constraints
This balances exploration (improve reward) with stability (stay on manifold).
Key Components
Micro-Level: Precise Manifold-Aware SDE
Standard approaches use first-order approximations for noise variance, missing geometric signal decay in diffusion.
Standard (Linear) Approximation:
sigma_t = sqrt(1 - alpha_cumprod_t)
Σ_t = eta^2 * sigma_t^2
Precise (Logarithmic) Correction:
Σ_t = η²[-(σ_t - σ_{t+1}) + log((1 - σ_{t+1})/(1 - σ_t))]
Geometric Intuition:
In diffusion, noise doesn't degrade linearly—signal decays exponentially. The log term captures this exponential decay that linear methods miss. At high noise (t→1), signal is nearly zero; at low noise (t→0), signal is dense. The logarithmic correction accounts for this non-linear relationship.
Implementation:
def ():
linear_term = sigma_t - sigma_t_next
signal_ratio = ( - sigma_t_next) / ( - sigma_t)
log_term = np.log(signal_ratio)
variance = eta** * (linear_term + log_term)
variance