feat: refactor transcription, translation, and distribution into multiple processes

This commit is contained in:
Adolfo Reyna
2026-03-15 19:46:03 -04:00
7 changed files with 608 additions and 92 deletions
+12
View File
@@ -0,0 +1,12 @@
{
"model": "mlx-community/whisper-base-mlx",
"device": null,
"lang": null,
"silence": 1000,
"max_buffer": 20,
"channels": 1,
"es": false,
"fr": false,
"ar": false,
"ingest": false
}
+49
View File
@@ -0,0 +1,49 @@
import requests
import time
import argparse
import multiprocessing
INGEST_URL = "https://emiapi.reynafamily.com/live-captions/ingest"
def run_distribution(in_queue, args):
print("[Distribute] Starting distribution service...")
while True:
try:
payload = in_queue.get()
if payload is None: break
if args.ingest:
delay = 1
max_delay = 15
success = 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 in {delay}s...")
except Exception as e:
print(f"[Distribute Error] {e}. Retrying in {delay}s...")
if not success:
time.sleep(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:
print(f"[Distribute] Error: {e}")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("-i", "--ingest", action="store_true")
args = parser.parse_args()
# Dummy queue for testing
in_q = multiprocessing.Queue()
run_distribution(in_q, args)
+143
View File
@@ -0,0 +1,143 @@
import sys
from unittest.mock import MagicMock
# Comprehensive workaround for missing _lzma in some Python builds
try:
import lzma
except ImportError:
mock_lzma = MagicMock()
mock_lzma.FORMAT_XZ = 1
mock_lzma.FORMAT_ALONE = 2
mock_lzma.FORMAT_RAW = 3
mock_lzma.CHECK_NONE = 0
mock_lzma.CHECK_CRC32 = 1
mock_lzma.CHECK_CRC64 = 4
mock_lzma.CHECK_SHA256 = 10
sys.modules["_lzma"] = MagicMock()
sys.modules["lzma"] = mock_lzma
import sys
import time
import numpy as np
import sounddevice as sd
import torch
import mlx_whisper
from silero_vad import load_silero_vad, get_speech_timestamps
import queue
import argparse
# Audio Constants
SAMPLERATE = 16000
BLOCK_SIZE = 512
VAD_THRESHOLD = 0.5
def transcribe_with_controls(audio, transcribe_kwargs):
"""Call mlx_whisper.transcribe with graceful fallback for unsupported args."""
try:
return mlx_whisper.transcribe(audio, **transcribe_kwargs)
except Exception as exc:
# Fallback logic for beam_size or other specific MLX implementation gaps
if "beam_size" in transcribe_kwargs:
transcribe_kwargs.pop("beam_size")
return mlx_whisper.transcribe(audio, **transcribe_kwargs)
raise exc
def audio_callback(indata, frames, time, status, audio_queue):
if status:
print(status, file=sys.stderr)
audio_queue.put(indata.copy())
def run_transcription(out_queue, args):
device = "mps" if torch.backends.mps.is_available() else "cpu"
print(f"[Transcribe] Loading Whisper model '{args.model}' on {device}...")
vad_model = load_silero_vad()
audio_queue = queue.Queue()
# Selection of device
device_index = args.device
def callback_wrapper(indata, frames, time, status):
audio_callback(indata, frames, time, status, audio_queue)
audio_buffer = []
speech_started = False
buffer_limit = SAMPLERATE * args.max_buffer
print(f"[Transcribe] Starting audio stream on device {device_index}...")
try:
with sd.InputStream(samplerate=SAMPLERATE, channels=args.channels, callback=callback_wrapper, blocksize=BLOCK_SIZE, device=device_index):
while True:
while not audio_queue.empty():
data = audio_queue.get()
if args.channels > 1:
data = np.mean(data, axis=1) # Mix to mono
audio_buffer.append(data.flatten())
if audio_buffer:
current_audio = np.concatenate(audio_buffer)
audio_tensor = torch.from_numpy(current_audio)
buffer_duration = len(current_audio) / SAMPLERATE
speech_timestamps = get_speech_timestamps(
audio_tensor,
vad_model,
sampling_rate=SAMPLERATE,
threshold=VAD_THRESHOLD,
min_silence_duration_ms=args.silence
)
if speech_timestamps:
speech_started = True
last_end = speech_timestamps[-1]['end']
buffer_len_samples = len(current_audio)
if (buffer_len_samples - last_end) > (SAMPLERATE * args.silence / 1000) or buffer_len_samples > buffer_limit:
# 1. Transcribe
transcribe_kwargs = {
"path_or_hf_repo": args.model,
"temperature": (0.0, 0.2, 0.4, 0.6, 0.8, 1.0),
}
if args.lang:
transcribe_kwargs["language"] = args.lang
result = transcribe_with_controls(current_audio, transcribe_kwargs)
text = result['text'].strip()
detected_lang = result.get('language', 'en')
if text:
# Send to Translation Process
out_queue.put({
"original": text,
"detected_lang": detected_lang,
"ts": time.time()
})
print(f"[Transcribe] {detected_lang.upper()}: {text}")
audio_buffer = []
speech_started = False
elif not speech_started and len(current_audio) > SAMPLERATE * 2:
audio_buffer = []
except KeyboardInterrupt:
print("[Transcribe] Stopped.")
except Exception as e:
print(f"[Transcribe] Error: {e}")
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
q = multiprocessing.Queue()
run_transcription(q, args)
+130
View File
@@ -0,0 +1,130 @@
import sys
from unittest.mock import MagicMock
# Comprehensive workaround for missing _lzma in some Python builds
try:
import lzma
except ImportError:
mock_lzma = MagicMock()
mock_lzma.FORMAT_XZ = 1
mock_lzma.FORMAT_ALONE = 2
mock_lzma.FORMAT_RAW = 3
mock_lzma.CHECK_NONE = 0
mock_lzma.CHECK_CRC32 = 1
mock_lzma.CHECK_CRC64 = 4
mock_lzma.CHECK_SHA256 = 10
sys.modules["_lzma"] = MagicMock()
sys.modules["lzma"] = mock_lzma
import torch
from transformers import MarianMTModel, MarianTokenizer
import re
import argparse
import multiprocessing
import time
TARGET_LANGS = {
"es": "Helsinki-NLP/opus-mt-en-es",
"fr": "Helsinki-NLP/opus-mt-en-fr",
"ar": "Helsinki-NLP/opus-mt-en-ar"
}
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
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 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)
paragraph_buffer = []
print("[Translate] Ready.")
while True:
try:
item = in_queue.get()
if item is None: break
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.
paragraph_buffer.append(text)
if 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)
for lang_key, (model, tokenizer) in translation_engines.items():
chunks = split_text_for_translation(full_text)
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_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}")
out_queue.put(payload)
paragraph_buffer = []
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)
+84
View File
@@ -0,0 +1,84 @@
import sys
from unittest.mock import MagicMock
# Comprehensive workaround for missing _lzma in some Python builds
try:
import lzma
except ImportError:
mock_lzma = MagicMock()
mock_lzma.FORMAT_XZ = 1
mock_lzma.FORMAT_ALONE = 2
mock_lzma.FORMAT_RAW = 3
mock_lzma.CHECK_NONE = 0
mock_lzma.CHECK_CRC32 = 1
mock_lzma.CHECK_CRC64 = 4
mock_lzma.CHECK_SHA256 = 10
sys.modules["_lzma"] = MagicMock()
sys.modules["lzma"] = mock_lzma
import multiprocessing
import argparse
import sys
import json
import os
from engine_transcribe import run_transcription
from engine_translate import run_translation
from engine_distribute import run_distribution
def main():
config_path = "config.json"
defaults = {}
if os.path.exists(config_path):
with open(config_path, "r") as f:
defaults = json.load(f)
parser = argparse.ArgumentParser(description="Multi-process Transcription and Translation.")
# Transcribe Args
parser.add_argument("--model", type=str, default=defaults.get("model", "mlx-community/whisper-base-mlx"))
parser.add_argument("--device", type=int, default=defaults.get("device"))
parser.add_argument("--lang", type=str, default=defaults.get("lang"))
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("--channels", type=int, default=defaults.get("channels", 1))
# Translate Args
parser.add_argument("-es", action="store_true", default=defaults.get("es", False), help="Enable Spanish translation")
parser.add_argument("-fr", action="store_true", default=defaults.get("fr", False), help="Enable French translation")
parser.add_argument("-ar", action="store_true", default=defaults.get("ar", False), help="Enable Arabic translation")
# Distribute Args
parser.add_argument("-i", "--ingest", action="store_true", default=defaults.get("ingest", False), help="Enable data transmission to server")
args = parser.parse_args()
# Queues for communication
# Transcribe -> Translate
q_trans_to_tl = multiprocessing.Queue()
# Translate -> Distribute
q_tl_to_dist = multiprocessing.Queue()
# Processes
p_transcribe = multiprocessing.Process(target=run_transcription, args=(q_trans_to_tl, args))
p_translate = multiprocessing.Process(target=run_translation, args=(q_trans_to_tl, q_tl_to_dist, args))
p_distribute = multiprocessing.Process(target=run_distribution, args=(q_tl_to_dist, args))
print("[Main] Starting processes...")
p_transcribe.start()
p_translate.start()
p_distribute.start()
try:
p_transcribe.join()
p_translate.join()
p_distribute.join()
except KeyboardInterrupt:
print("\n[Main] Stopping processes...")
p_transcribe.terminate()
p_translate.terminate()
p_distribute.terminate()
sys.exit(0)
if __name__ == "__main__":
multiprocessing.freeze_support()
main()
+4
View File
@@ -9,3 +9,7 @@ pyinstaller
sacremoses sacremoses
joblib joblib
sentencepiece sentencepiece
pyannote.audio
soundfile
langdetect
python-dotenv
+186 -92
View File
@@ -33,6 +33,18 @@ import queue
import torch import torch
from silero_vad import load_silero_vad, get_speech_timestamps from silero_vad import load_silero_vad, get_speech_timestamps
from transformers import MarianMTModel, MarianTokenizer from transformers import MarianMTModel, MarianTokenizer
from pyannote.audio import Pipeline
from dotenv import load_dotenv
from langdetect import detect, DetectorFactory
from langdetect.lang_detect_exception import LangDetectException
# Enforce deterministic langdetect results
DetectorFactory.seed = 0
# Load environment variables from .env file
load_dotenv()
# Parameters # Parameters
WHISPER_MODEL = "mlx-community/whisper-base-mlx" WHISPER_MODEL = "mlx-community/whisper-base-mlx"
@@ -168,51 +180,6 @@ def quality_score(mean_logprob, max_compression):
compression_penalty = 0.25 * max(0.0, (max_compression or 1.5) - 1.5) compression_penalty = 0.25 * max(0.0, (max_compression or 1.5) - 1.5)
return mean_logprob - compression_penalty return mean_logprob - compression_penalty
def cosine_similarity(vec_a, vec_b):
if vec_a is None or vec_b is None:
return None
denom = (np.linalg.norm(vec_a) * np.linalg.norm(vec_b))
if denom == 0:
return None
return float(np.dot(vec_a, vec_b) / denom)
def compute_speaker_signature(audio, samplerate):
"""Compute a lightweight speaker signature from averaged log spectrum."""
if audio is None or len(audio) < int(0.5 * samplerate):
return None
data = np.asarray(audio, dtype=np.float32)
if data.ndim > 1:
data = np.mean(data, axis=1)
# Keep only voiced-ish samples to reduce silence/noise influence.
voiced = data[np.abs(data) > 0.01]
if len(voiced) < int(0.3 * samplerate):
voiced = data
frame_size = 1024
hop = 512
if len(voiced) < frame_size:
return None
window = np.hanning(frame_size).astype(np.float32)
spectra = []
for i in range(0, len(voiced) - frame_size + 1, hop):
frame = voiced[i:i + frame_size] * window
mag = np.abs(np.fft.rfft(frame))
if mag.size > 1:
mag = mag[1:257] # Focus on lower-frequency envelope.
spectra.append(np.log1p(mag))
if not spectra:
return None
signature = np.mean(np.stack(spectra, axis=0), axis=0)
norm = np.linalg.norm(signature)
if norm == 0:
return None
return signature / norm
def should_retry_transcription(text, mean_logprob, max_compression, args): def should_retry_transcription(text, mean_logprob, max_compression, args):
if not text: if not text:
return False return False
@@ -373,6 +340,34 @@ def post_correct_with_local_llm(text, args, llm_state, context_lines=None):
print(f"[POST-LLM][ERROR]: {exc}", flush=True) print(f"[POST-LLM][ERROR]: {exc}", flush=True)
return text return text
def structure_paragraph_with_llm(text, args, llm_state):
"""Use a local LLM to structure a series of captions into a paragraph."""
prompt = (
"You are a real-time English caption editor. Your task is to take a series of transcribed audio segments and format them into a single, coherent paragraph.\n"
"Rules:\n"
"1. Combine the segments into a flowing paragraph.\n"
"2. Correct grammar, punctuation, and casing.\n"
"3. Remove disfluencies and false starts (e.g., 'um', 'uh').\n"
"4. Do not change the meaning of the text.\n"
"5. If there are speaker labels (e.g., [SPEAKER_01]:), preserve them.\n"
"6. Output ONLY the formatted paragraph.\n\n"
"Here are the segments:\n"
f"{text}\n\n"
"Formatted paragraph:"
)
try:
response = call_ollama_generate(args, prompt, args.post_correct_llm_timeout, temperature=0.2)
response.raise_for_status()
raw = response.json().get("response", "").strip()
return raw
except Exception as exc:
if not llm_state.get("paragraph_warned", False):
print(f"\n[SYSTEM]: Local paragraph structuring LLM unavailable ({exc}).")
llm_state["paragraph_warned"] = True
return text # Fallback to the original text
def maybe_merge_recent_caption(recent_en_lines, current_text, now_ts, window_sec): def maybe_merge_recent_caption(recent_en_lines, current_text, now_ts, window_sec):
if not recent_en_lines: if not recent_en_lines:
return None return None
@@ -507,6 +502,38 @@ def is_hallucination(text):
return False return False
def assign_speakers_to_segments(segments, diarization):
"""Assign speaker labels to transcription segments based on diarization."""
if not diarization or not segments:
return segments
for segment in segments:
segment_start = segment['start']
segment_end = 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:
intersection_duration = intersection_end - intersection_start
if speaker not in speaker_intersections:
speaker_intersections[speaker] = 0
speaker_intersections[speaker] += intersection_duration
if speaker_intersections:
dominant_speaker = max(speaker_intersections, key=speaker_intersections.get)
segment['speaker'] = dominant_speaker
else:
# Keep speaker from previous segment if segment is short and follows closely
# This can help with short utterances that are missed by the diarizer
# For now, we'll just label as unknown
segment['speaker'] = 'UNKNOWN'
return segments
def callback(indata, frames, time, status): def callback(indata, frames, time, status):
if status: if status:
print(status, file=sys.stderr) print(status, file=sys.stderr)
@@ -569,6 +596,7 @@ def main():
parser.add_argument("-s", "--stream", action="store_true", help="Enable real-time streaming transcription (Draft mode)") parser.add_argument("-s", "--stream", action="store_true", help="Enable real-time streaming transcription (Draft mode)")
parser.add_argument("-c", "--context", action="store_true", help="Enable prompt caching/rolling context for better continuity") parser.add_argument("-c", "--context", action="store_true", help="Enable prompt caching/rolling context for better continuity")
parser.add_argument("--lang", type=str, help="Hardcode source language (e.g. 'en', 'es') to bypass detection") parser.add_argument("--lang", type=str, help="Hardcode source language (e.g. 'en', 'es') to bypass detection")
parser.add_argument("--detect-auto-translation", action="store_true", help="Enable heuristic to detect and leverage auto-translated English segments.")
parser.add_argument("--silence", type=int, default=1000, help="Minimum silence duration in ms to end a chunk (default: 1000)") parser.add_argument("--silence", type=int, default=1000, help="Minimum silence duration in ms to end a chunk (default: 1000)")
parser.add_argument("--max-buffer", type=int, default=20, help="Maximum buffer duration in seconds before forcing a flush (default: 20)") parser.add_argument("--max-buffer", type=int, default=20, help="Maximum buffer duration in seconds before forcing a flush (default: 20)")
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)")
@@ -594,7 +622,7 @@ def main():
parser.add_argument("--glossary-pair", action="append", type=parse_glossary_pair, default=[], help="Term replacement pair SOURCE=TARGET (repeatable)") parser.add_argument("--glossary-pair", action="append", type=parse_glossary_pair, default=[], help="Term replacement pair SOURCE=TARGET (repeatable)")
parser.add_argument("--post-correct", action="store_true", help="Enable v2 finalized English caption post-correction") parser.add_argument("--post-correct", action="store_true", help="Enable v2 finalized English caption post-correction")
parser.add_argument("--post-correct-llm", action="store_true", help="Use local LLM for post-correction (fallback to rules on failure)") parser.add_argument("--post-correct-llm", action="store_true", help="Use local LLM for post-correction (fallback to rules on failure)")
parser.add_argument("--post-correct-model", type=str, default="llama3.1:8b-instruct", help="Local Ollama model for post-correction") parser.add_argument("--post-correct-model", type=str, default="qwen:2b", help="Local Ollama model for post-correction")
parser.add_argument("--post-correct-ollama-url", type=str, default="http://127.0.0.1:11434/api/generate", help="Ollama generate endpoint for local post-correction") parser.add_argument("--post-correct-ollama-url", type=str, default="http://127.0.0.1:11434/api/generate", help="Ollama generate endpoint for local post-correction")
parser.add_argument("--post-correct-llm-timeout", type=float, default=8.0, help="Timeout seconds for local LLM post-correction (default: 8.0)") parser.add_argument("--post-correct-llm-timeout", type=float, default=8.0, help="Timeout seconds for local LLM post-correction (default: 8.0)")
parser.add_argument("--post-correct-keep-alive", type=str, default="30m", help="Ollama keep_alive duration to keep the model loaded (default: 30m)") parser.add_argument("--post-correct-keep-alive", type=str, default="30m", help="Ollama keep_alive duration to keep the model loaded (default: 30m)")
@@ -602,12 +630,10 @@ def main():
parser.add_argument("--skip-post-correct-warmup", action="store_true", help="Skip startup warmup call for the post-correct LLM") parser.add_argument("--skip-post-correct-warmup", action="store_true", help="Skip startup warmup call for the post-correct LLM")
parser.add_argument("--post-correct-min-overlap", type=float, default=0.45, help="Minimum token overlap ratio required to accept LLM rewrite (default: 0.45)") parser.add_argument("--post-correct-min-overlap", type=float, default=0.45, help="Minimum token overlap ratio required to accept LLM rewrite (default: 0.45)")
parser.add_argument("--post-correct-debug", action="store_true", help="Print pre/post LLM correction text for debugging") parser.add_argument("--post-correct-debug", action="store_true", help="Print pre/post LLM correction text for debugging")
parser.add_argument("--llm-paragraph", action="store_true", help="Enable LLM-based paragraph structuring.")
parser.add_argument("--llm-merge-decider", action="store_true", help="Use local LLM to validate/refine smart-correct line merges") parser.add_argument("--llm-merge-decider", action="store_true", help="Use local LLM to validate/refine smart-correct line merges")
parser.add_argument("--llm-merge-timeout", type=float, default=1.5, help="Timeout seconds for LLM merge decision (default: 1.5)") parser.add_argument("--llm-merge-timeout", type=float, default=1.5, help="Timeout seconds for LLM merge decision (default: 1.5)")
parser.add_argument("--speaker-change-detect", action="store_true", help="Enable lightweight speaker-change detection to force caption line cuts") parser.add_argument("--speaker-diarization", action="store_true", help="Enable speaker diarization using pyannote.audio.")
parser.add_argument("--speaker-sim-threshold", type=float, default=0.72, help="Cosine similarity threshold below which a speaker change is assumed (default: 0.72)")
parser.add_argument("--speaker-min-buffer", type=float, default=1.6, help="Minimum buffered speech seconds before speaker-change checks (default: 1.6)")
parser.add_argument("--speaker-check-interval", type=float, default=0.6, help="Seconds between speaker-change checks while buffering (default: 0.6)")
parser.add_argument("--session-log-file", type=str, default="transcribe_session.log", help="Path to session debug log file (default: transcribe_session.log)") parser.add_argument("--session-log-file", type=str, default="transcribe_session.log", help="Path to session debug log file (default: transcribe_session.log)")
parser.add_argument("--disable-session-log", action="store_true", help="Disable writing session debug logs to file") parser.add_argument("--disable-session-log", action="store_true", help="Disable writing session debug logs to file")
@@ -646,15 +672,6 @@ def main():
print("[SYSTEM]: --llm-merge-decider requires --post-correct-llm. Falling back to heuristic merge.") print("[SYSTEM]: --llm-merge-decider requires --post-correct-llm. Falling back to heuristic merge.")
args.llm_merge_decider = False args.llm_merge_decider = False
log_session("[SYSTEM]: --llm-merge-decider requires --post-correct-llm. Falling back to heuristic merge.") log_session("[SYSTEM]: --llm-merge-decider requires --post-correct-llm. Falling back to heuristic merge.")
if args.speaker_change_detect:
print(
"Speaker-change detection enabled "
f"(threshold={args.speaker_sim_threshold}, min_buffer={args.speaker_min_buffer}s)."
)
log_session(
"Speaker-change detection enabled "
f"(threshold={args.speaker_sim_threshold}, min_buffer={args.speaker_min_buffer}s)."
)
if args.quantize: if args.quantize:
WHISPER_MODEL = "mlx-community/whisper-small-mlx-4bit" WHISPER_MODEL = "mlx-community/whisper-small-mlx-4bit"
@@ -670,6 +687,26 @@ def main():
print(f"Using device: {device}") print(f"Using device: {device}")
log_session(f"Using device: {device}") log_session(f"Using device: {device}")
diarization_pipeline = None
if args.speaker_diarization:
print("Loading speaker diarization pipeline...")
hf_token = os.environ.get("HF_TOKEN")
if not hf_token:
print("[WARN] Hugging Face token not found (HF_TOKEN env variable). Diarization may fail if model is private.")
log_session("[WARN] Hugging Face token not found (HF_TOKEN env variable).")
try:
diarization_pipeline = Pipeline.from_pretrained(
"pyannote/speaker-diarization-3.1"
)
diarization_pipeline.to(torch.device(device))
print("Speaker diarization pipeline loaded.")
log_session("Speaker diarization pipeline loaded.")
except Exception as e:
print(f"[ERROR] Could not load speaker diarization pipeline: {e}")
log_session(f"[ERROR] Could not load speaker diarization pipeline: {e}")
diarization_pipeline = None
# Only start ingest thread if enabled # Only start ingest thread if enabled
if args.ingest: if args.ingest:
threading.Thread(target=ingest_worker, daemon=True).start() threading.Thread(target=ingest_worker, daemon=True).start()
@@ -743,9 +780,8 @@ def main():
last_draft_text = "" last_draft_text = ""
draft_line_active = False draft_line_active = False
recent_en_lines = deque(maxlen=2) recent_en_lines = deque(maxlen=2)
llm_post_correct_state = {"available": True, "warned": False, "merge_warned": False} llm_post_correct_state = {"available": True, "warned": False, "merge_warned": False, "paragraph_warned": False}
speaker_reference_signature = None english_paragraph_buffer = []
last_speaker_check_time = 0.0
try: try:
with sd.InputStream(samplerate=SAMPLERATE, channels=args.channels, callback=callback, blocksize=BLOCK_SIZE, device=device_index): with sd.InputStream(samplerate=SAMPLERATE, channels=args.channels, callback=callback, blocksize=BLOCK_SIZE, device=device_index):
@@ -785,7 +821,6 @@ def main():
rolling_context = "" rolling_context = ""
rolling_context_en = "" rolling_context_en = ""
recent_en_lines.clear() recent_en_lines.clear()
speaker_reference_signature = None
last_draft_text = "" last_draft_text = ""
last_change_time = time.time() last_change_time = time.time()
if args.stream: if args.stream:
@@ -797,30 +832,10 @@ def main():
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)
force_flush_for_speaker_change = False
if args.speaker_change_detect and buffer_duration >= args.speaker_min_buffer:
now = time.time()
if (now - last_speaker_check_time) >= args.speaker_check_interval:
last_speaker_check_time = now
tail_secs = min(2.0, buffer_duration)
tail_samples = int(tail_secs * SAMPLERATE)
recent_signature = compute_speaker_signature(current_audio[-tail_samples:], SAMPLERATE)
if speaker_reference_signature is None and recent_signature is not None:
speaker_reference_signature = recent_signature
elif recent_signature is not None:
similarity = cosine_similarity(speaker_reference_signature, recent_signature)
if similarity is not None and similarity < args.speaker_sim_threshold:
force_flush_for_speaker_change = True
print(
f"\n[SYSTEM]: Speaker change detected (similarity={similarity:.2f}). "
"Cutting current line."
)
if ( if (
(buffer_len_samples - last_end) > (SAMPLERATE * args.silence / 1000) (buffer_len_samples - last_end) > (SAMPLERATE * args.silence / 1000)
or buffer_len_samples > buffer_limit or buffer_len_samples > buffer_limit
or force_flush_for_speaker_change
): ):
# Clear draft line if it was used # Clear draft line if it was used
@@ -830,7 +845,7 @@ def main():
draft_line_active = False draft_line_active = False
# Prepare transcription kwargs # Prepare transcription kwargs
transcribe_kwargs = {"path_or_hf_repo": WHISPER_MODEL} transcribe_kwargs = {"path_or_hf_repo": WHISPER_MODEL, "word_timestamps": args.speaker_diarization}
if args.lang: if args.lang:
transcribe_kwargs["language"] = args.lang transcribe_kwargs["language"] = args.lang
if args.context and rolling_context: if args.context and rolling_context:
@@ -841,9 +856,46 @@ def main():
transcribe_kwargs["logprob_threshold"] = args.logprob_threshold transcribe_kwargs["logprob_threshold"] = args.logprob_threshold
transcribe_kwargs["compression_ratio_threshold"] = args.compression_threshold transcribe_kwargs["compression_ratio_threshold"] = args.compression_threshold
# Speaker diarization
diarization_result = None
if diarization_pipeline:
log_session("Performing speaker diarization...")
audio_for_diarization = torch.from_numpy(current_audio).float().unsqueeze(0)
try:
diarization_result = diarization_pipeline({"waveform": audio_for_diarization, "sample_rate": SAMPLERATE})
except Exception as e:
print(f"\n[ERROR] Diarization failed: {e}")
log_session(f"Diarization failed: {e}")
# 1. Transcribe & Detect Language # 1. Transcribe & Detect Language
transcription_result = transcribe_with_controls(current_audio, transcribe_kwargs) transcription_result = transcribe_with_controls(current_audio, transcribe_kwargs)
original_text = transcription_result['text'].strip()
original_text_for_print = ""
if diarization_result:
segments_with_speakers = assign_speakers_to_segments(
transcription_result.get('segments', []),
diarization_result
)
transcription_result['segments'] = segments_with_speakers
# Reconstruct text with speaker labels for printing
text_parts_for_print = []
for i, segment in enumerate(segments_with_speakers):
speaker = segment.get('speaker', 'UNKNOWN')
text = segment['text'].strip()
if i > 0 and speaker == segments_with_speakers[i-1].get('speaker', 'UNKNOWN'):
text_parts_for_print.append(f" {text}")
else:
if i > 0: text_parts_for_print.append("\n")
text_parts_for_print.append(f"[{speaker}]: {text}")
original_text_for_print = "".join(text_parts_for_print)
# Reconstruct text without speaker labels for processing
original_text = " ".join([seg['text'].strip() for seg in transcription_result['segments']])
else:
original_text = transcription_result['text'].strip()
original_text_for_print = original_text
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')
avg_logprob, max_compression = extract_whisper_quality(transcription_result) avg_logprob, max_compression = extract_whisper_quality(transcription_result)
@@ -873,10 +925,20 @@ def main():
rolling_context = "" rolling_context = ""
rolling_context_en = "" rolling_context_en = ""
recent_en_lines.clear() recent_en_lines.clear()
speaker_reference_signature = None
if original_text: if original_text:
print_caption(detected_lang, original_text, leading_newline=True) if args.detect_auto_translation and args.lang and args.lang != 'en' and len(original_text.split()) > 3:
try:
detected_lang_for_segment = detect(original_text)
if detected_lang_for_segment == 'en':
print_caption("AUTO-EN", original_text, leading_newline=True)
if args.context:
rolling_context_en = (rolling_context_en + " " + original_text)[-200:].strip()
except LangDetectException:
# Could not detect language, probably too short
pass
print_caption(detected_lang, original_text_for_print, leading_newline=True)
last_change_time = time.time() # Successfully transcribed full segment last_change_time = time.time() # Successfully transcribed full segment
# Prepare payload # Prepare payload
@@ -957,10 +1019,44 @@ def main():
else: else:
recent_en_lines.append({"text": english_for_caption, "ts": now_ts}) recent_en_lines.append({"text": english_for_caption, "ts": now_ts})
if args.llm_paragraph:
english_paragraph_buffer.append(english_for_caption)
if english_for_caption.strip().endswith((".", "?", "!")):
full_text = "\n".join(english_paragraph_buffer)
structured_paragraph = structure_paragraph_with_llm(full_text, args, llm_post_correct_state)
print(f"\n[LLM Paragraph]:\n{structured_paragraph}")
log_session(f"[LLM-P]: {structured_paragraph}")
# Translate the paragraph here
if structured_paragraph and translation_engines:
english_chunks = split_text_for_translation(structured_paragraph, max_chars=args.mt_max_chars)
for lang_key, (model, tokenizer) in translation_engines.items():
translated_parts = []
for chunk in english_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,
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()
print_caption(lang_key, translated_text)
english_paragraph_buffer = [] # Reset buffer
if args.en and english_for_caption: if args.en and english_for_caption:
payload["en"] = english_for_caption payload["en"] = english_for_caption
print(f"[EN]: {english_for_caption}") # When llm_paragraph is enabled, we don't print the line-by-line English captions.
log_session(f"[EN]: {english_for_caption}") if not args.llm_paragraph:
print(f"[EN]: {english_for_caption}")
log_session(f"[EN]: {english_for_caption}")
# Update rolling context for next segment # Update rolling context for next segment
if args.context: if args.context:
@@ -968,7 +1064,7 @@ def main():
rolling_context = (rolling_context + " " + original_text)[-200:].strip() rolling_context = (rolling_context + " " + original_text)[-200:].strip()
# 3. Translate from English to other languages # 3. Translate from English to other languages
if english_for_translation and translation_engines: 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) english_chunks = split_text_for_translation(english_for_translation, max_chars=args.mt_max_chars)
for lang_key, (model, tokenizer) in translation_engines.items(): for lang_key, (model, tokenizer) in translation_engines.items():
@@ -1007,7 +1103,6 @@ def main():
last_stream_time = time.time() last_stream_time = time.time()
stuck_draft_count = 0 stuck_draft_count = 0
last_draft_text = "" last_draft_text = ""
speaker_reference_signature = None
draft_line_active = False draft_line_active = False
elif args.stream and (time.time() - last_stream_time) > 1.0: elif args.stream and (time.time() - last_stream_time) > 1.0:
@@ -1064,7 +1159,6 @@ def main():
elif not speech_started and len(current_audio) > SAMPLERATE * 2: elif not speech_started and len(current_audio) > SAMPLERATE * 2:
audio_buffer = [] audio_buffer = []
speaker_reference_signature = None
except KeyboardInterrupt: except KeyboardInterrupt:
print("\nStopped by user.") print("\nStopped by user.")