Files
foxhunt/ml/src/dqn/gae.rs
jgrusewski 2df1ea92e1 feat(ml): WAVE 29 DQN Codebase Cleanup & Refactoring Campaign
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>
2025-11-27 23:46:13 +01:00

490 lines
16 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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
}
}