use super::state::{FormatOptions, NumericMetricState};
use super::{MetricMetadata, SerializedEntry};
use crate::metric::{
Metric, MetricAttributes, MetricName, Numeric, NumericAttributes, NumericEntry,
};
use burn_core::tensor::{Int, Tensor};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Clone, Debug, Default)]
pub enum BleuSmoothing {
#[default]
None,
AddEpsilon(f64),
Exponential,
}
#[derive(Clone)]
pub struct BleuScore {
name: MetricName,
state: NumericMetricState,
max_n: usize,
pad_token: Option<usize>,
smoothing: BleuSmoothing,
}
#[derive(new)]
pub struct BleuInput {
pub outputs: Tensor<2, Int>,
pub targets: Tensor<2, Int>,
}
impl Default for BleuScore {
fn default() -> Self {
Self::with_max_n(4)
}
}
impl BleuScore {
pub fn with_max_n(max_n: usize) -> Self {
assert!(max_n >= 1, "max_n must be at least 1");
Self {
name: Arc::new(format!("BLEU-{max_n}")),
state: NumericMetricState::default(),
max_n,
pad_token: None,
smoothing: BleuSmoothing::default(),
}
}
pub fn new() -> Self {
Self::default()
}
pub fn with_pad_token(mut self, index: usize) -> Self {
self.pad_token = Some(index);
self
}
pub fn with_smoothing(mut self, smoothing: BleuSmoothing) -> Self {
self.smoothing = smoothing;
self
}
}
fn ngram_counts(tokens: &[i32], n: usize) -> HashMap<Vec<i32>, usize> {
let mut counts = HashMap::new();
if tokens.len() >= n {
for window in tokens.windows(n) {
*counts.entry(window.to_vec()).or_insert(0) += 1;
}
}
counts
}
fn corpus_bleu(
clipped_counts: &[usize],
total_counts: &[usize],
candidate_len: usize,
reference_len: usize,
max_n: usize,
smoothing: &BleuSmoothing,
) -> f64 {
if candidate_len == 0 {
return 0.0;
}
let bp = if candidate_len < reference_len {
(1.0 - reference_len as f64 / candidate_len as f64).exp()
} else {
1.0
};
let mut log_avg = 0.0;
let mut counted_orders = 0;
let mut smooth_mult = 1.0_f64;
for n in 0..max_n {
let total = total_counts[n];
let clipped = clipped_counts[n];
if total == 0 {
return 0.0;
}
let precision = if clipped == 0 {
match smoothing {
BleuSmoothing::None => return 0.0,
BleuSmoothing::AddEpsilon(eps) => *eps / total as f64,
BleuSmoothing::Exponential => {
smooth_mult *= 2.0;
1.0 / (smooth_mult * total as f64)
}
}
} else {
clipped as f64 / total as f64
};
log_avg += precision.ln();
counted_orders += 1;
}
if counted_orders == 0 {
return 0.0;
}
let score = bp * (log_avg / counted_orders as f64).exp();
score * 100.0
}
impl Metric for BleuScore {
type Input = BleuInput;
fn update(&mut self, input: &BleuInput, _metadata: &MetricMetadata) -> SerializedEntry {
let outputs = &input.outputs;
let targets = &input.targets;
let [batch_size, seq_len] = targets.dims();
let outputs_data = outputs.to_data().iter::<i32>().collect::<Vec<_>>();
let targets_data = targets.to_data().iter::<i32>().collect::<Vec<_>>();
let pad_token = self.pad_token.map(|p| p as i32);
let mut clipped_counts = vec![0usize; self.max_n];
let mut total_counts = vec![0usize; self.max_n];
let mut total_candidate_len = 0usize;
let mut total_reference_len = 0usize;
for i in 0..batch_size {
let start = i * seq_len;
let end = (i + 1) * seq_len;
let output_seq = &outputs_data[start..end];
let target_seq = &targets_data[start..end];
let output_seq = match pad_token {
Some(pad) => {
let len = output_seq
.iter()
.position(|&x| x == pad)
.unwrap_or(output_seq.len());
&output_seq[..len]
}
None => output_seq,
};
let target_seq = match pad_token {
Some(pad) => {
let len = target_seq
.iter()
.position(|&x| x == pad)
.unwrap_or(target_seq.len());
&target_seq[..len]
}
None => target_seq,
};
total_candidate_len += output_seq.len();
total_reference_len += target_seq.len();
for n in 1..=self.max_n {
let cand_ngrams = ngram_counts(output_seq, n);
let ref_ngrams = ngram_counts(target_seq, n);
for (ngram, &count) in &cand_ngrams {
let ref_count = ref_ngrams.get(ngram).copied().unwrap_or(0);
clipped_counts[n - 1] += count.min(ref_count);
total_counts[n - 1] += count;
}
}
}
let value = corpus_bleu(
&clipped_counts,
&total_counts,
total_candidate_len,
total_reference_len,
self.max_n,
&self.smoothing,
);
self.state.update(value, batch_size);
self.state
.compute_update(FormatOptions::new(self.name()).unit("%").precision(2))
}
fn compute(&mut self) -> SerializedEntry {
self.state
.compute_final(FormatOptions::new(self.name()).unit("%").precision(2))
}
fn clear(&mut self) {
self.state.reset();
}
fn name(&self) -> MetricName {
self.name.clone()
}
fn attributes(&self) -> MetricAttributes {
NumericAttributes {
unit: Some("%".to_string()),
higher_is_better: true,
}
.into()
}
}
impl Numeric for BleuScore {
fn value(&self) -> Option<NumericEntry> {
Some(self.state.current_value())
}
fn running_value(&self) -> Option<NumericEntry> {
Some(self.state.running_value())
}
fn final_value(&self) -> NumericEntry {
self.state.final_value()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_bleu_perfect_match() {
let device = Default::default();
let mut metric = BleuScore::new();
let preds = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!((metric.value().unwrap().current() - 100.0).abs() < 1e-6);
}
#[test]
fn test_bleu_no_match() {
let device = Default::default();
let mut metric = BleuScore::new();
let preds = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
let tgts = Tensor::from_data([[6, 7, 8, 9, 10]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert_eq!(0.0, metric.value().unwrap().current());
}
#[test]
fn test_bleu1_partial_match() {
let device = Default::default();
let mut metric = BleuScore::with_max_n(1);
let preds = Tensor::from_data([[1, 2, 3, 6, 7]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!((metric.value().unwrap().current() - 60.0).abs() < 1e-6);
}
#[test]
fn test_bleu_brevity_penalty() {
let device = Default::default();
let pad = 0_i64;
let mut metric = BleuScore::with_max_n(1).with_pad_token(pad as usize);
let preds = Tensor::from_data([[1, 2, 3, pad, pad]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
let expected = 100.0 * (1.0 - 5.0 / 3.0_f64).exp();
assert!((metric.value().unwrap().current() - expected).abs() < 0.1);
}
#[test]
fn test_bleu_with_padding() {
let device = Default::default();
let pad = 99_i64;
let mut metric = BleuScore::new().with_pad_token(pad as usize);
let preds = Tensor::from_data([[1, 2, 3, 4, 5, pad, pad]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5, pad, pad]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!((metric.value().unwrap().current() - 100.0).abs() < 1e-6);
}
#[test]
fn test_bleu_batch_corpus_style() {
let device = Default::default();
let mut metric = BleuScore::with_max_n(1);
let preds = Tensor::from_data([[1, 2, 3, 4, 5], [6, 7, 8, 9, 10]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5], [11, 12, 13, 14, 15]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!((metric.value().unwrap().current() - 50.0).abs() < 1e-6);
}
#[test]
fn test_clear_resets_state() {
let device = Default::default();
let mut metric = BleuScore::new();
let preds = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!(metric.value().unwrap().current() > 0.0);
metric.clear();
assert!(metric.value().unwrap().current().is_nan());
}
#[test]
fn test_bleu2_bigrams() {
let device = Default::default();
let mut metric = BleuScore::with_max_n(2);
let preds = Tensor::from_data([[1, 2, 3, 4]], &device);
let tgts = Tensor::from_data([[1, 2, 5, 6]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
let expected = 100.0 * ((0.5_f64.ln() + (1.0 / 3.0_f64).ln()) / 2.0).exp();
assert!((metric.value().unwrap().current() - expected).abs() < 0.1);
}
#[test]
fn test_bleu_custom_name() {
let metric = BleuScore::with_max_n(2);
assert_eq!(*metric.name(), "BLEU-2");
}
#[test]
fn test_bleu_default_name() {
let metric = BleuScore::new();
assert_eq!(*metric.name(), "BLEU-4");
}
#[test]
fn test_bleu_short_candidate_no_smoothing() {
let device = Default::default();
let pad = 0_i64;
let mut metric = BleuScore::new().with_pad_token(pad as usize);
let preds = Tensor::from_data([[1, 2, 3, pad, pad]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert_eq!(0.0, metric.value().unwrap().current());
}
#[test]
fn test_bleu_exponential_smoothing() {
let device = Default::default();
let mut metric = BleuScore::with_max_n(2).with_smoothing(BleuSmoothing::Exponential);
let preds = Tensor::from_data([[1, 3, 5, 7, 9]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
let mut metric_no_smooth = BleuScore::with_max_n(2);
metric_no_smooth.update(
&BleuInput::new(preds.clone(), tgts.clone()),
&MetricMetadata::fake(),
);
assert_eq!(0.0, metric_no_smooth.value().unwrap().current());
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!(
metric.value().unwrap().current() > 0.0,
"smoothing should produce non-zero score"
);
}
#[test]
fn test_bleu_add_epsilon_smoothing() {
let device = Default::default();
let mut metric = BleuScore::with_max_n(2).with_smoothing(BleuSmoothing::AddEpsilon(0.1));
let preds = Tensor::from_data([[1, 3, 5, 7, 9]], &device);
let tgts = Tensor::from_data([[1, 2, 3, 4, 5]], &device);
metric.update(&BleuInput::new(preds, tgts), &MetricMetadata::fake());
assert!(
metric.value().unwrap().current() > 0.0,
"epsilon smoothing should produce non-zero score"
);
}
}