Files
rustytorch/demos/rtx-embodied-demo/src/actor_critic.rs
T
2026-03-04 00:08:42 +00:00

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);
}
}