Skip to content

Commit 6d6f6bf

Browse files
authored
Add locale to target_lang_name in wmt24pp prepare.py (#1980)
Adds locale to target_lang_name in wmt24pp prepare.py, consistent with current nemo-skills implementation (https://github.com/NVIDIA-NeMo/Skills/blob/7b3fb04e131b3a31ec00d28b65ff7eea776b5249/nemo_skills/dataset/wmt24pp/prepare.py#L33). Also fixes issue of non-default languages crashing because they were not in _LANG_DISPLAY_NAMES. Added checks that user supplied languages are actually in wmt24pp. Languages codes and names are hardcoded to avoid network calls and dependence on additional packages - this should be fine as wmt24pp will not change languages. Signed-off-by: Brian Thompson <3534106+thompsonb@users.noreply.github.com> Co-authored-by: Brian Thompson <3534106+thompsonb@users.noreply.github.com>
1 parent 4c8e10d commit 6d6f6bf

1 file changed

Lines changed: 78 additions & 9 deletions

File tree

benchmarks/wmt24pp/prepare.py

Lines changed: 78 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,10 @@
2929
``wmt_translation`` resource server: ``text``, ``translation``,
3030
``source_language``, ``target_language``, ``source_lang_name``,
3131
``target_lang_name``.
32+
33+
Note that we need to somehow specify the dialect/locale we want the
34+
model to generate in. Currently, this is done by setting
35+
``target_lang_name`` to <language> (<country>) (e.g. "Spanish (Mexico)")
3236
"""
3337

3438
import json
@@ -54,14 +58,71 @@
5458
# stable is what makes the interleaved JSONL byte-comparable.
5559
DEFAULT_TARGET_LANGUAGES = ["de_DE", "es_MX", "fr_FR", "it_IT", "ja_JP"]
5660

57-
# Display names for the targets above. Matches NeMo-Skills'
58-
# `langcodes.Language(tgt[:2]).display_name()` output for these codes.
59-
_LANG_DISPLAY_NAMES = {
60-
"de": "German",
61-
"es": "Spanish",
62-
"fr": "French",
63-
"it": "Italian",
64-
"ja": "Japanese",
61+
# Hardcoding to avoid both dependency on https://github.com/rspeer/langcodes
62+
# and network call in prepare.py
63+
# import langcodes #uv pip install langcodes[data]
64+
# from datasets import get_dataset_config_names #uv pip install datasets
65+
# lang_pairs = get_dataset_config_names("google/wmt24pp")
66+
# _WMT24PP_LANG_MAP = {}
67+
# for lang_pair in lang_pairs:
68+
# _, tgt = lang_pair.split('-')
69+
# _WMT24PP_LANG_MAP[tgt] = langcodes.Language.get(tgt).display_name()
70+
_WMT24PP_LANG_MAP = {
71+
"ar_EG": "Arabic (Egypt)",
72+
"ar_SA": "Arabic (Saudi Arabia)",
73+
"bg_BG": "Bulgarian (Bulgaria)",
74+
"bn_IN": "Bangla (India)",
75+
"ca_ES": "Catalan (Spain)",
76+
"cs_CZ": "Czech (Czechia)",
77+
"da_DK": "Danish (Denmark)",
78+
"de_DE": "German (Germany)",
79+
"el_GR": "Greek (Greece)",
80+
"es_MX": "Spanish (Mexico)",
81+
"et_EE": "Estonian (Estonia)",
82+
"fa_IR": "Persian (Iran)",
83+
"fi_FI": "Finnish (Finland)",
84+
"fil_PH": "Filipino (Philippines)",
85+
"fr_CA": "French (Canada)",
86+
"fr_FR": "French (France)",
87+
"gu_IN": "Gujarati (India)",
88+
"he_IL": "Hebrew (Israel)",
89+
"hi_IN": "Hindi (India)",
90+
"hr_HR": "Croatian (Croatia)",
91+
"hu_HU": "Hungarian (Hungary)",
92+
"id_ID": "Indonesian (Indonesia)",
93+
"is_IS": "Icelandic (Iceland)",
94+
"it_IT": "Italian (Italy)",
95+
"ja_JP": "Japanese (Japan)",
96+
"kn_IN": "Kannada (India)",
97+
"ko_KR": "Korean (South Korea)",
98+
"lt_LT": "Lithuanian (Lithuania)",
99+
"lv_LV": "Latvian (Latvia)",
100+
"ml_IN": "Malayalam (India)",
101+
"mr_IN": "Marathi (India)",
102+
"nl_NL": "Dutch (Netherlands)",
103+
"no_NO": "Norwegian (Norway)",
104+
"pa_IN": "Punjabi (India)",
105+
"pl_PL": "Polish (Poland)",
106+
"pt_BR": "Portuguese (Brazil)",
107+
"pt_PT": "Portuguese (Portugal)",
108+
"ro_RO": "Romanian (Romania)",
109+
"ru_RU": "Russian (Russia)",
110+
"sk_SK": "Slovak (Slovakia)",
111+
"sl_SI": "Slovenian (Slovenia)",
112+
"sr_RS": "Serbian (Serbia)",
113+
"sv_SE": "Swedish (Sweden)",
114+
"sw_KE": "Swahili (Kenya)",
115+
"sw_TZ": "Swahili (Tanzania)",
116+
"ta_IN": "Tamil (India)",
117+
"te_IN": "Telugu (India)",
118+
"th_TH": "Thai (Thailand)",
119+
"tr_TR": "Turkish (Türkiye)",
120+
"uk_UA": "Ukrainian (Ukraine)",
121+
"ur_PK": "Urdu (Pakistan)",
122+
"vi_VN": "Vietnamese (Vietnam)",
123+
"zh_CN": "Chinese (China)",
124+
"zh_TW": "Chinese (Taiwan)",
125+
"zu_ZA": "Zulu (South Africa)",
65126
}
66127

67128

@@ -104,6 +165,14 @@ def prepare(target_languages: list[str] | None = None, prefetch_comet: bool = Tr
104165
"""
105166
if target_languages is None:
106167
target_languages = DEFAULT_TARGET_LANGUAGES
168+
else:
169+
# check user passed in langs
170+
unknown = [lang for lang in target_languages if lang not in _WMT24PP_LANG_MAP]
171+
if unknown:
172+
raise ValueError(
173+
f"requested target languages [{','.join(unknown)}] are not in wmt24pp. "
174+
f"Available languages: [{','.join(list(_WMT24PP_LANG_MAP))}]"
175+
)
107176

108177
DATA_DIR.mkdir(parents=True, exist_ok=True)
109178

@@ -126,7 +195,7 @@ def prepare(target_languages: list[str] | None = None, prefetch_comet: bool = Tr
126195
"source_language": "en",
127196
"target_language": tgt_lang,
128197
"source_lang_name": "English",
129-
"target_lang_name": _LANG_DISPLAY_NAMES[tgt_lang[:2]],
198+
"target_lang_name": _WMT24PP_LANG_MAP[tgt_lang],
130199
}
131200
fout.write(json.dumps(row) + "\n")
132201
count += 1

0 commit comments

Comments
 (0)