use rlkit::algs::Algorithm;
use rlkit::types::{EnvTrait, Status, Reward, Action};
use rand::{rng, seq::SliceRandom};
#[derive(Clone)]
pub struct GridWorld {
size: u16,
state_space: Vec<u16>,
position: (u16, u16),
goal: (u16, u16),
obstacles: Vec<(u16, u16)>, }
impl GridWorld {
pub fn new(size: usize, obstacles: Vec<(u16, u16)>) -> Self {
let start = (0, 0);
let goal = (size as u16 - 1, size as u16 - 1);
let filtered_obstacles: Vec<_> = obstacles
.into_iter()
.filter(|&pos| pos != start && pos != goal)
.collect();
Self {
size: size as u16,
state_space: vec![size as u16, size as u16],
position: start,
goal,
obstacles: filtered_obstacles,
}
}
fn is_obstacle(&self, x: u16, y: u16) -> bool {
self.obstacles.contains(&(x, y))
}
fn distance_reward(&self) -> f32 {
let (x, y) = self.position;
let (goal_x, goal_y) = self.goal;
let distance = ((x as i32 - goal_x as i32).abs() + (y as i32 - goal_y as i32).abs()) as f32;
(self.size as f32 / 10.0 - distance) / (self.size as f32 * 2.0)
}
pub fn generate_random_obstacles(size: usize, obstacle_density: f32) -> Vec<(u16, u16)> {
let mut rng = rng();
let mut obstacles = Vec::new();
let start = (0, 0);
let goal = (size as u16 - 1, size as u16 - 1);
let total_cells = (size as u16) * (size as u16);
let obstacle_count = (total_cells as f32 * obstacle_density) as usize;
let mut positions = Vec::new();
for x in 0..size as u16 {
for y in 0..size as u16 {
let pos = (x, y);
if pos != start && pos != goal {
positions.push(pos);
}
}
}
positions.shuffle(&mut rng);
for pos in positions.iter().take(obstacle_count) {
obstacles.push(*pos);
}
obstacles
}
fn pos_to_state(&self) -> Status<u16> {
Status::new(vec![self.position.0, self.position.1], vec![self.size, self.size])
}
}
impl EnvTrait<u16, u16> for GridWorld {
fn step(&mut self, _state: &Status<u16>, action: &Action<u16>) -> (Status<u16>, Reward, bool) {
let (x, y) = self.position;
let (mut new_x, mut new_y) = (x, y);
let action_idx = action.as_slice()[0];
let mut attempted_out_of_bounds = false;
let mut collided_with_obstacle = false;
match action_idx {
0 => if x > 0 { new_x = x - 1; } else { attempted_out_of_bounds = true; },
1 => if y < self.size - 1 { new_y = y + 1; } else { attempted_out_of_bounds = true; },
2 => if x < self.size - 1 { new_x = x + 1; } else { attempted_out_of_bounds = true; },
3 => if y > 0 { new_y = y - 1; } else { attempted_out_of_bounds = true; },
_ => {},
}
if !attempted_out_of_bounds && self.is_obstacle(new_x, new_y) {
collided_with_obstacle = true;
(new_x, new_y) = (x, y);
}
self.position = (new_x, new_y);
let new_state = self.pos_to_state();
let done = self.position == self.goal;
let mut reward_value = if done {
self.size as f32 * 10.0
} else {
-1.0
};
if attempted_out_of_bounds {
reward_value -= self.size as f32;
}
if collided_with_obstacle {
reward_value -= self.size as f32 / 2.0;
}
reward_value += self.distance_reward();
let reward = Reward(reward_value);
(new_state, reward, done)
}
fn reset(&mut self) -> Status<u16> {
self.position = (0, 0);
self.pos_to_state()
}
fn action_space(&self) -> &[u16] {
&[4]
}
fn state_space(&self) -> &[u16] {
&self.state_space
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
}
pub fn test_model<Env, T>(
q_learning: &mut impl Algorithm<Env, T, T>,
policy: &mut impl rlkit::policies::Policy<T>,
env: &mut Env,
max_steps: usize,
description: &str
) -> (f32, usize, bool)
where
Env: EnvTrait<T, T>,
T: Copy + Clone + std::iter::Product<T> + TryInto<u16> + TryFrom<u16> + 'static,
{
let env = env.as_any_mut().downcast_mut::<GridWorld>().unwrap();
println!("\n{}测试开始...", description);
let mut state = env.reset();
let mut total_reward: f32 = 0.0;
let mut steps = 0;
let mut reached_goal = false;
let mut path = Vec::new();
path.push(env.position);
loop {
steps += 1;
let action = q_learning.get_action(&Status::from_status(state.clone()).unwrap(), policy).unwrap();
let (next_state, reward, done) = env.step(&Status::from_status(state).unwrap(), &Action::from_actions(action).unwrap());
total_reward += reward.0;
state = Status::from_status(next_state).unwrap();
path.push(env.position);
if done {
reached_goal = true;
break;
}
if steps >= max_steps {
break;
}
}
println!("\n{}路径可视化:", description);
println!("总奖励: {}, 步数: {}, 是否到达目标: {}",
total_reward, steps, reached_goal);
let size = env.size as usize;
let mut grid: Vec<Vec<char>> = vec![vec!['░'; size]; size];
for &(x, y) in &env.obstacles {
grid[x as usize][y as usize] = '▓'; }
for &(x, y) in &path {
grid[x as usize][y as usize] = '●'; }
if let Some(&(start_x, start_y)) = path.first() {
grid[start_x as usize][start_y as usize] = 'S'; }
if let Some(&(end_x, end_y)) = path.last() {
grid[end_x as usize][end_y as usize] = if reached_goal { 'G' } else { 'E' }; }
println!("{}", "─".repeat(size * 2 + 2));
for row in &grid {
print!("│");
for &cell in row {
print!("{}", cell);
print!(" "); }
println!("│");
}
println!("{}", "─".repeat(size * 2 + 2));
(total_reward, steps, reached_goal)
}