From ffeb996d7e632d28c4afd182346f4b4a891df89c Mon Sep 17 00:00:00 2001 From: Adolfo Reyna Date: Mon, 16 Mar 2026 21:11:51 -0400 Subject: [PATCH] Checkpoint current transcription pipeline state --- config.json | 1 + engine_distribute.py | 42 ++++++++++++++++++++++++++++++------------ engine_translate.py | 38 ++++++++++++++++++++++---------------- main_v2.py | 8 ++++++++ transcribe.py | 10 +++++++--- 5 files changed, 68 insertions(+), 31 deletions(-) diff --git a/config.json b/config.json index 5c6d39b..5f17899 100644 --- a/config.json +++ b/config.json @@ -19,6 +19,7 @@ "post_correct_llm": false, "post_correct_model": "qwen:2b", "llm_paragraph": false, + "only_translate_llm": false, "temperature_fallback": [0.0, 0.2, 0.4, 0.6, 0.8, 1.0], "logprob_threshold": -0.8, "compression_threshold": 2.2, diff --git a/engine_distribute.py b/engine_distribute.py index 8329f2a..f585903 100644 --- a/engine_distribute.py +++ b/engine_distribute.py @@ -57,18 +57,36 @@ def run_distribution(in_queue, args): log_debug(lang_msg) if args.ingest: - delay, max_delay, success = 1, 15, False - while not success: - try: - response = requests.post(INGEST_URL, json=payload, timeout=5) - if response.status_code == 200: success = True - else: print(f"[Distribute Error] {response.status_code}. Retrying...") - except Exception as e: - print(f"[Distribute Error] {e}. Retrying...") - - if not success: - time.sleep(delay) - delay = min(delay * 2, max_delay) + # Flat JSON Schema: Keep only what's needed for the server + ingest_payload = { + "original": payload.get("original"), + "speaker": payload.get("speaker"), + "ts": payload.get("ts") + } + has_any_translation = False + for lang in ["es", "fr", "ar", "en"]: + if payload.get(lang): + ingest_payload[lang] = payload[lang] + has_any_translation = True + + # If only-translate-llm is on, ONLY send when we have a translation (paragraph-level) + should_send = True + if getattr(args, "only_translate_llm", False) and not has_any_translation: + should_send = False + + if should_send: + delay, max_delay, success = 1, 15, False + while not success: + try: + response = requests.post(INGEST_URL, json=ingest_payload, timeout=5) + if response.status_code == 200: success = True + else: print(f"[Distribute Error] {response.status_code}. Retrying...") + except Exception as e: + print(f"[Distribute Error] {e}. Retrying...") + + if not success: + time.sleep(delay) + delay = min(delay * 2, max_delay) except Exception as e: print(f"[Distribute] Error: {e}") diff --git a/engine_translate.py b/engine_translate.py index 7d70ab7..0cb1fc4 100644 --- a/engine_translate.py +++ b/engine_translate.py @@ -77,23 +77,29 @@ def run_translation(in_queue, out_queue, args): "speaker": item.get("speaker"), "ts": item.get("ts", time.time()) } - if args.en: payload["en"] = corrected_text or raw_text + + 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"] = corrected_text or raw_text - text_to_translate = paragraph_text or corrected_text or raw_text - if text_to_translate: - if 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) + 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__": diff --git a/main_v2.py b/main_v2.py index 242dc0e..987cad1 100644 --- a/main_v2.py +++ b/main_v2.py @@ -57,6 +57,7 @@ def main(): parser.add_argument("-fr", action="store_true", default=defaults.get("fr", False)) parser.add_argument("-ar", action="store_true", default=defaults.get("ar", False)) parser.add_argument("-en", action="store_true", default=defaults.get("en", False)) + parser.add_argument("--only-translate-llm", action="store_true", default=defaults.get("only_translate_llm", False), help="Only translate when the LLM produces a refined paragraph.") parser.add_argument("--mt-max-chars", type=int, default=defaults.get("mt_max_chars", 250)) parser.add_argument("--mt-max-new-tokens", type=int, default=defaults.get("mt_max_new_tokens", 150)) parser.add_argument("--mt-num-beams", type=int, default=defaults.get("mt_num_beams", 4)) @@ -66,6 +67,13 @@ def main(): args = parser.parse_args() + # Auto-enable required flags + if args.post_correct_llm and not args.post_correct: + args.post_correct = True + if args.only_translate_llm and not args.llm_paragraph: + print("[Main] Enabling --llm-paragraph because --only-translate-llm was requested.") + args.llm_paragraph = True + # Handle device listing if args.list_devices: list_audio_devices() diff --git a/transcribe.py b/transcribe.py index 5ef2fa1..74b3b5b 100644 --- a/transcribe.py +++ b/transcribe.py @@ -610,6 +610,7 @@ def main(): parser.add_argument("--mt-max-chars", type=int, default=250, help="Max chars per translation chunk before sentence-aware splitting (default: 250)") parser.add_argument("--mt-max-new-tokens", type=int, default=150, help="Max new tokens per translation chunk (default: 150)") parser.add_argument("--mt-num-beams", type=int, default=4, help="Beam size for Marian translation generation (default: 4)") + parser.add_argument("--only-translate-llm", action="store_true", help="Only translate into target languages when a refined LLM paragraph is available.") parser.add_argument("--mt-no-repeat-ngram-size", type=int, default=3, help="No-repeat n-gram size for Marian generation (default: 3)") parser.add_argument("--mt-length-penalty", type=float, default=1.0, help="Length penalty for Marian generation (default: 1.0)") parser.add_argument("--mt-repetition-penalty", type=float, default=1.05, help="Repetition penalty for Marian generation (default: 1.05)") @@ -648,6 +649,10 @@ def main(): args.post_correct = True print("[SYSTEM]: Enabling --post-correct because --post-correct-llm was requested.") log_session("[SYSTEM]: Enabling --post-correct because --post-correct-llm was requested.") + if args.only_translate_llm and not args.llm_paragraph: + args.llm_paragraph = True + print("[SYSTEM]: Enabling --llm-paragraph because --only-translate-llm was requested.") + log_session("[SYSTEM]: Enabling --llm-paragraph because --only-translate-llm was requested.") file_glossary_pairs = load_glossary_file(args.glossary_file) merged_glossary_pairs = file_glossary_pairs + args.glossary_pair if file_glossary_pairs: @@ -1064,9 +1069,8 @@ def main(): rolling_context = (rolling_context + " " + original_text)[-200:].strip() # 3. Translate from English to other languages - if not args.llm_paragraph and english_for_translation and translation_engines: - english_chunks = split_text_for_translation(english_for_translation, max_chars=args.mt_max_chars) - + if not args.llm_paragraph and not args.only_translate_llm and english_for_translation and translation_engines: + english_chunks = split_text_for_translation(english_for_translation, max_chars=args.mt_max_chars) for lang_key, (model, tokenizer) in translation_engines.items(): # Skip if we already filled this (e.g. detected lang was 'es') if lang_key in payload: