107 lines
2.5 KiB
Rust
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]);
|
|
}
|