366 lines
10 KiB
Rust
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,
|
|
})
|
|
}
|
|
}
|