use thiserror::Error;
#[derive(Error, Debug)]
pub enum PerplexityError {
#[error("Empty logits sequence provided")]
EmptySequence,
#[error("Token count mismatch: expected {expected} logit sets, got {actual}")]
TokenCountMismatch { expected: usize, actual: usize },
#[error("Invalid logits: contains NaN or infinite values")]
InvalidLogits,
#[error("Vocabulary size mismatch: logits have {logits_vocab} entries but target token {token_id} is out of range")]
VocabMismatch { logits_vocab: usize, token_id: u32 },
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PerplexityResult {
pub pre_quant: f64,
pub post_quant: f64,
pub delta: f64,
}
pub fn compute_perplexity(
logits_sequence: &[Vec<f32>],
targets: &[u32],
) -> Result<f64, PerplexityError> {
if logits_sequence.is_empty() || targets.is_empty() {
return Err(PerplexityError::EmptySequence);
}
if logits_sequence.len() != targets.len() {
return Err(PerplexityError::TokenCountMismatch {
expected: targets.len(),
actual: logits_sequence.len(),
});
}
let n = logits_sequence.len() as f64;
let mut total_log_prob = 0.0_f64;
for (logits, &target) in logits_sequence.iter().zip(targets.iter()) {
if logits.is_empty() {
return Err(PerplexityError::EmptySequence);
}
if target as usize >= logits.len() {
return Err(PerplexityError::VocabMismatch {
logits_vocab: logits.len(),
token_id: target,
});
}
if logits.iter().any(|x| x.is_nan() || x.is_infinite()) {
return Err(PerplexityError::InvalidLogits);
}
let log_prob = log_softmax_at(logits, target as usize);
total_log_prob += log_prob;
}
let avg_neg_log_prob = -total_log_prob / n;
Ok(avg_neg_log_prob.exp())
}
pub fn perplexity_delta(
original_logits: &[Vec<f32>],
quantized_logits: &[Vec<f32>],
targets: &[u32],
) -> Result<PerplexityResult, PerplexityError> {
let pre_quant = compute_perplexity(original_logits, targets)?;
let post_quant = compute_perplexity(quantized_logits, targets)?;
Ok(PerplexityResult {
pre_quant,
post_quant,
delta: post_quant - pre_quant,
})
}
fn log_softmax_at(logits: &[f32], target_idx: usize) -> f64 {
let max_val = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64;
let log_sum_exp: f64 = logits
.iter()
.map(|&x| ((x as f64) - max_val).exp())
.sum::<f64>()
.ln()
+ max_val;
(logits[target_idx] as f64) - log_sum_exp
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_perplexity_perfect_prediction() {
let logits = vec![
vec![-100.0, 100.0, -100.0], vec![100.0, -100.0, -100.0], ];
let targets = vec![1, 0];
let ppl = compute_perplexity(&logits, &targets).unwrap();
assert!(
ppl < 1.01,
"Perfect prediction should give perplexity ~1, got {}",
ppl
);
}
#[test]
fn test_perplexity_uniform_prediction() {
let logits = vec![vec![0.0, 0.0, 0.0], vec![0.0, 0.0, 0.0]];
let targets = vec![0, 1];
let ppl = compute_perplexity(&logits, &targets).unwrap();
assert!(
(ppl - 3.0).abs() < 0.01,
"Uniform prediction over 3 tokens should give perplexity ~3, got {}",
ppl
);
}
#[test]
fn test_perplexity_delta_identical_models() {
let logits = vec![vec![1.0, 2.0, 3.0], vec![3.0, 2.0, 1.0]];
let targets = vec![2, 0];
let result = perplexity_delta(&logits, &logits, &targets).unwrap();
assert!(
(result.delta).abs() < 1e-10,
"Identical models should have zero delta"
);
assert!((result.pre_quant - result.post_quant).abs() < 1e-10);
}
#[test]
fn test_perplexity_delta_degraded_model() {
let original = vec![
vec![-100.0, 100.0, -100.0], ];
let quantized = vec![vec![0.0, 0.0, 0.0]];
let targets = vec![1];
let result = perplexity_delta(&original, &quantized, &targets).unwrap();
assert!(
result.delta > 0.0,
"Degraded model should have positive delta"
);
assert!(result.post_quant > result.pre_quant);
}
#[test]
fn test_empty_sequence_error() {
let result = compute_perplexity(&[], &[]);
assert!(result.is_err());
}
#[test]
fn test_token_count_mismatch() {
let logits = vec![vec![1.0, 2.0, 3.0]];
let targets = vec![0, 1]; let result = compute_perplexity(&logits, &targets);
assert!(result.is_err());
}
#[test]
fn test_out_of_range_token() {
let logits = vec![vec![1.0, 2.0, 3.0]];
let targets = vec![5]; let result = compute_perplexity(&logits, &targets);
assert!(result.is_err());
}
#[test]
fn test_nan_logits() {
let logits = vec![vec![1.0, f32::NAN, 3.0]];
let targets = vec![0];
let result = compute_perplexity(&logits, &targets);
assert!(result.is_err());
}
#[test]
fn test_log_softmax_numerical_stability() {
let logits = vec![1000.0, 1001.0, 1002.0];
let log_prob = log_softmax_at(&logits, 2);
assert!(log_prob < 0.0, "Log probability should be negative");
assert!(
log_prob > -5.0,
"Log probability should not be extremely negative"
);
}
#[test]
fn test_perplexity_always_positive() {
let logits = vec![vec![0.5, 1.0, 0.3], vec![1.0, 0.2, 0.8]];
let targets = vec![1, 2];
let ppl = compute_perplexity(&logits, &targets).unwrap();
assert!(ppl > 0.0, "Perplexity should always be positive");
}
}