diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 0de36f0..66244e4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -35,7 +35,7 @@ jobs: ${{ runner.os }}-pip-${{ matrix.python-version }}- - name: Install system dependencies - run: sudo apt-get update && sudo apt-get install -y ffmpeg portaudio19-dev + run: sudo apt-get update && sudo apt-get install -y portaudio19-dev - name: Install Python dependencies run: | @@ -180,7 +180,7 @@ jobs: ubuntu-latest-pip-3.8- - name: Install system dependencies - run: sudo apt-get update && sudo apt-get install -y ffmpeg portaudio19-dev + run: sudo apt-get update && sudo apt-get install -y portaudio19-dev - name: Install Python dependencies run: | diff --git a/README.md b/README.md index dd93eb7..173306f 100644 --- a/README.md +++ b/README.md @@ -12,7 +12,7 @@ to convert speech input into text output. It can be used to transcribe both live input from microphone and pre-recorded audio files. ## Installation -- Install PyAudio and ffmpeg +- Install PyAudio ```bash bash scripts/setup.sh ``` diff --git a/docker/Dockerfile.tensorrt b/docker/Dockerfile.tensorrt index 3b9b125..c46a2d2 100644 --- a/docker/Dockerfile.tensorrt +++ b/docker/Dockerfile.tensorrt @@ -25,6 +25,7 @@ RUN apt update && bash setup.sh && rm setup.sh COPY requirements/server.txt . RUN pip install --no-cache-dir -r server.txt && rm server.txt +RUN pip install pynvml==11.5.0 COPY whisper_live ./whisper_live COPY scripts/build_whisper_tensorrt.sh . COPY run_server.py . \ No newline at end of file diff --git a/requirements/client.txt b/requirements/client.txt index 359cc04..2576140 100644 --- a/requirements/client.txt +++ b/requirements/client.txt @@ -1,4 +1,4 @@ PyAudio -ffmpeg-python +av scipy websocket-client \ No newline at end of file diff --git a/requirements/server.txt b/requirements/server.txt index 5298473..f6cd0b6 100644 --- a/requirements/server.txt +++ b/requirements/server.txt @@ -4,8 +4,8 @@ onnxruntime==1.16.0 numba kaldialign soundfile -ffmpeg-python scipy +av jiwer evaluate numpy<2 diff --git a/scripts/setup.sh b/scripts/setup.sh index 1b5cb53..8ae7036 100644 --- a/scripts/setup.sh +++ b/scripts/setup.sh @@ -1,3 +1,3 @@ #! /bin/bash -apt-get install portaudio19-dev ffmpeg wget -y +apt-get install portaudio19-dev wget -y diff --git a/setup.py b/setup.py index 714322a..96e4c4f 100644 --- a/setup.py +++ b/setup.py @@ -48,7 +48,6 @@ setup( "torchaudio", "websockets", "onnxruntime==1.16.0", - "ffmpeg-python", "scipy", "websocket-client", "numba", diff --git a/whisper_live/client.py b/whisper_live/client.py index 4045ce2..7cf8fca 100644 --- a/whisper_live/client.py +++ b/whisper_live/client.py @@ -10,7 +10,7 @@ import json import websocket import uuid import time -import ffmpeg +import av import whisper_live.utils as utils @@ -421,84 +421,83 @@ class TranscriptionTeeClient: def process_rtsp_stream(self, rtsp_url): """ - Connect to an RTSP source, process the audio stream, and send it for trascription. + Connect to an RTSP source, process the audio stream, and send it for transcription. Args: rtsp_url (str): The URL of the RTSP stream source. """ - process = self.get_rtsp_ffmpeg_process(rtsp_url) - self.handle_ffmpeg_process(process, stream_type='RTSP') + print("[INFO]: Connecting to RTSP stream...") + try: + container = av.open(rtsp_url, format="rtsp", options={"rtsp_transport": "tcp"}) + self.process_av_stream(container, stream_type="RTSP") + except Exception as e: + print(f"[ERROR]: Failed to process RTSP stream: {e}") + finally: + for client in self.clients: + client.wait_before_disconnect() + self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True) + self.close_all_clients() + self.write_all_clients_srt() + print("[INFO]: RTSP stream processing finished.") - def process_hls_stream(self, hls_url, save_file): + def process_hls_stream(self, hls_url, save_file=None): """ Connect to an HLS source, process the audio stream, and send it for transcription. Args: hls_url (str): The URL of the HLS stream source. - save_file (str, optional): Local path to save the network stream. + save_file (str, optional): Local path to save the network stream. """ - process = self.get_hls_ffmpeg_process(hls_url, save_file) - self.handle_ffmpeg_process(process, stream_type='HLS') - - def handle_ffmpeg_process(self, process, stream_type): - print(f"[INFO]: Connecting to {stream_type} stream...") - stderr_thread = threading.Thread(target=self.consume_stderr, args=(process,)) - stderr_thread.start() + print("[INFO]: Connecting to HLS stream...") try: - # Process the stream - while True: - in_bytes = process.stdout.read(self.chunk * 2) # 2 bytes per sample - if not in_bytes: - break - audio_array = self.bytes_to_float_array(in_bytes) - self.multicast_packet(audio_array.tobytes()) - + container = av.open(hls_url, format="hls") + self.process_av_stream(container, stream_type="HLS", save_file=save_file) except Exception as e: - print(f"[ERROR]: Failed to connect to {stream_type} stream: {e}") + print(f"[ERROR]: Failed to process HLS stream: {e}") finally: + for client in self.clients: + client.wait_before_disconnect() + self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True) self.close_all_clients() self.write_all_clients_srt() - if process: - process.kill() + print("[INFO]: HLS stream processing finished.") - print(f"[INFO]: {stream_type} stream processing finished.") - - def get_rtsp_ffmpeg_process(self, rtsp_url): - return ( - ffmpeg - .input(rtsp_url, threads=0) - .output('-', format='s16le', acodec='pcm_s16le', ac=1, ar=self.rate) - .run_async(pipe_stdout=True, pipe_stderr=True) - ) - - def get_hls_ffmpeg_process(self, hls_url, save_file): - if save_file is None: - process = ( - ffmpeg - .input(hls_url, threads=0) - .output('-', format='s16le', acodec='pcm_s16le', ac=1, ar=self.rate) - .run_async(pipe_stdout=True, pipe_stderr=True) - ) - else: - input = ffmpeg.input(hls_url, threads=0) - output_file = input.output(save_file, acodec='copy', vcodec='copy').global_args('-loglevel', 'quiet') - output_std = input.output('-', format='s16le', acodec='pcm_s16le', ac=1, ar=self.rate) - process = ( - ffmpeg.merge_outputs(output_file, output_std) - .run_async(pipe_stdout=True, pipe_stderr=True) - ) - - return process - - def consume_stderr(self, process): + def process_av_stream(self, container, stream_type, save_file=None): """ - Consume and log the stderr output of a process in a separate thread. + Process an AV container stream and send audio packets to the server. Args: - process (subprocess.Popen): The process whose stderr output will be logged. + container (av.container.InputContainer): The input container to process. + stream_type (str): The type of stream being processed ("RTSP" or "HLS"). + save_file (str, optional): Local path to save the stream. Default is None. """ - for line in iter(process.stderr.readline, b""): - logging.debug(f'[STDERR]: {line.decode()}') + audio_stream = next((s for s in container.streams if s.type == "audio"), None) + if not audio_stream: + print(f"[ERROR]: No audio stream found in {stream_type} source.") + return + + output_container = None + if save_file: + output_container = av.open(save_file, mode="w") + output_audio_stream = output_container.add_stream(codec_name="pcm_s16le", rate=self.rate) + + try: + for packet in container.demux(audio_stream): + for frame in packet.decode(): + audio_data = frame.to_ndarray().tobytes() + self.multicast_packet(audio_data) + + if save_file: + output_container.mux(frame) + except Exception as e: + print(f"[ERROR]: Error during {stream_type} stream processing: {e}") + finally: + # Wait for server to send any leftover transcription. + time.sleep(5) + self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True) + if output_container: + output_container.close() + container.close() def save_chunk(self, n_audio_file): """ diff --git a/whisper_live/tensorrt_utils.py b/whisper_live/tensorrt_utils.py index 9752e7a..01631d9 100644 --- a/whisper_live/tensorrt_utils.py +++ b/whisper_live/tensorrt_utils.py @@ -23,8 +23,12 @@ from typing import Dict, Iterable, List, Optional, TextIO, Tuple, Union import kaldialign import numpy as np import soundfile +import av +import wave import torch import torch.nn.functional as F +from whisper_live.utils import resample + Pathlike = Union[str, Path] @@ -35,38 +39,33 @@ CHUNK_LENGTH = 30 N_SAMPLES = CHUNK_LENGTH * SAMPLE_RATE # 480000 samples in a 30-second chunk -def load_audio(file: str, sr: int = SAMPLE_RATE): +def load_audio(file: str, sr: int = 16000): """ - Open an audio file and read as mono waveform, resampling as necessary + Open an audio file, resample it, and read as a mono waveform. Parameters ---------- file: str - The audio file to open + The audio file to open. sr: int - The sample rate to resample the audio if necessary + The sample rate to resample the audio if necessary. Returns ------- A NumPy array containing the audio waveform, in float32 dtype. """ + resampled_file = resample(file, sr) - # This launches a subprocess to decode audio while down-mixing - # and resampling as necessary. Requires the ffmpeg CLI in PATH. - # fmt: off - cmd = [ - "ffmpeg", "-nostdin", "-threads", "0", "-i", file, "-f", "s16le", "-ac", - "1", "-acodec", "pcm_s16le", "-ar", - str(sr), "-" - ] - # fmt: on - try: - out = run(cmd, capture_output=True, check=True).stdout - except CalledProcessError as e: - raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e + with wave.open(resampled_file, "rb") as wav_file: + num_frames = wav_file.getnframes() + raw_data = wav_file.readframes(num_frames) - return np.frombuffer(out, np.int16).flatten().astype(np.float32) / 32768.0 + audio_data = np.frombuffer(raw_data, dtype=np.int16) + + audio_data = audio_data.astype(np.float32) / 32768.0 + + return audio_data def load_audio_wav_format(wav_path): diff --git a/whisper_live/utils.py b/whisper_live/utils.py index d32105a..1f9b2ad 100644 --- a/whisper_live/utils.py +++ b/whisper_live/utils.py @@ -1,8 +1,9 @@ import os import textwrap import scipy -import ffmpeg import numpy as np +import av +from pathlib import Path def clear_screen(): @@ -26,8 +27,8 @@ def format_time(s): return f"{hours:02}:{minutes:02}:{seconds:02},{milliseconds:03}" -def create_srt_file(segments, output_file): - with open(output_file, 'w', encoding='utf-8') as srt_file: +def create_srt_file(segments, resampled_file): + with open(resampled_file, 'w', encoding='utf-8') as srt_file: segment_number = 1 for segment in segments: start_time = format_time(float(segment['start'])) @@ -43,9 +44,7 @@ def create_srt_file(segments, output_file): def resample(file: str, sr: int = 16000): """ - # https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22 - Open an audio file and read as mono waveform, resampling as necessary, - save the resampled audio + Resample the audio file to 16kHz. Args: file (str): The audio file to open @@ -54,18 +53,30 @@ def resample(file: str, sr: int = 16000): Returns: resampled_file (str): The resampled audio file """ - try: - # This launches a subprocess to decode audio while down-mixing and resampling as necessary. - # Requires the ffmpeg CLI and `ffmpeg-python` package to be installed. - out, _ = ( - ffmpeg.input(file, threads=0) - .output("-", format="s16le", acodec="pcm_s16le", ac=1, ar=sr) - .run(cmd=["ffmpeg", "-nostdin"], capture_stdout=True, capture_stderr=True) - ) - except ffmpeg.Error as e: - raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e - np_buffer = np.frombuffer(out, dtype=np.int16) + container = av.open(file) + stream = next(s for s in container.streams if s.type == 'audio') - resampled_file = f"{file.split('.')[0]}_resampled.wav" - scipy.io.wavfile.write(resampled_file, sr, np_buffer.astype(np.int16)) + resampler = av.AudioResampler( + format='s16', + layout='mono', + rate=sr, + ) + + resampled_file = Path(file).stem + "_resampled.wav" + output_container = av.open(resampled_file, mode='w') + output_stream = output_container.add_stream('pcm_s16le', rate=sr) + output_stream.layout = 'mono' + + for frame in container.decode(audio=0): + frame.pts = None + resampled_frames = resampler.resample(frame) + if resampled_frames is not None: + for resampled_frame in resampled_frames: + for packet in output_stream.encode(resampled_frame): + output_container.mux(packet) + + for packet in output_stream.encode(None): + output_container.mux(packet) + + output_container.close() return resampled_file