Compare commits
120 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4baccf75a7 | |||
| b7acb8c872 | |||
| fe7b55efe4 | |||
| c1b249ad0d | |||
| 5e4589cfe1 | |||
| b6b73730fb | |||
| 953a88c7da | |||
| 182b5cbd6d | |||
| 32ba924d8c | |||
| 450433b07b | |||
| 38bff6a901 | |||
| c936e5f727 | |||
| 18de63c649 | |||
| 7bcd8b9520 | |||
| 19c05c8231 | |||
| 49e232bc4d | |||
| 30617dfd44 | |||
| a55b99c11e | |||
| 2725f1aed9 | |||
| 53c31f3570 | |||
| e65fbcd9fc | |||
| 7f0c7a6791 | |||
| 2eff360b9e | |||
| c25a036c02 | |||
| 446fc6e835 | |||
| a1650eaa4f | |||
| a6523b6b71 | |||
| e275d34943 | |||
| 778a9c5903 | |||
| 0e89573798 | |||
| 8d89de22d8 | |||
| 81c57ae40c | |||
| 00f0ff1112 | |||
| 8b87a0562d | |||
| 617fda2864 | |||
| 0d74790c67 | |||
| 1322dd3c27 | |||
| be71657397 | |||
| a317597f01 | |||
| aaa47cfab5 | |||
| bc070d6688 | |||
| 8e7e329a39 | |||
| 380f07394b | |||
| 30f78a2cc6 | |||
| 01c6bc1ecd | |||
| bdaed45820 | |||
| 4870e9fb9e | |||
| ccb183b4d8 | |||
| fac62aaccc | |||
| aade67736a | |||
| abfe830eee | |||
| cb392cbb93 | |||
| 42733da59a | |||
| 26c517021f | |||
| cf721e8b53 | |||
| 5985ec82b6 | |||
| 2f1c934ea2 | |||
| b220ccb330 | |||
| 5e3906fc7b | |||
| a8b9275013 | |||
| 815441e8bb | |||
| 5b9bc2bc0e | |||
| ee132517fa | |||
| 761bb61e87 | |||
| 14077315ae | |||
| ab17c4dbc6 | |||
| 5e2421118d | |||
| d1de2ec3ce | |||
| 22a37e7843 | |||
| e4579ef291 | |||
| f73a146eb9 | |||
| cfba5b3e54 | |||
| 1ac7a278bb | |||
| 3a96f60006 | |||
| 3c09289dea | |||
| e1a42c22d2 | |||
| 3d043dc906 | |||
| 399e9e7efe | |||
| 9d2ea75247 | |||
| 225a98be0c | |||
| 8f373c3537 | |||
| 61d07edabb | |||
| 03e30e1fed | |||
| c0a947a8f6 | |||
| 819ab35b28 | |||
| 615c9c7aed | |||
| a9683319e0 | |||
| 0a2d92c5b8 | |||
| dccfce2a3c | |||
| f78fc473c5 | |||
| 0dfbdb2477 | |||
| e171b9c460 | |||
| 66b5dc7c15 | |||
| 0f1d36fc06 | |||
| d24c53198c | |||
| fe1640695c | |||
| 8d77f0fa5a | |||
| 2e37216282 | |||
| 7c7a446478 | |||
| 37d7f2ed66 | |||
| 754f22dfae | |||
| ebd2dc9568 | |||
| 9b2e17ec4d | |||
| 5b32dc4130 | |||
| c0f37c77e9 | |||
| 3b15dc76b4 | |||
| 4d477e35e7 | |||
| a17f4041de | |||
| 8a06ba802b | |||
| 02d4566289 | |||
| acd4902bec | |||
| a495a49b06 | |||
| 9e5ab408cd | |||
| 5e6c26c3a0 | |||
| 18b6168807 | |||
| ec1349360a | |||
| a41e714801 | |||
| 2d16ee552f | |||
| 9699611000 | |||
| ea64d47899 |
@@ -77,7 +77,7 @@ jobs:
|
||||
build-and-push-docker-cpu:
|
||||
needs: [run-tests, check-code-format]
|
||||
runs-on: ubuntu-22.04
|
||||
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags')
|
||||
if: github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/'))
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
@@ -99,11 +99,40 @@ jobs:
|
||||
push: true
|
||||
tags: ghcr.io/collabora/whisperlive-cpu:latest
|
||||
|
||||
build-and-push-docker-tensorrt:
|
||||
needs: [run-tests, check-code-format]
|
||||
timeout-minutes: 60
|
||||
runs-on: ubuntu-22.04
|
||||
if: github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/'))
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
- name: Log in to GitHub Container Registry
|
||||
uses: docker/login-action@v1
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
password: ${{ secrets.GHCR_TOKEN }}
|
||||
|
||||
- name: Docker Prune
|
||||
run: docker system prune -af
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v1
|
||||
|
||||
- name: Build and push Docker GPU image
|
||||
uses: docker/build-push-action@v2
|
||||
with:
|
||||
context: .
|
||||
file: docker/Dockerfile.tensorrt
|
||||
push: true
|
||||
tags: ghcr.io/collabora/whisperlive-tensorrt:latest
|
||||
|
||||
build-and-push-docker-gpu:
|
||||
needs: [run-tests, check-code-format, build-and-push-docker-cpu]
|
||||
timeout-minutes: 20
|
||||
runs-on: ubuntu-22.04
|
||||
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags')
|
||||
if: github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/'))
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ python3 run_server.py --port 9090 \
|
||||
|
||||
# running with custom model
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend faster_whisper
|
||||
--backend faster_whisper \
|
||||
-fw "/path/to/custom/faster/whisper/model"
|
||||
```
|
||||
|
||||
@@ -53,10 +53,33 @@ python3 run_server.py -p 9090 \
|
||||
-trt /home/TensorRT-LLM/examples/whisper/whisper_small \
|
||||
-m
|
||||
```
|
||||
#### Controlling OpenMP Threads
|
||||
To control the number of threads used by OpenMP, you can set the `OMP_NUM_THREADS` environment variable. This is useful for managing CPU resources and ensuring consistent performance. If not specified, `OMP_NUM_THREADS` is set to `1` by default. You can change this by using the `--omp_num_threads` argument:
|
||||
```bash
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend faster_whisper \
|
||||
--omp_num_threads 4
|
||||
```
|
||||
|
||||
#### Single model mode
|
||||
By default, when running the server without specifying a model, the server will instantiate a new whisper model for every client connection. This has the advantage, that the server can use different model sizes, based on the client's requested model size. On the other hand, it also means you have to wait for the model to be loaded upon client connection and you will have increased (V)RAM usage.
|
||||
|
||||
When serving a custom TensorRT model using the `-trt` or a custom faster_whisper model using the `-fw` option, the server will instead only instantiate the custom model once and then reuse it for all client connections.
|
||||
|
||||
If you don't want this, set `--no_single_model`.
|
||||
|
||||
|
||||
### Running the Client
|
||||
- Initializing the client:
|
||||
- Initializing the client with below parameters:
|
||||
- `lang`: Language of the input audio, applicable only if using a multilingual model.
|
||||
- `translate`: If set to `True` then translate from any language to `en`.
|
||||
- `model`: Whisper model size.
|
||||
- `use_vad`: Whether to use `Voice Activity Detection` on the server.
|
||||
- `save_output_recording`: Set to True to save the microphone input as a `.wav` file during live transcription. This option is helpful for recording sessions for later playback or analysis. Defaults to `False`.
|
||||
- `output_recording_filename`: Specifies the `.wav` file path where the microphone input will be saved if `save_output_recording` is set to `True`.
|
||||
- `max_clients`: Specifies the maximum number of clients the server should allow. Defaults to 4.
|
||||
- `max_connection_time`: Maximum connection time for each client in seconds. Defaults to 600.
|
||||
|
||||
```python
|
||||
from whisper_live.client import TranscriptionClient
|
||||
client = TranscriptionClient(
|
||||
@@ -64,13 +87,17 @@ client = TranscriptionClient(
|
||||
9090,
|
||||
lang="en",
|
||||
translate=False,
|
||||
model="small",
|
||||
model="small", # also support hf_model => `Systran/faster-whisper-small`
|
||||
use_vad=False,
|
||||
save_output_recording=True, # Only used for microphone input, False by Default
|
||||
output_recording_filename="./output_recording.wav", # Only used for microphone input
|
||||
max_clients=4,
|
||||
max_connection_time=600
|
||||
)
|
||||
```
|
||||
It connects to the server running on localhost at port 9090. Using a multilingual model, language for the transcription will be automatically detected. You can also use the language option to specify the target language for the transcription, in this case, English ("en"). The translate option should be set to `True` if we want to translate from the source language to English and `False` if we want to transcribe in the source language.
|
||||
|
||||
- Trancribe an audio file:
|
||||
- Transcribe an audio file:
|
||||
```python
|
||||
client("tests/jfk.wav")
|
||||
```
|
||||
@@ -80,9 +107,14 @@ client("tests/jfk.wav")
|
||||
client()
|
||||
```
|
||||
|
||||
- To transcribe from a RTSP stream:
|
||||
```python
|
||||
client(rtsp_url="rtsp://admin:admin@192.168.0.1/rtsp")
|
||||
```
|
||||
|
||||
- To transcribe from a HLS stream:
|
||||
```python
|
||||
client(hls_url="http://as-hls-ww-live.akamaized.net/pool_904/live/ww/bbc_1xtra/bbc_1xtra.isml/bbc_1xtra-audio%3d96000.norewind.m3u8")
|
||||
client(hls_url="http://as-hls-ww-live.akamaized.net/pool_904/live/ww/bbc_1xtra/bbc_1xtra.isml/bbc_1xtra-audio%3d96000.norewind.m3u8")
|
||||
```
|
||||
|
||||
## Browser Extensions
|
||||
@@ -96,7 +128,22 @@ client(hls_url="http://as-hls-ww-live.akamaized.net/pool_904/live/ww/bbc_1xtra/b
|
||||
docker run -it --gpus all -p 9090:9090 ghcr.io/collabora/whisperlive-gpu:latest
|
||||
```
|
||||
|
||||
- TensorRT. Follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) in order to setup docker and use TensorRT backend. We provide a pre-built docker image which has TensorRT-LLM built and ready to use.
|
||||
- TensorRT.
|
||||
```bash
|
||||
docker run -p 9090:9090 --runtime=nvidia --gpus all --entrypoint /bin/bash -it ghcr.io/collabora/whisperlive-tensorrt
|
||||
|
||||
# Build small.en engine
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en # float16
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en int8 # int8 weight only quantization
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en int4 # int4 weight only quantization
|
||||
|
||||
# Run server with small.en
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend tensorrt \
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_en_float16"
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_en_int8"
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_en_int4"
|
||||
```
|
||||
|
||||
- CPU
|
||||
```bash
|
||||
|
||||
+11
-40
@@ -1,67 +1,38 @@
|
||||
# Whisper-TensorRT
|
||||
# WhisperLive-TensorRT
|
||||
We have only tested the TensorRT backend in docker so, we recommend docker for a smooth TensorRT backend setup.
|
||||
**Note**: We use [our fork to setup TensorRT](https://github.com/makaveli10/TensorRT-LLM)
|
||||
**Note**: We use `tensorrt_llm==0.15.0.dev2024111200`
|
||||
|
||||
## Installation
|
||||
- Install [docker](https://docs.docker.com/engine/install/)
|
||||
- Install [nvidia-container-toolkit](https://docs.nvidia.com/datacenter/cloud-native/container-toolkit/latest/install-guide.html)
|
||||
|
||||
- Clone this repo.
|
||||
- Run WhisperLive TensorRT in docker
|
||||
```bash
|
||||
git clone https://github.com/collabora/WhisperLive.git
|
||||
cd WhisperLive
|
||||
docker run -p 9090:9090 --runtime=nvidia --gpus all --entrypoint /bin/bash -it ghcr.io/collabora/whisperlive-tensorrt:latest
|
||||
```
|
||||
|
||||
- Pull the TensorRT-LLM docker image which we prebuilt for WhisperLive TensorRT backend.
|
||||
```bash
|
||||
docker pull ghcr.io/collabora/whisperbot-base:latest
|
||||
```
|
||||
|
||||
- Next, we run the docker image and mount WhisperLive repo to the containers `/home` directory.
|
||||
```bash
|
||||
docker run -it --gpus all --shm-size=8g \
|
||||
--ipc=host --ulimit memlock=-1 --ulimit stack=67108864 \
|
||||
-p 9090:9090 -v /path/to/WhisperLive:/home/WhisperLive \
|
||||
ghcr.io/collabora/whisperbot-base:latest
|
||||
```
|
||||
|
||||
- Make sure to test the installation.
|
||||
```bash
|
||||
# export ENV=${ENV:-/etc/shinit_v2}
|
||||
# source $ENV
|
||||
python -c "import torch; import tensorrt; import tensorrt_llm"
|
||||
```
|
||||
**NOTE**: Uncomment and update library paths if imports fail.
|
||||
|
||||
## Whisper TensorRT Engine
|
||||
- We build `small.en` and `small` multilingual TensorRT engine. The script logs the path of the directory with Whisper TensorRT engine. We need the model_path to run the server.
|
||||
- We build `small.en` and `small` multilingual TensorRT engine as examples below. The script logs the path of the directory with Whisper TensorRT engine. We need that model_path to run the server.
|
||||
```bash
|
||||
# convert small.en
|
||||
bash scripts/build_whisper_tensorrt.sh /root/TensorRT-LLM-examples small.en
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en # float16
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en int8 # int8 weight only quantization
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small.en int4 # int4 weight only quantization
|
||||
|
||||
# convert small multilingual model
|
||||
bash scripts/build_whisper_tensorrt.sh /root/TensorRT-LLM-examples small
|
||||
bash build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small
|
||||
```
|
||||
|
||||
## Run WhisperLive Server with TensorRT Backend
|
||||
```bash
|
||||
cd /home/WhisperLive
|
||||
|
||||
# Install requirements
|
||||
apt update && bash scripts/setup.sh
|
||||
pip install -r requirements/server.txt
|
||||
|
||||
# Required to create mel spectogram
|
||||
wget --directory-prefix=assets assets/mel_filters.npz https://raw.githubusercontent.com/openai/whisper/main/whisper/assets/mel_filters.npz
|
||||
|
||||
# Run English only model
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend tensorrt \
|
||||
--trt_model_path "path/to/whisper_trt/from/build/step"
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_en_float16"
|
||||
|
||||
# Run Multilingual model
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend tensorrt \
|
||||
--trt_model_path "path/to/whisper_trt/from/build/step" \
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_float16" \
|
||||
--trt_multilingual
|
||||
```
|
||||
|
||||
+13
-12
@@ -1,22 +1,23 @@
|
||||
FROM python:3.8-slim-buster
|
||||
FROM python:3.10-bookworm
|
||||
|
||||
ARG DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
ca-certificates \
|
||||
sudo \
|
||||
git \
|
||||
bzip2 \
|
||||
libx11-6 \
|
||||
&& apt-get clean \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
# install lib required for pyaudio
|
||||
RUN apt update && apt install -y portaudio19-dev && apt-get clean && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# update pip to support for whl.metadata -> less downloading
|
||||
RUN pip install --no-cache-dir -U "pip>=24"
|
||||
|
||||
# create a working directory
|
||||
RUN mkdir /app
|
||||
WORKDIR /app
|
||||
|
||||
COPY scripts/setup.sh requirements/server.txt /app/
|
||||
# install pytorch, but without the nvidia-libs that are only necessary for gpu
|
||||
RUN pip install --no-cache-dir torch --index-url https://download.pytorch.org/whl/cpu
|
||||
|
||||
RUN apt update && bash setup.sh && pip install -r server.txt
|
||||
# install the requirements for running the whisper-live server
|
||||
COPY requirements/server.txt /app/
|
||||
RUN pip install --no-cache-dir -r server.txt && rm server.txt
|
||||
|
||||
COPY whisper_live /app/whisper_live
|
||||
COPY run_server.py /app
|
||||
|
||||
+13
-20
@@ -1,33 +1,26 @@
|
||||
FROM nvidia/cuda:11.8.0-cudnn8-runtime-ubuntu22.04
|
||||
FROM python:3.10-bookworm
|
||||
|
||||
ARG DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Remove any third-party apt sources to avoid issues with expiring keys.
|
||||
RUN rm -f /etc/apt/sources.list.d/*.list
|
||||
# install lib required for pyaudio
|
||||
RUN apt update && apt install -y portaudio19-dev && apt-get clean && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install some basic utilities.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
ca-certificates \
|
||||
sudo \
|
||||
git \
|
||||
bzip2 \
|
||||
libx11-6 \
|
||||
python3-dev \
|
||||
python3-pip \
|
||||
&& python3 -m pip install --upgrade pip \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
# update pip to support for whl.metadata -> less downloading
|
||||
RUN pip install --no-cache-dir -U "pip>=24"
|
||||
|
||||
# Create a working directory.
|
||||
# create a working directory
|
||||
RUN mkdir /app
|
||||
WORKDIR /app
|
||||
|
||||
COPY scripts/setup.sh requirements/server.txt /app
|
||||
# install the requirements for running the whisper-live server
|
||||
COPY requirements/server.txt /app/
|
||||
RUN pip install --no-cache-dir -r server.txt && rm server.txt
|
||||
|
||||
RUN apt update && bash setup.sh && rm setup.sh
|
||||
RUN pip install -r server.txt && rm server.txt
|
||||
# make the paths of the nvidia libs installed as wheels visible. equivalent to:
|
||||
# export LD_LIBRARY_PATH=`python3 -c 'import os; import nvidia.cublas.lib; import nvidia.cudnn.lib; print(os.path.dirname(nvidia.cublas.lib.__file__) + ":" + os.path.dirname(nvidia.cudnn.lib.__file__))'`
|
||||
ENV LD_LIBRARY_PATH="/usr/local/lib/python3.10/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.10/site-packages/nvidia/cudnn/lib"
|
||||
|
||||
COPY whisper_live /app/whisper_live
|
||||
|
||||
COPY run_server.py /app
|
||||
|
||||
CMD ["python", "run_server.py"]
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
FROM nvidia/cuda:12.5.1-runtime-ubuntu22.04 AS base
|
||||
|
||||
ARG DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
RUN apt-get update && apt-get install -y \
|
||||
python3.10 python3-pip openmpi-bin libopenmpi-dev git git-lfs wget \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
FROM base AS devel
|
||||
RUN pip3 install --no-cache-dir -U tensorrt_llm==0.15.0.dev2024111200 --extra-index-url https://pypi.nvidia.com
|
||||
WORKDIR /app
|
||||
RUN git clone https://github.com/NVIDIA/TensorRT-LLM.git && cd TensorRT-LLM && \
|
||||
git checkout c629546ce429623c8a163633095230154a6f0574 && cd ../ && \
|
||||
mv TensorRT-LLM/examples ./TensorRT-LLM-examples && \
|
||||
rm -rf TensorRT-LLM
|
||||
|
||||
|
||||
FROM devel AS release
|
||||
WORKDIR /app
|
||||
COPY assets/ ./assets
|
||||
RUN wget -nc -P assets/ https://raw.githubusercontent.com/openai/whisper/main/whisper/assets/mel_filters.npz
|
||||
|
||||
COPY scripts/setup.sh ./
|
||||
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
|
||||
COPY whisper_live ./whisper_live
|
||||
COPY scripts/build_whisper_tensorrt.sh .
|
||||
COPY run_server.py .
|
||||
@@ -1,12 +1,13 @@
|
||||
faster-whisper==0.10.0
|
||||
torch
|
||||
faster-whisper==1.1.0
|
||||
websockets
|
||||
onnxruntime==1.16.0
|
||||
numba
|
||||
openai-whisper
|
||||
kaldialign
|
||||
soundfile
|
||||
ffmpeg-python
|
||||
scipy
|
||||
jiwer
|
||||
evaluate
|
||||
evaluate
|
||||
numpy<2
|
||||
openai-whisper==20240930
|
||||
tokenizers==0.20.3
|
||||
+14
-2
@@ -1,5 +1,5 @@
|
||||
import argparse
|
||||
from whisper_live.server import TranscriptionServer
|
||||
import os
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
@@ -21,12 +21,23 @@ if __name__ == "__main__":
|
||||
parser.add_argument('--trt_multilingual', '-m',
|
||||
action="store_true",
|
||||
help='Boolean only for TensorRT model. True if multilingual.')
|
||||
parser.add_argument('--omp_num_threads', '-omp',
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of threads to use for OpenMP")
|
||||
parser.add_argument('--no_single_model', '-nsm',
|
||||
action='store_true',
|
||||
help='Set this if every connection should instantiate its own model. Only relevant for custom model, passed using -trt or -fw.')
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.backend == "tensorrt":
|
||||
if args.trt_model_path is None:
|
||||
raise ValueError("Please Provide a valid tensorrt model path")
|
||||
|
||||
if "OMP_NUM_THREADS" not in os.environ:
|
||||
os.environ["OMP_NUM_THREADS"] = str(args.omp_num_threads)
|
||||
|
||||
from whisper_live.server import TranscriptionServer
|
||||
server = TranscriptionServer()
|
||||
server.run(
|
||||
"0.0.0.0",
|
||||
@@ -34,5 +45,6 @@ if __name__ == "__main__":
|
||||
backend=args.backend,
|
||||
faster_whisper_custom_model_path=args.faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path=args.trt_model_path,
|
||||
trt_multilingual=args.trt_multilingual
|
||||
trt_multilingual=args.trt_multilingual,
|
||||
single_model=not args.no_single_model,
|
||||
)
|
||||
|
||||
@@ -38,12 +38,24 @@ download_and_build_model() {
|
||||
"large-v3" | "large")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/e5b1a55b89c1367dacf97e3e19bfd829a01529dbfdeefa8caeb59b3f1b81dadb/large-v3.pt"
|
||||
;;
|
||||
"large-v3-turbo" | "turbo")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/aff26ae408abcba5fbf8813c21e62b0941638c5f6eebfb145be0c9839262a19a/large-v3-turbo.pt"
|
||||
;;
|
||||
*)
|
||||
echo "Invalid model name: $model_name"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
if [ "$model_name" == "turbo" ]; then
|
||||
model_name="large-v3-turbo"
|
||||
fi
|
||||
|
||||
local inference_precision="float16"
|
||||
local weight_only_precision="${2:-float16}"
|
||||
local max_beam_width=4
|
||||
local max_batch_size=1
|
||||
|
||||
echo "Downloading $model_name..."
|
||||
# wget --directory-prefix=assets "$model_url"
|
||||
# echo "Download completed: ${model_name}.pt"
|
||||
@@ -54,11 +66,43 @@ download_and_build_model() {
|
||||
echo "${model_name}.pt already exists in assets directory."
|
||||
fi
|
||||
|
||||
local output_dir="whisper_${model_name//./_}"
|
||||
local sanitized_model_name="${model_name//./_}"
|
||||
local checkpoint_dir="whisper_${sanitized_model_name}_weights_${weight_only_precision}"
|
||||
local output_dir="whisper_${sanitized_model_name}_${weight_only_precision}"
|
||||
echo "$output_dir"
|
||||
echo "Running build script for $model_name with output directory $output_dir"
|
||||
python3 build.py --output_dir "$output_dir" --use_gpt_attention_plugin --use_gemm_plugin --use_bert_attention_plugin --model_name "$model_name"
|
||||
echo "Whisper $model_name TensorRT engine built."
|
||||
echo "Converting model weights for $model_name..."
|
||||
python3 convert_checkpoint.py \
|
||||
$( [[ "$weight_only_precision" == "int8" || "$weight_only_precision" == "int4" ]] && echo "--use_weight_only --weight_only_precision $weight_only_precision" ) \
|
||||
--output_dir "$checkpoint_dir" --model_name "$model_name"
|
||||
|
||||
echo "Building encoder for $model_name..."
|
||||
trtllm-build \
|
||||
--checkpoint_dir "${checkpoint_dir}/encoder" \
|
||||
--output_dir "${output_dir}/encoder" \
|
||||
--moe_plugin disable \
|
||||
--enable_xqa disable \
|
||||
--max_batch_size "$max_batch_size" \
|
||||
--gemm_plugin disable \
|
||||
--bert_attention_plugin "$inference_precision" \
|
||||
--max_input_len 3000 \
|
||||
--max_seq_len 3000
|
||||
|
||||
echo "Building decoder for $model_name..."
|
||||
trtllm-build \
|
||||
--checkpoint_dir "${checkpoint_dir}/decoder" \
|
||||
--output_dir "${output_dir}/decoder" \
|
||||
--moe_plugin disable \
|
||||
--enable_xqa disable \
|
||||
--max_beam_width "$max_beam_width" \
|
||||
--max_batch_size "$max_batch_size" \
|
||||
--max_seq_len 200 \
|
||||
--max_input_len 14 \
|
||||
--max_encoder_input_len 3000 \
|
||||
--gemm_plugin "$inference_precision" \
|
||||
--bert_attention_plugin "$inference_precision" \
|
||||
--gpt_attention_plugin "$inference_precision"
|
||||
|
||||
echo "TensorRT LLM engine built for $model_name."
|
||||
echo "========================================="
|
||||
echo "Model is located at: $(pwd)/$output_dir"
|
||||
}
|
||||
@@ -70,8 +114,9 @@ fi
|
||||
|
||||
tensorrt_examples_dir="$1"
|
||||
model_name="${2:-small.en}"
|
||||
weight_only_precision="${3:-float16}" # Default to float16 if not provided
|
||||
|
||||
cd $1/whisper
|
||||
cd $tensorrt_examples_dir/whisper
|
||||
pip install --no-deps -r requirements.txt
|
||||
|
||||
download_and_build_model "$model_name"
|
||||
download_and_build_model "$model_name" "$weight_only_precision"
|
||||
|
||||
@@ -11,7 +11,7 @@ README = (HERE / "README.md").read_text()
|
||||
|
||||
# This call to setup() does all the work
|
||||
setup(
|
||||
name="whisper-live",
|
||||
name="whisper_live",
|
||||
version=__version__,
|
||||
description="A nearly-live implementation of OpenAI's Whisper.",
|
||||
long_description=README,
|
||||
@@ -43,7 +43,7 @@ setup(
|
||||
),
|
||||
install_requires=[
|
||||
"PyAudio",
|
||||
"faster-whisper==0.10.0",
|
||||
"faster-whisper==1.1.0",
|
||||
"torch",
|
||||
"torchaudio",
|
||||
"websockets",
|
||||
@@ -52,9 +52,10 @@ setup(
|
||||
"scipy",
|
||||
"websocket-client",
|
||||
"numba",
|
||||
"openai-whisper",
|
||||
"openai-whisper==20240930",
|
||||
"kaldialign",
|
||||
"soundfile",
|
||||
"tokenizers==0.20.3"
|
||||
],
|
||||
python_requires=">=3.8"
|
||||
)
|
||||
|
||||
+57
-10
@@ -2,10 +2,12 @@ import json
|
||||
import os
|
||||
import scipy
|
||||
import websocket
|
||||
import copy
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
from whisper_live.client import TranscriptionClient
|
||||
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
|
||||
from whisper_live.utils import resample
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class BaseTestCase(unittest.TestCase):
|
||||
@@ -24,6 +26,7 @@ class BaseTestCase(unittest.TestCase):
|
||||
|
||||
self.mock_pyaudio = mock_pyaudio
|
||||
self.mock_websocket = mock_websocket
|
||||
self.mock_audio_packet = b'\x00\x01\x02\x03'
|
||||
|
||||
def tearDown(self):
|
||||
self.client.close_websocket()
|
||||
@@ -31,7 +34,6 @@ class BaseTestCase(unittest.TestCase):
|
||||
self.mock_websocket.stop()
|
||||
del self.client
|
||||
|
||||
|
||||
class TestClientWebSocketCommunication(BaseTestCase):
|
||||
def test_websocket_communication(self):
|
||||
expected_url = 'ws://localhost:9090'
|
||||
@@ -46,7 +48,9 @@ class TestClientCallbacks(BaseTestCase):
|
||||
"language": self.client.language,
|
||||
"task": self.client.task,
|
||||
"model": self.client.model,
|
||||
"use_vad": True
|
||||
"use_vad": True,
|
||||
"max_clients": 4,
|
||||
"max_connection_time": 600,
|
||||
})
|
||||
self.client.on_open(self.mock_ws_app)
|
||||
self.mock_ws_app.send.assert_called_with(expected_message)
|
||||
@@ -64,15 +68,15 @@ class TestClientCallbacks(BaseTestCase):
|
||||
message = json.dumps({
|
||||
"uid": self.client.uid,
|
||||
"segments": [
|
||||
{"start": 0, "end": 1, "text": "Test transcript"},
|
||||
{"start": 1, "end": 2, "text": "Test transcript 2"},
|
||||
{"start": 2, "end": 3, "text": "Test transcript 3"}
|
||||
{"start": 0, "end": 1, "text": "Test transcript", "completed": True},
|
||||
{"start": 1, "end": 2, "text": "Test transcript 2", "completed": True},
|
||||
{"start": 2, "end": 3, "text": "Test transcript 3", "completed": True}
|
||||
]
|
||||
})
|
||||
self.client.on_message(self.mock_ws_app, message)
|
||||
|
||||
# Assert that the transcript was updated correctly
|
||||
self.assertEqual(len(self.client.transcript), 2)
|
||||
self.assertEqual(len(self.client.transcript), 3)
|
||||
self.assertEqual(self.client.transcript[1]['text'], "Test transcript 2")
|
||||
|
||||
def test_on_close(self):
|
||||
@@ -106,6 +110,49 @@ class TestAudioResampling(unittest.TestCase):
|
||||
|
||||
class TestSendingAudioPacket(BaseTestCase):
|
||||
def test_send_packet(self):
|
||||
mock_audio_packet = b'\x00\x01\x02\x03'
|
||||
self.client.send_packet_to_server(mock_audio_packet)
|
||||
self.client.client_socket.send.assert_called_with(mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
self.client.send_packet_to_server(self.mock_audio_packet)
|
||||
self.client.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
class TestTee(BaseTestCase):
|
||||
@patch('whisper_live.client.websocket.WebSocketApp')
|
||||
@patch('whisper_live.client.pyaudio.PyAudio')
|
||||
def setUp(self, mock_audio, mock_websocket):
|
||||
super().setUp()
|
||||
self.client2 = Client(host='localhost', port=9090, lang="es", translate=False, srt_file_path="transcript.srt")
|
||||
self.client3 = Client(host='localhost', port=9090, lang="es", translate=True, srt_file_path="translation.srt")
|
||||
# need a separate mock for each websocket
|
||||
self.client3.client_socket = copy.deepcopy(self.client3.client_socket)
|
||||
self.tee = TranscriptionTeeClient([self.client2, self.client3])
|
||||
|
||||
def tearDown(self):
|
||||
self.tee.close_all_clients()
|
||||
del self.tee
|
||||
super().tearDown()
|
||||
|
||||
def test_invalid_constructor(self):
|
||||
with self.assertRaises(Exception) as context:
|
||||
TranscriptionTeeClient([])
|
||||
|
||||
def test_multicast_unconditional(self):
|
||||
self.tee.multicast_packet(self.mock_audio_packet, True)
|
||||
for client in self.tee.clients:
|
||||
client.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
def test_multicast_conditional(self):
|
||||
self.client2.recording = False
|
||||
self.client3.recording = True
|
||||
self.tee.multicast_packet(self.mock_audio_packet, False)
|
||||
self.client2.client_socket.send.assert_not_called()
|
||||
self.client3.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
def test_close_all(self):
|
||||
self.tee.close_all_clients()
|
||||
for client in self.tee.clients:
|
||||
client.client_socket.close.assert_called()
|
||||
|
||||
def test_write_all_srt(self):
|
||||
for client in self.tee.clients:
|
||||
client.server_backend = "faster_whisper"
|
||||
self.tee.write_all_clients_srt()
|
||||
self.assertTrue(Path("transcript.srt").is_file())
|
||||
self.assertTrue(Path("translation.srt").is_file())
|
||||
|
||||
+36
-25
@@ -5,17 +5,18 @@ import unittest
|
||||
from unittest import mock
|
||||
|
||||
import numpy as np
|
||||
import evaluate
|
||||
import jiwer
|
||||
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from whisper_live.server import TranscriptionServer
|
||||
from whisper_live.client import TranscriptionClient
|
||||
from whisper_live.server import TranscriptionServer, BackendType, ClientManager
|
||||
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
|
||||
from whisper.normalizers import EnglishTextNormalizer
|
||||
|
||||
|
||||
class TestTranscriptionServerInitialization(unittest.TestCase):
|
||||
def test_initialization(self):
|
||||
server = TranscriptionServer()
|
||||
server.client_manager = ClientManager(max_clients=4, max_connection_time=600)
|
||||
self.assertEqual(server.client_manager.max_clients, 4)
|
||||
self.assertEqual(server.client_manager.max_connection_time, 600)
|
||||
self.assertDictEqual(server.client_manager.clients, {})
|
||||
@@ -25,6 +26,7 @@ class TestTranscriptionServerInitialization(unittest.TestCase):
|
||||
class TestGetWaitTime(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600)
|
||||
self.server.client_manager.start_times = {
|
||||
'client1': time.time() - 120,
|
||||
'client2': time.time() - 300
|
||||
@@ -49,7 +51,7 @@ class TestServerConnection(unittest.TestCase):
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
})
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_recv_audio_exception_handling(self, mock_websocket):
|
||||
@@ -61,7 +63,7 @@ class TestServerConnection(unittest.TestCase):
|
||||
}), np.array([1, 2, 3]).tobytes()]
|
||||
|
||||
with self.assertLogs(level="ERROR"):
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
|
||||
self.assertNotIn(mock_websocket, self.server.client_manager.clients)
|
||||
|
||||
@@ -69,6 +71,10 @@ class TestServerConnection(unittest.TestCase):
|
||||
class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.mock_pyaudio_patch = mock.patch('pyaudio.PyAudio')
|
||||
cls.mock_pyaudio = cls.mock_pyaudio_patch.start()
|
||||
cls.mock_pyaudio.return_value.open.return_value = mock.MagicMock()
|
||||
|
||||
cls.server_process = subprocess.Popen(["python", "run_server.py"])
|
||||
time.sleep(2)
|
||||
|
||||
@@ -77,32 +83,37 @@ class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
cls.server_process.terminate()
|
||||
cls.server_process.wait()
|
||||
|
||||
@mock.patch('pyaudio.PyAudio')
|
||||
def setUp(self, mock_pyaudio):
|
||||
self.mock_pyaudio = mock_pyaudio.return_value
|
||||
self.mock_stream = mock.MagicMock()
|
||||
self.mock_pyaudio.open.return_value = self.mock_stream
|
||||
self.metric = evaluate.load("wer")
|
||||
def setUp(self):
|
||||
self.normalizer = EnglishTextNormalizer()
|
||||
self.client = TranscriptionClient(
|
||||
"localhost", "9090", model="base.en", lang="en",
|
||||
)
|
||||
|
||||
def test_inference(self):
|
||||
def check_prediction(self, srt_path):
|
||||
gt = "And so my fellow Americans, ask not, what your country can do for you. Ask what you can do for your country!"
|
||||
self.client("assets/jfk.flac")
|
||||
with open("output.srt", "r") as f:
|
||||
with open(srt_path, "r") as f:
|
||||
lines = f.readlines()
|
||||
prediction = " ".join([line.strip() for line in lines[2::4]])
|
||||
prediction_normalized = self.normalizer(prediction)
|
||||
gt_normalized = self.normalizer(gt)
|
||||
|
||||
# calculate WER
|
||||
wer = self.metric.compute(
|
||||
predictions=[prediction_normalized],
|
||||
references=[gt_normalized]
|
||||
wer_score = jiwer.wer(gt_normalized, prediction_normalized)
|
||||
self.assertLess(wer_score, 0.05)
|
||||
|
||||
def test_inference(self):
|
||||
client = TranscriptionClient(
|
||||
"localhost", "9090", model="base.en", lang="en",
|
||||
)
|
||||
self.assertLess(wer, 0.05)
|
||||
client("assets/jfk.flac")
|
||||
self.check_prediction("output.srt")
|
||||
|
||||
def test_simultaneous_inference(self):
|
||||
client1 = Client(
|
||||
"localhost", "9090", model="base.en", lang="en", srt_file_path="transcript1.srt")
|
||||
client2 = Client(
|
||||
"localhost", "9090", model="base.en", lang="en", srt_file_path="transcript2.srt")
|
||||
tee = TranscriptionTeeClient([client1, client2])
|
||||
tee("assets/jfk.flac")
|
||||
self.check_prediction("transcript1.srt")
|
||||
self.check_prediction("transcript2.srt")
|
||||
|
||||
|
||||
class TestExceptionHandling(unittest.TestCase):
|
||||
@@ -111,10 +122,10 @@ class TestExceptionHandling(unittest.TestCase):
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection_closed_exception(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = ConnectionClosed(1001, "testing connection closed")
|
||||
mock_websocket.recv.side_effect = ConnectionClosed(1001, "testing connection closed", rcvd_then_sent=mock.Mock())
|
||||
|
||||
with self.assertLogs(level="INFO") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
self.assertTrue(any("Connection closed by client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
@@ -122,7 +133,7 @@ class TestExceptionHandling(unittest.TestCase):
|
||||
mock_websocket.recv.return_value = "invalid json"
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
self.assertTrue(any("Failed to decode JSON from client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
@@ -130,7 +141,7 @@ class TestExceptionHandling(unittest.TestCase):
|
||||
mock_websocket.recv.side_effect = RuntimeError("Unexpected error")
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
for message in log.output:
|
||||
print(message)
|
||||
print()
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.2.1"
|
||||
__version__ = "0.6.1"
|
||||
|
||||
+410
-237
@@ -1,6 +1,8 @@
|
||||
import os
|
||||
import shutil
|
||||
import wave
|
||||
|
||||
import logging
|
||||
import numpy as np
|
||||
import pyaudio
|
||||
import threading
|
||||
@@ -14,7 +16,7 @@ import whisper_live.utils as utils
|
||||
|
||||
class Client:
|
||||
"""
|
||||
Handles audio recording, streaming, and communication with a server using WebSocket.
|
||||
Handles communication with a server using WebSocket.
|
||||
"""
|
||||
INSTANCES = {}
|
||||
END_OF_AUDIO = "END_OF_AUDIO"
|
||||
@@ -27,7 +29,10 @@ class Client:
|
||||
translate=False,
|
||||
model="small",
|
||||
srt_file_path="output.srt",
|
||||
use_vad=True
|
||||
use_vad=True,
|
||||
log_transcription=True,
|
||||
max_clients=4,
|
||||
max_connection_time=600,
|
||||
):
|
||||
"""
|
||||
Initializes a Client instance for audio recording and streaming to a server.
|
||||
@@ -42,37 +47,27 @@ class Client:
|
||||
lang (str, optional): The selected language for transcription. Default is None.
|
||||
translate (bool, optional): Specifies if the task is translation. Default is False.
|
||||
"""
|
||||
self.chunk = 4096
|
||||
self.format = pyaudio.paInt16
|
||||
self.channels = 1
|
||||
self.rate = 16000
|
||||
self.record_seconds = 60000
|
||||
self.recording = False
|
||||
self.task = "transcribe"
|
||||
self.uid = str(uuid.uuid4())
|
||||
self.waiting = False
|
||||
self.last_response_recieved = None
|
||||
self.last_response_received = None
|
||||
self.disconnect_if_no_response_for = 15
|
||||
self.language = lang
|
||||
self.model = model
|
||||
self.server_error = False
|
||||
self.srt_file_path = srt_file_path
|
||||
self.use_vad = use_vad
|
||||
self.last_recieved_segment = None
|
||||
self.last_segment = None
|
||||
self.last_received_segment = None
|
||||
self.log_transcription = log_transcription
|
||||
self.max_clients = max_clients
|
||||
self.max_connection_time = max_connection_time
|
||||
|
||||
if translate:
|
||||
self.task = "translate"
|
||||
|
||||
self.timestamp_offset = 0.0
|
||||
self.audio_bytes = None
|
||||
self.p = pyaudio.PyAudio()
|
||||
self.stream = self.p.open(
|
||||
format=self.format,
|
||||
channels=self.channels,
|
||||
rate=self.rate,
|
||||
input=True,
|
||||
frames_per_buffer=self.chunk,
|
||||
)
|
||||
|
||||
if host is not None and port is not None:
|
||||
socket_url = f"ws://{host}:{port}"
|
||||
@@ -96,7 +91,6 @@ class Client:
|
||||
self.ws_thread.setDaemon(True)
|
||||
self.ws_thread.start()
|
||||
|
||||
self.frames = b""
|
||||
self.transcript = []
|
||||
print("[INFO]: * recording")
|
||||
|
||||
@@ -118,21 +112,22 @@ class Client:
|
||||
for i, seg in enumerate(segments):
|
||||
if not text or text[-1] != seg["text"]:
|
||||
text.append(seg["text"])
|
||||
if i == len(segments) - 1:
|
||||
if i == len(segments) - 1 and not seg.get("completed", False):
|
||||
self.last_segment = seg
|
||||
elif (self.server_backend == "faster_whisper" and
|
||||
elif (self.server_backend == "faster_whisper" and seg.get("completed", False) and
|
||||
(not self.transcript or
|
||||
float(seg['start']) >= float(self.transcript[-1]['end']))):
|
||||
self.transcript.append(seg)
|
||||
# update last received segment and last valild responsne time
|
||||
if self.last_recieved_segment is None or self.last_recieved_segment != segments[-1]["text"]:
|
||||
self.last_response_recieved = time.time()
|
||||
self.last_recieved_segment = segments[-1]["text"]
|
||||
# update last received segment and last valid response time
|
||||
if self.last_received_segment is None or self.last_received_segment != segments[-1]["text"]:
|
||||
self.last_response_received = time.time()
|
||||
self.last_received_segment = segments[-1]["text"]
|
||||
|
||||
# Truncate to last 3 entries for brevity.
|
||||
text = text[-3:]
|
||||
utils.clear_screen()
|
||||
utils.print_transcript(text)
|
||||
if self.log_transcription:
|
||||
# Truncate to last 3 entries for brevity.
|
||||
text = text[-3:]
|
||||
utils.clear_screen()
|
||||
utils.print_transcript(text)
|
||||
|
||||
def on_message(self, ws, message):
|
||||
"""
|
||||
@@ -162,7 +157,7 @@ class Client:
|
||||
self.recording = False
|
||||
|
||||
if "message" in message.keys() and message["message"] == "SERVER_READY":
|
||||
self.last_response_recieved = time.time()
|
||||
self.last_response_received = time.time()
|
||||
self.recording = True
|
||||
self.server_backend = message["backend"]
|
||||
print(f"[INFO]: Server Running with backend {self.server_backend}")
|
||||
@@ -187,7 +182,6 @@ class Client:
|
||||
def on_close(self, ws, close_status_code, close_msg):
|
||||
print(f"[INFO]: Websocket connection closed: {close_status_code}: {close_msg}")
|
||||
self.recording = False
|
||||
self.server_error = False
|
||||
self.waiting = False
|
||||
|
||||
def on_open(self, ws):
|
||||
@@ -209,28 +203,13 @@ class Client:
|
||||
"language": self.language,
|
||||
"task": self.task,
|
||||
"model": self.model,
|
||||
"use_vad": self.use_vad
|
||||
"use_vad": self.use_vad,
|
||||
"max_clients": self.max_clients,
|
||||
"max_connection_time": self.max_connection_time,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def bytes_to_float_array(audio_bytes):
|
||||
"""
|
||||
Convert audio data from bytes to a NumPy float array.
|
||||
|
||||
It assumes that the audio data is in 16-bit PCM format. The audio data is normalized to
|
||||
have values between -1 and 1.
|
||||
|
||||
Args:
|
||||
audio_bytes (bytes): Audio data in bytes.
|
||||
|
||||
Returns:
|
||||
np.ndarray: A NumPy array containing the audio data as float values normalized between -1 and 1.
|
||||
"""
|
||||
raw_data = np.frombuffer(buffer=audio_bytes, dtype=np.int16)
|
||||
return raw_data.astype(np.float32) / 32768.0
|
||||
|
||||
def send_packet_to_server(self, message):
|
||||
"""
|
||||
Send an audio packet to the server using WebSocket.
|
||||
@@ -244,62 +223,6 @@ class Client:
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
def play_file(self, filename):
|
||||
"""
|
||||
Play an audio file and send it to the server for processing.
|
||||
|
||||
Reads an audio file, plays it through the audio output, and simultaneously sends
|
||||
the audio data to the server for processing. It uses PyAudio to create an audio
|
||||
stream for playback. The audio data is read from the file in chunks, converted to
|
||||
floating-point format, and sent to the server using WebSocket communication.
|
||||
This method is typically used when you want to process pre-recorded audio and send it
|
||||
to the server in real-time.
|
||||
|
||||
Args:
|
||||
filename (str): The path to the audio file to be played and sent to the server.
|
||||
"""
|
||||
|
||||
# read audio and create pyaudio stream
|
||||
with wave.open(filename, "rb") as wavfile:
|
||||
self.stream = self.p.open(
|
||||
format=self.p.get_format_from_width(wavfile.getsampwidth()),
|
||||
channels=wavfile.getnchannels(),
|
||||
rate=wavfile.getframerate(),
|
||||
input=True,
|
||||
output=True,
|
||||
frames_per_buffer=self.chunk,
|
||||
)
|
||||
try:
|
||||
while self.recording:
|
||||
data = wavfile.readframes(self.chunk)
|
||||
if data == b"":
|
||||
break
|
||||
|
||||
audio_array = self.bytes_to_float_array(data)
|
||||
self.send_packet_to_server(audio_array.tobytes())
|
||||
self.stream.write(data)
|
||||
|
||||
wavfile.close()
|
||||
|
||||
assert self.last_response_recieved
|
||||
while time.time() - self.last_response_recieved < self.disconnect_if_no_response_for:
|
||||
continue
|
||||
self.send_packet_to_server(Client.END_OF_AUDIO.encode('utf-8'))
|
||||
if self.server_backend == "faster_whisper":
|
||||
self.write_srt_file(self.srt_file_path)
|
||||
self.stream.close()
|
||||
self.close_websocket()
|
||||
|
||||
except KeyboardInterrupt:
|
||||
wavfile.close()
|
||||
self.stream.stop_stream()
|
||||
self.stream.close()
|
||||
self.p.terminate()
|
||||
self.close_websocket()
|
||||
if self.server_backend == "faster_whisper":
|
||||
self.write_srt_file(self.srt_file_path)
|
||||
print("[INFO]: Keyboard interrupt.")
|
||||
|
||||
def close_websocket(self):
|
||||
"""
|
||||
Close the WebSocket connection and join the WebSocket thread.
|
||||
@@ -327,6 +250,332 @@ class Client:
|
||||
"""
|
||||
return self.client_socket
|
||||
|
||||
def write_srt_file(self, output_path="output.srt"):
|
||||
"""
|
||||
Writes out the transcript in .srt format.
|
||||
|
||||
Args:
|
||||
message (output_path, optional): The path to the target file. Default is "output.srt".
|
||||
|
||||
"""
|
||||
if self.server_backend == "faster_whisper":
|
||||
if not self.transcript and self.last_segment is not None:
|
||||
self.transcript.append(self.last_segment)
|
||||
elif self.last_segment and self.transcript[-1]["text"] != self.last_segment["text"]:
|
||||
self.transcript.append(self.last_segment)
|
||||
utils.create_srt_file(self.transcript, output_path)
|
||||
|
||||
def wait_before_disconnect(self):
|
||||
"""Waits a bit before disconnecting in order to process pending responses."""
|
||||
assert self.last_response_received
|
||||
while time.time() - self.last_response_received < self.disconnect_if_no_response_for:
|
||||
continue
|
||||
|
||||
|
||||
class TranscriptionTeeClient:
|
||||
"""
|
||||
Client for handling audio recording, streaming, and transcription tasks via one or more
|
||||
WebSocket connections.
|
||||
|
||||
Acts as a high-level client for audio transcription tasks using a WebSocket connection. It can be used
|
||||
to send audio data for transcription to one or more servers, and receive transcribed text segments.
|
||||
Args:
|
||||
clients (list): one or more previously initialized Client instances
|
||||
|
||||
Attributes:
|
||||
clients (list): the underlying Client instances responsible for handling WebSocket connections.
|
||||
"""
|
||||
def __init__(self, clients, save_output_recording=False, output_recording_filename="./output_recording.wav"):
|
||||
self.clients = clients
|
||||
if not self.clients:
|
||||
raise Exception("At least one client is required.")
|
||||
self.chunk = 4096
|
||||
self.format = pyaudio.paInt16
|
||||
self.channels = 1
|
||||
self.rate = 16000
|
||||
self.record_seconds = 60000
|
||||
self.save_output_recording = save_output_recording
|
||||
self.output_recording_filename = output_recording_filename
|
||||
self.frames = b""
|
||||
self.p = pyaudio.PyAudio()
|
||||
try:
|
||||
self.stream = self.p.open(
|
||||
format=self.format,
|
||||
channels=self.channels,
|
||||
rate=self.rate,
|
||||
input=True,
|
||||
frames_per_buffer=self.chunk,
|
||||
)
|
||||
except OSError as error:
|
||||
print(f"[WARN]: Unable to access microphone. {error}")
|
||||
self.stream = None
|
||||
|
||||
def __call__(self, audio=None, rtsp_url=None, hls_url=None, save_file=None):
|
||||
"""
|
||||
Start the transcription process.
|
||||
|
||||
Initiates the transcription process by connecting to the server via a WebSocket. It waits for the server
|
||||
to be ready to receive audio data and then sends audio for transcription. If an audio file is provided, it
|
||||
will be played and streamed to the server; otherwise, it will perform live recording.
|
||||
|
||||
Args:
|
||||
audio (str, optional): Path to an audio file for transcription. Default is None, which triggers live recording.
|
||||
|
||||
"""
|
||||
assert sum(
|
||||
source is not None for source in [audio, rtsp_url, hls_url]
|
||||
) <= 1, 'You must provide only one selected source'
|
||||
|
||||
print("[INFO]: Waiting for server ready ...")
|
||||
for client in self.clients:
|
||||
while not client.recording:
|
||||
if client.waiting or client.server_error:
|
||||
self.close_all_clients()
|
||||
return
|
||||
|
||||
print("[INFO]: Server Ready!")
|
||||
if hls_url is not None:
|
||||
self.process_hls_stream(hls_url, save_file)
|
||||
elif audio is not None:
|
||||
resampled_file = utils.resample(audio)
|
||||
self.play_file(resampled_file)
|
||||
elif rtsp_url is not None:
|
||||
self.process_rtsp_stream(rtsp_url)
|
||||
else:
|
||||
self.record()
|
||||
|
||||
def close_all_clients(self):
|
||||
"""Closes all client websockets."""
|
||||
for client in self.clients:
|
||||
client.close_websocket()
|
||||
|
||||
def write_all_clients_srt(self):
|
||||
"""Writes out .srt files for all clients."""
|
||||
for client in self.clients:
|
||||
client.write_srt_file(client.srt_file_path)
|
||||
|
||||
def multicast_packet(self, packet, unconditional=False):
|
||||
"""
|
||||
Sends an identical packet via all clients.
|
||||
|
||||
Args:
|
||||
packet (bytes): The audio data packet in bytes to be sent.
|
||||
unconditional (bool, optional): If true, send regardless of whether clients are recording. Default is False.
|
||||
"""
|
||||
for client in self.clients:
|
||||
if (unconditional or client.recording):
|
||||
client.send_packet_to_server(packet)
|
||||
|
||||
def play_file(self, filename):
|
||||
"""
|
||||
Play an audio file and send it to the server for processing.
|
||||
|
||||
Reads an audio file, plays it through the audio output, and simultaneously sends
|
||||
the audio data to the server for processing. It uses PyAudio to create an audio
|
||||
stream for playback. The audio data is read from the file in chunks, converted to
|
||||
floating-point format, and sent to the server using WebSocket communication.
|
||||
This method is typically used when you want to process pre-recorded audio and send it
|
||||
to the server in real-time.
|
||||
|
||||
Args:
|
||||
filename (str): The path to the audio file to be played and sent to the server.
|
||||
"""
|
||||
|
||||
# read audio and create pyaudio stream
|
||||
with wave.open(filename, "rb") as wavfile:
|
||||
self.stream = self.p.open(
|
||||
format=self.p.get_format_from_width(wavfile.getsampwidth()),
|
||||
channels=wavfile.getnchannels(),
|
||||
rate=wavfile.getframerate(),
|
||||
input=True,
|
||||
output=True,
|
||||
frames_per_buffer=self.chunk,
|
||||
)
|
||||
try:
|
||||
while any(client.recording for client in self.clients):
|
||||
data = wavfile.readframes(self.chunk)
|
||||
if data == b"":
|
||||
break
|
||||
|
||||
audio_array = self.bytes_to_float_array(data)
|
||||
self.multicast_packet(audio_array.tobytes())
|
||||
self.stream.write(data)
|
||||
|
||||
wavfile.close()
|
||||
|
||||
for client in self.clients:
|
||||
client.wait_before_disconnect()
|
||||
self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True)
|
||||
self.write_all_clients_srt()
|
||||
self.stream.close()
|
||||
self.close_all_clients()
|
||||
|
||||
except KeyboardInterrupt:
|
||||
wavfile.close()
|
||||
self.stream.stop_stream()
|
||||
self.stream.close()
|
||||
self.p.terminate()
|
||||
self.close_all_clients()
|
||||
self.write_all_clients_srt()
|
||||
print("[INFO]: Keyboard interrupt.")
|
||||
|
||||
def process_rtsp_stream(self, rtsp_url):
|
||||
"""
|
||||
Connect to an RTSP source, process the audio stream, and send it for trascription.
|
||||
|
||||
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')
|
||||
|
||||
def process_hls_stream(self, hls_url, save_file):
|
||||
"""
|
||||
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.
|
||||
"""
|
||||
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()
|
||||
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())
|
||||
|
||||
except Exception as e:
|
||||
print(f"[ERROR]: Failed to connect to {stream_type} stream: {e}")
|
||||
finally:
|
||||
self.close_all_clients()
|
||||
self.write_all_clients_srt()
|
||||
if process:
|
||||
process.kill()
|
||||
|
||||
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):
|
||||
"""
|
||||
Consume and log the stderr output of a process in a separate thread.
|
||||
|
||||
Args:
|
||||
process (subprocess.Popen): The process whose stderr output will be logged.
|
||||
"""
|
||||
for line in iter(process.stderr.readline, b""):
|
||||
logging.debug(f'[STDERR]: {line.decode()}')
|
||||
|
||||
def save_chunk(self, n_audio_file):
|
||||
"""
|
||||
Saves the current audio frames to a WAV file in a separate thread.
|
||||
|
||||
Args:
|
||||
n_audio_file (int): The index of the audio file which determines the filename.
|
||||
This helps in maintaining the order and uniqueness of each chunk.
|
||||
"""
|
||||
t = threading.Thread(
|
||||
target=self.write_audio_frames_to_file,
|
||||
args=(self.frames[:], f"chunks/{n_audio_file}.wav",),
|
||||
)
|
||||
t.start()
|
||||
|
||||
def finalize_recording(self, n_audio_file):
|
||||
"""
|
||||
Finalizes the recording process by saving any remaining audio frames,
|
||||
closing the audio stream, and terminating the process.
|
||||
|
||||
Args:
|
||||
n_audio_file (int): The file index to be used if there are remaining audio frames to be saved.
|
||||
This index is incremented before use if the last chunk is saved.
|
||||
"""
|
||||
if self.save_output_recording and len(self.frames):
|
||||
self.write_audio_frames_to_file(
|
||||
self.frames[:], f"chunks/{n_audio_file}.wav"
|
||||
)
|
||||
n_audio_file += 1
|
||||
self.stream.stop_stream()
|
||||
self.stream.close()
|
||||
self.p.terminate()
|
||||
self.close_all_clients()
|
||||
if self.save_output_recording:
|
||||
self.write_output_recording(n_audio_file)
|
||||
self.write_all_clients_srt()
|
||||
|
||||
def record(self):
|
||||
"""
|
||||
Record audio data from the input stream and save it to a WAV file.
|
||||
|
||||
Continuously records audio data from the input stream, sends it to the server via a WebSocket
|
||||
connection, and simultaneously saves it to multiple WAV files in chunks. It stops recording when
|
||||
the `RECORD_SECONDS` duration is reached or when the `RECORDING` flag is set to `False`.
|
||||
|
||||
Audio data is saved in chunks to the "chunks" directory. Each chunk is saved as a separate WAV file.
|
||||
The recording will continue until the specified duration is reached or until the `RECORDING` flag is set to `False`.
|
||||
The recording process can be interrupted by sending a KeyboardInterrupt (e.g., pressing Ctrl+C). After recording,
|
||||
the method combines all the saved audio chunks into the specified `out_file`.
|
||||
"""
|
||||
n_audio_file = 0
|
||||
if self.save_output_recording:
|
||||
if os.path.exists("chunks"):
|
||||
shutil.rmtree("chunks")
|
||||
os.makedirs("chunks")
|
||||
try:
|
||||
for _ in range(0, int(self.rate / self.chunk * self.record_seconds)):
|
||||
if not any(client.recording for client in self.clients):
|
||||
break
|
||||
data = self.stream.read(self.chunk, exception_on_overflow=False)
|
||||
self.frames += data
|
||||
|
||||
audio_array = self.bytes_to_float_array(data)
|
||||
|
||||
self.multicast_packet(audio_array.tobytes())
|
||||
|
||||
# save frames if more than a minute
|
||||
if len(self.frames) > 60 * self.rate:
|
||||
if self.save_output_recording:
|
||||
self.save_chunk(n_audio_file)
|
||||
n_audio_file += 1
|
||||
self.frames = b""
|
||||
self.write_all_clients_srt()
|
||||
|
||||
except KeyboardInterrupt:
|
||||
self.finalize_recording(n_audio_file)
|
||||
|
||||
def write_audio_frames_to_file(self, frames, file_name):
|
||||
"""
|
||||
Write audio frames to a WAV file.
|
||||
@@ -346,104 +595,7 @@ class Client:
|
||||
wavfile.setframerate(self.rate)
|
||||
wavfile.writeframes(frames)
|
||||
|
||||
def process_hls_stream(self, hls_url):
|
||||
"""
|
||||
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.
|
||||
"""
|
||||
print("[INFO]: Connecting to HLS stream...")
|
||||
process = None # Initialize process to None
|
||||
|
||||
try:
|
||||
# Connecting to the HLS stream using ffmpeg-python
|
||||
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)
|
||||
)
|
||||
|
||||
# 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.send_packet_to_server(audio_array.tobytes())
|
||||
|
||||
except Exception as e:
|
||||
print(f"[ERROR]: Failed to connect to HLS stream: {e}")
|
||||
finally:
|
||||
if process:
|
||||
process.kill()
|
||||
|
||||
print("[INFO]: HLS stream processing finished.")
|
||||
|
||||
def record(self, out_file="output_recording.wav"):
|
||||
"""
|
||||
Record audio data from the input stream and save it to a WAV file.
|
||||
|
||||
Continuously records audio data from the input stream, sends it to the server via a WebSocket
|
||||
connection, and simultaneously saves it to multiple WAV files in chunks. It stops recording when
|
||||
the `RECORD_SECONDS` duration is reached or when the `RECORDING` flag is set to `False`.
|
||||
|
||||
Audio data is saved in chunks to the "chunks" directory. Each chunk is saved as a separate WAV file.
|
||||
The recording will continue until the specified duration is reached or until the `RECORDING` flag is set to `False`.
|
||||
The recording process can be interrupted by sending a KeyboardInterrupt (e.g., pressing Ctrl+C). After recording,
|
||||
the method combines all the saved audio chunks into the specified `out_file`.
|
||||
|
||||
Args:
|
||||
out_file (str, optional): The name of the output WAV file to save the entire recording.
|
||||
Default is "output_recording.wav".
|
||||
|
||||
"""
|
||||
n_audio_file = 0
|
||||
if not os.path.exists("chunks"):
|
||||
os.makedirs("chunks", exist_ok=True)
|
||||
try:
|
||||
for _ in range(0, int(self.rate / self.chunk * self.record_seconds)):
|
||||
if not self.recording:
|
||||
break
|
||||
data = self.stream.read(self.chunk, exception_on_overflow=False)
|
||||
self.frames += data
|
||||
|
||||
audio_array = Client.bytes_to_float_array(data)
|
||||
|
||||
self.send_packet_to_server(audio_array.tobytes())
|
||||
|
||||
# save frames if more than a minute
|
||||
if len(self.frames) > 60 * self.rate:
|
||||
t = threading.Thread(
|
||||
target=self.write_audio_frames_to_file,
|
||||
args=(
|
||||
self.frames[:],
|
||||
f"chunks/{n_audio_file}.wav",
|
||||
),
|
||||
)
|
||||
t.start()
|
||||
n_audio_file += 1
|
||||
self.frames = b""
|
||||
if self.server_backend == "faster_whisper":
|
||||
self.write_srt_file(self.srt_file_path)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
if len(self.frames):
|
||||
self.write_audio_frames_to_file(
|
||||
self.frames[:], f"chunks/{n_audio_file}.wav"
|
||||
)
|
||||
n_audio_file += 1
|
||||
self.stream.stop_stream()
|
||||
self.stream.close()
|
||||
self.p.terminate()
|
||||
self.close_websocket()
|
||||
|
||||
self.write_output_recording(n_audio_file, out_file)
|
||||
if self.server_backend == "faster_whisper":
|
||||
self.write_srt_file(self.srt_file_path)
|
||||
|
||||
def write_output_recording(self, n_audio_file, out_file):
|
||||
def write_output_recording(self, n_audio_file):
|
||||
"""
|
||||
Combine and save recorded audio chunks into a single WAV file.
|
||||
|
||||
@@ -462,7 +614,7 @@ class Client:
|
||||
for i in range(n_audio_file)
|
||||
if os.path.exists(f"chunks/{i}.wav")
|
||||
]
|
||||
with wave.open(out_file, "wb") as wavfile:
|
||||
with wave.open(self.output_recording_filename, "wb") as wavfile:
|
||||
wavfile: wave.Wave_write
|
||||
wavfile.setnchannels(self.channels)
|
||||
wavfile.setsampwidth(2)
|
||||
@@ -477,15 +629,31 @@ class Client:
|
||||
# remove this file
|
||||
os.remove(in_file)
|
||||
wavfile.close()
|
||||
# clean up temporary directory to store chunks
|
||||
if os.path.exists("chunks"):
|
||||
shutil.rmtree("chunks")
|
||||
|
||||
def write_srt_file(self, output_path="output.srt"):
|
||||
self.transcript.append(self.last_segment)
|
||||
utils.create_srt_file(self.transcript, output_path)
|
||||
@staticmethod
|
||||
def bytes_to_float_array(audio_bytes):
|
||||
"""
|
||||
Convert audio data from bytes to a NumPy float array.
|
||||
|
||||
It assumes that the audio data is in 16-bit PCM format. The audio data is normalized to
|
||||
have values between -1 and 1.
|
||||
|
||||
Args:
|
||||
audio_bytes (bytes): Audio data in bytes.
|
||||
|
||||
Returns:
|
||||
np.ndarray: A NumPy array containing the audio data as float values normalized between -1 and 1.
|
||||
"""
|
||||
raw_data = np.frombuffer(buffer=audio_bytes, dtype=np.int16)
|
||||
return raw_data.astype(np.float32) / 32768.0
|
||||
|
||||
|
||||
class TranscriptionClient:
|
||||
class TranscriptionClient(TranscriptionTeeClient):
|
||||
"""
|
||||
Client for handling audio transcription tasks via a WebSocket connection.
|
||||
Client for handling audio transcription tasks via a single WebSocket connection.
|
||||
|
||||
Acts as a high-level client for audio transcription tasks using a WebSocket connection. It can be used
|
||||
to send audio data for transcription to a server and receive transcribed text segments.
|
||||
@@ -495,6 +663,9 @@ class TranscriptionClient:
|
||||
port (int): The port number to connect to on the server.
|
||||
lang (str, optional): The primary language for transcription. Default is None, which defaults to English ('en').
|
||||
translate (bool, optional): Indicates whether translation tasks are required (default is False).
|
||||
save_output_recording (bool, optional): Indicates whether to save recording from microphone.
|
||||
output_recording_filename (str, optional): File to save the output recording.
|
||||
output_transcription_path (str, optional): File to save the output transcription.
|
||||
|
||||
Attributes:
|
||||
client (Client): An instance of the underlying Client class responsible for handling the WebSocket connection.
|
||||
@@ -506,32 +677,34 @@ class TranscriptionClient:
|
||||
transcription_client()
|
||||
```
|
||||
"""
|
||||
def __init__(self, host, port, lang=None, translate=False, model="small", use_vad=True):
|
||||
self.client = Client(host, port, lang, translate, model, srt_file_path="output.srt", use_vad=use_vad)
|
||||
def __init__(
|
||||
self,
|
||||
host,
|
||||
port,
|
||||
lang=None,
|
||||
translate=False,
|
||||
model="small",
|
||||
use_vad=True,
|
||||
save_output_recording=False,
|
||||
output_recording_filename="./output_recording.wav",
|
||||
output_transcription_path="./output.srt",
|
||||
log_transcription=True,
|
||||
max_clients=4,
|
||||
max_connection_time=600,
|
||||
):
|
||||
self.client = Client(
|
||||
host, port, lang, translate, model, srt_file_path=output_transcription_path,
|
||||
use_vad=use_vad, log_transcription=log_transcription, max_clients=max_clients,
|
||||
max_connection_time=max_connection_time
|
||||
)
|
||||
|
||||
def __call__(self, audio=None, hls_url=None):
|
||||
"""
|
||||
Start the transcription process.
|
||||
|
||||
Initiates the transcription process by connecting to the server via a WebSocket. It waits for the server
|
||||
to be ready to receive audio data and then sends audio for transcription. If an audio file is provided, it
|
||||
will be played and streamed to the server; otherwise, it will perform live recording.
|
||||
|
||||
Args:
|
||||
audio (str, optional): Path to an audio file for transcription. Default is None, which triggers live recording.
|
||||
|
||||
"""
|
||||
print("[INFO]: Waiting for server ready ...")
|
||||
while not self.client.recording:
|
||||
if self.client.waiting or self.client.server_error:
|
||||
self.client.close_websocket()
|
||||
return
|
||||
|
||||
print("[INFO]: Server Ready!")
|
||||
if hls_url is not None:
|
||||
self.client.process_hls_stream(hls_url)
|
||||
elif audio is not None:
|
||||
resampled_file = utils.resample(audio)
|
||||
self.client.play_file(resampled_file)
|
||||
else:
|
||||
self.client.record()
|
||||
if save_output_recording and not output_recording_filename.endswith(".wav"):
|
||||
raise ValueError(f"Please provide a valid `output_recording_filename`: {output_recording_filename}")
|
||||
if not output_transcription_path.endswith(".srt"):
|
||||
raise ValueError(f"Please provide a valid `output_transcription_path`: {output_transcription_path}. The file extension should be `.srt`.")
|
||||
TranscriptionTeeClient.__init__(
|
||||
self,
|
||||
[self.client],
|
||||
save_output_recording=save_output_recording,
|
||||
output_recording_filename=output_recording_filename
|
||||
)
|
||||
|
||||
+237
-84
@@ -4,6 +4,9 @@ import threading
|
||||
import json
|
||||
import functools
|
||||
import logging
|
||||
from enum import Enum
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from websockets.sync.server import serve
|
||||
@@ -121,19 +124,41 @@ class ClientManager:
|
||||
return False
|
||||
|
||||
|
||||
class BackendType(Enum):
|
||||
FASTER_WHISPER = "faster_whisper"
|
||||
TENSORRT = "tensorrt"
|
||||
|
||||
@staticmethod
|
||||
def valid_types() -> List[str]:
|
||||
return [backend_type.value for backend_type in BackendType]
|
||||
|
||||
@staticmethod
|
||||
def is_valid(backend: str) -> bool:
|
||||
return backend in BackendType.valid_types()
|
||||
|
||||
def is_faster_whisper(self) -> bool:
|
||||
return self == BackendType.FASTER_WHISPER
|
||||
|
||||
def is_tensorrt(self) -> bool:
|
||||
return self == BackendType.TENSORRT
|
||||
|
||||
|
||||
class TranscriptionServer:
|
||||
RATE = 16000
|
||||
|
||||
def __init__(self):
|
||||
self.client_manager = ClientManager()
|
||||
self.client_manager = None
|
||||
self.no_voice_activity_chunks = 0
|
||||
self.use_vad = True
|
||||
self.single_model = False
|
||||
|
||||
def initialize_client(
|
||||
self, websocket, options, faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path, trt_multilingual
|
||||
):
|
||||
if self.backend == "tensorrt":
|
||||
client: Optional[ServeClientBase] = None
|
||||
|
||||
if self.backend.is_tensorrt():
|
||||
try:
|
||||
client = ServeClientTensorRT(
|
||||
websocket,
|
||||
@@ -141,7 +166,8 @@ class TranscriptionServer:
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"],
|
||||
model=whisper_tensorrt_path
|
||||
model=whisper_tensorrt_path,
|
||||
single_model=self.single_model,
|
||||
)
|
||||
logging.info("Running TensorRT backend.")
|
||||
except Exception as e:
|
||||
@@ -153,23 +179,31 @@ class TranscriptionServer:
|
||||
"message": "TensorRT-LLM not supported on Server yet. "
|
||||
"Reverting to available backend: 'faster_whisper'"
|
||||
}))
|
||||
self.backend = "faster_whisper"
|
||||
self.backend = BackendType.FASTER_WHISPER
|
||||
|
||||
if self.backend == "faster_whisper":
|
||||
if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path):
|
||||
logging.info(f"Using custom model {faster_whisper_custom_model_path}")
|
||||
options["model"] = faster_whisper_custom_model_path
|
||||
client = ServeClientFasterWhisper(
|
||||
websocket,
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"],
|
||||
model=options["model"],
|
||||
initial_prompt=options.get("initial_prompt"),
|
||||
vad_parameters=options.get("vad_parameters"),
|
||||
use_vad=self.use_vad,
|
||||
)
|
||||
logging.info("Running faster_whisper backend.")
|
||||
try:
|
||||
if self.backend.is_faster_whisper():
|
||||
if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path):
|
||||
logging.info(f"Using custom model {faster_whisper_custom_model_path}")
|
||||
options["model"] = faster_whisper_custom_model_path
|
||||
client = ServeClientFasterWhisper(
|
||||
websocket,
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"],
|
||||
model=options["model"],
|
||||
initial_prompt=options.get("initial_prompt"),
|
||||
vad_parameters=options.get("vad_parameters"),
|
||||
use_vad=self.use_vad,
|
||||
single_model=self.single_model,
|
||||
)
|
||||
|
||||
logging.info("Running faster_whisper backend.")
|
||||
except Exception as e:
|
||||
return
|
||||
|
||||
if client is None:
|
||||
raise ValueError(f"Backend type {self.backend.value} not recognised or not handled.")
|
||||
|
||||
self.client_manager.add_client(websocket, client)
|
||||
|
||||
@@ -194,12 +228,18 @@ class TranscriptionServer:
|
||||
logging.info("New client connected")
|
||||
options = websocket.recv()
|
||||
options = json.loads(options)
|
||||
|
||||
if self.client_manager is None:
|
||||
max_clients = options.get('max_clients', 4)
|
||||
max_connection_time = options.get('max_connection_time', 600)
|
||||
self.client_manager = ClientManager(max_clients, max_connection_time)
|
||||
|
||||
self.use_vad = options.get('use_vad')
|
||||
if self.client_manager.is_server_full(websocket, options):
|
||||
websocket.close()
|
||||
return False # Indicates that the connection should not continue
|
||||
|
||||
if self.backend == "tensorrt":
|
||||
if self.backend.is_tensorrt():
|
||||
self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE)
|
||||
self.initialize_client(websocket, options, faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path, trt_multilingual)
|
||||
@@ -218,11 +258,11 @@ class TranscriptionServer:
|
||||
frame_np = self.get_audio_from_websocket(websocket)
|
||||
client = self.client_manager.get_client(websocket)
|
||||
if frame_np is False:
|
||||
if self.backend == "tensorrt":
|
||||
if self.backend.is_tensorrt():
|
||||
client.set_eos(True)
|
||||
return False
|
||||
|
||||
if self.backend == "tensorrt":
|
||||
if self.backend.is_tensorrt():
|
||||
voice_active = self.voice_activity(websocket, frame_np)
|
||||
if voice_active:
|
||||
self.no_voice_activity_chunks = 0
|
||||
@@ -235,7 +275,7 @@ class TranscriptionServer:
|
||||
|
||||
def recv_audio(self,
|
||||
websocket,
|
||||
backend="faster_whisper",
|
||||
backend: BackendType = BackendType.FASTER_WHISPER,
|
||||
faster_whisper_custom_model_path=None,
|
||||
whisper_tensorrt_path=None,
|
||||
trt_multilingual=False):
|
||||
@@ -288,7 +328,8 @@ class TranscriptionServer:
|
||||
backend="tensorrt",
|
||||
faster_whisper_custom_model_path=None,
|
||||
whisper_tensorrt_path=None,
|
||||
trt_multilingual=False):
|
||||
trt_multilingual=False,
|
||||
single_model=False):
|
||||
"""
|
||||
Run the transcription server.
|
||||
|
||||
@@ -296,10 +337,23 @@ class TranscriptionServer:
|
||||
host (str): The host address to bind the server.
|
||||
port (int): The port number to bind the server.
|
||||
"""
|
||||
if faster_whisper_custom_model_path is not None and not os.path.exists(faster_whisper_custom_model_path):
|
||||
raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path.")
|
||||
if whisper_tensorrt_path is not None and not os.path.exists(whisper_tensorrt_path):
|
||||
raise ValueError(f"TensorRT model '{whisper_tensorrt_path}' is not a valid path.")
|
||||
if single_model:
|
||||
if faster_whisper_custom_model_path or whisper_tensorrt_path:
|
||||
logging.info("Custom model option was provided. Switching to single model mode.")
|
||||
self.single_model = True
|
||||
# TODO: load model initially
|
||||
else:
|
||||
logging.info("Single model mode currently only works with custom models.")
|
||||
if not BackendType.is_valid(backend):
|
||||
raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}")
|
||||
with serve(
|
||||
functools.partial(
|
||||
self.recv_audio,
|
||||
backend=backend,
|
||||
backend=BackendType(backend),
|
||||
faster_whisper_custom_model_path=faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path=whisper_tensorrt_path,
|
||||
trt_multilingual=trt_multilingual
|
||||
@@ -367,7 +421,7 @@ class ServeClientBase(object):
|
||||
self.prev_out = ''
|
||||
self.t_start = None
|
||||
self.exit = False
|
||||
self.same_output_threshold = 0
|
||||
self.same_output_count = 0
|
||||
self.show_prev_out_thresh = 5 # if pause(no output from whisper) show previous output for 5 seconds
|
||||
self.add_pause_thresh = 3 # add a blank to segment list as a pause(no speech) for 3 seconds
|
||||
self.transcript = []
|
||||
@@ -408,6 +462,11 @@ class ServeClientBase(object):
|
||||
if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE:
|
||||
self.frames_offset += 30.0
|
||||
self.frames_np = self.frames_np[int(30*self.RATE):]
|
||||
# check timestamp offset(should be >= self.frame_offset)
|
||||
# this basically means that there is no speech as timestamp offset hasnt updated
|
||||
# and is less than frame_offset
|
||||
if self.timestamp_offset < self.frames_offset:
|
||||
self.timestamp_offset = self.frames_offset
|
||||
if self.frames_np is None:
|
||||
self.frames_np = frame_np.copy()
|
||||
else:
|
||||
@@ -420,9 +479,10 @@ class ServeClientBase(object):
|
||||
Clip audio if the current chunk exceeds 30 seconds, this basically implies that
|
||||
no valid segment for the last 30 seconds from whisper
|
||||
"""
|
||||
if self.frames_np[int((self.timestamp_offset - self.frames_offset)*self.RATE):].shape[0] > 25 * self.RATE:
|
||||
duration = self.frames_np.shape[0] / self.RATE
|
||||
self.timestamp_offset = self.frames_offset + duration - 5
|
||||
with self.lock:
|
||||
if self.frames_np[int((self.timestamp_offset - self.frames_offset)*self.RATE):].shape[0] > 25 * self.RATE:
|
||||
duration = self.frames_np.shape[0] / self.RATE
|
||||
self.timestamp_offset = self.frames_offset + duration - 5
|
||||
|
||||
def get_audio_chunk_for_processing(self):
|
||||
"""
|
||||
@@ -438,8 +498,9 @@ class ServeClientBase(object):
|
||||
- input_bytes (np.ndarray): The next chunk of audio data to be processed.
|
||||
- duration (float): The duration of the audio chunk in seconds.
|
||||
"""
|
||||
samples_take = max(0, (self.timestamp_offset - self.frames_offset) * self.RATE)
|
||||
input_bytes = self.frames_np[int(samples_take):].copy()
|
||||
with self.lock:
|
||||
samples_take = max(0, (self.timestamp_offset - self.frames_offset) * self.RATE)
|
||||
input_bytes = self.frames_np[int(samples_take):].copy()
|
||||
duration = input_bytes.shape[0] / self.RATE
|
||||
return input_bytes, duration
|
||||
|
||||
@@ -527,7 +588,11 @@ class ServeClientBase(object):
|
||||
|
||||
|
||||
class ServeClientTensorRT(ServeClientBase):
|
||||
def __init__(self, websocket, task="transcribe", multilingual=False, language=None, client_uid=None, model=None):
|
||||
|
||||
SINGLE_MODEL = None
|
||||
SINGLE_MODEL_LOCK = threading.Lock()
|
||||
|
||||
def __init__(self, websocket, task="transcribe", multilingual=False, language=None, client_uid=None, model=None, single_model=False):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
@@ -541,21 +606,22 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
multilingual (bool, optional): Whether the client supports multilingual transcription. Defaults to False.
|
||||
language (str, optional): The language for transcription. Defaults to None.
|
||||
client_uid (str, optional): A unique identifier for the client. Defaults to None.
|
||||
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
||||
|
||||
"""
|
||||
super().__init__(client_uid, websocket)
|
||||
self.language = language if multilingual else "en"
|
||||
self.task = task
|
||||
self.eos = False
|
||||
self.transcriber = WhisperTRTLLM(
|
||||
model,
|
||||
assets_dir="assets",
|
||||
device="cuda",
|
||||
is_multilingual=multilingual,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
self.warmup()
|
||||
|
||||
if single_model:
|
||||
if ServeClientTensorRT.SINGLE_MODEL is None:
|
||||
self.create_model(model, multilingual)
|
||||
ServeClientTensorRT.SINGLE_MODEL = self.transcriber
|
||||
else:
|
||||
self.transcriber = ServeClientTensorRT.SINGLE_MODEL
|
||||
else:
|
||||
self.create_model(model, multilingual)
|
||||
|
||||
# threading
|
||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||
@@ -567,6 +633,21 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
"backend": "tensorrt"
|
||||
}))
|
||||
|
||||
def create_model(self, model, multilingual, warmup=True):
|
||||
"""
|
||||
Instantiates a new model, sets it as the transcriber and does warmup if desired.
|
||||
"""
|
||||
self.transcriber = WhisperTRTLLM(
|
||||
model,
|
||||
assets_dir="assets",
|
||||
device="cuda",
|
||||
is_multilingual=multilingual,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
if warmup:
|
||||
self.warmup()
|
||||
|
||||
def warmup(self, warmup_steps=10):
|
||||
"""
|
||||
Warmup TensorRT since first few inferences are slow.
|
||||
@@ -575,7 +656,7 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
warmup_steps (int): Number of steps to warm up the model for.
|
||||
"""
|
||||
logging.info("[INFO:] Warming up TensorRT engine..")
|
||||
mel, _ = self.transcriber.log_mel_spectrogram("tests/jfk.flac")
|
||||
mel, _ = self.transcriber.log_mel_spectrogram("assets/jfk.flac")
|
||||
for i in range(warmup_steps):
|
||||
self.transcriber.transcribe(mel)
|
||||
|
||||
@@ -611,9 +692,16 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
Args:
|
||||
input_bytes (np.array): The audio chunk to transcribe.
|
||||
"""
|
||||
if ServeClientTensorRT.SINGLE_MODEL:
|
||||
ServeClientTensorRT.SINGLE_MODEL_LOCK.acquire()
|
||||
logging.info(f"[WhisperTensorRT:] Processing audio with duration: {input_bytes.shape[0] / self.RATE}")
|
||||
mel, duration = self.transcriber.log_mel_spectrogram(input_bytes)
|
||||
last_segment = self.transcriber.transcribe(mel)
|
||||
last_segment = self.transcriber.transcribe(
|
||||
mel,
|
||||
text_prefix=f"<|startoftranscript|><|{self.language}|><|{self.task}|><|notimestamps|>"
|
||||
)
|
||||
if ServeClientTensorRT.SINGLE_MODEL:
|
||||
ServeClientTensorRT.SINGLE_MODEL_LOCK.release()
|
||||
if last_segment:
|
||||
self.handle_transcription_output(last_segment, duration)
|
||||
|
||||
@@ -629,7 +717,9 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
self.transcript.append({"text": last_segment + " "})
|
||||
elif self.transcript[-1]["text"].strip() != last_segment:
|
||||
self.transcript.append({"text": last_segment + " "})
|
||||
self.timestamp_offset += duration
|
||||
|
||||
with self.lock:
|
||||
self.timestamp_offset += duration
|
||||
|
||||
def speech_to_text(self):
|
||||
"""
|
||||
@@ -673,8 +763,12 @@ class ServeClientTensorRT(ServeClientBase):
|
||||
|
||||
|
||||
class ServeClientFasterWhisper(ServeClientBase):
|
||||
|
||||
SINGLE_MODEL = None
|
||||
SINGLE_MODEL_LOCK = threading.Lock()
|
||||
|
||||
def __init__(self, websocket, task="transcribe", device=None, language=None, client_uid=None, model="small.en",
|
||||
initial_prompt=None, vad_parameters=None, use_vad=True):
|
||||
initial_prompt=None, vad_parameters=None, use_vad=True, single_model=False):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
@@ -689,33 +783,55 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
client_uid (str, optional): A unique identifier for the client. Defaults to None.
|
||||
model (str, optional): The whisper model size. Defaults to 'small.en'
|
||||
initial_prompt (str, optional): Prompt for whisper inference. Defaults to None.
|
||||
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
||||
"""
|
||||
super().__init__(client_uid, websocket)
|
||||
self.model_sizes = [
|
||||
"tiny", "tiny.en", "base", "base.en", "small", "small.en",
|
||||
"medium", "medium.en", "large-v2", "large-v3",
|
||||
"medium", "medium.en", "large-v2", "large-v3", "distil-small.en",
|
||||
"distil-medium.en", "distil-large-v2", "distil-large-v3",
|
||||
"large-v3-turbo", "turbo"
|
||||
]
|
||||
if not os.path.exists(model):
|
||||
self.model_size_or_path = self.check_valid_model(model)
|
||||
else:
|
||||
self.model_size_or_path = model
|
||||
|
||||
self.model_size_or_path = model
|
||||
self.language = "en" if self.model_size_or_path.endswith("en") else language
|
||||
self.task = task
|
||||
self.initial_prompt = initial_prompt
|
||||
self.vad_parameters = vad_parameters or {"threshold": 0.5}
|
||||
self.vad_parameters = vad_parameters or {"onset": 0.5}
|
||||
self.no_speech_thresh = 0.45
|
||||
self.same_output_threshold = 10
|
||||
self.end_time_for_same_output = None
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if device == "cuda":
|
||||
major, _ = torch.cuda.get_device_capability(device)
|
||||
self.compute_type = "float16" if major >= 7 else "float32"
|
||||
else:
|
||||
self.compute_type = "int8"
|
||||
|
||||
if self.model_size_or_path is None:
|
||||
return
|
||||
logging.info(f"Using Device={device} with precision {self.compute_type}")
|
||||
|
||||
try:
|
||||
if single_model:
|
||||
if ServeClientFasterWhisper.SINGLE_MODEL is None:
|
||||
self.create_model(device)
|
||||
ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber
|
||||
else:
|
||||
self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL
|
||||
else:
|
||||
self.create_model(device)
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to load model: {e}")
|
||||
self.websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"status": "ERROR",
|
||||
"message": f"Failed to load model: {str(self.model_size_or_path)}"
|
||||
}))
|
||||
self.websocket.close()
|
||||
return
|
||||
|
||||
self.transcriber = WhisperModel(
|
||||
self.model_size_or_path,
|
||||
device=device,
|
||||
compute_type="int8" if device == "cpu" else "float16",
|
||||
local_files_only=False,
|
||||
)
|
||||
self.use_vad = use_vad
|
||||
|
||||
# threading
|
||||
@@ -731,6 +847,17 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
)
|
||||
)
|
||||
|
||||
def create_model(self, device):
|
||||
"""
|
||||
Instantiates a new model, sets it as the transcriber.
|
||||
"""
|
||||
self.transcriber = WhisperModel(
|
||||
self.model_size_or_path,
|
||||
device=device,
|
||||
compute_type=self.compute_type,
|
||||
local_files_only=False,
|
||||
)
|
||||
|
||||
def check_valid_model(self, model_size):
|
||||
"""
|
||||
Check if it's a valid whisper model size.
|
||||
@@ -786,6 +913,8 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
depends on the implementation of the `transcriber.transcribe` method but typically
|
||||
includes the transcribed text.
|
||||
"""
|
||||
if ServeClientFasterWhisper.SINGLE_MODEL:
|
||||
ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire()
|
||||
result, info = self.transcriber.transcribe(
|
||||
input_sample,
|
||||
initial_prompt=self.initial_prompt,
|
||||
@@ -793,7 +922,10 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
task=self.task,
|
||||
vad_filter=self.use_vad,
|
||||
vad_parameters=self.vad_parameters if self.use_vad else None)
|
||||
if self.language is None:
|
||||
if ServeClientFasterWhisper.SINGLE_MODEL:
|
||||
ServeClientFasterWhisper.SINGLE_MODEL_LOCK.release()
|
||||
|
||||
if self.language is None and info is not None:
|
||||
self.set_language(info)
|
||||
return result
|
||||
|
||||
@@ -873,12 +1005,15 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
|
||||
input_bytes, duration = self.get_audio_chunk_for_processing()
|
||||
if duration < 1.0:
|
||||
time.sleep(0.1) # wait for audio chunks to arrive
|
||||
continue
|
||||
try:
|
||||
input_sample = input_bytes.copy()
|
||||
result = self.transcribe_audio(input_sample)
|
||||
|
||||
if self.language is None:
|
||||
if result is None or self.language is None:
|
||||
self.timestamp_offset += duration
|
||||
time.sleep(0.25) # wait for voice activity, result is None when no voice activity
|
||||
continue
|
||||
self.handle_transcription_output(result, duration)
|
||||
|
||||
@@ -886,7 +1021,7 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
logging.error(f"[ERROR]: Failed to transcribe audio chunk: {e}")
|
||||
time.sleep(0.01)
|
||||
|
||||
def format_segment(self, start, end, text):
|
||||
def format_segment(self, start, end, text, completed=False):
|
||||
"""
|
||||
Formats a transcription segment with precise start and end times alongside the transcribed text.
|
||||
|
||||
@@ -903,7 +1038,8 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
return {
|
||||
'start': "{:.3f}".format(start),
|
||||
'end': "{:.3f}".format(end),
|
||||
'text': text
|
||||
'text': text,
|
||||
'completed': completed
|
||||
}
|
||||
|
||||
def update_segments(self, segments, duration):
|
||||
@@ -930,52 +1066,69 @@ class ServeClientFasterWhisper(ServeClientBase):
|
||||
offset = None
|
||||
self.current_out = ''
|
||||
last_segment = None
|
||||
|
||||
# process complete segments
|
||||
if len(segments) > 1:
|
||||
if len(segments) > 1 and segments[-1].no_speech_prob <= self.no_speech_thresh:
|
||||
for i, s in enumerate(segments[:-1]):
|
||||
text_ = s.text
|
||||
self.text.append(text_)
|
||||
start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end)
|
||||
with self.lock:
|
||||
start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end)
|
||||
|
||||
if start >= end:
|
||||
continue
|
||||
if s.no_speech_prob > self.no_speech_thresh:
|
||||
continue
|
||||
|
||||
self.transcript.append(self.format_segment(start, end, text_))
|
||||
self.transcript.append(self.format_segment(start, end, text_, completed=True))
|
||||
offset = min(duration, s.end)
|
||||
|
||||
self.current_out += segments[-1].text
|
||||
last_segment = self.format_segment(
|
||||
self.timestamp_offset + segments[-1].start,
|
||||
self.timestamp_offset + min(duration, segments[-1].end),
|
||||
self.current_out
|
||||
)
|
||||
# only process the last segment if it satisfies the no_speech_thresh
|
||||
if segments[-1].no_speech_prob <= self.no_speech_thresh:
|
||||
self.current_out += segments[-1].text
|
||||
with self.lock:
|
||||
last_segment = self.format_segment(
|
||||
self.timestamp_offset + segments[-1].start,
|
||||
self.timestamp_offset + min(duration, segments[-1].end),
|
||||
self.current_out,
|
||||
completed=False
|
||||
)
|
||||
|
||||
if self.current_out.strip() == self.prev_out.strip() and self.current_out != '':
|
||||
self.same_output_count += 1
|
||||
|
||||
# if we remove the audio because of same output on the nth reptition we might remove the
|
||||
# audio thats not yet transcribed so, capturing the time when it was repeated for the first time
|
||||
if self.end_time_for_same_output is None:
|
||||
self.end_time_for_same_output = segments[-1].end
|
||||
time.sleep(0.1) # wait for some voice activity just in case there is an unitended pause from the speaker for better punctuations.
|
||||
else:
|
||||
self.same_output_count = 0
|
||||
self.end_time_for_same_output = None
|
||||
|
||||
# if same incomplete segment is seen multiple times then update the offset
|
||||
# and append the segment to the list
|
||||
if self.current_out.strip() == self.prev_out.strip() and self.current_out != '':
|
||||
self.same_output_threshold += 1
|
||||
else:
|
||||
self.same_output_threshold = 0
|
||||
|
||||
if self.same_output_threshold > 5:
|
||||
if self.same_output_count > self.same_output_threshold:
|
||||
if not len(self.text) or self.text[-1].strip().lower() != self.current_out.strip().lower():
|
||||
self.text.append(self.current_out)
|
||||
self.transcript.append(self.format_segment(
|
||||
self.timestamp_offset,
|
||||
self.timestamp_offset + duration,
|
||||
self.current_out
|
||||
))
|
||||
with self.lock:
|
||||
self.transcript.append(self.format_segment(
|
||||
self.timestamp_offset,
|
||||
self.timestamp_offset + min(duration, self.end_time_for_same_output),
|
||||
self.current_out,
|
||||
completed=True
|
||||
))
|
||||
self.current_out = ''
|
||||
offset = duration
|
||||
self.same_output_threshold = 0
|
||||
offset = min(duration, self.end_time_for_same_output)
|
||||
self.same_output_count = 0
|
||||
last_segment = None
|
||||
self.end_time_for_same_output = None
|
||||
else:
|
||||
self.prev_out = self.current_out
|
||||
|
||||
# update offset
|
||||
if offset is not None:
|
||||
self.timestamp_offset += offset
|
||||
with self.lock:
|
||||
self.timestamp_offset += offset
|
||||
|
||||
return last_segment
|
||||
|
||||
+1146
-296
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import re
|
||||
import math
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Union
|
||||
@@ -14,7 +15,8 @@ import tensorrt_llm
|
||||
import tensorrt_llm.logger as logger
|
||||
from tensorrt_llm._utils import (str_dtype_to_torch, str_dtype_to_trt,
|
||||
trt_dtype_to_torch)
|
||||
from tensorrt_llm.runtime import ModelConfig, SamplingConfig
|
||||
from tensorrt_llm.bindings import GptJsonConfig, KVCacheType
|
||||
from tensorrt_llm.runtime import PYTHON_BINDINGS, ModelConfig, SamplingConfig
|
||||
from tensorrt_llm.runtime.session import Session, TensorInfo
|
||||
|
||||
|
||||
@@ -24,39 +26,102 @@ HOP_LENGTH = 160
|
||||
CHUNK_LENGTH = 30
|
||||
N_SAMPLES = CHUNK_LENGTH * SAMPLE_RATE # 480000 samples in a 30-second chunk
|
||||
|
||||
def read_config(component, engine_dir):
|
||||
config_path = engine_dir / component / 'config.json'
|
||||
with open(config_path, 'r') as f:
|
||||
config = json.load(f)
|
||||
model_config = OrderedDict()
|
||||
model_config.update(config['pretrained_config'])
|
||||
model_config.update(config['build_config'])
|
||||
return model_config
|
||||
|
||||
|
||||
def remove_tensor_padding(input_tensor,
|
||||
input_tensor_lengths=None,
|
||||
pad_value=None):
|
||||
if pad_value:
|
||||
assert input_tensor_lengths is None, "input_tensor_lengths should be None when pad_value is provided"
|
||||
# Text tensor case: batch, seq_len
|
||||
assert torch.all(
|
||||
input_tensor[:, 0] != pad_value
|
||||
), "First token in each sequence should not be pad_value"
|
||||
assert input_tensor_lengths is None
|
||||
|
||||
# Create a mask for all non-pad tokens
|
||||
mask = input_tensor != pad_value
|
||||
|
||||
# Apply the mask to input_tensor to remove pad tokens
|
||||
output_tensor = input_tensor[mask].view(1, -1)
|
||||
|
||||
else:
|
||||
# Audio tensor case: batch, seq_len, feature_len
|
||||
# position_ids case: batch, seq_len
|
||||
assert input_tensor_lengths is not None, "input_tensor_lengths must be provided for 3D input_tensor"
|
||||
|
||||
# Initialize a list to collect valid sequences
|
||||
valid_sequences = []
|
||||
|
||||
for i in range(input_tensor.shape[0]):
|
||||
valid_length = input_tensor_lengths[i]
|
||||
valid_sequences.append(input_tensor[i, :valid_length])
|
||||
|
||||
# Concatenate all valid sequences along the batch dimension
|
||||
output_tensor = torch.cat(valid_sequences, dim=0)
|
||||
return output_tensor
|
||||
|
||||
|
||||
class WhisperEncoding:
|
||||
|
||||
def __init__(self, engine_dir):
|
||||
self.session = self.get_session(engine_dir)
|
||||
config = read_config('encoder', engine_dir)
|
||||
self.n_mels = config['n_mels']
|
||||
self.dtype = config['dtype']
|
||||
self.num_languages = config['num_languages']
|
||||
self.encoder_config = config
|
||||
|
||||
def get_session(self, engine_dir):
|
||||
config_path = engine_dir / 'encoder_config.json'
|
||||
with open(config_path, 'r') as f:
|
||||
config = json.load(f)
|
||||
|
||||
dtype = config['builder_config']['precision']
|
||||
n_mels = config['builder_config']['n_mels']
|
||||
num_languages = config['builder_config']['num_languages']
|
||||
|
||||
self.dtype = dtype
|
||||
self.n_mels = n_mels
|
||||
self.num_languages = num_languages
|
||||
|
||||
serialize_path = engine_dir / f'whisper_encoder_{self.dtype}_tp1_rank0.engine'
|
||||
|
||||
serialize_path = engine_dir / 'encoder' / 'rank0.engine'
|
||||
with open(serialize_path, 'rb') as f:
|
||||
session = Session.from_serialized_engine(f.read())
|
||||
|
||||
return session
|
||||
|
||||
def get_audio_features(self, mel):
|
||||
def get_audio_features(self,
|
||||
mel,
|
||||
mel_input_lengths,
|
||||
encoder_downsampling_factor=2):
|
||||
if isinstance(mel, list):
|
||||
longest_mel = max([f.shape[-1] for f in mel])
|
||||
mel = [
|
||||
torch.nn.functional.pad(f, (0, longest_mel - f.shape[-1]),
|
||||
mode='constant') for f in mel
|
||||
]
|
||||
mel = torch.cat(mel, dim=0).type(
|
||||
str_dtype_to_torch("float16")).contiguous()
|
||||
bsz, seq_len = mel.shape[0], mel.shape[2]
|
||||
position_ids = torch.arange(
|
||||
math.ceil(seq_len / encoder_downsampling_factor),
|
||||
dtype=torch.int32,
|
||||
device=mel.device).expand(bsz, -1).contiguous()
|
||||
if self.encoder_config['plugin_config']['remove_input_padding']:
|
||||
# mel B,D,T -> B,T,D -> BxT, D
|
||||
mel = mel.transpose(1, 2)
|
||||
mel = remove_tensor_padding(mel, mel_input_lengths)
|
||||
position_ids = remove_tensor_padding(
|
||||
position_ids, mel_input_lengths // encoder_downsampling_factor)
|
||||
inputs = OrderedDict()
|
||||
output_list = []
|
||||
inputs['input_features'] = mel
|
||||
inputs['input_lengths'] = mel_input_lengths
|
||||
inputs['position_ids'] = position_ids
|
||||
|
||||
inputs.update({'x': mel})
|
||||
output_list.append(
|
||||
TensorInfo('x', str_dtype_to_trt(self.dtype), mel.shape))
|
||||
output_list = [
|
||||
TensorInfo('input_features', str_dtype_to_trt(self.dtype),
|
||||
mel.shape),
|
||||
TensorInfo('input_lengths', str_dtype_to_trt('int32'),
|
||||
mel_input_lengths.shape),
|
||||
TensorInfo('position_ids', str_dtype_to_trt('int32'),
|
||||
inputs['position_ids'].shape)
|
||||
]
|
||||
|
||||
output_info = (self.session).infer_shapes(output_list)
|
||||
|
||||
@@ -73,46 +138,44 @@ class WhisperEncoding:
|
||||
stream=stream.cuda_stream)
|
||||
assert ok, 'Engine execution failed'
|
||||
stream.synchronize()
|
||||
audio_features = outputs['output']
|
||||
return audio_features
|
||||
encoder_output = outputs['encoder_output']
|
||||
encoder_output_lengths = mel_input_lengths // encoder_downsampling_factor
|
||||
return encoder_output, encoder_output_lengths
|
||||
|
||||
|
||||
class WhisperDecoding:
|
||||
|
||||
def __init__(self, engine_dir, runtime_mapping, debug_mode=False):
|
||||
|
||||
self.decoder_config = self.get_config(engine_dir)
|
||||
self.decoder_config = read_config('decoder', engine_dir)
|
||||
self.decoder_generation_session = self.get_session(
|
||||
engine_dir, runtime_mapping, debug_mode)
|
||||
|
||||
def get_config(self, engine_dir):
|
||||
config_path = engine_dir / 'decoder_config.json'
|
||||
with open(config_path, 'r') as f:
|
||||
config = json.load(f)
|
||||
decoder_config = OrderedDict()
|
||||
decoder_config.update(config['plugin_config'])
|
||||
decoder_config.update(config['builder_config'])
|
||||
return decoder_config
|
||||
|
||||
def get_session(self, engine_dir, runtime_mapping, debug_mode=False):
|
||||
dtype = self.decoder_config['precision']
|
||||
serialize_path = engine_dir / f'whisper_decoder_{dtype}_tp1_rank0.engine'
|
||||
serialize_path = engine_dir / 'decoder' / 'rank0.engine'
|
||||
with open(serialize_path, "rb") as f:
|
||||
decoder_engine_buffer = f.read()
|
||||
|
||||
decoder_model_config = ModelConfig(
|
||||
num_heads=self.decoder_config['num_heads'],
|
||||
num_kv_heads=self.decoder_config['num_heads'],
|
||||
max_batch_size=self.decoder_config['max_batch_size'],
|
||||
max_beam_width=self.decoder_config['max_beam_width'],
|
||||
num_heads=self.decoder_config['num_attention_heads'],
|
||||
num_kv_heads=self.decoder_config['num_attention_heads'],
|
||||
hidden_size=self.decoder_config['hidden_size'],
|
||||
vocab_size=self.decoder_config['vocab_size'],
|
||||
num_layers=self.decoder_config['num_layers'],
|
||||
gpt_attention_plugin=self.decoder_config['gpt_attention_plugin'],
|
||||
remove_input_padding=self.decoder_config['remove_input_padding'],
|
||||
cross_attention=self.decoder_config['cross_attention'],
|
||||
cross_attention=True,
|
||||
num_layers=self.decoder_config['num_hidden_layers'],
|
||||
gpt_attention_plugin=self.decoder_config['plugin_config']
|
||||
['gpt_attention_plugin'],
|
||||
remove_input_padding=self.decoder_config['plugin_config']
|
||||
['remove_input_padding'],
|
||||
kv_cache_type=KVCacheType.PAGED
|
||||
if self.decoder_config['plugin_config']['paged_kv_cache'] == True
|
||||
else KVCacheType.CONTINUOUS,
|
||||
has_position_embedding=self.
|
||||
decoder_config['has_position_embedding'],
|
||||
has_token_type_embedding=self.
|
||||
decoder_config['has_token_type_embedding'],
|
||||
dtype=self.decoder_config['dtype'],
|
||||
has_token_type_embedding=False,
|
||||
)
|
||||
decoder_generation_session = tensorrt_llm.runtime.GenerationSession(
|
||||
decoder_model_config,
|
||||
@@ -125,14 +188,12 @@ class WhisperDecoding:
|
||||
def generate(self,
|
||||
decoder_input_ids,
|
||||
encoder_outputs,
|
||||
encoder_max_input_length,
|
||||
encoder_input_lengths,
|
||||
eot_id,
|
||||
max_new_tokens=40,
|
||||
num_beams=1):
|
||||
encoder_input_lengths = torch.tensor(
|
||||
[encoder_outputs.shape[1] for x in range(encoder_outputs.shape[0])],
|
||||
dtype=torch.int32,
|
||||
device='cuda')
|
||||
|
||||
batch_size = decoder_input_ids.shape[0]
|
||||
decoder_input_lengths = torch.tensor([
|
||||
decoder_input_ids.shape[-1]
|
||||
for _ in range(decoder_input_ids.shape[0])
|
||||
@@ -141,6 +202,10 @@ class WhisperDecoding:
|
||||
device='cuda')
|
||||
decoder_max_input_length = torch.max(decoder_input_lengths).item()
|
||||
|
||||
cross_attention_mask = torch.ones([
|
||||
batch_size, decoder_max_input_length + max_new_tokens,
|
||||
encoder_max_input_length
|
||||
]).int().cuda()
|
||||
# generation config
|
||||
sampling_config = SamplingConfig(end_id=eot_id,
|
||||
pad_id=eot_id,
|
||||
@@ -150,17 +215,31 @@ class WhisperDecoding:
|
||||
decoder_max_input_length,
|
||||
max_new_tokens,
|
||||
beam_width=num_beams,
|
||||
encoder_max_input_length=encoder_outputs.shape[1])
|
||||
encoder_max_input_length=encoder_max_input_length)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
|
||||
decoder_input_ids = decoder_input_ids.type(torch.int32).cuda()
|
||||
if self.decoder_config['plugin_config']['remove_input_padding']:
|
||||
# 50256 is the index of <pad> for all whisper models' decoder
|
||||
WHISPER_PAD_TOKEN_ID = 50256
|
||||
decoder_input_ids = remove_tensor_padding(
|
||||
decoder_input_ids, pad_value=WHISPER_PAD_TOKEN_ID)
|
||||
if encoder_outputs.dim() == 3:
|
||||
encoder_output_lens = torch.full((encoder_outputs.shape[0], ),
|
||||
encoder_outputs.shape[1],
|
||||
dtype=torch.int32,
|
||||
device='cuda')
|
||||
|
||||
encoder_outputs = remove_tensor_padding(encoder_outputs,
|
||||
encoder_output_lens)
|
||||
output_ids = self.decoder_generation_session.decode(
|
||||
decoder_input_ids,
|
||||
decoder_input_lengths,
|
||||
sampling_config,
|
||||
encoder_output=encoder_outputs,
|
||||
encoder_input_lengths=encoder_input_lengths,
|
||||
cross_attention_mask=cross_attention_mask,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
@@ -178,18 +257,23 @@ class WhisperTRTLLM(object):
|
||||
runtime_mapping = tensorrt_llm.Mapping(world_size, runtime_rank)
|
||||
torch.cuda.set_device(runtime_rank % runtime_mapping.gpus_per_node)
|
||||
engine_dir = Path(engine_dir)
|
||||
encoder_config = read_config('encoder', engine_dir)
|
||||
decoder_config = read_config('decoder', engine_dir)
|
||||
self.n_mels = encoder_config['n_mels']
|
||||
self.num_languages = encoder_config['num_languages']
|
||||
is_multilingual = (decoder_config['vocab_size'] >= 51865)
|
||||
|
||||
self.encoder = WhisperEncoding(engine_dir)
|
||||
self.decoder = WhisperDecoding(engine_dir,
|
||||
runtime_mapping,
|
||||
debug_mode=False)
|
||||
runtime_mapping,
|
||||
debug_mode=False)
|
||||
self.n_mels = self.encoder.n_mels
|
||||
# self.tokenizer = get_tokenizer(num_languages=self.encoder.num_languages,
|
||||
# tokenizer_dir=assets_dir)
|
||||
self.device = device
|
||||
self.tokenizer = get_tokenizer(
|
||||
is_multilingual,
|
||||
num_languages=self.encoder.num_languages,
|
||||
num_languages=self.num_languages,
|
||||
language=language,
|
||||
task=task,
|
||||
)
|
||||
@@ -256,8 +340,10 @@ class WhisperTRTLLM(object):
|
||||
def process_batch(
|
||||
self,
|
||||
mel,
|
||||
mel_input_lengths,
|
||||
text_prefix="<|startoftranscript|><|en|><|transcribe|><|notimestamps|>",
|
||||
num_beams=1):
|
||||
num_beams=1,
|
||||
max_new_tokens=96):
|
||||
prompt_id = self.tokenizer.encode(
|
||||
text_prefix, allowed_special=set(self.tokenizer.special_tokens.keys()))
|
||||
|
||||
@@ -265,11 +351,14 @@ class WhisperTRTLLM(object):
|
||||
batch_size = mel.shape[0]
|
||||
decoder_input_ids = prompt_id.repeat(batch_size, 1)
|
||||
|
||||
encoder_output = self.encoder.get_audio_features(mel)
|
||||
encoder_output, encoder_output_lengths = self.encoder.get_audio_features(mel, mel_input_lengths)
|
||||
encoder_max_input_length = torch.max(encoder_output_lengths).item()
|
||||
output_ids = self.decoder.generate(decoder_input_ids,
|
||||
encoder_output,
|
||||
encoder_max_input_length,
|
||||
encoder_output_lengths,
|
||||
self.tokenizer.eot,
|
||||
max_new_tokens=96,
|
||||
max_new_tokens=max_new_tokens,
|
||||
num_beams=num_beams)
|
||||
texts = []
|
||||
for i in range(len(output_ids)):
|
||||
@@ -284,10 +373,22 @@ class WhisperTRTLLM(object):
|
||||
dtype='float16',
|
||||
batch_size=1,
|
||||
num_beams=1,
|
||||
padding_strategy="max",
|
||||
):
|
||||
mel = mel.type(str_dtype_to_torch(dtype))
|
||||
mel = mel.unsqueeze(0)
|
||||
predictions = self.process_batch(mel, text_prefix, num_beams)
|
||||
# repeat the mel spectrogram to match the batch size
|
||||
mel = mel.repeat(batch_size, 1, 1)
|
||||
if padding_strategy == "longest":
|
||||
pass
|
||||
else:
|
||||
mel = torch.nn.functional.pad(mel, (0, 3000 - mel.shape[2]))
|
||||
features_input_lengths = torch.full((mel.shape[0], ),
|
||||
mel.shape[2],
|
||||
dtype=torch.int32,
|
||||
device=mel.device)
|
||||
|
||||
predictions = self.process_batch(mel, features_input_lengths, text_prefix, num_beams)
|
||||
prediction = predictions[0]
|
||||
|
||||
# remove all special tokens in the prediction
|
||||
|
||||
+30
-15
@@ -1,10 +1,9 @@
|
||||
# original: https://github.com/snakers4/silero-vad/blob/master/utils_vad.py
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import torch
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
import warnings
|
||||
|
||||
|
||||
class VoiceActivityDetection():
|
||||
@@ -24,7 +23,11 @@ class VoiceActivityDetection():
|
||||
self.session = onnxruntime.InferenceSession(path, providers=['CUDAExecutionProvider'], sess_options=opts)
|
||||
|
||||
self.reset_states()
|
||||
self.sample_rates = [8000, 16000]
|
||||
if '16k' in path:
|
||||
warnings.warn('This model support only 16000 sampling rate!')
|
||||
self.sample_rates = [16000]
|
||||
else:
|
||||
self.sample_rates = [8000, 16000]
|
||||
|
||||
def _validate_input(self, x, sr: int):
|
||||
if x.dim() == 1:
|
||||
@@ -34,27 +37,32 @@ class VoiceActivityDetection():
|
||||
|
||||
if sr != 16000 and (sr % 16000 == 0):
|
||||
step = sr // 16000
|
||||
x = x[:, ::step]
|
||||
x = x[:,::step]
|
||||
sr = 16000
|
||||
|
||||
if sr not in self.sample_rates:
|
||||
raise ValueError(f"Supported sampling rates: {self.sample_rates} (or multiply of 16000)")
|
||||
|
||||
if sr / x.shape[1] > 31.25:
|
||||
raise ValueError("Input audio chunk is too short")
|
||||
|
||||
return x, sr
|
||||
|
||||
def reset_states(self, batch_size=1):
|
||||
self._h = np.zeros((2, batch_size, 64)).astype('float32')
|
||||
self._c = np.zeros((2, batch_size, 64)).astype('float32')
|
||||
self._state = torch.zeros((2, batch_size, 128)).float()
|
||||
self._context = torch.zeros(0)
|
||||
self._last_sr = 0
|
||||
self._last_batch_size = 0
|
||||
|
||||
def __call__(self, x, sr: int):
|
||||
|
||||
x, sr = self._validate_input(x, sr)
|
||||
num_samples = 512 if sr == 16000 else 256
|
||||
|
||||
if x.shape[-1] != num_samples:
|
||||
raise ValueError(f"Provided number of samples is {x.shape[-1]} (Supported values: 256 for 8000 sample rate, 512 for 16000)")
|
||||
|
||||
batch_size = x.shape[0]
|
||||
context_size = 64 if sr == 16000 else 32
|
||||
|
||||
if not self._last_batch_size:
|
||||
self.reset_states(batch_size)
|
||||
@@ -63,28 +71,35 @@ class VoiceActivityDetection():
|
||||
if (self._last_batch_size) and (self._last_batch_size != batch_size):
|
||||
self.reset_states(batch_size)
|
||||
|
||||
if not len(self._context):
|
||||
self._context = torch.zeros(batch_size, context_size)
|
||||
|
||||
x = torch.cat([self._context, x], dim=1)
|
||||
if sr in [8000, 16000]:
|
||||
ort_inputs = {'input': x.numpy(), 'h': self._h, 'c': self._c, 'sr': np.array(sr, dtype='int64')}
|
||||
ort_inputs = {'input': x.numpy(), 'state': self._state.numpy(), 'sr': np.array(sr, dtype='int64')}
|
||||
ort_outs = self.session.run(None, ort_inputs)
|
||||
out, self._h, self._c = ort_outs
|
||||
out, state = ort_outs
|
||||
self._state = torch.from_numpy(state)
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self._context = x[..., -context_size:]
|
||||
self._last_sr = sr
|
||||
self._last_batch_size = batch_size
|
||||
|
||||
out = torch.tensor(out)
|
||||
out = torch.from_numpy(out)
|
||||
return out
|
||||
|
||||
def audio_forward(self, x, sr: int, num_samples: int = 512):
|
||||
def audio_forward(self, x, sr: int):
|
||||
outs = []
|
||||
x, sr = self._validate_input(x, sr)
|
||||
self.reset_states()
|
||||
num_samples = 512 if sr == 16000 else 256
|
||||
|
||||
if x.shape[1] % num_samples:
|
||||
pad_num = num_samples - (x.shape[1] % num_samples)
|
||||
x = torch.nn.functional.pad(x, (0, pad_num), 'constant', value=0.0)
|
||||
|
||||
self.reset_states(x.shape[0])
|
||||
for i in range(0, x.shape[1], num_samples):
|
||||
wavs_batch = x[:, i:i+num_samples]
|
||||
out_chunk = self.__call__(wavs_batch, sr)
|
||||
@@ -94,7 +109,7 @@ class VoiceActivityDetection():
|
||||
return stacked.cpu()
|
||||
|
||||
@staticmethod
|
||||
def download(model_url="https://github.com/snakers4/silero-vad/raw/master/files/silero_vad.onnx"):
|
||||
def download(model_url="https://github.com/snakers4/silero-vad/raw/v5.0/files/silero_vad.onnx"):
|
||||
target_dir = os.path.expanduser("~/.cache/whisper-live/")
|
||||
|
||||
# Ensure the target directory exists
|
||||
@@ -138,5 +153,5 @@ class VoiceActivityDetector:
|
||||
bool: True if the speech probability exceeds the threshold, indicating the presence of voice activity;
|
||||
False otherwise.
|
||||
"""
|
||||
speech_prob = self.model(torch.from_numpy(audio_frame), self.frame_rate).item()
|
||||
return speech_prob > self.threshold
|
||||
speech_probs = self.model.audio_forward(torch.from_numpy(audio_frame.copy()), self.frame_rate)[0]
|
||||
return torch.any(speech_probs > self.threshold).item()
|
||||
|
||||
Reference in New Issue
Block a user