use crate::storage::{AccumCell, StorageBackend};
use crate::update_rule::UpdateRule;
pub trait RegretLane<B: StorageBackend>: Sized {
fn new(num_rows: usize, num_actions: usize) -> Self;
fn read_row(&self, row: usize, num_actions: usize, out: &mut [f32]);
fn write_row(&self, row: usize, num_actions: usize, regret: &[f32]);
}
pub trait StrategyLane<R: UpdateRule, B: StorageBackend>: Sized {
fn new(num_rows: usize, num_actions: usize) -> Self;
fn accumulate(
&self,
row: usize,
num_actions: usize,
step: &R::Step,
strategy: &[f32],
update_count: usize,
);
fn average_into(&self, row: usize, num_actions: usize, out: &mut [f32]);
fn reset(&self);
}
pub struct F32Regret<B: StorageBackend> {
cells: Vec<B::Cell<u32>>,
}
impl<B: StorageBackend> RegretLane<B> for F32Regret<B> {
fn new(num_rows: usize, num_actions: usize) -> Self {
Self {
cells: (0..num_rows * num_actions)
.map(|_| B::Cell::<u32>::default())
.collect(),
}
}
fn read_row(&self, row: usize, n: usize, out: &mut [f32]) {
for (i, slot) in out[..n].iter_mut().enumerate() {
*slot = f32::from_bits(self.cells[row * n + i].load());
}
}
fn write_row(&self, row: usize, n: usize, regret: &[f32]) {
for (i, &v) in regret[..n].iter().enumerate() {
self.cells[row * n + i].store(v.to_bits());
}
}
}
pub struct F32SumStrategy<B: StorageBackend> {
cells: Vec<B::Cell<u32>>,
}
impl<B: StorageBackend> F32SumStrategy<B> {
pub(crate) fn new(num_rows: usize, num_actions: usize) -> Self {
Self {
cells: (0..num_rows * num_actions)
.map(|_| B::Cell::<u32>::default())
.collect(),
}
}
#[inline]
fn load(&self, idx: usize) -> f32 {
f32::from_bits(self.cells[idx].load())
}
#[inline]
fn store(&self, idx: usize, v: f32) {
self.cells[idx].store(v.to_bits());
}
}
impl<R: UpdateRule, B: StorageBackend> StrategyLane<R, B> for F32SumStrategy<B> {
fn new(num_rows: usize, num_actions: usize) -> Self {
F32SumStrategy::new(num_rows, num_actions)
}
fn accumulate(
&self,
row: usize,
n: usize,
step: &R::Step,
strategy: &[f32],
_update_count: usize,
) {
let (discount, weight) = R::strategy_accumulation(step);
for (i, &s) in strategy[..n].iter().enumerate() {
let idx = row * n + i;
self.store(idx, self.load(idx) * discount + weight * s);
}
}
fn average_into(&self, row: usize, n: usize, out: &mut [f32]) {
for (i, slot) in out[..n].iter_mut().enumerate() {
*slot = self.load(row * n + i);
}
crate::probability::normalize_inplace(&mut out[..n]);
}
fn reset(&self) {
self.reset_cells();
}
}
#[allow(dead_code)] impl<B: StorageBackend> F32SumStrategy<B> {
pub(crate) fn reset_cells(&self) {
for cell in &self.cells {
cell.store(0u32); }
}
pub(crate) fn accumulate<R: UpdateRule>(
&self,
row: usize,
n: usize,
step: &R::Step,
s: &[f32],
update_count: usize,
) {
<Self as StrategyLane<R, B>>::accumulate(self, row, n, step, s, update_count);
}
pub(crate) fn average_into<R: UpdateRule>(&self, row: usize, n: usize, out: &mut [f32]) {
<Self as StrategyLane<R, B>>::average_into(self, row, n, out);
}
}
#[cfg(test)]
impl<B: StorageBackend> F32SumStrategy<B> {
pub(crate) fn strategy_raw_cell(&self, row: usize, i: usize, num_actions: usize) -> f32 {
self.load(row * num_actions + i)
}
}
pub struct U16AvgStrategy<B: StorageBackend> {
cells: Vec<B::Cell<u16>>,
weight: Vec<B::Cell<u32>>, }
impl<B: StorageBackend> U16AvgStrategy<B> {
const MAX: u32 = u16::MAX as u32;
pub(crate) fn new(num_rows: usize, num_actions: usize) -> Self {
Self {
cells: (0..num_rows * num_actions)
.map(|_| B::Cell::<u16>::default())
.collect(),
weight: (0..num_rows).map(|_| B::Cell::<u32>::default()).collect(),
}
}
#[inline]
fn sigma(&self, idx: usize) -> f32 {
crate::unit_fixed::decode(self.cells[idx].load() as u32, Self::MAX)
}
#[inline]
fn set_sigma_stochastic(&self, idx: usize, v: f32, u01: f32) {
#[allow(clippy::cast_possible_truncation)]
self.cells[idx].store(crate::unit_fixed::encode_stochastic(v, Self::MAX, u01) as u16);
}
#[inline]
fn w_load(&self, row: usize) -> f32 {
f32::from_bits(self.weight[row].load())
}
#[inline]
fn w_store(&self, row: usize, v: f32) {
self.weight[row].store(v.to_bits());
}
}
impl<R: UpdateRule, B: StorageBackend> StrategyLane<R, B> for U16AvgStrategy<B> {
fn new(num_rows: usize, num_actions: usize) -> Self {
U16AvgStrategy::new(num_rows, num_actions)
}
fn accumulate(
&self,
row: usize,
n: usize,
step: &R::Step,
strategy: &[f32],
update_count: usize,
) {
let (discount, weight) = R::strategy_accumulation(step);
let w_new = discount * self.w_load(row) + weight;
self.w_store(row, w_new);
let frac = if w_new > 0.0 { weight / w_new } else { 0.0 };
for (i, &s) in strategy[..n].iter().enumerate() {
let idx = row * n + i;
let cur = self.sigma(idx);
let v = cur + frac * (s - cur);
let u = crate::unit_fixed::u01(row, i, update_count);
self.set_sigma_stochastic(idx, v, u);
}
}
fn average_into(&self, row: usize, n: usize, out: &mut [f32]) {
for (i, slot) in out[..n].iter_mut().enumerate() {
*slot = self.sigma(row * n + i);
}
crate::probability::normalize_inplace(&mut out[..n]);
}
fn reset(&self) {
self.reset_cells();
}
}
#[allow(dead_code)] impl<B: StorageBackend> U16AvgStrategy<B> {
pub(crate) fn reset_cells(&self) {
for cell in &self.cells {
cell.store(0u16); }
for w in &self.weight {
w.store(0u32); }
}
pub(crate) fn accumulate<R: UpdateRule>(
&self,
row: usize,
n: usize,
step: &R::Step,
s: &[f32],
update_count: usize,
) {
<Self as StrategyLane<R, B>>::accumulate(self, row, n, step, s, update_count);
}
pub(crate) fn average_into<R: UpdateRule>(&self, row: usize, n: usize, out: &mut [f32]) {
<Self as StrategyLane<R, B>>::average_into(self, row, n, out);
}
}
pub struct U16AvgStrategyShared<B: StorageBackend> {
cells: Vec<B::Cell<u16>>,
weight: B::Cell<u32>, }
impl<B: StorageBackend> U16AvgStrategyShared<B> {
const MAX: u32 = u16::MAX as u32;
pub(crate) fn new(num_rows: usize, num_actions: usize) -> Self {
Self {
cells: (0..num_rows * num_actions)
.map(|_| B::Cell::<u16>::default())
.collect(),
weight: B::Cell::<u32>::default(),
}
}
#[inline]
fn sigma(&self, idx: usize) -> f32 {
crate::unit_fixed::decode(self.cells[idx].load() as u32, Self::MAX)
}
#[inline]
fn set_sigma_stochastic(&self, idx: usize, v: f32, u01: f32) {
#[allow(clippy::cast_possible_truncation)]
self.cells[idx].store(crate::unit_fixed::encode_stochastic(v, Self::MAX, u01) as u16);
}
#[inline]
fn w_load(&self) -> f32 {
f32::from_bits(self.weight.load())
}
#[inline]
fn w_store(&self, v: f32) {
self.weight.store(v.to_bits());
}
}
impl<R: UpdateRule, B: StorageBackend> StrategyLane<R, B> for U16AvgStrategyShared<B> {
fn new(num_rows: usize, num_actions: usize) -> Self {
U16AvgStrategyShared::new(num_rows, num_actions)
}
fn accumulate(
&self,
row: usize,
n: usize,
step: &R::Step,
strategy: &[f32],
update_count: usize,
) {
let (discount, weight) = R::strategy_accumulation(step);
if row == 0 {
let w_new = discount * self.w_load() + weight;
self.w_store(w_new);
}
let w = self.w_load();
let frac = if w > 0.0 { weight / w } else { 0.0 };
for (i, &s) in strategy[..n].iter().enumerate() {
let idx = row * n + i;
let cur = self.sigma(idx);
let v = cur + frac * (s - cur);
let u = crate::unit_fixed::u01(row, i, update_count);
self.set_sigma_stochastic(idx, v, u);
}
}
fn average_into(&self, row: usize, n: usize, out: &mut [f32]) {
for (i, slot) in out[..n].iter_mut().enumerate() {
*slot = self.sigma(row * n + i);
}
crate::probability::normalize_inplace(&mut out[..n]);
}
fn reset(&self) {
self.reset_cells();
}
}
#[allow(dead_code)] impl<B: StorageBackend> U16AvgStrategyShared<B> {
pub(crate) fn reset_cells(&self) {
for cell in &self.cells {
cell.store(0u16); }
self.weight.store(0u32); }
pub(crate) fn accumulate<R: UpdateRule>(
&self,
row: usize,
n: usize,
step: &R::Step,
s: &[f32],
update_count: usize,
) {
<Self as StrategyLane<R, B>>::accumulate(self, row, n, step, s, update_count);
}
pub(crate) fn average_into<R: UpdateRule>(&self, row: usize, n: usize, out: &mut [f32]) {
<Self as StrategyLane<R, B>>::average_into(self, row, n, out);
}
}
pub trait Layout<R: UpdateRule, B: StorageBackend> {
type Regret: RegretLane<B>;
type Strategy: StrategyLane<R, B>;
}
pub struct F32Full;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for F32Full {
type Regret = F32Regret<B>;
type Strategy = F32SumStrategy<B>;
}
pub struct HalfStrategy;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for HalfStrategy {
type Regret = F32Regret<B>;
type Strategy = U16AvgStrategy<B>;
}
pub struct Int16Regret<B: StorageBackend> {
cells: Vec<B::Cell<u16>>, scale: Vec<B::Cell<u32>>, }
impl<B: StorageBackend> Int16Regret<B> {
#[inline]
fn scale_load(&self, row: usize) -> f32 {
f32::from_bits(self.scale[row].load())
}
#[inline]
fn scale_store(&self, row: usize, s: f32) {
self.scale[row].store(s.to_bits());
}
#[inline]
fn code(&self, idx: usize) -> i16 {
#[allow(clippy::cast_possible_truncation)]
let v = self.cells[idx].load() as i16;
v
}
#[inline]
fn set_code(&self, idx: usize, q: i16) {
#[allow(clippy::cast_possible_truncation)]
self.cells[idx].store(q as u16);
}
}
impl<B: StorageBackend> RegretLane<B> for Int16Regret<B> {
fn new(num_rows: usize, num_actions: usize) -> Self {
Self {
cells: (0..num_rows * num_actions)
.map(|_| B::Cell::<u16>::default())
.collect(),
scale: (0..num_rows).map(|_| B::Cell::<u32>::default()).collect(),
}
}
fn read_row(&self, row: usize, n: usize, out: &mut [f32]) {
let s = self.scale_load(row);
let s = if s > 0.0 { s } else { 0.0 }; for (i, slot) in out[..n].iter_mut().enumerate() {
*slot = crate::scaled_int::decode(self.code(row * n + i), s);
}
}
fn write_row(&self, row: usize, n: usize, regret: &[f32]) {
let s = crate::scaled_int::choose_scale(®ret[..n]);
self.scale_store(row, s);
for (i, &r) in regret[..n].iter().enumerate() {
self.set_code(row * n + i, crate::scaled_int::encode(r, s));
}
}
}
pub struct HalfRegret;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for HalfRegret {
type Regret = Int16Regret<B>;
type Strategy = F32SumStrategy<B>;
}
pub struct HalfBoth;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for HalfBoth {
type Regret = Int16Regret<B>;
type Strategy = U16AvgStrategy<B>;
}
pub struct HalfStrategyShared;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for HalfStrategyShared {
type Regret = F32Regret<B>;
type Strategy = U16AvgStrategyShared<B>;
}
pub struct HalfBothShared;
impl<R: UpdateRule, B: StorageBackend> Layout<R, B> for HalfBothShared {
type Regret = Int16Regret<B>;
type Strategy = U16AvgStrategyShared<B>;
}
#[cfg(test)]
mod tests {
use super::*;
use crate::discount::DiscountParams;
use crate::rules::Dcfr;
use crate::storage::Local;
use crate::update_rule::UpdateRule;
#[test]
fn f32_regret_lane_round_trips_exactly() {
let lane = F32Regret::<Local>::new(2, 3);
let vals = [1.5f32, -2.0, 1e9];
lane.write_row(1, 3, &vals);
let mut out = [0.0f32; 3];
lane.read_row(1, 3, &mut out);
for (a, b) in vals.iter().zip(&out) {
assert_eq!(a.to_bits(), b.to_bits(), "exact f32 round-trip");
}
lane.read_row(0, 3, &mut out);
assert!(out.iter().all(|&v| v == 0.0));
}
#[test]
fn f32_sum_strategy_matches_accumulate_then_normalize() {
let lane = F32SumStrategy::<Local>::new(1, 3);
let params = DiscountParams::RECOMMENDED;
let step = Dcfr::step(¶ms, 1);
let strat = [0.2f32, 0.3, 0.5];
lane.accumulate::<Dcfr>(0, 3, &step, &strat, 1);
let mut out = [0.0f32; 3];
lane.average_into::<Dcfr>(0, 3, &mut out);
for (a, b) in strat.iter().zip(&out) {
assert!((a - b).abs() < 1e-6, "{a} vs {b}");
}
}
#[test]
fn u16_average_keeps_moving_to_horizon() {
use crate::rules::LinearCfr; let lane = U16AvgStrategy::<Local>::new(1, 2);
let mut last = [0.0f32; 2];
let mut moved_late = false;
for t in 1..=20_000usize {
let step = LinearCfr::step(&(), t);
let strat = if t % 2 == 0 {
[0.7f32, 0.3]
} else {
[0.3f32, 0.7]
};
lane.accumulate::<LinearCfr>(0, 2, &step, &strat, t);
if t > 15_000 {
let mut cur = [0.0f32; 2];
lane.average_into::<LinearCfr>(0, 2, &mut cur);
if (cur[0] - last[0]).abs() > 1e-6 {
moved_late = true;
}
last = cur;
}
}
assert!(moved_late, "u16 average froze before the horizon");
}
#[test]
fn u16_avg_matches_f32_sum_average_within_quantum() {
use crate::rules::Dcfr;
use crate::storage::Local;
let params = DiscountParams::RECOMMENDED;
let f32_lane = F32SumStrategy::<Local>::new(1, 3);
let u16_lane = U16AvgStrategy::<Local>::new(1, 3);
let seq = [
[0.5f32, 0.3, 0.2],
[0.1, 0.8, 0.1],
[0.33, 0.33, 0.34],
[0.6, 0.1, 0.3],
];
for (t, strat) in seq.iter().enumerate() {
let step = Dcfr::step(¶ms, t + 1);
f32_lane.accumulate::<Dcfr>(0, 3, &step, strat, t + 1);
u16_lane.accumulate::<Dcfr>(0, 3, &step, strat, t + 1);
}
let mut a = [0.0f32; 3];
let mut b = [0.0f32; 3];
f32_lane.average_into::<Dcfr>(0, 3, &mut a);
u16_lane.average_into::<Dcfr>(0, 3, &mut b);
for (x, y) in a.iter().zip(&b) {
assert!(
(x - y).abs() < 2.0 / u16::MAX as f32,
"u16 avg within a couple quanta: {x} vs {y}"
);
}
}
#[test]
fn int16_regret_round_trips_within_row_quantum() {
let lane = Int16Regret::<Local>::new(1, 3);
let regret = [1000.0f32, -250.0, 30.0];
lane.write_row(0, 3, ®ret);
let mut out = [0.0f32; 3];
lane.read_row(0, 3, &mut out);
let s = 1000.0f32 / i16::MAX as f32;
for (a, b) in regret.iter().zip(&out) {
assert!((a - b).abs() <= s + 1e-2, "{a} vs {b}");
}
}
#[test]
fn int16_regret_rescales_on_growth_without_overflow() {
let lane = Int16Regret::<Local>::new(1, 2);
lane.write_row(0, 2, &[1.0, -1.0]); lane.write_row(0, 2, &[1.0e6, -5.0e5]); let mut out = [0.0f32; 2];
lane.read_row(0, 2, &mut out);
assert!(
(out[0] - 1.0e6).abs() / 1.0e6 < 1e-3,
"rescaled high value preserved: {out:?}"
);
assert!(out[1] < 0.0, "sign preserved");
}
#[test]
fn u16_reset_then_first_accumulate_lands_sigma1() {
let lane = U16AvgStrategy::<Local>::new(1, 3);
let params = DiscountParams::RECOMMENDED;
for t in 1..=5usize {
let step = Dcfr::step(¶ms, t);
lane.accumulate::<Dcfr>(0, 3, &step, &[0.5, 0.3, 0.2], t);
}
lane.reset_cells();
let sigma1 = [0.6f32, 0.25, 0.15];
let step = Dcfr::step(¶ms, 6);
lane.accumulate::<Dcfr>(0, 3, &step, &sigma1, 6);
let mut out = [0.0f32; 3];
lane.average_into::<Dcfr>(0, 3, &mut out);
let quantum = 2.0 / u16::MAX as f32;
for (a, b) in sigma1.iter().zip(&out) {
assert!(
(a - b).abs() < quantum,
"σ̄ should equal σ1 within u16 quantum: {a} vs {b}"
);
}
}
#[test]
fn shared_weight_matches_per_row_under_batch() {
let params = DiscountParams::RECOMMENDED;
let per_row = U16AvgStrategy::<Local>::new(2, 3);
let shared = U16AvgStrategyShared::<Local>::new(2, 3);
let strategies = [
[0.5f32, 0.3, 0.2],
[0.1f32, 0.8, 0.1],
[0.4f32, 0.4, 0.2],
[0.6f32, 0.1, 0.3],
[0.33f32, 0.33, 0.34],
[0.2f32, 0.5, 0.3],
[0.7f32, 0.15, 0.15],
[0.25f32, 0.5, 0.25],
];
for (t, sigma) in strategies.iter().enumerate() {
let step = Dcfr::step(¶ms, t + 1);
per_row.accumulate::<Dcfr>(0, 3, &step, sigma, t + 1);
per_row.accumulate::<Dcfr>(1, 3, &step, sigma, t + 1);
shared.accumulate::<Dcfr>(0, 3, &step, sigma, t + 1);
shared.accumulate::<Dcfr>(1, 3, &step, sigma, t + 1);
}
let quantum = 2.0 / u16::MAX as f32;
for row in 0..2 {
let mut a = [0.0f32; 3];
let mut b = [0.0f32; 3];
per_row.average_into::<Dcfr>(row, 3, &mut a);
shared.average_into::<Dcfr>(row, 3, &mut b);
for (x, y) in a.iter().zip(&b) {
assert!(
(x - y).abs() < quantum,
"row {row}: per-row {x} vs shared {y} — diff exceeds u16 quantum"
);
}
}
}
#[test]
fn u16_stochastic_tracks_f32_where_round_to_nearest_freezes() {
let params = DiscountParams::new(2.3, 0.0, 10.0); let max = u16::MAX as u32;
let quantum = 1.0 / max as f32;
let f32_lane = F32SumStrategy::<Local>::new(1, 2);
let u16_lane = U16AvgStrategy::<Local>::new(1, 2);
let mut rn_code: u32 = 0;
let mut rn_w: f32 = 0.0;
let n = 1_200_000usize;
let probe = 700_000usize; let (mut f32_at_probe, mut u16_at_probe, mut rn_at_probe) = (0.0f32, 0.0f32, 0u32);
for t in 1..=n {
let step = Dcfr::step(¶ms, t);
let g = t as f32 / n as f32;
let sigma = [0.3 + 0.4 * g, 0.7 - 0.4 * g];
f32_lane.accumulate::<Dcfr>(0, 2, &step, &sigma, t);
u16_lane.accumulate::<Dcfr>(0, 2, &step, &sigma, t);
let (discount, weight) = Dcfr::strategy_accumulation(&step);
rn_w = discount * rn_w + weight;
let frac = if rn_w > 0.0 { weight / rn_w } else { 0.0 };
let cur = rn_code as f32 / max as f32;
rn_code = crate::unit_fixed::encode(cur + frac * (sigma[0] - cur), max);
if t == probe {
let mut tmp = [0.0f32; 2];
f32_lane.average_into::<Dcfr>(0, 2, &mut tmp);
f32_at_probe = tmp[0];
u16_lane.average_into::<Dcfr>(0, 2, &mut tmp);
u16_at_probe = tmp[0];
rn_at_probe = rn_code;
}
}
let mut f = [0.0f32; 2];
let mut u = [0.0f32; 2];
f32_lane.average_into::<Dcfr>(0, 2, &mut f);
u16_lane.average_into::<Dcfr>(0, 2, &mut u);
let rn_end = rn_code as f32 / max as f32;
let f32_late_move = (f[0] - f32_at_probe).abs();
assert!(
f32_late_move > 10.0 * quantum,
"not exercising late movement: {f32_late_move}"
);
let rn_late_move = (rn_end - rn_at_probe as f32 / max as f32).abs();
assert!(
rn_late_move < 2.0 * quantum,
"round-to-nearest unexpectedly moved: {rn_late_move}"
);
let u16_late_move = (u[0] - u16_at_probe).abs();
assert!(
u16_late_move > 5.0 * quantum,
"stochastic u16 froze: {u16_late_move}"
);
assert!(
(u[0] - f[0]).abs() < (rn_end - f[0]).abs(),
"stochastic should track f32 better than frozen RN: u16={u:?} f32={f:?} rn={rn_end}"
);
}
}