Skip to content

Commit 9c99ca8

Browse files
committed
Add support for S2TT, modeling side
1 parent 0dd2203 commit 9c99ca8

1 file changed

Lines changed: 11 additions & 1 deletion

File tree

src/fairseq2/recipes/asr/_common.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717

1818
from fairseq2.logging import log
1919
from fairseq2.metrics import Mean
20-
from fairseq2.metrics.text import WerMetric
20+
from fairseq2.metrics.text import BleuMetric, WerMetric
2121
from fairseq2.models.asr import AsrModel, AsrModelOutput
2222
from fairseq2.models.seq2seq import Seq2SeqBatch
2323
from 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__(
160163
class 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

Comments
 (0)