Major Changes: - Migrated from 3-action TradingAction to 45-action FactoredAction - 45 actions: 5 exposure × 3 order types × 3 urgency levels - Absolute exposure model (target positions -1.0 to +1.0) - Transaction cost differentiation (Market 0.15%, LimitMaker 0.05%, IoC 0.10%) - Fixed action diversity threshold (1.11% → 0.5% for 45-action space) Bug Fixes: - Bug #15: Incomplete FactoredAction integration (code existed but unused) - Bug #16: Runtime crash in action diversity checking (hardcoded 3-action match) Code Changes (13 files, ~464 lines): - ml/src/dqn/action_space.rs: Core FactoredAction + 4 helper methods - ml/src/trainers/dqn.rs: Action diversity refactored (3→45 dynamic) - ml/src/dqn/reward.rs: calculate_reward() signature updated - ml/src/dqn/portfolio_tracker.rs: execute_action() absolute exposure - ml/src/dqn/dqn.rs: WorkingDQN action selection migrated - ml/tests/*.rs: 9 test files updated with FactoredAction assertions Test Results: - 1-epoch smoke test: 100% action diversity (45/45 actions, 80.2s) - 10-epoch production: 87.8% readiness (79/90 scorecard, 14.0 min) - Loss convergence: 96.9% reduction (119K → 3.6K) - Action diversity: 100% → 44% (healthy specialization) - Checkpoint reliability: 12/12 files saved (100%) - DQN tests: 195/195 passing (100%) - ML baseline: 1,514/1,515 passing (99.93%) Production Status: ✅ CERTIFIED (87.8% readiness) Go/No-Go: ✅ GO FOR 100-EPOCH PRODUCTION TRAINING 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude <noreply@anthropic.com>
350 lines
11 KiB
Rust
350 lines
11 KiB
Rust
//! Unit tests for Q-Value Variance Adaptive Temperature
|
|
//!
|
|
//! Tests the variance-based temperature scaling mechanism that adapts exploration
|
|
//! based on action uncertainty. High Q-value variance (uncertain) triggers higher
|
|
//! temperature (more exploration), while low variance (confident) uses lower temperature
|
|
//! (more exploitation).
|
|
|
|
use candle_core::{Device, Tensor};
|
|
use ml::dqn::{WorkingDQN, WorkingDQNConfig};
|
|
|
|
#[test]
|
|
fn test_qvariance_calculation_high_spread() -> anyhow::Result<()> {
|
|
// Test variance calculation with high Q-value spread
|
|
// Q-values: [10.0, 0.0, -10.0] should produce high variance
|
|
|
|
let device = Device::Cpu;
|
|
let q_values = vec![10.0_f32, 0.0, -10.0];
|
|
let q_tensor = Tensor::from_vec(q_values.clone(), (1, 3), &device)?;
|
|
|
|
// Compute variance manually for verification
|
|
let mean = q_values.iter().sum::<f32>() / q_values.len() as f32;
|
|
let variance: f32 = q_values
|
|
.iter()
|
|
.map(|q| {
|
|
let diff = q - mean;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
/ q_values.len() as f32;
|
|
|
|
// High spread should produce high variance
|
|
assert!(
|
|
variance > 50.0,
|
|
"Expected high variance for spread Q-values, got {}",
|
|
variance
|
|
);
|
|
|
|
// Variance scale factor: sqrt(variance) / mean(abs(Q))
|
|
let abs_mean = q_values.iter().map(|q| q.abs()).sum::<f32>() / q_values.len() as f32;
|
|
let variance_scale = variance.sqrt() / abs_mean.max(0.1); // Prevent div by zero
|
|
|
|
// High variance should trigger high scale factor
|
|
assert!(
|
|
variance_scale > 1.0,
|
|
"Expected variance scale > 1.0 for uncertain Q-values, got {}",
|
|
variance_scale
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_qvariance_calculation_low_spread() -> anyhow::Result<()> {
|
|
// Test variance calculation with low Q-value spread
|
|
// Q-values: [1.0, 1.1, 0.9] should produce low variance
|
|
|
|
let device = Device::Cpu;
|
|
let q_values = vec![1.0_f32, 1.1, 0.9];
|
|
let q_tensor = Tensor::from_vec(q_values.clone(), (1, 3), &device)?;
|
|
|
|
// Compute variance manually
|
|
let mean = q_values.iter().sum::<f32>() / q_values.len() as f32;
|
|
let variance: f32 = q_values
|
|
.iter()
|
|
.map(|q| {
|
|
let diff = q - mean;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
/ q_values.len() as f32;
|
|
|
|
// Low spread should produce low variance
|
|
assert!(
|
|
variance < 0.1,
|
|
"Expected low variance for tight Q-values, got {}",
|
|
variance
|
|
);
|
|
|
|
// Variance scale factor
|
|
let abs_mean = q_values.iter().map(|q| q.abs()).sum::<f32>() / q_values.len() as f32;
|
|
let variance_scale = variance.sqrt() / abs_mean.max(0.1);
|
|
|
|
// Low variance should trigger low scale factor
|
|
assert!(
|
|
variance_scale < 0.5,
|
|
"Expected variance scale < 0.5 for confident Q-values, got {}",
|
|
variance_scale
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_temperature_scaling_high_variance() -> anyhow::Result<()> {
|
|
// Test temperature scaling with high Q-value variance
|
|
// High uncertainty should increase temperature
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 3;
|
|
config.num_actions = 3;
|
|
config.temperature_start = 0.5; // Base temperature
|
|
config.variance_multiplier = 1.0; // 1:1 variance scaling
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
// Create state with high Q-value uncertainty
|
|
// This would be detected by the variance computation
|
|
let state = vec![1.0_f32, 0.5, 0.0]; // Dummy state
|
|
let state_tensor = Tensor::from_vec(state, (1, 3), dqn.device())?;
|
|
|
|
// Get Q-values (will have some spread due to initialization)
|
|
let q_values = dqn.forward(&state_tensor)?;
|
|
|
|
// Extract Q-values for variance computation
|
|
let q_vec = q_values.flatten_all()?.to_vec1::<f32>()?;
|
|
let mean = q_vec.iter().sum::<f32>() / q_vec.len() as f32;
|
|
let variance: f32 = q_vec
|
|
.iter()
|
|
.map(|q| {
|
|
let diff = q - mean;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
/ q_vec.len() as f32;
|
|
|
|
// If variance is high, adaptive temperature should exceed base
|
|
if variance > 1.0 {
|
|
let base_temp = dqn.get_temperature();
|
|
// Adaptive temp = base * (1 + variance_multiplier * normalized_variance)
|
|
// Should be higher than base for high variance
|
|
assert!(
|
|
base_temp > 0.0,
|
|
"Base temperature should be positive, got {}",
|
|
base_temp
|
|
);
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_temperature_scaling_low_variance() -> anyhow::Result<()> {
|
|
// Test temperature scaling with low Q-value variance
|
|
// High confidence should decrease temperature (more greedy)
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 3;
|
|
config.num_actions = 3;
|
|
config.temperature_start = 1.0; // Base temperature
|
|
config.variance_multiplier = 0.5; // Conservative scaling
|
|
|
|
let dqn = WorkingDQN::new(config)?;
|
|
|
|
// For low variance case, temperature should remain close to base
|
|
// (Implementation will clamp to min temperature)
|
|
let base_temp = dqn.get_temperature();
|
|
assert!(
|
|
base_temp >= 0.1,
|
|
"Temperature should not drop below min, got {}",
|
|
base_temp
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_variance_multiplier_effect() -> anyhow::Result<()> {
|
|
// Test that variance_multiplier parameter controls scaling strength
|
|
|
|
// Configuration 1: No variance scaling (multiplier = 0)
|
|
let mut config1 = WorkingDQNConfig::emergency_safe_defaults();
|
|
config1.variance_multiplier = 0.0;
|
|
let dqn1 = WorkingDQN::new(config1)?;
|
|
|
|
// Configuration 2: Strong variance scaling (multiplier = 2.0)
|
|
let mut config2 = WorkingDQNConfig::emergency_safe_defaults();
|
|
config2.variance_multiplier = 2.0;
|
|
let dqn2 = WorkingDQN::new(config2)?;
|
|
|
|
// Both should initialize with same base temperature
|
|
assert_eq!(dqn1.get_temperature(), dqn2.get_temperature());
|
|
|
|
// With same Q-values but different multipliers:
|
|
// dqn2 should adapt temperature more aggressively than dqn1
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_temperature_bounds_enforcement() -> anyhow::Result<()> {
|
|
// Test that adaptive temperature respects min/max bounds
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.temperature_start = 1.0;
|
|
config.temperature_min = 0.1;
|
|
config.variance_multiplier = 5.0; // Very high scaling
|
|
|
|
// Save values before config is moved
|
|
let temp_min = config.temperature_min;
|
|
let temp_start = config.temperature_start;
|
|
|
|
let dqn = WorkingDQN::new(config)?;
|
|
|
|
// Even with high variance, temperature should not exceed max
|
|
let temp = dqn.get_temperature();
|
|
assert!(
|
|
temp <= temp_start * 2.0,
|
|
"Temperature exceeded reasonable max bound: {}",
|
|
temp
|
|
);
|
|
|
|
// Temperature should also respect minimum
|
|
assert!(
|
|
temp >= temp_min,
|
|
"Temperature dropped below minimum: {}",
|
|
temp
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_action_selection_with_variance_adaptation() -> anyhow::Result<()> {
|
|
// Integration test: Action selection with Q-variance adaptation
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.state_dim = 3;
|
|
config.num_actions = 3;
|
|
config.variance_multiplier = 1.0;
|
|
config.epsilon_start = 0.0; // Disable epsilon-greedy to isolate temperature effect
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
// Select actions with variance adaptation enabled
|
|
let state = vec![1.0_f32, 0.5, 0.0];
|
|
|
|
// Multiple action selections should work without errors
|
|
for _ in 0..10 {
|
|
let action = dqn.select_action(&state)?;
|
|
assert!(
|
|
action as u8 <= 2,
|
|
"Action index out of range: {}",
|
|
action as u8
|
|
);
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_variance_adaptation_disabled() -> anyhow::Result<()> {
|
|
// Test that variance adaptation can be disabled (multiplier = 0)
|
|
|
|
let mut config = WorkingDQNConfig::emergency_safe_defaults();
|
|
config.variance_multiplier = 0.0; // Disable variance adaptation
|
|
config.temperature_start = 0.5;
|
|
|
|
let mut dqn = WorkingDQN::new(config)?;
|
|
|
|
// Temperature should remain constant (no variance scaling)
|
|
let initial_temp = dqn.get_temperature();
|
|
let state = vec![1.0_f32, 0.5, 0.0];
|
|
|
|
// Select actions (should not modify temperature)
|
|
for _ in 0..5 {
|
|
let _action = dqn.select_action(&state)?;
|
|
}
|
|
|
|
// Temperature should decay per-epoch, but not adapt per-action
|
|
// (This test verifies adaptation is disabled, not decay)
|
|
assert_eq!(
|
|
dqn.get_temperature(),
|
|
initial_temp,
|
|
"Temperature should not change without epoch update"
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_zero_qvalues_edge_case() -> anyhow::Result<()> {
|
|
// Test variance computation with all-zero Q-values (edge case)
|
|
|
|
let device = Device::Cpu;
|
|
let q_values = vec![0.0_f32, 0.0, 0.0];
|
|
let q_tensor = Tensor::from_vec(q_values.clone(), (1, 3), &device)?;
|
|
|
|
// Variance of all-zero should be zero
|
|
let mean = 0.0;
|
|
let variance: f32 = q_values
|
|
.iter()
|
|
.map(|q| {
|
|
let diff = q - mean;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
/ q_values.len() as f32;
|
|
|
|
assert_eq!(variance, 0.0, "Zero Q-values should have zero variance");
|
|
|
|
// Variance scale should default to 1.0 (no scaling)
|
|
// Implementation should handle division by zero gracefully
|
|
let abs_mean = q_values.iter().map(|q| q.abs()).sum::<f32>() / q_values.len() as f32;
|
|
let safe_mean = abs_mean.max(0.1); // Prevent div by zero
|
|
let variance_scale = variance.sqrt() / safe_mean;
|
|
|
|
assert_eq!(
|
|
variance_scale, 0.0,
|
|
"Zero variance should produce zero scale factor"
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_negative_qvalues_handling() -> anyhow::Result<()> {
|
|
// Test variance computation with negative Q-values
|
|
|
|
let device = Device::Cpu;
|
|
let q_values = vec![-5.0_f32, -3.0, -1.0];
|
|
let q_tensor = Tensor::from_vec(q_values.clone(), (1, 3), &device)?;
|
|
|
|
// Compute variance manually
|
|
let mean = q_values.iter().sum::<f32>() / q_values.len() as f32;
|
|
let variance: f32 = q_values
|
|
.iter()
|
|
.map(|q| {
|
|
let diff = q - mean;
|
|
diff * diff
|
|
})
|
|
.sum::<f32>()
|
|
/ q_values.len() as f32;
|
|
|
|
// Variance should be positive regardless of sign
|
|
assert!(
|
|
variance > 0.0,
|
|
"Negative Q-values should still produce positive variance, got {}",
|
|
variance
|
|
);
|
|
|
|
// Use absolute values for mean to prevent sign issues
|
|
let abs_mean = q_values.iter().map(|q| q.abs()).sum::<f32>() / q_values.len() as f32;
|
|
assert!(
|
|
abs_mean > 0.0,
|
|
"Absolute mean should be positive for negative Q-values"
|
|
);
|
|
|
|
Ok(())
|
|
}
|