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, QLearning, TrainArgs};
use rlkit::policies::{EpsilonGreedy};
use rlkit::types::{EnvTrait, Status, Reward, Action};
use rand::Rng;
use std::{collections::HashSet, fmt::Display};

#[derive(Debug, Clone)]
struct Node(HashSet<usize>);

impl Node {
    fn if_possible(&self, s2: usize) -> bool {
        self.0.contains(&s2)
    }
}

impl Display for Node {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(f, "{:?}", self.0)
    }
}

// 1. 实现一个抓兔子环境
#[derive(Clone)]
struct RabbitCatchEnv {
    node_size: usize,
    nodes: Vec<Node>,
    rabbit_pos: usize,
    agent_count: usize,
    agent_positions: Vec<usize>,
    state_space: Vec<usize>,
    action_space: Vec<usize>,
    max_steps_per_episode: usize,
    steps_taken: usize,
}

impl RabbitCatchEnv {
    fn new(node_size: usize, edge_size: usize, agent_count: usize) -> Self {
        if node_size <= agent_count {
            panic!("节点数量必须大于智能体数量")
        }
        if edge_size < node_size {
            panic!("边的数量不能少于节点数量")
        }

        let mut rng = rand::rng();
        let mut nodes: Vec<Node> = Vec::new();

        nodes.push(Node(HashSet::new()));
        for i in 1..node_size {
            let max_possible_edges = i.min(edge_size);
            let new_edge_size = rng.random_range(1..=max_possible_edges);
            let mut new_edges = HashSet::new();

            let mut j = 0;
            while new_edges.len() < new_edge_size && j < 100 {
                j += 1;
                let new_edge = rng.random_range(0..i);
                new_edges.insert(new_edge);
            }
            for low_node in &new_edges {
                nodes[*low_node].0.insert(i);
            }
            nodes.push(Node(new_edges));
        }

        let rabbit_pos = rng.random_range(0..node_size);
        let mut agent_positions = vec![rabbit_pos; agent_count];
        for pos in agent_positions.iter_mut() {
            while *pos == rabbit_pos {
                *pos = rng.random_range(0..node_size);
            }
        }

        // 状态空间是智能体数量 + 1(兔子位置)
        let state_space = vec![node_size; agent_count + 1];
        // 动作空间存储为字段
        let action_space = vec![node_size];

        RabbitCatchEnv {
            node_size,
            nodes,
            rabbit_pos,
            agent_count,
            agent_positions,
            state_space,
            action_space,
            max_steps_per_episode: 100,
            steps_taken: 0,
        }
    }
    
    // 将环境状态转换为Status<usize>格式
    fn to_status(&self) -> Status<usize> {
        let mut state_values = self.agent_positions.clone();
        state_values.push(self.rabbit_pos);
        Status::new(state_values, self.state_space.clone())
    }
    
    // 将动作索引解析为目标节点
    fn parse_action(&self, action_idx: usize) -> Vec<usize> {
        // 为每个智能体分配一个目标节点
        let mut targets = Vec::with_capacity(self.agent_count);
        for i in 0..self.agent_count {
            // 将动作索引映射到节点索引
            let node_idx = (action_idx + i) % self.node_size;
            targets.push(node_idx);
        }
        targets
    }
}

impl Display for RabbitCatchEnv {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.write_str("环境图结构:\n")?;
        for (i, node) in self.nodes.iter().enumerate() {
            f.write_str(&format!("节点 {i} 连接到: {:?}\n", node))?;
        }
        f.write_str(&format!("\n智能体位置: {:?}\n", self.agent_positions))?;
        f.write_str(&format!("兔子位置: {}\n", self.rabbit_pos))?;
        Ok(())
    }
}

impl EnvTrait<usize, usize> for RabbitCatchEnv {
    fn step(&mut self, _state: &Status<usize>, action: &Action<usize>) -> (Status<usize>, Reward, bool) {
        let mut reward = 0.0;
        let mut caught_rabbit = false;
        
        // 获取动作索引
        let action_idx = action.as_slice()[0];
        
        // 解析动作
        let target_positions = self.parse_action(action_idx);

        // 执行动作
        for (i, &target) in target_positions.iter().enumerate() {
            let current_pos = self.agent_positions[i];
            
            // 判断动作是否可行(检查当前节点是否连接到目标节点)
            if let Some(current_node) = self.nodes.get(current_pos) {
                if current_node.if_possible(target) {
                    // 移动到新位置
                    self.agent_positions[i] = target;
                    
                    // 检查是否抓住兔子
                    if target == self.rabbit_pos {
                        reward += 100.0;
                        caught_rabbit = true;
                        
                        // 兔子被抓住,重置位置
                        let mut rng = rand::rng();
                        let mut new_rabbit_pos;
                        loop {
                            new_rabbit_pos = rng.random_range(0..self.node_size);
                            if !self.agent_positions.contains(&new_rabbit_pos) {
                                break;
                            }
                        }
                        self.rabbit_pos = new_rabbit_pos;
                    }
                } else {
                    // 无效移动惩罚
                    reward -= 10.0;
                }
            }
        }

        // 兔子随机移动
        if !caught_rabbit {
            let mut rng = rand::rng();
            if let Some(rabbit_node) = self.nodes.get(self.rabbit_pos) {
                if !rabbit_node.0.is_empty() {
                    let possible_moves: Vec<_> = rabbit_node.0.iter().collect();
                    let move_idx = rng.random_range(0..possible_moves.len());
                    let new_rabbit_pos = *possible_moves[move_idx];
                    
                    // 确保新位置不与任何智能体重叠
                    if !self.agent_positions.contains(&new_rabbit_pos) {
                        self.rabbit_pos = new_rabbit_pos;
                    }
                }
            }
        }

        // 每步基础惩罚,鼓励快速抓住兔子
        reward -= 1.0;
        
        // 增加步数计数
        self.steps_taken += 1;
        
        // 达到最大步数则结束回合
        let done = (self.steps_taken >= self.max_steps_per_episode) || caught_rabbit;

        (self.to_status(), Reward(reward), done)
    }

    fn reset(&mut self) -> Status<usize> {
        let mut rng = rand::rng();

        // 重置步数
        self.steps_taken = 0;
        
        // 重置兔子位置
        let rabbit_pos = rng.random_range(0..self.node_size);
        
        // 重置智能体位置,确保不与兔子重叠
        let mut agent_positions = vec![rabbit_pos; self.agent_count];
        for pos in agent_positions.iter_mut() {
            while *pos == rabbit_pos {
                *pos = rng.random_range(0..self.node_size);
            }
        }

        self.agent_positions = agent_positions;
        self.rabbit_pos = rabbit_pos;

        self.to_status()
    }
    
    fn action_space(&self) -> &[usize] {
        &self.action_space
    }

    fn state_space(&self) -> &[usize] {
        &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: std::fmt::Debug + Copy + 'static,
{
    println!("\n{}测试开始...", description);
    let mut state = env.reset();
    let mut total_reward: f32 = 0.0;
    let mut steps = 0;
    let mut caught_rabbit = false;
    
    loop {
        steps += 1;
        let action = q_learning.get_action(&state, policy).expect("获取动作失败");
        let (next_state, reward, done) = env.step(&state, &action);
        
        total_reward += reward.0;
        state = next_state;
        
        // 打印动作
        let action_str = format!("{:?}", action.as_slice());
        
        println!("步骤 {}: 动作 = {}, 奖励 = {}, 完成 = {}", 
                 steps, action_str, reward.0, done);
        
        // 检查是否抓住了兔子(通过奖励判断)
        if reward.0 > 50.0 {
            caught_rabbit = true;
        }
        
        if done || steps >= max_steps {
            break;
        }
    }
    
    println!("{}测试完成! 总奖励: {}, 步数: {}, 是否抓住兔子: {}", 
             description, total_reward, steps, caught_rabbit);
    
    (total_reward, steps, caught_rabbit)
}

fn main() {
    println!("开始兔子抓捕环境的Q-Learning算法测试...");
    
    // 2. 创建环境
    let node_size = 5;
    let edge_size = 10;
    let agent_count = 2;
    let mut env = RabbitCatchEnv::new(node_size, edge_size, agent_count);
    println!("环境初始化完成:");
    println!("{}", env);
    
    // 3. 创建Q-Learning算法
    println!("创建Q-Learning算法实例...");
    let mut q_learning = QLearning::new(&env, 10000).unwrap();
    let mut policy = EpsilonGreedy::new(1.0, 0.01, 0.995);

    // 5. 训练前测试未训练模型
    let max_steps = 100;
    let (pre_train_reward, pre_train_steps, pre_train_caught_rabbit) = 
        test_model(&mut q_learning, &mut policy, &mut env, max_steps, "训练前");
    
    // 6. 配置训练参数
    let train_args = TrainArgs {
        epochs: 1000,
        max_steps,
        update_interval: 100,
        update_freq: 10,
        batch_size: 64,
        learning_rate: 0.1,
        gamma: 0.99,
    };
    
    // 7. 训练模型
    println!("\n开始训练模型...");
    if let Err(e) = q_learning.train(&mut env, &mut policy, train_args) {
        println!("训练失败: {}", e);
        return;
    }

    // 8. 重置策略的epsilon值以便测试(可选)
    policy.epsilon = 0.01; // 设置很低的探索率用于测试
    println!("\n测试时设置epsilon值为0.01以减少随机探索");
    
    // 9. 测试训练后的模型
    let (post_train_reward, post_train_steps, post_train_caught_rabbit) = 
        test_model(&mut q_learning, &mut policy, &mut env, max_steps, "训练后");
    
    // 10. 比较训练前后的效果
    println!("\n=== 训练前后效果对比 ===");
    println!("训练前总奖励: {:.2}, 训练后总奖励: {:.2}, 提升: {:.2}%", 
             pre_train_reward, post_train_reward, 
             if pre_train_reward != 0.0 {
                 (post_train_reward - pre_train_reward) / pre_train_reward.abs() * 100.0
             } else {
                 0.0
             });
    println!("训练前步数: {}, 训练后步数: {}, 减少: {} 步 ({:.2}%)", 
             pre_train_steps, post_train_steps, 
             if pre_train_steps > post_train_steps {
                 pre_train_steps - post_train_steps
             } else {
                 0
             },
             if pre_train_steps != 0 {
                 ((pre_train_steps as f32 - post_train_steps as f32) / pre_train_steps as f32).abs() * 100.0
             } else {
                 0.0
             });
    println!("训练前是否抓住兔子: {}, 训练后是否抓住兔子: {}", 
             pre_train_caught_rabbit, post_train_caught_rabbit);
    
    println!("\n测试结束");
}