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

366 lines
10 KiB
Rust

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<Vec<u8>> {
// Dummy serialization for testing
Ok(vec![1, 2, 3, 4])
}
pub async fn deserialize(_data: &[u8]) -> Result<Self> {
// 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<Tensor>,
},
MetricsUpdate {
learner_id: usize,
loss: f32,
entropy: f32,
},
Shutdown {
learner_id: usize,
},
}
impl LearnerMessage {
pub async fn serialize(&self) -> Result<Vec<u8>> {
// Dummy serialization for testing
Ok(vec![5, 6, 7, 8])
}
pub async fn deserialize(_data: &[u8]) -> Result<Self> {
// 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<Tensor>,
pub actions: Vec<Tensor>,
pub rewards: Vec<f32>,
pub next_states: Vec<Tensor>,
pub dones: Vec<bool>,
}
#[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<E: Environment>(
&mut self,
_actor_id: usize,
env: &mut E,
) -> Result<Rollout>
where
E::State: Into<Vec<f32>>,
E::Action: From<Vec<f32>>,
E::Reward: Into<f32>,
{
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<f32> = 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<f32> = 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<UpdateResult> {
// 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<Vec<Tensor>> {
// 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<Tensor>) -> Result<Tensor> {
// 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<Tensor> {
let (actions, _, _) = self.ppo.forward(states).await?;
Ok(actions)
}
pub async fn training_step<E: Environment>(&mut self, envs: &mut [E]) -> Result<TrainingMetrics>
where
E::State: Into<Vec<f32>>,
E::Action: From<Vec<f32>>,
E::Reward: Into<f32>,
{
// 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::<f32>();
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::<Vec<f32>>(),
&[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,
})
}
}