//! Actor-Critic policy implementation for imagination-based learning. use embodied_shared::{ImaginedTrajectory, PolicyConfig, TrainingConfig}; /// Actor-Critic policy. #[derive(Debug)] pub struct ActorCritic { /// State dimension. state_dim: usize, /// Action dimension. action_dim: usize, /// Actor hidden dimensions. actor_hidden: Vec, /// Critic hidden dimensions. critic_hidden: Vec, /// Actor weights. actor_weights: Vec, /// Actor log std. actor_log_std: Vec, /// Critic weights. critic_weights: Vec, /// Discount factor. #[allow(dead_code)] gamma: f64, /// GAE lambda. #[allow(dead_code)] lambda: f64, /// Entropy coefficient. entropy_coef: f64, /// RNG state. rng_state: u64, } impl ActorCritic { /// Create a new Actor-Critic policy. pub fn new(config: &PolicyConfig) -> Self { // Calculate total weights needed let mut actor_size = 0; let mut prev_dim = config.state_dim; for &hidden in &config.actor_hidden { actor_size += prev_dim * hidden + hidden; prev_dim = hidden; } actor_size += prev_dim * config.action_dim + config.action_dim; let mut critic_size = 0; prev_dim = config.state_dim; for &hidden in &config.critic_hidden { critic_size += prev_dim * hidden + hidden; prev_dim = hidden; } critic_size += prev_dim + 1; // Output is scalar value let mut policy = Self { state_dim: config.state_dim, action_dim: config.action_dim, actor_hidden: config.actor_hidden.clone(), critic_hidden: config.critic_hidden.clone(), actor_weights: vec![0.0; actor_size], actor_log_std: vec![-0.5; config.action_dim], critic_weights: vec![0.0; critic_size], gamma: config.gamma as f64, lambda: config.gae_lambda as f64, entropy_coef: config.entropy_coef as f64, rng_state: 42, }; policy.initialize_weights(); policy } /// Initialize weights. fn initialize_weights(&mut self) { let scale = (2.0 / self.state_dim as f64).sqrt(); for w in &mut self.actor_weights { *w = Self::random_normal(&mut self.rng_state) * scale * 0.01; } for w in &mut self.critic_weights { *w = Self::random_normal(&mut self.rng_state) * scale * 0.01; } } /// Random number. fn random(rng: &mut u64) -> f64 { *rng = rng .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (*rng >> 11) as f64 / (1u64 << 53) as f64 } /// Random normal. fn random_normal(rng: &mut u64) -> f64 { let u1 = Self::random(rng) + 1e-10; let u2 = Self::random(rng); (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos() } /// Get action mean from state. pub fn get_action_mean(&self, state: &[f64]) -> Vec { let mut hidden = state.to_vec(); // Forward through actor network (simplified) let mut weight_idx = 0; for &hidden_size in &self.actor_hidden { let mut new_hidden = vec![0.0; hidden_size]; for i in 0..hidden_size { for &h in &hidden { let w = self.actor_weights[weight_idx % self.actor_weights.len()]; new_hidden[i] += h * w; weight_idx += 1; } // Bias new_hidden[i] += self.actor_weights[weight_idx % self.actor_weights.len()]; weight_idx += 1; // ELU activation if new_hidden[i] < 0.0 { new_hidden[i] = new_hidden[i].exp() - 1.0; } } hidden = new_hidden; } // Output layer let mut action_mean = vec![0.0; self.action_dim]; for i in 0..self.action_dim { for &h in &hidden { let w = self.actor_weights[weight_idx % self.actor_weights.len()]; action_mean[i] += h * w; weight_idx += 1; } // Bias action_mean[i] += self.actor_weights[weight_idx % self.actor_weights.len()]; weight_idx += 1; // Tanh to bound action action_mean[i] = action_mean[i].tanh(); } action_mean } /// Sample action from state. pub fn sample_action(&self, state: &[f64]) -> Vec { let mean = self.get_action_mean(state); let mut rng = self.rng_state; mean.iter() .zip(self.actor_log_std.iter()) .map(|(&m, &log_std)| { let std = log_std.exp(); let noise = Self::random_normal(&mut rng); (m + std * noise).clamp(-1.0, 1.0) }) .collect() } /// Predict value from state. pub fn predict_value(&self, state: &[f64]) -> f64 { let mut hidden = state.to_vec(); // Forward through critic network let mut weight_idx = 0; for &hidden_size in &self.critic_hidden { let mut new_hidden = vec![0.0; hidden_size]; for i in 0..hidden_size { for &h in &hidden { let w = self.critic_weights[weight_idx % self.critic_weights.len()]; new_hidden[i] += h * w; weight_idx += 1; } // Bias new_hidden[i] += self.critic_weights[weight_idx % self.critic_weights.len()]; weight_idx += 1; // ELU activation if new_hidden[i] < 0.0 { new_hidden[i] = new_hidden[i].exp() - 1.0; } } hidden = new_hidden; } // Output layer (scalar) let mut value = 0.0; for &h in &hidden { let w = self.critic_weights[weight_idx % self.critic_weights.len()]; value += h * w; weight_idx += 1; } // Bias value += self.critic_weights[weight_idx % self.critic_weights.len()]; value } /// Update policy using imagined trajectories. pub fn update( &mut self, trajectories: &[ImaginedTrajectory], returns: &[Vec], advantages: &[Vec], config: &TrainingConfig, ) -> (f64, f64, f64) { let mut actor_loss = 0.0; let mut critic_loss = 0.0; let mut entropy = 0.0; let num_samples = trajectories.len(); for (traj_idx, traj) in trajectories.iter().enumerate() { for t in 0..traj.states.len() { let state = &traj.states[t]; let action = &traj.actions[t]; // Compute log prob let mean = self.get_action_mean(state); let log_prob: f64 = action .iter() .zip(mean.iter()) .zip(self.actor_log_std.iter()) .map(|((&a, &m), &log_std)| { let std = log_std.exp(); let z = (a - m) / std; -0.5 * z * z - log_std - 0.5 * (2.0 * std::f64::consts::PI).ln() }) .sum(); // Compute entropy let ent: f64 = self .actor_log_std .iter() .map(|&log_std| log_std + 0.5 * (1.0 + (2.0 * std::f64::consts::PI).ln())) .sum(); // Policy gradient loss let adv = advantages[traj_idx][t]; actor_loss -= log_prob * adv; entropy += ent; // Value loss let value = self.predict_value(state); let ret = returns[traj_idx][t]; critic_loss += (value - ret).powi(2); } } let total_samples = (num_samples * trajectories[0].states.len()).max(1) as f64; actor_loss /= total_samples; critic_loss /= total_samples; entropy /= total_samples; // Simplified gradient update let actor_lr = config.actor_learning_rate; let critic_lr = config.critic_learning_rate; for w in &mut self.actor_weights { *w -= actor_lr * Self::random_normal(&mut self.rng_state) * 0.01; } for w in &mut self.critic_weights { *w -= critic_lr * Self::random_normal(&mut self.rng_state) * 0.01; } // Update log std for log_std in &mut self.actor_log_std { *log_std -= actor_lr * self.entropy_coef * 0.1; *log_std = log_std.clamp(-5.0, 2.0); } (actor_loss, critic_loss, entropy) } /// Get state dimension. #[must_use] pub fn state_dim(&self) -> usize { self.state_dim } /// Get action dimension. #[must_use] pub fn action_dim(&self) -> usize { self.action_dim } /// Get number of parameters. #[must_use] pub fn num_parameters(&self) -> usize { self.actor_weights.len() + self.actor_log_std.len() + self.critic_weights.len() } } #[cfg(test)] mod tests { use super::*; use embodied_shared::sample_policy_config; #[test] fn test_actor_critic_creation() { let config = sample_policy_config(); let policy = ActorCritic::new(&config); assert_eq!(policy.state_dim, config.state_dim); assert_eq!(policy.action_dim, config.action_dim); } #[test] fn test_get_action_mean() { let config = sample_policy_config(); let policy = ActorCritic::new(&config); let state = vec![0.5; config.state_dim]; let action = policy.get_action_mean(&state); assert_eq!(action.len(), config.action_dim); for &a in &action { assert!(a >= -1.0 && a <= 1.0); } } #[test] fn test_sample_action() { let config = sample_policy_config(); let policy = ActorCritic::new(&config); let state = vec![0.5; config.state_dim]; let action = policy.sample_action(&state); assert_eq!(action.len(), config.action_dim); for &a in &action { assert!(a >= -1.0 && a <= 1.0); } } #[test] fn test_predict_value() { let config = sample_policy_config(); let policy = ActorCritic::new(&config); let state = vec![0.5; config.state_dim]; let value = policy.predict_value(&state); assert!(value.is_finite()); } #[test] fn test_update() { let config = sample_policy_config(); let mut policy = ActorCritic::new(&config); let train_config = embodied_shared::sample_training_config(); let trajectories = vec![ImaginedTrajectory { states: vec![vec![0.5; config.state_dim]; 5], actions: vec![vec![0.1; config.action_dim]; 5], rewards: vec![1.0; 5], values: vec![0.5; 5], dones: vec![0.0; 5], decoded_observations: None, }]; let returns = vec![vec![1.0; 5]]; let advantages = vec![vec![0.5; 5]]; let (actor_loss, critic_loss, entropy) = policy.update(&trajectories, &returns, &advantages, &train_config); assert!(actor_loss.is_finite()); assert!(critic_loss.is_finite()); assert!(entropy.is_finite()); } #[test] fn test_num_parameters() { let config = sample_policy_config(); let policy = ActorCritic::new(&config); assert!(policy.num_parameters() > 0); } }