feat: add advanced tuning, diarization, and LLM support to multi-process engine

This commit is contained in:
Adolfo Reyna
2026-03-15 19:53:36 -04:00
parent c474184169
commit ca07d96ff0
5 changed files with 188 additions and 113 deletions
+43 -40
View File
@@ -22,6 +22,7 @@ import re
import argparse
import multiprocessing
import time
import requests
TARGET_LANGS = {
"es": "Helsinki-NLP/opus-mt-en-es",
@@ -33,98 +34,100 @@ 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 = sentence.split()
word_chunk = ""
for word in words:
candidate = f"{word_chunk} {word}".strip()
if len(candidate) <= max_chars:
word_chunk = candidate
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
if candidate and len(candidate) <= max_chars: current = candidate
else:
if current: chunks.append(current)
current = sentence
if current: chunks.append(current)
return chunks
def call_ollama_generate(args, prompt):
payload = {"model": args.post_correct_model, "prompt": prompt, "stream": False}
try:
resp = requests.post("http://127.0.0.1:11434/api/generate", json=payload, timeout=10)
return resp.json().get("response", "").strip()
except: return ""
def structure_paragraph_with_llm(text, args):
prompt = f"Format the following transcribed segments into a coherent paragraph. Output ONLY the paragraph:\n{text}"
return call_ollama_generate(args, prompt) or text
def post_correct_with_llm(text, args):
prompt = f"Correct grammar and typos in this caption. Output ONLY the corrected text:\n{text}"
return call_ollama_generate(args, prompt) or text
def run_translation(in_queue, out_queue, args):
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...")
tokenizer = MarianTokenizer.from_pretrained(model_id)
model = MarianMTModel.from_pretrained(model_id).to(device)
translation_engines[lang_key] = (model, tokenizer)
translation_engines[lang_key] = (
MarianMTModel.from_pretrained(model_id).to(device),
MarianTokenizer.from_pretrained(model_id)
)
paragraph_buffer = []
print("[Translate] Ready.")
while True:
try:
item = in_queue.get()
if item is None: break
# Draft support
if "draft" in item:
out_queue.put({"draft": item["draft"]})
continue
text = item.get("original", "").strip()
detected_lang = item.get("detected_lang", "en")
# Simple paragraph logic:
# If the segment ends with sentence-terminal punctuation, flush the paragraph.
if args.post_correct and args.post_correct_llm:
text = post_correct_with_llm(text, args)
paragraph_buffer.append(text)
if text.endswith((".", "?", "!")):
# Paragraph logic
if not args.llm_paragraph or text.endswith((".", "?", "!")):
full_text = " ".join(paragraph_buffer)
payload = {"original": full_text, "en": full_text if detected_lang == "en" else None}
# If detected language is not English, we'd normally bridge to English first.
# For now, let's assume direct translation for simplicity or bridge if needed.
# (Refining bridge logic can come later)
if args.llm_paragraph:
full_text = structure_paragraph_with_llm(full_text, args)
payload = {"original": full_text, "ts": item.get("ts", time.time())}
if args.en: payload["en"] = full_text
for lang_key, (model, tokenizer) in translation_engines.items():
chunks = split_text_for_translation(full_text)
chunks = split_text_for_translation(full_text, 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=150)
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())
translated_text = " ".join(part for part in translated_parts if part).strip()
payload[lang_key] = translated_text
print(f"[Translate] {lang_key.upper()}: {translated_text}")
payload[lang_key] = " ".join(translated_parts)
out_queue.put(payload)
paragraph_buffer = []
except Exception as e:
print(f"[Translate] Error: {e}")
except Exception as e: print(f"[Translate] Error: {e}")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("-es", action="store_true")
parser.add_argument("-fr", action="store_true")
parser.add_argument("-ar", action="store_true")
args = parser.parse_args()
# Dummy queues for testing
in_q = multiprocessing.Queue()
out_q = multiprocessing.Queue()
run_translation(in_q, out_q, args)
import multiprocessing
run_translation(multiprocessing.Queue(), multiprocessing.Queue(), argparse.Namespace())