#![allow(clippy::cast_precision_loss)]
use ndarray::prelude::*;
use rand::distributions::Distribution;
use rand::thread_rng;
use rand_distr::WeightedAliasIndex;
use std::vec::Vec;
use crate::errors::LittleError;
#[derive(Debug, Clone)]
pub struct RegretMatcher {
p: Array1<f32>,
sum_p: Array1<f32>,
expert_reward: Array1<f32>,
cumulative_reward: f32,
dist: WeightedAliasIndex<f32>,
num_updates: usize,
}
impl RegretMatcher {
fn init_weights(num_experts: usize) -> Vec<f32> {
vec![1.0 / num_experts as f32; num_experts]
}
pub fn new(num_experts: usize) -> Result<Self, LittleError> {
let p = Self::init_weights(num_experts);
Self::new_from_p(p)
}
pub fn new_from_p(p: Vec<f32>) -> Result<Self, LittleError> {
let num_experts = p.len();
let dist = WeightedAliasIndex::new(p.clone())?;
Ok(Self {
p: Array1::from(p),
sum_p: Array1::zeros(num_experts),
cumulative_reward: 0.0_f32,
expert_reward: Array1::from(vec![0.0_f32; num_experts]),
dist,
num_updates: 0,
})
}
pub fn next_action(&self) -> usize {
self.dist.sample(&mut thread_rng())
}
pub fn update_regret(&mut self, reward_array: ArrayView1<f32>) -> Result<(), LittleError> {
let num_experts = self.p.len();
let r = self.p.dot(&reward_array);
self.cumulative_reward += r;
self.expert_reward += &reward_array;
let regret = &self.expert_reward - self.cumulative_reward;
let capped_regret: Array1<f32> = regret.iter().map(|v: &f32| f32::max(0.0, *v)).collect();
let regret_sum = capped_regret.sum();
if regret_sum <= 0.0 {
self.p = Array1::from(Self::init_weights(num_experts));
self.cumulative_reward = 0.0;
self.expert_reward = Array1::zeros(num_experts);
self.num_updates = 0;
} else {
self.p = capped_regret / regret_sum;
self.sum_p += &self.p;
self.num_updates += 1;
}
self.dist = WeightedAliasIndex::new(self.p.to_vec())?;
Ok(())
}
#[must_use]
pub fn best_weight(&self) -> Vec<f32> {
(self.sum_p.clone() / self.num_updates as f32).to_vec()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn it_works() {
assert_eq!(2 + 2, 4);
}
#[test]
fn test_regret_gen_new() {
let _rg = RegretMatcher::new(3);
}
#[test]
fn test_next_action() {
let rg = RegretMatcher::new(100).unwrap();
for _i in 0..500 {
let a = rg.next_action();
assert!(a < 100);
}
}
}