feat: refactor transcription, translation, and distribution into multiple processes
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user