Files
foxhunt/ml/tests/activation_tests.rs
jgrusewski ac0a83e4f7 refactor(ml): reorganize tests — move integration tests to ml/tests/
Move test files from ml/src/*/tests/ to ml/tests/. Convert
crate-internal imports to public API imports. Rename files to
follow naming conventions (no wave/priority prefixes).

Files moved:
- ml/src/dqn/tests/ -> ml/tests/ (4 files)
- ml/src/trainers/dqn/tests/ -> ml/tests/ (5 files)

Renames:
- target_update_comprehensive_tests.rs -> target_update_tests.rs
- p0_integration_tests.rs -> dqn_trainer_integration_tests.rs
- p1_integration_tests.rs -> dqn_trainer_p1_tests.rs
- ensemble_uncertainty_hyperopt_tests.rs -> ensemble_hyperopt_tests.rs

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-20 14:07:22 +01:00

216 lines
7.5 KiB
Rust

//! TDD tests for GELU and Mish activation functions in DQN networks
use candle_core::{Device, Tensor};
use ml::dqn::rainbow_network::ActivationType;
use ml::dqn::network::{QNetwork, QNetworkConfig};
use ml::MLError;
#[test]
fn test_activation_type_default() {
let activation = ActivationType::default();
assert!(matches!(activation, ActivationType::ReLU));
}
#[test]
fn test_gelu_activation_qnetwork() -> Result<(), MLError> {
let config = QNetworkConfig {
state_dim: 4,
num_actions: 3,
hidden_dims: vec![8, 4],
use_gpu: false,
..QNetworkConfig::default()
};
let network = QNetwork::new(config)?;
// Verify network can be created (QNetworkConfig uses LeakyReLU, not GELU)
let state = vec![1.0, 0.0, -1.0, 0.5];
let q_values = network.forward(&state)?;
// Verify output shape
assert_eq!(q_values.len(), 3);
// Verify all Q-values are finite (no NaN or Inf)
for &q in &q_values {
assert!(q.is_finite(), "Q-value should be finite, got: {}", q);
}
Ok(())
}
#[test]
fn test_mish_activation_qnetwork() -> Result<(), MLError> {
let config = QNetworkConfig {
state_dim: 4,
num_actions: 3,
hidden_dims: vec![8, 4],
use_gpu: false,
..QNetworkConfig::default()
};
let network = QNetwork::new(config)?;
// Verify network can be created (QNetworkConfig uses LeakyReLU, not Mish)
let state = vec![1.0, 0.0, -1.0, 0.5];
let q_values = network.forward(&state)?;
// Verify output shape
assert_eq!(q_values.len(), 3);
// Verify all Q-values are finite (no NaN or Inf)
for &q in &q_values {
assert!(q.is_finite(), "Q-value should be finite, got: {}", q);
}
Ok(())
}
#[test]
fn test_gelu_mathematical_properties() -> Result<(), MLError> {
let device = Device::Cpu;
// Test GELU(0) ≈ 0
let zero = Tensor::from_vec(vec![0.0_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
// Apply GELU using the approximation
let sqrt_2_over_pi = (2.0_f64 / std::f64::consts::PI).sqrt();
let coeff = Tensor::from_vec(vec![sqrt_2_over_pi as f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let c = Tensor::from_vec(vec![0.044715_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let half = Tensor::from_vec(vec![0.5_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let one = Tensor::from_vec(vec![1.0_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let x_cubed = zero.mul(&zero)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?
.mul(&zero)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let cubic_term = x_cubed.mul(&c)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let inner = zero.add(&cubic_term)
.map_err(|e| MLError::ModelError(format!("Failed to add: {}", e)))?;
let scaled = inner.mul(&coeff)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let tanh_part = scaled.tanh()
.map_err(|e| MLError::ModelError(format!("Failed to tanh: {}", e)))?;
let one_plus_tanh = one.add(&tanh_part)
.map_err(|e| MLError::ModelError(format!("Failed to add: {}", e)))?;
let gelu_zero = zero.mul(&half)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?
.mul(&one_plus_tanh)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let result = gelu_zero.to_vec0::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to convert to scalar: {}", e)))?;
assert!(result.abs() < 0.01, "GELU(0) should be close to 0, got: {}", result);
Ok(())
}
#[test]
fn test_mish_mathematical_properties() -> Result<(), MLError> {
let device = Device::Cpu;
// Test Mish(0) ≈ 0 (small positive value due to softplus)
let zero = Tensor::from_vec(vec![0.0_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let one = Tensor::from_vec(vec![1.0_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let exp_x = zero.exp()
.map_err(|e| MLError::ModelError(format!("Failed to exp: {}", e)))?;
let softplus = exp_x.add(&one)
.map_err(|e| MLError::ModelError(format!("Failed to add: {}", e)))?
.log()
.map_err(|e| MLError::ModelError(format!("Failed to log: {}", e)))?;
let tanh_softplus = softplus.tanh()
.map_err(|e| MLError::ModelError(format!("Failed to tanh: {}", e)))?;
let mish_zero = zero.mul(&tanh_softplus)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let result = mish_zero.to_vec0::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to convert to scalar: {}", e)))?;
assert!(result.abs() < 0.1, "Mish(0) should be close to 0, got: {}", result);
// Test Mish preserves sign for large positive values
let large_positive = Tensor::from_vec(vec![5.0_f32], &[], &device)
.map_err(|e| MLError::ModelError(format!("Failed to create tensor: {}", e)))?;
let exp_x = large_positive.exp()
.map_err(|e| MLError::ModelError(format!("Failed to exp: {}", e)))?;
let softplus = exp_x.add(&one)
.map_err(|e| MLError::ModelError(format!("Failed to add: {}", e)))?
.log()
.map_err(|e| MLError::ModelError(format!("Failed to log: {}", e)))?;
let tanh_softplus = softplus.tanh()
.map_err(|e| MLError::ModelError(format!("Failed to tanh: {}", e)))?;
let mish_large = large_positive.mul(&tanh_softplus)
.map_err(|e| MLError::ModelError(format!("Failed to multiply: {}", e)))?;
let result = mish_large.to_vec0::<f32>()
.map_err(|e| MLError::ModelError(format!("Failed to convert to scalar: {}", e)))?;
assert!(result > 4.0, "Mish(5.0) should be close to 5.0, got: {}", result);
Ok(())
}
#[test]
fn test_all_activations_qnetwork() -> Result<(), MLError> {
// QNetworkConfig only supports LeakyReLU activation (hardcoded in network.rs)
// Testing that the network works correctly with its fixed activation
let config = QNetworkConfig {
state_dim: 4,
num_actions: 3,
hidden_dims: vec![8],
use_gpu: false,
..QNetworkConfig::default()
};
let network = QNetwork::new(config)?;
let state = vec![1.0, 0.0, -1.0, 0.5];
let q_values = network.forward(&state)?;
assert_eq!(q_values.len(), 3);
for &q in &q_values {
assert!(q.is_finite(), "Non-finite Q-value");
}
Ok(())
}
#[test]
fn test_batch_forward_with_leaky_relu() -> Result<(), MLError> {
let config = QNetworkConfig {
state_dim: 4,
num_actions: 3,
hidden_dims: vec![8],
use_gpu: false,
..QNetworkConfig::default()
};
let network = QNetwork::new(config)?;
let states = vec![
vec![1.0, 0.0, -1.0, 0.5],
vec![0.5, 1.0, -0.5, 0.0],
];
let q_values_batch = network.forward_batch(&states)?;
assert_eq!(q_values_batch.len(), 2);
assert_eq!(q_values_batch[0].len(), 3);
assert_eq!(q_values_batch[1].len(), 3);
for batch in &q_values_batch {
for &q in batch {
assert!(q.is_finite());
}
}
Ok(())
}