dynamic_weighted_sampler/
dynamic_weighted_sampler.rs1use rand::{Rng, RngExt, distr::{weighted::WeightedIndex, uniform::SampleUniform}, seq::IteratorRandom};
2use rand_distr::Distribution;
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize};
5
6const DEFAULT_CAPACITY: usize = 1000;
7
8pub trait Float:
9 Copy
10 + Default
11 + PartialEq
12 + PartialOrd
13 + std::ops::Add<Output = Self>
14 + std::ops::Sub<Output = Self>
15 + std::ops::AddAssign
16 + std::ops::SubAssign
17 + std::fmt::Debug
18 + std::fmt::Display
19 + std::ops::Mul<Output = Self>
20 + rand::distr::weighted::Weight + SampleUniform + 'static
23{
24 fn log2_ceil_bits(self) -> usize;
26
27 fn two_pow(exp: usize) -> Self;
29
30 fn random_unit<R: Rng + ?Sized>(rng: &mut R) -> Self;
32}
33
34impl Float for f32 {
35 #[inline] fn log2_ceil_bits(self) -> usize { log2_ceil2_f32(self) }
36 #[inline] fn two_pow(exp: usize) -> Self { 2.0f32.powi(exp as i32) }
37 #[inline] fn random_unit<R: Rng + ?Sized>(rng: &mut R) -> Self { rng.random::<f32>() }
38}
39
40impl Float for f64 {
41 #[inline] fn log2_ceil_bits(self) -> usize { log2_ceil2_f64(self) }
42 #[inline] fn two_pow(exp: usize) -> Self { 2.0f64.powi(exp as i32) }
43 #[inline] fn random_unit<R: Rng + ?Sized>(rng: &mut R) -> Self { rng.random::<f64>() }
44}
45
46#[derive(Debug, Clone, Copy)]
51#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
52struct Slot<W> {
53 weight: W,
54 idx_in_level: u32,
55 level: u8,
56 _pad: [u8; 3],
57}
58
59impl<W: Float> Default for Slot<W> {
60 fn default() -> Self {
61 Self { weight: W::default(), idx_in_level: 0, level: 0, _pad: [0; 3] }
63 }
64}
65
66#[derive(Debug, Clone)]
69#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
70pub struct DynamicWeightedSampler<W: Float = f32> {
71 max_value: W,
72 n_levels: usize,
73 total_weight: W,
74 slots: Vec<Slot<W>>,
75 level_weight: Vec<W>,
76 level_bucket: Vec<Vec<u32>>, level_max: Vec<W>,
78}
79
80impl<W: Float> DynamicWeightedSampler<W> {
81 pub fn new(max_value: W) -> Self {
82 Self::new_with_capacity(max_value, DEFAULT_CAPACITY)
83 }
84
85 pub fn new_with_capacity(max_value: W, physical_capacity: usize) -> Self {
86 assert!(physical_capacity > 0);
87 let n_levels = max_value.log2_ceil_bits() + 1;
88 let max_value = W::two_pow(max_value.log2_ceil_bits());
89 let slots = vec![Slot::default(); physical_capacity];
90 let level_weight = vec![W::default(); n_levels];
91 let level_bucket = vec![vec![]; n_levels];
92 let top_level = n_levels - 1;
93 let level_max: Vec<W> = (0..n_levels).map(|i| W::two_pow(top_level - i)).collect();
94 Self {
95 max_value,
96 n_levels,
97 total_weight: W::default(),
98 slots,
99 level_weight,
100 level_bucket,
101 level_max,
102 }
103 }
104
105 pub fn insert(&mut self, id: usize, weight: W) {
106 assert!(weight > W::default());
107 if id > self.slots.len() - 1 {
108 self.slots.resize(id + 1, Slot::default());
109 }
110 assert!(self.slots[id].weight == W::default(), "Inserting element id {id} with weight {weight}, but it already existed with weight {}", self.slots[id].weight);
111 assert!(weight <= self.max_value, "Adding element {id} with weight {weight} exceeds the maximum weight capacity of {}", self.max_value);
112 let level = self.level(weight);
113 self.slots[id].weight = weight;
114 self.total_weight += weight;
115 self.insert_to_level(id, level, weight);
116 }
117
118 #[inline]
119 fn level(&self, weight: W) -> usize {
120 debug_assert!(weight <= self.max_value, "{weight} > {}", self.max_value);
121 debug_assert!(weight > W::default());
122 let top_level = self.n_levels - 1;
123 let level_from_top = weight.log2_ceil_bits();
124 debug_assert!(top_level >= level_from_top);
125 top_level - level_from_top
126 }
127
128 #[inline]
129 fn insert_to_level(&mut self, id: usize, level: usize, weight: W) {
130 self.level_weight[level] += weight;
131 let idx = self.level_bucket[level].len() as u32;
132 self.level_bucket[level].push(id as u32);
133 self.slots[id].idx_in_level = idx;
134 self.slots[id].level = level as u8;
135 }
136
137 #[inline]
138 fn remove_from_level(&mut self, id: usize, level: usize, weight: W) {
139 debug_assert_eq!(self.level_bucket[level][self.slots[id].idx_in_level as usize] as usize, id);
140 self.level_weight[level] -= weight;
141 let idx_in_level = self.slots[id].idx_in_level as usize;
142 let last_idx_in_level = self.level_bucket[level].len() - 1;
143 if idx_in_level != last_idx_in_level {
144 let id_in_last = self.level_bucket[level][last_idx_in_level] as usize;
145 self.level_bucket[level].swap(idx_in_level, last_idx_in_level);
146 self.slots[id_in_last].idx_in_level = idx_in_level as u32;
147 }
148 self.level_bucket[level].pop();
149 self.slots[id].idx_in_level = 0;
150 self.slots[id].level = 0;
151 }
152
153 pub fn remove(&mut self, id: usize) -> W {
154 let slot = self.slots[id]; debug_assert!(slot.weight > W::default(), "removing element {id} with 0 weight");
156 self.slots[id].weight = W::default();
157 self.total_weight -= slot.weight;
158 self.remove_from_level(id, slot.level as usize, slot.weight);
159 slot.weight
160 }
161
162 pub fn update(&mut self, id: usize, new_weight: W) {
163 let slot = self.slots[id]; if slot.weight == new_weight {
165 return;
166 }
167 if new_weight == W::default() {
168 debug_assert!(slot.weight > W::default(), "removing element {id} with 0 weight");
170 self.slots[id].weight = W::default();
171 self.total_weight -= slot.weight;
172 self.remove_from_level(id, slot.level as usize, slot.weight);
173 return;
174 }
175 if slot.weight == W::default() {
176 self.insert(id, new_weight);
177 return;
178 }
179 let new_level = self.level(new_weight);
180 self.total_weight += new_weight;
181 self.total_weight -= slot.weight;
182 self.slots[id].weight = new_weight;
183 if slot.level as usize == new_level {
184 self.level_weight[slot.level as usize] += new_weight;
186 self.level_weight[slot.level as usize] -= slot.weight;
187 } else {
188 self.remove_from_level(id, slot.level as usize, slot.weight);
189 self.insert_to_level(id, new_level, new_weight);
190 }
191 }
192
193 pub fn update_delta(&mut self, id: usize, delta: W) {
194 let slot = self.slots[id]; let new_weight = slot.weight + delta;
196 if new_weight == slot.weight {
197 return;
198 }
199 if new_weight <= W::default() {
200 if slot.weight > W::default() {
201 self.slots[id].weight = W::default();
202 self.total_weight -= slot.weight;
203 self.remove_from_level(id, slot.level as usize, slot.weight);
204 }
205 return;
206 }
207 if slot.weight == W::default() {
208 self.insert(id, new_weight);
209 return;
210 }
211 let new_level = self.level(new_weight);
212 self.total_weight += delta;
213 self.slots[id].weight = new_weight;
214 if slot.level as usize == new_level {
215 self.level_weight[slot.level as usize] += delta;
216 } else {
217 self.remove_from_level(id, slot.level as usize, slot.weight);
218 self.insert_to_level(id, new_level, new_weight);
219 }
220 }
221
222 #[inline]
223 pub fn get_weight(&self, id: usize) -> W {
224 self.slots[id].weight
225 }
226
227 #[inline]
228 pub fn get_total_weight(&self) -> W {
229 self.total_weight
230 }
231
232 pub fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> Option<usize> {
233 assert!(self.total_weight <= self.max_value, "weighted sampler total weight {} is bigger than max weight {}.", self.total_weight, self.max_value);
234 let levels_sampler = WeightedIndex::new(self.level_weight.iter().copied()).ok()?;
235 let level = levels_sampler.sample(rng);
236
237 loop {
238 let idx_in_level = (0..self.level_bucket[level].len()).choose(rng).unwrap();
239 let sampled_id = self.level_bucket[level][idx_in_level] as usize;
240 let weight = self.slots[sampled_id].weight;
241 debug_assert!(weight <= self.level_max[level]);
243 debug_assert!(level == self.n_levels - 1 || self.level_max[level + 1] < weight);
244 let u = W::random_unit(rng) * self.level_max[level];
245 if u <= weight {
246 break Some(sampled_id);
247 }
248 }
249 }
250
251 pub fn check_invariant(&self) -> bool {
252 let sum_level = self.level_weight.iter().fold(W::default(), |acc, &w| acc + w);
253 let sum_slots = self.slots.iter().fold(W::default(), |acc, s| acc + s.weight);
254 sum_level == self.total_weight
255 && sum_slots == self.total_weight
256 && self.total_weight <= self.max_value
257 }
258}
259
260fn log2_ceil2_f32(weight: f32) -> usize {
263 let b: u32 = weight.to_bits();
264 let e = (b >> 23) & 0xFF; let frac = b & ((1 << 23) - 1); let z = if frac == 0 { e as i32 - 127 } else { e as i32 - 126 };
267 z as usize
268}
269
270fn log2_ceil2_f64(weight: f64) -> usize {
271 let b: u64 = weight.to_bits();
272 let e = (b >> 52) & ((1 << 11) - 1);
273 let frac = b & ((1 << 52) - 1);
274 let z = if frac == 0 { e as i64 - 1023 } else { e as i64 - 1022 };
275 z as usize
276}
277
278#[cfg(test)]
279mod test_weighted_sampler {
280 use std::time::Instant;
281 use std::collections::HashMap;
282 use rand::rng;
283 use super::*;
284
285 #[test]
286 fn test_distr() {
287 let mut sampler = DynamicWeightedSampler::new_with_capacity(1000., 5);
288 let mut samples: HashMap<usize, usize> = HashMap::new();
289
290 sampler.insert(1, 999.);
291 sampler.insert(2, 1.);
292
293 let n_samples = 1_000_000;
294 let start = Instant::now();
295 for _ in 1..n_samples {
296 let sample = sampler.sample(&mut rng()).unwrap();
297 *samples.entry(sample).or_default() += 1;
298 }
299 let duration = start.elapsed();
300
301 assert!(duration.as_secs() <= 3);
302 approx::assert_abs_diff_eq!(samples[&1] as f32 / n_samples as f32, 0.999, epsilon = 1e-4);
303 approx::assert_abs_diff_eq!(samples[&2] as f32 / n_samples as f32, 0.001, epsilon = 1e-4);
304
305 sampler.update(1, 99.);
306
307 samples.drain();
308 let n_samples = 1_000;
309 for _ in 1..n_samples {
310 let sample = sampler.sample(&mut rng()).unwrap();
311 *samples.entry(sample).or_default() += 1;
312 }
313
314 approx::assert_abs_diff_eq!(samples[&1] as f32 / n_samples as f32, 0.99, epsilon = 1e-2);
315 approx::assert_abs_diff_eq!(samples[&2] as f32 / n_samples as f32, 0.01, epsilon = 1e-2);
316 }
317
318 #[test]
319 fn test_distr_f64() {
320 let mut sampler = DynamicWeightedSampler::<f64>::new_with_capacity(1000., 5);
321 let mut samples: HashMap<usize, usize> = HashMap::new();
322
323 sampler.insert(1, 999.);
324 sampler.insert(2, 1.);
325
326 let n_samples = 1_000_000;
327 for _ in 1..n_samples {
328 let sample = sampler.sample(&mut rng()).unwrap();
329 *samples.entry(sample).or_default() += 1;
330 }
331
332 approx::assert_abs_diff_eq!(samples[&1] as f64 / n_samples as f64, 0.999, epsilon = 1e-4);
333 approx::assert_abs_diff_eq!(samples[&2] as f64 / n_samples as f64, 0.001, epsilon = 1e-4);
334 }
335
336 #[test]
337 fn test_remove() {
338 let mut sampler = DynamicWeightedSampler::new_with_capacity(1000., 5);
339 let level = sampler.level(500.);
340 sampler.insert(1, 500.);
341 assert_eq!(Some(&1u32), sampler.level_bucket[level].get(0));
342 sampler.insert(2, 510.);
343 assert_eq!(Some(&2u32), sampler.level_bucket[level].get(1));
344 sampler.remove(1);
345 assert_eq!(Some(&2u32), sampler.level_bucket[level].get(0));
346 sampler.insert(1, 500.);
347 assert_eq!(Some(&1u32), sampler.level_bucket[level].get(1));
348 sampler.remove(1);
349 }
350
351 #[test]
352 fn test_level() {
353 let sampler = DynamicWeightedSampler::new_with_capacity(1000., 5);
354 assert_eq!(11, sampler.n_levels);
355 assert_eq!(11 - 1, sampler.level(1.));
356 assert_eq!(11 - 2, sampler.level(2.));
357 assert_eq!(11 - 3, sampler.level(3.));
358 assert_eq!(11 - 3, sampler.level(4.));
359 }
360}