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 candle_core::Device;
use rlkit::algs::dqn::DNQStateMode;
use rlkit::algs::{Algorithm, DQN, TrainArgs};
use rlkit::policies::EpsilonGreedy;

pub fn train_with_dqn() {
    println!("开始DQN算法测试...");

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

    // 3. 创建DQN算法
    println!("创建DQN算法实例...");
    let device = Device::new_cuda(0).unwrap();
    println!("使用设备: {:?}", device.location());

    let mut dqn = DQN::new(
        &env,
        10000,
        &[128, 64, 16],
        DNQStateMode::OneHot,
        &device
    ).expect("创建DQN算法失败");
    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 dqn, &mut policy, &mut env, max_steps, "训练前");

    // 6. 配置训练参数
    let train_args = TrainArgs {
        epochs: 1000,
        max_steps,
        batch_size: 32,
        learning_rate: 1e-3,
        gamma: 0.99,
        update_freq: 5,
        update_interval: 100,
        ..Default::default()
    };

    // 7. 训练模型
    println!("\n开始训练模型...");
    if let Err(e) = dqn.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 dqn, &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测试结束");
}