385 lines
12 KiB
Rust
385 lines
12 KiB
Rust
//! 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<usize>,
|
|
/// Critic hidden dimensions.
|
|
critic_hidden: Vec<usize>,
|
|
/// Actor weights.
|
|
actor_weights: Vec<f64>,
|
|
/// Actor log std.
|
|
actor_log_std: Vec<f64>,
|
|
/// Critic weights.
|
|
critic_weights: Vec<f64>,
|
|
/// 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<f64> {
|
|
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<f64> {
|
|
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<f64>],
|
|
advantages: &[Vec<f64>],
|
|
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);
|
|
}
|
|
}
|