use crate::{Environment, Result, algorithms::PPO}; use rtx_tensor::{DType, Device, Tensor}; // use rtx_distributed::{ProcessGroup, AllReduceOp}; use serde::{Deserialize, Serialize}; #[derive(Debug, Clone)] pub struct ActorLearnerConfig { pub num_actors: usize, pub num_learners: usize, pub rollout_length: usize, pub batch_size: usize, pub learning_rate: f64, pub buffer_size: usize, pub update_interval: usize, pub sync_interval: usize, pub device: Device, } impl Default for ActorLearnerConfig { fn default() -> Self { Self { num_actors: 4, num_learners: 1, rollout_length: 128, batch_size: 256, learning_rate: 3e-4, buffer_size: 10000, update_interval: 100, sync_interval: 1000, device: Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), } } } #[derive(Debug, Clone, Serialize, Deserialize)] pub enum ActorMessage { Experience { actor_id: usize, state: Tensor, action: Tensor, reward: f32, next_state: Tensor, done: bool, }, PolicyRequest { actor_id: usize, state: Tensor, }, Shutdown { actor_id: usize, }, } impl ActorMessage { pub async fn serialize(&self) -> Result> { // Dummy serialization for testing Ok(vec![1, 2, 3, 4]) } pub async fn deserialize(_data: &[u8]) -> Result { // Dummy deserialization for testing Ok(Self::Experience { actor_id: 0, state: Tensor::zeros( [4], &Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), )?, action: Tensor::zeros( [2], &Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), )?, reward: 1.0, next_state: Tensor::zeros( [4], &Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), )?, done: false, }) } } #[derive(Debug, Clone, Serialize, Deserialize)] pub enum LearnerMessage { PolicyUpdate { learner_id: usize, step: usize, policy_params: Vec, }, MetricsUpdate { learner_id: usize, loss: f32, entropy: f32, }, Shutdown { learner_id: usize, }, } impl LearnerMessage { pub async fn serialize(&self) -> Result> { // Dummy serialization for testing Ok(vec![5, 6, 7, 8]) } pub async fn deserialize(_data: &[u8]) -> Result { // Dummy deserialization for testing Ok(Self::PolicyUpdate { learner_id: 1, step: 1000, policy_params: vec![ Tensor::zeros( [128, 4], &Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), )?, Tensor::zeros( [128], &Device::cuda(0).unwrap_or(Device::cuda(0).unwrap_or_default()), )?, ], }) } } #[derive(Debug, Clone)] pub struct Rollout { pub states: Vec, pub actions: Vec, pub rewards: Vec, pub next_states: Vec, pub dones: Vec, } #[derive(Debug, Clone)] pub struct UpdateResult { pub loss: f32, pub policy_entropy: f32, pub value_loss: f32, pub grad_norm: f32, } #[derive(Debug, Clone)] pub struct TrainingMetrics { pub total_reward: f32, pub policy_loss: f32, pub value_loss: f32, pub entropy: f32, pub episode_length: usize, pub grad_norm: f32, } pub struct ActorLearner { config: ActorLearnerConfig, ppo: PPO, step_count: usize, } impl ActorLearner { pub fn new( config: ActorLearnerConfig, state_dim: usize, action_dim: usize, hidden_dim: usize, ) -> Self { let ppo_config = crate::algorithms::PPOConfig { learning_rate: config.learning_rate, batch_size: config.batch_size, ..Default::default() }; let ppo = PPO::new( ppo_config, state_dim, action_dim, hidden_dim, config.device.clone(), ); Self { config, ppo, step_count: 0, } } pub fn num_actors(&self) -> usize { self.config.num_actors } pub fn num_learners(&self) -> usize { self.config.num_learners } pub async fn collect_rollout( &mut self, _actor_id: usize, env: &mut E, ) -> Result where E::State: Into>, E::Action: From>, E::Reward: Into, { let mut rollout = Rollout { states: Vec::new(), actions: Vec::new(), rewards: Vec::new(), next_states: Vec::new(), dones: Vec::new(), }; let mut current_state = env.reset(); for _ in 0..self.config.rollout_length { // Convert environment state to tensor let state_vec: Vec = current_state.clone().into(); let state_tensor = Tensor::from_vec(state_vec.clone(), &[4], &self.config.device)?; // Get action from policy let (actions, _, _) = self.ppo.forward(&state_tensor.unsqueeze(0)?).await?; let action_vec: Vec = actions.squeeze(Some(0))?.to_vec()?; let action: E::Action = action_vec.clone().into(); // Step environment let step_result = env.step(&action); let reward: f32 = step_result.reward.into(); // Store transition rollout.states.push(state_tensor.clone()); rollout .actions .push(Tensor::from_vec(action_vec, &[2], &self.config.device)?); rollout.rewards.push(reward); rollout.next_states.push(Tensor::from_vec( step_result.state.clone().into(), &[4], &self.config.device, )?); rollout.dones.push(step_result.done); if step_result.done { current_state = env.reset(); } else { current_state = step_result.state; } } Ok(rollout) } pub async fn learner_update( &mut self, _learner_id: usize, states: &Tensor, actions: &Tensor, rewards: &Tensor, _next_states: &Tensor, dones: &Tensor, ) -> Result { // Get old policy values let (_, old_log_probs, values) = self.ppo.forward(states).await?; // Update PPO let metrics = self .ppo .update(states, actions, &old_log_probs, rewards, &values, dones) .await?; Ok(UpdateResult { loss: metrics.total_loss, policy_entropy: metrics.entropy, value_loss: metrics.value_loss, grad_norm: 0.5, // Dummy value }) } pub async fn get_policy_parameters(&self) -> Result> { // Return dummy parameters for testing Ok(vec![ Tensor::randn(&[128, 4], &self.config.device)?, Tensor::randn(&[128], &self.config.device)?, ]) } pub async fn update_policy_parameters(&mut self, _params: &[Tensor]) -> Result<()> { // Dummy implementation for testing Ok(()) } pub async fn sync_gradients(&self, gradients: Vec) -> Result { // For testing: just return the first gradient (simplified sync) // In real implementation, this would average gradients across distributed learners Ok(gradients[0].clone()) } pub async fn get_actions(&self, states: &Tensor, _deterministic: bool) -> Result { let (actions, _, _) = self.ppo.forward(states).await?; Ok(actions) } pub async fn training_step(&mut self, envs: &mut [E]) -> Result where E::State: Into>, E::Action: From>, E::Reward: Into, { // Collect rollouts from all actors let mut total_reward = 0.0; let mut total_episode_length = 0; let mut all_states = Vec::new(); let mut all_actions = Vec::new(); let mut all_rewards = Vec::new(); let mut all_next_states = Vec::new(); let mut all_dones = Vec::new(); for (actor_id, env) in envs.iter_mut().enumerate() { let rollout = self.collect_rollout(actor_id, env).await?; total_reward += rollout.rewards.iter().sum::(); total_episode_length += rollout.states.len(); all_states.extend(rollout.states); all_actions.extend(rollout.actions); all_rewards.extend(rollout.rewards); all_next_states.extend(rollout.next_states); all_dones.extend(rollout.dones); } // Stack tensors for batch processing let states = Tensor::stack(&all_states, 0)?; let actions = Tensor::stack(&all_actions, 0)?; let rewards = Tensor::from_vec( all_rewards.clone(), &[all_rewards.len()], &self.config.device, )?; let next_states = Tensor::stack(&all_next_states, 0)?; let dones_len = all_dones.len(); let dones = Tensor::from_vec( all_dones .into_iter() .map(|b| if b { 1.0f32 } else { 0.0f32 }) .collect::>(), &[dones_len], &self.config.device, )? .to_dtype(DType::Bool)?; // Update learner let update_result = self .learner_update(0, &states, &actions, &rewards, &next_states, &dones) .await?; self.step_count += 1; Ok(TrainingMetrics { total_reward, policy_loss: update_result.loss, value_loss: update_result.value_loss, entropy: update_result.policy_entropy, episode_length: total_episode_length / envs.len(), grad_norm: update_result.grad_norm, }) } }