Files
rustytorch/crates/training/rtx-rl/tests/sac_tests.rs
T
2026-03-04 00:08:42 +00:00

310 lines
10 KiB
Rust

use rtx_rl::algorithms::{SAC, SACConfig};
use rtx_tensor::{DType, Device, Tensor};
#[tokio::test]
async fn test_sac_creation() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig {
learning_rate: 3e-4,
gamma: 0.99,
tau: 0.005,
alpha: 0.2,
target_update_interval: 1,
automatic_entropy_tuning: true,
target_entropy: None,
replay_buffer_size: 1_000_000,
batch_size: 256,
};
let state_dim = 8;
let action_dim = 3;
let hidden_dim = 256;
let sac = SAC::new(config, state_dim, action_dim, hidden_dim, device);
assert_eq!(sac.state_dim(), state_dim);
assert_eq!(sac.action_dim(), action_dim);
}
#[tokio::test]
#[ignore = "rtx-rl SAC actor tanh bounds incomplete"]
async fn test_sac_actor_forward() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let state_dim = 8;
let action_dim = 3;
let hidden_dim = 256;
let batch_size = 32;
let sac = SAC::new(config, state_dim, action_dim, hidden_dim, device.clone());
let states = Tensor::randn(&[batch_size, state_dim], &device).unwrap();
let (actions, log_probs) = sac
.actor_forward(&states, false)
.await
.expect("Actor forward should work");
assert_eq!(actions.shape(), &[batch_size, action_dim]);
assert_eq!(log_probs.shape(), &[batch_size]);
// Actions should be in [-1, 1] after tanh
let action_data: Vec<f32> = actions.to_vec().expect("Should convert");
for action in action_data {
assert!(action >= -1.0 && action <= 1.0);
}
}
#[tokio::test]
async fn test_sac_critic_forward() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 8, 3, 256, device.clone());
let batch_size = 32;
let states = Tensor::randn(&[batch_size, 8], &device).unwrap();
let actions = Tensor::randn(&[batch_size, 3], &device).unwrap();
let q_values = sac
.critic_forward(&states, &actions, 0)
.await
.expect("Critic forward should work");
assert_eq!(q_values.shape(), &[batch_size]);
}
#[tokio::test]
async fn test_sac_target_critic_forward() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 8, 3, 256, device.clone());
let batch_size = 32;
let states = Tensor::randn(&[batch_size, 8], &device).unwrap();
let actions = Tensor::randn(&[batch_size, 3], &device).unwrap();
let q_values = sac
.target_critic_forward(&states, &actions, 0)
.await
.expect("Target critic forward should work");
assert_eq!(q_values.shape(), &[batch_size]);
}
#[tokio::test]
#[ignore = "rtx-rl SAC loss shape mismatch"]
async fn test_sac_compute_actor_loss() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 8, 3, 256, device.clone());
let batch_size = 64;
let states = Tensor::randn(&[batch_size, 8], &device).unwrap();
let actor_loss = sac
.compute_actor_loss(&states)
.await
.expect("Should compute actor loss");
assert_eq!(actor_loss.shape(), &[]);
let loss_val: f32 = actor_loss.item().expect("Should get scalar");
assert!(loss_val.is_finite());
}
#[tokio::test]
#[ignore = "rtx-tensor Bool dtype CPU copy not supported"]
async fn test_sac_compute_critic_loss() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 8, 3, 256, device.clone());
let batch_size = 64;
let states = Tensor::randn(&[batch_size, 8], &device).unwrap();
let actions = Tensor::randn(&[batch_size, 3], &device).unwrap();
let rewards = Tensor::randn(&[batch_size], &device).unwrap();
let next_states = Tensor::randn(&[batch_size, 8], &device).unwrap();
let dones = Tensor::zeros(&[batch_size], &device)
.unwrap()
.to_dtype(DType::Bool)
.unwrap();
let (critic1_loss, critic2_loss) = sac
.compute_critic_loss(&states, &actions, &rewards, &next_states, &dones)
.await
.expect("Should compute critic losses");
assert_eq!(critic1_loss.shape(), &[]);
assert_eq!(critic2_loss.shape(), &[]);
let loss1_val: f32 = critic1_loss.item().expect("Should get scalar");
let loss2_val: f32 = critic2_loss.item().expect("Should get scalar");
assert!(loss1_val >= 0.0); // MSE loss should be non-negative
assert!(loss2_val >= 0.0);
}
#[tokio::test]
#[ignore = "rtx-rl SAC temperature implementation incomplete"]
async fn test_sac_compute_temperature_loss() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig {
automatic_entropy_tuning: true,
target_entropy: Some(-3.0),
..SACConfig::default()
};
let sac = SAC::new(config, 8, 3, 256, device.clone());
let batch_size = 64;
let log_probs = Tensor::randn(&[batch_size], &device)
.unwrap()
.scalar_mul(-2.0)
.unwrap(); // Negative log probs
let temperature_loss = sac
.compute_temperature_loss(&log_probs)
.await
.expect("Should compute temperature loss");
assert_eq!(temperature_loss.shape(), &[]);
let loss_val: f32 = temperature_loss.item().expect("Should get scalar");
assert!(loss_val.is_finite());
}
#[tokio::test]
#[ignore = "rtx-tensor Bool dtype CPU copy not supported"]
async fn test_sac_soft_update() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig {
tau: 0.1, // Larger tau for more noticeable updates
..SACConfig::default()
};
let mut sac = SAC::new(config, 4, 2, 64, device.clone());
// Get initial target parameters
let initial_target_params = sac
.get_target_critic_params()
.await
.expect("Should get initial params");
// Modify main critic parameters by training on some dummy data
let batch_size = 32;
let states = Tensor::randn(&[batch_size, 4], &device).unwrap();
let actions = Tensor::randn(&[batch_size, 2], &device).unwrap();
let rewards = Tensor::ones(&[batch_size], &device).unwrap();
let next_states = Tensor::randn(&[batch_size, 4], &device).unwrap();
let dones = Tensor::zeros(&[batch_size], &device)
.unwrap()
.to_dtype(DType::Bool)
.unwrap();
// Perform one update step
let _metrics = sac
.update(&states, &actions, &rewards, &next_states, &dones)
.await
.expect("Update should succeed");
// Verify target networks were soft-updated
let updated_target_params = sac
.get_target_critic_params()
.await
.expect("Should get updated params");
// Target parameters should have changed but not completely
assert_ne!(initial_target_params.len(), 0);
assert_eq!(initial_target_params.len(), updated_target_params.len());
}
#[tokio::test]
#[ignore = "rtx-rl SAC deterministic action incomplete"]
async fn test_sac_deterministic_action() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 4, 2, 64, device.clone());
let state = Tensor::randn(&[1, 4], &device).unwrap();
let (action1, _) = sac
.actor_forward(&state, true)
.await
.expect("Should get deterministic action");
let (action2, _) = sac
.actor_forward(&state, true)
.await
.expect("Should get deterministic action");
// Deterministic actions should be identical for same state
let action1_data: Vec<f32> = action1.to_vec().expect("Should convert");
let action2_data: Vec<f32> = action2.to_vec().expect("Should convert");
for (a1, a2) in action1_data.iter().zip(action2_data.iter()) {
assert!((a1 - a2).abs() < 1e-6);
}
}
#[tokio::test]
async fn test_sac_stochastic_action() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig::default();
let sac = SAC::new(config, 4, 2, 64, device.clone());
let state = Tensor::randn(&[1, 4], &device).unwrap();
let (action1, _) = sac
.actor_forward(&state, false)
.await
.expect("Should get stochastic action");
let (action2, _) = sac
.actor_forward(&state, false)
.await
.expect("Should get stochastic action");
// Stochastic actions should be different for same state
let action1_data: Vec<f32> = action1.to_vec().expect("Should convert");
let action2_data: Vec<f32> = action2.to_vec().expect("Should convert");
let mut different = false;
for (a1, a2) in action1_data.iter().zip(action2_data.iter()) {
if (a1 - a2).abs() > 1e-3 {
different = true;
break;
}
}
assert!(different, "Stochastic actions should be different");
}
#[tokio::test]
#[ignore = "rtx-tensor Bool dtype CPU copy not supported"]
async fn test_sac_full_update_cycle() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = SACConfig {
batch_size: 32,
target_update_interval: 1,
..SACConfig::default()
};
let mut sac = SAC::new(config, 4, 2, 64, device.clone());
let batch_size = 32;
let states = Tensor::randn(&[batch_size, 4], &device).unwrap();
let actions = Tensor::randn(&[batch_size, 2], &device).unwrap();
let rewards = Tensor::randn(&[batch_size], &device).unwrap();
let next_states = Tensor::randn(&[batch_size, 4], &device).unwrap();
let dones = Tensor::zeros(&[batch_size], &device)
.unwrap()
.to_dtype(DType::Bool)
.unwrap();
let metrics = sac
.update(&states, &actions, &rewards, &next_states, &dones)
.await
.expect("Full update should succeed");
assert!(metrics.actor_loss.is_finite());
assert!(metrics.critic1_loss.is_finite());
assert!(metrics.critic2_loss.is_finite());
assert!(metrics.temperature_loss.is_finite());
assert!(metrics.alpha > 0.0);
assert!(metrics.target_entropy.is_finite());
}