use crate::EvaluationError;
use scirs2_core::random::{Normal, Rng, RngExt, SeedableRng};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BayesianABTestResult {
pub prob_a_better: f64,
pub prob_b_better: f64,
pub expected_difference: f64,
pub credible_interval_95: (f64, f64),
pub credible_interval_99: (f64, f64),
pub posterior_mean_a: f64,
pub posterior_mean_b: f64,
pub bayes_factor: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BayesianEstimation {
pub posterior_mean: f64,
pub posterior_std: f64,
pub credible_interval_95: (f64, f64),
pub credible_interval_99: (f64, f64),
pub posterior_mode: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BayesianModelComparison {
pub model_names: Vec<String>,
pub log_marginal_likelihoods: Vec<f64>,
pub bayes_factors: Vec<f64>,
pub model_probabilities: Vec<f64>,
pub best_model_index: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub enum PriorType {
Uniform,
Normal,
Beta,
Gamma,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PriorParameters {
pub prior_type: PriorType,
pub param1: f64,
pub param2: f64,
}
pub struct BayesianAnalyzer {
n_samples: usize,
rng: scirs2_core::random::ChaCha8Rng,
}
impl BayesianAnalyzer {
#[must_use]
pub fn new(n_samples: usize) -> Self {
Self {
n_samples,
rng: scirs2_core::random::ChaCha8Rng::seed_from_u64(fastrand::u64(..)),
}
}
#[must_use]
pub fn default() -> Self {
Self::new(10_000)
}
pub fn ab_test(
&mut self,
group_a: &[f64],
group_b: &[f64],
) -> Result<BayesianABTestResult, EvaluationError> {
if group_a.is_empty() || group_b.is_empty() {
return Err(EvaluationError::InvalidInput {
message: "Groups cannot be empty".to_string(),
});
}
if group_a.len() < 2 || group_b.len() < 2 {
return Err(EvaluationError::InvalidInput {
message: "Groups must have at least 2 samples".to_string(),
});
}
let mean_a = group_a.iter().sum::<f64>() / group_a.len() as f64;
let mean_b = group_b.iter().sum::<f64>() / group_b.len() as f64;
let var_a =
group_a.iter().map(|x| (x - mean_a).powi(2)).sum::<f64>() / (group_a.len() - 1) as f64;
let var_b =
group_b.iter().map(|x| (x - mean_b).powi(2)).sum::<f64>() / (group_b.len() - 1) as f64;
let posterior_mean_a = mean_a;
let posterior_mean_b = mean_b;
let posterior_var_a = var_a / group_a.len() as f64;
let posterior_var_b = var_b / group_b.len() as f64;
let mut samples_a = Vec::with_capacity(self.n_samples);
let mut samples_b = Vec::with_capacity(self.n_samples);
for _ in 0..self.n_samples {
let sample_a = self.rng.sample(
Normal::new(posterior_mean_a, posterior_var_a.sqrt())
.expect("value should be present"),
);
let sample_b = self.rng.sample(
Normal::new(posterior_mean_b, posterior_var_b.sqrt())
.expect("value should be present"),
);
samples_a.push(sample_a);
samples_b.push(sample_b);
}
let mut count_a_better = 0;
let mut differences = Vec::with_capacity(self.n_samples);
for i in 0..self.n_samples {
let diff = samples_a[i] - samples_b[i];
differences.push(diff);
if diff > 0.0 {
count_a_better += 1;
}
}
let prob_a_better = count_a_better as f64 / self.n_samples as f64;
let prob_b_better = 1.0 - prob_a_better;
let expected_difference = differences.iter().sum::<f64>() / differences.len() as f64;
let mut sorted_diffs = differences.clone();
sorted_diffs.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let ci_95_lower = sorted_diffs[(self.n_samples as f64 * 0.025) as usize];
let ci_95_upper = sorted_diffs[(self.n_samples as f64 * 0.975) as usize];
let ci_99_lower = sorted_diffs[(self.n_samples as f64 * 0.005) as usize];
let ci_99_upper = sorted_diffs[(self.n_samples as f64 * 0.995) as usize];
let bayes_factor = if prob_a_better > 0.5 {
prob_a_better / (1.0 - prob_a_better)
} else {
prob_b_better / (1.0 - prob_b_better)
};
Ok(BayesianABTestResult {
prob_a_better,
prob_b_better,
expected_difference,
credible_interval_95: (ci_95_lower, ci_95_upper),
credible_interval_99: (ci_99_lower, ci_99_upper),
posterior_mean_a,
posterior_mean_b,
bayes_factor,
})
}
pub fn estimate_parameter(
&mut self,
data: &[f64],
prior: &PriorParameters,
) -> Result<BayesianEstimation, EvaluationError> {
if data.is_empty() {
return Err(EvaluationError::InvalidInput {
message: "Data cannot be empty".to_string(),
});
}
let mean = data.iter().sum::<f64>() / data.len() as f64;
let variance = data.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / data.len() as f64;
let (posterior_mean, posterior_var) = match prior.prior_type {
PriorType::Normal => {
let prior_mean = prior.param1;
let prior_var = prior.param2.powi(2);
let likelihood_var = variance / data.len() as f64;
let post_var = 1.0 / (1.0 / prior_var + data.len() as f64 / variance);
let post_mean =
post_var * (prior_mean / prior_var + data.len() as f64 * mean / variance);
(post_mean, post_var)
}
PriorType::Uniform => {
(mean, variance / data.len() as f64)
}
_ => {
(mean, variance / data.len() as f64)
}
};
let mut samples = Vec::with_capacity(self.n_samples);
for _ in 0..self.n_samples {
let sample = self.rng.sample(
Normal::new(posterior_mean, posterior_var.sqrt()).expect("value should be present"),
);
samples.push(sample);
}
samples.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let ci_95_lower = samples[(self.n_samples as f64 * 0.025) as usize];
let ci_95_upper = samples[(self.n_samples as f64 * 0.975) as usize];
let ci_99_lower = samples[(self.n_samples as f64 * 0.005) as usize];
let ci_99_upper = samples[(self.n_samples as f64 * 0.995) as usize];
let posterior_mode = posterior_mean;
Ok(BayesianEstimation {
posterior_mean,
posterior_std: posterior_var.sqrt(),
credible_interval_95: (ci_95_lower, ci_95_upper),
credible_interval_99: (ci_99_lower, ci_99_upper),
posterior_mode,
})
}
pub fn compare_models(
&self,
model_names: Vec<String>,
log_likelihoods: Vec<f64>,
) -> Result<BayesianModelComparison, EvaluationError> {
if model_names.len() != log_likelihoods.len() {
return Err(EvaluationError::InvalidInput {
message: "Number of model names must match number of log likelihoods".to_string(),
});
}
if model_names.is_empty() {
return Err(EvaluationError::InvalidInput {
message: "At least one model must be provided".to_string(),
});
}
let reference_ll = log_likelihoods[0];
let bayes_factors: Vec<f64> = log_likelihoods
.iter()
.map(|ll| (ll - reference_ll).exp())
.collect();
let total: f64 = bayes_factors.iter().sum();
let model_probabilities: Vec<f64> = bayes_factors.iter().map(|bf| bf / total).collect();
let best_model_index = model_probabilities
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx)
.expect("value should be present");
Ok(BayesianModelComparison {
model_names,
log_marginal_likelihoods: log_likelihoods,
bayes_factors,
model_probabilities,
best_model_index,
})
}
#[must_use]
pub fn interpret_bayes_factor(bf: f64) -> &'static str {
match bf {
bf if bf < 1.0 => "Negative (supports alternative)",
bf if bf < 3.0 => "Barely worth mentioning",
bf if bf < 10.0 => "Substantial",
bf if bf < 30.0 => "Strong",
bf if bf < 100.0 => "Very strong",
_ => "Decisive",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_bayesian_analyzer_creation() {
let analyzer = BayesianAnalyzer::default();
assert_eq!(analyzer.n_samples, 10_000);
}
#[test]
fn test_bayesian_ab_test() {
let mut analyzer = BayesianAnalyzer::new(1000);
let group_a = vec![3.9, 4.0, 4.1, 4.0, 3.8, 4.2, 4.0, 3.9];
let group_b = vec![3.4, 3.5, 3.6, 3.5, 3.4, 3.6, 3.5, 3.4];
let result = analyzer.ab_test(&group_a, &group_b).unwrap();
assert!(result.prob_a_better > 0.5);
assert!(result.expected_difference > 0.0);
assert!(result.posterior_mean_a > result.posterior_mean_b);
}
#[test]
fn test_bayesian_ab_test_equal_groups() {
let mut analyzer = BayesianAnalyzer::new(1000);
let group_a = vec![4.0, 4.1, 3.9, 4.0, 4.1];
let group_b = vec![4.0, 3.9, 4.1, 4.0, 3.9];
let result = analyzer.ab_test(&group_a, &group_b).unwrap();
assert!((result.prob_a_better - 0.5).abs() < 0.3);
assert!((result.prob_b_better - 0.5).abs() < 0.3);
}
#[test]
fn test_bayesian_ab_test_empty_input() {
let mut analyzer = BayesianAnalyzer::new(1000);
let empty: Vec<f64> = vec![];
let group_b = vec![1.0, 2.0, 3.0];
assert!(analyzer.ab_test(&empty, &group_b).is_err());
}
#[test]
fn test_parameter_estimation() {
let mut analyzer = BayesianAnalyzer::new(1000);
let data = vec![4.0, 4.1, 3.9, 4.0, 4.2, 3.8, 4.1, 3.9, 4.0];
let prior = PriorParameters {
prior_type: PriorType::Uniform,
param1: 0.0,
param2: 1.0,
};
let result = analyzer.estimate_parameter(&data, &prior).unwrap();
assert!((result.posterior_mean - 4.0).abs() < 0.1);
assert!(result.credible_interval_95.0 < result.posterior_mean);
assert!(result.credible_interval_95.1 > result.posterior_mean);
}
#[test]
fn test_parameter_estimation_normal_prior() {
let mut analyzer = BayesianAnalyzer::new(1000);
let data = vec![4.0, 4.1, 3.9, 4.0, 4.2];
let prior = PriorParameters {
prior_type: PriorType::Normal,
param1: 3.5, param2: 1.0, };
let result = analyzer.estimate_parameter(&data, &prior).unwrap();
assert!(result.posterior_mean > 3.5);
assert!(result.posterior_mean < 4.2);
}
#[test]
fn test_model_comparison() {
let analyzer = BayesianAnalyzer::default();
let model_names = vec![
"Model A".to_string(),
"Model B".to_string(),
"Model C".to_string(),
];
let log_likelihoods = vec![-100.0, -90.0, -95.0];
let result = analyzer
.compare_models(model_names.clone(), log_likelihoods)
.unwrap();
assert_eq!(result.best_model_index, 1); assert_eq!(result.model_names[1], "Model B");
assert!(result.model_probabilities[1] > result.model_probabilities[0]);
assert!(result.model_probabilities[1] > result.model_probabilities[2]);
}
#[test]
fn test_bayes_factor_interpretation() {
assert_eq!(
BayesianAnalyzer::interpret_bayes_factor(0.5),
"Negative (supports alternative)"
);
assert_eq!(
BayesianAnalyzer::interpret_bayes_factor(2.0),
"Barely worth mentioning"
);
assert_eq!(BayesianAnalyzer::interpret_bayes_factor(5.0), "Substantial");
assert_eq!(BayesianAnalyzer::interpret_bayes_factor(15.0), "Strong");
assert_eq!(
BayesianAnalyzer::interpret_bayes_factor(50.0),
"Very strong"
);
assert_eq!(BayesianAnalyzer::interpret_bayes_factor(150.0), "Decisive");
}
#[test]
fn test_credible_intervals() {
let mut analyzer = BayesianAnalyzer::new(1000);
let data = vec![4.0, 4.1, 3.9, 4.0, 4.2, 3.8, 4.1];
let prior = PriorParameters {
prior_type: PriorType::Uniform,
param1: 0.0,
param2: 1.0,
};
let result = analyzer.estimate_parameter(&data, &prior).unwrap();
let width_95 = result.credible_interval_95.1 - result.credible_interval_95.0;
let width_99 = result.credible_interval_99.1 - result.credible_interval_99.0;
assert!(width_99 > width_95);
}
}