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 super::env::*;
use rlkit::algs::{Algorithm, QLearning, TrainArgs};
use rlkit::policies::EpsilonGreedy;

pub fn train_with_q_learning() {
    println!("开始Q-Learning算法测试...");

    // 2. 创建环境
    let size = 50;
    // 设置障碍物密度(0.0-1.0之间,推荐0.2-0.3)
    let obstacle_density = 0.1;
    // 随机生成障碍物
    let obstacles = GridWorld::generate_random_obstacles(size, obstacle_density);
    println!("生成了 {} 个随机障碍物", obstacles.len());
    // 创建包含随机障碍物的环境
    let mut env = GridWorld::new(size, obstacles);

    // 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 = (size * size) as usize;
    let (pre_train_reward, pre_train_steps, pre_train_reached_goal) =
        test_model(&mut q_learning, &mut policy, &mut env, max_steps, "训练前");

    // 6. 配置训练参数
    let train_args = TrainArgs {
        epochs: 1000,
        max_steps,
        batch_size: 64,
        learning_rate: 0.5,
        gamma: 0.99,
        ..Default::default()
    };

    // 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_reached_goal) =
        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_reached_goal, post_train_reached_goal
    );

    println!("\n测试结束");
}