Add local MarianMT translation service
This commit is contained in:
@@ -0,0 +1,118 @@
|
||||
import json
|
||||
import os
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
|
||||
from langdetect import DetectorFactory, LangDetectException, detect
|
||||
from transformers import MarianMTModel, MarianTokenizer
|
||||
|
||||
DetectorFactory.seed = 0
|
||||
|
||||
HOST = os.getenv("MARIAN_HOST", "127.0.0.1")
|
||||
PORT = int(os.getenv("MARIAN_PORT", "8000"))
|
||||
MAX_INPUT_LENGTH = int(os.getenv("MARIAN_MAX_INPUT_LENGTH", "1000"))
|
||||
DEFAULT_SOURCE_LANGUAGE = os.getenv("MARIAN_DEFAULT_SOURCE_LANGUAGE", "en")
|
||||
SUPPORTED_LANGUAGES = {"en", "es", "fr", "da", "ar"}
|
||||
MODEL_CACHE = {}
|
||||
|
||||
|
||||
def normalize_language(value):
|
||||
language = str(value or "").strip().lower().split(",")[0].split("-")[0]
|
||||
return language
|
||||
|
||||
|
||||
def detect_source_language(text):
|
||||
try:
|
||||
detected = normalize_language(detect(text))
|
||||
if detected in SUPPORTED_LANGUAGES:
|
||||
return detected
|
||||
except LangDetectException:
|
||||
pass
|
||||
return DEFAULT_SOURCE_LANGUAGE
|
||||
|
||||
|
||||
def get_model(source, target):
|
||||
model_name = f"Helsinki-NLP/opus-mt-{source}-{target}"
|
||||
if model_name not in MODEL_CACHE:
|
||||
MODEL_CACHE[model_name] = (
|
||||
MarianTokenizer.from_pretrained(model_name),
|
||||
MarianMTModel.from_pretrained(model_name),
|
||||
)
|
||||
return model_name, MODEL_CACHE[model_name]
|
||||
|
||||
|
||||
def translate_once(text, source, target):
|
||||
model_name, (tokenizer, model) = get_model(source, target)
|
||||
encoded = tokenizer([text], return_tensors="pt", truncation=True)
|
||||
generated = model.generate(**encoded)
|
||||
return tokenizer.batch_decode(generated, skip_special_tokens=True)[0], model_name
|
||||
|
||||
|
||||
def translate(text, source, target):
|
||||
if source == "auto":
|
||||
source = detect_source_language(text)
|
||||
if source not in SUPPORTED_LANGUAGES or target not in SUPPORTED_LANGUAGES:
|
||||
raise ValueError("Only en, es, fr, da, and ar are supported")
|
||||
if source == target:
|
||||
return text, source, "none"
|
||||
if source == "en" or target == "en":
|
||||
translated, model_name = translate_once(text, source, target)
|
||||
return translated, source, model_name
|
||||
|
||||
english, first_model = translate_once(text, source, "en")
|
||||
translated, second_model = translate_once(english, "en", target)
|
||||
return translated, source, f"{first_model},{second_model}"
|
||||
|
||||
|
||||
class TranslationHandler(BaseHTTPRequestHandler):
|
||||
def send_json(self, status, body):
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(payload)))
|
||||
self.end_headers()
|
||||
self.wfile.write(payload)
|
||||
|
||||
def do_GET(self):
|
||||
if self.path != "/health":
|
||||
self.send_json(404, {"status": "not found"})
|
||||
return
|
||||
self.send_json(200, {"status": "ok", "provider": "marianmt", "loadedModels": list(MODEL_CACHE)})
|
||||
|
||||
def do_POST(self):
|
||||
if self.path != "/translate":
|
||||
self.send_json(404, {"status": "not found"})
|
||||
return
|
||||
try:
|
||||
content_length = int(self.headers.get("Content-Length", "0"))
|
||||
body = json.loads(self.rfile.read(content_length).decode("utf-8"))
|
||||
text = str(body.get("text") or "").strip()
|
||||
source = normalize_language(body.get("sourceLang")) or "auto"
|
||||
target = normalize_language(body.get("targetLang"))
|
||||
if not text or not target:
|
||||
self.send_json(400, {"status": "text and targetLang are required"})
|
||||
return
|
||||
if len(text) > MAX_INPUT_LENGTH:
|
||||
self.send_json(400, {"status": f"text exceeds {MAX_INPUT_LENGTH} characters"})
|
||||
return
|
||||
translated, detected_source, model_name = translate(text, source, target)
|
||||
self.send_json(200, {
|
||||
"status": "ok",
|
||||
"translatedText": translated,
|
||||
"sourceLang": detected_source,
|
||||
"targetLang": target,
|
||||
"provider": "marianmt",
|
||||
"model": model_name,
|
||||
})
|
||||
except (ValueError, json.JSONDecodeError) as error:
|
||||
self.send_json(400, {"status": str(error)})
|
||||
except Exception as error:
|
||||
print(f"Translation failed: {error}", flush=True)
|
||||
self.send_json(502, {"status": "Translation failed"})
|
||||
|
||||
def log_message(self, format_string, *args):
|
||||
print(f"[marianmt] {self.address_string()} {format_string % args}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(f"MarianMT translation service listening on {HOST}:{PORT}", flush=True)
|
||||
ThreadingHTTPServer((HOST, PORT), TranslationHandler).serve_forever()
|
||||
Reference in New Issue
Block a user