diff --git a/src/fairseq2/recipes/asr/_common.py b/src/fairseq2/recipes/asr/_common.py index b023b714b..9c12c0900 100644 --- a/src/fairseq2/recipes/asr/_common.py +++ b/src/fairseq2/recipes/asr/_common.py @@ -16,7 +16,7 @@ from fairseq2.data.text.tokenizers import TextTokenDecoder, TextTokenizer from fairseq2.gang import Gang from fairseq2.metrics import Mean -from fairseq2.metrics.text import WerMetric +from fairseq2.metrics.text import BleuMetric, WerMetric from fairseq2.models.asr import AsrModel, AsrModelOutput from fairseq2.models.seq2seq import Seq2SeqBatch from fairseq2.models.sequence import SequenceBatch @@ -121,6 +121,8 @@ def __call__( refs, ref_seqs, ref_padding_mask, hyps, hyp_seqs, hyp_padding_mask ) + metric_bag.bleu.update(refs, hyps) + try: # Dump references. stream = self._ref_output_stream @@ -148,6 +150,7 @@ def __call__( class AsrMetricBag(BaseMetricBag): ctc_loss: Mean wer: WerMetric + bleu: BleuMetric def __init__(self, gang: Gang, train: bool = True) -> None: super().__init__(gang, train=train) @@ -158,6 +161,12 @@ def __init__(self, gang: Gang, train: bool = True) -> None: self.register_metric("wer", WerMetric(device=self.device), persistent=False) + self.register_metric( + "bleu", + BleuMetric(tokenizer="flores200", device=self.device), + persistent=False, + ) + @torch.inference_mode() def update_ctc_loss(self, batch: Seq2SeqBatch, loss: Tensor) -> None: n = batch.batch_size