use crate::splitmix::{SplitmixHasher, mix, splitmix};
use crate::{Game, Moves, NodeType, Outcomes, PlayerNum, RegretParams, SolveError};
use dashmap::DashMap;
use rayon::iter::{IntoParallelIterator, ParallelIterator};
use std::hash::{BuildHasherDefault, Hash};
fn hash_of<H: Hash>(value: &H) -> u64 {
mix(0, value)
}
fn sample_key(iter: u64, context: u64, salt: u64) -> u64 {
splitmix(splitmix(iter.wrapping_mul(0x100_0001).wrapping_add(salt)) ^ context)
}
fn pick(weights: &[f64], r: u64) -> usize {
let total: f64 = weights.iter().sum();
#[allow(clippy::cast_precision_loss)]
let mut x = (r as f64 / u64::MAX as f64) * total;
for (i, &w) in weights.iter().enumerate() {
x -= w;
if x < 0.0 {
return i;
}
}
weights.len() - 1
}
#[derive(Debug, Clone, Default)]
struct Stats {
regret: f64,
strat_sum: f64,
skip_until: u64,
}
#[derive(Debug, Clone)]
struct Entry {
stats: Box<[Stats]>,
}
impl Entry {
fn new(actions: usize) -> Self {
Entry {
stats: vec![Stats::default(); actions].into_boxed_slice(),
}
}
fn average(&self) -> Vec<f64> {
let sum: f64 = self.stats.iter().map(|stat| stat.strat_sum).sum();
if sum > 0.0 {
self.stats.iter().map(|stat| stat.strat_sum / sum).collect()
} else {
#[allow(clippy::cast_precision_loss)]
let uniform = 1.0 / self.stats.len() as f64;
vec![uniform; self.stats.len()]
}
}
}
fn matched_strategy(stats: &[Stats], params: &RegretParams) -> Vec<f64> {
let count = stats.len();
let norm: f64 = stats
.iter()
.map(|stat| stat.regret)
.filter(|&value| value > 0.0)
.sum();
if norm > 0.0 {
return stats
.iter()
.map(|stat| {
if stat.regret > 0.0 {
stat.regret / norm
} else {
0.0
}
})
.collect();
}
let no_positive = params.no_positive;
if no_positive == 0.0 {
#[allow(clippy::cast_precision_loss)]
let uniform = 1.0 / count as f64;
vec![uniform; count]
} else if no_positive.is_infinite() {
let want_max = no_positive.is_sign_positive();
let pick = (0..count)
.max_by(|&left, &right| {
let order = stats[left].regret.total_cmp(&stats[right].regret);
if want_max { order } else { order.reverse() }
})
.unwrap_or(0);
let mut strat = vec![0.0; count];
strat[pick] = 1.0;
strat
} else {
let max = stats
.iter()
.map(|stat| stat.regret)
.fold(f64::NEG_INFINITY, f64::max);
let weights: Vec<f64> = stats
.iter()
.map(|stat| ((stat.regret - max) * no_positive).exp())
.collect();
let total: f64 = weights.iter().sum();
weights.iter().map(|&weight| weight / total).collect()
}
}
fn is_pruned(strat: &[f64], skip_until: &[u64], i: usize, iter: u64) -> bool {
strat[i] == 0.0 && iter < skip_until[i]
}
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] fn commit_update(
entry: &mut Entry,
regret_delta: &[f64],
strat_delta: &[f64],
explored: &[bool],
(pos, neg, strat): (f64, f64, f64),
iter: u64,
swing: f64,
) {
for (i, stat) in entry.stats.iter_mut().enumerate() {
if !explored[i] {
continue;
}
stat.regret += regret_delta[i];
stat.regret *= if stat.regret > 0.0 { pos } else { neg };
stat.strat_sum += strat_delta[i];
stat.strat_sum *= strat;
stat.skip_until = if stat.regret < 0.0 && swing > 0.0 {
iter + (stat.regret.abs() / swing) as u64
} else {
0
};
}
}
type Table<I> = DashMap<I, Entry, BuildHasherDefault<SplitmixHasher>>;
pub struct LazySolver<G: Game> {
table: [Table<G::Infoset>; 2],
iter: u64,
params: RegretParams,
factors: (f64, f64, f64), swing: f64, }
impl<G: Game> std::fmt::Debug for LazySolver<G>
where
G::Infoset: Eq + Hash,
{
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
fmt.debug_struct("LazySolver")
.field("iter", &self.iter)
.field("params", &self.params)
.field("factors", &self.factors)
.field("swing", &self.swing)
.field("infosets", &self.infosets())
.finish_non_exhaustive()
}
}
impl<G: Game> LazySolver<G>
where
G::Infoset: Eq + Hash,
{
#[must_use]
pub fn new(min_payoff: f64, max_payoff: f64) -> Self {
Self::with_params(RegretParams::dcfr(), min_payoff, max_payoff)
}
#[must_use]
pub fn with_params(params: RegretParams, min_payoff: f64, max_payoff: f64) -> Self {
LazySolver {
table: [Table::default(), Table::default()],
iter: 0,
params,
factors: (1.0, 1.0, 1.0),
swing: max_payoff - min_payoff,
}
}
#[must_use]
pub fn infosets(&self) -> usize {
let [one, two] = &self.table;
one.len() + two.len()
}
#[must_use]
pub fn player_regret_bound(&self, player: PlayerNum) -> f64 {
if self.iter == 0 {
return f64::INFINITY;
}
self.table[player.index()]
.iter()
.map(|entry| {
#[allow(clippy::cast_precision_loss)]
let iters = self.iter as f64;
let max_regret = entry
.value()
.stats
.iter()
.map(|stat| stat.regret)
.fold(0.0, f64::max);
2.0 * max_regret / iters
})
.sum()
}
pub fn run(&mut self, root: &G, iters: u64)
where
G: Clone + Sync,
G::Infoset: Send + Sync,
G::ChanceInfoset: Hash,
G::Player: Sync,
{
for _ in 0..iters {
self.iter += 1;
self.factors = self.params.iteration_factors(self.iter);
for traverser in 0..2 {
Traversal {
table: &self.table,
traverser,
cut: 0,
iter: self.iter,
params: &self.params,
factors: self.factors,
swing: self.swing,
}
.run(root.clone(), 1.0, 0, 0);
}
}
}
pub fn run_parallel(
&mut self,
root: &G,
iters: u64,
num_threads: usize,
cut: u32,
) -> Result<(), SolveError>
where
G: Clone + Sync,
G::Infoset: Send + Sync,
G::ChanceInfoset: Hash,
G::Player: Sync,
{
let mut builder = rayon::ThreadPoolBuilder::new();
if num_threads > 0 {
builder = builder.num_threads(num_threads);
}
let pool = builder.build()?;
let params = self.params;
let swing = self.swing;
pool.install(|| {
for _ in 0..iters {
self.iter += 1;
let factors = params.iteration_factors(self.iter);
for traverser in 0..2 {
Traversal {
table: &self.table,
traverser,
cut,
iter: self.iter,
params: ¶ms,
factors,
swing,
}
.run(root.clone(), 1.0, 0, 0);
}
}
});
Ok(())
}
#[must_use]
pub fn average(&self, player: PlayerNum, info: &G::Infoset) -> Option<Vec<f64>> {
self.table[player.index()]
.get(info)
.map(|entry| entry.average())
}
#[must_use]
pub fn estimate_value(&self, root: &G, samples: u64) -> f64
where
G: Clone,
{
let mut total = 0.0;
for seed in 0..samples {
total += self.playout(root, splitmix(seed.wrapping_add(0xabcd)));
}
#[allow(clippy::cast_precision_loss)]
let samples = samples as f64;
total / samples
}
fn playout(&self, root: &G, mut rng: u64) -> f64
where
G: Clone,
{
let mut state = root.clone();
loop {
match state.into_node() {
NodeType::Terminal(payoff) => return payoff,
NodeType::Chance(_info, outcomes) => {
let weights: Vec<f64> = outcomes.iter().map(|(prob, _)| prob).collect();
rng = splitmix(rng);
state = outcomes.get(pick(&weights, rng)).1;
}
NodeType::Player(num, info, moves) => {
let count = moves.len();
let strat = self.average(num, &info).unwrap_or_else(|| {
#[allow(clippy::cast_precision_loss)]
let uniform = 1.0 / count as f64;
vec![uniform; count]
});
rng = splitmix(rng);
state = moves.apply(pick(&strat, rng));
}
}
}
}
}
fn strategy_and_skip<G: Game>(
table: &[Table<G::Infoset>; 2],
player: usize,
info: &G::Infoset,
count: usize,
params: &RegretParams,
) -> (Vec<f64>, Vec<u64>)
where
G::Infoset: Eq + Hash,
{
table[player].get(info).map_or_else(
|| {
(
matched_strategy(&vec![Stats::default(); count], params),
vec![0; count],
)
},
|entry| {
let skip_until = entry.stats.iter().map(|stat| stat.skip_until).collect();
(matched_strategy(&entry.stats, params), skip_until)
},
)
}
fn descend(path: u64, index: usize) -> u64 {
splitmix(path ^ (index as u64).wrapping_add(1))
}
struct Traversal<'a, G: Game> {
table: &'a [Table<G::Infoset>; 2],
traverser: usize,
cut: u32, iter: u64,
params: &'a RegretParams,
factors: (f64, f64, f64), swing: f64,
}
impl<G> Traversal<'_, G>
where
G: Game + Sync,
G::Infoset: Eq + Hash + Send + Sync,
G::ChanceInfoset: Hash,
G::Player: Sync,
{
fn run(&self, state: G, reach: f64, depth: u32, path: u64) -> f64 {
match state.into_node() {
NodeType::Terminal(payoff) => {
if self.traverser == 0 {
payoff
} else {
-payoff
}
}
NodeType::Chance(info, outcomes) => {
let weights: Vec<f64> = outcomes.iter().map(|(prob, _)| prob).collect();
let context = info.as_ref().map_or(path, hash_of);
let choice = pick(&weights, sample_key(self.iter, context, 1));
self.run(outcomes.get(choice).1, reach, depth, descend(path, choice))
}
NodeType::Player(num, info, moves) => {
let player = num.index();
let count = moves.len();
let (strat, skip_until) =
strategy_and_skip::<G>(self.table, player, &info, count, self.params);
if player != self.traverser {
let choice = pick(&strat, sample_key(self.iter, hash_of(&info), 2));
return self.run(moves.apply(choice), reach, depth, descend(path, choice));
}
let explored: Vec<bool> = (0..count)
.map(|i| !is_pruned(&strat, &skip_until, i, self.iter))
.collect();
let explore = |i: usize| -> f64 {
if explored[i] {
self.run(moves.apply(i), reach * strat[i], depth + 1, descend(path, i))
} else {
0.0
}
};
let util: Vec<f64> = if depth < self.cut {
(0..count).into_par_iter().map(explore).collect()
} else {
(0..count).map(explore).collect()
};
let node_util: f64 = (0..count).map(|i| strat[i] * util[i]).sum();
let regret_delta: Vec<f64> = (0..count).map(|i| util[i] - node_util).collect();
let strat_delta: Vec<f64> = (0..count).map(|i| reach * strat[i]).collect();
let mut entry = self.table[player]
.entry(info)
.or_insert_with(|| Entry::new(count));
commit_update(
entry.value_mut(),
®ret_delta,
&strat_delta,
&explored,
self.factors,
self.iter,
self.swing,
);
node_util
}
}
}
}
#[cfg(test)]
mod tests {
use super::LazySolver;
use crate::{
Game, GameTree, Moves, NodeType, PlayerNum, RegretParams, SolveMethod, SolveParams,
};
use std::convert::Infallible;
#[derive(Debug, Clone)]
enum Pennies {
Start,
Mid(bool),
End(bool, bool),
}
impl Game for Pennies {
type Action = bool;
type Infoset = u8;
type ChanceInfoset = Infallible;
type Chance = Infallible;
type Player = Pennies;
fn into_node(self) -> NodeType<Self> {
match self {
Pennies::Start => NodeType::Player(PlayerNum::One, 0, Pennies::Start),
Pennies::Mid(first) => NodeType::Player(PlayerNum::Two, 1, Pennies::Mid(first)),
Pennies::End(a, b) => NodeType::Terminal(if a == b { 1.0 } else { -1.0 }),
}
}
}
impl Moves<Pennies> for Pennies {
fn len(&self) -> usize {
2
}
fn action(&self, index: usize) -> bool {
index == 0
}
fn apply(&self, index: usize) -> Pennies {
let action = index == 0;
match self {
Pennies::Start => Pennies::Mid(action),
Pennies::Mid(first) => Pennies::End(*first, action),
Pennies::End(..) => unreachable!(),
}
}
}
#[test]
fn matching_pennies_is_balanced() {
let mut solver = LazySolver::new(-1.0, 1.0);
solver.run(&Pennies::Start, 200_000);
let value = solver.estimate_value(&Pennies::Start, 200_000);
assert!(value.abs() < 0.05, "value not ~0: {value}");
let one = solver.average(PlayerNum::One, &0).unwrap();
let two = solver.average(PlayerNum::Two, &1).unwrap();
assert!((one[0] - 0.5).abs() < 0.1, "player one not ~50/50: {one:?}");
assert!((two[0] - 0.5).abs() < 0.1, "player two not ~50/50: {two:?}");
}
#[test]
fn materialized_pennies_matches() {
let game = GameTree::from_game(Pennies::Start).unwrap();
let (strats, bound) = game
.solve(SolveMethod::Full, 50_000, 0.0, 1, SolveParams::default())
.unwrap();
for player in [PlayerNum::One, PlayerNum::Two] {
assert!(
bound.player_regret_bound(player) < 0.02,
"player {player:?} not converged"
);
}
let [one, two] = strats.as_named();
for named in [one, two] {
for (_info, actions) in named {
for (_action, prob) in actions {
assert!((prob - 0.5).abs() < 0.05, "not ~50/50: {prob}");
}
}
}
}
#[test]
fn reports_low_regret_bound() {
let mut solver: LazySolver<Pennies> = LazySolver::new(-1.0, 1.0);
assert!(solver.player_regret_bound(PlayerNum::One).is_infinite());
solver.run(&Pennies::Start, 50_000);
let total =
solver.player_regret_bound(PlayerNum::One) + solver.player_regret_bound(PlayerNum::Two);
assert!(total < 0.05, "regret bound not small: {total}");
}
#[test]
fn vanilla_params_also_converge() {
let mut solver = LazySolver::with_params(RegretParams::vanilla(), -1.0, 1.0);
solver.run(&Pennies::Start, 200_000);
let value = solver.estimate_value(&Pennies::Start, 200_000);
assert!(value.abs() < 0.05, "value not ~0 under vanilla: {value}");
}
#[test]
fn pruning_params_stay_sound() {
let mut solver = LazySolver::with_params(RegretParams::dcfr_prune(), -1.0, 1.0);
solver.run(&Pennies::Start, 200_000);
let value = solver.estimate_value(&Pennies::Start, 200_000);
assert!(value.abs() < 0.05, "value not ~0 with pruning: {value}");
let one = solver.average(PlayerNum::One, &0).unwrap();
assert!(
(one[0] - 0.5).abs() < 0.1,
"player one not ~50/50 with pruning: {one:?}"
);
}
#[test]
fn parallel_matches_serial() {
let mut serial = LazySolver::new(-1.0, 1.0);
serial.run(&Pennies::Start, 50_000);
let mut parallel = LazySolver::new(-1.0, 1.0);
parallel
.run_parallel(&Pennies::Start, 50_000, 4, 2)
.unwrap();
for (player, info) in [(PlayerNum::One, 0u8), (PlayerNum::Two, 1u8)] {
assert_eq!(
serial.average(player, &info),
parallel.average(player, &info),
"serial vs parallel diverged for {player:?}"
);
}
}
}