Skip to main content

rill_ml/bandit/
thompson.rs

1//! Thompson Sampling bandit algorithm (Bernoulli rewards).
2//!
3//! Thompson Sampling maintains a Beta distribution for each arm and selects
4//! the arm with the highest sampled value. For Bernoulli rewards (0 or 1),
5//! each arm's posterior is `Beta(alpha, beta)` where `alpha = successes + prior`
6//! and `beta = failures + prior`.
7//!
8//! This implementation uses the Marsaglia-Tsang method for Gamma distribution
9//! sampling, combined via `Beta(a, b) = Gamma(a) / (Gamma(a) + Gamma(b))`.
10//! No external statistics crate is required.
11//!
12//! ## Complexity
13//!
14//! - `select`: `O(arm_count)` — samples one Beta value per arm.
15//! - `update`: `O(1)`.
16//! - Space: `O(arm_count)`.
17//!
18//! ## Reference
19//!
20//! Russo, Van Roy, Kazerouni, Osband, Wen. "A Tutorial on Thompson Sampling."
21//! Foundations and Trends in Machine Learning, 2018.
22
23use crate::bandit::stats::ArmStats;
24use crate::bandit::{
25    Bandit, checked_finite_add, checked_increment, validate_arm, validate_reward_01,
26    validate_sample_count,
27};
28use crate::error::RillError;
29#[cfg(feature = "serde")]
30use crate::persistence::ValidateState;
31use rand::Rng;
32
33/// Configuration for [`ThompsonSampling`].
34///
35/// # Examples
36///
37/// ```
38/// use rill_ml::bandit::ThompsonConfig;
39///
40/// let mut config = ThompsonConfig::default();
41/// config.alpha_prior = 1.0;
42/// config.beta_prior = 1.0;
43/// assert!((config.alpha_prior - 1.0).abs() < 1e-12);
44/// ```
45#[derive(Debug, Clone, PartialEq)]
46#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
47#[non_exhaustive]
48pub struct ThompsonConfig {
49    /// Prior alpha (success) parameter for the Beta distribution.
50    ///
51    /// Must be finite and positive. The default `1.0` gives a uniform prior
52    /// `Beta(1, 1)`.
53    pub alpha_prior: f64,
54
55    /// Prior beta (failure) parameter for the Beta distribution.
56    ///
57    /// Must be finite and positive. The default `1.0` gives a uniform prior
58    /// `Beta(1, 1)`.
59    pub beta_prior: f64,
60}
61
62impl Default for ThompsonConfig {
63    fn default() -> Self {
64        Self {
65            alpha_prior: 1.0,
66            beta_prior: 1.0,
67        }
68    }
69}
70
71impl ThompsonConfig {
72    /// Validate the configuration without constructing a bandit.
73    pub fn validate(&self) -> Result<(), RillError> {
74        if !self.alpha_prior.is_finite() || self.alpha_prior <= 0.0 {
75            return Err(RillError::InvalidParameter {
76                name: "alpha_prior",
77                value: self.alpha_prior,
78            });
79        }
80        if !self.beta_prior.is_finite() || self.beta_prior <= 0.0 {
81            return Err(RillError::InvalidParameter {
82                name: "beta_prior",
83                value: self.beta_prior,
84            });
85        }
86        Ok(())
87    }
88}
89
90/// Thompson Sampling multi-armed bandit (Bernoulli rewards).
91///
92/// Maintains a Beta posterior for each arm. On `select`, samples from each
93/// arm's posterior and returns the arm with the highest sample. On `update`,
94/// applies a soft update to the arm's alpha (success) and beta (failure)
95/// parameters.
96///
97/// Rewards must be in `[0, 1]`. The update is `alpha += reward` and
98/// `beta += 1 - reward`; for strict Bernoulli rewards this is the standard
99/// success/failure update, while fractional rewards produce a weighted update.
100///
101/// # Examples
102///
103/// ```
104/// use rill_ml::bandit::{Bandit, ThompsonSampling, ThompsonConfig};
105/// use rand::SeedableRng;
106/// use rand_chacha::ChaCha8Rng;
107///
108/// let mut rng = ChaCha8Rng::seed_from_u64(0);
109/// let mut bandit = ThompsonSampling::new(3, ThompsonConfig::default()).unwrap();
110///
111/// let arm = bandit.select(&mut rng).unwrap();
112/// bandit.update(arm, 1.0).unwrap();
113/// assert_eq!(bandit.samples_seen(), 1);
114/// ```
115#[derive(Debug, Clone)]
116#[cfg_attr(feature = "serde", derive(serde::Serialize))]
117pub struct ThompsonSampling {
118    arm_count: usize,
119    config: ThompsonConfig,
120    /// Per-arm alpha (successes + prior).
121    alphas: Vec<f64>,
122    /// Per-arm beta (failures + prior).
123    betas: Vec<f64>,
124    /// Per-arm pull counts.
125    pulls: Vec<u64>,
126    /// Per-arm total rewards (for diagnostics).
127    total_rewards: Vec<f64>,
128    /// Total number of updates.
129    samples_seen: u64,
130}
131
132impl ThompsonSampling {
133    /// Create a new Thompson Sampling bandit.
134    ///
135    /// # Errors
136    ///
137    /// Returns `RillError::InvalidArmCount` if `arm_count` is zero.
138    /// Returns `RillError::InvalidParameter` if priors are not finite and positive.
139    pub fn new(arm_count: usize, config: ThompsonConfig) -> Result<Self, RillError> {
140        if arm_count == 0 {
141            return Err(RillError::InvalidArmCount(arm_count));
142        }
143        config.validate()?;
144
145        let alpha_prior = config.alpha_prior;
146        let beta_prior = config.beta_prior;
147        Ok(Self {
148            arm_count,
149            config,
150            alphas: vec![alpha_prior; arm_count],
151            betas: vec![beta_prior; arm_count],
152            pulls: vec![0; arm_count],
153            total_rewards: vec![0.0; arm_count],
154            samples_seen: 0,
155        })
156    }
157
158    /// Per-arm alpha parameters (diagnostic).
159    pub fn alphas(&self) -> &[f64] {
160        &self.alphas
161    }
162
163    /// Per-arm beta parameters (diagnostic).
164    pub fn betas(&self) -> &[f64] {
165        &self.betas
166    }
167
168    /// Per-arm pull counts (diagnostic).
169    pub fn pulls(&self) -> &[u64] {
170        &self.pulls
171    }
172
173    /// Validate all persisted state invariants.
174    ///
175    /// This is also run automatically during deserialization.
176    pub fn validate(&self) -> Result<(), RillError> {
177        if self.arm_count == 0 {
178            return Err(RillError::InvalidArmCount(self.arm_count));
179        }
180        self.config.validate()?;
181        if self.alphas.len() != self.arm_count
182            || self.betas.len() != self.arm_count
183            || self.pulls.len() != self.arm_count
184            || self.total_rewards.len() != self.arm_count
185        {
186            return Err(RillError::InvalidState(
187                "arm_count does not match per-arm state lengths".to_owned(),
188            ));
189        }
190        validate_sample_count(&self.pulls, self.samples_seen)?;
191
192        for arm in 0..self.arm_count {
193            let pulls = self.pulls[arm] as f64;
194            let total = self.total_rewards[arm];
195            let alpha = self.alphas[arm];
196            let beta = self.betas[arm];
197            if !total.is_finite() || total < 0.0 || total > pulls {
198                return Err(RillError::InvalidState(format!(
199                    "total reward for arm {arm} is inconsistent with [0, 1] rewards"
200                )));
201            }
202            let expected_alpha = self.config.alpha_prior + total;
203            let expected_beta = self.config.beta_prior + pulls - total;
204            let alpha_tolerance = 1e-9 * expected_alpha.abs().max(1.0);
205            let beta_tolerance = 1e-9 * expected_beta.abs().max(1.0);
206            if !alpha.is_finite()
207                || !beta.is_finite()
208                || (alpha - expected_alpha).abs() > alpha_tolerance
209                || (beta - expected_beta).abs() > beta_tolerance
210            {
211                return Err(RillError::InvalidState(format!(
212                    "posterior parameters for arm {arm} are inconsistent with observations"
213                )));
214            }
215        }
216        Ok(())
217    }
218
219    /// Sample from a Beta(alpha, beta) distribution using the
220    /// Gamma ratio method.
221    ///
222    /// Beta(a, b) = Gamma(a) / (Gamma(a) + Gamma(b))
223    fn sample_beta(rng: &mut impl Rng, alpha: f64, beta: f64) -> f64 {
224        let x = Self::sample_gamma(rng, alpha);
225        let y = Self::sample_gamma(rng, beta);
226        // Handle degenerate case where both samples are 0.
227        let denom = x + y;
228        if denom <= 0.0 {
229            // Fall back to 0.5 for the degenerate case.
230            0.5
231        } else {
232            x / denom
233        }
234    }
235
236    /// Sample from a Gamma(shape, scale=1) distribution using the
237    /// Marsaglia-Tsang method.
238    ///
239    /// For shape >= 1, uses the standard acceptance-rejection method.
240    /// For shape < 1, uses the boosting trick: sample Gamma(shape+1) then
241    /// multiply by U^(1/shape).
242    fn sample_gamma(rng: &mut impl Rng, shape: f64) -> f64 {
243        if shape < 1.0 {
244            // Boosting: Gamma(shape) = Gamma(shape + 1) * U^(1/shape)
245            let u: f64 = rng.gen_range(1e-10..1.0);
246            let g = Self::sample_gamma(rng, shape + 1.0);
247            return g * u.powf(1.0 / shape);
248        }
249
250        // Marsaglia-Tsang for shape >= 1.
251        let d = shape - 1.0 / 3.0;
252        let c = 1.0 / (9.0 * d).sqrt();
253
254        loop {
255            // Sample from Normal(0, 1) using Box-Muller.
256            let (x, _unused) = Self::box_muller(rng);
257            let v = (1.0 + c * x).powi(3);
258            if v <= 0.0 {
259                continue;
260            }
261            let u: f64 = rng.gen_range(0.0..1.0);
262            if u < 1.0 - 0.0331 * x.powi(4) {
263                return d * v;
264            }
265            if u.ln() < 0.5 * x * x + d * (1.0 - v + v.ln()) {
266                return d * v;
267            }
268        }
269    }
270
271    /// Generate a pair of standard normal random variables using the
272    /// Box-Muller transform. Returns (z0, z1).
273    fn box_muller(rng: &mut impl Rng) -> (f64, f64) {
274        let u1: f64 = rng.gen_range(1e-10..1.0);
275        let u2: f64 = rng.gen_range(0.0..1.0);
276        let mag = (-2.0 * u1.ln()).sqrt();
277        let z0 = mag * (2.0 * std::f64::consts::PI * u2).cos();
278        let z1 = mag * (2.0 * std::f64::consts::PI * u2).sin();
279        (z0, z1)
280    }
281}
282
283impl Bandit for ThompsonSampling {
284    fn arm_count(&self) -> usize {
285        self.arm_count
286    }
287
288    fn samples_seen(&self) -> u64 {
289        self.samples_seen
290    }
291
292    fn select(&self, rng: &mut impl Rng) -> Result<usize, RillError> {
293        let mut best_arm = 0usize;
294        let mut best_sample = f64::NEG_INFINITY;
295
296        for arm in 0..self.arm_count {
297            let sample = Self::sample_beta(rng, self.alphas[arm], self.betas[arm]);
298            if sample > best_sample {
299                best_sample = sample;
300                best_arm = arm;
301            }
302        }
303
304        Ok(best_arm)
305    }
306
307    fn update(&mut self, arm: usize, reward: f64) -> Result<(), RillError> {
308        validate_arm(self.arm_count, arm)?;
309        validate_reward_01(reward)?;
310
311        // Soft update: alpha += reward, beta += (1 - reward).
312        // For strict Bernoulli (0 or 1), this is equivalent to the standard
313        // success/failure counting. For continuous rewards in [0, 1], this
314        // provides a weighted update.
315        let next_alpha = checked_finite_add(self.alphas[arm], reward, "alpha")?;
316        let next_beta = checked_finite_add(self.betas[arm], 1.0 - reward, "beta")?;
317        let next_pulls = checked_increment(self.pulls[arm], "pulls")?;
318        let next_total = checked_finite_add(self.total_rewards[arm], reward, "total_rewards")?;
319        let next_samples = checked_increment(self.samples_seen, "samples_seen")?;
320
321        self.alphas[arm] = next_alpha;
322        self.betas[arm] = next_beta;
323        self.pulls[arm] = next_pulls;
324        self.total_rewards[arm] = next_total;
325        self.samples_seen = next_samples;
326        Ok(())
327    }
328
329    fn reset(&mut self) {
330        self.alphas.fill(self.config.alpha_prior);
331        self.betas.fill(self.config.beta_prior);
332        self.pulls.fill(0);
333        self.total_rewards.fill(0.0);
334        self.samples_seen = 0;
335    }
336
337    fn arm_stats(&self, arm: usize) -> Result<ArmStats, RillError> {
338        validate_arm(self.arm_count, arm)?;
339        ArmStats::new(self.pulls[arm], self.total_rewards[arm])
340    }
341}
342
343#[cfg(feature = "serde")]
344#[derive(serde::Deserialize)]
345struct ThompsonSamplingState {
346    arm_count: usize,
347    config: ThompsonConfig,
348    alphas: Vec<f64>,
349    betas: Vec<f64>,
350    pulls: Vec<u64>,
351    total_rewards: Vec<f64>,
352    samples_seen: u64,
353}
354
355#[cfg(feature = "serde")]
356impl<'de> serde::Deserialize<'de> for ThompsonSampling {
357    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
358    where
359        D: serde::Deserializer<'de>,
360    {
361        let state = ThompsonSamplingState::deserialize(deserializer)?;
362        let bandit = Self {
363            arm_count: state.arm_count,
364            config: state.config,
365            alphas: state.alphas,
366            betas: state.betas,
367            pulls: state.pulls,
368            total_rewards: state.total_rewards,
369            samples_seen: state.samples_seen,
370        };
371        bandit.validate().map_err(serde::de::Error::custom)?;
372        Ok(bandit)
373    }
374}
375
376#[cfg(feature = "serde")]
377impl ValidateState for ThompsonSampling {
378    fn validate_state(&self) -> Result<(), RillError> {
379        ThompsonSampling::validate(self)
380    }
381}
382
383#[cfg(test)]
384mod tests {
385    use super::*;
386    use rand::SeedableRng;
387    use rand_chacha::ChaCha8Rng;
388
389    fn make_bandit() -> ThompsonSampling {
390        ThompsonSampling::new(3, ThompsonConfig::default()).unwrap()
391    }
392
393    #[test]
394    fn rejects_zero_arm_count() {
395        let result = ThompsonSampling::new(0, ThompsonConfig::default());
396        assert!(matches!(result, Err(RillError::InvalidArmCount(0))));
397    }
398
399    #[test]
400    fn rejects_invalid_priors() {
401        for &bad in &[0.0, -1.0, f64::NAN, f64::INFINITY] {
402            let result = ThompsonSampling::new(
403                3,
404                ThompsonConfig {
405                    alpha_prior: bad,
406                    beta_prior: 1.0,
407                },
408            );
409            assert!(matches!(result, Err(RillError::InvalidParameter { .. })));
410
411            let result = ThompsonSampling::new(
412                3,
413                ThompsonConfig {
414                    alpha_prior: 1.0,
415                    beta_prior: bad,
416                },
417            );
418            assert!(matches!(result, Err(RillError::InvalidParameter { .. })));
419        }
420    }
421
422    #[test]
423    fn initial_state() {
424        let b = make_bandit();
425        assert_eq!(b.arm_count(), 3);
426        assert_eq!(b.samples_seen(), 0);
427        // Alpha and beta should be initialized to priors.
428        for &a in b.alphas() {
429            assert!((a - 1.0).abs() < 1e-12);
430        }
431        for &be in b.betas() {
432            assert!((be - 1.0).abs() < 1e-12);
433        }
434    }
435
436    #[test]
437    fn select_returns_valid_arm() {
438        let b = make_bandit();
439        let mut rng = ChaCha8Rng::seed_from_u64(42);
440        let arm = b.select(&mut rng).unwrap();
441        assert!(arm < 3);
442    }
443
444    #[test]
445    fn update_with_success_increases_alpha() {
446        let mut b = make_bandit();
447        b.update(0, 1.0).unwrap();
448        // alpha += 1.0, beta += 0.0
449        assert!((b.alphas()[0] - 2.0).abs() < 1e-12);
450        assert!((b.betas()[0] - 1.0).abs() < 1e-12);
451    }
452
453    #[test]
454    fn update_with_failure_increases_beta() {
455        let mut b = make_bandit();
456        b.update(0, 0.0).unwrap();
457        // alpha += 0.0, beta += 1.0
458        assert!((b.alphas()[0] - 1.0).abs() < 1e-12);
459        assert!((b.betas()[0] - 2.0).abs() < 1e-12);
460    }
461
462    #[test]
463    fn update_with_continuous_reward() {
464        let mut b = make_bandit();
465        b.update(0, 0.7).unwrap();
466        // alpha += 0.7, beta += 0.3
467        assert!((b.alphas()[0] - 1.7).abs() < 1e-12);
468        assert!((b.betas()[0] - 1.3).abs() < 1e-12);
469    }
470
471    #[test]
472    fn update_rejects_invalid_arm() {
473        let mut b = make_bandit();
474        assert!(b.update(3, 1.0).is_err());
475    }
476
477    #[test]
478    fn update_rejects_reward_out_of_range() {
479        let mut b = make_bandit();
480        assert!(b.update(0, 1.5).is_err());
481        assert!(b.update(0, -0.1).is_err());
482        assert!(b.update(0, f64::NAN).is_err());
483    }
484
485    #[test]
486    fn reset_clears_state() {
487        let mut b = make_bandit();
488        b.update(0, 1.0).unwrap();
489        b.update(1, 0.0).unwrap();
490        assert_eq!(b.samples_seen(), 2);
491
492        b.reset();
493        assert_eq!(b.samples_seen(), 0);
494        for &a in b.alphas() {
495            assert!((a - 1.0).abs() < 1e-12);
496        }
497        for &be in b.betas() {
498            assert!((be - 1.0).abs() < 1e-12);
499        }
500        for &p in b.pulls() {
501            assert_eq!(p, 0);
502        }
503    }
504
505    #[test]
506    fn arm_stats_after_updates() {
507        let mut b = make_bandit();
508        b.update(0, 1.0).unwrap();
509        b.update(0, 0.0).unwrap();
510        b.update(0, 1.0).unwrap();
511        let stats = b.arm_stats(0).unwrap();
512        assert_eq!(stats.pulls, 3);
513        assert!((stats.total_reward - 2.0).abs() < 1e-12);
514    }
515
516    #[test]
517    fn arm_stats_rejects_invalid_arm() {
518        let b = make_bandit();
519        assert!(b.arm_stats(5).is_err());
520    }
521
522    #[test]
523    fn finds_best_arm_in_simulation() {
524        let mut b = make_bandit();
525        let mut rng = ChaCha8Rng::seed_from_u64(42);
526
527        // Simulate Bernoulli rewards:
528        // arm 0: p=0.8, arm 1: p=0.2, arm 2: p=0.5
529        for _ in 0..1000 {
530            let arm = b.select(&mut rng).unwrap();
531            let p = match arm {
532                0 => 0.8,
533                1 => 0.2,
534                _ => 0.5,
535            };
536            let reward = if rng.gen_range(0.0..1.0) < p {
537                1.0
538            } else {
539                0.0
540            };
541            b.update(arm, reward).unwrap();
542        }
543
544        // Arm 0 should be pulled most often.
545        let stats0 = b.arm_stats(0).unwrap();
546        let stats1 = b.arm_stats(1).unwrap();
547        let stats2 = b.arm_stats(2).unwrap();
548        assert!(stats0.pulls > stats1.pulls);
549        assert!(stats0.pulls > stats2.pulls);
550        // Arm 0's mean reward should be close to 0.8.
551        assert!(stats0.mean_reward > 0.6);
552    }
553
554    #[test]
555    fn sample_beta_returns_value_in_unit_interval() {
556        let mut rng = ChaCha8Rng::seed_from_u64(99);
557        for _ in 0..1000 {
558            let v = ThompsonSampling::sample_beta(&mut rng, 2.0, 5.0);
559            assert!((0.0..=1.0).contains(&v), "Beta sample {v} out of [0, 1]");
560        }
561    }
562
563    #[test]
564    fn sample_gamma_returns_positive_value() {
565        let mut rng = ChaCha8Rng::seed_from_u64(7);
566        for shape in &[0.5, 1.0, 2.0, 5.0, 10.0] {
567            for _ in 0..100 {
568                let v = ThompsonSampling::sample_gamma(&mut rng, *shape);
569                assert!(v > 0.0, "Gamma sample {v} not positive for shape {shape}");
570            }
571        }
572    }
573
574    #[test]
575    fn sample_gamma_mean_converges() {
576        // Gamma(shape, 1) has mean = shape.
577        let mut rng = ChaCha8Rng::seed_from_u64(42);
578        let shape = 5.0;
579        let n = 10000;
580        let mut sum = 0.0;
581        for _ in 0..n {
582            sum += ThompsonSampling::sample_gamma(&mut rng, shape);
583        }
584        let mean = sum / n as f64;
585        // Allow 10% tolerance.
586        assert!(
587            (mean - shape).abs() / shape < 0.1,
588            "Gamma mean {mean} too far from {shape}"
589        );
590    }
591
592    #[test]
593    fn sample_beta_mean_converges() {
594        // Beta(2, 5) has mean = 2 / (2 + 5) ≈ 0.2857.
595        let mut rng = ChaCha8Rng::seed_from_u64(42);
596        let alpha = 2.0;
597        let beta = 5.0;
598        let n = 10000;
599        let mut sum = 0.0;
600        for _ in 0..n {
601            sum += ThompsonSampling::sample_beta(&mut rng, alpha, beta);
602        }
603        let mean = sum / n as f64;
604        let expected = alpha / (alpha + beta);
605        // Allow 10% tolerance.
606        assert!(
607            (mean - expected).abs() / expected < 0.1,
608            "Beta mean {mean} too far from {expected}"
609        );
610    }
611
612    #[cfg(feature = "serde")]
613    #[test]
614    fn serde_roundtrip() {
615        let mut b = ThompsonSampling::new(
616            3,
617            ThompsonConfig {
618                alpha_prior: 2.0,
619                beta_prior: 3.0,
620            },
621        )
622        .unwrap();
623        b.update(0, 1.0).unwrap();
624        b.update(1, 0.0).unwrap();
625
626        let json = serde_json::to_string(&b).unwrap();
627        let restored: ThompsonSampling = serde_json::from_str(&json).unwrap();
628        assert_eq!(restored.arm_count(), b.arm_count());
629        assert_eq!(restored.samples_seen(), b.samples_seen());
630        assert_eq!(restored.alphas(), b.alphas());
631        assert_eq!(restored.betas(), b.betas());
632    }
633
634    #[cfg(feature = "serde")]
635    #[test]
636    fn serde_rejects_malformed_state() {
637        let json = r#"{
638            "arm_count": 2,
639            "config": {"alpha_prior": 1.0, "beta_prior": 1.0},
640            "alphas": [2.0],
641            "betas": [1.0],
642            "pulls": [1],
643            "total_rewards": [1.0],
644            "samples_seen": 1
645        }"#;
646        assert!(serde_json::from_str::<ThompsonSampling>(json).is_err());
647    }
648}