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
+17 -1
View File
@@ -8,5 +8,21 @@
"es": false, "es": false,
"fr": false, "fr": false,
"ar": false, "ar": false,
"ingest": false "en": false,
"ingest": false,
"filter_lang": false,
"stream": false,
"context": false,
"quantize": false,
"speaker_diarization": false,
"post_correct": false,
"post_correct_llm": false,
"post_correct_model": "qwen:2b",
"llm_paragraph": false,
"temperature_fallback": [0.0, 0.2, 0.4, 0.6, 0.8, 1.0],
"logprob_threshold": -0.8,
"compression_threshold": 2.2,
"mt_max_chars": 250,
"mt_max_new_tokens": 150,
"mt_num_beams": 4
} }
+21 -19
View File
@@ -2,6 +2,7 @@ import requests
import time import time
import argparse import argparse
import multiprocessing import multiprocessing
import sys
INGEST_URL = "https://emiapi.reynafamily.com/live-captions/ingest" INGEST_URL = "https://emiapi.reynafamily.com/live-captions/ingest"
@@ -13,37 +14,38 @@ def run_distribution(in_queue, args):
payload = in_queue.get() payload = in_queue.get()
if payload is None: break if payload is None: break
# Draft support
if "draft" in payload:
draft_text = payload["draft"]
sys.stdout.write(f"\r\033[K[DRAFT]: {draft_text}")
sys.stdout.flush()
if args.ingest: if args.ingest:
delay = 1 requests.post(INGEST_URL, json={"draft": draft_text}, timeout=2)
max_delay = 15 continue
success = False
# Finalized segment
print(f"\n[Final]: {payload.get('original', '')}")
for lang in ["es", "fr", "ar", "en"]:
if payload.get(lang):
print(f"[{lang.upper()}]: {payload[lang]}")
if args.ingest:
delay, max_delay, success = 1, 15, False
while not success: while not success:
try: try:
response = requests.post(INGEST_URL, json=payload, timeout=5) response = requests.post(INGEST_URL, json=payload, timeout=5)
if response.status_code == 200: if response.status_code == 200: success = True
success = True else: print(f"[Distribute Error] {response.status_code}. Retrying...")
else:
print(f"[Distribute Error] {response.status_code}. Retrying in {delay}s...")
except Exception as e: except Exception as e:
print(f"[Distribute Error] {e}. Retrying in {delay}s...") print(f"[Distribute Error] {e}. Retrying...")
if not success: if not success:
time.sleep(delay) time.sleep(delay)
delay = min(delay * 2, max_delay) delay = min(delay * 2, max_delay)
print(f"[Distribute] Sent payload: {list(payload.keys())}")
else:
print(f"[Distribute] Local display (ingest disabled): {payload}")
except Exception as e: except Exception as e:
print(f"[Distribute] Error: {e}") print(f"[Distribute] Error: {e}")
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser() import multiprocessing
parser.add_argument("-i", "--ingest", action="store_true") run_distribution(multiprocessing.Queue(), argparse.Namespace())
args = parser.parse_args()
# Dummy queue for testing
in_q = multiprocessing.Queue()
run_distribution(in_q, args)
+79 -44
View File
@@ -16,13 +16,14 @@ except ImportError:
sys.modules["_lzma"] = MagicMock() sys.modules["_lzma"] = MagicMock()
sys.modules["lzma"] = mock_lzma sys.modules["lzma"] = mock_lzma
import sys
import time import time
import os
import numpy as np import numpy as np
import sounddevice as sd import sounddevice as sd
import torch import torch
import mlx_whisper import mlx_whisper
from silero_vad import load_silero_vad, get_speech_timestamps from silero_vad import load_silero_vad, get_speech_timestamps
from pyannote.audio import Pipeline
import queue import queue
import argparse import argparse
@@ -36,7 +37,7 @@ def transcribe_with_controls(audio, transcribe_kwargs):
try: try:
return mlx_whisper.transcribe(audio, **transcribe_kwargs) return mlx_whisper.transcribe(audio, **transcribe_kwargs)
except Exception as exc: except Exception as exc:
# Fallback logic for beam_size or other specific MLX implementation gaps message = str(exc)
if "beam_size" in transcribe_kwargs: if "beam_size" in transcribe_kwargs:
transcribe_kwargs.pop("beam_size") transcribe_kwargs.pop("beam_size")
return mlx_whisper.transcribe(audio, **transcribe_kwargs) return mlx_whisper.transcribe(audio, **transcribe_kwargs)
@@ -47,14 +48,44 @@ def audio_callback(indata, frames, time, status, audio_queue):
print(status, file=sys.stderr) print(status, file=sys.stderr)
audio_queue.put(indata.copy()) audio_queue.put(indata.copy())
def assign_speakers_to_segments(segments, diarization):
if not diarization or not segments:
return segments
for segment in segments:
segment_start, segment_end = segment['start'], segment['end']
speaker_intersections = {}
for turn, _, speaker in diarization.itertracks(yield_label=True):
intersection_start = max(segment_start, turn.start)
intersection_end = min(segment_end, turn.end)
if intersection_end > intersection_start:
duration = intersection_end - intersection_start
speaker_intersections[speaker] = speaker_intersections.get(speaker, 0) + duration
if speaker_intersections:
segment['speaker'] = max(speaker_intersections, key=speaker_intersections.get)
else:
segment['speaker'] = 'UNKNOWN'
return segments
def run_transcription(out_queue, args): def run_transcription(out_queue, args):
device = "mps" if torch.backends.mps.is_available() else "cpu" device = "mps" if torch.backends.mps.is_available() else "cpu"
print(f"[Transcribe] Loading Whisper model '{args.model}' on {device}...") model_name = args.model
if args.quantize and "4bit" not in model_name:
model_name = "mlx-community/whisper-small-mlx-4bit"
print(f"[Transcribe] Loading Whisper model '{model_name}' on {device}...")
vad_model = load_silero_vad() vad_model = load_silero_vad()
audio_queue = queue.Queue()
# Selection of device diarization_pipeline = None
if args.speaker_diarization:
print("[Transcribe] Loading speaker diarization pipeline...")
hf_token = os.environ.get("HF_TOKEN")
try:
diarization_pipeline = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1", use_auth_token=hf_token)
diarization_pipeline.to(torch.device(device))
except Exception as e:
print(f"[Transcribe] Diarization load failed: {e}")
audio_queue = queue.Queue()
device_index = args.device device_index = args.device
def callback_wrapper(indata, frames, time, status): def callback_wrapper(indata, frames, time, status):
@@ -63,16 +94,17 @@ def run_transcription(out_queue, args):
audio_buffer = [] audio_buffer = []
speech_started = False speech_started = False
buffer_limit = SAMPLERATE * args.max_buffer buffer_limit = SAMPLERATE * args.max_buffer
rolling_context = ""
last_stream_time = time.time()
print(f"[Transcribe] Starting audio stream on device {device_index}...") print(f"[Transcribe] Audio stream ready on device {device_index}.")
try: try:
with sd.InputStream(samplerate=SAMPLERATE, channels=args.channels, callback=callback_wrapper, blocksize=BLOCK_SIZE, device=device_index): with sd.InputStream(samplerate=SAMPLERATE, channels=args.channels, callback=callback_wrapper, blocksize=BLOCK_SIZE, device=device_index):
while True: while True:
while not audio_queue.empty(): while not audio_queue.empty():
data = audio_queue.get() data = audio_queue.get()
if args.channels > 1: if args.channels > 1: data = np.mean(data, axis=1)
data = np.mean(data, axis=1) # Mix to mono
audio_buffer.append(data.flatten()) audio_buffer.append(data.flatten())
if audio_buffer: if audio_buffer:
@@ -80,64 +112,67 @@ def run_transcription(out_queue, args):
audio_tensor = torch.from_numpy(current_audio) audio_tensor = torch.from_numpy(current_audio)
buffer_duration = len(current_audio) / SAMPLERATE buffer_duration = len(current_audio) / SAMPLERATE
speech_timestamps = get_speech_timestamps( speech_timestamps = get_speech_timestamps(audio_tensor, vad_model, sampling_rate=SAMPLERATE, threshold=VAD_THRESHOLD, min_silence_duration_ms=args.silence)
audio_tensor,
vad_model,
sampling_rate=SAMPLERATE,
threshold=VAD_THRESHOLD,
min_silence_duration_ms=args.silence
)
if speech_timestamps: if speech_timestamps:
speech_started = True speech_started = True
last_end = speech_timestamps[-1]['end'] last_end = speech_timestamps[-1]['end']
buffer_len_samples = len(current_audio) buffer_len_samples = len(current_audio)
# Check for flush
if (buffer_len_samples - last_end) > (SAMPLERATE * args.silence / 1000) or buffer_len_samples > buffer_limit: if (buffer_len_samples - last_end) > (SAMPLERATE * args.silence / 1000) or buffer_len_samples > buffer_limit:
# 1. Transcribe
transcribe_kwargs = { transcribe_kwargs = {
"path_or_hf_repo": args.model, "path_or_hf_repo": model_name,
"temperature": (0.0, 0.2, 0.4, 0.6, 0.8, 1.0), "temperature": args.temperature_fallback,
"word_timestamps": args.speaker_diarization,
"logprob_threshold": args.logprob_threshold,
"compression_ratio_threshold": args.compression_threshold,
} }
if args.lang: if args.lang: transcribe_kwargs["language"] = args.lang
transcribe_kwargs["language"] = args.lang if args.context and rolling_context: transcribe_kwargs["initial_prompt"] = rolling_context
result = transcribe_with_controls(current_audio, transcribe_kwargs) result = transcribe_with_controls(current_audio, transcribe_kwargs)
text = result['text'].strip() text = result['text'].strip()
detected_lang = result.get('language', 'en') detected_lang = result.get('language', args.lang or 'en')
if args.filter_lang and args.lang and detected_lang != args.lang:
print(f"[Transcribe] Filtered {detected_lang}")
text = ""
if text: if text:
# Send to Translation Process segments = result.get('segments', [])
out_queue.put({ if diarization_pipeline:
"original": text, audio_for_diarization = torch.from_numpy(current_audio).float().unsqueeze(0)
"detected_lang": detected_lang, try:
"ts": time.time() diarization_result = diarization_pipeline({"waveform": audio_for_diarization, "sample_rate": SAMPLERATE})
}) segments = assign_speakers_to_speakers(segments, diarization_result)
except: pass
out_queue.put({"original": text, "detected_lang": detected_lang, "segments": segments, "ts": time.time()})
print(f"[Transcribe] {detected_lang.upper()}: {text}") print(f"[Transcribe] {detected_lang.upper()}: {text}")
if args.context: rolling_context = (rolling_context + " " + text)[-200:].strip()
audio_buffer = [] audio_buffer = []
speech_started = False speech_started = False
elif args.stream and (time.time() - last_stream_time) > 1.5:
# Draft Mode
draft_kwargs = {"path_or_hf_repo": model_name, "temperature": 0.0}
if args.lang: draft_kwargs["language"] = args.lang
if args.context and rolling_context: draft_kwargs["initial_prompt"] = rolling_context
draft_result = transcribe_with_controls(current_audio, draft_kwargs)
draft_text = draft_result['text'].strip()
if draft_text:
out_queue.put({"draft": draft_text, "ts": time.time()})
last_stream_time = time.time()
elif not speech_started and len(current_audio) > SAMPLERATE * 2: elif not speech_started and len(current_audio) > SAMPLERATE * 2:
audio_buffer = [] audio_buffer = []
except KeyboardInterrupt: except KeyboardInterrupt: pass
print("[Transcribe] Stopped.") except Exception as e: print(f"[Transcribe] Error: {e}")
except Exception as e:
print(f"[Transcribe] Error: {e}")
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, default="mlx-community/whisper-base-mlx")
parser.add_argument("--device", type=int, default=None)
parser.add_argument("--lang", type=str, default=None)
parser.add_argument("--silence", type=int, default=1000)
parser.add_argument("--max-buffer", type=int, default=20)
parser.add_argument("--channels", type=int, default=1)
args = parser.parse_args()
# This part would normally be called by the main coordinator
# For testing, we can use a dummy queue
import multiprocessing import multiprocessing
q = multiprocessing.Queue() run_transcription(multiprocessing.Queue(), argparse.Namespace())
run_transcription(q, args)
+43 -40
View File
@@ -22,6 +22,7 @@ import re
import argparse import argparse
import multiprocessing import multiprocessing
import time import time
import requests
TARGET_LANGS = { TARGET_LANGS = {
"es": "Helsinki-NLP/opus-mt-en-es", "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() normalized = " ".join(text.split()).strip()
if not normalized or len(normalized) <= max_chars: if not normalized or len(normalized) <= max_chars:
return [normalized] if normalized else [] return [normalized] if normalized else []
chunks = [] chunks = []
current = "" current = ""
sentences = [s for s in re.split(r"(?<=[.!?])\s+", normalized) if s] sentences = [s for s in re.split(r"(?<=[.!?])\s+", normalized) if s]
for sentence in sentences: for sentence in sentences:
if len(sentence) > max_chars: if len(sentence) > max_chars:
words = sentence.split() words = sentence.split()
word_chunk = "" word_chunk = ""
for word in words: for word in words:
candidate = f"{word_chunk} {word}".strip() candidate = f"{word_chunk} {word}".strip()
if len(candidate) <= max_chars: if len(candidate) <= max_chars: word_chunk = candidate
word_chunk = candidate
else: else:
if word_chunk: chunks.append(word_chunk) if word_chunk: chunks.append(word_chunk)
word_chunk = word word_chunk = word
if word_chunk: chunks.append(word_chunk) if word_chunk: chunks.append(word_chunk)
continue continue
candidate = f"{current} {sentence}".strip() candidate = f"{current} {sentence}".strip()
if candidate and len(candidate) <= max_chars: if candidate and len(candidate) <= max_chars: current = candidate
current = candidate
else: else:
if current: chunks.append(current) if current: chunks.append(current)
current = sentence current = sentence
if current: chunks.append(current) if current: chunks.append(current)
return chunks 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): def run_translation(in_queue, out_queue, args):
device = "mps" if torch.backends.mps.is_available() else "cpu" device = "mps" if torch.backends.mps.is_available() else "cpu"
print(f"[Translate] Loading translation models on {device}...")
translation_engines = {} translation_engines = {}
for lang_key, model_id in TARGET_LANGS.items(): for lang_key, model_id in TARGET_LANGS.items():
if getattr(args, lang_key, False): if getattr(args, lang_key, False):
print(f"[Translate] Loading {lang_key} model...") print(f"[Translate] Loading {lang_key} model...")
tokenizer = MarianTokenizer.from_pretrained(model_id) translation_engines[lang_key] = (
model = MarianMTModel.from_pretrained(model_id).to(device) MarianMTModel.from_pretrained(model_id).to(device),
translation_engines[lang_key] = (model, tokenizer) MarianTokenizer.from_pretrained(model_id)
)
paragraph_buffer = [] paragraph_buffer = []
print("[Translate] Ready.")
while True: while True:
try: try:
item = in_queue.get() item = in_queue.get()
if item is None: break if item is None: break
# Draft support
if "draft" in item:
out_queue.put({"draft": item["draft"]})
continue
text = item.get("original", "").strip() text = item.get("original", "").strip()
detected_lang = item.get("detected_lang", "en") detected_lang = item.get("detected_lang", "en")
# Simple paragraph logic: if args.post_correct and args.post_correct_llm:
# If the segment ends with sentence-terminal punctuation, flush the paragraph. text = post_correct_with_llm(text, args)
paragraph_buffer.append(text) paragraph_buffer.append(text)
if text.endswith((".", "?", "!")): # Paragraph logic
if not args.llm_paragraph or text.endswith((".", "?", "!")):
full_text = " ".join(paragraph_buffer) 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. if args.llm_paragraph:
# For now, let's assume direct translation for simplicity or bridge if needed. full_text = structure_paragraph_with_llm(full_text, args)
# (Refining bridge logic can come later)
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(): 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 = [] translated_parts = []
for chunk in chunks: for chunk in chunks:
inputs = tokenizer(chunk, return_tensors="pt", padding=True).to(device) inputs = tokenizer(chunk, return_tensors="pt", padding=True).to(device)
with torch.no_grad(): 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_parts.append(tokenizer.decode(translated_tokens[0], skip_special_tokens=True).strip())
payload[lang_key] = " ".join(translated_parts)
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}")
out_queue.put(payload) out_queue.put(payload)
paragraph_buffer = [] paragraph_buffer = []
except Exception as e: except Exception as e: print(f"[Translate] Error: {e}")
print(f"[Translate] Error: {e}")
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser() import multiprocessing
parser.add_argument("-es", action="store_true") run_translation(multiprocessing.Queue(), multiprocessing.Queue(), argparse.Namespace())
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)
+27 -8
View File
@@ -18,13 +18,18 @@ except ImportError:
import multiprocessing import multiprocessing
import argparse import argparse
import sys
import json import json
import os import os
from engine_transcribe import run_transcription from engine_transcribe import run_transcription
from engine_translate import run_translation from engine_translate import run_translation
from engine_distribute import run_distribution from engine_distribute import run_distribution
def parse_temperature_fallback(value):
try:
return tuple(float(x.strip()) for x in value.split(",") if x.strip())
except:
return (0.0, 0.2, 0.4, 0.6, 0.8, 1.0)
def main(): def main():
config_path = "config.json" config_path = "config.json"
defaults = {} defaults = {}
@@ -41,21 +46,35 @@ def main():
parser.add_argument("--silence", type=int, default=defaults.get("silence", 1000)) parser.add_argument("--silence", type=int, default=defaults.get("silence", 1000))
parser.add_argument("--max-buffer", type=int, default=defaults.get("max_buffer", 20)) parser.add_argument("--max-buffer", type=int, default=defaults.get("max_buffer", 20))
parser.add_argument("--channels", type=int, default=defaults.get("channels", 1)) parser.add_argument("--channels", type=int, default=defaults.get("channels", 1))
parser.add_argument("--filter-lang", action="store_true", default=defaults.get("filter_lang", False))
parser.add_argument("-q", "--quantize", action="store_true", default=defaults.get("quantize", False))
parser.add_argument("-s", "--stream", action="store_true", default=defaults.get("stream", False))
parser.add_argument("-c", "--context", action="store_true", default=defaults.get("context", False))
parser.add_argument("--speaker-diarization", action="store_true", default=defaults.get("speaker_diarization", False))
parser.add_argument("--temperature-fallback", type=parse_temperature_fallback, default=tuple(defaults.get("temperature_fallback", [0.0, 0.2, 0.4, 0.6, 0.8, 1.0])))
parser.add_argument("--logprob-threshold", type=float, default=defaults.get("logprob_threshold", -0.8))
parser.add_argument("--compression-threshold", type=float, default=defaults.get("compression_threshold", 2.2))
# Translate Args # Translate Args
parser.add_argument("-es", action="store_true", default=defaults.get("es", False), help="Enable Spanish translation") parser.add_argument("-es", action="store_true", default=defaults.get("es", False))
parser.add_argument("-fr", action="store_true", default=defaults.get("fr", False), help="Enable French translation") parser.add_argument("-fr", action="store_true", default=defaults.get("fr", False))
parser.add_argument("-ar", action="store_true", default=defaults.get("ar", False), help="Enable Arabic translation") 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("--post-correct", action="store_true", default=defaults.get("post_correct", False))
parser.add_argument("--post-correct-llm", action="store_true", default=defaults.get("post_correct_llm", False))
parser.add_argument("--post-correct-model", type=str, default=defaults.get("post_correct_model", "qwen:2b"))
parser.add_argument("--llm-paragraph", action="store_true", default=defaults.get("llm_paragraph", False))
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))
# Distribute Args # Distribute Args
parser.add_argument("-i", "--ingest", action="store_true", default=defaults.get("ingest", False), help="Enable data transmission to server") parser.add_argument("-i", "--ingest", action="store_true", default=defaults.get("ingest", False))
args = parser.parse_args() args = parser.parse_args()
# Queues for communication # Queues
# Transcribe -> Translate
q_trans_to_tl = multiprocessing.Queue() q_trans_to_tl = multiprocessing.Queue()
# Translate -> Distribute
q_tl_to_dist = multiprocessing.Queue() q_tl_to_dist = multiprocessing.Queue()
# Processes # Processes