Files
whisper-translation/engine_translate.py
T
2026-03-17 00:06:47 -04:00

110 lines
4.6 KiB
Python

import argparse
import multiprocessing
import time
import re
def run_translation(in_queue, out_queue, args):
import sys
from unittest.mock import MagicMock
# Comprehensive workaround for missing _lzma in some Python builds
try:
import lzma
except ImportError:
mock_lzma = MagicMock()
mock_lzma.FORMAT_XZ, mock_lzma.FORMAT_ALONE, mock_lzma.FORMAT_RAW = 1, 2, 3
mock_lzma.CHECK_NONE, mock_lzma.CHECK_CRC32, mock_lzma.CHECK_CRC64, mock_lzma.CHECK_SHA256 = 0, 1, 4, 10
sys.modules["_lzma"] = MagicMock()
sys.modules["lzma"] = mock_lzma
import torch
from transformers import MarianMTModel, MarianTokenizer
TARGET_LANGS = {
"es": "Helsinki-NLP/opus-mt-en-es",
"fr": "Helsinki-NLP/opus-mt-en-fr",
"ar": "Helsinki-NLP/opus-mt-en-ar"
}
def split_text_for_translation(text, max_chars=250):
normalized = " ".join(text.split()).strip()
if not normalized or len(normalized) <= max_chars: return [normalized] if normalized else []
chunks, current = [], ""
sentences = [s for s in re.split(r"(?<=[.!?])\s+", normalized) if s]
for sentence in sentences:
if len(sentence) > max_chars:
words, word_chunk = sentence.split(), ""
for word in words:
candidate = f"{word_chunk} {word}".strip()
if len(candidate) <= max_chars: word_chunk = candidate
else:
if word_chunk: chunks.append(word_chunk)
word_chunk = word
if word_chunk: chunks.append(word_chunk)
continue
candidate = f"{current} {sentence}".strip()
if candidate and len(candidate) <= max_chars: current = candidate
else:
if current: chunks.append(current)
current = sentence
if current: chunks.append(current)
return chunks
device = "mps" if torch.backends.mps.is_available() else "cpu"
print(f"[Translate] Loading translation models on {device}...")
translation_engines = {}
for lang_key, model_id in TARGET_LANGS.items():
if getattr(args, lang_key, False):
print(f"[Translate] Loading {lang_key} model...")
translation_engines[lang_key] = (
MarianMTModel.from_pretrained(model_id).to(device),
MarianTokenizer.from_pretrained(model_id)
)
while True:
try:
item = in_queue.get()
if item is None: break
if "draft" in item:
out_queue.put(item)
continue
raw_text, corrected_text, paragraph_text = item.get("raw"), item.get("corrected"), item.get("paragraph")
payload = {
"original": raw_text,
"corrected": corrected_text,
"paragraph": paragraph_text,
"en_bridge": item.get("en_bridge"),
"used_llm_line": item.get("used_llm_line", False),
"used_llm_paragraph": item.get("used_llm_paragraph", False),
"speaker": item.get("speaker"),
"ts": item.get("ts", time.time())
}
if args.only_translate_llm:
text_to_translate = paragraph_text
if args.en and paragraph_text:
payload["en"] = paragraph_text
else:
text_to_translate = paragraph_text or corrected_text or raw_text
if args.en:
payload["en"] = text_to_translate
if text_to_translate and translation_engines:
for lang_key, (model, tokenizer) in translation_engines.items():
chunks = split_text_for_translation(text_to_translate, args.mt_max_chars)
translated_parts = []
for chunk in chunks:
inputs = tokenizer(chunk, return_tensors="pt", padding=True).to(device)
with torch.no_grad():
translated_tokens = model.generate(**inputs, max_new_tokens=args.mt_max_new_tokens, num_beams=args.mt_num_beams)
translated_parts.append(tokenizer.decode(translated_tokens[0], skip_special_tokens=True).strip())
payload[lang_key] = " ".join(translated_parts)
# ALWAYS put in queue, even if no translations were done
out_queue.put(payload)
except Exception as e: print(f"[Translate] Error: {e}")
if __name__ == "__main__":
run_translation(multiprocessing.Queue(), multiprocessing.Queue(), argparse.Namespace())