rlkit 0.0.3

A deep reinforcement learning library based on Rust and Candle, providing complete implementations of Q-Learning and DQN algorithms, supporting custom environments, various policy choices, and flexible training configurations. Future support will include more reinforcement learning algorithms, such as DDPG, PPO, A2C, etc.
Documentation
use rlkit::algs::Algorithm;
use rlkit::types::{EnvTrait, Status, Reward, Action};
use rand::{rng, seq::SliceRandom};

// 1. 实现一个简单的网格世界环境
#[derive(Clone)]
pub struct GridWorld {
    size: u16,
    state_space: Vec<u16>,
    position: (u16, u16),
    goal: (u16, u16),
    obstacles: Vec<(u16, u16)>, // 障碍物位置集合
}

impl GridWorld {
    // 创建包含障碍物的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) {
        // 动作:0=上, 1=右, 2=下, 3=左
        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)
}