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)
}
}
#[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);
}
}
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,
}
}
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算法测试...");
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);
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 = 100;
let (pre_train_reward, pre_train_steps, pre_train_caught_rabbit) =
test_model(&mut q_learning, &mut policy, &mut env, max_steps, "训练前");
let train_args = TrainArgs {
epochs: 1000,
max_steps,
update_interval: 100,
update_freq: 10,
batch_size: 64,
learning_rate: 0.1,
gamma: 0.99,
};
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_caught_rabbit) =
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_caught_rabbit, post_train_caught_rabbit);
println!("\n测试结束");
}