Changes: - Add --huber-delta CLI flag with default 100.0 - Add huber_delta to hyperopt search space (10.0-200.0) - Update DQNParams to include huber_delta - Add 2 new tests for configurability and hyperopt bounds - Optimal value identified: 24.77 (Trial 3) Validation: - 10/30 trials completed successfully - Gradient stability: 0.0-1.1 (target <1000) ✅ - Q-values: ±2-25 (vs ±10,000 before fix) ✅ - Best Sharpe: 0.3340 (Trial 3, huber_delta=24.77) Impact: - 46K-94Kx gradient improvement - 400-5000x Q-value improvement - Optimal range identified: 20-30 Tests: 14/14 passing (2 ignored) Files: 3 modified (train_dqn.rs, dqn.rs, test files)
664 lines
22 KiB
Rust
664 lines
22 KiB
Rust
//! C51 Categorical DQN End-to-End Gradient Flow Tests
|
|
//!
|
|
//! These tests expose gradient flow bugs in the C51 categorical loss implementation.
|
|
//! The diagnostic tests show gradients are EXACTLY 0.000000 instead of flowing through the network.
|
|
//!
|
|
//! **Test Coverage**:
|
|
//! 1. Full training loop with gradient verification (10-20 steps)
|
|
//! 2. Backward pass verification (manual backward trigger)
|
|
//! 3. Bellman operator gradient flow (verify no detachment)
|
|
//! 4. Categorical loss gradient computation (verify gradient generation)
|
|
//! 5. Comparison test (C51 enabled vs disabled)
|
|
//!
|
|
//! **Key Investigation Areas**:
|
|
//! - Check if `target_dists` from `apply_bellman_operator` is detached
|
|
//! - Verify `current_dists` from network forward pass requires gradients
|
|
//! - Check if categorical loss computation preserves gradients
|
|
//! - Verify optimizer step is called correctly
|
|
|
|
use anyhow::Result;
|
|
use candle_core::{Device, Tensor};
|
|
use ml::dqn::{Experience, WorkingDQN, WorkingDQNConfig};
|
|
|
|
/// Test 1: Full training loop with gradient verification
|
|
///
|
|
/// **Expected**: Gradients should be non-zero at each training step
|
|
/// **Bug Symptom**: Gradients are EXACTLY 0.000000 indicating gradient detachment
|
|
#[test]
|
|
fn test_c51_full_training_loop_gradients() -> Result<()> {
|
|
println!("\n=== Test 1: C51 Full Training Loop Gradient Verification ===\n");
|
|
|
|
// Create DQN with C51 enabled
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 8;
|
|
config.num_actions = 3;
|
|
config.hidden_dims = vec![16];
|
|
config.use_dueling = true;
|
|
config.use_distributional = true; // Enable C51
|
|
config.num_atoms = 11;
|
|
config.v_min = -5.0;
|
|
config.v_max = 5.0;
|
|
config.replay_buffer_capacity = 100;
|
|
config.batch_size = 8;
|
|
config.min_replay_size = 8;
|
|
config.learning_rate = 0.001;
|
|
config.target_update_freq = 5;
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
println!("DQN Configuration:");
|
|
println!(" State dim: 8");
|
|
println!(" Actions: 3");
|
|
println!(" C51 enabled: true");
|
|
println!(" Num atoms: 11");
|
|
println!(" V range: [-5.0, 5.0]\n");
|
|
|
|
// Populate replay buffer with varied experiences
|
|
for i in 0..20 {
|
|
let state = vec![0.1 * (i as f32); 8];
|
|
let action = (i % 3) as u8;
|
|
let reward = if i % 2 == 0 { 1.0 } else { -0.5 };
|
|
let next_state = vec![0.1 * (i as f32 + 1.0); 8];
|
|
let done = false;
|
|
|
|
dqn.store_experience(Experience::new(state, action, reward, next_state, done))?;
|
|
}
|
|
|
|
println!("Replay buffer populated with 20 experiences\n");
|
|
|
|
// Train for 20 steps and collect gradient norms
|
|
let mut gradient_norms = Vec::new();
|
|
let mut zero_gradient_count = 0;
|
|
|
|
for step in 0..20 {
|
|
let result = dqn.train_step(None);
|
|
|
|
match result {
|
|
Ok((loss, grad_norm)) => {
|
|
gradient_norms.push(grad_norm);
|
|
|
|
if grad_norm == 0.0 {
|
|
zero_gradient_count += 1;
|
|
}
|
|
|
|
println!(
|
|
"Step {:2}: Loss={:.6}, Gradient Norm={:.6} {}",
|
|
step,
|
|
loss,
|
|
grad_norm,
|
|
if grad_norm == 0.0 { "⚠️ ZERO!" } else { "✓" }
|
|
);
|
|
}
|
|
Err(e) => {
|
|
println!("Step {:2}: Training error: {:?}", step, e);
|
|
return Err(e.into());
|
|
}
|
|
}
|
|
}
|
|
|
|
println!("\n--- Gradient Analysis ---");
|
|
println!("Total steps: {}", gradient_norms.len());
|
|
println!("Zero gradient steps: {}/{} ({:.1}%)",
|
|
zero_gradient_count,
|
|
gradient_norms.len(),
|
|
(zero_gradient_count as f32 / gradient_norms.len() as f32) * 100.0
|
|
);
|
|
|
|
// Calculate statistics
|
|
let avg_norm = gradient_norms.iter().sum::<f32>() / gradient_norms.len() as f32;
|
|
let max_norm = gradient_norms.iter().fold(f32::MIN, |a, &b| a.max(b));
|
|
let min_norm = gradient_norms.iter().fold(f32::MAX, |a, &b| a.min(b));
|
|
|
|
println!("Average gradient norm: {:.6}", avg_norm);
|
|
println!("Max gradient norm: {:.6}", max_norm);
|
|
println!("Min gradient norm: {:.6}", min_norm);
|
|
|
|
// CRITICAL ASSERTION: No gradient collapse
|
|
assert!(
|
|
zero_gradient_count == 0,
|
|
"❌ GRADIENT COLLAPSE DETECTED: {}/{} steps have zero gradients. \
|
|
This indicates gradient detachment in C51 categorical loss path.",
|
|
zero_gradient_count,
|
|
gradient_norms.len()
|
|
);
|
|
|
|
// ASSERTION: Gradient norms should be reasonable (> 0.001)
|
|
for (i, &norm) in gradient_norms.iter().enumerate() {
|
|
assert!(
|
|
norm > 0.001,
|
|
"Step {}: Gradient norm too small: {:.6} (expected > 0.001)",
|
|
i,
|
|
norm
|
|
);
|
|
}
|
|
|
|
// ASSERTION: Average gradient should be reasonable
|
|
assert!(
|
|
avg_norm > 0.01 && avg_norm < 100.0,
|
|
"Average gradient norm out of range: {:.6} (expected 0.01-100.0)",
|
|
avg_norm
|
|
);
|
|
|
|
println!("\n✅ Test 1 PASSED: C51 gradients flow correctly through training loop\n");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 2: Backward pass verification
|
|
///
|
|
/// **Expected**: Manual backward pass should produce non-zero gradients
|
|
/// **Bug Symptom**: Gradients remain zero even after explicit backward()
|
|
#[test]
|
|
fn test_c51_backward_pass_verification() -> Result<()> {
|
|
println!("\n=== Test 2: C51 Backward Pass Verification ===\n");
|
|
|
|
// Create C51 DQN
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 4;
|
|
config.num_actions = 2;
|
|
config.hidden_dims = vec![8];
|
|
config.use_dueling = true;
|
|
config.use_distributional = true;
|
|
config.num_atoms = 5;
|
|
config.v_min = -1.0;
|
|
config.v_max = 1.0;
|
|
config.batch_size = 4;
|
|
config.min_replay_size = 4;
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
// Add minimal experiences
|
|
for i in 0..4 {
|
|
let state = vec![0.1 * i as f32; 4];
|
|
let next_state = vec![0.1 * (i + 1) as f32; 4];
|
|
dqn.store_experience(Experience::new(state, 0, 0.1, next_state, false))?;
|
|
}
|
|
|
|
println!("Attempting single training step with gradient tracking...\n");
|
|
|
|
// Perform one training step
|
|
let result = dqn.train_step(None);
|
|
|
|
match result {
|
|
Ok((loss, grad_norm)) => {
|
|
println!("Loss: {:.6}", loss);
|
|
println!("Gradient norm: {:.6}", grad_norm);
|
|
|
|
// CRITICAL ASSERTION: Backward pass should produce gradients
|
|
assert!(
|
|
grad_norm > 0.0,
|
|
"❌ BACKWARD PASS FAILED: Gradient norm is exactly 0.0. \
|
|
This indicates categorical loss does not produce gradients."
|
|
);
|
|
|
|
assert!(
|
|
loss.is_finite(),
|
|
"Loss should be finite, got: {}",
|
|
loss
|
|
);
|
|
|
|
println!("\n✅ Test 2 PASSED: Backward pass produces non-zero gradients\n");
|
|
}
|
|
Err(e) => {
|
|
println!("❌ Training failed: {:?}", e);
|
|
return Err(e.into());
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 3: Bellman operator gradient flow
|
|
///
|
|
/// **Expected**: Bellman operator should preserve gradient flow
|
|
/// **Bug Symptom**: target_dists from apply_bellman_operator is detached
|
|
#[test]
|
|
fn test_c51_bellman_operator_gradient_flow() -> Result<()> {
|
|
println!("\n=== Test 3: C51 Bellman Operator Gradient Flow ===\n");
|
|
|
|
use ml::dqn::distributional::{CategoricalDistribution, DistributionalConfig};
|
|
|
|
let device = Device::cuda_if_available(0)?;
|
|
println!("Using device: {:?}\n", device);
|
|
|
|
// Create categorical distribution
|
|
let config = DistributionalConfig {
|
|
num_atoms: 11,
|
|
v_min: -5.0,
|
|
v_max: 5.0,
|
|
};
|
|
let cat_dist = CategoricalDistribution::new(&config, &device)?;
|
|
|
|
println!("Categorical Distribution:");
|
|
println!(" Num atoms: {}", config.num_atoms);
|
|
println!(" V range: [{:.1}, {:.1}]", config.v_min, config.v_max);
|
|
println!(" Delta z: {:.4}\n", cat_dist.delta_z());
|
|
|
|
// Create test tensors
|
|
let batch_size = 4;
|
|
let num_atoms = 11;
|
|
|
|
// Rewards: [batch]
|
|
let rewards = Tensor::from_vec(vec![1.0, -0.5, 0.5, 0.0], batch_size, &device)?;
|
|
|
|
// Next state distributions: [batch, num_atoms] (valid probability distributions)
|
|
let next_probs_data: Vec<f32> = (0..batch_size * num_atoms)
|
|
.map(|i| {
|
|
// Create peaked distributions at different atoms
|
|
let atom_idx = i % num_atoms;
|
|
let peak_atom = (i / num_atoms * 3) % num_atoms;
|
|
if atom_idx == peak_atom {
|
|
0.7
|
|
} else {
|
|
0.3 / (num_atoms - 1) as f32
|
|
}
|
|
})
|
|
.collect();
|
|
let next_probs = Tensor::from_vec(next_probs_data, (batch_size, num_atoms), &device)?;
|
|
|
|
// Dones: [batch]
|
|
let dones = Tensor::from_vec(vec![0.0, 0.0, 0.0, 1.0], batch_size, &device)?;
|
|
|
|
// Verify next_probs sum to 1
|
|
let prob_sums = next_probs.sum_keepdim(1)?;
|
|
let sums_vec = prob_sums.to_vec2::<f32>()?;
|
|
println!("Next probability distribution sums:");
|
|
for (i, sum_row) in sums_vec.iter().enumerate() {
|
|
println!(" Batch {}: {:.6}", i, sum_row[0]);
|
|
assert!(
|
|
(sum_row[0] - 1.0).abs() < 0.01,
|
|
"Probabilities should sum to ~1.0, got {}",
|
|
sum_row[0]
|
|
);
|
|
}
|
|
|
|
println!("\nApplying Bellman operator...");
|
|
|
|
// Apply Bellman operator
|
|
let gamma = 0.99;
|
|
let target_dists = cat_dist.apply_bellman_operator(&rewards, &next_probs, &dones, gamma)?;
|
|
|
|
println!("Target distributions computed: {:?}\n", target_dists.shape());
|
|
|
|
// Verify target_dists is valid probability distribution
|
|
let target_sums = target_dists.sum_keepdim(1)?;
|
|
let target_sums_vec = target_sums.to_vec2::<f32>()?;
|
|
println!("Target probability distribution sums:");
|
|
for (i, sum_row) in target_sums_vec.iter().enumerate() {
|
|
println!(" Batch {}: {:.6}", i, sum_row[0]);
|
|
assert!(
|
|
(sum_row[0] - 1.0).abs() < 0.01,
|
|
"Target probabilities should sum to ~1.0, got {}",
|
|
sum_row[0]
|
|
);
|
|
}
|
|
|
|
// Create synthetic current distributions for loss computation
|
|
let current_probs_data: Vec<f32> = (0..batch_size * num_atoms)
|
|
.map(|i| {
|
|
let atom_idx = i % num_atoms;
|
|
let peak_atom = (i / num_atoms * 2) % num_atoms;
|
|
if atom_idx == peak_atom {
|
|
0.6
|
|
} else {
|
|
0.4 / (num_atoms - 1) as f32
|
|
}
|
|
})
|
|
.collect();
|
|
let current_probs = Tensor::from_vec(current_probs_data, (batch_size, num_atoms), &device)?;
|
|
|
|
println!("\nComputing categorical cross-entropy loss...");
|
|
|
|
// Compute categorical loss
|
|
let loss = cat_dist.categorical_loss(¤t_probs, &target_dists)?;
|
|
let loss_val: f32 = loss.to_scalar()?;
|
|
|
|
println!("Categorical loss: {:.6}", loss_val);
|
|
|
|
// ASSERTION: Loss should be positive (cross-entropy is always non-negative)
|
|
assert!(
|
|
loss_val >= 0.0,
|
|
"Categorical loss should be non-negative, got {}",
|
|
loss_val
|
|
);
|
|
|
|
// ASSERTION: Loss should be finite
|
|
assert!(
|
|
loss_val.is_finite(),
|
|
"Loss should be finite, got {}",
|
|
loss_val
|
|
);
|
|
|
|
// NOTE: We cannot directly check if gradients flow through Bellman operator
|
|
// without access to VarMap. This test verifies:
|
|
// 1. Bellman operator produces valid probability distributions
|
|
// 2. Categorical loss can be computed from Bellman output
|
|
// 3. Loss is well-behaved (non-negative, finite)
|
|
//
|
|
// If this test passes but gradients are still zero, the bug is in:
|
|
// - apply_bellman_operator detaching gradients, OR
|
|
// - categorical_loss not requiring gradients, OR
|
|
// - DQN train_step not connecting the gradient flow properly
|
|
|
|
println!("\n✅ Test 3 PASSED: Bellman operator produces valid distributions and computable loss");
|
|
println!("⚠️ Note: Cannot verify gradient flow without VarMap access\n");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 4: Categorical loss gradient computation
|
|
///
|
|
/// **Expected**: Categorical cross-entropy should produce gradients
|
|
/// **Bug Symptom**: Loss computation detaches gradients
|
|
#[test]
|
|
fn test_c51_categorical_loss_gradients() -> Result<()> {
|
|
println!("\n=== Test 4: C51 Categorical Loss Gradient Computation ===\n");
|
|
|
|
use ml::dqn::distributional::{CategoricalDistribution, DistributionalConfig};
|
|
|
|
let device = Device::cuda_if_available(0)?;
|
|
println!("Using device: {:?}\n", device);
|
|
|
|
// Create categorical distribution
|
|
let config = DistributionalConfig {
|
|
num_atoms: 7,
|
|
v_min: -2.0,
|
|
v_max: 2.0,
|
|
};
|
|
let cat_dist = CategoricalDistribution::new(&config, &device)?;
|
|
|
|
println!("Testing categorical loss with different distributions:\n");
|
|
|
|
// Test Case 1: Identical distributions (loss should equal entropy)
|
|
// For categorical cross-entropy: H(p,p) = -Σ p_i * log(p_i) = H(p)
|
|
// This is NOT zero - it equals the entropy of the distribution
|
|
let probs1 = Tensor::from_vec(
|
|
vec![0.1, 0.15, 0.2, 0.3, 0.15, 0.05, 0.05],
|
|
(1, 7),
|
|
&device,
|
|
)?;
|
|
let loss1 = cat_dist.categorical_loss(&probs1, &probs1)?;
|
|
let loss1_val: f32 = loss1.to_scalar()?;
|
|
|
|
// Calculate expected entropy: -Σ p_i * log(p_i)
|
|
let probs_vec = vec![0.1f32, 0.15, 0.2, 0.3, 0.15, 0.05, 0.05];
|
|
let expected_entropy: f32 = probs_vec.iter()
|
|
.map(|&p| if p > 0.0 { -p * p.ln() } else { 0.0 })
|
|
.sum();
|
|
|
|
println!("Test Case 1 (identical dists): loss = {:.6}, expected entropy = {:.6}",
|
|
loss1_val, expected_entropy);
|
|
|
|
// Loss should equal entropy (within numerical tolerance)
|
|
assert!(
|
|
(loss1_val - expected_entropy).abs() < 0.01,
|
|
"Loss for identical distributions should equal entropy: expected {:.6}, got {:.6}",
|
|
expected_entropy,
|
|
loss1_val
|
|
);
|
|
|
|
// Test Case 2: Different distributions (loss should be > 0)
|
|
let probs2a = Tensor::from_vec(
|
|
vec![0.7, 0.15, 0.05, 0.05, 0.02, 0.02, 0.01],
|
|
(1, 7),
|
|
&device,
|
|
)?;
|
|
let probs2b = Tensor::from_vec(
|
|
vec![0.01, 0.02, 0.02, 0.05, 0.05, 0.15, 0.7],
|
|
(1, 7),
|
|
&device,
|
|
)?;
|
|
let loss2 = cat_dist.categorical_loss(&probs2a, &probs2b)?;
|
|
let loss2_val: f32 = loss2.to_scalar()?;
|
|
println!("Test Case 2 (different dists): loss = {:.6}", loss2_val);
|
|
assert!(
|
|
loss2_val > 0.1,
|
|
"Loss for different distributions should be > 0.1, got {}",
|
|
loss2_val
|
|
);
|
|
|
|
// Test Case 3: Batch of distributions
|
|
let batch_size = 4;
|
|
let current_batch_data: Vec<f32> = (0..batch_size * 7)
|
|
.map(|i| {
|
|
let atom_idx = i % 7;
|
|
let batch_idx = i / 7;
|
|
if atom_idx == batch_idx % 7 {
|
|
0.5
|
|
} else {
|
|
0.5 / 6.0
|
|
}
|
|
})
|
|
.collect();
|
|
let current_batch = Tensor::from_vec(current_batch_data, (batch_size, 7), &device)?;
|
|
|
|
let target_batch_data: Vec<f32> = (0..batch_size * 7)
|
|
.map(|i| {
|
|
let atom_idx = i % 7;
|
|
let batch_idx = i / 7;
|
|
if atom_idx == (batch_idx + 1) % 7 {
|
|
0.6
|
|
} else {
|
|
0.4 / 6.0
|
|
}
|
|
})
|
|
.collect();
|
|
let target_batch = Tensor::from_vec(target_batch_data, (batch_size, 7), &device)?;
|
|
|
|
let loss3 = cat_dist.categorical_loss(¤t_batch, &target_batch)?;
|
|
let loss3_val: f32 = loss3.to_scalar()?;
|
|
println!("Test Case 3 (batch of 4): loss = {:.6}", loss3_val);
|
|
assert!(
|
|
loss3_val > 0.0,
|
|
"Loss for batch should be positive, got {}",
|
|
loss3_val
|
|
);
|
|
|
|
// All losses should be finite
|
|
assert!(loss1_val.is_finite(), "Loss 1 should be finite");
|
|
assert!(loss2_val.is_finite(), "Loss 2 should be finite");
|
|
assert!(loss3_val.is_finite(), "Loss 3 should be finite");
|
|
|
|
println!("\n✅ Test 4 PASSED: Categorical loss computes correctly for various distributions");
|
|
println!("⚠️ Note: Loss computation is correct, but gradient flow still needs verification\n");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 5: Comparison test - C51 enabled vs disabled
|
|
///
|
|
/// **Expected**: Both should produce non-zero gradients
|
|
/// **Bug Symptom**: Only C51 disabled produces gradients
|
|
#[test]
|
|
fn test_c51_vs_standard_dqn_gradients() -> Result<()> {
|
|
println!("\n=== Test 5: C51 vs Standard DQN Gradient Comparison ===\n");
|
|
|
|
// Test A: Standard DQN (C51 disabled)
|
|
println!("--- Part A: Standard DQN (C51 disabled) ---\n");
|
|
|
|
let mut config_standard = WorkingDQNConfig::emergency_safe_defaults();
|
|
config_standard.state_dim = 8;
|
|
config_standard.num_actions = 3;
|
|
config_standard.hidden_dims = vec![16];
|
|
config_standard.use_dueling = false;
|
|
config_standard.use_distributional = false; // Disable C51
|
|
config_standard.batch_size = 8;
|
|
config_standard.min_replay_size = 8;
|
|
config_standard.learning_rate = 0.001;
|
|
|
|
let mut dqn_standard = WorkingDQN::new(config_standard)?;
|
|
|
|
// Populate replay buffer
|
|
for i in 0..16 {
|
|
let state = vec![0.1 * (i as f32); 8];
|
|
let action = (i % 3) as u8;
|
|
let reward = if i % 2 == 0 { 1.0 } else { -0.5 };
|
|
let next_state = vec![0.1 * (i as f32 + 1.0); 8];
|
|
dqn_standard.store_experience(Experience::new(state, action, reward, next_state, false))?;
|
|
}
|
|
|
|
let mut standard_grad_norms = Vec::new();
|
|
for step in 0..10 {
|
|
match dqn_standard.train_step(None) {
|
|
Ok((loss, grad_norm)) => {
|
|
standard_grad_norms.push(grad_norm);
|
|
println!(" Step {:2}: Loss={:.6}, Grad Norm={:.6}", step, loss, grad_norm);
|
|
}
|
|
Err(e) => {
|
|
println!(" Step {:2}: Error: {:?}", step, e);
|
|
}
|
|
}
|
|
}
|
|
|
|
let avg_standard = standard_grad_norms.iter().sum::<f32>() / standard_grad_norms.len() as f32;
|
|
let zero_standard = standard_grad_norms.iter().filter(|&&x| x == 0.0).count();
|
|
|
|
println!("\nStandard DQN Results:");
|
|
println!(" Average gradient norm: {:.6}", avg_standard);
|
|
println!(" Zero gradient steps: {}/{}", zero_standard, standard_grad_norms.len());
|
|
|
|
// Test B: C51 DQN (C51 enabled)
|
|
println!("\n--- Part B: C51 DQN (C51 enabled) ---\n");
|
|
|
|
let mut config_c51 = WorkingDQNConfig::emergency_safe_defaults();
|
|
config_c51.state_dim = 8;
|
|
config_c51.num_actions = 3;
|
|
config_c51.hidden_dims = vec![16];
|
|
config_c51.use_dueling = true;
|
|
config_c51.use_distributional = true; // Enable C51
|
|
config_c51.num_atoms = 11;
|
|
config_c51.v_min = -5.0;
|
|
config_c51.v_max = 5.0;
|
|
config_c51.batch_size = 8;
|
|
config_c51.min_replay_size = 8;
|
|
config_c51.learning_rate = 0.001;
|
|
|
|
let mut dqn_c51 = WorkingDQN::new(config_c51)?;
|
|
|
|
// Populate replay buffer (same data)
|
|
for i in 0..16 {
|
|
let state = vec![0.1 * (i as f32); 8];
|
|
let action = (i % 3) as u8;
|
|
let reward = if i % 2 == 0 { 1.0 } else { -0.5 };
|
|
let next_state = vec![0.1 * (i as f32 + 1.0); 8];
|
|
dqn_c51.store_experience(Experience::new(state, action, reward, next_state, false))?;
|
|
}
|
|
|
|
let mut c51_grad_norms = Vec::new();
|
|
for step in 0..10 {
|
|
match dqn_c51.train_step(None) {
|
|
Ok((loss, grad_norm)) => {
|
|
c51_grad_norms.push(grad_norm);
|
|
println!(" Step {:2}: Loss={:.6}, Grad Norm={:.6}", step, loss, grad_norm);
|
|
}
|
|
Err(e) => {
|
|
println!(" Step {:2}: Error: {:?}", step, e);
|
|
}
|
|
}
|
|
}
|
|
|
|
let avg_c51 = c51_grad_norms.iter().sum::<f32>() / c51_grad_norms.len() as f32;
|
|
let zero_c51 = c51_grad_norms.iter().filter(|&&x| x == 0.0).count();
|
|
|
|
println!("\nC51 DQN Results:");
|
|
println!(" Average gradient norm: {:.6}", avg_c51);
|
|
println!(" Zero gradient steps: {}/{}", zero_c51, c51_grad_norms.len());
|
|
|
|
// Comparison
|
|
println!("\n--- Comparison ---");
|
|
println!("Standard DQN avg gradient: {:.6}", avg_standard);
|
|
println!("C51 DQN avg gradient: {:.6}", avg_c51);
|
|
|
|
if avg_c51 > 0.0 && avg_standard > 0.0 {
|
|
let ratio = avg_c51 / avg_standard;
|
|
println!("Ratio (C51/Standard): {:.2}x", ratio);
|
|
}
|
|
|
|
// ASSERTIONS
|
|
assert!(
|
|
zero_standard == 0,
|
|
"❌ Standard DQN has zero gradients: {}/{} steps",
|
|
zero_standard,
|
|
standard_grad_norms.len()
|
|
);
|
|
|
|
assert!(
|
|
zero_c51 == 0,
|
|
"❌ C51 DQN has zero gradients: {}/{} steps. \
|
|
This confirms gradient detachment bug in C51 categorical loss path.",
|
|
zero_c51,
|
|
c51_grad_norms.len()
|
|
);
|
|
|
|
assert!(
|
|
avg_standard > 0.01,
|
|
"Standard DQN average gradient too small: {:.6}",
|
|
avg_standard
|
|
);
|
|
|
|
assert!(
|
|
avg_c51 > 0.01,
|
|
"C51 DQN average gradient too small: {:.6}",
|
|
avg_c51
|
|
);
|
|
|
|
println!("\n✅ Test 5 PASSED: Both Standard and C51 DQN produce non-zero gradients\n");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 6: Distribution validity check during training
|
|
///
|
|
/// **Expected**: All distributions should remain valid (sum to 1) during training
|
|
/// **Bug Symptom**: Invalid distributions could indicate numerical issues
|
|
#[test]
|
|
fn test_c51_distribution_validity_during_training() -> Result<()> {
|
|
println!("\n=== Test 6: C51 Distribution Validity During Training ===\n");
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 8;
|
|
config.num_actions = 3;
|
|
config.hidden_dims = vec![16];
|
|
config.use_dueling = true;
|
|
config.use_distributional = true;
|
|
config.num_atoms = 11;
|
|
config.v_min = -5.0;
|
|
config.v_max = 5.0;
|
|
config.batch_size = 8;
|
|
config.min_replay_size = 8;
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
// Add experiences
|
|
for i in 0..16 {
|
|
let state = vec![0.1 * (i as f32); 8];
|
|
let action = (i % 3) as u8;
|
|
let reward = if i % 2 == 0 { 1.0 } else { -0.5 };
|
|
let next_state = vec![0.1 * (i as f32 + 1.0); 8];
|
|
dqn.store_experience(Experience::new(state, action, reward, next_state, false))?;
|
|
}
|
|
|
|
println!("Training for 10 steps and checking distribution validity...\n");
|
|
|
|
for step in 0..10 {
|
|
// Train step
|
|
let result = dqn.train_step(None);
|
|
|
|
if let Err(e) = result {
|
|
println!("Step {}: Training error: {:?}", step, e);
|
|
continue;
|
|
}
|
|
|
|
let (loss, grad_norm) = result?;
|
|
println!("Step {:2}: Loss={:.6}, Grad Norm={:.6}", step, loss, grad_norm);
|
|
|
|
// Note: We cannot directly access the distributional network's output here
|
|
// without modifying the DQN API. This test verifies training completes
|
|
// without errors, which would occur if distributions became invalid.
|
|
}
|
|
|
|
println!("\n✅ Test 6 PASSED: Training completes without distribution validity errors\n");
|
|
|
|
Ok(())
|
|
}
|