use super::cer::edit_distance;
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::sync::Arc;
#[derive(Clone)]
pub struct WordErrorRate {
name: MetricName,
state: NumericMetricState,
pad_token: Option<usize>,
}
#[derive(new)]
pub struct WerInput {
pub outputs: Tensor<2, Int>,
pub targets: Tensor<2, Int>,
}
impl Default for WordErrorRate {
fn default() -> Self {
Self::new()
}
}
impl WordErrorRate {
pub fn new() -> Self {
Self {
name: Arc::new("WER".to_string()),
state: NumericMetricState::default(),
pad_token: None,
}
}
pub fn with_pad_token(mut self, index: usize) -> Self {
self.pad_token = Some(index);
self
}
}
impl Metric for WordErrorRate {
type Input = WerInput;
fn update(&mut self, input: &WerInput, _metadata: &MetricMetadata) -> SerializedEntry {
let outputs = input.outputs.clone();
let targets = input.targets.clone();
let [batch_size, seq_len] = targets.dims();
let outputs_data = outputs
.to_data()
.convert::<i32>()
.to_vec()
.expect("Failed to convert outputs to Vec");
let targets_data = targets
.to_data()
.convert::<i32>()
.to_vec()
.expect("Failed to convert targets to Vec");
let pad_token = self.pad_token.map(|p| p as i32);
let mut total_edit_distance = 0.0;
let mut total_target_length = 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 (ed, target_len) = match pad_token {
Some(pad) => {
let output_seq_no_pad = output_seq
.iter()
.take_while(|&&x| x != pad)
.copied()
.collect::<Vec<_>>();
let target_seq_no_pad = target_seq
.iter()
.take_while(|&&x| x != pad)
.copied()
.collect::<Vec<_>>();
(
edit_distance(&target_seq_no_pad, &output_seq_no_pad),
target_seq_no_pad.len(),
)
}
None => (edit_distance(target_seq, output_seq), target_seq.len()),
};
total_edit_distance += ed as f64;
total_target_length += target_len;
}
let value = if total_target_length > 0 {
100.0 * total_edit_distance / total_target_length as f64
} else {
0.0
};
self.state.update(value, total_target_length);
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 name(&self) -> MetricName {
self.name.clone()
}
fn clear(&mut self) {
self.state.reset();
}
fn attributes(&self) -> MetricAttributes {
NumericAttributes {
unit: Some("%".to_string()),
higher_is_better: false,
}
.into()
}
}
impl Numeric for WordErrorRate {
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_wer_without_padding() {
let device = Default::default();
let mut metric = WordErrorRate::new();
let preds = Tensor::from_data([[1, 2], [3, 4]], &device);
let tgts = Tensor::from_data([[1, 2], [3, 4]], &device);
metric.update(&WerInput::new(preds, tgts), &MetricMetadata::fake());
assert_eq!(0.0, metric.value().unwrap().current());
}
#[test]
fn test_wer_without_padding_two_errors() {
let device = Default::default();
let mut metric = WordErrorRate::new();
let preds = Tensor::from_data([[1, 2], [3, 5]], &device);
let tgts = Tensor::from_data([[1, 3], [3, 4]], &device);
metric.update(&WerInput::new(preds, tgts), &MetricMetadata::fake());
assert_eq!(50.0, metric.value().unwrap().current());
}
#[test]
fn test_wer_with_padding() {
let device = Default::default();
let pad = 9_i64;
let mut metric = WordErrorRate::new().with_pad_token(pad as usize);
let preds = Tensor::from_data([[1, 2, pad], [3, 5, pad]], &device);
let tgts = Tensor::from_data([[1, 3, pad], [3, 4, pad]], &device);
metric.update(&WerInput::new(preds, tgts), &MetricMetadata::fake());
assert_eq!(50.0, metric.value().unwrap().current());
}
#[test]
fn test_clear_resets_state() {
let device = Default::default();
let mut metric = WordErrorRate::new();
let preds = Tensor::from_data([[1, 2]], &device);
let tgts = Tensor::from_data([[1, 3]], &device);
metric.update(
&WerInput::new(preds.clone(), tgts.clone()),
&MetricMetadata::fake(),
);
assert!(metric.value().unwrap().current() > 0.0);
metric.clear();
assert!(metric.value().unwrap().current().is_nan());
}
}