#![allow(clippy::cast_precision_loss)]
pub trait RegretMinimizer: Clone {
fn new(num_experts: usize) -> Self
where
Self: Sized;
fn update_regret(&mut self, rewards: &[f32]);
#[must_use]
fn num_updates(&self) -> usize;
#[must_use]
fn current_strategy(&self) -> &[f32];
#[must_use]
fn cumulative_strategy(&self) -> &[f32];
fn next_action<R: rand::Rng>(&self, rng: &mut R) -> usize {
crate::probability::sample_action(self.current_strategy(), rng)
}
#[must_use]
fn best_weight(&self) -> Vec<f32> {
crate::probability::normalize_by_sum(self.cumulative_strategy())
}
#[must_use]
fn cumulative_regret(&self) -> &[f32];
#[must_use]
fn regret_weight_total(&self) -> f32 {
self.num_updates() as f32
}
#[must_use]
fn average_regret(&self) -> f32 {
let w = self.regret_weight_total();
if w <= 0.0 {
return 0.0;
}
let max_pos = self
.cumulative_regret()
.iter()
.fold(0.0_f32, |m, &r| m.max(r.max(0.0)));
max_pos / w
}
}
pub(crate) fn regret_match(regrets: &[f32], p: &mut [f32]) {
for (pi, &r) in p.iter_mut().zip(regrets) {
*pi = r.max(0.0);
}
crate::probability::normalize_inplace(p);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_regret_match_positive_regrets() {
let regrets = [2.0, 0.0, 8.0];
let mut p = vec![0.0; 3];
regret_match(®rets, &mut p);
assert!((p[0] - 0.2).abs() < 1e-6);
assert!((p[1]).abs() < 1e-6);
assert!((p[2] - 0.8).abs() < 1e-6);
}
#[test]
fn test_regret_match_all_negative_falls_back_to_uniform() {
let regrets = [-1.0, -2.0, -3.0];
let mut p = vec![0.0; 3];
regret_match(®rets, &mut p);
let expected = 1.0 / 3.0;
assert!(p.iter().all(|&v| (v - expected).abs() < 1e-6));
}
#[test]
fn test_regret_match_mixed_regrets() {
let regrets = [-5.0, 3.0, 7.0];
let mut p = vec![0.0; 3];
regret_match(®rets, &mut p);
assert!((p[0]).abs() < 1e-6); assert!((p[1] - 0.3).abs() < 1e-6);
assert!((p[2] - 0.7).abs() < 1e-6);
}
fn assert_default_formula<M: RegretMinimizer>(m: &M) {
let w = m.regret_weight_total();
let expected = if w <= 0.0 {
0.0
} else {
m.cumulative_regret()
.iter()
.fold(0.0_f32, |acc, &r| acc.max(r.max(0.0)))
/ w
};
assert!(
(m.average_regret() - expected).abs() < 1e-6,
"average_regret {} != formula {}",
m.average_regret(),
expected
);
}
#[test]
fn test_average_regret_matches_formula() {
use crate::{DiscountedRegretMatcher, PcfrPlusRegretMatcher};
let mut m = PcfrPlusRegretMatcher::new(3);
assert_eq!(m.average_regret(), 0.0);
assert_default_formula(&m);
m.update_regret(&[3.0, 0.0, 0.0]);
assert!((m.average_regret() - 2.0).abs() < 1e-6);
assert_default_formula(&m);
let mut d = DiscountedRegretMatcher::new(3);
for _ in 0..5 {
d.update_regret(&[1.0, 0.0, -1.0]);
}
assert_default_formula(&d);
}
}