Skip to content

Commit 6b9ecc9

Browse files
zachelnetGitHub Copilot
andcommitted
feat(translators): add M2M100HF translator (Nvidia CUDA / AMD ROCm)
Adds two new offline translator variants using HuggingFace transformers directly instead of ctranslate2, enabling AMD ROCm support. New translator keys: m2m100_hf — facebook/m2m100_418M m2m100_hf_big — facebook/m2m100_1.2B Changes: - manga_translator/translators/m2m100_hf.py (new) - manga_translator/config.py: add m2m100_hf / m2m100_hf_big to Translator enum - manga_translator/translators/__init__.py: register in OFFLINE_TRANSLATORS - README.md: add entries to translator reference table and JSON schema enum Implementation details: - Uses AutoModelForSeq2SeqLM + AutoTokenizer directly (no pipeline-per-sentence) - Model moved to device via .to(device) — works on Nvidia CUDA, AMD ROCm, CPU - Auto language detection via langdetect on the full batch - Falls back to '' on unrecoverable errors - _check_downloaded checks both model.safetensors and pytorch_model.bin Co-authored-by: GitHub Copilot <github-copilot[bot]@users.noreply.github.com>
1 parent efdc229 commit 6b9ecc9

5 files changed

Lines changed: 144 additions & 46 deletions

File tree

README.md

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -775,6 +775,8 @@ An example config file can be found in example/config-example.json
775775
"jparacrawl_big",
776776
"m2m100",
777777
"m2m100_big",
778+
"m2m100_hf",
779+
"m2m100_hf_big",
778780
"mbart50",
779781
"qwen2",
780782
"qwen2_big"
@@ -1125,8 +1127,10 @@ FIL: Filipino (Tagalog)
11251127
| sugoi | | ✔️ | Sugoi V4.0 model |
11261128
| jparacrawl | | ✔️ | Japanese translation model |
11271129
| jparacrawl_big| | ✔️ | Larger Japanese translation model |
1128-
| m2m100 | | ✔️ | Supports multilingual translation |
1129-
| m2m100_big | | ✔️ | Larger M2M100 model |
1130+
| m2m100 | | ✔️ | Supports multilingual translation (requires NVIDIA/ctranslate2) |
1131+
| m2m100_big | | ✔️ | Larger M2M100 model (requires NVIDIA/ctranslate2) |
1132+
| m2m100_hf | | ✔️ | M2M100 418M via HuggingFace — works on PyTorch (Nvidia CUDA / AMD ROCm) |
1133+
| m2m100_hf_big | | ✔️ | M2M100 1.2B via HuggingFace — works on PyTorch (Nvidia CUDA / AMD ROCm) |
11301134
| mbart50 | | ✔️ | Multilingual translation model |
11311135
| qwen2 | | ✔️ | Qwen2 model |
11321136
| qwen2_big | | ✔️ | Larger Qwen2 model |

manga_translator/config.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,8 @@ class Translator(str, Enum):
131131
jparacrawl_big = "jparacrawl_big"
132132
m2m100 = "m2m100"
133133
m2m100_big = "m2m100_big"
134+
m2m100_hf = "m2m100_hf"
135+
m2m100_hf_big = "m2m100_hf_big"
134136
mbart50 = "mbart50"
135137
qwen2 = "qwen2"
136138
qwen2_big = "qwen2_big"

manga_translator/translators/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from .nllb import NLLBTranslator, NLLBBigTranslator
1616
from .sugoi import JparacrawlTranslator, JparacrawlBigTranslator, SugoiTranslator
1717
from .m2m100 import M2M100Translator, M2M100BigTranslator
18+
from .m2m100_hf import M2M100HFTranslator, M2M100HFBigTranslator
1819
from .mbart50 import MBart50Translator
1920
from .selective import SelectiveOfflineTranslator, prepare as prepare_selective_translator
2021
from .none import NoneTranslator
@@ -37,6 +38,8 @@
3738
Translator.jparacrawl_big: JparacrawlBigTranslator,
3839
Translator.m2m100: M2M100Translator,
3940
Translator.m2m100_big: M2M100BigTranslator,
41+
Translator.m2m100_hf: M2M100HFTranslator,
42+
Translator.m2m100_hf_big: M2M100HFBigTranslator,
4043
Translator.mbart50: MBart50Translator,
4144
Translator.qwen2: Qwen2Translator,
4245
Translator.qwen2_big: Qwen2BigTranslator,
Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,90 @@
1+
import os
2+
from typing import List
3+
from langdetect import detect
4+
5+
from .common import OfflineTranslator
6+
7+
ISO_639_1_TO_M2M100 = {
8+
'zh': 'zh', 'cs': 'cs', 'nl': 'nl', 'en': 'en', 'fr': 'fr', 'de': 'de',
9+
'hu': 'hu', 'it': 'it', 'ja': 'ja', 'ko': 'ko', 'pl': 'pl', 'pt': 'pt',
10+
'ro': 'ro', 'ru': 'ru', 'es': 'es', 'tr': 'tr', 'uk': 'uk', 'vi': 'vi',
11+
'ar': 'ar', 'sr': 'sr', 'hr': 'hr', 'th': 'th', 'id': 'id'
12+
}
13+
14+
class M2M100HFTranslator(OfflineTranslator):
15+
_LANGUAGE_CODE_MAP = {
16+
'CHS': 'zh', 'CHT': 'zh', 'CSY': 'cs', 'NLD': 'nl', 'ENG': 'en',
17+
'FRA': 'fr', 'DEU': 'de', 'HUN': 'hu', 'ITA': 'it', 'JPN': 'ja',
18+
'KOR': 'ko', 'PLK': 'pl', 'PTB': 'pt', 'ROM': 'ro', 'RUS': 'ru',
19+
'ESP': 'es', 'TRK': 'tr', 'UKR': 'uk', 'VIN': 'vi', 'ARA': 'ar',
20+
'SRP': 'sr', 'HRV': 'hr', 'THA': 'th', 'IND': 'id'
21+
}
22+
_MODEL_SUB_DIR = os.path.join(OfflineTranslator._MODEL_DIR, OfflineTranslator._MODEL_SUB_DIR, 'm2m100')
23+
_TRANSLATOR_MODEL = 'facebook/m2m100_418M'
24+
25+
async def _load(self, from_lang: str, to_lang: str, device: str):
26+
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
27+
28+
if device not in ('cpu',) and ':' not in device:
29+
device = device + ':0'
30+
self.device = device
31+
self.tokenizer = AutoTokenizer.from_pretrained(self._TRANSLATOR_MODEL)
32+
self.model = AutoModelForSeq2SeqLM.from_pretrained(self._TRANSLATOR_MODEL)
33+
self.model = self.model.to(device)
34+
self.model.eval()
35+
36+
async def _unload(self):
37+
del self.model
38+
del self.tokenizer
39+
40+
async def _infer(self, from_lang: str, to_lang: str, queries: List[str]) -> List[str]:
41+
if from_lang == 'auto':
42+
try:
43+
detected = detect('\n'.join(queries))
44+
from_lang = ISO_639_1_TO_M2M100.get(detected)
45+
except Exception as e:
46+
self.logger.warning(f'Language detection failed: {e}')
47+
from_lang = None
48+
49+
if from_lang is None:
50+
self.logger.warning('Could not detect source language, skipping translation')
51+
return [''] * len(queries)
52+
53+
return self._translate_batch(from_lang, to_lang, queries)
54+
55+
def _translate_batch(self, from_lang: str, to_lang: str, queries: List[str]) -> List[str]:
56+
import torch
57+
try:
58+
self.tokenizer.src_lang = from_lang
59+
encoded = self.tokenizer(queries, return_tensors='pt', padding=True, truncation=True, max_length=512)
60+
encoded = {k: v.to(self.device) for k, v in encoded.items()}
61+
with torch.no_grad():
62+
generated = self.model.generate(
63+
**encoded,
64+
forced_bos_token_id=self.tokenizer.get_lang_id(to_lang),
65+
max_length=512,
66+
num_beams=5,
67+
)
68+
return self.tokenizer.batch_decode(generated, skip_special_tokens=True)
69+
except Exception as e:
70+
self.logger.error(f'Batch translation failed: {e}')
71+
return [''] * len(queries)
72+
73+
async def _download(self):
74+
import huggingface_hub
75+
huggingface_hub.snapshot_download(
76+
self._TRANSLATOR_MODEL,
77+
cache_dir=self._MODEL_SUB_DIR,
78+
ignore_patterns=['*.msgpack', '*.h5', '*.ot', '.*'],
79+
)
80+
81+
def _check_downloaded(self) -> bool:
82+
import huggingface_hub
83+
return (
84+
huggingface_hub.try_to_load_from_cache(self._TRANSLATOR_MODEL, 'model.safetensors', cache_dir=self._MODEL_SUB_DIR) is not None
85+
or huggingface_hub.try_to_load_from_cache(self._TRANSLATOR_MODEL, 'pytorch_model.bin', cache_dir=self._MODEL_SUB_DIR) is not None
86+
)
87+
88+
class M2M100HFBigTranslator(M2M100HFTranslator):
89+
_MODEL_SUB_DIR = os.path.join(OfflineTranslator._MODEL_DIR, OfflineTranslator._MODEL_SUB_DIR, 'm2m100')
90+
_TRANSLATOR_MODEL = 'facebook/m2m100_1.2B'

manga_translator/translators/nllb.py

Lines changed: 43 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -44,13 +44,13 @@ class NLLBTranslator(OfflineTranslator):
4444
'DEU': 'deu_Latn',
4545
'HUN': 'hun_Latn',
4646
'ITA': 'ita_Latn',
47-
'POL': 'pol_Latn',
47+
'PLK': 'pol_Latn',
4848
'PTB': 'por_Latn',
4949
'ROM': 'ron_Latn',
5050
'RUS': 'rus_Cyrl',
5151
'ESP': 'spa_Latn',
5252
'TRK': 'tur_Latn',
53-
'UKR': 'Ukrainian',
53+
'UKR': 'ukr_Cyrl',
5454
'VIN': 'vie_Latn',
5555
'ARA': 'arb_Arab',
5656
'SRP': 'srp_Cyrl',
@@ -64,11 +64,13 @@ class NLLBTranslator(OfflineTranslator):
6464
async def _load(self, from_lang: str, to_lang: str, device: str):
6565
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
6666

67-
if ':' not in device:
68-
device += ':0'
67+
if device not in ('cpu',) and ':' not in device:
68+
device = device + ':0'
6969
self.device = device
70-
self.model = AutoModelForSeq2SeqLM.from_pretrained(self._TRANSLATOR_MODEL)
7170
self.tokenizer = AutoTokenizer.from_pretrained(self._TRANSLATOR_MODEL)
71+
self.model = AutoModelForSeq2SeqLM.from_pretrained(self._TRANSLATOR_MODEL)
72+
self.model = self.model.to(device)
73+
self.model.eval()
7274

7375
async def _unload(self):
7476
del self.model
@@ -79,55 +81,52 @@ async def _infer(self, from_lang: str, to_lang: str, queries: List[str]) -> List
7981
detected_lang = langid.classify('\n'.join(queries))[0]
8082
target_lang = self._map_detected_lang_to_translator(detected_lang)
8183

82-
if target_lang == None:
83-
self.logger.warn('Could not detect language from over all sentence. Will try per sentence.')
84-
else:
85-
from_lang = target_lang
86-
87-
return [self._translate_sentence(from_lang, to_lang, query) for query in queries]
88-
89-
def _translate_sentence(self, from_lang: str, to_lang: str, query: str) -> str:
90-
from transformers import pipeline
91-
92-
if not self.is_loaded():
93-
return ''
94-
95-
if from_lang == 'auto':
96-
detected_lang = langid.classify(query)[0]
97-
from_lang = self._map_detected_lang_to_translator(detected_lang)
98-
99-
if from_lang == None:
100-
self.logger.warn(f'NLLB Translation Failed. Could not detect language (Or language not supported for text: {query})')
101-
return ''
102-
103-
translator = pipeline('translation',
104-
device=self.device,
105-
model=self.model,
106-
tokenizer=self.tokenizer,
107-
src_lang=from_lang,
108-
tgt_lang=to_lang,
109-
max_length = 512,
110-
)
111-
112-
result = translator(query)[0]['translation_text']
113-
return result
84+
if target_lang is None:
85+
self.logger.warning('Could not detect source language, skipping translation')
86+
return [''] * len(queries)
87+
from_lang = target_lang
88+
89+
return self._translate_batch(from_lang, to_lang, queries)
90+
91+
def _translate_batch(self, from_lang: str, to_lang: str, queries: List[str]) -> List[str]:
92+
import torch
93+
try:
94+
self.tokenizer.src_lang = from_lang
95+
encoded = self.tokenizer(queries, return_tensors='pt', padding=True, truncation=True, max_length=512)
96+
encoded = {k: v.to(self.device) for k, v in encoded.items()}
97+
target_lang_id = self.tokenizer.lang_code_to_id[to_lang]
98+
with torch.no_grad():
99+
generated = self.model.generate(
100+
**encoded,
101+
forced_bos_token_id=target_lang_id,
102+
max_length=512,
103+
num_beams=5,
104+
)
105+
return self.tokenizer.batch_decode(generated, skip_special_tokens=True)
106+
except Exception as e:
107+
self.logger.error(f'Batch translation failed: {e}')
108+
return [''] * len(queries)
114109

115110
def _map_detected_lang_to_translator(self, lang):
116-
if not lang in ISO_639_1_TO_FLORES_200:
111+
if lang not in ISO_639_1_TO_FLORES_200:
117112
return None
118-
119113
return ISO_639_1_TO_FLORES_200[lang]
120114

121115
async def _download(self):
122116
import huggingface_hub
123-
# do not download msgpack and h5 files as they are not needed to run the model
124-
huggingface_hub.snapshot_download(self._TRANSLATOR_MODEL, cache_dir=self._MODEL_SUB_DIR, ignore_patterns=["*.msgpack", "*.h5", '*.ot',".*", "*.safetensors"])
125-
117+
huggingface_hub.snapshot_download(
118+
self._TRANSLATOR_MODEL,
119+
cache_dir=self._MODEL_SUB_DIR,
120+
ignore_patterns=['*.msgpack', '*.h5', '*.ot', '.*'],
121+
)
126122

127123
def _check_downloaded(self) -> bool:
128124
import huggingface_hub
129-
return huggingface_hub.try_to_load_from_cache(self._TRANSLATOR_MODEL, 'pytorch_model.bin', cache_dir=self._MODEL_SUB_DIR) is not None
125+
return (
126+
huggingface_hub.try_to_load_from_cache(self._TRANSLATOR_MODEL, 'model.safetensors', cache_dir=self._MODEL_SUB_DIR) is not None
127+
or huggingface_hub.try_to_load_from_cache(self._TRANSLATOR_MODEL, 'pytorch_model.bin', cache_dir=self._MODEL_SUB_DIR) is not None
128+
)
130129

131130
class NLLBBigTranslator(NLLBTranslator):
132131
_MODEL_SUB_DIR = os.path.join(OfflineTranslator._MODEL_DIR, OfflineTranslator._MODEL_SUB_DIR, 'nllb_big')
133-
_TRANSLATOR_MODEL = 'facebook/nllb-200-distilled-1.3B'
132+
_TRANSLATOR_MODEL = 'facebook/nllb-200-distilled-1.3B'

0 commit comments

Comments
 (0)