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

107 lines
2.5 KiB
Rust

use rtx_rl::{Environment, Step};
#[derive(Clone, Debug)]
struct TestEnv {
state: Vec<f32>,
step_count: usize,
max_steps: usize,
}
impl Environment for TestEnv {
type State = Vec<f32>;
type Action = Vec<f32>;
type Reward = f32;
fn reset(&mut self) -> Self::State {
self.state = vec![0.0, 0.0];
self.step_count = 0;
self.state.clone()
}
fn step(&mut self, action: &Self::Action) -> Step<Self::State, Self::Reward> {
self.state[0] += action[0];
self.state[1] += action[1];
self.step_count += 1;
let reward = -(self.state[0].powi(2) + self.state[1].powi(2));
let done = self.step_count >= self.max_steps;
Step {
state: self.state.clone(),
reward,
done,
info: std::collections::HashMap::new(),
}
}
fn action_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-1.0, -1.0], vec![1.0, 1.0])
}
fn observation_space(&self) -> (Vec<f32>, Vec<f32>) {
(vec![-10.0, -10.0], vec![10.0, 10.0])
}
}
#[tokio::test]
async fn test_environment_reset() {
let mut env = TestEnv {
state: vec![5.0, 5.0],
step_count: 10,
max_steps: 100,
};
let state = env.reset();
assert_eq!(state, vec![0.0, 0.0]);
assert_eq!(env.step_count, 0);
}
#[tokio::test]
async fn test_environment_step() {
let mut env = TestEnv {
state: vec![0.0, 0.0],
step_count: 0,
max_steps: 100,
};
let action = vec![0.5, -0.3];
let step_result = env.step(&action);
assert_eq!(step_result.state, vec![0.5, -0.3]);
assert!(step_result.reward < 0.0); // Negative quadratic reward
assert!(!step_result.done);
assert_eq!(env.step_count, 1);
}
#[tokio::test]
async fn test_environment_termination() {
let mut env = TestEnv {
state: vec![0.0, 0.0],
step_count: 99,
max_steps: 100,
};
let action = vec![0.1, 0.1];
let step_result = env.step(&action);
assert!(step_result.done);
assert_eq!(env.step_count, 100);
}
#[tokio::test]
async fn test_environment_spaces() {
let env = TestEnv {
state: vec![0.0, 0.0],
step_count: 0,
max_steps: 100,
};
let (action_low, action_high) = env.action_space();
assert_eq!(action_low, vec![-1.0, -1.0]);
assert_eq!(action_high, vec![1.0, 1.0]);
let (obs_low, obs_high) = env.observation_space();
assert_eq!(obs_low, vec![-10.0, -10.0]);
assert_eq!(obs_high, vec![10.0, 10.0]);
}