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

409 lines
13 KiB
Rust

use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main};
use rtx_rl::{
ActorLearner, ActorLearnerConfig, Environment, Experience, PrioritizedReplayBuffer,
ReplayBuffer, Step,
algorithms::{PPO, PPOConfig, SAC, SACConfig},
};
use rtx_tensor::{DType, Device, Tensor};
use std::collections::HashMap;
#[derive(Clone, Debug)]
struct BenchExperience {
state: Vec<f32>,
action: Vec<f32>,
reward: f32,
next_state: Vec<f32>,
done: bool,
}
impl Experience for BenchExperience {
type State = Vec<f32>;
type Action = Vec<f32>;
type Reward = f32;
fn state(&self) -> &Self::State {
&self.state
}
fn action(&self) -> &Self::Action {
&self.action
}
fn reward(&self) -> Self::Reward {
self.reward
}
fn next_state(&self) -> &Self::State {
&self.next_state
}
fn done(&self) -> bool {
self.done
}
}
#[derive(Clone)]
struct BenchEnv {
state_dim: usize,
action_dim: usize,
step_count: usize,
max_steps: usize,
}
impl Environment for BenchEnv {
type State = Vec<f32>;
type Action = Vec<f32>;
type Reward = f32;
fn reset(&mut self) -> Self::State {
self.step_count = 0;
vec![0.0; self.state_dim]
}
fn step(&mut self, action: &Self::Action) -> Step<Self::State, Self::Reward> {
self.step_count += 1;
Step {
state: vec![0.5; self.state_dim],
reward: 1.0,
done: self.step_count >= self.max_steps,
info: HashMap::new(),
}
}
fn action_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-1.0; self.action_dim], vec![1.0; self.action_dim])
}
fn observation_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-10.0; self.state_dim], vec![10.0; self.state_dim])
}
}
fn bench_replay_buffer_operations(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("replay_buffer");
for buffer_size in [1000, 10000, 100000].iter() {
group.throughput(Throughput::Elements(*buffer_size as u64));
group.bench_with_input(
BenchmarkId::new("add_experience", buffer_size),
buffer_size,
|b, &size| {
b.iter(|| {
rt.block_on(async {
let mut buffer = ReplayBuffer::<BenchExperience>::new(size, device.clone());
for i in 0..size {
let exp = BenchExperience {
state: vec![i as f32; 8],
action: vec![0.5; 4],
reward: 1.0,
next_state: vec![i as f32 + 1.0; 8],
done: false,
};
buffer.add(exp);
}
})
});
},
);
group.bench_with_input(
BenchmarkId::new("sample_batch", buffer_size),
buffer_size,
|b, &size| {
let mut buffer = rt.block_on(async {
let mut buf = ReplayBuffer::<BenchExperience>::new(size, device.clone());
for i in 0..size {
let exp = BenchExperience {
state: vec![i as f32; 8],
action: vec![0.5; 4],
reward: 1.0,
next_state: vec![i as f32 + 1.0; 8],
done: false,
};
buf.add(exp);
}
buf
});
b.iter(|| {
rt.block_on(async {
let _batch = buffer.sample(256).await.unwrap();
})
});
},
);
}
group.finish();
}
fn bench_prioritized_replay_buffer(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("prioritized_replay_buffer");
for buffer_size in [10000, 100000].iter() {
group.throughput(Throughput::Elements(*buffer_size as u64));
group.bench_with_input(
BenchmarkId::new("prioritized_sample", buffer_size),
buffer_size,
|b, &size| {
let mut buffer = rt.block_on(async {
let mut buf = PrioritizedReplayBuffer::<BenchExperience>::new(
size,
0.6,
0.4,
device.clone(),
);
for i in 0..size {
let exp = BenchExperience {
state: vec![i as f32; 8],
action: vec![0.5; 4],
reward: 1.0,
next_state: vec![i as f32 + 1.0; 8],
done: false,
};
buf.add_with_priority(exp, 0.5);
}
buf
});
b.iter(|| {
rt.block_on(async {
let (_batch, _weights, indices) =
buffer.sample_with_weights(256).await.unwrap();
let new_priorities = vec![0.8; indices.len()];
buffer.update_priorities(&indices, &new_priorities);
})
});
},
);
}
group.finish();
}
fn bench_ppo_algorithm(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("ppo");
for batch_size in [64, 256, 1024].iter() {
group.throughput(Throughput::Elements(*batch_size as u64));
group.bench_with_input(
BenchmarkId::new("forward_pass", batch_size),
batch_size,
|b, &size| {
let config = PPOConfig::default();
let ppo = PPO::new(config, 8, 4, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let _output = ppo.forward(&states).await.unwrap();
})
});
},
);
group.bench_with_input(
BenchmarkId::new("update_step", batch_size),
batch_size,
|b, &size| {
let config = PPOConfig {
n_epochs: 1,
batch_size: size,
minibatch_size: size.min(64),
..PPOConfig::default()
};
let mut ppo = PPO::new(config, 8, 4, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let actions = Tensor::randn(&[size, 4], DType::F32, device.clone());
let old_log_probs = Tensor::randn(&[size], DType::F32, device.clone());
let rewards = Tensor::ones(&[size], DType::F32, device.clone());
let values = Tensor::zeros(&[size], DType::F32, device.clone());
let dones = Tensor::zeros(&[size], DType::Bool, device.clone());
let _metrics = ppo
.update(&states, &actions, &old_log_probs, &rewards, &values, &dones)
.await
.unwrap();
})
});
},
);
}
group.finish();
}
fn bench_sac_algorithm(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("sac");
for batch_size in [64, 256, 1024].iter() {
group.throughput(Throughput::Elements(*batch_size as u64));
group.bench_with_input(
BenchmarkId::new("actor_forward", batch_size),
batch_size,
|b, &size| {
let config = SACConfig::default();
let sac = SAC::new(config, 8, 4, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let _output = sac.actor_forward(&states, false).await.unwrap();
})
});
},
);
group.bench_with_input(
BenchmarkId::new("critic_forward", batch_size),
batch_size,
|b, &size| {
let config = SACConfig::default();
let sac = SAC::new(config, 8, 4, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let actions = Tensor::randn(&[size, 4], DType::F32, device.clone());
let _q_vals = sac.critic_forward(&states, &actions, 0).await.unwrap();
})
});
},
);
group.bench_with_input(
BenchmarkId::new("full_update", batch_size),
batch_size,
|b, &size| {
let config = SACConfig {
batch_size: size,
..SACConfig::default()
};
let mut sac = SAC::new(config, 8, 4, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let actions = Tensor::randn(&[size, 4], DType::F32, device.clone());
let rewards = Tensor::randn(&[size], DType::F32, device.clone());
let next_states = Tensor::randn(&[size, 8], DType::F32, device.clone());
let dones = Tensor::zeros(&[size], DType::Bool, device.clone());
let _metrics = sac
.update(&states, &actions, &rewards, &next_states, &dones)
.await
.unwrap();
})
});
},
);
}
group.finish();
}
fn bench_actor_learner(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("actor_learner");
for num_actors in [1, 4, 8].iter() {
group.throughput(Throughput::Elements(*num_actors as u64));
group.bench_with_input(
BenchmarkId::new("rollout_collection", num_actors),
num_actors,
|b, &actors| {
let config = ActorLearnerConfig {
num_actors: actors,
num_learners: 1,
rollout_length: 128,
batch_size: 256,
device: device.clone(),
..Default::default()
};
let mut actor_learner = ActorLearner::new(config, 8, 4, 256);
b.iter(|| {
rt.block_on(async {
let mut envs: Vec<BenchEnv> = (0..actors)
.map(|_| BenchEnv {
state_dim: 8,
action_dim: 4,
step_count: 0,
max_steps: 200,
})
.collect();
let _metrics = actor_learner.training_step(&mut envs).await.unwrap();
})
});
},
);
}
group.finish();
}
fn bench_throughput_scaling(c: &mut Criterion) {
let rt = tokio::runtime::Runtime::new().unwrap();
let device = Device::cuda(0);
let mut group = c.benchmark_group("throughput_scaling");
// Test how throughput scales with different parameters
for (state_dim, action_dim) in [(4, 2), (16, 8), (64, 32), (256, 128)].iter() {
group.throughput(Throughput::Elements(1000));
group.bench_with_input(
BenchmarkId::new(
"inference_throughput",
format!("{}x{}", state_dim, action_dim),
),
&(*state_dim, *action_dim),
|b, &(s_dim, a_dim)| {
let config = PPOConfig::default();
let ppo = PPO::new(config, s_dim, a_dim, 256, device.clone());
b.iter(|| {
rt.block_on(async {
let states = Tensor::randn(&[1000, s_dim], DType::F32, device.clone());
let _output = ppo.forward(&states).await.unwrap();
})
});
},
);
}
group.finish();
}
criterion_group!(
benches,
bench_replay_buffer_operations,
bench_prioritized_replay_buffer,
bench_ppo_algorithm,
bench_sac_algorithm,
bench_actor_learner,
bench_throughput_scaling
);
criterion_main!(benches);