use segment_tree::{
ops::{MaxIgnoreNaN, MinIgnoreNaN},
SegmentPoint,
};
use serde::{Deserialize, Serialize};
#[derive(Copy, Debug, Clone, Deserialize, Serialize, PartialEq)]
pub enum WeightNormalizer {
All,
Batch,
}
#[derive(Debug)]
pub struct SumTree {
eps: f32,
alpha: f32,
capacity: usize,
n_samples: usize,
tree: Vec<f32>,
min_tree: SegmentPoint<f32, MinIgnoreNaN>,
max_tree: SegmentPoint<f32, MaxIgnoreNaN>,
normalize: WeightNormalizer,
}
impl SumTree {
pub fn new(capacity: usize, alpha: f32, normalize: WeightNormalizer) -> Self {
Self {
eps: 1e-8,
alpha,
capacity,
n_samples: 0,
tree: vec![0f32; 2 * capacity - 1],
min_tree: SegmentPoint::build(vec![f32::MAX; capacity], MinIgnoreNaN),
max_tree: SegmentPoint::build(vec![1e-8f32; capacity], MaxIgnoreNaN),
normalize,
}
}
fn propagate(&mut self, ix: usize, change: f32) {
let parent = (ix - 1) / 2;
self.tree[parent] += change;
if parent != 0 {
self.propagate(parent, change);
}
}
fn retrieve(&self, ix: usize, s: f32) -> usize {
let left = 2 * ix + 1;
let right = left + 1;
if left >= self.tree.len() {
return ix;
}
if s <= self.tree[left] || self.tree[right] == 0f32 {
return self.retrieve(left, s);
} else {
return self.retrieve(right, s - self.tree[left]);
}
}
pub fn total(&self) -> f32 {
return self.tree[0];
}
pub fn max(&self) -> f32 {
self.max_tree
.query(0, self.max_tree.len())
.powf(1.0 / self.alpha)
}
pub fn add(&mut self, ix: usize, p: f32) {
debug_assert!(ix <= self.n_samples);
self.update(ix, p);
if self.n_samples < self.capacity {
self.n_samples += 1;
}
}
pub fn update(&mut self, ix: usize, p: f32) {
debug_assert!(ix < self.capacity);
let p = (p + self.eps).powf(self.alpha);
self.min_tree.modify(ix, p);
self.max_tree.modify(ix, p);
let ix = ix + self.capacity - 1;
let change = p - self.tree[ix];
if change.is_nan() {
println!("{:?}, {:?}", p, self.tree[ix]);
panic!();
}
self.tree[ix] = p;
self.propagate(ix, change);
}
pub fn get(&self, s: f32) -> usize {
let ix = self.retrieve(0, s);
debug_assert!(ix >= (self.capacity - 1));
ix + 1 - self.capacity
}
pub fn sample(&self, batch_size: usize, beta: f32) -> (Vec<i64>, Vec<f32>) {
let p_sum = &self.total();
let ps = (0..batch_size)
.map(|_| p_sum * fastrand::f32())
.collect::<Vec<_>>();
let indices = ps.iter().map(|&p| self.get(p)).collect::<Vec<_>>();
let n = self.n_samples as f32 / p_sum;
let ws = indices
.iter()
.map(|ix| self.tree[ix + self.capacity - 1])
.map(|p| (n * p).powf(-beta))
.collect::<Vec<_>>();
let w_max_inv = match self.normalize {
WeightNormalizer::All => (n * self.min_tree.query(0, self.n_samples)).powf(beta),
WeightNormalizer::Batch => 1f32 / ws.iter().fold(0.0 / 0.0, |m, v| v.max(m)),
};
let ws = ws.iter().map(|w| w * w_max_inv).collect::<Vec<f32>>();
if p_sum.is_nan() || w_max_inv.is_nan() || ws.iter().sum::<f32>().is_nan() {
println!("self.n_samples: {:?}", self.n_samples);
println!("p_sum: {:?}", p_sum);
println!("w_max_inv: {:?}", w_max_inv);
println!("ps: {:?}", ps);
println!("indices: {:?}", indices);
println!("{:?}", ws);
panic!();
}
let ixs = indices.iter().map(|&ix| ix as i64).collect();
(ixs, ws)
}
#[allow(dead_code)]
pub fn print_tree(&self) {
let mut nl = 1;
for i in 0..self.tree.len() {
print!("{} ", self.tree[i]);
if i == 2 * nl - 2 {
println!();
nl *= 2;
}
}
println!("max = {}", self.max());
println!("total = {}", self.total());
}
}
#[cfg(test)]
mod tests {
use super::{SumTree, WeightNormalizer::Batch};
#[test]
fn test_sum_tree_odd() {
let data = vec![0.5f32, 0.2, 0.8, 0.3, 1.1, 2.5, 3.9];
let mut sum_tree = SumTree::new(8, 1.0, Batch);
for ix in 0..data.len() {
sum_tree.add(ix, data[ix]);
}
sum_tree.print_tree();
println!();
assert_eq!(sum_tree.get(0.0), 0);
assert_eq!(sum_tree.get(0.4), 0);
assert_eq!(sum_tree.get(0.5), 0);
assert_eq!(sum_tree.get(0.6), 1);
assert_eq!(sum_tree.get(1.2), 2);
assert_eq!(sum_tree.get(1.6), 3);
assert_eq!(sum_tree.get(2.0), 4);
assert_eq!(sum_tree.get(2.8), 4);
sum_tree.update(7, 2.0);
sum_tree.print_tree();
println!();
}
}