use crate::{Result, TextError};
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum SmoothingMethod {
None,
AddOne,
AddK(f64),
GoodTuring,
KneserNey,
}
#[derive(Debug, Clone)]
pub struct PerplexityConfig {
pub smoothing: SmoothingMethod,
pub min_probability: f64,
pub use_log_space: bool,
pub vocabulary_size: Option<usize>,
}
impl Default for PerplexityConfig {
fn default() -> Self {
Self {
smoothing: SmoothingMethod::None,
min_probability: 1e-12,
use_log_space: true,
vocabulary_size: None,
}
}
}
#[derive(Debug, Clone)]
pub struct PerplexityCalculator {
config: PerplexityConfig,
}
impl Default for PerplexityCalculator {
fn default() -> Self {
Self::new()
}
}
impl PerplexityCalculator {
pub fn new() -> Self {
Self {
config: PerplexityConfig::default(),
}
}
pub fn with_config(config: PerplexityConfig) -> Self {
Self { config }
}
pub fn with_smoothing(mut self, smoothing: SmoothingMethod) -> Self {
self.config.smoothing = smoothing;
self
}
pub fn with_min_probability(mut self, min_prob: f64) -> Self {
self.config.min_probability = min_prob.max(0.0);
self
}
pub fn with_log_space(mut self, use_log_space: bool) -> Self {
self.config.use_log_space = use_log_space;
self
}
pub fn with_vocabulary_size(mut self, vocab_size: usize) -> Self {
self.config.vocabulary_size = Some(vocab_size);
self
}
pub fn calculate_from_probabilities(&self, probabilities: &[f64]) -> Result<f64> {
if probabilities.is_empty() {
return Err(TextError::Other(anyhow::anyhow!(
"No probabilities provided for perplexity calculation"
)));
}
let smoothed_probs = self.apply_smoothing(probabilities)?;
self.compute_perplexity(&smoothed_probs)
}
pub fn calculate_from_logits(&self, logits: &[f64]) -> Result<f64> {
if logits.is_empty() {
return Err(TextError::Other(anyhow::anyhow!(
"No logits provided for perplexity calculation"
)));
}
let log_probs = self.logits_to_log_probabilities(logits)?;
let avg_neg_log_prob = -log_probs.iter().sum::<f64>() / log_probs.len() as f64;
Ok(avg_neg_log_prob.exp())
}
pub fn calculate_from_cross_entropy(&self, cross_entropy: f64) -> Result<f64> {
if cross_entropy < 0.0 {
return Err(TextError::Other(anyhow::anyhow!(
"Cross-entropy must be non-negative, got: {}",
cross_entropy
)));
}
Ok(cross_entropy.exp())
}
pub fn calculate_sequence_level(
&self,
sequences_probabilities: &[Vec<f64>],
) -> Result<SequencePerplexityMetrics> {
if sequences_probabilities.is_empty() {
return Err(TextError::Other(anyhow::anyhow!(
"No sequences provided for perplexity calculation"
)));
}
let mut sequence_perplexities = Vec::new();
let mut total_log_prob = 0.0;
let mut total_tokens = 0;
for sequence_probs in sequences_probabilities {
if sequence_probs.is_empty() {
continue;
}
let seq_perplexity = self.calculate_from_probabilities(sequence_probs)?;
sequence_perplexities.push(seq_perplexity);
let smoothed_probs = self.apply_smoothing(sequence_probs)?;
for &prob in &smoothed_probs {
total_log_prob += prob.max(self.config.min_probability).ln();
total_tokens += 1;
}
}
if sequence_perplexities.is_empty() {
return Err(TextError::Other(anyhow::anyhow!(
"No valid sequences found for perplexity calculation"
)));
}
let corpus_perplexity = if total_tokens > 0 {
(-total_log_prob / total_tokens as f64).exp()
} else {
0.0
};
let average_perplexity =
sequence_perplexities.iter().sum::<f64>() / sequence_perplexities.len() as f64;
let min_perplexity = sequence_perplexities
.iter()
.fold(f64::INFINITY, |a, &b| a.min(b));
let max_perplexity = sequence_perplexities
.iter()
.fold(f64::NEG_INFINITY, |a, &b| a.max(b));
let variance = sequence_perplexities
.iter()
.map(|&x| (x - average_perplexity).powi(2))
.sum::<f64>()
/ sequence_perplexities.len() as f64;
let std_deviation = variance.sqrt();
Ok(SequencePerplexityMetrics {
sequence_perplexities,
corpus_perplexity,
average_perplexity,
min_perplexity,
max_perplexity,
std_deviation,
total_sequences: sequences_probabilities.len(),
total_tokens,
})
}
pub fn compare_models(
&self,
model1_probabilities: &[Vec<f64>],
model2_probabilities: &[Vec<f64>],
model1_name: &str,
model2_name: &str,
) -> Result<ModelComparisonMetrics> {
if model1_probabilities.len() != model2_probabilities.len() {
return Err(TextError::Other(anyhow::anyhow!(
"Number of sequences must match between models: {} vs {}",
model1_probabilities.len(),
model2_probabilities.len()
)));
}
let metrics1 = self.calculate_sequence_level(model1_probabilities)?;
let metrics2 = self.calculate_sequence_level(model2_probabilities)?;
let relative_improvement = if metrics1.corpus_perplexity > 0.0 {
(metrics1.corpus_perplexity - metrics2.corpus_perplexity) / metrics1.corpus_perplexity
} else {
0.0
};
let better_model = if metrics2.corpus_perplexity < metrics1.corpus_perplexity {
model2_name
} else {
model1_name
};
let mut model1_wins = 0;
let mut model2_wins = 0;
let mut ties = 0;
for (&perp1, &perp2) in metrics1
.sequence_perplexities
.iter()
.zip(metrics2.sequence_perplexities.iter())
{
if (perp1 - perp2).abs() < 1e-6 {
ties += 1;
} else if perp1 < perp2 {
model1_wins += 1;
} else {
model2_wins += 1;
}
}
Ok(ModelComparisonMetrics {
model1_name: model1_name.to_string(),
model2_name: model2_name.to_string(),
model1_metrics: metrics1,
model2_metrics: metrics2,
relative_improvement,
better_model: better_model.to_string(),
model1_wins,
model2_wins,
ties,
})
}
pub fn calculate_confidence_interval(
&self,
probabilities: &[f64],
confidence_level: f64,
bootstrap_samples: usize,
) -> Result<ConfidenceInterval> {
if confidence_level <= 0.0 || confidence_level >= 1.0 {
return Err(TextError::Other(anyhow::anyhow!(
"Confidence level must be between 0 and 1, got: {}",
confidence_level
)));
}
if bootstrap_samples < 10 {
return Err(TextError::Other(anyhow::anyhow!(
"Bootstrap samples must be at least 10, got: {}",
bootstrap_samples
)));
}
let base_perplexity = self.calculate_from_probabilities(probabilities)?;
let mut bootstrap_perplexities = Vec::new();
for _ in 0..bootstrap_samples {
let mut resampled_probs = Vec::with_capacity(probabilities.len());
for _ in 0..probabilities.len() {
let idx = (rand::random::<f64>() * probabilities.len() as f64) as usize;
resampled_probs.push(probabilities[idx.min(probabilities.len() - 1)]);
}
if let Ok(perp) = self.calculate_from_probabilities(&resampled_probs) {
bootstrap_perplexities.push(perp);
}
}
if bootstrap_perplexities.is_empty() {
return Err(TextError::Other(anyhow::anyhow!(
"No valid bootstrap samples generated"
)));
}
bootstrap_perplexities.sort_by(|a, b| {
a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)
});
let alpha = 1.0 - confidence_level;
let lower_idx = ((alpha / 2.0) * bootstrap_perplexities.len() as f64) as usize;
let upper_idx = ((1.0 - alpha / 2.0) * bootstrap_perplexities.len() as f64) as usize;
let lower_bound = bootstrap_perplexities[lower_idx.min(bootstrap_perplexities.len() - 1)];
let upper_bound = bootstrap_perplexities[upper_idx.min(bootstrap_perplexities.len() - 1)];
Ok(ConfidenceInterval {
base_perplexity,
confidence_level,
lower_bound,
upper_bound,
bootstrap_samples: bootstrap_perplexities.len(),
})
}
fn apply_smoothing(&self, probabilities: &[f64]) -> Result<Vec<f64>> {
match self.config.smoothing {
SmoothingMethod::None => {
Ok(probabilities
.iter()
.map(|&p| p.max(self.config.min_probability))
.collect())
}
SmoothingMethod::AddOne => {
let vocab_size = self.config.vocabulary_size.unwrap_or(probabilities.len());
Ok(probabilities
.iter()
.map(|&p| {
(p * probabilities.len() as f64 + 1.0)
/ (probabilities.len() + vocab_size) as f64
})
.collect())
}
SmoothingMethod::AddK(k) => {
let vocab_size = self.config.vocabulary_size.unwrap_or(probabilities.len());
Ok(probabilities
.iter()
.map(|&p| {
(p * probabilities.len() as f64 + k)
/ (probabilities.len() as f64 + k * vocab_size as f64)
})
.collect())
}
SmoothingMethod::GoodTuring | SmoothingMethod::KneserNey => {
Ok(probabilities
.iter()
.map(|&p| {
if p == 0.0 {
self.config.min_probability
} else {
p
}
})
.collect())
}
}
}
fn logits_to_log_probabilities(&self, logits: &[f64]) -> Result<Vec<f64>> {
let max_logit = logits.iter().fold(f64::NEG_INFINITY, |acc, &x| acc.max(x));
let exp_sum: f64 = logits.iter().map(|&x| (x - max_logit).exp()).sum();
if exp_sum <= 0.0 {
return Err(TextError::Other(anyhow::anyhow!(
"Invalid logits resulted in zero or negative sum"
)));
}
let log_sum_exp = max_logit + exp_sum.ln();
Ok(logits.iter().map(|&logit| logit - log_sum_exp).collect())
}
fn compute_perplexity(&self, probabilities: &[f64]) -> Result<f64> {
let mut log_sum = 0.0;
for &prob in probabilities {
if prob <= 0.0 {
return Err(TextError::Other(anyhow::anyhow!(
"Invalid probability encountered: {}",
prob
)));
}
log_sum += prob.ln();
}
let average_log_prob = log_sum / probabilities.len() as f64;
Ok((-average_log_prob).exp())
}
}
#[derive(Debug, Clone)]
pub struct SequencePerplexityMetrics {
pub sequence_perplexities: Vec<f64>,
pub corpus_perplexity: f64,
pub average_perplexity: f64,
pub min_perplexity: f64,
pub max_perplexity: f64,
pub std_deviation: f64,
pub total_sequences: usize,
pub total_tokens: usize,
}
#[derive(Debug, Clone)]
pub struct ModelComparisonMetrics {
pub model1_name: String,
pub model2_name: String,
pub model1_metrics: SequencePerplexityMetrics,
pub model2_metrics: SequencePerplexityMetrics,
pub relative_improvement: f64,
pub better_model: String,
pub model1_wins: usize,
pub model2_wins: usize,
pub ties: usize,
}
#[derive(Debug, Clone)]
pub struct ConfidenceInterval {
pub base_perplexity: f64,
pub confidence_level: f64,
pub lower_bound: f64,
pub upper_bound: f64,
pub bootstrap_samples: usize,
}
impl SequencePerplexityMetrics {
pub fn is_consistent(&self, threshold: f64) -> bool {
self.std_deviation <= threshold
}
pub fn perplexity_range(&self) -> f64 {
self.max_perplexity - self.min_perplexity
}
pub fn coefficient_of_variation(&self) -> f64 {
if self.average_perplexity > 0.0 {
self.std_deviation / self.average_perplexity
} else {
0.0
}
}
}
impl ModelComparisonMetrics {
pub fn is_model2_better(&self) -> bool {
self.model2_metrics.corpus_perplexity < self.model1_metrics.corpus_perplexity
}
pub fn model2_win_rate(&self) -> f64 {
if self.total_comparisons() > 0 {
self.model2_wins as f64 / self.total_comparisons() as f64
} else {
0.0
}
}
pub fn total_comparisons(&self) -> usize {
self.model1_wins + self.model2_wins + self.ties
}
pub fn improvement_percentage(&self) -> f64 {
self.relative_improvement * 100.0
}
}
impl ConfidenceInterval {
pub fn contains(&self, value: f64) -> bool {
value >= self.lower_bound && value <= self.upper_bound
}
pub fn width(&self) -> f64 {
self.upper_bound - self.lower_bound
}
pub fn margin_of_error(&self) -> f64 {
self.width() / 2.0
}
}
pub fn calculate(probabilities: &[f64]) -> Result<f64> {
PerplexityCalculator::new().calculate_from_probabilities(probabilities)
}
pub fn calculate_from_logits(logits: &[f64]) -> Result<f64> {
PerplexityCalculator::new().calculate_from_logits(logits)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_perplexity_calculation() {
let calculator = PerplexityCalculator::new();
let probabilities = &[0.25, 0.25, 0.25, 0.25]; let perplexity = calculator
.calculate_from_probabilities(probabilities)
.expect("operation should succeed");
assert!((perplexity - 4.0).abs() < 1e-10);
}
#[test]
fn test_perplexity_from_logits() {
let calculator = PerplexityCalculator::new();
let logits = &[1.0, 1.0, 1.0, 1.0]; let perplexity = calculator.calculate_from_logits(logits).expect("calculation from logits should succeed");
assert!((perplexity - 4.0).abs() < 1e-10);
}
#[test]
fn test_perfect_prediction() {
let calculator = PerplexityCalculator::new();
let probabilities = &[1.0]; let perplexity = calculator
.calculate_from_probabilities(probabilities)
.expect("operation should succeed");
assert!((perplexity - 1.0).abs() < 1e-10);
}
#[test]
fn test_empty_input() {
let calculator = PerplexityCalculator::new();
let result = calculator.calculate_from_probabilities(&[]);
assert!(result.is_err());
let result = calculator.calculate_from_logits(&[]);
assert!(result.is_err());
}
#[test]
fn test_invalid_probabilities() {
let calculator = PerplexityCalculator::new();
let probabilities = &[0.5, -0.1]; let result = calculator.calculate_from_probabilities(probabilities);
assert!(result.is_ok());
}
#[test]
fn test_cross_entropy_conversion() {
let calculator = PerplexityCalculator::new();
let cross_entropy = 2.0;
let perplexity = calculator
.calculate_from_cross_entropy(cross_entropy)
.expect("operation should succeed");
assert!((perplexity - cross_entropy.exp()).abs() < 1e-10);
}
#[test]
fn test_sequence_level_analysis() {
let calculator = PerplexityCalculator::new();
let sequences = vec![vec![0.5, 0.5], vec![0.25, 0.25, 0.25, 0.25], vec![0.1, 0.9]];
let metrics = calculator.calculate_sequence_level(&sequences).expect("sequence-level calculation should succeed");
assert_eq!(metrics.total_sequences, 3);
assert_eq!(metrics.sequence_perplexities.len(), 3);
assert!(metrics.corpus_perplexity > 0.0);
assert!(metrics.average_perplexity > 0.0);
assert!(metrics.min_perplexity <= metrics.max_perplexity);
}
#[test]
fn test_model_comparison() {
let calculator = PerplexityCalculator::new();
let model1_probs = vec![
vec![0.1, 0.9], vec![0.3, 0.7], ];
let model2_probs = vec![
vec![0.01, 0.99], vec![0.05, 0.95], ];
let comparison = calculator
.compare_models(&model1_probs, &model2_probs, "Model1", "Model2")
.expect("operation should succeed");
assert_eq!(comparison.better_model, "Model2");
assert!(comparison.relative_improvement > 0.0);
assert_eq!(comparison.model2_wins, 2);
assert_eq!(comparison.model1_wins, 0);
}
#[test]
fn test_smoothing_methods() {
let calculator_none = PerplexityCalculator::new().with_smoothing(SmoothingMethod::None);
let calculator_add_one = PerplexityCalculator::new()
.with_smoothing(SmoothingMethod::AddOne)
.with_vocabulary_size(10);
let probabilities = &[0.0, 0.5, 0.5];
let perp_none = calculator_none
.calculate_from_probabilities(probabilities)
.expect("operation should succeed");
let perp_add_one = calculator_add_one
.calculate_from_probabilities(probabilities)
.expect("operation should succeed");
assert!(perp_none > 0.0);
assert!(perp_add_one > 0.0);
assert_ne!(perp_none, perp_add_one);
}
#[test]
fn test_confidence_interval() {
let calculator = PerplexityCalculator::new();
let probabilities = vec![0.2, 0.3, 0.1, 0.4];
let ci = calculator
.calculate_confidence_interval(&probabilities, 0.95, 100)
.expect("operation should succeed");
assert!(ci.confidence_level == 0.95);
assert!(ci.lower_bound <= ci.base_perplexity);
assert!(ci.upper_bound >= ci.base_perplexity);
assert!(ci.lower_bound <= ci.upper_bound);
assert!(ci.bootstrap_samples > 0);
}
#[test]
fn test_legacy_functions() {
let probabilities = &[0.25, 0.25, 0.25, 0.25];
let perplexity = calculate(probabilities).expect("calculation should succeed");
assert!((perplexity - 4.0).abs() < 1e-10);
let logits = &[1.0, 1.0, 1.0, 1.0];
let perplexity = calculate_from_logits(logits).expect("calculation from logits should succeed");
assert!((perplexity - 4.0).abs() < 1e-10);
}
#[test]
fn test_numerical_stability() {
let calculator = PerplexityCalculator::new();
let small_probs = &[1e-10, 1e-10, 1.0 - 2e-10];
let perplexity = calculator
.calculate_from_probabilities(small_probs)
.expect("operation should succeed");
assert!(perplexity.is_finite());
assert!(perplexity > 0.0);
let large_logits = &[100.0, 101.0, 99.0];
let perplexity = calculator.calculate_from_logits(large_logits).expect("calculation from logits should succeed");
assert!(perplexity.is_finite());
assert!(perplexity > 0.0);
}
}