use super::env::*;
use rlkit::algs::{Algorithm, QLearning, TrainArgs};
use rlkit::policies::EpsilonGreedy;
pub fn train_with_q_learning() {
println!("开始Q-Learning算法测试...");
let size = 50;
let obstacle_density = 0.1;
let obstacles = GridWorld::generate_random_obstacles(size, obstacle_density);
println!("生成了 {} 个随机障碍物", obstacles.len());
let mut env = GridWorld::new(size, obstacles);
println!("创建Q-Learning算法实例...");
let mut q_learning = QLearning::new(&env, 10000).unwrap();
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 q_learning, &mut policy, &mut env, max_steps, "训练前");
let train_args = TrainArgs {
epochs: 1000,
max_steps,
batch_size: 64,
learning_rate: 0.5,
gamma: 0.99,
..Default::default()
};
println!("\n开始训练模型...");
if let Err(e) = q_learning.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 q_learning, &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测试结束");
}