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

435 lines
12 KiB
Rust

//! Standalone RL implementation to demonstrate TDD principles
//! This shows the Red-Green-Refactor cycle for core RL concepts
use rand::{Rng, thread_rng};
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
/// Simple tensor-like structure for demonstration
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Matrix {
data: Vec<f32>,
shape: Vec<usize>,
}
impl Matrix {
pub fn new(data: Vec<f32>, shape: Vec<usize>) -> Self {
assert_eq!(data.len(), shape.iter().product::<usize>());
Self { data, shape }
}
pub fn zeros(shape: Vec<usize>) -> Self {
let size = shape.iter().product::<usize>();
Self {
data: vec![0.0; size],
shape,
}
}
pub fn randn(shape: Vec<usize>) -> Self {
let size = shape.iter().product::<usize>();
let mut rng = thread_rng();
let data = (0..size).map(|_| rng.gen_range(-1.0..1.0)).collect();
Self { data, shape }
}
pub fn shape(&self) -> &[usize] {
&self.shape
}
pub fn data(&self) -> &[f32] {
&self.data
}
pub fn add(&self, other: &Self) -> Self {
assert_eq!(self.shape, other.shape);
let data = self
.data
.iter()
.zip(&other.data)
.map(|(a, b)| a + b)
.collect();
Self::new(data, self.shape.clone())
}
pub fn mul(&self, scalar: f32) -> Self {
let data = self.data.iter().map(|x| x * scalar).collect();
Self::new(data, self.shape.clone())
}
pub fn mean(&self) -> f32 {
self.data.iter().sum::<f32>() / self.data.len() as f32
}
}
/// Experience tuple for RL
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Experience {
pub state: Vec<f32>,
pub action: Vec<f32>,
pub reward: f32,
pub next_state: Vec<f32>,
pub done: bool,
}
/// Replay buffer for experience storage
pub struct ReplayBuffer {
buffer: VecDeque<Experience>,
capacity: usize,
}
impl ReplayBuffer {
pub fn new(capacity: usize) -> Self {
Self {
buffer: VecDeque::with_capacity(capacity),
capacity,
}
}
pub fn push(&mut self, experience: Experience) {
if self.buffer.len() >= self.capacity {
self.buffer.pop_front();
}
self.buffer.push_back(experience);
}
pub fn sample(&self, batch_size: usize) -> Option<Vec<Experience>> {
if self.buffer.len() < batch_size {
return None;
}
let mut rng = thread_rng();
let mut samples = Vec::with_capacity(batch_size);
for _ in 0..batch_size {
let idx = rng.gen_range(0..self.buffer.len());
samples.push(self.buffer[idx].clone());
}
Some(samples)
}
pub fn len(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
}
/// Simple policy network
pub struct PolicyNetwork {
weights: Matrix,
bias: Matrix,
}
impl PolicyNetwork {
pub fn new(input_dim: usize, output_dim: usize) -> Self {
Self {
weights: Matrix::randn(vec![output_dim, input_dim]),
bias: Matrix::randn(vec![output_dim]),
}
}
pub fn forward(&self, input: &[f32]) -> Vec<f32> {
// Simple linear transformation: output = weights * input + bias
let mut output = self.bias.data().to_vec();
for (i, bias_val) in output.iter_mut().enumerate() {
let mut weighted_sum = *bias_val;
for (j, &input_val) in input.iter().enumerate() {
let weight_idx = i * input.len() + j;
if weight_idx < self.weights.data().len() {
weighted_sum += self.weights.data()[weight_idx] * input_val;
}
}
*bias_val = weighted_sum.tanh(); // Tanh activation
}
output
}
}
/// PPO Algorithm implementation
pub struct PPO {
policy: PolicyNetwork,
value_network: PolicyNetwork,
learning_rate: f32,
clip_epsilon: f32,
}
impl PPO {
pub fn new(state_dim: usize, action_dim: usize) -> Self {
Self {
policy: PolicyNetwork::new(state_dim, action_dim),
value_network: PolicyNetwork::new(state_dim, 1),
learning_rate: 3e-4,
clip_epsilon: 0.2,
}
}
pub fn get_action(&self, state: &[f32]) -> Vec<f32> {
self.policy.forward(state)
}
pub fn get_value(&self, state: &[f32]) -> f32 {
self.value_network.forward(state)[0]
}
pub fn compute_advantages(&self, experiences: &[Experience]) -> Vec<f32> {
let mut advantages = Vec::new();
let gamma = 0.99;
for exp in experiences {
let value = self.get_value(&exp.state);
let next_value = if exp.done {
0.0
} else {
self.get_value(&exp.next_state)
};
let advantage = exp.reward + gamma * next_value - value;
advantages.push(advantage);
}
advantages
}
pub fn update(&mut self, experiences: &[Experience]) -> f32 {
let advantages = self.compute_advantages(experiences);
// Simplified policy loss computation
let mut policy_loss = 0.0;
for (exp, advantage) in experiences.iter().zip(&advantages) {
let action = self.get_action(&exp.state);
// Simplified loss: advantage * action magnitude
let action_magnitude: f32 = action.iter().map(|x| x.abs()).sum();
policy_loss += advantage * action_magnitude;
}
policy_loss / experiences.len() as f32
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_matrix_creation() {
let matrix = Matrix::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]);
assert_eq!(matrix.shape(), &[2, 2]);
assert_eq!(matrix.data(), &[1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_matrix_zeros() {
let matrix = Matrix::zeros(vec![3, 3]);
assert_eq!(matrix.shape(), &[3, 3]);
assert_eq!(matrix.data(), &[0.0; 9]);
}
#[test]
fn test_matrix_randn() {
let matrix = Matrix::randn(vec![2, 2]);
assert_eq!(matrix.shape(), &[2, 2]);
assert_eq!(matrix.data().len(), 4);
// Values should be different (very unlikely to be all equal for random)
assert_ne!(matrix.data()[0], matrix.data()[1]);
}
#[test]
fn test_matrix_operations() {
let a = Matrix::new(vec![1.0, 2.0], vec![2]);
let b = Matrix::new(vec![3.0, 4.0], vec![2]);
let sum = a.add(&b);
assert_eq!(sum.data(), &[4.0, 6.0]);
let scaled = a.mul(2.0);
assert_eq!(scaled.data(), &[2.0, 4.0]);
let mean_val = a.mean();
assert_eq!(mean_val, 1.5);
}
#[test]
fn test_replay_buffer() {
let mut buffer = ReplayBuffer::new(3);
assert!(buffer.is_empty());
assert_eq!(buffer.len(), 0);
let exp1 = Experience {
state: vec![1.0, 2.0],
action: vec![0.5],
reward: 1.0,
next_state: vec![1.5, 2.5],
done: false,
};
buffer.push(exp1.clone());
assert_eq!(buffer.len(), 1);
assert!(!buffer.is_empty());
// Fill buffer beyond capacity
buffer.push(exp1.clone());
buffer.push(exp1.clone());
buffer.push(exp1.clone()); // Should evict first element
assert_eq!(buffer.len(), 3);
// Test sampling
let samples = buffer.sample(2);
assert!(samples.is_some());
assert_eq!(samples.unwrap().len(), 2);
// Test insufficient samples
let samples = buffer.sample(5);
assert!(samples.is_none());
}
#[test]
fn test_policy_network() {
let network = PolicyNetwork::new(2, 3);
let input = vec![1.0, -1.0];
let output = network.forward(&input);
assert_eq!(output.len(), 3);
// Output should be in [-1, 1] range due to tanh
for &val in &output {
assert!(val >= -1.0 && val <= 1.0);
}
}
#[test]
fn test_ppo_creation() {
let ppo = PPO::new(4, 2);
let state = vec![0.5, -0.5, 1.0, 0.0];
let action = ppo.get_action(&state);
assert_eq!(action.len(), 2);
let value = ppo.get_value(&state);
assert!(value.is_finite());
}
#[test]
fn test_ppo_advantages() {
let ppo = PPO::new(2, 1);
let experiences = vec![
Experience {
state: vec![0.0, 0.0],
action: vec![0.1],
reward: 1.0,
next_state: vec![0.1, 0.1],
done: false,
},
Experience {
state: vec![0.1, 0.1],
action: vec![-0.1],
reward: -0.5,
next_state: vec![0.0, 0.0],
done: true,
},
];
let advantages = ppo.compute_advantages(&experiences);
assert_eq!(advantages.len(), 2);
for advantage in &advantages {
assert!(advantage.is_finite());
}
}
#[test]
fn test_ppo_update() {
let mut ppo = PPO::new(2, 1);
let experiences = vec![
Experience {
state: vec![1.0, 0.0],
action: vec![0.5],
reward: 1.0,
next_state: vec![1.5, 0.0],
done: false,
},
Experience {
state: vec![1.5, 0.0],
action: vec![-0.3],
reward: 0.5,
next_state: vec![1.2, 0.0],
done: false,
},
];
let loss = ppo.update(&experiences);
assert!(loss.is_finite());
}
#[test]
fn test_experience_serialization() {
let exp = Experience {
state: vec![1.0, 2.0],
action: vec![0.5],
reward: 1.5,
next_state: vec![1.5, 2.5],
done: false,
};
let serialized = serde_json::to_string(&exp).unwrap();
let deserialized: Experience = serde_json::from_str(&serialized).unwrap();
assert_eq!(exp, deserialized);
}
#[test]
fn test_integration_episode() {
let mut ppo = PPO::new(2, 1);
let mut buffer = ReplayBuffer::new(100);
// Simulate a simple episode
let mut state = vec![0.0, 0.0];
let mut total_reward = 0.0;
for step in 0..10 {
let action = ppo.get_action(&state);
// Simple environment: move towards zero, get reward for being close
let next_state = vec![
state[0] + action[0] * 0.1,
state[1] - state[0] * 0.1, // Simple dynamics
];
let distance = (next_state[0].powi(2) + next_state[1].powi(2)).sqrt();
let reward = 1.0 - distance; // Reward for being close to origin
let done = step >= 9;
let experience = Experience {
state: state.clone(),
action: action.clone(),
reward,
next_state: next_state.clone(),
done,
};
buffer.push(experience);
total_reward += reward;
state = next_state;
if done {
break;
}
}
assert!(total_reward.is_finite());
assert_eq!(buffer.len(), 10);
// Train on collected experience
if let Some(batch) = buffer.sample(5) {
let loss = ppo.update(&batch);
assert!(loss.is_finite());
}
}
}