use crate::NeuralResult;
use scirs2_core::ndarray::{Array1, Array2, Array3, ScalarOperand};
use scirs2_core::random::thread_rng;
use sklears_core::types::FloatBounds;
use std::collections::VecDeque;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum CoordinationMode {
Cooperative,
Competitive,
Mixed,
Independent,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum CommunicationProtocol {
None,
Broadcast,
Targeted,
Learned,
}
#[derive(Debug, Clone)]
pub struct MultiAgentExperience<T: FloatBounds> {
pub joint_state: Array1<T>,
pub agent_states: Vec<Array1<T>>,
pub actions: Vec<usize>,
pub rewards: Vec<T>,
pub next_joint_state: Array1<T>,
pub next_agent_states: Vec<Array1<T>>,
pub done: bool,
pub messages: Option<Vec<Array1<T>>>,
}
#[derive(Debug)]
#[allow(dead_code)] pub struct MultiAgentReplayBuffer<T: FloatBounds> {
buffer: VecDeque<MultiAgentExperience<T>>,
capacity: usize,
n_agents: usize,
}
impl<T: FloatBounds> MultiAgentReplayBuffer<T> {
pub fn new(capacity: usize, n_agents: usize) -> Self {
Self {
buffer: VecDeque::with_capacity(capacity),
capacity,
n_agents,
}
}
pub fn push(&mut self, experience: MultiAgentExperience<T>) {
if self.buffer.len() >= self.capacity {
self.buffer.pop_front();
}
self.buffer.push_back(experience);
}
pub fn sample(&self, batch_size: usize) -> Option<Vec<MultiAgentExperience<T>>> {
if self.buffer.len() < batch_size {
return None;
}
let mut rng = thread_rng();
let mut samples = Vec::with_capacity(batch_size);
for _ in 0..batch_size {
let idx = rng.random_range(0..self.buffer.len());
samples.push(self.buffer[idx].clone());
}
Some(samples)
}
pub fn len(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
}
#[derive(Debug)]
#[allow(dead_code)] pub struct IndependentQLearner<T: FloatBounds> {
agent_networks: Vec<Array2<T>>,
learning_rate: T,
gamma: T,
epsilon: T,
n_agents: usize,
state_dim: usize,
n_actions: usize,
}
impl<T: FloatBounds + ScalarOperand + std::iter::Sum + std::ops::AddAssign> IndependentQLearner<T> {
pub fn new(
n_agents: usize,
state_dim: usize,
n_actions: usize,
learning_rate: T,
gamma: T,
epsilon: T,
) -> Self {
let mut rng = thread_rng();
let mut agent_networks = Vec::new();
for _ in 0..n_agents {
let std_dev = (2.0 / (state_dim + n_actions) as f64).sqrt();
let weights = Array2::from_shape_fn((n_actions, state_dim), |_| {
T::from(rng.gen_range(-std_dev..std_dev)).unwrap_or(T::zero())
});
agent_networks.push(weights);
}
Self {
agent_networks,
learning_rate,
gamma,
epsilon,
n_agents,
state_dim,
n_actions,
}
}
pub fn select_actions(&self, states: &[Array1<T>]) -> Vec<usize> {
let mut rng = thread_rng();
let mut actions = Vec::new();
for (agent_idx, state) in states.iter().enumerate() {
if rng.random::<f64>() < self.epsilon.to_f64().unwrap_or(0.1) {
actions.push(rng.gen_range(0..self.n_actions));
} else {
let q_values = self.agent_networks[agent_idx].dot(state);
let best_action = q_values
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx)
.unwrap_or(0);
actions.push(best_action);
}
}
actions
}
pub fn update(&mut self, experience: &MultiAgentExperience<T>) -> NeuralResult<()> {
for agent_idx in 0..self.n_agents {
let state = &experience.agent_states[agent_idx];
let action = experience.actions[agent_idx];
let reward = experience.rewards[agent_idx];
let next_state = &experience.next_agent_states[agent_idx];
let q_values = self.agent_networks[agent_idx].dot(state);
let q_value = q_values[action];
let next_q_values = self.agent_networks[agent_idx].dot(next_state);
let max_next_q = next_q_values
.iter()
.max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.cloned()
.unwrap_or(T::zero());
let target = if experience.done {
reward
} else {
reward + self.gamma * max_next_q
};
let td_error = target - q_value;
let gradient = state.mapv(|s| s * td_error * T::from(2.0).unwrap_or_else(|| T::zero()));
let mut row = self.agent_networks[agent_idx].row_mut(action);
for (w, &g) in row.iter_mut().zip(gradient.iter()) {
*w += g * self.learning_rate;
}
}
Ok(())
}
pub fn decay_epsilon(&mut self, decay_rate: T) {
self.epsilon *= decay_rate;
let min_epsilon = T::from(0.01).unwrap_or_else(|| T::zero());
if self.epsilon < min_epsilon {
self.epsilon = min_epsilon;
}
}
}
#[derive(Debug)]
pub struct VDN<T: FloatBounds> {
agent_networks: Vec<Array2<T>>,
learning_rate: T,
gamma: T,
epsilon: T,
n_agents: usize,
}
impl<T: FloatBounds + ScalarOperand + std::iter::Sum + std::ops::AddAssign> VDN<T> {
pub fn new(
n_agents: usize,
state_dim: usize,
n_actions: usize,
learning_rate: T,
gamma: T,
epsilon: T,
) -> Self {
let mut rng = thread_rng();
let mut agent_networks = Vec::new();
for _ in 0..n_agents {
let std_dev = (2.0 / (state_dim + n_actions) as f64).sqrt();
let weights = Array2::from_shape_fn((n_actions, state_dim), |_| {
T::from(rng.gen_range(-std_dev..std_dev)).unwrap_or(T::zero())
});
agent_networks.push(weights);
}
Self {
agent_networks,
learning_rate,
gamma,
epsilon,
n_agents,
}
}
pub fn compute_total_q(&self, states: &[Array1<T>], actions: &[usize]) -> T {
let mut total_q = T::zero();
for (agent_idx, (state, &action)) in states.iter().zip(actions.iter()).enumerate() {
let q_values = self.agent_networks[agent_idx].dot(state);
total_q += q_values[action];
}
total_q
}
pub fn select_joint_actions(&self, states: &[Array1<T>]) -> Vec<usize> {
let mut rng = thread_rng();
let mut actions = Vec::new();
for (agent_idx, state) in states.iter().enumerate() {
if rng.random::<f64>() < self.epsilon.to_f64().unwrap_or(0.1) {
actions.push(rng.gen_range(0..self.agent_networks[agent_idx].nrows()));
} else {
let q_values = self.agent_networks[agent_idx].dot(state);
let best_action = q_values
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx)
.unwrap_or(0);
actions.push(best_action);
}
}
actions
}
pub fn update_cooperative(&mut self, experience: &MultiAgentExperience<T>) -> NeuralResult<()> {
let shared_reward = experience.rewards.iter().cloned().sum::<T>()
/ T::from(self.n_agents as f64).unwrap_or_else(|| T::zero());
let current_q = self.compute_total_q(&experience.agent_states, &experience.actions);
let best_next_actions: Vec<usize> = experience
.next_agent_states
.iter()
.enumerate()
.map(|(agent_idx, state)| {
let q_values = self.agent_networks[agent_idx].dot(state);
q_values
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx)
.unwrap_or(0)
})
.collect();
let next_q = self.compute_total_q(&experience.next_agent_states, &best_next_actions);
let target = if experience.done {
shared_reward
} else {
shared_reward + self.gamma * next_q
};
let td_error = target - current_q;
for (agent_idx, state) in experience.agent_states.iter().enumerate() {
let action = experience.actions[agent_idx];
let gradient = state.mapv(|s| s * td_error * T::from(2.0).unwrap_or_else(|| T::zero()));
let mut row = self.agent_networks[agent_idx].row_mut(action);
for (w, &g) in row.iter_mut().zip(gradient.iter()) {
*w += g * self.learning_rate;
}
}
Ok(())
}
}
#[derive(Debug)]
pub struct QMIXMixingNetwork<T: FloatBounds> {
hyper_w1: Array2<T>,
hyper_b1: Array1<T>,
hyper_w_final: Array2<T>,
hyper_b_final: Array1<T>,
hidden_dim: usize,
}
impl<T: FloatBounds + ScalarOperand + std::iter::Sum + std::ops::AddAssign> QMIXMixingNetwork<T> {
pub fn new(n_agents: usize, state_dim: usize, hidden_dim: usize) -> Self {
let mut rng = thread_rng();
let std_w1 = (2.0 / (state_dim + hidden_dim) as f64).sqrt();
let hyper_w1 = Array2::from_shape_fn((hidden_dim, state_dim), |_| {
T::from(rng.gen_range(-std_w1..std_w1)).unwrap_or(T::zero())
});
let hyper_b1 = Array1::zeros(hidden_dim);
let std_w_final = (2.0 / (state_dim + 1) as f64).sqrt();
let hyper_w_final = Array2::from_shape_fn((n_agents, state_dim), |_| {
T::from(rng.gen_range(-std_w_final..std_w_final)).unwrap_or(T::zero())
});
let hyper_b_final = Array1::zeros(n_agents);
Self {
hyper_w1,
hyper_b1,
hyper_w_final,
hyper_b_final,
hidden_dim,
}
}
pub fn mix(&self, agent_q_values: &[T], global_state: &Array1<T>) -> T {
let w1 = (self.hyper_w1.dot(global_state) + &self.hyper_b1).mapv(|x| {
if x.to_f64().unwrap_or(0.0).abs() < x.to_f64().unwrap_or(0.0) {
-x
} else {
x
}
});
let agent_q_array = Array1::from(agent_q_values.to_vec());
let h = w1
.iter()
.zip(agent_q_array.iter().cycle())
.take(self.hidden_dim)
.map(|(&w, &q)| w * q)
.sum::<T>();
let w_final = (self.hyper_w_final.row(0).dot(global_state) + self.hyper_b_final[0]).abs();
h * w_final
}
}
#[derive(Debug)]
#[allow(dead_code)] pub struct MADDPG<T: FloatBounds> {
actors: Vec<Array2<T>>,
critic: Array2<T>,
target_actors: Vec<Array2<T>>,
target_critic: Array2<T>,
actor_lr: T,
critic_lr: T,
gamma: T,
tau: T,
n_agents: usize,
}
impl<T: FloatBounds + ScalarOperand + std::iter::Sum + std::ops::AddAssign> MADDPG<T> {
pub fn new(
n_agents: usize,
state_dim: usize,
action_dim: usize,
actor_lr: T,
critic_lr: T,
gamma: T,
tau: T,
) -> Self {
let mut rng = thread_rng();
let mut actors = Vec::new();
let mut target_actors = Vec::new();
for _ in 0..n_agents {
let std_dev = (2.0 / (state_dim + action_dim) as f64).sqrt();
let actor = Array2::from_shape_fn((action_dim, state_dim), |_| {
T::from(rng.gen_range(-std_dev..std_dev)).unwrap_or(T::zero())
});
target_actors.push(actor.clone());
actors.push(actor);
}
let total_input_dim = n_agents * (state_dim + action_dim);
let std_dev = (2.0 / total_input_dim as f64).sqrt();
let critic = Array2::from_shape_fn((1, total_input_dim), |_| {
T::from(rng.gen_range(-std_dev..std_dev)).unwrap_or(T::zero())
});
let target_critic = critic.clone();
Self {
actors,
critic,
target_actors,
target_critic,
actor_lr,
critic_lr,
gamma,
tau,
n_agents,
}
}
pub fn get_actions(&self, states: &[Array1<T>]) -> Vec<Array1<T>> {
states
.iter()
.enumerate()
.map(|(agent_idx, state)| self.actors[agent_idx].dot(state))
.collect()
}
pub fn soft_update_targets(&mut self) {
for (target, online) in self.target_actors.iter_mut().zip(&self.actors) {
*target = target.mapv(|t| t * (T::one() - self.tau)) + online.mapv(|o| o * self.tau);
}
self.target_critic = self.target_critic.mapv(|t| t * (T::one() - self.tau))
+ self.critic.mapv(|o| o * self.tau);
}
}
#[derive(Debug)]
pub struct CommunicationLayer<T: FloatBounds> {
protocol: CommunicationProtocol,
message_dim: usize,
attention_weights: Option<Array3<T>>,
}
impl<T: FloatBounds + ScalarOperand + std::iter::Sum> CommunicationLayer<T> {
pub fn new(protocol: CommunicationProtocol, message_dim: usize, n_agents: usize) -> Self {
let attention_weights = if protocol == CommunicationProtocol::Learned {
let mut rng = thread_rng();
let std_dev = (2.0 / (message_dim + message_dim) as f64).sqrt();
Some(Array3::from_shape_fn(
(n_agents, n_agents, message_dim),
|_| T::from(rng.gen_range(-std_dev..std_dev)).unwrap_or(T::zero()),
))
} else {
None
};
Self {
protocol,
message_dim,
attention_weights,
}
}
pub fn communicate(&self, agent_states: &[Array1<T>]) -> Vec<Array1<T>> {
match self.protocol {
CommunicationProtocol::None => {
agent_states.to_vec()
}
CommunicationProtocol::Broadcast => {
let mean_state = self.compute_mean_state(agent_states);
agent_states
.iter()
.map(|state| self.concatenate_state_message(state, &mean_state))
.collect()
}
CommunicationProtocol::Targeted => {
self.targeted_communication(agent_states)
}
CommunicationProtocol::Learned => {
self.learned_communication(agent_states)
}
}
}
fn compute_mean_state(&self, states: &[Array1<T>]) -> Array1<T> {
if states.is_empty() {
return Array1::zeros(self.message_dim);
}
let sum = states
.iter()
.fold(Array1::zeros(states[0].len()), |acc, state| acc + state);
sum / T::from(states.len() as f64).unwrap_or_else(|| T::zero())
}
fn concatenate_state_message(&self, state: &Array1<T>, message: &Array1<T>) -> Array1<T> {
let mut result = Vec::new();
result.extend_from_slice(state.as_slice().unwrap_or(&[]));
result.extend_from_slice(message.as_slice().unwrap_or(&[]));
Array1::from(result)
}
fn targeted_communication(&self, agent_states: &[Array1<T>]) -> Vec<Array1<T>> {
agent_states
.iter()
.enumerate()
.map(|(i, state)| {
let prev_idx = if i > 0 { i - 1 } else { agent_states.len() - 1 };
let next_idx = (i + 1) % agent_states.len();
let neighbor_message = (&agent_states[prev_idx] + &agent_states[next_idx])
/ T::from(2.0).unwrap_or_else(|| T::zero());
self.concatenate_state_message(state, &neighbor_message)
})
.collect()
}
fn learned_communication(&self, agent_states: &[Array1<T>]) -> Vec<Array1<T>> {
if let Some(ref attention) = self.attention_weights {
agent_states
.iter()
.enumerate()
.map(|(i, state)| {
let mut messages = Vec::new();
for (j, other_state) in agent_states.iter().enumerate() {
if i != j {
let att_weight = &attention.slice(s![i, j, ..]);
let att_sum = att_weight.iter().cloned().sum::<T>();
let weighted_message = other_state.mapv(|x| x) * att_sum
/ T::from(att_weight.len() as f64).unwrap_or_else(|| T::zero());
messages.push(weighted_message);
}
}
let aggregated = if !messages.is_empty() {
messages
.iter()
.fold(Array1::zeros(state.len()), |acc, msg| acc + msg)
/ T::from(messages.len() as f64).unwrap_or_else(|| T::zero())
} else {
Array1::zeros(state.len())
};
self.concatenate_state_message(state, &aggregated)
})
.collect()
} else {
agent_states.to_vec()
}
}
}
use scirs2_core::ndarray::s;
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ScientificNumber;
#[test]
fn test_independent_q_learner_creation() {
let iql: IndependentQLearner<f64> = IndependentQLearner::new(3, 4, 5, 0.01, 0.99, 0.1);
assert_eq!(iql.n_agents, 3);
assert_eq!(iql.state_dim, 4);
assert_eq!(iql.n_actions, 5);
}
#[test]
fn test_iql_action_selection() {
let iql: IndependentQLearner<f64> = IndependentQLearner::new(3, 4, 5, 0.01, 0.99, 0.0);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 0.0, 1.0, 0.0]),
];
let actions = iql.select_actions(&states);
assert_eq!(actions.len(), 3);
for &action in &actions {
assert!(action < 5);
}
}
#[test]
fn test_multi_agent_replay_buffer() {
let mut buffer: MultiAgentReplayBuffer<f64> = MultiAgentReplayBuffer::new(100, 3);
let experience = MultiAgentExperience {
joint_state: Array1::from(vec![1.0, 2.0, 3.0]),
agent_states: vec![
Array1::from(vec![1.0]),
Array1::from(vec![2.0]),
Array1::from(vec![3.0]),
],
actions: vec![0, 1, 2],
rewards: vec![1.0, 0.5, 0.8],
next_joint_state: Array1::from(vec![1.1, 2.1, 3.1]),
next_agent_states: vec![
Array1::from(vec![1.1]),
Array1::from(vec![2.1]),
Array1::from(vec![3.1]),
],
done: false,
messages: None,
};
buffer.push(experience);
assert_eq!(buffer.len(), 1);
}
#[test]
fn test_iql_update() {
let mut iql: IndependentQLearner<f64> = IndependentQLearner::new(2, 3, 4, 0.01, 0.99, 0.1);
let experience = MultiAgentExperience {
joint_state: Array1::from(vec![1.0, 0.0, 0.0]),
agent_states: vec![Array1::from(vec![1.0, 0.0, 0.0]); 2],
actions: vec![0, 1],
rewards: vec![1.0, 0.5],
next_joint_state: Array1::from(vec![0.0, 1.0, 0.0]),
next_agent_states: vec![Array1::from(vec![0.0, 1.0, 0.0]); 2],
done: false,
messages: None,
};
let result = iql.update(&experience);
assert!(result.is_ok());
}
#[test]
fn test_iql_epsilon_decay() {
let mut iql: IndependentQLearner<f64> = IndependentQLearner::new(2, 3, 4, 0.01, 0.99, 0.5);
let initial_epsilon = iql.epsilon;
iql.decay_epsilon(0.99);
assert!(iql.epsilon < initial_epsilon);
assert!(iql.epsilon >= 0.01); }
#[test]
fn test_vdn_creation() {
let vdn: VDN<f64> = VDN::new(3, 4, 5, 0.01, 0.99, 0.1);
assert_eq!(vdn.n_agents, 3);
}
#[test]
fn test_vdn_total_q_computation() {
let vdn: VDN<f64> = VDN::new(2, 3, 4, 0.01, 0.99, 0.1);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0]),
];
let actions = vec![0, 1];
let total_q = vdn.compute_total_q(&states, &actions);
assert!(total_q.to_f64().is_some());
}
#[test]
fn test_vdn_cooperative_update() {
let mut vdn: VDN<f64> = VDN::new(2, 3, 4, 0.01, 0.99, 0.1);
let experience = MultiAgentExperience {
joint_state: Array1::from(vec![1.0, 0.0, 0.0]),
agent_states: vec![
Array1::from(vec![1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0]),
],
actions: vec![0, 1],
rewards: vec![1.0, 0.8],
next_joint_state: Array1::from(vec![0.0, 1.0, 0.0]),
next_agent_states: vec![
Array1::from(vec![0.0, 1.0, 0.0]),
Array1::from(vec![1.0, 0.0, 0.0]),
],
done: false,
messages: None,
};
let result = vdn.update_cooperative(&experience);
assert!(result.is_ok());
}
#[test]
fn test_qmix_mixing_network() {
let qmix: QMIXMixingNetwork<f64> = QMIXMixingNetwork::new(3, 4, 32);
assert_eq!(qmix.hidden_dim, 32);
let agent_q_values = vec![0.5, 0.3, 0.7];
let global_state = Array1::from(vec![1.0, 0.0, 0.0, 1.0]);
let total_q = qmix.mix(&agent_q_values, &global_state);
assert!(total_q.to_f64().is_some());
}
#[test]
fn test_maddpg_creation() {
let maddpg: MADDPG<f64> = MADDPG::new(3, 4, 2, 0.001, 0.01, 0.99, 0.001);
assert_eq!(maddpg.n_agents, 3);
}
#[test]
fn test_maddpg_action_generation() {
let maddpg: MADDPG<f64> = MADDPG::new(2, 3, 2, 0.001, 0.01, 0.99, 0.001);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0]),
];
let actions = maddpg.get_actions(&states);
assert_eq!(actions.len(), 2);
}
#[test]
fn test_maddpg_soft_update() {
let mut maddpg: MADDPG<f64> = MADDPG::new(2, 3, 2, 0.001, 0.01, 0.99, 0.1);
let initial_target = maddpg.target_actors[0].clone();
maddpg.actors[0] = maddpg.actors[0].mapv(|x| x + 1.0);
maddpg.soft_update_targets();
let updated_target = &maddpg.target_actors[0];
let diff: f64 = (&initial_target - updated_target).mapv(|x| x.abs()).sum();
assert!(diff > 0.0);
}
#[test]
fn test_communication_layer_none() {
let comm: CommunicationLayer<f64> =
CommunicationLayer::new(CommunicationProtocol::None, 4, 3);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 0.0, 1.0, 0.0]),
];
let communicated = comm.communicate(&states);
assert_eq!(communicated.len(), 3);
}
#[test]
fn test_communication_layer_broadcast() {
let comm: CommunicationLayer<f64> =
CommunicationLayer::new(CommunicationProtocol::Broadcast, 4, 3);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 0.0, 1.0, 0.0]),
];
let communicated = comm.communicate(&states);
assert_eq!(communicated.len(), 3);
assert!(communicated[0].len() > states[0].len());
}
#[test]
fn test_communication_layer_learned() {
let comm: CommunicationLayer<f64> =
CommunicationLayer::new(CommunicationProtocol::Learned, 4, 3);
let states = vec![
Array1::from(vec![1.0, 0.0, 0.0, 0.0]),
Array1::from(vec![0.0, 1.0, 0.0, 0.0]),
Array1::from(vec![0.0, 0.0, 1.0, 0.0]),
];
let communicated = comm.communicate(&states);
assert_eq!(communicated.len(), 3);
}
}