feat: improve bridge context and translation decoding controls

This commit is contained in:
Adolfo Reyna
2026-03-06 18:56:21 -05:00
parent 6e94210105
commit dc86ddedbd
3 changed files with 151 additions and 12 deletions
+6 -2
View File
@@ -1,3 +1,7 @@
__pycache__/\n*.pyc\n.DS_Store __pycache__/
build/\ndist/\n*.spec *.pyc
.DS_Store
build/
dist/
*.spec
venv/ venv/
+10
View File
@@ -126,3 +126,13 @@ The project now features a high-performance, Apple Silicon-optimized pipeline th
- **Heuristic Thresholds:** Utilize `logprob_threshold` (confidence) and `compression_ratio_threshold` (repetition) to programmatically identify and discard bad transcriptions before they reach the user. - **Heuristic Thresholds:** Utilize `logprob_threshold` (confidence) and `compression_ratio_threshold` (repetition) to programmatically identify and discard bad transcriptions before they reach the user.
- **Structural Features:** - **Structural Features:**
- **Word-Level Timestamps:** Enable `word_timestamps=True` to provide granular timing data for front-end caption highlighting. - **Word-Level Timestamps:** Enable `word_timestamps=True` to provide granular timing data for front-end caption highlighting.
## Phase 16: Decoder Controls + Bridge Context Refinements
- **Goal:** Improve English caption stability and translation quality without sacrificing runtime compatibility.
- **Approach:**
- **Whisper Decoder Controls (implemented):** Added configurable decoding flags for final segments and drafts, including temperature fallback and quality thresholds (`logprob_threshold`, `compression_ratio_threshold`), with compatibility fallback when unsupported.
- **Beam Compatibility Hardening:** Defaulted beam sizes to greedy-safe values and auto-retry without `beam_size` when the runtime reports "Beam search decoder is not yet implemented."
- **Dedicated English Bridge Context:** Split context handling into source-language context and a separate English context used specifically for `task="translate"` bridge generation.
- **Sentence-Aware MT Chunking:** Replaced hard truncation with sentence-aware splitting for long English bridge text before translation.
- **Marian Generation Tuning:** Added configurable generation controls (`num_beams`, `no_repeat_ngram_size`, `length_penalty`, `repetition_penalty`, `early_stopping`) to reduce repetitive or unstable outputs.
- **Outcome:** Better continuity for non-English to English bridging, fewer clipped translations on long segments, and cleaner target-language output with safer defaults for current `mlx_whisper` builds.
+135 -10
View File
@@ -50,6 +50,84 @@ VAD_THRESHOLD = 0.5
audio_queue = queue.Queue() audio_queue = queue.Queue()
ingest_queue = queue.Queue() ingest_queue = queue.Queue()
def parse_temperature_fallback(value):
"""Parse comma-separated temperatures into a tuple of floats."""
try:
temps = tuple(float(x.strip()) for x in value.split(",") if x.strip())
except ValueError as exc:
raise argparse.ArgumentTypeError("Invalid --temperature-fallback value.") from exc
if not temps:
raise argparse.ArgumentTypeError("--temperature-fallback requires at least one value.")
return temps
def split_text_for_translation(text, max_chars=250):
"""Split long English text into sentence-aware chunks for MT."""
normalized = " ".join(text.split()).strip()
if not normalized:
return []
if len(normalized) <= max_chars:
return [normalized]
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
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
def transcribe_with_controls(audio, transcribe_kwargs):
"""Call mlx_whisper.transcribe and gracefully fallback if a decoder arg is unsupported."""
optional_keys = [
"beam_size",
"temperature",
"logprob_threshold",
"compression_ratio_threshold",
]
try:
return mlx_whisper.transcribe(audio, **transcribe_kwargs)
except TypeError as exc:
message = str(exc)
unsupported = [k for k in optional_keys if f"'{k}'" in message]
if not unsupported:
raise
retry_kwargs = {k: v for k, v in transcribe_kwargs.items() if k not in unsupported}
print(f"\n[SYSTEM]: Decoder args not supported by current mlx_whisper build: {', '.join(unsupported)}. Retrying without them.")
return mlx_whisper.transcribe(audio, **retry_kwargs)
except Exception as exc:
message = str(exc).lower()
if "beam search decoder is not yet implemented" in message and "beam_size" in transcribe_kwargs:
retry_kwargs = dict(transcribe_kwargs)
retry_kwargs.pop("beam_size", None)
print("\n[SYSTEM]: Beam search is not implemented in this mlx_whisper build. Retrying with greedy decoding.")
return mlx_whisper.transcribe(audio, **retry_kwargs)
raise
def is_hallucination(text): def is_hallucination(text):
"""Detect common Whisper hallucinations or high repetition.""" """Detect common Whisper hallucinations or high repetition."""
if not text: return False if not text: return False
@@ -129,6 +207,18 @@ def main():
parser.add_argument("--channels", type=int, default=1, help="Number of input channels (default: 1)") parser.add_argument("--channels", type=int, default=1, help="Number of input channels (default: 1)")
parser.add_argument("--pick-channel", type=int, choices=[0, 1], help="Pick a specific channel (0 or 1) from stereo input") parser.add_argument("--pick-channel", type=int, choices=[0, 1], help="Pick a specific channel (0 or 1) from stereo input")
parser.add_argument("--filter-lang", action="store_true", help="Discard segments where detected language does not match --lang") parser.add_argument("--filter-lang", action="store_true", help="Discard segments where detected language does not match --lang")
parser.add_argument("--draft-beam-size", type=int, default=1, help="Beam size for draft mode transcriptions (default: 1)")
parser.add_argument("--final-beam-size", type=int, default=1, help="Beam size for final segment transcriptions (default: 1)")
parser.add_argument("--temperature-fallback", type=parse_temperature_fallback, default=(0.0, 0.2, 0.4, 0.6, 0.8, 1.0), help="Comma-separated temperatures for fallback decoding (default: 0.0,0.2,0.4,0.6,0.8,1.0)")
parser.add_argument("--logprob-threshold", type=float, default=-0.8, help="Reject low-confidence tokens below this avg logprob (default: -0.8)")
parser.add_argument("--compression-threshold", type=float, default=2.2, help="Reject repetitive outputs above this compression ratio (default: 2.2)")
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("--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)")
parser.add_argument("--mt-no-early-stopping", action="store_true", help="Disable early stopping in Marian beam search")
args = parser.parse_args() args = parser.parse_args()
@@ -209,6 +299,7 @@ def main():
last_stream_time = time.time() last_stream_time = time.time()
last_change_time = time.time() last_change_time = time.time()
rolling_context = "" rolling_context = ""
rolling_context_en = ""
last_draft_text = "" last_draft_text = ""
try: try:
@@ -247,6 +338,7 @@ def main():
audio_buffer = [] audio_buffer = []
speech_started = False speech_started = False
rolling_context = "" rolling_context = ""
rolling_context_en = ""
last_draft_text = "" last_draft_text = ""
last_change_time = time.time() last_change_time = time.time()
if args.stream: if args.stream:
@@ -272,9 +364,14 @@ def main():
transcribe_kwargs["language"] = args.lang transcribe_kwargs["language"] = args.lang
if args.context and rolling_context: if args.context and rolling_context:
transcribe_kwargs["initial_prompt"] = rolling_context transcribe_kwargs["initial_prompt"] = rolling_context
if args.final_beam_size > 1:
transcribe_kwargs["beam_size"] = args.final_beam_size
transcribe_kwargs["temperature"] = args.temperature_fallback
transcribe_kwargs["logprob_threshold"] = args.logprob_threshold
transcribe_kwargs["compression_ratio_threshold"] = args.compression_threshold
# 1. Transcribe & Detect Language # 1. Transcribe & Detect Language
transcription_result = mlx_whisper.transcribe(current_audio, **transcribe_kwargs) transcription_result = transcribe_with_controls(current_audio, transcribe_kwargs)
original_text = transcription_result['text'].strip() original_text = transcription_result['text'].strip()
detected_lang = transcription_result.get('language', args.lang if args.lang else 'en') detected_lang = transcription_result.get('language', args.lang if args.lang else 'en')
@@ -287,6 +384,7 @@ def main():
print(f"\n[SYSTEM]: Hallucination detected, ignoring and resetting context.") print(f"\n[SYSTEM]: Hallucination detected, ignoring and resetting context.")
original_text = "" original_text = ""
rolling_context = "" rolling_context = ""
rolling_context_en = ""
if original_text: if original_text:
print(f"\n[{detected_lang.upper()}]: {original_text}") print(f"\n[{detected_lang.upper()}]: {original_text}")
@@ -305,13 +403,25 @@ def main():
bridge_kwargs = {"path_or_hf_repo": WHISPER_MODEL, "task": "translate"} bridge_kwargs = {"path_or_hf_repo": WHISPER_MODEL, "task": "translate"}
if args.lang: if args.lang:
bridge_kwargs["language"] = args.lang bridge_kwargs["language"] = args.lang
bridge_result = mlx_whisper.transcribe(current_audio, **bridge_kwargs) if args.context and rolling_context_en:
# Keep a dedicated English context for translate mode.
bridge_kwargs["initial_prompt"] = rolling_context_en
if args.final_beam_size > 1:
bridge_kwargs["beam_size"] = args.final_beam_size
bridge_kwargs["temperature"] = args.temperature_fallback
bridge_kwargs["logprob_threshold"] = args.logprob_threshold
bridge_kwargs["compression_ratio_threshold"] = args.compression_threshold
bridge_result = transcribe_with_controls(current_audio, bridge_kwargs)
english_text = bridge_result['text'].strip() english_text = bridge_result['text'].strip()
if args.context and english_text:
rolling_context_en = (rolling_context_en + " " + english_text)[-200:].strip()
if args.en: if args.en:
payload["en"] = english_text payload["en"] = english_text
print(f"[EN]: {english_text}") print(f"[EN]: {english_text}")
else: else:
english_text = original_text english_text = original_text
if args.context and english_text:
rolling_context_en = (rolling_context_en + " " + english_text)[-200:].strip()
if args.en: if args.en:
payload["en"] = english_text payload["en"] = english_text
@@ -322,8 +432,7 @@ def main():
# 3. Translate from English to other languages # 3. Translate from English to other languages
if english_text and translation_engines: if english_text and translation_engines:
# Limit input length english_chunks = split_text_for_translation(english_text, max_chars=args.mt_max_chars)
clean_en = english_text[:247] + "..." if len(english_text) > 250 else english_text
for lang_key, (model, tokenizer) in translation_engines.items(): for lang_key, (model, tokenizer) in translation_engines.items():
# Skip if we already filled this (e.g. detected lang was 'es') # Skip if we already filled this (e.g. detected lang was 'es')
@@ -331,11 +440,23 @@ def main():
if lang_key != detected_lang: # Already printed original if lang_key != detected_lang: # Already printed original
print(f"[{lang_key.upper()}]: {payload[lang_key]}") print(f"[{lang_key.upper()}]: {payload[lang_key]}")
continue continue
inputs = tokenizer(clean_en, return_tensors="pt", padding=True).to(device) translated_parts = []
with torch.no_grad(): for chunk in english_chunks:
translated_tokens = model.generate(**inputs, max_new_tokens=150) inputs = tokenizer(chunk, return_tensors="pt", padding=True).to(device)
translated_text = tokenizer.decode(translated_tokens[0], skip_special_tokens=True) with torch.no_grad():
translated_tokens = model.generate(
**inputs,
max_new_tokens=args.mt_max_new_tokens,
num_beams=args.mt_num_beams,
no_repeat_ngram_size=args.mt_no_repeat_ngram_size,
length_penalty=args.mt_length_penalty,
repetition_penalty=args.mt_repetition_penalty,
early_stopping=not args.mt_no_early_stopping,
)
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 payload[lang_key] = translated_text
print(f"[{lang_key.upper()}]: {translated_text}") print(f"[{lang_key.upper()}]: {translated_text}")
@@ -355,8 +476,11 @@ def main():
draft_kwargs = {"path_or_hf_repo": WHISPER_MODEL} draft_kwargs = {"path_or_hf_repo": WHISPER_MODEL}
if args.lang: draft_kwargs["language"] = args.lang if args.lang: draft_kwargs["language"] = args.lang
if args.context and rolling_context: draft_kwargs["initial_prompt"] = rolling_context if args.context and rolling_context: draft_kwargs["initial_prompt"] = rolling_context
if args.draft_beam_size > 1:
draft_kwargs["beam_size"] = args.draft_beam_size
draft_kwargs["temperature"] = 0.0
draft_result = mlx_whisper.transcribe(current_audio, **draft_kwargs) draft_result = transcribe_with_controls(current_audio, draft_kwargs)
draft_text = draft_result['text'].strip() draft_text = draft_result['text'].strip()
draft_lang = draft_result.get('language', args.lang if args.lang else 'en') draft_lang = draft_result.get('language', args.lang if args.lang else 'en')
@@ -383,6 +507,7 @@ def main():
audio_buffer = [] audio_buffer = []
speech_started = False speech_started = False
rolling_context = "" rolling_context = ""
rolling_context_en = ""
last_change_time = time.time() last_change_time = time.time()
last_draft_text = "" last_draft_text = ""
sys.stdout.write("\r\033[K") sys.stdout.write("\r\033[K")