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算法测试...");
let size = 20;
let obstacle_density = 0.05;
let obstacles = GridWorld::generate_random_obstacles(size, obstacle_density);
println!("生成了 {} 个随机障碍物", obstacles.len());
let mut env = GridWorld::new(size, obstacles);
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);
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, "训练前");
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()
};
println!("\n开始训练模型...");
if let Err(e) = dqn.train(&mut env, &mut policy, train_args) {
println!("训练失败: {}", e);
return;
}
policy.epsilon = 0.01; println!("\n测试时设置epsilon值为0.01以减少随机探索");
let (post_train_reward, post_train_steps, post_train_reached_goal) =
test_model(&mut dqn, &mut policy, &mut env, max_steps, "训练后");
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测试结束");
}