#![allow(clippy::cast_precision_loss)]
use crate::discount::DiscountParams;
use crate::regret_minimizer::RegretMinimizer;
use crate::{probability, vector_ops};
#[derive(Debug, Clone)]
pub struct PdcfrPlusRegretMatcher {
alpha: f32,
gamma: f32,
p: Vec<f32>,
sum_p: Vec<f32>,
cumulative_regret: Vec<f32>,
last_instantaneous_regret: Vec<f32>,
regret_weight: f32,
num_updates: usize,
}
impl PdcfrPlusRegretMatcher {
#[must_use]
pub fn new_with_params(num_experts: usize, alpha: f32, gamma: f32) -> Self {
let p = probability::uniform_weights(num_experts);
Self {
alpha,
gamma,
p,
sum_p: vec![0.0; num_experts],
cumulative_regret: vec![0.0; num_experts],
last_instantaneous_regret: vec![0.0; num_experts],
regret_weight: 0.0,
num_updates: 0,
}
}
#[must_use]
pub fn recommended(num_experts: usize) -> Self {
Self::new_with_params(num_experts, 2.3, 5.0)
}
#[must_use]
pub fn alpha(&self) -> f32 {
self.alpha
}
#[must_use]
pub fn gamma(&self) -> f32 {
self.gamma
}
}
impl RegretMinimizer for PdcfrPlusRegretMatcher {
fn new(num_experts: usize) -> Self {
Self::recommended(num_experts)
}
fn update_regret(&mut self, rewards: &[f32]) {
let t = self.num_updates + 1;
let prev_discount = if t > 1 {
DiscountParams::discount_factor(t - 1, self.alpha)
} else {
0.0
};
let curr_discount = DiscountParams::discount_factor(t, self.alpha);
let strategy_discount = if t > 1 {
((t - 1) as f32 / t as f32).powf(self.gamma)
} else {
0.0
};
let expected = vector_ops::dot(&self.p, rewards);
for ((cr, lr), &rw) in self
.cumulative_regret
.iter_mut()
.zip(self.last_instantaneous_regret.iter_mut())
.zip(rewards)
{
let inst = rw - expected;
*cr = (*cr * prev_discount + inst).max(0.0);
*lr = inst;
}
for ((pi, &cr), &lr) in self
.p
.iter_mut()
.zip(self.cumulative_regret.iter())
.zip(self.last_instantaneous_regret.iter())
{
*pi = (cr * curr_discount + lr).max(0.0);
}
probability::normalize_inplace(&mut self.p);
self.regret_weight = self.regret_weight * prev_discount + 1.0;
vector_ops::discounted_accumulate(&mut self.sum_p, strategy_discount, &self.p);
self.num_updates += 1;
}
fn num_updates(&self) -> usize {
self.num_updates
}
fn current_strategy(&self) -> &[f32] {
&self.p
}
fn cumulative_strategy(&self) -> &[f32] {
&self.sum_p
}
fn regret_weight_total(&self) -> f32 {
self.regret_weight
}
fn cumulative_regret(&self) -> &[f32] {
&self.cumulative_regret
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rng;
#[test]
fn test_pdcfr_plus_new() {
let rm = PdcfrPlusRegretMatcher::new(3);
assert!((rm.alpha() - 2.3).abs() < 1e-6);
assert!((rm.gamma() - 5.0).abs() < 1e-6);
}
#[test]
fn test_pdcfr_plus_custom_params() {
let rm = PdcfrPlusRegretMatcher::new_with_params(3, 2.0, 4.0);
assert!((rm.alpha() - 2.0).abs() < 1e-6);
assert!((rm.gamma() - 4.0).abs() < 1e-6);
}
#[test]
fn test_next_action() {
let rm = PdcfrPlusRegretMatcher::new(100);
let mut rng = rng();
for _ in 0..500 {
let a = rm.next_action(&mut rng);
assert!(a < 100);
}
}
#[test]
fn test_best_weight_sums_to_one() {
let mut rm = PdcfrPlusRegretMatcher::new(3);
for _ in 0..10 {
rm.update_regret(&[1.0, 0.0, -1.0]);
}
let weights = rm.best_weight();
let sum: f32 = weights.iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
}
#[test]
fn test_num_updates_increments() {
let mut rm = PdcfrPlusRegretMatcher::new(3);
assert_eq!(rm.num_updates(), 0);
rm.update_regret(&[1.0, 0.0, -1.0]);
assert_eq!(rm.num_updates(), 1);
}
#[test]
fn test_cumulative_regret_len() {
let mut rm = PdcfrPlusRegretMatcher::new(4);
assert_eq!(rm.cumulative_regret().len(), 4);
rm.update_regret(&[1.0, 0.0, -1.0, 0.5]);
assert_eq!(rm.cumulative_regret().len(), 4);
}
#[test]
fn test_average_regret_zero_before_updates() {
let rm = PdcfrPlusRegretMatcher::new(3);
assert_eq!(rm.average_regret(), 0.0);
}
#[test]
fn test_average_regret_positive_after_dominant_action() {
let mut rm = PdcfrPlusRegretMatcher::new(3);
rm.update_regret(&[1.0, 0.0, -1.0]);
assert!(rm.average_regret() > 0.0);
}
}