BREAKING CHANGES: - Removed orphaned dqn.rs monolithic trainer (4,975 lines) - Removed orphaned dqn_ensemble.rs module (816 lines) - Removed orphaned tft.rs and tft_complete_int8_integration_test.rs - TFT trainer split into modular directory structure DQN Module Refactoring: - Split trainers/dqn.rs into modular structure (config.rs, statistics.rs, trainer.rs) - Fixed hyperopt 39D search space (continuous params only) - Boolean flags (use_dueling, use_double_dqn, use_per, use_noisy_nets) are now FIXED architectural decisions - use_distributional defaults to false (Candle BUG #36 - scatter_add gradient issues) Clean Module Structure: - ml/src/trainers/dqn/ directory with proper mod.rs exports - ml/src/trainers/tft/ directory with config.rs, types.rs, model.rs, trainer.rs, tests.rs - All P0 features validated: TD-error clamping, batch diversity, LR scheduler, priority staleness Documentation: - Added comprehensive docs in docs/codebase-cleanup/ - ADR-001 for DQN refactoring decisions - Rainbow DQN component matrix and quick reference guides Build Status: Compiles with zero errors 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude <noreply@anthropic.com>
490 lines
16 KiB
Rust
490 lines
16 KiB
Rust
//! Generalized Advantage Estimation (GAE) for lower variance returns
|
||
//!
|
||
//! GAE combines TD(λ) and Monte Carlo returns to provide a bias-variance tradeoff
|
||
//! in advantage estimation. This is particularly useful for policy gradient methods
|
||
//! and can improve DQN training stability.
|
||
//!
|
||
//! Reference: "High-Dimensional Continuous Control Using Generalized Advantage Estimation"
|
||
//! Schulman et al., 2016 (https://arxiv.org/abs/1506.02438)
|
||
|
||
use serde::{Deserialize, Serialize};
|
||
|
||
/// Configuration for Generalized Advantage Estimation
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub struct GAEConfig {
|
||
/// Discount factor (gamma), typically 0.99
|
||
pub gamma: f64,
|
||
/// GAE lambda parameter for bias-variance tradeoff, typically 0.95
|
||
/// lambda = 0: high bias, low variance (TD(0))
|
||
/// lambda = 1: low bias, high variance (Monte Carlo)
|
||
pub lambda: f64,
|
||
}
|
||
|
||
impl Default for GAEConfig {
|
||
fn default() -> Self {
|
||
Self {
|
||
gamma: 0.99,
|
||
lambda: 0.95,
|
||
}
|
||
}
|
||
}
|
||
|
||
/// Generalized Advantage Estimation calculator
|
||
///
|
||
/// Computes advantage estimates and returns using GAE(λ) algorithm.
|
||
/// This provides a smooth interpolation between TD and Monte Carlo methods.
|
||
#[derive(Debug)]
|
||
pub struct GAECalculator {
|
||
/// Discount factor (gamma)
|
||
gamma: f64,
|
||
/// GAE lambda parameter
|
||
lambda: f64,
|
||
}
|
||
|
||
impl GAECalculator {
|
||
/// Create a new GAE calculator with specified parameters
|
||
///
|
||
/// # Arguments
|
||
/// * `gamma` - Discount factor (0.0 to 1.0), typically 0.99
|
||
/// * `lambda` - GAE lambda parameter (0.0 to 1.0), typically 0.95
|
||
///
|
||
/// # Panics
|
||
/// Panics if gamma or lambda are outside [0, 1] range
|
||
pub fn new(gamma: f64, lambda: f64) -> Self {
|
||
assert!(
|
||
(0.0..=1.0).contains(&gamma),
|
||
"Gamma must be in [0, 1], got {}",
|
||
gamma
|
||
);
|
||
assert!(
|
||
(0.0..=1.0).contains(&lambda),
|
||
"Lambda must be in [0, 1], got {}",
|
||
lambda
|
||
);
|
||
Self { gamma, lambda }
|
||
}
|
||
|
||
/// Create a new GAE calculator from configuration
|
||
pub fn from_config(config: &GAEConfig) -> Self {
|
||
Self::new(config.gamma, config.lambda)
|
||
}
|
||
|
||
/// Compute GAE returns from rewards, value estimates, and done flags
|
||
///
|
||
/// # Arguments
|
||
/// * `rewards` - Rewards for each timestep
|
||
/// * `values` - Value function estimates V(s_t) for each state
|
||
/// * `dones` - Episode termination flags (true if episode ended)
|
||
///
|
||
/// # Returns
|
||
/// Vector of GAE-based returns (advantages + values)
|
||
///
|
||
/// # Algorithm
|
||
/// 1. Compute TD errors: δ_t = r_t + γ V(s_{t+1}) - V(s_t)
|
||
/// 2. Compute GAE advantages: A^GAE_t = Σ_{l=0}^∞ (γλ)^l δ_{t+l}
|
||
/// 3. Return: R_t = A^GAE_t + V(s_t)
|
||
///
|
||
/// # Panics
|
||
/// Panics if rewards, values, and dones have different lengths
|
||
pub fn compute_returns(
|
||
&self,
|
||
rewards: &[f64],
|
||
values: &[f64],
|
||
dones: &[bool],
|
||
) -> Vec<f64> {
|
||
let n = rewards.len();
|
||
assert_eq!(
|
||
values.len(),
|
||
n,
|
||
"Values length {} must match rewards length {}",
|
||
values.len(),
|
||
n
|
||
);
|
||
assert_eq!(
|
||
dones.len(),
|
||
n,
|
||
"Dones length {} must match rewards length {}",
|
||
dones.len(),
|
||
n
|
||
);
|
||
|
||
if n == 0 {
|
||
return Vec::new();
|
||
}
|
||
|
||
let mut advantages = vec![0.0; n];
|
||
let mut gae = 0.0;
|
||
|
||
// Backward pass: compute GAE advantages
|
||
for t in (0..n).rev() {
|
||
// Next value is 0 if episode ended or at trajectory end
|
||
let next_value = if t == n - 1 || dones[t] {
|
||
0.0
|
||
} else {
|
||
values[t + 1]
|
||
};
|
||
|
||
// TD error: δ_t = r_t + γ V(s_{t+1}) - V(s_t)
|
||
let delta = rewards[t] + self.gamma * next_value - values[t];
|
||
|
||
// GAE: A^GAE_t = δ_t + γλ A^GAE_{t+1}
|
||
// Reset GAE if episode ended
|
||
gae = if dones[t] {
|
||
delta // Don't accumulate past episode boundary
|
||
} else {
|
||
delta + self.gamma * self.lambda * gae
|
||
};
|
||
|
||
advantages[t] = gae;
|
||
}
|
||
|
||
// Returns = advantages + values
|
||
advantages
|
||
.iter()
|
||
.zip(values)
|
||
.map(|(a, v)| a + v)
|
||
.collect()
|
||
}
|
||
|
||
/// Compute only advantages (without adding back values)
|
||
///
|
||
/// Useful when you want advantages separately from returns.
|
||
pub fn compute_advantages(
|
||
&self,
|
||
rewards: &[f64],
|
||
values: &[f64],
|
||
dones: &[bool],
|
||
) -> Vec<f64> {
|
||
let n = rewards.len();
|
||
assert_eq!(values.len(), n);
|
||
assert_eq!(dones.len(), n);
|
||
|
||
if n == 0 {
|
||
return Vec::new();
|
||
}
|
||
|
||
let mut advantages = vec![0.0; n];
|
||
let mut gae = 0.0;
|
||
|
||
for t in (0..n).rev() {
|
||
let next_value = if t == n - 1 || dones[t] {
|
||
0.0
|
||
} else {
|
||
values[t + 1]
|
||
};
|
||
|
||
let delta = rewards[t] + self.gamma * next_value - values[t];
|
||
gae = if dones[t] {
|
||
delta
|
||
} else {
|
||
delta + self.gamma * self.lambda * gae
|
||
};
|
||
|
||
advantages[t] = gae;
|
||
}
|
||
|
||
advantages
|
||
}
|
||
|
||
/// Get gamma parameter
|
||
pub fn gamma(&self) -> f64 {
|
||
self.gamma
|
||
}
|
||
|
||
/// Get lambda parameter
|
||
pub fn lambda(&self) -> f64 {
|
||
self.lambda
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
#[test]
|
||
fn test_gae_config_default() {
|
||
let config = GAEConfig::default();
|
||
assert_eq!(config.gamma, 0.99);
|
||
assert_eq!(config.lambda, 0.95);
|
||
}
|
||
|
||
#[test]
|
||
fn test_gae_calculator_creation() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
assert_eq!(gae.gamma(), 0.99);
|
||
assert_eq!(gae.lambda(), 0.95);
|
||
}
|
||
|
||
#[test]
|
||
fn test_gae_from_config() {
|
||
let config = GAEConfig {
|
||
gamma: 0.98,
|
||
lambda: 0.9,
|
||
};
|
||
let gae = GAECalculator::from_config(&config);
|
||
assert_eq!(gae.gamma(), 0.98);
|
||
assert_eq!(gae.lambda(), 0.9);
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Gamma must be in [0, 1]")]
|
||
fn test_invalid_gamma_high() {
|
||
GAECalculator::new(1.5, 0.95);
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Gamma must be in [0, 1]")]
|
||
fn test_invalid_gamma_low() {
|
||
GAECalculator::new(-0.1, 0.95);
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Lambda must be in [0, 1]")]
|
||
fn test_invalid_lambda_high() {
|
||
GAECalculator::new(0.99, 1.5);
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Lambda must be in [0, 1]")]
|
||
fn test_invalid_lambda_low() {
|
||
GAECalculator::new(0.99, -0.1);
|
||
}
|
||
|
||
#[test]
|
||
fn test_empty_trajectory() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let returns = gae.compute_returns(&[], &[], &[]);
|
||
assert!(returns.is_empty());
|
||
|
||
let advantages = gae.compute_advantages(&[], &[], &[]);
|
||
assert!(advantages.is_empty());
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Values length")]
|
||
fn test_mismatched_values_length() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6]; // Wrong length
|
||
let dones = vec![false, false, false];
|
||
gae.compute_returns(&rewards, &values, &dones);
|
||
}
|
||
|
||
#[test]
|
||
#[should_panic(expected = "Dones length")]
|
||
fn test_mismatched_dones_length() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, false]; // Wrong length
|
||
gae.compute_returns(&rewards, &values, &dones);
|
||
}
|
||
|
||
#[test]
|
||
fn test_single_step_trajectory() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0];
|
||
let values = vec![0.5];
|
||
let dones = vec![true];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
// δ = r + γ*0 - V = 1.0 + 0 - 0.5 = 0.5
|
||
// A = δ = 0.5 (episode ended)
|
||
// Return = A + V = 0.5 + 0.5 = 1.0
|
||
assert_eq!(returns.len(), 1);
|
||
assert!((returns[0] - 1.0).abs() < 1e-6);
|
||
}
|
||
|
||
#[test]
|
||
fn test_two_step_trajectory_no_termination() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0];
|
||
let values = vec![0.5, 0.6];
|
||
let dones = vec![false, false];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
|
||
// Step 1 (t=1): δ₁ = 2.0 + 0.99*0 - 0.6 = 1.4, A₁ = 1.4
|
||
// Step 0 (t=0): δ₀ = 1.0 + 0.99*0.6 - 0.5 = 1.094, A₀ = 1.094 + 0.99*0.95*1.4 ≈ 2.41
|
||
// Returns: [A₀+V₀, A₁+V₁]
|
||
assert_eq!(returns.len(), 2);
|
||
|
||
let expected_delta_1 = 2.0 + 0.99 * 0.0 - 0.6; // 1.4
|
||
let expected_a1 = expected_delta_1; // 1.4
|
||
let expected_return_1 = expected_a1 + 0.6; // 2.0
|
||
|
||
let expected_delta_0 = 1.0 + 0.99 * 0.6 - 0.5; // 1.094
|
||
let expected_a0 = expected_delta_0 + 0.99 * 0.95 * expected_a1; // 1.094 + 1.31814 = 2.41214
|
||
let expected_return_0 = expected_a0 + 0.5; // 2.91214
|
||
|
||
assert!((returns[1] - expected_return_1).abs() < 1e-6);
|
||
assert!((returns[0] - expected_return_0).abs() < 1e-6);
|
||
}
|
||
|
||
#[test]
|
||
fn test_trajectory_with_episode_boundary() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, true, false]; // Episode ends at step 1
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
assert_eq!(returns.len(), 3);
|
||
|
||
// Step 2 (t=2): δ₂ = 3.0 + 0 - 0.7 = 2.3, A₂ = 2.3
|
||
// Step 1 (t=1): δ₁ = 2.0 + 0 - 0.6 = 1.4, A₁ = 1.4 (done, don't accumulate)
|
||
// Step 0 (t=0): δ₀ = 1.0 + 0.99*0.6 - 0.5 = 1.094, A₀ = 1.094 + 0.99*0.95*1.4
|
||
let expected_delta_2 = 3.0 - 0.7; // 2.3
|
||
let expected_a2 = expected_delta_2; // 2.3
|
||
let expected_return_2 = expected_a2 + 0.7; // 3.0
|
||
|
||
let expected_delta_1 = 2.0 - 0.6; // 1.4 (done, next_value=0)
|
||
let expected_a1 = expected_delta_1; // 1.4 (done, no accumulation)
|
||
let expected_return_1 = expected_a1 + 0.6; // 2.0
|
||
|
||
let expected_delta_0 = 1.0 + 0.99 * 0.6 - 0.5; // 1.094
|
||
let expected_a0 = expected_delta_0 + 0.99 * 0.95 * expected_a1; // 1.094 + 1.31814
|
||
let expected_return_0 = expected_a0 + 0.5;
|
||
|
||
assert!((returns[2] - expected_return_2).abs() < 1e-6);
|
||
assert!((returns[1] - expected_return_1).abs() < 1e-6);
|
||
assert!((returns[0] - expected_return_0).abs() < 1e-6);
|
||
}
|
||
|
||
#[test]
|
||
fn test_lambda_zero_equals_td() {
|
||
// Lambda = 0 should give TD(0) returns (no bootstrapping beyond one step)
|
||
let gae = GAECalculator::new(0.99, 0.0);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, false, false];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
|
||
// With λ=0, GAE reduces to TD(0): A_t = δ_t
|
||
// Step 2: δ₂ = 3.0 - 0.7 = 2.3, Return = 2.3 + 0.7 = 3.0
|
||
// Step 1: δ₁ = 2.0 + 0.99*0.7 - 0.6 = 2.093, Return = 2.093 + 0.6 = 2.693
|
||
// Step 0: δ₀ = 1.0 + 0.99*0.6 - 0.5 = 1.094, Return = 1.094 + 0.5 = 1.594
|
||
assert!((returns[2] - 3.0).abs() < 1e-6);
|
||
assert!((returns[1] - 2.693).abs() < 1e-6);
|
||
assert!((returns[0] - 1.594).abs() < 1e-6);
|
||
}
|
||
|
||
#[test]
|
||
fn test_lambda_one_accumulates_fully() {
|
||
// Lambda = 1 should give full Monte Carlo-like accumulation
|
||
let gae = GAECalculator::new(0.99, 1.0);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, false, true];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
|
||
// With λ=1, GAE accumulates all future TD errors
|
||
// This should give higher returns due to full bootstrapping
|
||
assert_eq!(returns.len(), 3);
|
||
// Returns should be increasing towards the end (accumulating rewards)
|
||
assert!(returns[0] > 1.0); // Should be > reward[0] + value[0]
|
||
}
|
||
|
||
#[test]
|
||
fn test_compute_advantages_separate() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0];
|
||
let values = vec![0.5, 0.6];
|
||
let dones = vec![false, false];
|
||
|
||
let advantages = gae.compute_advantages(&rewards, &values, &dones);
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
|
||
assert_eq!(advantages.len(), returns.len());
|
||
|
||
// Verify: returns = advantages + values
|
||
for i in 0..advantages.len() {
|
||
assert!((returns[i] - (advantages[i] + values[i])).abs() < 1e-6);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn test_all_zero_trajectory() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![0.0, 0.0, 0.0];
|
||
let values = vec![0.0, 0.0, 0.0];
|
||
let dones = vec![false, false, false];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
assert_eq!(returns.len(), 3);
|
||
for r in returns {
|
||
assert!((r - 0.0).abs() < 1e-10);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn test_negative_rewards() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![-1.0, -2.0, -3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, false, true];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
assert_eq!(returns.len(), 3);
|
||
|
||
// Negative rewards should produce negative advantages
|
||
// Returns can still be computed correctly
|
||
for r in returns {
|
||
assert!(r.is_finite());
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn test_constant_value_estimates() {
|
||
// If value estimates are constant, TD errors depend only on rewards
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![1.0, 1.0, 1.0]; // Constant values
|
||
let dones = vec![false, false, true];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
assert_eq!(returns.len(), 3);
|
||
|
||
// With constant values, advantages primarily track reward differences
|
||
for r in returns {
|
||
assert!(r.is_finite());
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn test_gamma_zero_no_bootstrapping() {
|
||
// Gamma = 0 means no bootstrapping from future states
|
||
let gae = GAECalculator::new(0.0, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0];
|
||
let values = vec![0.5, 0.6, 0.7];
|
||
let dones = vec![false, false, false];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
|
||
// With γ=0: δ_t = r_t - V(s_t), no future value
|
||
// Returns should be close to rewards (since advantages ≈ rewards - values)
|
||
assert_eq!(returns.len(), 3);
|
||
for i in 0..3 {
|
||
// Return should be close to reward (as values are subtracted then added back)
|
||
assert!((returns[i] - rewards[i]).abs() < 0.1);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn test_increasing_rewards_trajectory() {
|
||
let gae = GAECalculator::new(0.99, 0.95);
|
||
let rewards = vec![1.0, 2.0, 3.0, 4.0, 5.0];
|
||
let values = vec![0.5, 0.6, 0.7, 0.8, 0.9];
|
||
let dones = vec![false, false, false, false, true];
|
||
|
||
let returns = gae.compute_returns(&rewards, &values, &dones);
|
||
assert_eq!(returns.len(), 5);
|
||
|
||
// With increasing rewards, earlier timesteps should have higher returns
|
||
// due to GAE accumulation
|
||
assert!(returns[0] > returns[4]); // Earlier steps accumulate more
|
||
}
|
||
}
|