feat: refactor transcription, translation, and distribution into multiple processes
This commit is contained in:
+12
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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
@@ -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()
|
||||||
@@ -9,3 +9,7 @@ pyinstaller
|
|||||||
sacremoses
|
sacremoses
|
||||||
joblib
|
joblib
|
||||||
sentencepiece
|
sentencepiece
|
||||||
|
pyannote.audio
|
||||||
|
soundfile
|
||||||
|
langdetect
|
||||||
|
python-dotenv
|
||||||
|
|||||||
+186
-92
@@ -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.")
|
||||||
|
|||||||
Reference in New Issue
Block a user