Skip to content

Commit 0a3c03f

Browse files
authored
fix bleu score calculation (#1426)
Signed-off-by: naymaraq <dkaramyan@nvidia.com> Co-authored-by: naymaraq <dkaramyan@nvidia.com>
1 parent 10f04f5 commit 0a3c03f

14 files changed

Lines changed: 99 additions & 75 deletions

File tree

nemo_skills/dataset/covost2/__init__.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,3 @@
1616
METRICS_TYPE = "audio"
1717
EVAL_ARGS = "++eval_type=audio ++eval_config.normalization_mode=multilingual"
1818
GENERATION_ARGS = "++prompt_format=openai ++enable_audio=true"
19-
JUDGE_PIPELINE_ARGS = {
20-
"source_key": "extra_fields.src_text",
21-
"reference_key": "extra_fields.tgt_text",
22-
}

nemo_skills/dataset/covost2/prepare.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -173,9 +173,11 @@ def _build_record(
173173
subset_for_metrics: str,
174174
task_type: str,
175175
extra_fields: dict,
176+
source: str | None = None,
177+
reference: str | None = None,
176178
) -> dict:
177179
audio_metadata = {"path": container_audio_path, "duration": duration}
178-
return {
180+
record = {
179181
"expected_answer": expected_answer,
180182
"audio_path": container_audio_path,
181183
"duration": duration,
@@ -187,6 +189,11 @@ def _build_record(
187189
"task_type": f"Multilingual-{task_type.upper()}",
188190
"extra_fields": extra_fields,
189191
}
192+
if source is not None:
193+
record["source"] = source
194+
if reference is not None:
195+
record["reference"] = reference
196+
return record
190197

191198

192199
def prepare_covost2(
@@ -269,6 +276,8 @@ def prepare_covost2(
269276
"src_lang": src_lang,
270277
"tgt_lang": tgt_lang,
271278
},
279+
source=item["sentence"],
280+
reference=item["translation"],
272281
)
273282
out.write(json.dumps(record, ensure_ascii=False) + "\n")
274283
print(f"CoVoST2 {task_type} dataset prepared: {output_jsonl}")

nemo_skills/dataset/fleurs/__init__.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,3 @@
1616
METRICS_TYPE = "audio"
1717
EVAL_ARGS = "++eval_type=audio ++eval_config.normalization_mode=multilingual"
1818
GENERATION_ARGS = "++prompt_format=openai ++enable_audio=true"
19-
JUDGE_PIPELINE_ARGS = {
20-
"source_key": "extra_fields.src_raw_text",
21-
"reference_key": "extra_fields.tgt_raw_text",
22-
}

nemo_skills/dataset/fleurs/prepare.py

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -151,9 +151,11 @@ def _build_record(
151151
subset_for_metrics: str,
152152
task_type: str,
153153
extra_fields: dict,
154+
source: str | None = None,
155+
reference: str | None = None,
154156
) -> dict:
155157
audio_metadata = {"path": container_audio_path, "duration": duration}
156-
return {
158+
record = {
157159
"expected_answer": expected_answer,
158160
"audio_path": container_audio_path,
159161
"duration": duration,
@@ -165,6 +167,11 @@ def _build_record(
165167
"task_type": f"Multilingual-{task_type.upper()}",
166168
"extra_fields": extra_fields,
167169
}
170+
if source is not None:
171+
record["source"] = source
172+
if reference is not None:
173+
record["reference"] = reference
174+
return record
168175

169176

170177
def prepare_fleurs(data_dir: Path, split: str, languages: list[str], no_audio: bool, task_type: str) -> None:
@@ -244,8 +251,12 @@ def prepare_fleurs(data_dir: Path, split: str, languages: list[str], no_audio: b
244251
"tgt_lang_group": FLEURS_LANG_TO_GROUP[tgt_locale],
245252
}
246253
)
254+
source_text = source_row["raw_transcription"]
255+
reference_text = target_row["raw_transcription"]
247256
else:
248257
expected_answer = source_row[gt_key]
258+
source_text = None
259+
reference_text = None
249260

250261
record = _build_record(
251262
expected_answer=expected_answer,
@@ -255,6 +266,8 @@ def prepare_fleurs(data_dir: Path, split: str, languages: list[str], no_audio: b
255266
subset_for_metrics=subset_for_metrics,
256267
task_type=task_type,
257268
extra_fields=extra_fields,
269+
source=source_text,
270+
reference=reference_text,
258271
)
259272
out.write(json.dumps(record, ensure_ascii=False) + "\n")
260273
print(f"Fleurs {task_type} dataset prepared: {output_jsonl}")

nemo_skills/dataset/flores200/__init__.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,3 @@
1616
METRICS_TYPE = "translation"
1717
GENERATION_ARGS = "++prompt_config=multilingual/segment-translation"
1818
EVAL_SPLIT = "devtest"
19-
JUDGE_PIPELINE_ARGS = {
20-
"source_key": "text",
21-
"reference_key": "translation",
22-
}

nemo_skills/dataset/flores200/prepare.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,10 +25,10 @@ def write_data_to_file(output_file, datasets, src_languages, tgt_languages):
2525
for src_lang in src_languages:
2626
for tgt_lang in tgt_languages:
2727
if src_lang != tgt_lang:
28-
for src, tgt in zip(datasets[src_lang], datasets[tgt_lang], strict=True):
28+
for src, ref in zip(datasets[src_lang], datasets[tgt_lang], strict=True):
2929
json_dict = {
30-
"text": src,
31-
"translation": tgt,
30+
"source": src,
31+
"reference": ref,
3232
"source_language": src_lang,
3333
"target_language": tgt_lang,
3434
"source_lang_name": Language(src_lang).display_name(),

nemo_skills/dataset/wmt24pp/__init__.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,3 @@
1515

1616
METRICS_TYPE = "translation"
1717
GENERATION_ARGS = "++prompt_config=multilingual/segment-translation"
18-
JUDGE_PIPELINE_ARGS = {
19-
"source_key": "text",
20-
"reference_key": "translation",
21-
}

nemo_skills/dataset/wmt24pp/prepare.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,10 +23,10 @@
2323
def write_data_to_file(output_file, datasets, tgt_languages):
2424
with open(output_file, "wt", encoding="utf-8") as fout:
2525
for tgt_lang in tgt_languages:
26-
for src, tgt in zip(datasets[tgt_lang]["source"], datasets[tgt_lang]["target"], strict=True):
26+
for src, ref in zip(datasets[tgt_lang]["source"], datasets[tgt_lang]["target"], strict=True):
2727
json_dict = {
28-
"text": src,
29-
"translation": tgt,
28+
"source": src,
29+
"reference": ref,
3030
"source_language": "en",
3131
"target_language": tgt_lang,
3232
"source_lang_name": "English",

nemo_skills/evaluation/evaluator/audio.py

Lines changed: 20 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -558,28 +558,33 @@ def evaluate_asr(
558558
return result
559559

560560

561+
_BLEU_TOKENIZE_BY_LANG = {
562+
"ja": "ja-mecab",
563+
"zh": "zh",
564+
"cmn": "zh",
565+
"yue": "zh",
566+
"ko": "ko-mecab",
567+
}
568+
569+
570+
def resolve_bleu_tokenize(tgt_lang: str | None) -> str:
571+
"""Resolve sacrebleu tokenize from a target language code."""
572+
if not isinstance(tgt_lang, str):
573+
return "13a"
574+
lang_code = tgt_lang.split("_")[0]
575+
return _BLEU_TOKENIZE_BY_LANG.get(lang_code, "13a")
576+
577+
561578
def evaluate_translation(
562579
reference: str,
563580
hypothesis: str,
564581
tgt_lang: str | None = None,
565582
) -> dict[str, Any]:
566583
"""Evaluate translation: computes sentence-level BLEU score."""
584+
tokenize = resolve_bleu_tokenize(tgt_lang)
567585
try:
568586
import sacrebleu
569587

570-
tokenize = "13a"
571-
if isinstance(tgt_lang, str):
572-
lang_code = tgt_lang.split("_")[0]
573-
if lang_code in ["cmn", "yue"]:
574-
lang_code = "zh"
575-
576-
if lang_code == "ja":
577-
tokenize = "ja-mecab"
578-
elif lang_code == "zh":
579-
tokenize = "zh"
580-
elif lang_code == "ko":
581-
tokenize = "ko-mecab"
582-
583588
text = reference.strip()
584589
pred_text = hypothesis.strip()
585590
bleu = sacrebleu.sentence_bleu(pred_text, [text], tokenize=tokenize)
@@ -590,6 +595,7 @@ def evaluate_translation(
590595
"is_correct": bleu_score > 0.3,
591596
"text": text,
592597
"pred_text": pred_text,
598+
"bleu_tokenize": tokenize,
593599
}
594600
except Exception as e:
595601
return {
@@ -598,6 +604,7 @@ def evaluate_translation(
598604
"error": str(e),
599605
"text": reference.strip(),
600606
"pred_text": hypothesis.strip(),
607+
"bleu_tokenize": tokenize,
601608
}
602609

603610

nemo_skills/evaluation/evaluator/comet.py

Lines changed: 3 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -46,22 +46,11 @@ def load_comet_model(model_path: str):
4646
return model
4747

4848

49-
def _get_nested(sample: dict, key: str):
50-
if "." in key:
51-
value = sample
52-
for part in key.split("."):
53-
value = value[part]
54-
return value
55-
return sample[key]
56-
57-
5849
def process_file(
5950
input_file: Path,
6051
output_file: Path,
6152
comet_model,
6253
batch_size: int = 16,
63-
source_key: str = "source",
64-
reference_key: str = "reference",
6554
):
6655
"""Copy input file to output location and run xCOMET-XXL evaluation."""
6756
LOG.info(f"Processing {input_file} -> {output_file}")
@@ -87,9 +76,9 @@ def process_file(
8776
try:
8877
comet_list.append(
8978
{
90-
"src": _get_nested(sample, source_key),
91-
"mt": _get_nested(sample, "generation"),
92-
"ref": _get_nested(sample, reference_key),
79+
"src": sample["source"],
80+
"mt": sample["generation"],
81+
"ref": sample["reference"],
9382
}
9483
)
9584
except KeyError as e:
@@ -150,18 +139,6 @@ def main():
150139
default=1,
151140
help="Number of random seeds (for multiple seeds mode)",
152141
)
153-
parser.add_argument(
154-
"--source-key",
155-
type=str,
156-
default="source",
157-
help="Sample field (supports dotted paths) holding the source text passed as COMET 'src'",
158-
)
159-
parser.add_argument(
160-
"--reference-key",
161-
type=str,
162-
default="reference",
163-
help="Sample field (supports dotted paths) holding the reference translation passed as COMET 'ref'",
164-
)
165142
args = parser.parse_args()
166143

167144
output_dir = Path(args.output_dir)
@@ -193,8 +170,6 @@ def main():
193170
output_file,
194171
comet_model,
195172
args.batch_size,
196-
source_key=args.source_key,
197-
reference_key=args.reference_key,
198173
)
199174

200175
LOG.info("All files processed successfully")

0 commit comments

Comments
 (0)