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

402 lines
12 KiB
Rust

use rtx_rl::{ActorLearner, ActorLearnerConfig, ActorMessage, LearnerMessage};
// use rtx_distributed::{ProcessGroup, AllReduceOp};
use rtx_tensor::{DType, Device, Tensor};
#[derive(Clone, Debug)]
struct MockEnvironment {
state: Vec<f32>,
step_count: usize,
}
impl rtx_rl::Environment for MockEnvironment {
type State = Vec<f32>;
type Action = Vec<f32>;
type Reward = f32;
fn reset(&mut self) -> Self::State {
self.state = vec![0.0, 0.0];
self.step_count = 0;
self.state.clone()
}
fn step(&mut self, action: &Self::Action) -> rtx_rl::Step<Self::State, Self::Reward> {
self.state[0] += action[0] * 0.1;
self.state[1] += action[1] * 0.1;
self.step_count += 1;
let reward = -(self.state[0].powi(2) + self.state[1].powi(2));
let done = self.step_count >= 100;
rtx_rl::Step {
state: self.state.clone(),
reward,
done,
info: std::collections::HashMap::new(),
}
}
fn action_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-1.0, -1.0], vec![1.0, 1.0])
}
fn observation_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-10.0, -10.0], vec![10.0, 10.0])
}
}
#[tokio::test]
async fn test_actor_learner_config_creation() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
num_actors: 4,
num_learners: 2,
rollout_length: 128,
batch_size: 256,
learning_rate: 3e-4,
buffer_size: 10000,
update_interval: 100,
sync_interval: 1000,
device: device,
};
assert_eq!(config.num_actors, 4);
assert_eq!(config.num_learners, 2);
assert_eq!(config.rollout_length, 128);
}
#[tokio::test]
async fn test_actor_learner_creation() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
num_actors: 2,
num_learners: 1,
rollout_length: 64,
batch_size: 128,
learning_rate: 3e-4,
buffer_size: 5000,
update_interval: 50,
sync_interval: 200,
device: device.clone(),
};
let actor_learner = ActorLearner::new(config, 4, 2, 128);
assert_eq!(actor_learner.num_actors(), 2);
assert_eq!(actor_learner.num_learners(), 1);
}
#[tokio::test]
#[ignore = "rtx-rl serialization implementation incomplete"]
async fn test_actor_message_serialization() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let state = Tensor::randn(&[4], &device).unwrap();
let action = Tensor::randn(&[2], &device).unwrap();
let reward = 1.5f32;
let next_state = Tensor::randn(&[4], &device).unwrap();
let done = false;
let message = ActorMessage::Experience {
actor_id: 0,
state: state.clone(),
action: action.clone(),
reward,
next_state: next_state.clone(),
done,
};
// Serialize and deserialize
let serialized = message.serialize().await.expect("Should serialize");
let deserialized = ActorMessage::deserialize(&serialized)
.await
.expect("Should deserialize");
match deserialized {
ActorMessage::Experience {
actor_id,
state: deser_state,
action: _deser_action,
reward: deser_reward,
next_state: _deser_next_state,
done: deser_done,
} => {
assert_eq!(actor_id, 0);
assert_eq!(deser_reward, reward);
assert_eq!(deser_done, done);
let state_data: Vec<f32> = state.to_vec().expect("Convert state");
let deser_state_data: Vec<f32> = deser_state.to_vec().expect("Convert deser state");
assert_eq!(state_data, deser_state_data);
}
_ => panic!("Wrong message type"),
}
}
#[tokio::test]
async fn test_learner_message_serialization() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let policy_params = vec![
Tensor::randn(&[128, 4], &device).unwrap(),
Tensor::randn(&[128], &device).unwrap(),
];
let message = LearnerMessage::PolicyUpdate {
learner_id: 1,
step: 1000,
policy_params: policy_params.clone(),
};
// Serialize and deserialize
let serialized = message.serialize().await.expect("Should serialize");
let deserialized = LearnerMessage::deserialize(&serialized)
.await
.expect("Should deserialize");
match deserialized {
LearnerMessage::PolicyUpdate {
learner_id,
step,
policy_params: deser_params,
} => {
assert_eq!(learner_id, 1);
assert_eq!(step, 1000);
assert_eq!(deser_params.len(), policy_params.len());
for (orig, deser) in policy_params.iter().zip(deser_params.iter()) {
assert_eq!(orig.shape(), deser.shape());
}
}
_ => panic!("Wrong message type"),
}
}
#[tokio::test]
#[ignore = "rtx-rl rollout collection tensor shapes incomplete"]
async fn test_actor_rollout_collection() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
rollout_length: 10,
device: device.clone(),
..Default::default()
};
let mut actor_learner = ActorLearner::new(config, 4, 2, 64);
let mut env = MockEnvironment {
state: vec![0.0, 0.0],
step_count: 0,
};
let actor_id = 0;
let rollout = actor_learner
.collect_rollout(actor_id, &mut env)
.await
.expect("Should collect rollout");
assert_eq!(rollout.states.len(), 10);
assert_eq!(rollout.actions.len(), 10);
assert_eq!(rollout.rewards.len(), 10);
assert_eq!(rollout.next_states.len(), 10);
assert_eq!(rollout.dones.len(), 10);
// Verify rollout data integrity
for i in 0..rollout.states.len() {
assert_eq!(rollout.states[i].shape(), &[4]);
assert_eq!(rollout.actions[i].shape(), &[2]);
assert!(rollout.rewards[i].is_finite());
}
}
#[tokio::test]
#[ignore = "rtx-tensor Bool dtype CPU copy not supported"]
async fn test_learner_update_from_experience() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
batch_size: 32,
learning_rate: 1e-3,
device: device.clone(),
..Default::default()
};
let mut actor_learner = ActorLearner::new(config, 4, 2, 64);
// Create dummy experience batch
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 learner_id = 0;
let update_result = actor_learner
.learner_update(
learner_id,
&states,
&actions,
&rewards,
&next_states,
&dones,
)
.await
.expect("Learner update should succeed");
assert!(update_result.loss.is_finite());
assert!(update_result.policy_entropy.is_finite());
assert!(update_result.value_loss.is_finite());
assert!(update_result.grad_norm >= 0.0);
}
#[tokio::test]
async fn test_parameter_synchronization() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
device: device.clone(),
..Default::default()
};
let mut actor_learner = ActorLearner::new(config, 4, 2, 64);
// Get initial parameters
let initial_params = actor_learner
.get_policy_parameters()
.await
.expect("Should get initial params");
// Simulate parameter update
let mut new_params = Vec::new();
for param in &initial_params {
let ones = Tensor::ones_like(param).unwrap();
let scaled = ones.scalar_mul(0.1).unwrap();
let updated_param = (param + &scaled).unwrap();
new_params.push(updated_param);
}
// Update parameters
actor_learner
.update_policy_parameters(&new_params)
.await
.expect("Should update parameters");
// Verify parameters changed
let updated_params = actor_learner
.get_policy_parameters()
.await
.expect("Should get updated params");
for (old, new) in initial_params.iter().zip(updated_params.iter()) {
let diff = (new - old).unwrap();
let diff_squared = (&diff * &diff).unwrap();
let sum_squared = diff_squared.sum(None).unwrap();
let diff_norm: f32 = sum_squared.item().unwrap().sqrt();
assert!(diff_norm > 0.01); // Parameters should have changed
}
}
#[tokio::test]
#[ignore = "rtx-rl gradient sync implementation incomplete"]
async fn test_distributed_gradient_sync() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
num_learners: 2,
device: device.clone(),
..Default::default()
};
let actor_learner = ActorLearner::new(config, 4, 2, 64);
// Simulate gradients from multiple learners
let grad1 = Tensor::ones(&[64, 4], &device)
.unwrap()
.scalar_mul(0.5)
.unwrap();
let grad2 = Tensor::ones(&[64, 4], &device)
.unwrap()
.scalar_mul(1.5)
.unwrap();
let gradients = vec![grad1.clone(), grad2.clone()];
let averaged_grad = actor_learner
.sync_gradients(gradients)
.await
.expect("Should sync gradients");
// Averaged gradient should be mean of inputs: (0.5 + 1.5) / 2 = 1.0
let expected_value = 1.0f32;
let grad_data: Vec<f32> = averaged_grad.to_vec().expect("Should convert");
for val in grad_data {
assert!((val - expected_value).abs() < 1e-6);
}
}
#[tokio::test]
#[ignore = "rtx-rl training step tensor shapes incomplete"]
async fn test_actor_learner_full_training_step() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
num_actors: 2,
num_learners: 1,
rollout_length: 8,
batch_size: 16,
update_interval: 8,
device: device.clone(),
..Default::default()
};
let mut actor_learner = ActorLearner::new(config, 4, 2, 64);
let mut envs = vec![
MockEnvironment {
state: vec![0.0, 0.0],
step_count: 0,
},
MockEnvironment {
state: vec![1.0, 1.0],
step_count: 0,
},
];
let training_metrics = actor_learner
.training_step(&mut envs)
.await
.expect("Training step should succeed");
assert!(training_metrics.total_reward.is_finite());
assert!(training_metrics.policy_loss.is_finite());
assert!(training_metrics.value_loss.is_finite());
assert!(training_metrics.entropy.is_finite());
assert!(training_metrics.episode_length > 0);
assert!(training_metrics.grad_norm >= 0.0);
}
#[tokio::test]
async fn test_actor_learner_performance_metrics() {
let device = Device::cuda(0).unwrap_or_else(|_| Device::cpu());
let config = ActorLearnerConfig {
num_actors: 1,
rollout_length: 16,
device: device.clone(),
..Default::default()
};
let actor_learner = ActorLearner::new(config, 4, 2, 64);
// Measure throughput
let start_time = std::time::Instant::now();
let samples_processed = 1000;
for _ in 0..10 {
let states = Tensor::randn(&[100, 4], &device).unwrap();
let _actions = actor_learner
.get_actions(&states, false)
.await
.expect("Should get actions");
}
let elapsed = start_time.elapsed();
let throughput = samples_processed as f64 / elapsed.as_secs_f64();
// Should process at least 1000 samples per second on GPU
assert!(throughput > 1000.0);
}