1717
1818from fairseq2 .logging import log
1919from fairseq2 .metrics import Mean
20- from fairseq2 .metrics .text import WerMetric
20+ from fairseq2 .metrics .text import BleuMetric , WerMetric
2121from fairseq2 .models .asr import AsrModel , AsrModelOutput
2222from fairseq2 .models .seq2seq import Seq2SeqBatch
2323from fairseq2 .models .sequence import SequenceBatch
@@ -46,6 +46,8 @@ def __call__(
4646 ) -> tuple [Tensor , int ]:
4747 log .info (f"s3: calling forward" )
4848 output = self ._forward (batch )
49+ # self._model.module._gang = metric_bag._gang
50+ # output = self._forward(batch, metric_bag._gang)
4951
5052 log .info (f"s4: calling loss" )
5153 loss , extra_metrics = output .compute_loss (
@@ -132,6 +134,7 @@ def __call__(
132134 metric_bag .wer .update (
133135 refs , ref_seqs , ref_padding_mask , hyps , hyp_seqs , hyp_padding_mask
134136 )
137+ metric_bag .bleu .update (refs , hyps )
135138
136139 try :
137140 # Dump references.
@@ -160,6 +163,7 @@ def __call__(
160163class AsrMetricBag (BaseMetricBag ):
161164 ctc_loss : Mean
162165 wer : WerMetric
166+ bleu : BleuMetric
163167
164168 def __init__ (self , gang : Gang , train : bool = True ) -> None :
165169 super ().__init__ (gang , train = train )
@@ -170,6 +174,12 @@ def __init__(self, gang: Gang, train: bool = True) -> None:
170174
171175 self .register_metric ("wer" , WerMetric (device = self .device ), persistent = False )
172176
177+ self .register_metric (
178+ "bleu" ,
179+ BleuMetric (tokenizer = "flores200" , device = self .device ),
180+ persistent = False ,
181+ )
182+
173183 @torch .inference_mode ()
174184 def update_ctc_loss (self , batch : Seq2SeqBatch , loss : Tensor ) -> None :
175185 n = batch .batch_size
0 commit comments