Compare commits
363 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4b46371dac | |||
| 1f4c918d01 | |||
| cd327bab50 | |||
| b91b3664c2 | |||
| 2375924b45 | |||
| ae169245a1 | |||
| 4ba576fb06 | |||
| a27ac16d1f | |||
| 188b21f1d0 | |||
| d29993048d | |||
| 41d9f683a8 | |||
| d9d8d511c7 | |||
| 275ed4e45b | |||
| 9cfd8f85b6 | |||
| 7fb2d356f9 | |||
| af50fed180 | |||
| a2271806c3 | |||
| 0abf8693ef | |||
| 444a1df740 | |||
| 47ee035f65 | |||
| d9cb4ffdd0 | |||
| 9b364f267a | |||
| 617f587699 | |||
| fb3deb2745 | |||
| 5e430f8154 | |||
| efb51bf0fa | |||
| 2abca69c9d | |||
| a62495b090 | |||
| c1ac71ada0 | |||
| f5bea0a693 | |||
| 5b3bef5845 | |||
| 2c761adc32 | |||
| 379bd146fc | |||
| e93c2823b1 | |||
| 87520498e9 | |||
| 23d71fdbce | |||
| ef7c32dc95 | |||
| 28be23340b | |||
| ba5aa5aa38 | |||
| 779baff9c3 | |||
| 5aa5826f36 | |||
| 893265bb3f | |||
| 5120afbc25 | |||
| 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 | |||
| c067224474 | |||
| e92f53cfd9 | |||
| 308ac1cff7 | |||
| 2fced08705 | |||
| 8bdaf9249d | |||
| b47a56ca6d | |||
| f975bd452e | |||
| cb963c4834 | |||
| 1db94ea96e | |||
| babe5de074 | |||
| 99af50208d | |||
| dc22b7da9f | |||
| 2e9f67ba0b | |||
| c919ba3501 | |||
| a38fdb494d | |||
| fddc244228 | |||
| 5e1174ff33 | |||
| e40414ab1b | |||
| 17873c66a0 | |||
| 1147f58225 | |||
| 0baa1dc0a6 | |||
| d1de4948ee | |||
| ff871ad485 | |||
| 5fe5e0c8ba | |||
| b42ced9816 | |||
| 06794470f8 | |||
| d530957b2c | |||
| 6cabbe441b | |||
| 78da3f6750 | |||
| b04cffc458 | |||
| fd7c5965b3 | |||
| 4471665085 | |||
| 147e97002e | |||
| e3c7666cf7 | |||
| 8266099ed0 | |||
| 01dc69e068 | |||
| 9bb92b9bb2 | |||
| 57c4b60e04 | |||
| 3cd96367fb | |||
| c1420cba0d | |||
| 4db91eed66 | |||
| 7bcb92c266 | |||
| 170ba22e5b | |||
| ac00e28b86 | |||
| ceb3cc8747 | |||
| eaec0ead08 | |||
| 9fbff47126 | |||
| b4abe95fc6 | |||
| 14974af951 | |||
| bc474b4a76 | |||
| 9ccf940f51 | |||
| 9a9972007e | |||
| b2ad6478f5 | |||
| 490efdeacc | |||
| cb570d28ce | |||
| 4e5e086c38 | |||
| cf78d5d608 | |||
| 9d29b08cea | |||
| 6071cc1cc5 | |||
| f98e309663 | |||
| 30b00d6c89 | |||
| da2992bcaf | |||
| 16c5ed8ce9 | |||
| e14fefb671 | |||
| 98399707a3 | |||
| 567ceb1246 | |||
| 28ea8a20f1 | |||
| acf6dfe5b7 | |||
| 84a97f5fdd | |||
| ca2634bbb6 | |||
| 444ce63440 | |||
| 8db063ee33 | |||
| 92cbc37e9c | |||
| 24fd835356 | |||
| 5409d14bcb | |||
| d6edf8e847 | |||
| cc3ed74c0e | |||
| 20a8a8ad3d | |||
| 07387abbc0 | |||
| 4ecc59783e | |||
| ec9074d712 | |||
| 3a25db4cb9 | |||
| 4924ec0adb | |||
| f35abc7f81 | |||
| f383121ec3 | |||
| 17d62272cf | |||
| 91e1b75bfc | |||
| 7aad2ae721 | |||
| 6a1b82f953 | |||
| dc84839873 | |||
| 56d19f5469 | |||
| 08fa183ba4 | |||
| d89b27b8aa | |||
| ce68cc6c87 | |||
| e697574870 | |||
| b098b52a4d | |||
| ad5543b03e | |||
| 60455b1583 | |||
| cb458fc207 | |||
| 32ed089a76 | |||
| 5b28ddefbd | |||
| e1f531eccf | |||
| 8200207530 | |||
| 08575a03c2 | |||
| f590446865 | |||
| 36d137888e | |||
| e64bc9f3d6 | |||
| 7cc945aded | |||
| 2c8a25d355 | |||
| f4027de343 | |||
| d1754d2c46 | |||
| 0e6b1c0632 | |||
| d6b51ccd7d | |||
| 2f3c1cd172 | |||
| 703263b375 | |||
| 30d2cffb93 | |||
| 4d94c6b38b | |||
| d5a0f5859e | |||
| 025873d2ca | |||
| 8c36768f7f | |||
| ce13e7b622 | |||
| 3498787ccd | |||
| 5cd59b1e4c | |||
| bd543295f3 | |||
| 8e2642283a | |||
| 3bf5b47947 | |||
| 634dae835b | |||
| 969a5aa9e5 | |||
| 44a2e20c68 | |||
| 986823dbef | |||
| 1e2faa3f2b | |||
| e3084b34cb | |||
| b955e63dc1 | |||
| f25ff1785a | |||
| 867ff522ae | |||
| 75001ae6b7 | |||
| 6f1d13f25b | |||
| 7a9dc6db40 | |||
| 735d6c7763 | |||
| 0942dc2cfd | |||
| 881fd55776 | |||
| 0c01d7b1e5 | |||
| c810369324 | |||
| 71d0fe69c6 | |||
| 67232fffd5 | |||
| 076aebf3b6 | |||
| 4cf9d95f73 | |||
| 389bb5ae37 | |||
| a7eedc5d84 | |||
| d91330d790 | |||
| 783d147316 | |||
| 058c93e55e | |||
| f06b9bc827 | |||
| 3c202bf836 | |||
| 647c576e6a | |||
| 71a062b726 | |||
| a26f990586 | |||
| ddb1e0947f | |||
| 6dff4fbdd3 | |||
| 244ca9e6ba | |||
| 0f9e93d203 | |||
| fd86340f30 | |||
| 2300eedc8b | |||
| cafcb04fbc | |||
| 72ead71eeb | |||
| 7b2f5cff72 | |||
| 32c6a565d7 | |||
| e30286c046 | |||
| 7c0b32b85e | |||
| 01665a54c1 | |||
| 02793a93f8 | |||
| db2e0bbcdd | |||
| e92ddd291a | |||
| 5918b5ed42 | |||
| 71d207a607 | |||
| 6ee4cd09f2 | |||
| 5de4de4b84 | |||
| e006722da7 | |||
| a52dc0cbf8 | |||
| 048ab0a8f4 | |||
| 261bb9e961 | |||
| 091f6179d4 | |||
| 1e1349cd80 | |||
| 14beb4f942 | |||
| 7ffcad64ba | |||
| 402fceb9f3 | |||
| 09b18e8ab8 | |||
| da72d03073 | |||
| a1a8d5f92a | |||
| b6dee4e46e | |||
| f3cd20fbf3 | |||
| da86c18205 | |||
| 8097e9b44a | |||
| 222852ff33 | |||
| 2de67ee02f | |||
| 073cfc20f3 | |||
| ee80bd21bd | |||
| 410b91d133 | |||
| a2b5220738 | |||
| 1938dfb490 |
+182
-36
@@ -1,4 +1,4 @@
|
||||
name: CI
|
||||
name: Test & Build CI/CD
|
||||
|
||||
on:
|
||||
push:
|
||||
@@ -7,46 +7,192 @@ on:
|
||||
tags:
|
||||
- v*
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
branches: [ main ]
|
||||
types: [opened, synchronize, reopened]
|
||||
|
||||
jobs:
|
||||
build-and-push-package:
|
||||
runs-on: ubuntu-latest
|
||||
run-tests:
|
||||
runs-on: ubuntu-22.04
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: [3.9, '3.10', 3.11, 3.12]
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Cache Python dependencies
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cache/pip
|
||||
!~/.cache/pip/log
|
||||
key: ${{ runner.os }}-pip-${{ matrix.python-version }}-${{ hashFiles('requirements/server.txt', 'requirements/client.txt') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-pip-${{ matrix.python-version }}-
|
||||
|
||||
- name: Install system dependencies
|
||||
run: sudo apt-get update && sudo apt-get install -y portaudio19-dev
|
||||
|
||||
- name: Install Python dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements/server.txt --extra-index-url https://download.pytorch.org/whl/cpu
|
||||
pip install -r requirements/client.txt
|
||||
|
||||
- name: Run tests
|
||||
run: |
|
||||
echo "Running tests with Python ${{ matrix.python-version }}"
|
||||
python -m unittest discover -s tests
|
||||
|
||||
check-code-format:
|
||||
runs-on: ubuntu-22.04
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: [3.9, '3.10', 3.11, 3.12]
|
||||
|
||||
steps:
|
||||
- name: Check Out Repository
|
||||
uses: actions/checkout@v2
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: 3.8
|
||||
|
||||
- name: Set up FFmpeg
|
||||
uses: FedericoCarboni/setup-ffmpeg@v2
|
||||
|
||||
- name: Install Additional requirements
|
||||
run: |
|
||||
sudo apt-get -y install portaudio19-dev wget
|
||||
shell: bash
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install Client Requirements
|
||||
run: pip install -r requirements/client.txt
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
python -m pip install flake8
|
||||
|
||||
- name: Install Server Requirements
|
||||
run: pip install -r requirements/server.txt
|
||||
- name: Lint with flake8
|
||||
run: |
|
||||
# stop the build if there are Python syntax errors or undefined names
|
||||
flake8 . --count --select=E9,F63,F7,F82 --show-source --statistics
|
||||
# exit-zero treats all errors as warnings. The GitHub editor is 127 chars wide
|
||||
flake8 . --count --exit-zero --max-complexity=10 --max-line-length=127 --statistics
|
||||
|
||||
- name: Install Wheel for build
|
||||
run: pip install wheel twine
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
python setup.py sdist bdist_wheel
|
||||
|
||||
- name: Push package on Test PyPI
|
||||
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags')
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
user: __token__
|
||||
password: ${{ secrets.PYPI_API_TOKEN }}
|
||||
build-and-push-docker-cpu:
|
||||
needs: [run-tests, check-code-format]
|
||||
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: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v1
|
||||
|
||||
- name: Build and push Docker image
|
||||
uses: docker/build-push-action@v2
|
||||
with:
|
||||
context: .
|
||||
file: docker/Dockerfile.cpu
|
||||
push: true
|
||||
tags: ghcr.io/collabora/whisperlive-cpu: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' && (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.gpu
|
||||
push: true
|
||||
tags: ghcr.io/collabora/whisperlive-gpu:latest
|
||||
|
||||
build-and-push-docker-openvino:
|
||||
needs: [run-tests, check-code-format, build-and-push-docker-cpu]
|
||||
timeout-minutes: 20
|
||||
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.openvino
|
||||
push: true
|
||||
tags: ghcr.io/collabora/whisperlive-openvino:latest
|
||||
|
||||
publish-to-pypi:
|
||||
needs: [run-tests, check-code-format]
|
||||
runs-on: ubuntu-22.04
|
||||
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags')
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
- name: Set up Python 3.9
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: 3.9
|
||||
|
||||
- name: Cache Python dependencies
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cache/pip
|
||||
!~/.cache/pip/log
|
||||
key: ubuntu-latest-pip-3.9-${{ hashFiles('requirements/server.txt', 'requirements/client.txt') }}
|
||||
restore-keys: |
|
||||
ubuntu-latest-pip-3.9-
|
||||
|
||||
- name: Install system dependencies
|
||||
run: sudo apt-get update && sudo apt-get install -y portaudio19-dev
|
||||
|
||||
- name: Install Python dependencies
|
||||
run: |
|
||||
pip install -r requirements/server.txt
|
||||
pip install -r requirements/client.txt
|
||||
pip install wheel
|
||||
|
||||
- name: Build package
|
||||
run: python setup.py sdist bdist_wheel
|
||||
|
||||
- name: Publish package to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
user: __token__
|
||||
password: ${{ secrets.PYPI_API_TOKEN }}
|
||||
|
||||
@@ -26,9 +26,9 @@ To capture the audio in the current tab, we used the chrome `tabCapture` API to
|
||||
### Options
|
||||
When using the Audio Transcription extension, you have the following options:
|
||||
- **Use Collabora Server**: We provide a demo server which runs the whisper small model.
|
||||
- **Use Multilingual Model**: Enable this option to utilize the multilingual capabilities of OpenAI-whisper.
|
||||
- **Language**: Select the target language for transcription or translation. You can choose from a variety of languages supported by OpenAI-whisper.
|
||||
- **Task:** Choose the specific task to perform on the audio. You can select either "transcribe" for transcription or "translate" to translate the audio to English.
|
||||
- **Model Size**: Select the whisper model size to run the server with.
|
||||
|
||||
### Getting Started
|
||||
- Make sure the transcription server is running properly. To know more about how to start the server, see the [documentation here](https://github.com/collabora/whisper-live).
|
||||
|
||||
@@ -156,7 +156,9 @@ async function startCapture(options) {
|
||||
port: options.port,
|
||||
multilingual: options.useMultilingual,
|
||||
language: options.language,
|
||||
task: options.task
|
||||
task: options.task,
|
||||
modelSize: options.modelSize,
|
||||
useVad: options.useVad,
|
||||
},
|
||||
});
|
||||
} else {
|
||||
@@ -207,13 +209,3 @@ chrome.runtime.onMessage.addListener(async (message) => {
|
||||
});
|
||||
|
||||
|
||||
/**
|
||||
* Listens for if the tab is reloaded.
|
||||
* @param {Object} message - The message received from the runtime.
|
||||
*/
|
||||
chrome.tabs.onUpdated.addListener(async (tabId, changeInfo, tab) => {
|
||||
if (changeInfo.status === 'complete') {
|
||||
await executeScriptInTab(tabId, "content.js");
|
||||
await delayExecution(500);
|
||||
}
|
||||
});
|
||||
|
||||
@@ -59,7 +59,7 @@ function init_element() {
|
||||
|
||||
elem_container = document.createElement('div');
|
||||
elem_container.id = "transcription";
|
||||
elem_container.style.cssText = 'padding-top:16px;font-size:18px;line-height:18px;top:0px;position:absolute;width:500px;height:90px;opacity:0.9;z-index:100;background:black;border-radius:10px;color:white;';
|
||||
elem_container.style.cssText = 'padding-top:16px;font-size:18px;position: fixed; top: 85%; left: 50%; transform: translate(-50%, -50%);line-height:18px;width:500px;height:90px;opacity:0.9;z-index:100;background:black;border-radius:10px;color:white;';
|
||||
|
||||
for (var i = 0; i < 4; i++) {
|
||||
elem_text = document.createElement('span');
|
||||
@@ -173,13 +173,13 @@ chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
|
||||
if (type === "STOP") {
|
||||
remove_element();
|
||||
sendResponse({data: "STOPPED"});
|
||||
return;
|
||||
return true;
|
||||
} else if (type === "showWaitPopup"){
|
||||
initPopupElement();
|
||||
|
||||
showPopup(`Estimated wait time ~ ${Math.round(data)} minutes`);
|
||||
sendResponse({data: "popup"});
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
|
||||
init_element();
|
||||
@@ -234,4 +234,5 @@ chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
|
||||
}
|
||||
|
||||
sendResponse({});
|
||||
return true;
|
||||
});
|
||||
|
||||
@@ -93,16 +93,14 @@ async function startRecord(option) {
|
||||
const socket = new WebSocket(`ws://${option.host}:${option.port}/`);
|
||||
let isServerReady = false;
|
||||
let language = option.language;
|
||||
if (language === null && !option.multilingual) {
|
||||
language = 'en';
|
||||
}
|
||||
socket.onopen = function(e) {
|
||||
socket.send(
|
||||
JSON.stringify({
|
||||
uid: uuid,
|
||||
multilingual: option.multilingual,
|
||||
language: option.language,
|
||||
task: option.task
|
||||
task: option.task,
|
||||
model: option.modelSize,
|
||||
use_vad: option.useVad
|
||||
})
|
||||
);
|
||||
};
|
||||
@@ -184,16 +182,17 @@ async function startRecord(option) {
|
||||
* @param {Object} sender - The sender object containing information about the message sender.
|
||||
* @param {Function} sendResponse - The function to send a response back to the message sender.
|
||||
*/
|
||||
chrome.runtime.onMessage.addListener(async (request, sender, sendResponse) => {
|
||||
chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
|
||||
const { type, data } = request;
|
||||
|
||||
switch (type) {
|
||||
case "start_capture":
|
||||
await startRecord(data);
|
||||
startRecord(data);
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
|
||||
sendResponse({});
|
||||
return true;
|
||||
});
|
||||
|
||||
@@ -16,120 +16,137 @@
|
||||
<label for="useServerCheckbox">Use Collabora Whisper-Live Server</label>
|
||||
</div>
|
||||
<div class="checkbox-container">
|
||||
<input type="checkbox" id="useMultilingualCheckbox">
|
||||
<label for="useMultilingualCheckbox">Use Multilingual Model</label>
|
||||
<input type="checkbox" id="useVadCheckbox">
|
||||
<label for="useVadCheckbox">Use Voice Activity Detection</label>
|
||||
</div>
|
||||
<div class="dropdown-container">
|
||||
<label for="languageDropdown">Select Language:</label>
|
||||
<select id="languageDropdown" disabled>
|
||||
<option value="">Select Language</option>
|
||||
<option value="zh">Chinese</option>
|
||||
<option value="de">German</option>
|
||||
<option value="es">Spanish</option>
|
||||
<option value="ru">Russian</option>
|
||||
<option value="ko">Korean</option>
|
||||
<option value="fr">French</option>
|
||||
<option value="ja">Japanese</option>
|
||||
<option value="pt">Portuguese</option>
|
||||
<option value="tr">Turkish</option>
|
||||
<option value="pl">Polish</option>
|
||||
<option value="ca">Catalan</option>
|
||||
<option value="nl">Dutch</option>
|
||||
<option value="ar">Arabic</option>
|
||||
<option value="sv">Swedish</option>
|
||||
<option value="it">Italian</option>
|
||||
<option value="id">Indonesian</option>
|
||||
<option value="hi">Hindi</option>
|
||||
<option value="fi">Finnish</option>
|
||||
<option value="vi">Vietnamese</option>
|
||||
<option value="he">Hebrew</option>
|
||||
<option value="uk">Ukrainian</option>
|
||||
<option value="el">Greek</option>
|
||||
<option value="ms">Malay</option>
|
||||
<option value="cs">Czech</option>
|
||||
<option value="ro">Romanian</option>
|
||||
<option value="da">Danish</option>
|
||||
<option value="hu">Hungarian</option>
|
||||
<option value="ta">Tamil</option>
|
||||
<option value="no">Norwegian</option>
|
||||
<option value="th">Thai</option>
|
||||
<option value="ur">Urdu</option>
|
||||
<option value="hr">Croatian</option>
|
||||
<option value="bg">Bulgarian</option>
|
||||
<option value="lt">Lithuanian</option>
|
||||
<option value="la">Latin</option>
|
||||
<option value="mi">Maori</option>
|
||||
<option value="ml">Malayalam</option>
|
||||
<option value="cy">Welsh</option>
|
||||
<option value="sk">Slovak</option>
|
||||
<option value="te">Telugu</option>
|
||||
<option value="fa">Persian</option>
|
||||
<option value="lv">Latvian</option>
|
||||
<option value="bn">Bengali</option>
|
||||
<option value="sr">Serbian</option>
|
||||
<option value="az">Azerbaijani</option>
|
||||
<option value="sl">Slovenian</option>
|
||||
<option value="kn">Kannada</option>
|
||||
<option value="et">Estonian</option>
|
||||
<option value="mk">Macedonian</option>
|
||||
<option value="br">Breton</option>
|
||||
<option value="eu">Basque</option>
|
||||
<option value="is">Icelandic</option>
|
||||
<option value="hy">Armenian</option>
|
||||
<option value="ne">Nepali</option>
|
||||
<option value="mn">Mongolian</option>
|
||||
<option value="bs">Bosnian</option>
|
||||
<option value="kk">Kazakh</option>
|
||||
<option value="sq">Albanian</option>
|
||||
<option value="sw">Swahili</option>
|
||||
<option value="gl">Galician</option>
|
||||
<option value="mr">Marathi</option>
|
||||
<option value="pa">Punjabi</option>
|
||||
<option value="si">Sinhala</option>
|
||||
<option value="km">Khmer</option>
|
||||
<option value="sn">Shona</option>
|
||||
<option value="yo">Yoruba</option>
|
||||
<option value="so">Somali</option>
|
||||
<select id="languageDropdown">
|
||||
<option value="" selected>Automatically detect</option>
|
||||
<option value="af">Afrikaans</option>
|
||||
<option value="oc">Occitan</option>
|
||||
<option value="ka">Georgian</option>
|
||||
<option value="be">Belarusian</option>
|
||||
<option value="tg">Tajik</option>
|
||||
<option value="sd">Sindhi</option>
|
||||
<option value="gu">Gujarati</option>
|
||||
<option value="sq">Albanian</option>
|
||||
<option value="am">Amharic</option>
|
||||
<option value="yi">Yiddish</option>
|
||||
<option value="lo">Lao</option>
|
||||
<option value="uz">Uzbek</option>
|
||||
<option value="fo">Faroese</option>
|
||||
<option value="ht">Haitian Creole</option>
|
||||
<option value="ps">Pashto</option>
|
||||
<option value="tk">Turkmen</option>
|
||||
<option value="nn">Nynorsk</option>
|
||||
<option value="mt">Maltese</option>
|
||||
<option value="sa">Sanskrit</option>
|
||||
<option value="lb">Luxembourgish</option>
|
||||
<option value="my">Myanmar</option>
|
||||
<option value="bo">Tibetan</option>
|
||||
<option value="tl">Tagalog</option>
|
||||
<option value="mg">Malagasy</option>
|
||||
<option value="ar">Arabic</option>
|
||||
<option value="hy">Armenian</option>
|
||||
<option value="as">Assamese</option>
|
||||
<option value="tt">Tatar</option>
|
||||
<option value="haw">Hawaiian</option>
|
||||
<option value="ln">Lingala</option>
|
||||
<option value="ha">Hausa</option>
|
||||
<option value="az">Azerbaijani</option>
|
||||
<option value="ba">Bashkir</option>
|
||||
<option value="eu">Basque</option>
|
||||
<option value="be">Belarusian</option>
|
||||
<option value="bn">Bengali</option>
|
||||
<option value="bs">Bosnian</option>
|
||||
<option value="br">Breton</option>
|
||||
<option value="bg">Bulgarian</option>
|
||||
<option value="ca">Catalan</option>
|
||||
<option value="zh">Chinese</option>
|
||||
<option value="hr">Croatian</option>
|
||||
<option value="cs">Czech</option>
|
||||
<option value="da">Danish</option>
|
||||
<option value="nl">Dutch</option>
|
||||
<option value="en">English</option>
|
||||
<option value="et">Estonian</option>
|
||||
<option value="fo">Faroese</option>
|
||||
<option value="fi">Finnish</option>
|
||||
<option value="fr">French</option>
|
||||
<option value="gl">Galician</option>
|
||||
<option value="ka">Georgian</option>
|
||||
<option value="de">German</option>
|
||||
<option value="el">Greek</option>
|
||||
<option value="gu">Gujarati</option>
|
||||
<option value="ht">Haitian Creole</option>
|
||||
<option value="ha">Hausa</option>
|
||||
<option value="haw">Hawaiian</option>
|
||||
<option value="he">Hebrew</option>
|
||||
<option value="hi">Hindi</option>
|
||||
<option value="hu">Hungarian</option>
|
||||
<option value="is">Icelandic</option>
|
||||
<option value="id">Indonesian</option>
|
||||
<option value="it">Italian</option>
|
||||
<option value="ja">Japanese</option>
|
||||
<option value="jw">Javanese</option>
|
||||
<option value="kn">Kannada</option>
|
||||
<option value="kk">Kazakh</option>
|
||||
<option value="km">Khmer</option>
|
||||
<option value="ko">Korean</option>
|
||||
<option value="lo">Lao</option>
|
||||
<option value="la">Latin</option>
|
||||
<option value="lv">Latvian</option>
|
||||
<option value="ln">Lingala</option>
|
||||
<option value="lt">Lithuanian</option>
|
||||
<option value="lb">Luxembourgish</option>
|
||||
<option value="mk">Macedonian</option>
|
||||
<option value="mg">Malagasy</option>
|
||||
<option value="ms">Malay</option>
|
||||
<option value="ml">Malayalam</option>
|
||||
<option value="mt">Maltese</option>
|
||||
<option value="mi">Maori</option>
|
||||
<option value="mr">Marathi</option>
|
||||
<option value="mn">Mongolian</option>
|
||||
<option value="my">Myanmar</option>
|
||||
<option value="ne">Nepali</option>
|
||||
<option value="no">Norwegian</option>
|
||||
<option value="nn">Nynorsk</option>
|
||||
<option value="oc">Occitan</option>
|
||||
<option value="ps">Pashto</option>
|
||||
<option value="fa">Persian</option>
|
||||
<option value="pl">Polish</option>
|
||||
<option value="pt">Portuguese</option>
|
||||
<option value="pa">Punjabi</option>
|
||||
<option value="ro">Romanian</option>
|
||||
<option value="ru">Russian</option>
|
||||
<option value="sa">Sanskrit</option>
|
||||
<option value="sr">Serbian</option>
|
||||
<option value="sn">Shona</option>
|
||||
<option value="sd">Sindhi</option>
|
||||
<option value="si">Sinhala</option>
|
||||
<option value="sk">Slovak</option>
|
||||
<option value="sl">Slovenian</option>
|
||||
<option value="so">Somali</option>
|
||||
<option value="es">Spanish</option>
|
||||
<option value="su">Sundanese</option>
|
||||
<option value="sw">Swahili</option>
|
||||
<option value="sv">Swedish</option>
|
||||
<option value="tl">Tagalog</option>
|
||||
<option value="tg">Tajik</option>
|
||||
<option value="ta">Tamil</option>
|
||||
<option value="tt">Tatar</option>
|
||||
<option value="te">Telugu</option>
|
||||
<option value="th">Thai</option>
|
||||
<option value="bo">Tibetan</option>
|
||||
<option value="tr">Turkish</option>
|
||||
<option value="tk">Turkmen</option>
|
||||
<option value="uk">Ukrainian</option>
|
||||
<option value="ur">Urdu</option>
|
||||
<option value="uz">Uzbek</option>
|
||||
<option value="vi">Vietnamese</option>
|
||||
<option value="cy">Welsh</option>
|
||||
<option value="yi">Yiddish</option>
|
||||
<option value="yo">Yoruba</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="dropdown-container">
|
||||
<label for="taskDropdown">Select task:</label>
|
||||
<select id="taskDropdown" disabled>
|
||||
<select id="taskDropdown" >
|
||||
<option value="">Select Task</option>
|
||||
<option value="transcribe" selected>Transcribe</option>
|
||||
<option value="translate">Translate</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="dropdown-container">
|
||||
<label for="modelSizeDropdown">Select Model Size:</label>
|
||||
<select id="modelSizeDropdown">
|
||||
<option value="">Select model</option>
|
||||
<option value="tiny">Tiny </option>
|
||||
<option value="tiny.en">Tiny (English-only)</option>
|
||||
<option value="base">Base</option>
|
||||
<option value="base.en">Base (English-only)</option>
|
||||
<option value="small" selected>Small</option>
|
||||
<option value="small.en">Small (English-only)</option>
|
||||
<option value="medium">Medium</option>
|
||||
<option value="medium.en">Medium (English-only)</option>
|
||||
<option value="large-v2">Large-v2</option>
|
||||
<option value="large-v3">Large-v3</option>
|
||||
</select>
|
||||
</div>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
@@ -4,11 +4,13 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
const stopButton = document.getElementById("stopCapture");
|
||||
|
||||
const useServerCheckbox = document.getElementById("useServerCheckbox");
|
||||
const useMultilingualCheckbox = document.getElementById('useMultilingualCheckbox');
|
||||
const useVadCheckbox = document.getElementById("useVadCheckbox");
|
||||
const languageDropdown = document.getElementById('languageDropdown');
|
||||
const taskDropdown = document.getElementById('taskDropdown');
|
||||
const modelSizeDropdown = document.getElementById('modelSizeDropdown');
|
||||
let selectedLanguage = null;
|
||||
let selectedTask = taskDropdown.value;
|
||||
let selectedModelSize = modelSizeDropdown.value;
|
||||
|
||||
// Add click event listeners to the buttons
|
||||
startButton.addEventListener("click", startCapture);
|
||||
@@ -30,11 +32,9 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
}
|
||||
});
|
||||
|
||||
chrome.storage.local.get("useMultilingualModelState", ({ useMultilingualModelState }) => {
|
||||
if (useMultilingualModelState !== undefined) {
|
||||
useMultilingualCheckbox.checked = useMultilingualModelState;
|
||||
languageDropdown.disabled = !useMultilingualModelState;
|
||||
taskDropdown.disabled = !useMultilingualModelState;
|
||||
chrome.storage.local.get("useVadState", ({ useVadState }) => {
|
||||
if (useVadState !== undefined) {
|
||||
useVadCheckbox.checked = useVadState;
|
||||
}
|
||||
});
|
||||
|
||||
@@ -52,6 +52,13 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
}
|
||||
});
|
||||
|
||||
chrome.storage.local.get("selectedModelSize", ({ selectedModelSize: storedModelSize }) => {
|
||||
if (storedModelSize !== undefined) {
|
||||
modelSizeDropdown.value = storedModelSize;
|
||||
selectedModelSize = storedModelSize;
|
||||
}
|
||||
});
|
||||
|
||||
// Function to handle the start capture button click event
|
||||
async function startCapture() {
|
||||
// Ignore click if the button is disabled
|
||||
@@ -77,9 +84,10 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
tabId: currentTab.id,
|
||||
host: host,
|
||||
port: port,
|
||||
useMultilingual: useMultilingualCheckbox.checked,
|
||||
language: selectedLanguage,
|
||||
task: selectedTask
|
||||
task: selectedTask,
|
||||
modelSize: selectedModelSize,
|
||||
useVad: useVadCheckbox.checked,
|
||||
}, () => {
|
||||
// Update capturing state in storage and toggle the buttons
|
||||
chrome.storage.local.set({ capturingState: { isCapturing: true } }, () => {
|
||||
@@ -118,9 +126,11 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
function toggleCaptureButtons(isCapturing) {
|
||||
startButton.disabled = isCapturing;
|
||||
stopButton.disabled = !isCapturing;
|
||||
useServerCheckbox.disabled = isCapturing;
|
||||
useMultilingualCheckbox.disabled = isCapturing;
|
||||
|
||||
useServerCheckbox.disabled = isCapturing;
|
||||
useVadCheckbox.disabled = isCapturing;
|
||||
modelSizeDropdown.disabled = isCapturing;
|
||||
languageDropdown.disabled = isCapturing;
|
||||
taskDropdown.disabled = isCapturing;
|
||||
startButton.classList.toggle("disabled", isCapturing);
|
||||
stopButton.classList.toggle("disabled", !isCapturing);
|
||||
}
|
||||
@@ -131,16 +141,9 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
chrome.storage.local.set({ useServerState });
|
||||
});
|
||||
|
||||
useMultilingualCheckbox.addEventListener('change', function() {
|
||||
const useMultilingualModelState = useMultilingualCheckbox.checked;
|
||||
if (useMultilingualModelState) {
|
||||
languageDropdown.disabled = false;
|
||||
taskDropdown.disabled = false;
|
||||
} else {
|
||||
languageDropdown.disabled = true;
|
||||
taskDropdown.disabled = true;
|
||||
}
|
||||
chrome.storage.local.set({ useMultilingualModelState });
|
||||
useVadCheckbox.addEventListener("change", () => {
|
||||
const useVadState = useVadCheckbox.checked;
|
||||
chrome.storage.local.set({ useVadState });
|
||||
});
|
||||
|
||||
languageDropdown.addEventListener('change', function() {
|
||||
@@ -157,6 +160,11 @@ document.addEventListener("DOMContentLoaded", function () {
|
||||
chrome.storage.local.set({ selectedTask });
|
||||
});
|
||||
|
||||
modelSizeDropdown.addEventListener('change', function() {
|
||||
selectedModelSize = modelSizeDropdown.value;
|
||||
chrome.storage.local.set({ selectedModelSize });
|
||||
});
|
||||
|
||||
chrome.runtime.onMessage.addListener(async (request, sender, sendResponse) => {
|
||||
if (request.action === "updateSelectedLanguage") {
|
||||
const detectedLanguage = request.detectedLanguage;
|
||||
|
||||
@@ -24,9 +24,9 @@ To capture the audio in the current tab, we used the chrome `tabCapture` API to
|
||||
### Options
|
||||
When using the Audio Transcription extension, you have the following options:
|
||||
- **Use Collabora Server**: We provide a demo server which runs the whisper small model.
|
||||
- **Use Multilingual Model**: Enable this option to utilize the multilingual capabilities of OpenAI-whisper.
|
||||
- **Language**: Select the target language for transcription or translation. You can choose from a variety of languages supported by OpenAI-whisper.
|
||||
- **Task:** Choose the specific task to perform on the audio. You can select either "transcribe" for transcription or "translate" to translate the audio to English.
|
||||
- **Model Size**: Select the whisper model size to run the server with.
|
||||
|
||||
### Getting Started
|
||||
- Make sure the transcription server is running properly. To know more about how to start the server, see the [documentation here](https://github.com/collabora/whisper-live).
|
||||
|
||||
@@ -66,18 +66,16 @@ function resampleTo16kHZ(audioData, origSampleRate = 44100) {
|
||||
function startRecording(data) {
|
||||
socket = new WebSocket(`ws://${data.host}:${data.port}/`);
|
||||
language = data.language;
|
||||
if (language === null && !data.useMultilingual) {
|
||||
language = 'en';
|
||||
}
|
||||
|
||||
const uuid = generateUUID();
|
||||
socket.onopen = function(e) {
|
||||
socket.send(
|
||||
JSON.stringify({
|
||||
uid: uuid,
|
||||
multilingual: data.useMultilingual,
|
||||
language: data.language,
|
||||
task: data.task
|
||||
task: data.task,
|
||||
model: data.modelSize,
|
||||
use_vad: data.useVad
|
||||
})
|
||||
);
|
||||
};
|
||||
@@ -200,7 +198,7 @@ function init_element() {
|
||||
|
||||
elem_container = document.createElement('div');
|
||||
elem_container.id = "transcription";
|
||||
elem_container.style.cssText = 'padding-top:16px;font-size:18px;line-height:18px;top:0px;position:absolute;width:500px;height:90px;opacity:0.9;z-index:100;background:black;border-radius:10px;color:white;';
|
||||
elem_container.style.cssText = 'padding-top:16px;font-size:18px;line-height:18px;position:fixed;top:85%;left:50%;transform:translate(-50%,-50%);width:500px;height:90px;opacity:0.9;z-index:100;background:black;border-radius:10px;color:white;';
|
||||
|
||||
for (var i = 0; i < 4; i++) {
|
||||
elem_text = document.createElement('span');
|
||||
|
||||
@@ -15,114 +15,114 @@
|
||||
<input type="checkbox" id="useServerCheckbox">
|
||||
<label for="useServerCheckbox">Use Collabora Whisper-Live Server</label>
|
||||
</div>
|
||||
<textarea id="waitTextBox" style="display: none;"></textarea>
|
||||
|
||||
<div class="checkbox-container">
|
||||
<input type="checkbox" id="useMultilingualCheckbox">
|
||||
<label for="useMultilingualCheckbox">Use Multilingual Model</label>
|
||||
<input type="checkbox" id="useVadCheckbox">
|
||||
<label for="useVadCheckbox">Use Voice Activity Detection</label>
|
||||
</div>
|
||||
<textarea id="waitTextBox" style="display: none;"></textarea>
|
||||
<div class="dropdown-container">
|
||||
<label for="languageDropdown">Select Language:</label>
|
||||
<select id="languageDropdown" disabled>
|
||||
<option value="">Select Language</option>
|
||||
<option value="zh">Chinese</option>
|
||||
<option value="de">German</option>
|
||||
<option value="es">Spanish</option>
|
||||
<option value="ru">Russian</option>
|
||||
<option value="ko">Korean</option>
|
||||
<option value="fr">French</option>
|
||||
<option value="ja">Japanese</option>
|
||||
<option value="pt">Portuguese</option>
|
||||
<option value="tr">Turkish</option>
|
||||
<option value="pl">Polish</option>
|
||||
<option value="ca">Catalan</option>
|
||||
<option value="nl">Dutch</option>
|
||||
<option value="ar">Arabic</option>
|
||||
<option value="sv">Swedish</option>
|
||||
<option value="it">Italian</option>
|
||||
<option value="id">Indonesian</option>
|
||||
<option value="hi">Hindi</option>
|
||||
<option value="fi">Finnish</option>
|
||||
<option value="vi">Vietnamese</option>
|
||||
<option value="he">Hebrew</option>
|
||||
<option value="uk">Ukrainian</option>
|
||||
<option value="el">Greek</option>
|
||||
<option value="ms">Malay</option>
|
||||
<option value="cs">Czech</option>
|
||||
<option value="ro">Romanian</option>
|
||||
<option value="da">Danish</option>
|
||||
<option value="hu">Hungarian</option>
|
||||
<option value="ta">Tamil</option>
|
||||
<option value="no">Norwegian</option>
|
||||
<option value="th">Thai</option>
|
||||
<option value="ur">Urdu</option>
|
||||
<option value="hr">Croatian</option>
|
||||
<option value="bg">Bulgarian</option>
|
||||
<option value="lt">Lithuanian</option>
|
||||
<option value="la">Latin</option>
|
||||
<option value="mi">Maori</option>
|
||||
<option value="ml">Malayalam</option>
|
||||
<option value="cy">Welsh</option>
|
||||
<option value="sk">Slovak</option>
|
||||
<option value="te">Telugu</option>
|
||||
<option value="fa">Persian</option>
|
||||
<option value="lv">Latvian</option>
|
||||
<option value="bn">Bengali</option>
|
||||
<option value="sr">Serbian</option>
|
||||
<option value="az">Azerbaijani</option>
|
||||
<option value="sl">Slovenian</option>
|
||||
<option value="kn">Kannada</option>
|
||||
<option value="et">Estonian</option>
|
||||
<option value="mk">Macedonian</option>
|
||||
<option value="br">Breton</option>
|
||||
<option value="eu">Basque</option>
|
||||
<option value="is">Icelandic</option>
|
||||
<option value="hy">Armenian</option>
|
||||
<option value="ne">Nepali</option>
|
||||
<option value="mn">Mongolian</option>
|
||||
<option value="bs">Bosnian</option>
|
||||
<option value="kk">Kazakh</option>
|
||||
<option value="sq">Albanian</option>
|
||||
<option value="sw">Swahili</option>
|
||||
<option value="gl">Galician</option>
|
||||
<option value="mr">Marathi</option>
|
||||
<option value="pa">Punjabi</option>
|
||||
<option value="si">Sinhala</option>
|
||||
<option value="km">Khmer</option>
|
||||
<option value="sn">Shona</option>
|
||||
<option value="yo">Yoruba</option>
|
||||
<option value="so">Somali</option>
|
||||
<select id="languageDropdown">
|
||||
<option value="" selected>Automatically detect</option>
|
||||
<option value="af">Afrikaans</option>
|
||||
<option value="oc">Occitan</option>
|
||||
<option value="ka">Georgian</option>
|
||||
<option value="be">Belarusian</option>
|
||||
<option value="tg">Tajik</option>
|
||||
<option value="sd">Sindhi</option>
|
||||
<option value="gu">Gujarati</option>
|
||||
<option value="sq">Albanian</option>
|
||||
<option value="am">Amharic</option>
|
||||
<option value="yi">Yiddish</option>
|
||||
<option value="lo">Lao</option>
|
||||
<option value="uz">Uzbek</option>
|
||||
<option value="fo">Faroese</option>
|
||||
<option value="ht">Haitian Creole</option>
|
||||
<option value="ps">Pashto</option>
|
||||
<option value="tk">Turkmen</option>
|
||||
<option value="nn">Nynorsk</option>
|
||||
<option value="mt">Maltese</option>
|
||||
<option value="sa">Sanskrit</option>
|
||||
<option value="lb">Luxembourgish</option>
|
||||
<option value="my">Myanmar</option>
|
||||
<option value="bo">Tibetan</option>
|
||||
<option value="tl">Tagalog</option>
|
||||
<option value="mg">Malagasy</option>
|
||||
<option value="ar">Arabic</option>
|
||||
<option value="hy">Armenian</option>
|
||||
<option value="as">Assamese</option>
|
||||
<option value="tt">Tatar</option>
|
||||
<option value="haw">Hawaiian</option>
|
||||
<option value="ln">Lingala</option>
|
||||
<option value="ha">Hausa</option>
|
||||
<option value="az">Azerbaijani</option>
|
||||
<option value="ba">Bashkir</option>
|
||||
<option value="eu">Basque</option>
|
||||
<option value="be">Belarusian</option>
|
||||
<option value="bn">Bengali</option>
|
||||
<option value="bs">Bosnian</option>
|
||||
<option value="br">Breton</option>
|
||||
<option value="bg">Bulgarian</option>
|
||||
<option value="ca">Catalan</option>
|
||||
<option value="zh">Chinese</option>
|
||||
<option value="hr">Croatian</option>
|
||||
<option value="cs">Czech</option>
|
||||
<option value="da">Danish</option>
|
||||
<option value="nl">Dutch</option>
|
||||
<option value="en">English</option>
|
||||
<option value="et">Estonian</option>
|
||||
<option value="fo">Faroese</option>
|
||||
<option value="fi">Finnish</option>
|
||||
<option value="fr">French</option>
|
||||
<option value="gl">Galician</option>
|
||||
<option value="ka">Georgian</option>
|
||||
<option value="de">German</option>
|
||||
<option value="el">Greek</option>
|
||||
<option value="gu">Gujarati</option>
|
||||
<option value="ht">Haitian Creole</option>
|
||||
<option value="ha">Hausa</option>
|
||||
<option value="haw">Hawaiian</option>
|
||||
<option value="he">Hebrew</option>
|
||||
<option value="hi">Hindi</option>
|
||||
<option value="hu">Hungarian</option>
|
||||
<option value="is">Icelandic</option>
|
||||
<option value="id">Indonesian</option>
|
||||
<option value="it">Italian</option>
|
||||
<option value="ja">Japanese</option>
|
||||
<option value="jw">Javanese</option>
|
||||
<option value="kn">Kannada</option>
|
||||
<option value="kk">Kazakh</option>
|
||||
<option value="km">Khmer</option>
|
||||
<option value="ko">Korean</option>
|
||||
<option value="lo">Lao</option>
|
||||
<option value="la">Latin</option>
|
||||
<option value="lv">Latvian</option>
|
||||
<option value="ln">Lingala</option>
|
||||
<option value="lt">Lithuanian</option>
|
||||
<option value="lb">Luxembourgish</option>
|
||||
<option value="mk">Macedonian</option>
|
||||
<option value="mg">Malagasy</option>
|
||||
<option value="ms">Malay</option>
|
||||
<option value="ml">Malayalam</option>
|
||||
<option value="mt">Maltese</option>
|
||||
<option value="mi">Maori</option>
|
||||
<option value="mr">Marathi</option>
|
||||
<option value="mn">Mongolian</option>
|
||||
<option value="my">Myanmar</option>
|
||||
<option value="ne">Nepali</option>
|
||||
<option value="no">Norwegian</option>
|
||||
<option value="nn">Nynorsk</option>
|
||||
<option value="oc">Occitan</option>
|
||||
<option value="ps">Pashto</option>
|
||||
<option value="fa">Persian</option>
|
||||
<option value="pl">Polish</option>
|
||||
<option value="pt">Portuguese</option>
|
||||
<option value="pa">Punjabi</option>
|
||||
<option value="ro">Romanian</option>
|
||||
<option value="ru">Russian</option>
|
||||
<option value="sa">Sanskrit</option>
|
||||
<option value="sr">Serbian</option>
|
||||
<option value="sn">Shona</option>
|
||||
<option value="sd">Sindhi</option>
|
||||
<option value="si">Sinhala</option>
|
||||
<option value="sk">Slovak</option>
|
||||
<option value="sl">Slovenian</option>
|
||||
<option value="so">Somali</option>
|
||||
<option value="es">Spanish</option>
|
||||
<option value="su">Sundanese</option>
|
||||
<option value="sw">Swahili</option>
|
||||
<option value="sv">Swedish</option>
|
||||
<option value="tl">Tagalog</option>
|
||||
<option value="tg">Tajik</option>
|
||||
<option value="ta">Tamil</option>
|
||||
<option value="tt">Tatar</option>
|
||||
<option value="te">Telugu</option>
|
||||
<option value="th">Thai</option>
|
||||
<option value="bo">Tibetan</option>
|
||||
<option value="tr">Turkish</option>
|
||||
<option value="tk">Turkmen</option>
|
||||
<option value="uk">Ukrainian</option>
|
||||
<option value="ur">Urdu</option>
|
||||
<option value="uz">Uzbek</option>
|
||||
<option value="vi">Vietnamese</option>
|
||||
<option value="cy">Welsh</option>
|
||||
<option value="yi">Yiddish</option>
|
||||
<option value="yo">Yoruba</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="dropdown-container">
|
||||
@@ -133,5 +133,21 @@
|
||||
<option value="translate">Translate</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="dropdown-container">
|
||||
<label for="modelSizeDropdown">Select Model Size:</label>
|
||||
<select id="modelSizeDropdown">
|
||||
<option value="">Select model</option>
|
||||
<option value="tiny">Tiny </option>
|
||||
<option value="tiny.en">Tiny (English-only)</option>
|
||||
<option value="base">Base</option>
|
||||
<option value="base.en">Base (English-only)</option>
|
||||
<option value="small" selected>Small</option>
|
||||
<option value="small.en">Small (English-only)</option>
|
||||
<option value="medium">Medium</option>
|
||||
<option value="medium.en">Medium (English-only)</option>
|
||||
<option value="large-v2">Large-v2</option>
|
||||
<option value="large-v3">Large-v3</option>
|
||||
</select>
|
||||
</div>
|
||||
</body>
|
||||
</html>
|
||||
</html>
|
||||
|
||||
@@ -3,11 +3,14 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
const stopButton = document.getElementById("stopCapture");
|
||||
|
||||
const useServerCheckbox = document.getElementById("useServerCheckbox");
|
||||
const useMultilingualCheckbox = document.getElementById('useMultilingualCheckbox');
|
||||
const useVadCheckbox = document.getElementById("useVadCheckbox");
|
||||
const languageDropdown = document.getElementById('languageDropdown');
|
||||
const taskDropdown = document.getElementById('taskDropdown');
|
||||
const modelSizeDropdown = document.getElementById('modelSizeDropdown');
|
||||
let selectedLanguage = null;
|
||||
let selectedTask = taskDropdown.value;
|
||||
let selectedModelSize = modelSizeDropdown.value;
|
||||
|
||||
|
||||
browser.storage.local.get("capturingState")
|
||||
.then(function(result) {
|
||||
@@ -32,11 +35,9 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
}
|
||||
});
|
||||
|
||||
browser.storage.local.get("useMultilingualModelState", ({ useMultilingualModelState }) => {
|
||||
if (useMultilingualModelState !== undefined) {
|
||||
useMultilingualCheckbox.checked = useMultilingualModelState;
|
||||
languageDropdown.disabled = !useMultilingualModelState;
|
||||
taskDropdown.disabled = !useMultilingualModelState;
|
||||
browser.storage.local.get("useVadState", ({ useVadState }) => {
|
||||
if (useVadState !== undefined) {
|
||||
useVadCheckbox.checked = useVadState;
|
||||
}
|
||||
});
|
||||
|
||||
@@ -54,6 +55,13 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
}
|
||||
});
|
||||
|
||||
browser.storage.local.get("selectedModelSize", ({ selectedModelSize: storedModelSize }) => {
|
||||
if (storedModelSize !== undefined) {
|
||||
modelSizeDropdown.value = storedModelSize;
|
||||
selectedModelSize = storedModelSize;
|
||||
}
|
||||
});
|
||||
|
||||
startButton.addEventListener("click", function() {
|
||||
let host = "localhost";
|
||||
let port = "9090";
|
||||
@@ -73,9 +81,10 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
data: {
|
||||
host: host,
|
||||
port: port,
|
||||
useMultilingual: useMultilingualCheckbox.checked,
|
||||
language: selectedLanguage,
|
||||
task: selectedTask
|
||||
task: selectedTask,
|
||||
modelSize: selectedModelSize,
|
||||
useVad: useVadCheckbox.checked,
|
||||
}
|
||||
});
|
||||
toggleCaptureButtons(true);
|
||||
@@ -114,8 +123,10 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
startButton.disabled = isCapturing;
|
||||
stopButton.disabled = !isCapturing;
|
||||
useServerCheckbox.disabled = isCapturing;
|
||||
useMultilingualCheckbox.disabled = isCapturing;
|
||||
|
||||
useVadCheckbox.disabled = isCapturing;
|
||||
modelSizeDropdown.disabled = isCapturing;
|
||||
languageDropdown.disabled = isCapturing;
|
||||
taskDropdown.disabled = isCapturing;
|
||||
startButton.classList.toggle("disabled", isCapturing);
|
||||
stopButton.classList.toggle("disabled", !isCapturing);
|
||||
}
|
||||
@@ -126,16 +137,9 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
browser.storage.local.set({ useServerState });
|
||||
});
|
||||
|
||||
useMultilingualCheckbox.addEventListener('change', function() {
|
||||
const useMultilingualModelState = useMultilingualCheckbox.checked;
|
||||
if (useMultilingualModelState) {
|
||||
languageDropdown.disabled = false;
|
||||
taskDropdown.disabled = false;
|
||||
} else {
|
||||
languageDropdown.disabled = true;
|
||||
taskDropdown.disabled = true;
|
||||
}
|
||||
browser.storage.local.set({ useMultilingualModelState });
|
||||
useVadCheckbox.addEventListener("change", () => {
|
||||
const useVadState = useVadCheckbox.checked;
|
||||
browser.storage.local.set({ useVadState });
|
||||
});
|
||||
|
||||
languageDropdown.addEventListener('change', function() {
|
||||
@@ -152,6 +156,11 @@ document.addEventListener("DOMContentLoaded", function() {
|
||||
browser.storage.local.set({ selectedTask });
|
||||
});
|
||||
|
||||
modelSizeDropdown.addEventListener('change', function() {
|
||||
selectedModelSize = modelSizeDropdown.value;
|
||||
browser.storage.local.set({ selectedModelSize });
|
||||
});
|
||||
|
||||
browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
|
||||
if (request.action === "updateSelectedLanguage") {
|
||||
const detectedLanguage = request.data;
|
||||
|
||||
@@ -108,4 +108,4 @@ label {
|
||||
|
||||
.dropdown-container {
|
||||
padding: 10px;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,14 +1,30 @@
|
||||
# whisper-live
|
||||
A nearly-live implementation of OpenAI's Whisper.
|
||||
# WhisperLive
|
||||
|
||||
This project is a real-time transcription application that uses the OpenAI Whisper model to convert speech input into text output. It can be used to transcribe both live audio input from microphone and pre-recorded audio files.
|
||||
<h2 align="center">
|
||||
<a href="https://www.youtube.com/watch?v=0PHWCApIcCI"><img
|
||||
src="https://img.youtube.com/vi/0PHWCApIcCI/0.jpg" style="background-color:rgba(0,0,0,0);" height=300 alt="WhisperLive"></a>
|
||||
<br><br>A nearly-live implementation of OpenAI's Whisper.
|
||||
<br><br>
|
||||
</h2>
|
||||
|
||||
Unlike traditional speech recognition systems that rely on continuous audio streaming, we use [voice activity detection (VAD)](https://github.com/snakers4/silero-vad) to detect the presence of speech and only send the audio data to whisper when speech is detected. This helps to reduce the amount of data sent to the whisper model and improves the accuracy of the transcription output.
|
||||
This project is a real-time transcription application that uses the OpenAI Whisper model
|
||||
to convert speech input into text output. It can be used to transcribe both live audio
|
||||
input from microphone and pre-recorded audio files.
|
||||
|
||||
- [Installation](#installation)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Running the Server](#running-the-server)
|
||||
- [Running the Client](#running-the-client)
|
||||
- [Browser Extensions](#browser-extensions)
|
||||
- [Whisper Live Server in Docker](#whisper-live-server-in-docker)
|
||||
- [Future Work](#future-work)
|
||||
- [Contact](#contact)
|
||||
- [Citations](#citations)
|
||||
|
||||
## Installation
|
||||
- Install PyAudio and ffmpeg
|
||||
- Install PyAudio
|
||||
```bash
|
||||
bash setup.sh
|
||||
bash scripts/setup.sh
|
||||
```
|
||||
|
||||
- Install whisper-live from pip
|
||||
@@ -16,64 +32,155 @@ Unlike traditional speech recognition systems that rely on continuous audio stre
|
||||
pip install whisper-live
|
||||
```
|
||||
|
||||
### Setting up NVIDIA/TensorRT-LLM for TensorRT backend
|
||||
- Please follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) for setup of [NVIDIA/TensorRT-LLM](https://github.com/NVIDIA/TensorRT-LLM) and for building Whisper-TensorRT engine.
|
||||
|
||||
## Getting Started
|
||||
- Run the server
|
||||
```python
|
||||
from whisper_live.server import TranscriptionServer
|
||||
server = TranscriptionServer()
|
||||
server.run("0.0.0.0", 9090)
|
||||
The server supports 3 backends `faster_whisper`, `tensorrt` and `openvino`. If running `tensorrt` backend follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md)
|
||||
|
||||
### Running the Server
|
||||
- [Faster Whisper](https://github.com/SYSTRAN/faster-whisper) backend
|
||||
```bash
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend faster_whisper
|
||||
|
||||
# running with custom model
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend faster_whisper \
|
||||
-fw "/path/to/custom/faster/whisper/model"
|
||||
```
|
||||
|
||||
- On the client side
|
||||
- To transcribe an audio file:
|
||||
```python
|
||||
from whisper_live.client import TranscriptionClient
|
||||
client = TranscriptionClient("localhost", 9090, is_multilingual=True, lang="hi", translate=True)
|
||||
client(audio_file_path)
|
||||
```
|
||||
This command transcribes the specified audio file (audio.wav) using the Whisper model. It connects to the server running on localhost at port 9090. It also enables the multilingual feature, allowing transcription in multiple languages. The language option specifies the target language for transcription, in this case, Hindi ("hi"). 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.
|
||||
- TensorRT backend. Currently, we recommend to only use the docker setup for TensorRT. Follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) which works as expected. Make sure to build your TensorRT Engines before running the server with TensorRT backend.
|
||||
```bash
|
||||
# Run English only model
|
||||
python3 run_server.py -p 9090 \
|
||||
-b tensorrt \
|
||||
-trt /home/TensorRT-LLM/examples/whisper/whisper_small_en
|
||||
|
||||
- To transcribe from microphone:
|
||||
```python
|
||||
from whisper_live.client import TranscriptionClient
|
||||
client = TranscriptionClient(host, port, is_multilingual=True, lang="hi", translate=True)
|
||||
client()
|
||||
```
|
||||
This command captures audio from the microphone and sends it to the server for transcription. It uses the same options as the previous command, enabling the multilingual feature and specifying the target language and task.
|
||||
|
||||
|
||||
## Transcribe audio from browser
|
||||
- Run the server
|
||||
```python
|
||||
from whisper_live.server import TranscriptionServer
|
||||
server = TranscriptionServer()
|
||||
server.run("0.0.0.0", 9090)
|
||||
# Run Multilingual model
|
||||
python3 run_server.py -p 9090 \
|
||||
-b tensorrt \
|
||||
-trt /home/TensorRT-LLM/examples/whisper/whisper_small \
|
||||
-m
|
||||
```
|
||||
This would start the websocket server on port ```9090```.
|
||||
|
||||
### Chrome Extension
|
||||
- Refer to [Audio-Transcription-Chrome](https://github.com/collabora/whisper-live/tree/main/Audio-Transcription-Chrome#readme) to use Chrome extension.
|
||||
- WhisperLive now supports the [OpenVINO](https://github.com/openvinotoolkit/openvino) backend for efficient inference on Intel CPUs, iGPU and dGPUs. Currently, we tested the models uploaded to [huggingface by OpenVINO](https://huggingface.co/OpenVINO?search_models=whisper).
|
||||
- > **Docker Recommended:** Running WhisperLive with OpenVINO inside Docker automatically enables GPU support (iGPU/dGPU) without requiring additional host setup.
|
||||
- > **Native (non-Docker) Use:** If you prefer running outside Docker, ensure the Intel drivers and OpenVINO runtime are installed and properly configured on your system. Refer to the documentation for [installing OpenVINO](https://docs.openvino.ai/2025/get-started/install-openvino.html?PACKAGE=OPENVINO_BASE&VERSION=v_2025_0_0&OP_SYSTEM=LINUX&DISTRIBUTION=PIP#).
|
||||
|
||||
### Firefox Extension
|
||||
- Refer to [Audio-Transcription-Firefox](https://github.com/collabora/whisper-live/tree/main/Audio-Transcription-Firefox#readme) to use Mozilla Firefox extension.
|
||||
```
|
||||
python3 run_server.py -p 9090 -b openvino
|
||||
```
|
||||
|
||||
|
||||
#### 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 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.
|
||||
- `mute_audio_playback`: Whether to mute audio playback when transcribing an audio file. Defaults to False.
|
||||
|
||||
```python
|
||||
from whisper_live.client import TranscriptionClient
|
||||
client = TranscriptionClient(
|
||||
"localhost",
|
||||
9090,
|
||||
lang="en",
|
||||
translate=False,
|
||||
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,
|
||||
mute_audio_playback=False, # Only used for file input, False by Default
|
||||
)
|
||||
```
|
||||
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.
|
||||
|
||||
- Transcribe an audio file:
|
||||
```python
|
||||
client("tests/jfk.wav")
|
||||
```
|
||||
|
||||
- To transcribe from microphone:
|
||||
```python
|
||||
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")
|
||||
```
|
||||
|
||||
## Browser Extensions
|
||||
- Run the server with your desired backend as shown [here](https://github.com/collabora/WhisperLive?tab=readme-ov-file#running-the-server).
|
||||
- Transcribe audio directly from your browser using our Chrome or Firefox extensions. Refer to [Audio-Transcription-Chrome](https://github.com/collabora/whisper-live/tree/main/Audio-Transcription-Chrome#readme) and https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md
|
||||
|
||||
## Whisper Live Server in Docker
|
||||
- GPU
|
||||
```bash
|
||||
docker build . -t whisper-live -f docker/Dockerfile.gpu
|
||||
docker run -it --gpus all -p 9090:9090 whisper-live:latest
|
||||
```
|
||||
- Faster-Whisper
|
||||
```bash
|
||||
docker run -it --gpus all -p 9090:9090 ghcr.io/collabora/whisperlive-gpu:latest
|
||||
```
|
||||
|
||||
- TensorRT. Refer to [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) for setup and more tensorrt backend configurations.
|
||||
```bash
|
||||
docker build . -f docker/Dockerfile.tensorrt -t whisperlive-tensorrt
|
||||
docker run -p 9090:9090 --runtime=nvidia --gpus all --entrypoint /bin/bash -it 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"
|
||||
```
|
||||
|
||||
- OpenVINO
|
||||
```
|
||||
docker run -it --device=/dev/dri -p 9090:9090 ghcr.io/collabora/whisperlive-openvino
|
||||
```
|
||||
|
||||
- CPU
|
||||
```bash
|
||||
docker build . -t whisper-live -f docker/Dockerfile.cpu
|
||||
docker run -it -p 9090:9090 whisper-live:latest
|
||||
```
|
||||
**Note**: By default we use "small" model size. To build docker image for a different model size, change the size in server.py and then build the docker image.
|
||||
- Faster-whisper
|
||||
```bash
|
||||
docker run -it -p 9090:9090 ghcr.io/collabora/whisperlive-cpu:latest
|
||||
```
|
||||
|
||||
## Future Work
|
||||
- [ ] Add translation to other languages on top of transcription.
|
||||
- [ ] TensorRT backend for Whisper.
|
||||
|
||||
## Contact
|
||||
|
||||
@@ -98,6 +205,5 @@ We are available to help you with both Open Source and proprietary AI projects.
|
||||
publisher = {GitHub},
|
||||
journal = {GitHub repository},
|
||||
howpublished = {\url{https://github.com/snakers4/silero-vad}},
|
||||
commit = {insert_some_commit_here},
|
||||
email = {hello@silero.ai}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
# WhisperLive-TensorRT
|
||||
We have only tested the TensorRT backend in docker so, we recommend docker for a smooth TensorRT backend setup.
|
||||
**Note**: We use `tensorrt_llm==0.18.2`
|
||||
|
||||
## 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)
|
||||
|
||||
- Run WhisperLive TensorRT in docker
|
||||
```bash
|
||||
docker build . -f docker/Dockerfile.tensorrt -t whisperlive-tensorrt
|
||||
docker run -p 9090:9090 --runtime=nvidia --gpus all --entrypoint /bin/bash -it whisperlive-tensorrt
|
||||
```
|
||||
|
||||
## Whisper TensorRT Engine
|
||||
- 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 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 build_whisper_tensorrt.sh /app/TensorRT-LLM-examples small
|
||||
```
|
||||
|
||||
## Run WhisperLive Server with TensorRT Backend
|
||||
```bash
|
||||
# Run English only model
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend tensorrt \
|
||||
--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 "/app/TensorRT-LLM-examples/whisper/whisper_small_float16" \
|
||||
--trt_multilingual
|
||||
```
|
||||
|
||||
By default trt_backend uses cpp_session, to use python session pass `--trt_py_session` to run_server.py
|
||||
```bash
|
||||
python3 run_server.py --port 9090 \
|
||||
--backend tensorrt \
|
||||
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_float16" \
|
||||
--trt_py_session
|
||||
```
|
||||
Binary file not shown.
+12
-32
@@ -1,45 +1,25 @@
|
||||
FROM ubuntu:focal
|
||||
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 \
|
||||
curl \
|
||||
ca-certificates \
|
||||
sudo \
|
||||
git \
|
||||
bzip2 \
|
||||
libx11-6 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
# update pip to support for whl.metadata -> less downloading
|
||||
RUN pip install --no-cache-dir -U "pip>=24"
|
||||
|
||||
RUN apt update
|
||||
|
||||
# install python
|
||||
RUN apt install software-properties-common -y && \
|
||||
add-apt-repository ppa:deadsnakes/ppa && \
|
||||
apt update
|
||||
|
||||
RUN apt install python3-dev -y && \
|
||||
apt install python-is-python3
|
||||
|
||||
|
||||
# install pip
|
||||
RUN apt install python3-pip -y
|
||||
|
||||
# Create a working directory.
|
||||
# create a working directory
|
||||
RUN mkdir /app
|
||||
WORKDIR /app
|
||||
|
||||
COPY setup.sh /app
|
||||
COPY requirements/ /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 bash setup.sh
|
||||
RUN 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
|
||||
|
||||
CMD ["python", "run_server.py"]
|
||||
|
||||
+12
-33
@@ -1,47 +1,26 @@
|
||||
FROM nvidia/cuda:11.2.2-cudnn8-runtime-ubuntu20.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 \
|
||||
curl \
|
||||
ca-certificates \
|
||||
sudo \
|
||||
git \
|
||||
bzip2 \
|
||||
libx11-6 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
# update pip to support for whl.metadata -> less downloading
|
||||
RUN pip install --no-cache-dir -U "pip>=24"
|
||||
|
||||
RUN apt update
|
||||
|
||||
# install python
|
||||
RUN apt install software-properties-common -y && \
|
||||
add-apt-repository ppa:deadsnakes/ppa && \
|
||||
apt update
|
||||
|
||||
RUN apt install python3-dev -y && \
|
||||
apt install python-is-python3
|
||||
|
||||
|
||||
# install pip
|
||||
RUN apt install python3-pip -y
|
||||
|
||||
# Create a working directory.
|
||||
# create a working directory
|
||||
RUN mkdir /app
|
||||
WORKDIR /app
|
||||
|
||||
COPY setup.sh /app
|
||||
COPY requirements/ /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 --fix-missing
|
||||
RUN bash setup.sh
|
||||
RUN pip install -r 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,19 @@
|
||||
FROM openvino/ubuntu22_runtime:latest
|
||||
|
||||
ARG DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
USER root
|
||||
|
||||
RUN apt update && apt install -y portaudio19-dev python-is-python3 && apt-get clean && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN pip install --no-cache-dir -U "pip>=24"
|
||||
|
||||
RUN mkdir /app
|
||||
WORKDIR /app
|
||||
|
||||
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
|
||||
CMD ["python", "run_server.py", "--backend", "openvino"]
|
||||
@@ -0,0 +1,30 @@
|
||||
FROM nvidia/cuda:12.8.1-base-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 \
|
||||
&& apt install python-is-python3 \
|
||||
&& pip install --upgrade pip setuptools \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
FROM base AS devel
|
||||
RUN pip install --no-cache-dir -U tensorrt_llm==0.18.2 --extra-index-url https://pypi.nvidia.com
|
||||
WORKDIR /app
|
||||
RUN git clone -b v0.18.2 https://github.com/NVIDIA/TensorRT-LLM.git \
|
||||
&& 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,4 +1,4 @@
|
||||
PyAudio
|
||||
ffmpeg-python
|
||||
av
|
||||
scipy
|
||||
websocket-client
|
||||
+20
-6
@@ -1,7 +1,21 @@
|
||||
PyAudio
|
||||
faster-whisper==0.6.0
|
||||
--extra-index-url https://download.pytorch.org/whl/cu111
|
||||
torch==1.10.1
|
||||
torchaudio==0.10.1
|
||||
faster-whisper==1.1.0
|
||||
websockets
|
||||
onnxruntime==1.16.0
|
||||
onnxruntime==1.17.0
|
||||
numba
|
||||
kaldialign
|
||||
soundfile
|
||||
scipy
|
||||
av
|
||||
jiwer
|
||||
evaluate
|
||||
numpy<2
|
||||
openai-whisper==20240930
|
||||
tokenizers==0.20.3
|
||||
|
||||
# openvino
|
||||
librosa
|
||||
openvino
|
||||
openvino-genai
|
||||
openvino-tokenizers
|
||||
optimum
|
||||
optimum-intel
|
||||
+51
-2
@@ -1,5 +1,54 @@
|
||||
from whisper_live.server import TranscriptionServer
|
||||
import argparse
|
||||
import os
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--port', '-p',
|
||||
type=int,
|
||||
default=9090,
|
||||
help="Websocket port to run the server on.")
|
||||
parser.add_argument('--backend', '-b',
|
||||
type=str,
|
||||
default='faster_whisper',
|
||||
help='Backends from ["tensorrt", "faster_whisper", "openvino"]')
|
||||
parser.add_argument('--faster_whisper_custom_model_path', '-fw',
|
||||
type=str, default=None,
|
||||
help="Custom Faster Whisper Model")
|
||||
parser.add_argument('--trt_model_path', '-trt',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Whisper TensorRT model path')
|
||||
parser.add_argument('--trt_multilingual', '-m',
|
||||
action="store_true",
|
||||
help='Boolean only for TensorRT model. True if multilingual.')
|
||||
parser.add_argument('--trt_py_session',
|
||||
action="store_true",
|
||||
help='Boolean only for TensorRT model. Use python session or cpp session, By default uses Cpp.')
|
||||
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")
|
||||
server.run(
|
||||
"0.0.0.0",
|
||||
port=args.port,
|
||||
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_py_session=args.trt_py_session,
|
||||
single_model=not args.no_single_model,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
#!/bin/bash
|
||||
|
||||
download_and_build_model() {
|
||||
local model_name="$1"
|
||||
local model_url=""
|
||||
|
||||
case "$model_name" in
|
||||
"tiny.en")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/d3dd57d32accea0b295c96e26691aa14d8822fac7d9d27d5dc00b4ca2826dd03/tiny.en.pt"
|
||||
;;
|
||||
"tiny")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/65147644a518d12f04e32d6f3b26facc3f8dd46e5390956a9424a650c0ce22b9/tiny.pt"
|
||||
;;
|
||||
"base.en")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/25a8566e1d0c1e2231d1c762132cd20e0f96a85d16145c3a00adf5d1ac670ead/base.en.pt"
|
||||
;;
|
||||
"base")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/ed3a0b6b1c0edf879ad9b11b1af5a0e6ab5db9205f891f668f8b0e6c6326e34e/base.pt"
|
||||
;;
|
||||
"small.en")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/f953ad0fd29cacd07d5a9eda5624af0f6bcf2258be67c92b79389873d91e0872/small.en.pt"
|
||||
;;
|
||||
"small")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/9ecf779972d90ba49c06d968637d720dd632c55bbf19d441fb42bf17a411e794/small.pt"
|
||||
;;
|
||||
"medium.en")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/d7440d1dc186f76616474e0ff0b3b6b879abc9d1a4926b7adfa41db2d497ab4f/medium.en.pt"
|
||||
;;
|
||||
"medium")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/345ae4da62f9b3d59415adc60127b97c714f32e89e936602e85993674d08dcb1/medium.pt"
|
||||
;;
|
||||
"large-v1")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large-v1.pt"
|
||||
;;
|
||||
"large-v2")
|
||||
model_url="https://openaipublic.azureedge.net/main/whisper/models/81f7c96c852ee8fc832187b0132e569d6c3065a3252ed18e56effd0b6a73e524/large-v2.pt"
|
||||
;;
|
||||
"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=4
|
||||
|
||||
echo "Downloading $model_name..."
|
||||
# wget --directory-prefix=assets "$model_url"
|
||||
# echo "Download completed: ${model_name}.pt"
|
||||
if [ ! -f "assets/${model_name}.pt" ]; then
|
||||
wget --directory-prefix=assets "$model_url"
|
||||
echo "Download completed: ${model_name}.pt"
|
||||
else
|
||||
echo "${model_name}.pt already exists in assets directory."
|
||||
fi
|
||||
|
||||
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 "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 \
|
||||
--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 \
|
||||
--max_beam_width "$max_beam_width" \
|
||||
--max_batch_size "$max_batch_size" \
|
||||
--max_seq_len 225 \
|
||||
--max_input_len 32 \
|
||||
--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"
|
||||
}
|
||||
|
||||
if [ "$#" -lt 1 ]; then
|
||||
echo "Usage: $0 <path-to-tensorrt-examples-dir> [model-name]"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
tensorrt_examples_dir="$1"
|
||||
model_name="${2:-small.en}"
|
||||
weight_only_precision="${3:-float16}" # Default to float16 if not provided
|
||||
|
||||
cd $tensorrt_examples_dir/whisper
|
||||
pip install --no-deps -r requirements.txt
|
||||
|
||||
download_and_build_model "$model_name" "$weight_only_precision"
|
||||
@@ -0,0 +1,3 @@
|
||||
#! /bin/bash
|
||||
|
||||
apt-get install portaudio19-dev wget -y
|
||||
@@ -10,45 +10,58 @@ HERE = pathlib.Path(__file__).parent
|
||||
README = (HERE / "README.md").read_text()
|
||||
|
||||
# This call to setup() does all the work
|
||||
setup(name="whisper-live",
|
||||
version=__version__,
|
||||
description="A nearly-live implementation of OpenAI's Whisper.",
|
||||
long_description=README,
|
||||
long_description_content_type="text/markdown",
|
||||
include_package_data=True,
|
||||
url="https://github.com/collabora/WhisperLive",
|
||||
author="Collabora Ltd",
|
||||
author_email="vineet.suryan@collabora.com",
|
||||
license="MIT",
|
||||
classifiers=[
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Intended Audience :: Science/Research",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3 :: Only",
|
||||
"Programming Language :: Python :: 3.8",
|
||||
"Programming Language :: Python :: 3.9",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
],
|
||||
packages=find_packages(
|
||||
exclude=("examples",
|
||||
"Audio-Transcription-Chrome",
|
||||
"Audio-Transcription-Firefox",
|
||||
"requirements",
|
||||
"whisper-finetuning"
|
||||
)
|
||||
),
|
||||
install_requires=[
|
||||
setup(
|
||||
name="whisper_live",
|
||||
version=__version__,
|
||||
description="A nearly-live implementation of OpenAI's Whisper.",
|
||||
long_description=README,
|
||||
long_description_content_type="text/markdown",
|
||||
include_package_data=True,
|
||||
url="https://github.com/collabora/WhisperLive",
|
||||
author="Collabora Ltd",
|
||||
author_email="vineet.suryan@collabora.com",
|
||||
license="MIT",
|
||||
classifiers=[
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Intended Audience :: Science/Research",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3 :: Only",
|
||||
"Programming Language :: Python :: 3.8",
|
||||
"Programming Language :: Python :: 3.9",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
],
|
||||
packages=find_packages(
|
||||
exclude=(
|
||||
"examples",
|
||||
"Audio-Transcription-Chrome",
|
||||
"Audio-Transcription-Firefox",
|
||||
"requirements",
|
||||
"whisper-finetuning"
|
||||
)
|
||||
),
|
||||
install_requires=[
|
||||
"PyAudio",
|
||||
"faster-whisper==0.6.0",
|
||||
"faster-whisper==1.1.0",
|
||||
"torch",
|
||||
"torchaudio",
|
||||
"websockets",
|
||||
"onnxruntime",
|
||||
"ffmpeg-python",
|
||||
"onnxruntime==1.17.0",
|
||||
"scipy",
|
||||
"websocket-client",
|
||||
],
|
||||
python_requires=">=3.8"
|
||||
)
|
||||
"numba",
|
||||
"openai-whisper==20240930",
|
||||
"kaldialign",
|
||||
"soundfile",
|
||||
"tokenizers==0.20.3",
|
||||
"librosa",
|
||||
"numpy==1.26.4",
|
||||
"openvino",
|
||||
"openvino-genai",
|
||||
"openvino-tokenizers",
|
||||
"optimum",
|
||||
"optimum-intel",
|
||||
],
|
||||
python_requires=">=3.9"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
import json
|
||||
import os
|
||||
import scipy
|
||||
import websocket
|
||||
import copy
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
|
||||
from whisper_live.utils import resample
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class BaseTestCase(unittest.TestCase):
|
||||
@patch('whisper_live.client.websocket.WebSocketApp')
|
||||
@patch('whisper_live.client.pyaudio.PyAudio')
|
||||
def setUp(self, mock_pyaudio, mock_websocket):
|
||||
self.mock_pyaudio_instance = MagicMock()
|
||||
mock_pyaudio.return_value = self.mock_pyaudio_instance
|
||||
self.mock_stream = MagicMock()
|
||||
self.mock_pyaudio_instance.open.return_value = self.mock_stream
|
||||
|
||||
self.mock_ws_app = mock_websocket.return_value
|
||||
self.mock_ws_app.send = MagicMock()
|
||||
|
||||
self.client = TranscriptionClient(host='localhost', port=9090, lang="en").client
|
||||
|
||||
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()
|
||||
self.mock_pyaudio.stop()
|
||||
self.mock_websocket.stop()
|
||||
del self.client
|
||||
|
||||
class TestClientWebSocketCommunication(BaseTestCase):
|
||||
def test_websocket_communication(self):
|
||||
expected_url = 'ws://localhost:9090'
|
||||
self.mock_websocket.assert_called()
|
||||
self.assertEqual(self.mock_websocket.call_args[0][0], expected_url)
|
||||
|
||||
|
||||
class TestClientCallbacks(BaseTestCase):
|
||||
def test_on_open(self):
|
||||
expected_message = json.dumps({
|
||||
"uid": self.client.uid,
|
||||
"language": self.client.language,
|
||||
"task": self.client.task,
|
||||
"model": self.client.model,
|
||||
"use_vad": True,
|
||||
"max_clients": 4,
|
||||
"max_connection_time": 600,
|
||||
"send_last_n_segments": 10,
|
||||
"no_speech_thresh": 0.45,
|
||||
"clip_audio": False,
|
||||
"same_output_threshold": 10,
|
||||
})
|
||||
self.client.on_open(self.mock_ws_app)
|
||||
self.mock_ws_app.send.assert_called_with(expected_message)
|
||||
|
||||
def test_on_message(self):
|
||||
message = json.dumps(
|
||||
{
|
||||
"uid": self.client.uid,
|
||||
"message": "SERVER_READY",
|
||||
"backend": "faster_whisper"
|
||||
}
|
||||
)
|
||||
self.client.on_message(self.mock_ws_app, message)
|
||||
|
||||
message = json.dumps({
|
||||
"uid": self.client.uid,
|
||||
"segments": [
|
||||
{"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), 3)
|
||||
self.assertEqual(self.client.transcript[1]['text'], "Test transcript 2")
|
||||
|
||||
def test_on_close(self):
|
||||
close_status_code = 1000
|
||||
close_msg = "Normal closure"
|
||||
self.client.on_close(self.mock_ws_app, close_status_code, close_msg)
|
||||
|
||||
self.assertFalse(self.client.recording)
|
||||
self.assertFalse(self.client.server_error)
|
||||
self.assertFalse(self.client.waiting)
|
||||
|
||||
def test_on_error(self):
|
||||
error_message = "Test Error"
|
||||
self.client.on_error(self.mock_ws_app, error_message)
|
||||
|
||||
self.assertTrue(self.client.server_error)
|
||||
self.assertEqual(self.client.error_message, error_message)
|
||||
|
||||
|
||||
class TestAudioResampling(unittest.TestCase):
|
||||
def test_resample_audio(self):
|
||||
original_audio = "assets/jfk.flac"
|
||||
expected_sr = 16000
|
||||
resampled_audio = resample(original_audio, expected_sr)
|
||||
|
||||
sr, _ = scipy.io.wavfile.read(resampled_audio)
|
||||
self.assertEqual(sr, expected_sr)
|
||||
|
||||
os.remove(resampled_audio)
|
||||
|
||||
|
||||
class TestSendingAudioPacket(BaseTestCase):
|
||||
def test_send_packet(self):
|
||||
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())
|
||||
@@ -0,0 +1,148 @@
|
||||
import subprocess
|
||||
import time
|
||||
import json
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import numpy as np
|
||||
import jiwer
|
||||
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
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, {})
|
||||
self.assertDictEqual(server.client_manager.start_times, {})
|
||||
|
||||
|
||||
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
|
||||
}
|
||||
self.server.client_manager.max_connection_time = 600
|
||||
|
||||
def test_get_wait_time(self):
|
||||
expected_wait_time = (600 - (time.time() - self.server.client_manager.start_times['client2'])) / 60
|
||||
print(self.server.client_manager.get_wait_time(), expected_wait_time)
|
||||
self.assertAlmostEqual(self.server.client_manager.get_wait_time(), expected_wait_time, places=2)
|
||||
|
||||
|
||||
class TestServerConnection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = json.dumps({
|
||||
'uid': 'test_client',
|
||||
'language': 'en',
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
})
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_recv_audio_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = [json.dumps({
|
||||
'uid': 'test_client',
|
||||
'language': 'en',
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
}), np.array([1, 2, 3]).tobytes()]
|
||||
|
||||
with self.assertLogs(level="ERROR"):
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
|
||||
self.assertNotIn(mock_websocket, self.server.client_manager.clients)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
cls.server_process.terminate()
|
||||
cls.server_process.wait()
|
||||
|
||||
def setUp(self):
|
||||
self.normalizer = EnglishTextNormalizer()
|
||||
|
||||
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!"
|
||||
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_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",
|
||||
)
|
||||
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):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection_closed_exception(self, mock_websocket):
|
||||
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, BackendType("faster_whisper"))
|
||||
self.assertTrue(any("Connection closed by client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_json_decode_exception(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = "invalid json"
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
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')
|
||||
def test_unexpected_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = RuntimeError("Unexpected error")
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, BackendType("faster_whisper"))
|
||||
for message in log.output:
|
||||
print(message)
|
||||
print()
|
||||
self.assertTrue(any("Unexpected error" in message for message in log.output))
|
||||
@@ -0,0 +1,26 @@
|
||||
import unittest
|
||||
import numpy as np
|
||||
from whisper_live.transcriber.tensorrt_utils import load_audio
|
||||
from whisper_live.vad import VoiceActivityDetector
|
||||
|
||||
|
||||
class TestVoiceActivityDetection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.vad = VoiceActivityDetector()
|
||||
self.sample_rate = 16000
|
||||
|
||||
def generate_silence(self, duration_seconds):
|
||||
return np.zeros(int(self.sample_rate * duration_seconds), dtype=np.float32)
|
||||
|
||||
def load_speech_segment(self, filepath):
|
||||
return load_audio(filepath)
|
||||
|
||||
def test_vad_silence_detection(self):
|
||||
silence = self.generate_silence(3)
|
||||
is_speech_present = self.vad(silence.copy())
|
||||
self.assertFalse(is_speech_present, "VAD incorrectly identified silence as speech.")
|
||||
|
||||
def test_vad_speech_detection(self):
|
||||
audio_tensor = load_audio("assets/jfk.flac")
|
||||
is_speech_present = self.vad(audio_tensor)
|
||||
self.assertTrue(is_speech_present, "VAD failed to identify speech segment.")
|
||||
@@ -1 +1 @@
|
||||
__version__="0.0.7"
|
||||
__version__ = "0.7.1"
|
||||
|
||||
@@ -0,0 +1,361 @@
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import numpy as np
|
||||
|
||||
|
||||
class ServeClientBase(object):
|
||||
RATE = 16000
|
||||
SERVER_READY = "SERVER_READY"
|
||||
DISCONNECT = "DISCONNECT"
|
||||
|
||||
client_uid: str
|
||||
"""A unique identifier for the client."""
|
||||
websocket: object
|
||||
"""The WebSocket connection for the client."""
|
||||
send_last_n_segments: int
|
||||
"""Number of most recent segments to send to the client."""
|
||||
no_speech_thresh: float
|
||||
"""Segments with no speech probability above this threshold will be discarded."""
|
||||
clip_audio: bool
|
||||
"""Whether to clip audio with no valid segments."""
|
||||
same_output_threshold: int
|
||||
"""Number of repeated outputs before considering it as a valid segment."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client_uid,
|
||||
websocket,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
):
|
||||
self.client_uid = client_uid
|
||||
self.websocket = websocket
|
||||
self.send_last_n_segments = send_last_n_segments
|
||||
self.no_speech_thresh = no_speech_thresh
|
||||
self.clip_audio = clip_audio
|
||||
self.same_output_threshold = same_output_threshold
|
||||
|
||||
self.frames = b""
|
||||
self.timestamp_offset = 0.0
|
||||
self.frames_np = None
|
||||
self.frames_offset = 0.0
|
||||
self.text = []
|
||||
self.current_out = ""
|
||||
self.prev_out = ""
|
||||
self.exit = False
|
||||
self.same_output_count = 0
|
||||
self.transcript = []
|
||||
self.end_time_for_same_output = None
|
||||
|
||||
# threading
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def speech_to_text(self):
|
||||
"""
|
||||
Process an audio stream in an infinite loop, continuously transcribing the speech.
|
||||
|
||||
This method continuously receives audio frames, performs real-time transcription, and sends
|
||||
transcribed segments to the client via a WebSocket connection.
|
||||
|
||||
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
|
||||
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
|
||||
are sent to the client in real-time, and a history of segments is maintained to provide context.
|
||||
|
||||
Raises:
|
||||
Exception: If there is an issue with audio processing or WebSocket communication.
|
||||
|
||||
"""
|
||||
while True:
|
||||
if self.exit:
|
||||
logging.info("Exiting speech to text thread")
|
||||
break
|
||||
|
||||
if self.frames_np is None:
|
||||
continue
|
||||
|
||||
if self.clip_audio:
|
||||
self.clip_audio_if_no_valid_segment()
|
||||
|
||||
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 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)
|
||||
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: Failed to transcribe audio chunk: {e}")
|
||||
time.sleep(0.01)
|
||||
|
||||
def transcribe_audio(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def handle_transcription_output(self, result, duration):
|
||||
raise NotImplementedError
|
||||
|
||||
def format_segment(self, start, end, text, completed=False):
|
||||
"""
|
||||
Formats a transcription segment with precise start and end times alongside the transcribed text.
|
||||
|
||||
Args:
|
||||
start (float): The start time of the transcription segment in seconds.
|
||||
end (float): The end time of the transcription segment in seconds.
|
||||
text (str): The transcribed text corresponding to the segment.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary representing the formatted transcription segment, including
|
||||
'start' and 'end' times as strings with three decimal places and the 'text'
|
||||
of the transcription.
|
||||
"""
|
||||
return {
|
||||
'start': "{:.3f}".format(start),
|
||||
'end': "{:.3f}".format(end),
|
||||
'text': text,
|
||||
'completed': completed
|
||||
}
|
||||
|
||||
def add_frames(self, frame_np):
|
||||
"""
|
||||
Add audio frames to the ongoing audio stream buffer.
|
||||
|
||||
This method is responsible for maintaining the audio stream buffer, allowing the continuous addition
|
||||
of audio frames as they are received. It also ensures that the buffer does not exceed a specified size
|
||||
to prevent excessive memory usage.
|
||||
|
||||
If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds
|
||||
of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided
|
||||
audio frame. The audio stream buffer is used for real-time processing of audio data for transcription.
|
||||
|
||||
Args:
|
||||
frame_np (numpy.ndarray): The audio frame data as a NumPy array.
|
||||
|
||||
"""
|
||||
self.lock.acquire()
|
||||
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:
|
||||
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
|
||||
self.lock.release()
|
||||
|
||||
def clip_audio_if_no_valid_segment(self):
|
||||
"""
|
||||
Update the timestamp offset based on audio buffer status.
|
||||
Clip audio if the current chunk exceeds 30 seconds, this basically implies that
|
||||
no valid segment for the last 30 seconds from whisper
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
Retrieves the next chunk of audio data for processing based on the current offsets.
|
||||
|
||||
Calculates which part of the audio data should be processed next, based on
|
||||
the difference between the current timestamp offset and the frame's offset, scaled by
|
||||
the audio sample rate (RATE). It then returns this chunk of audio data along with its
|
||||
duration in seconds.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing:
|
||||
- input_bytes (np.ndarray): The next chunk of audio data to be processed.
|
||||
- duration (float): The duration of the audio chunk in seconds.
|
||||
"""
|
||||
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
|
||||
|
||||
def prepare_segments(self, last_segment=None):
|
||||
"""
|
||||
Prepares the segments of transcribed text to be sent to the client.
|
||||
|
||||
This method compiles the recent segments of transcribed text, ensuring that only the
|
||||
specified number of the most recent segments are included. It also appends the most
|
||||
recent segment of text if provided (which is considered incomplete because of the possibility
|
||||
of the last word being truncated in the audio chunk).
|
||||
|
||||
Args:
|
||||
last_segment (str, optional): The most recent segment of transcribed text to be added
|
||||
to the list of segments. Defaults to None.
|
||||
|
||||
Returns:
|
||||
list: A list of transcribed text segments to be sent to the client.
|
||||
"""
|
||||
segments = []
|
||||
if len(self.transcript) >= self.send_last_n_segments:
|
||||
segments = self.transcript[-self.send_last_n_segments:].copy()
|
||||
else:
|
||||
segments = self.transcript.copy()
|
||||
if last_segment is not None:
|
||||
segments = segments + [last_segment]
|
||||
return segments
|
||||
|
||||
def get_audio_chunk_duration(self, input_bytes):
|
||||
"""
|
||||
Calculates the duration of the provided audio chunk.
|
||||
|
||||
Args:
|
||||
input_bytes (numpy.ndarray): The audio chunk for which to calculate the duration.
|
||||
|
||||
Returns:
|
||||
float: The duration of the audio chunk in seconds.
|
||||
"""
|
||||
return input_bytes.shape[0] / self.RATE
|
||||
|
||||
def send_transcription_to_client(self, segments):
|
||||
"""
|
||||
Sends the specified transcription segments to the client over the websocket connection.
|
||||
|
||||
This method formats the transcription segments into a JSON object and attempts to send
|
||||
this object to the client. If an error occurs during the send operation, it logs the error.
|
||||
|
||||
Returns:
|
||||
segments (list): A list of transcription segments to be sent to the client.
|
||||
"""
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments,
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: Sending data to client: {e}")
|
||||
|
||||
def disconnect(self):
|
||||
"""
|
||||
Notify the client of disconnection and send a disconnect message.
|
||||
|
||||
This method sends a disconnect message to the client via the WebSocket connection to notify them
|
||||
that the transcription service is disconnecting gracefully.
|
||||
|
||||
"""
|
||||
self.websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"message": self.DISCONNECT
|
||||
}))
|
||||
|
||||
def cleanup(self):
|
||||
"""
|
||||
Perform cleanup tasks before exiting the transcription service.
|
||||
|
||||
This method performs necessary cleanup tasks, including stopping the transcription thread, marking
|
||||
the exit flag to indicate the transcription thread should exit gracefully, and destroying resources
|
||||
associated with the transcription process.
|
||||
|
||||
"""
|
||||
logging.info("Cleaning up.")
|
||||
self.exit = True
|
||||
|
||||
def get_segment_no_speech_prob(self, segment):
|
||||
return getattr(segment, "no_speech_prob", 0)
|
||||
|
||||
def get_segment_start(self, segment):
|
||||
return getattr(segment, "start", getattr(segment, "start_ts", 0))
|
||||
|
||||
def get_segment_end(self, segment):
|
||||
return getattr(segment, "end", getattr(segment, "end_ts", 0))
|
||||
|
||||
def update_segments(self, segments, duration):
|
||||
"""
|
||||
Processes the segments from Whisper and updates the transcript.
|
||||
Uses helper methods to account for differences between backends.
|
||||
|
||||
Args:
|
||||
segments (list): List of segments returned by the transcriber.
|
||||
duration (float): Duration of the current audio chunk.
|
||||
|
||||
Returns:
|
||||
dict or None: The last processed segment (if any).
|
||||
"""
|
||||
offset = None
|
||||
self.current_out = ''
|
||||
last_segment = None
|
||||
|
||||
# Process complete segments only if there are more than one
|
||||
# and if the last segment's no_speech_prob is below the threshold.
|
||||
if len(segments) > 1 and self.get_segment_no_speech_prob(segments[-1]) <= self.no_speech_thresh:
|
||||
for s in segments[:-1]:
|
||||
text_ = s.text
|
||||
self.text.append(text_)
|
||||
with self.lock:
|
||||
start = self.timestamp_offset + self.get_segment_start(s)
|
||||
end = self.timestamp_offset + min(duration, self.get_segment_end(s))
|
||||
if start >= end:
|
||||
continue
|
||||
if self.get_segment_no_speech_prob(s) > self.no_speech_thresh:
|
||||
continue
|
||||
self.transcript.append(self.format_segment(start, end, text_, completed=True))
|
||||
offset = min(duration, self.get_segment_end(s))
|
||||
|
||||
# Process the last segment if its no_speech_prob is acceptable.
|
||||
if self.get_segment_no_speech_prob(segments[-1]) <= self.no_speech_thresh:
|
||||
self.current_out += segments[-1].text
|
||||
with self.lock:
|
||||
last_segment = self.format_segment(
|
||||
self.timestamp_offset + self.get_segment_start(segments[-1]),
|
||||
self.timestamp_offset + min(duration, self.get_segment_end(segments[-1])),
|
||||
self.current_out,
|
||||
completed=False
|
||||
)
|
||||
|
||||
# Handle repeated output logic.
|
||||
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 = self.get_segment_end(segments[-1])
|
||||
time.sleep(0.1) # wait briefly for any new voice activity
|
||||
else:
|
||||
self.same_output_count = 0
|
||||
self.end_time_for_same_output = None
|
||||
|
||||
# If the same incomplete segment is repeated too many times,
|
||||
# append it to the transcript and update the offset.
|
||||
if self.same_output_count > self.same_output_threshold:
|
||||
if not self.text or self.text[-1].strip().lower() != self.current_out.strip().lower():
|
||||
self.text.append(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 = 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
|
||||
|
||||
if offset is not None:
|
||||
with self.lock:
|
||||
self.timestamp_offset += offset
|
||||
|
||||
return last_segment
|
||||
@@ -0,0 +1,216 @@
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import torch
|
||||
|
||||
from whisper_live.transcriber.transcriber_faster_whisper import WhisperModel
|
||||
from whisper_live.backend.base import 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,
|
||||
single_model=False,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
The transcription thread is started upon initialization. A "SERVER_READY" message is sent
|
||||
to the client to indicate that the server is ready.
|
||||
|
||||
Args:
|
||||
websocket (WebSocket): The WebSocket connection for the client.
|
||||
task (str, optional): The task type, e.g., "transcribe". Defaults to "transcribe".
|
||||
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
||||
language (str, optional): The language for transcription. Defaults to None.
|
||||
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.
|
||||
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||
|
||||
"""
|
||||
super().__init__(
|
||||
client_uid,
|
||||
websocket,
|
||||
send_last_n_segments,
|
||||
no_speech_thresh,
|
||||
clip_audio,
|
||||
same_output_threshold,
|
||||
)
|
||||
self.model_sizes = [
|
||||
"tiny", "tiny.en", "base", "base.en", "small", "small.en",
|
||||
"medium", "medium.en", "large-v2", "large-v3", "distil-small.en",
|
||||
"distil-medium.en", "distil-large-v2", "distil-large-v3",
|
||||
"large-v3-turbo", "turbo"
|
||||
]
|
||||
|
||||
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 {"onset": 0.5}
|
||||
|
||||
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.use_vad = use_vad
|
||||
|
||||
# threading
|
||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||
self.trans_thread.start()
|
||||
self.websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
"uid": self.client_uid,
|
||||
"message": self.SERVER_READY,
|
||||
"backend": "faster_whisper"
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
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.
|
||||
|
||||
Args:
|
||||
model_size (str): The name of the model size to check.
|
||||
|
||||
Returns:
|
||||
str: The model size if valid, None otherwise.
|
||||
"""
|
||||
if model_size not in self.model_sizes:
|
||||
self.websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
"uid": self.client_uid,
|
||||
"status": "ERROR",
|
||||
"message": f"Invalid model size {model_size}. Available choices: {self.model_sizes}"
|
||||
}
|
||||
)
|
||||
)
|
||||
return None
|
||||
return model_size
|
||||
|
||||
def set_language(self, info):
|
||||
"""
|
||||
Updates the language attribute based on the detected language information.
|
||||
|
||||
Args:
|
||||
info (object): An object containing the detected language and its probability. This object
|
||||
must have at least two attributes: `language`, a string indicating the detected
|
||||
language, and `language_probability`, a float representing the confidence level
|
||||
of the language detection.
|
||||
"""
|
||||
if info.language_probability > 0.5:
|
||||
self.language = info.language
|
||||
logging.info(f"Detected language {self.language} with probability {info.language_probability}")
|
||||
self.websocket.send(json.dumps(
|
||||
{"uid": self.client_uid, "language": self.language, "language_prob": info.language_probability}))
|
||||
|
||||
def transcribe_audio(self, input_sample):
|
||||
"""
|
||||
Transcribes the provided audio sample using the configured transcriber instance.
|
||||
|
||||
If the language has not been set, it updates the session's language based on the transcription
|
||||
information.
|
||||
|
||||
Args:
|
||||
input_sample (np.array): The audio chunk to be transcribed. This should be a NumPy
|
||||
array representing the audio data.
|
||||
|
||||
Returns:
|
||||
The transcription result from the transcriber. The exact format of this result
|
||||
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,
|
||||
language=self.language,
|
||||
task=self.task,
|
||||
vad_filter=self.use_vad,
|
||||
vad_parameters=self.vad_parameters if self.use_vad else 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
|
||||
|
||||
def handle_transcription_output(self, result, duration):
|
||||
"""
|
||||
Handle the transcription output, updating the transcript and sending data to the client.
|
||||
|
||||
Args:
|
||||
result (str): The result from whisper inference i.e. the list of segments.
|
||||
duration (float): Duration of the transcribed audio chunk.
|
||||
"""
|
||||
segments = []
|
||||
if len(result):
|
||||
self.t_start = None
|
||||
last_segment = self.update_segments(result, duration)
|
||||
segments = self.prepare_segments(last_segment)
|
||||
|
||||
if len(segments):
|
||||
self.send_transcription_to_client(segments)
|
||||
@@ -0,0 +1,148 @@
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
from openvino import Core
|
||||
from whisper_live.backend.base import ServeClientBase
|
||||
from whisper_live.transcriber.transcriber_openvino import WhisperOpenVINO
|
||||
|
||||
|
||||
class ServeClientOpenVINO(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,
|
||||
single_model=False,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
The transcription thread is started upon initialization. A "SERVER_READY" message is sent
|
||||
to the client to indicate that the server is ready.
|
||||
|
||||
Args:
|
||||
websocket (WebSocket): The WebSocket connection for the client.
|
||||
task (str, optional): The task type, e.g., "transcribe." Defaults to "transcribe".
|
||||
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
||||
language (str, optional): The language for transcription. Defaults to None.
|
||||
client_uid (str, optional): A unique identifier for the client. Defaults to None.
|
||||
model (str, optional): Huggingface model_id for a valid OpenVINO model.
|
||||
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.
|
||||
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||
"""
|
||||
super().__init__(
|
||||
client_uid,
|
||||
websocket,
|
||||
send_last_n_segments,
|
||||
no_speech_thresh,
|
||||
clip_audio,
|
||||
same_output_threshold,
|
||||
)
|
||||
self.language = "en" if language is None else language
|
||||
if not self.language.startswith("<|"):
|
||||
self.language = f"<|{self.language}|>"
|
||||
|
||||
self.task = "transcribe" if task is None else task
|
||||
|
||||
self.clip_audio = True
|
||||
|
||||
core = Core()
|
||||
available_devices = core.available_devices
|
||||
if 'GPU' in available_devices:
|
||||
selected_device = 'GPU'
|
||||
else:
|
||||
gpu_devices = [d for d in available_devices if d.startswith('GPU')]
|
||||
selected_device = gpu_devices[0] if gpu_devices else 'CPU'
|
||||
self.device = selected_device
|
||||
|
||||
|
||||
if single_model:
|
||||
if ServeClientOpenVINO.SINGLE_MODEL is None:
|
||||
self.create_model(model)
|
||||
ServeClientOpenVINO.SINGLE_MODEL = self.transcriber
|
||||
else:
|
||||
self.transcriber = ServeClientOpenVINO.SINGLE_MODEL
|
||||
else:
|
||||
self.create_model(model)
|
||||
|
||||
# threading
|
||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||
self.trans_thread.start()
|
||||
|
||||
self.websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"message": self.SERVER_READY,
|
||||
"backend": "openvino"
|
||||
}))
|
||||
logging.info(f"Using OpenVINO device: {self.device}")
|
||||
logging.info(f"Running OpenVINO backend with language: {self.language} and task: {self.task}")
|
||||
|
||||
def create_model(self, model_id):
|
||||
"""
|
||||
Instantiates a new model, sets it as the transcriber.
|
||||
"""
|
||||
self.transcriber = WhisperOpenVINO(
|
||||
model_id,
|
||||
device=self.device,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
|
||||
def transcribe_audio(self, input_sample):
|
||||
"""
|
||||
Transcribes the provided audio sample using the configured transcriber instance.
|
||||
|
||||
If the language has not been set, it updates the session's language based on the transcription
|
||||
information.
|
||||
|
||||
Args:
|
||||
input_sample (np.array): The audio chunk to be transcribed. This should be a NumPy
|
||||
array representing the audio data.
|
||||
|
||||
Returns:
|
||||
The transcription result from the transcriber. The exact format of this result
|
||||
depends on the implementation of the `transcriber.transcribe` method but typically
|
||||
includes the transcribed text.
|
||||
"""
|
||||
if ServeClientOpenVINO.SINGLE_MODEL:
|
||||
ServeClientOpenVINO.SINGLE_MODEL_LOCK.acquire()
|
||||
result = self.transcriber.transcribe(input_sample)
|
||||
if ServeClientOpenVINO.SINGLE_MODEL:
|
||||
ServeClientOpenVINO.SINGLE_MODEL_LOCK.release()
|
||||
return result
|
||||
|
||||
def handle_transcription_output(self, result, duration):
|
||||
"""
|
||||
Handle the transcription output, updating the transcript and sending data to the client.
|
||||
|
||||
Args:
|
||||
result (str): The result from whisper inference i.e. the list of segments.
|
||||
duration (float): Duration of the transcribed audio chunk.
|
||||
"""
|
||||
segments = []
|
||||
if len(result):
|
||||
self.t_start = None
|
||||
last_segment = self.update_segments(result, duration)
|
||||
segments = self.prepare_segments(last_segment)
|
||||
|
||||
if len(segments):
|
||||
self.send_transcription_to_client(segments)
|
||||
@@ -0,0 +1,210 @@
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
from whisper_live.backend.base import ServeClientBase
|
||||
from whisper_live.transcriber.transcriber_tensorrt import WhisperTRTLLM
|
||||
|
||||
|
||||
class ServeClientTensorRT(ServeClientBase):
|
||||
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,
|
||||
use_py_session=False,
|
||||
max_new_tokens=225,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
The transcription thread is started upon initialization. A "SERVER_READY" message is sent
|
||||
to the client to indicate that the server is ready.
|
||||
|
||||
Args:
|
||||
websocket (WebSocket): The WebSocket connection for the client.
|
||||
task (str, optional): The task type, e.g., "transcribe." Defaults to "transcribe".
|
||||
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
||||
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.
|
||||
use_py_session (bool, optional): Use python session or cpp session. Defaults to Cpp Session.
|
||||
max_new_tokens (int, optional): Max number of tokens to generate.
|
||||
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||
"""
|
||||
super().__init__(
|
||||
client_uid,
|
||||
websocket,
|
||||
send_last_n_segments,
|
||||
no_speech_thresh,
|
||||
clip_audio,
|
||||
same_output_threshold,
|
||||
)
|
||||
|
||||
self.language = language if multilingual else "en"
|
||||
self.task = task
|
||||
self.eos = False
|
||||
self.max_new_tokens = max_new_tokens
|
||||
|
||||
if single_model:
|
||||
if ServeClientTensorRT.SINGLE_MODEL is None:
|
||||
self.create_model(model, multilingual, use_py_session=use_py_session)
|
||||
ServeClientTensorRT.SINGLE_MODEL = self.transcriber
|
||||
else:
|
||||
self.transcriber = ServeClientTensorRT.SINGLE_MODEL
|
||||
else:
|
||||
self.create_model(model, multilingual, use_py_session=use_py_session)
|
||||
|
||||
# threading
|
||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||
self.trans_thread.start()
|
||||
|
||||
self.websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"message": self.SERVER_READY,
|
||||
"backend": "tensorrt"
|
||||
}))
|
||||
|
||||
def create_model(self, model, multilingual, warmup=True, use_py_session=False):
|
||||
"""
|
||||
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,
|
||||
use_py_session=use_py_session,
|
||||
max_output_len=self.max_new_tokens,
|
||||
)
|
||||
if warmup:
|
||||
self.warmup()
|
||||
|
||||
def warmup(self, warmup_steps=10):
|
||||
"""
|
||||
Warmup TensorRT since first few inferences are slow.
|
||||
|
||||
Args:
|
||||
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("assets/jfk.flac")
|
||||
for i in range(warmup_steps):
|
||||
self.transcriber.transcribe(mel)
|
||||
|
||||
def set_eos(self, eos):
|
||||
"""
|
||||
Sets the End of Speech (EOS) flag.
|
||||
|
||||
Args:
|
||||
eos (bool): The value to set for the EOS flag.
|
||||
"""
|
||||
self.lock.acquire()
|
||||
self.eos = eos
|
||||
self.lock.release()
|
||||
|
||||
def handle_transcription_output(self, last_segment, duration):
|
||||
"""
|
||||
Handle the transcription output, updating the transcript and sending data to the client.
|
||||
|
||||
Args:
|
||||
last_segment (str): The last segment from the whisper output which is considered to be incomplete because
|
||||
of the possibility of word being truncated.
|
||||
duration (float): Duration of the transcribed audio chunk.
|
||||
"""
|
||||
segments = self.prepare_segments({"text": last_segment})
|
||||
self.send_transcription_to_client(segments)
|
||||
if self.eos:
|
||||
self.update_timestamp_offset(last_segment, duration)
|
||||
|
||||
def transcribe_audio(self, input_bytes):
|
||||
"""
|
||||
Transcribe the audio chunk and send the results to the client.
|
||||
|
||||
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,
|
||||
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)
|
||||
|
||||
def update_timestamp_offset(self, last_segment, duration):
|
||||
"""
|
||||
Update timestamp offset and transcript.
|
||||
|
||||
Args:
|
||||
last_segment (str): Last transcribed audio from the whisper model.
|
||||
duration (float): Duration of the last audio chunk.
|
||||
"""
|
||||
if not len(self.transcript):
|
||||
self.transcript.append({"text": last_segment + " "})
|
||||
elif self.transcript[-1]["text"].strip() != last_segment:
|
||||
self.transcript.append({"text": last_segment + " "})
|
||||
|
||||
with self.lock:
|
||||
self.timestamp_offset += duration
|
||||
|
||||
def speech_to_text(self):
|
||||
"""
|
||||
Process an audio stream in an infinite loop, continuously transcribing the speech.
|
||||
|
||||
This method continuously receives audio frames, performs real-time transcription, and sends
|
||||
transcribed segments to the client via a WebSocket connection.
|
||||
|
||||
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
|
||||
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
|
||||
are sent to the client in real-time, and a history of segments is maintained to provide context.
|
||||
|
||||
Raises:
|
||||
Exception: If there is an issue with audio processing or WebSocket communication.
|
||||
|
||||
"""
|
||||
while True:
|
||||
if self.exit:
|
||||
logging.info("Exiting speech to text thread")
|
||||
break
|
||||
|
||||
if self.frames_np is None:
|
||||
time.sleep(0.02) # wait for any audio to arrive
|
||||
continue
|
||||
|
||||
self.clip_audio_if_no_valid_segment()
|
||||
|
||||
input_bytes, duration = self.get_audio_chunk_for_processing()
|
||||
if duration < 0.4:
|
||||
continue
|
||||
|
||||
try:
|
||||
input_sample = input_bytes.copy()
|
||||
logging.info(f"[WhisperTensorRT:] Processing audio with duration: {duration}")
|
||||
self.transcribe_audio(input_sample)
|
||||
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
+551
-259
@@ -1,56 +1,44 @@
|
||||
import os
|
||||
import shutil
|
||||
import wave
|
||||
|
||||
import logging
|
||||
import numpy as np
|
||||
import scipy
|
||||
import ffmpeg
|
||||
import pyaudio
|
||||
import threading
|
||||
import textwrap
|
||||
import json
|
||||
import websocket
|
||||
import uuid
|
||||
import time
|
||||
|
||||
|
||||
def resample(file: str, sr: int = 16000):
|
||||
"""
|
||||
# https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22
|
||||
Open an audio file and read as mono waveform, resampling as necessary,
|
||||
save the resampled audio
|
||||
|
||||
Args:
|
||||
file (str): The audio file to open
|
||||
sr (int): The sample rate to resample the audio if necessary
|
||||
|
||||
Returns:
|
||||
resampled_file (str): The resampled audio file
|
||||
"""
|
||||
try:
|
||||
# This launches a subprocess to decode audio while down-mixing and resampling as necessary.
|
||||
# Requires the ffmpeg CLI and `ffmpeg-python` package to be installed.
|
||||
out, _ = (
|
||||
ffmpeg.input(file, threads=0)
|
||||
.output("-", format="s16le", acodec="pcm_s16le", ac=1, ar=sr)
|
||||
.run(cmd=["ffmpeg", "-nostdin"], capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
except ffmpeg.Error as e:
|
||||
raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e
|
||||
np_buffer = np.frombuffer(out, dtype=np.int16)
|
||||
|
||||
resampled_file = f"{file.split('.')[0]}_resampled.wav"
|
||||
scipy.io.wavfile.write(resampled_file, sr, np_buffer.astype(np.int16))
|
||||
return resampled_file
|
||||
import av
|
||||
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"
|
||||
|
||||
def __init__(
|
||||
self, host=None, port=None, is_multilingual=False, lang=None, translate=False
|
||||
self,
|
||||
host=None,
|
||||
port=None,
|
||||
lang=None,
|
||||
translate=False,
|
||||
model="small",
|
||||
srt_file_path="output.srt",
|
||||
use_vad=True,
|
||||
use_wss=False,
|
||||
log_transcription=True,
|
||||
max_clients=4,
|
||||
max_connection_time=600,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
transcription_callback=None,
|
||||
):
|
||||
"""
|
||||
Initializes a Client instance for audio recording and streaming to a server.
|
||||
@@ -62,41 +50,51 @@ class Client:
|
||||
Args:
|
||||
host (str): The hostname or IP address of the server.
|
||||
port (int): The port number for the WebSocket server.
|
||||
is_multilingual (bool, optional): Specifies if multilingual transcription is enabled. Default is False.
|
||||
lang (str, optional): The selected language for transcription when multilingual is disabled. Default is None.
|
||||
lang (str, optional): The selected language for transcription. Default is None.
|
||||
translate (bool, optional): Specifies if the task is translation. Default is False.
|
||||
model (str, optional): The whisper model to use (e.g., "small", "medium", "large"). Default is "small".
|
||||
srt_file_path (str, optional): The file path to save the output SRT file. Default is "output.srt".
|
||||
use_vad (bool, optional): Whether to enable voice activity detection. Default is True.
|
||||
log_transcription (bool, optional): Whether to log transcription output to the console. Default is True.
|
||||
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
||||
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
||||
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||
transcription_callback (callable, optional): A callback function to handle transcription results. Default is None.
|
||||
"""
|
||||
self.chunk = 1024
|
||||
self.format = pyaudio.paInt16
|
||||
self.channels = 1
|
||||
self.rate = 16000
|
||||
self.record_seconds = 60000
|
||||
self.recording = False
|
||||
self.multilingual = False
|
||||
self.language = None
|
||||
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.multilingual = is_multilingual
|
||||
self.language = lang if is_multilingual else "en"
|
||||
self.language = lang
|
||||
self.model = model
|
||||
self.server_error = False
|
||||
self.srt_file_path = srt_file_path
|
||||
self.use_vad = use_vad
|
||||
self.use_wss = use_wss
|
||||
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
|
||||
self.send_last_n_segments = send_last_n_segments
|
||||
self.no_speech_thresh = no_speech_thresh
|
||||
self.clip_audio = clip_audio
|
||||
self.same_output_threshold = same_output_threshold
|
||||
self.transcription_callback = transcription_callback
|
||||
|
||||
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}"
|
||||
socket_protocol = 'wss' if self.use_wss else "ws"
|
||||
socket_url = f"{socket_protocol}://{host}:{port}"
|
||||
self.client_socket = websocket.WebSocketApp(
|
||||
socket_url,
|
||||
on_open=lambda ws: self.on_open(ws),
|
||||
@@ -114,16 +112,59 @@ class Client:
|
||||
|
||||
# start websocket client in a thread
|
||||
self.ws_thread = threading.Thread(target=self.client_socket.run_forever)
|
||||
self.ws_thread.setDaemon(True)
|
||||
self.ws_thread.daemon = True
|
||||
self.ws_thread.start()
|
||||
|
||||
self.frames = b""
|
||||
self.transcript = []
|
||||
print("[INFO]: * recording")
|
||||
|
||||
def handle_status_messages(self, message_data):
|
||||
"""Handles server status messages."""
|
||||
status = message_data["status"]
|
||||
if status == "WAIT":
|
||||
self.waiting = True
|
||||
print(f"[INFO]: Server is full. Estimated wait time {round(message_data['message'])} minutes.")
|
||||
elif status == "ERROR":
|
||||
print(f"Message from Server: {message_data['message']}")
|
||||
self.server_error = True
|
||||
elif status == "WARNING":
|
||||
print(f"Message from Server: {message_data['message']}")
|
||||
|
||||
def process_segments(self, segments):
|
||||
"""Processes transcript segments."""
|
||||
text = []
|
||||
for i, seg in enumerate(segments):
|
||||
if not text or text[-1] != seg["text"]:
|
||||
text.append(seg["text"])
|
||||
if i == len(segments) - 1 and not seg.get("completed", False):
|
||||
self.last_segment = seg
|
||||
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 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"]
|
||||
|
||||
# call the transcription callback if provided
|
||||
if self.transcription_callback and callable(self.transcription_callback):
|
||||
try:
|
||||
self.transcription_callback(" ".join(text), segments) # string, list
|
||||
except Exception as e:
|
||||
print(f"[WARN] transcription_callback raised: {e}")
|
||||
return
|
||||
|
||||
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):
|
||||
"""
|
||||
Callback function called when a message is received from the server.
|
||||
|
||||
|
||||
It updates various attributes of the client based on the received message, including
|
||||
recording status, language detection, and server messages. If a disconnect message
|
||||
is received, it sets the recording status to False.
|
||||
@@ -133,25 +174,25 @@ class Client:
|
||||
message (str): The received message from the server.
|
||||
|
||||
"""
|
||||
self.last_response_recieved = time.time()
|
||||
message = json.loads(message)
|
||||
|
||||
if self.uid != message.get("uid"):
|
||||
print("[ERROR]: invalid client uid")
|
||||
return
|
||||
|
||||
if "status" in message.keys() and message["status"] == "WAIT":
|
||||
self.waiting = True
|
||||
print(
|
||||
f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes."
|
||||
)
|
||||
if "status" in message.keys():
|
||||
self.handle_status_messages(message)
|
||||
return
|
||||
|
||||
if "message" in message.keys() and message["message"] == "DISCONNECT":
|
||||
print("[INFO]: Server overtime disconnected.")
|
||||
print("[INFO]: Server disconnected due to overtime.")
|
||||
self.recording = False
|
||||
|
||||
if "message" in message.keys() and message["message"] == "SERVER_READY":
|
||||
self.last_response_received = time.time()
|
||||
self.recording = True
|
||||
self.server_backend = message["backend"]
|
||||
print(f"[INFO]: Server Running with backend {self.server_backend}")
|
||||
return
|
||||
|
||||
if "language" in message.keys():
|
||||
@@ -162,78 +203,49 @@ class Client:
|
||||
)
|
||||
return
|
||||
|
||||
if "segments" not in message.keys():
|
||||
return
|
||||
|
||||
message = message["segments"]
|
||||
text = []
|
||||
if len(message):
|
||||
for seg in message:
|
||||
if text and text[-1] == seg["text"]:
|
||||
# already got it
|
||||
continue
|
||||
text.append(seg["text"])
|
||||
# keep only last 3
|
||||
if len(text) > 3:
|
||||
text = text[-3:]
|
||||
wrapper = textwrap.TextWrapper(width=60)
|
||||
word_list = wrapper.wrap(text="".join(text))
|
||||
# Print each line.
|
||||
if os.name == "nt":
|
||||
os.system("cls")
|
||||
else:
|
||||
os.system("clear")
|
||||
for element in word_list:
|
||||
print(element)
|
||||
if "segments" in message.keys():
|
||||
self.process_segments(message["segments"])
|
||||
|
||||
def on_error(self, ws, error):
|
||||
print(error)
|
||||
print(f"[ERROR] WebSocket Error: {error}")
|
||||
self.server_error = True
|
||||
self.error_message = error
|
||||
|
||||
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.waiting = False
|
||||
|
||||
def on_open(self, ws):
|
||||
"""
|
||||
Callback function called when the WebSocket connection is successfully opened.
|
||||
|
||||
Sends an initial configuration message to the server, including client UID, multilingual mode,
|
||||
|
||||
Sends an initial configuration message to the server, including client UID,
|
||||
language selection, and task type.
|
||||
|
||||
Args:
|
||||
ws (websocket.WebSocketApp): The WebSocket client instance.
|
||||
|
||||
"""
|
||||
print(self.multilingual, self.language, self.task)
|
||||
|
||||
print("[INFO]: Opened connection")
|
||||
ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"uid": self.uid,
|
||||
"multilingual": self.multilingual,
|
||||
"language": self.language,
|
||||
"task": self.task,
|
||||
"model": self.model,
|
||||
"use_vad": self.use_vad,
|
||||
"max_clients": self.max_clients,
|
||||
"max_connection_time": self.max_connection_time,
|
||||
"send_last_n_segments": self.send_last_n_segments,
|
||||
"no_speech_thresh": self.no_speech_thresh,
|
||||
"clip_audio": self.clip_audio,
|
||||
"same_output_threshold": self.same_output_threshold,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@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.
|
||||
@@ -247,62 +259,11 @@ 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.stream.close()
|
||||
self.close_websocket()
|
||||
|
||||
except KeyboardInterrupt:
|
||||
wavfile.close()
|
||||
self.stream.stop_stream()
|
||||
self.stream.close()
|
||||
self.p.terminate()
|
||||
self.close_websocket()
|
||||
print("[INFO]: Keyboard interrupt.")
|
||||
|
||||
def close_websocket(self):
|
||||
"""
|
||||
Close the WebSocket connection and join the WebSocket thread.
|
||||
|
||||
First attempts to close the WebSocket connection using `self.client_socket.close()`. After
|
||||
First attempts to close the WebSocket connection using `self.client_socket.close()`. After
|
||||
closing the connection, it joins the WebSocket thread to ensure proper termination.
|
||||
|
||||
"""
|
||||
@@ -325,11 +286,341 @@ 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", mute_audio_playback=False):
|
||||
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.mute_audio_playback = mute_audio_playback
|
||||
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,
|
||||
)
|
||||
chunk_duration = self.chunk / float(wavfile.getframerate())
|
||||
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())
|
||||
if self.mute_audio_playback:
|
||||
time.sleep(chunk_duration)
|
||||
else:
|
||||
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 transcription.
|
||||
|
||||
Args:
|
||||
rtsp_url (str): The URL of the RTSP stream source.
|
||||
"""
|
||||
print("[INFO]: Connecting to RTSP stream...")
|
||||
try:
|
||||
container = av.open(rtsp_url, format="rtsp", options={"rtsp_transport": "tcp"})
|
||||
self.process_av_stream(container, stream_type="RTSP")
|
||||
except Exception as e:
|
||||
print(f"[ERROR]: Failed to process RTSP stream: {e}")
|
||||
finally:
|
||||
for client in self.clients:
|
||||
client.wait_before_disconnect()
|
||||
self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True)
|
||||
self.close_all_clients()
|
||||
self.write_all_clients_srt()
|
||||
print("[INFO]: RTSP stream processing finished.")
|
||||
|
||||
def process_hls_stream(self, hls_url, save_file=None):
|
||||
"""
|
||||
Connect to an HLS source, process the audio stream, and send it for transcription.
|
||||
|
||||
Args:
|
||||
hls_url (str): The URL of the HLS stream source.
|
||||
save_file (str, optional): Local path to save the network stream.
|
||||
"""
|
||||
print("[INFO]: Connecting to HLS stream...")
|
||||
try:
|
||||
container = av.open(hls_url, format="hls")
|
||||
self.process_av_stream(container, stream_type="HLS", save_file=save_file)
|
||||
except Exception as e:
|
||||
print(f"[ERROR]: Failed to process HLS stream: {e}")
|
||||
finally:
|
||||
for client in self.clients:
|
||||
client.wait_before_disconnect()
|
||||
self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True)
|
||||
self.close_all_clients()
|
||||
self.write_all_clients_srt()
|
||||
print("[INFO]: HLS stream processing finished.")
|
||||
|
||||
def process_av_stream(self, container, stream_type, save_file=None):
|
||||
"""
|
||||
Process an AV container stream and send audio packets to the server.
|
||||
|
||||
Args:
|
||||
container (av.container.InputContainer): The input container to process.
|
||||
stream_type (str): The type of stream being processed ("RTSP" or "HLS").
|
||||
save_file (str, optional): Local path to save the stream. Default is None.
|
||||
"""
|
||||
audio_stream = next((s for s in container.streams if s.type == "audio"), None)
|
||||
if not audio_stream:
|
||||
print(f"[ERROR]: No audio stream found in {stream_type} source.")
|
||||
return
|
||||
|
||||
output_container = None
|
||||
if save_file:
|
||||
output_container = av.open(save_file, mode="w")
|
||||
output_audio_stream = output_container.add_stream(codec_name="pcm_s16le", rate=self.rate)
|
||||
|
||||
try:
|
||||
for packet in container.demux(audio_stream):
|
||||
for frame in packet.decode():
|
||||
audio_data = frame.to_ndarray().tobytes()
|
||||
self.multicast_packet(audio_data)
|
||||
|
||||
if save_file:
|
||||
output_container.mux(frame)
|
||||
except Exception as e:
|
||||
print(f"[ERROR]: Error during {stream_type} stream processing: {e}")
|
||||
finally:
|
||||
# Wait for server to send any leftover transcription.
|
||||
time.sleep(5)
|
||||
self.multicast_packet(Client.END_OF_AUDIO.encode('utf-8'), True)
|
||||
if output_container:
|
||||
output_container.close()
|
||||
container.close()
|
||||
|
||||
def save_chunk(self, n_audio_file):
|
||||
"""
|
||||
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.
|
||||
|
||||
The WAV file is created or overwritten with the specified name. The audio frames should be
|
||||
The WAV file is created or overwritten with the specified name. The audio frames should be
|
||||
in the correct format and match the specified channel, sample width, and sample rate.
|
||||
|
||||
Args:
|
||||
@@ -344,68 +635,11 @@ class Client:
|
||||
wavfile.setframerate(self.rate)
|
||||
wavfile.writeframes(frames)
|
||||
|
||||
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)
|
||||
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""
|
||||
|
||||
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)
|
||||
|
||||
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.
|
||||
|
||||
The individual audio chunk files are expected to be located in the "chunks" directory. Reads each chunk
|
||||
|
||||
The individual audio chunk files are expected to be located in the "chunks" directory. Reads each chunk
|
||||
file, appends its audio data to the final recording, and then deletes the chunk file. After combining
|
||||
and saving, the final recording is stored in the specified `out_file`.
|
||||
|
||||
@@ -420,7 +654,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)
|
||||
@@ -435,11 +669,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")
|
||||
|
||||
@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.
|
||||
@@ -447,9 +701,22 @@ class TranscriptionClient:
|
||||
Args:
|
||||
host (str): The hostname or IP address of the server.
|
||||
port (int): The port number to connect to on the server.
|
||||
is_multilingual (bool, optional): Indicates whether the transcription should support multiple languages (default is False).
|
||||
lang (str, optional): The primary language for transcription (used if `is_multilingual` is False). Default is None, which defaults to English ('en').
|
||||
translate (bool, optional): Indicates whether translation tasks are required (default is False).
|
||||
lang (str, optional): The primary language for transcription. Default is None, which defaults to English ('en').
|
||||
translate (bool, optional): If True, the task will be translation instead of transcription. Default is False.
|
||||
model (str, optional): The whisper model to use (e.g., "small", "base"). Default is "small".
|
||||
use_vad (bool, optional): Whether to enable voice activity detection. Default is True.
|
||||
save_output_recording (bool, optional): Whether to save the microphone recording. Default is False.
|
||||
output_recording_filename (str, optional): Path to save the output recording WAV file. Default is "./output_recording.wav".
|
||||
output_transcription_path (str, optional): File path to save the output transcription (SRT file). Default is "./output.srt".
|
||||
log_transcription (bool, optional): Whether to log transcription output to the console. Default is True.
|
||||
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
||||
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
||||
mute_audio_playback (bool, optional): If True, mutes audio playback during file playback. Default is False.
|
||||
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||
transcription_callback (callable, optional): A callback function to handle transcription results. Default is None.
|
||||
|
||||
Attributes:
|
||||
client (Client): An instance of the underlying Client class responsible for handling the WebSocket connection.
|
||||
@@ -457,34 +724,59 @@ class TranscriptionClient:
|
||||
Example:
|
||||
To create a TranscriptionClient and start transcription on microphone audio:
|
||||
```python
|
||||
transcription_client = TranscriptionClient(host="localhost", port=9090, is_multilingual=True)
|
||||
transcription_client = TranscriptionClient(host="localhost", port=9090)
|
||||
transcription_client()
|
||||
```
|
||||
"""
|
||||
def __init__(self, host, port, is_multilingual=False, lang=None, translate=False):
|
||||
self.client = Client(host, port, is_multilingual, lang, translate)
|
||||
def __init__(
|
||||
self,
|
||||
host,
|
||||
port,
|
||||
lang=None,
|
||||
translate=False,
|
||||
model="small",
|
||||
use_vad=True,
|
||||
use_wss=False,
|
||||
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,
|
||||
mute_audio_playback=False,
|
||||
send_last_n_segments=10,
|
||||
no_speech_thresh=0.45,
|
||||
clip_audio=False,
|
||||
same_output_threshold=10,
|
||||
transcription_callback=None,
|
||||
):
|
||||
self.client = Client(
|
||||
host,
|
||||
port,
|
||||
lang,
|
||||
translate,
|
||||
model,
|
||||
srt_file_path=output_transcription_path,
|
||||
use_vad=use_vad,
|
||||
use_wss=use_wss,
|
||||
log_transcription=log_transcription,
|
||||
max_clients=max_clients,
|
||||
max_connection_time=max_connection_time,
|
||||
send_last_n_segments=send_last_n_segments,
|
||||
no_speech_thresh=no_speech_thresh,
|
||||
clip_audio=clip_audio,
|
||||
same_output_threshold=same_output_threshold,
|
||||
transcription_callback=transcription_callback,
|
||||
)
|
||||
|
||||
def __call__(self, audio=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:
|
||||
self.client.close_websocket()
|
||||
return
|
||||
pass
|
||||
print("[INFO]: Server Ready!")
|
||||
if audio is not None:
|
||||
resampled_file = 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,
|
||||
mute_audio_playback=mute_audio_playback
|
||||
)
|
||||
|
||||
+383
-468
@@ -1,70 +1,326 @@
|
||||
import websockets
|
||||
import os
|
||||
import time
|
||||
import threading
|
||||
import json
|
||||
import textwrap
|
||||
|
||||
import functools
|
||||
import logging
|
||||
# logging.basicConfig(level = logging.INFO)
|
||||
from enum import Enum
|
||||
from typing import List, Optional
|
||||
|
||||
from websockets.sync.server import serve
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import time
|
||||
from whisper_live.transcriber import WhisperModel
|
||||
from whisper_live.vad import VoiceActivityDetection
|
||||
from websockets.sync.server import serve
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from whisper_live.vad import VoiceActivityDetector
|
||||
from whisper_live.backend.base import ServeClientBase
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
|
||||
class TranscriptionServer:
|
||||
"""
|
||||
Represents a transcription server that handles incoming audio from clients.
|
||||
|
||||
Attributes:
|
||||
RATE (int): The audio sampling rate (constant) set to 16000.
|
||||
vad_model (torch.Module): The voice activity detection model.
|
||||
vad_threshold (float): The voice activity detection threshold.
|
||||
clients (dict): A dictionary to store connected clients.
|
||||
websockets (dict): A dictionary to store WebSocket connections.
|
||||
clients_start_time (dict): A dictionary to track client start times.
|
||||
max_clients (int): Maximum allowed connected clients.
|
||||
max_connection_time (int): Maximum allowed connection time in seconds.
|
||||
"""
|
||||
|
||||
RATE = 16000
|
||||
|
||||
def __init__(self):
|
||||
# voice activity detection model
|
||||
self.vad_model = VoiceActivityDetection()
|
||||
self.vad_threshold = 0.4
|
||||
class ClientManager:
|
||||
def __init__(self, max_clients=4, max_connection_time=600):
|
||||
"""
|
||||
Initializes the ClientManager with specified limits on client connections and connection durations.
|
||||
|
||||
Args:
|
||||
max_clients (int, optional): The maximum number of simultaneous client connections allowed. Defaults to 4.
|
||||
max_connection_time (int, optional): The maximum duration (in seconds) a client can stay connected. Defaults
|
||||
to 600 seconds (10 minutes).
|
||||
"""
|
||||
self.clients = {}
|
||||
self.websockets = {}
|
||||
self.clients_start_time = {}
|
||||
self.max_clients = 4
|
||||
self.max_connection_time = 600
|
||||
self.start_times = {}
|
||||
self.max_clients = max_clients
|
||||
self.max_connection_time = max_connection_time
|
||||
|
||||
def add_client(self, websocket, client):
|
||||
"""
|
||||
Adds a client and their connection start time to the tracking dictionaries.
|
||||
|
||||
Args:
|
||||
websocket: The websocket associated with the client to add.
|
||||
client: The client object to be added and tracked.
|
||||
"""
|
||||
self.clients[websocket] = client
|
||||
self.start_times[websocket] = time.time()
|
||||
|
||||
def get_client(self, websocket):
|
||||
"""
|
||||
Retrieves a client associated with the given websocket.
|
||||
|
||||
Args:
|
||||
websocket: The websocket associated with the client to retrieve.
|
||||
|
||||
Returns:
|
||||
The client object if found, False otherwise.
|
||||
"""
|
||||
if websocket in self.clients:
|
||||
return self.clients[websocket]
|
||||
return False
|
||||
|
||||
def remove_client(self, websocket):
|
||||
"""
|
||||
Removes a client and their connection start time from the tracking dictionaries. Performs cleanup on the
|
||||
client if necessary.
|
||||
|
||||
Args:
|
||||
websocket: The websocket associated with the client to be removed.
|
||||
"""
|
||||
client = self.clients.pop(websocket, None)
|
||||
if client:
|
||||
client.cleanup()
|
||||
self.start_times.pop(websocket, None)
|
||||
|
||||
def get_wait_time(self):
|
||||
"""
|
||||
Calculate and return the estimated wait time for clients.
|
||||
Calculates the estimated wait time for new clients based on the remaining connection times of current clients.
|
||||
|
||||
Returns:
|
||||
float: The estimated wait time in minutes.
|
||||
The estimated wait time in minutes for new clients to connect. Returns 0 if there are available slots.
|
||||
"""
|
||||
wait_time = None
|
||||
|
||||
for k, v in self.clients_start_time.items():
|
||||
current_client_time_remaining = self.max_connection_time - (time.time() - v)
|
||||
|
||||
for start_time in self.start_times.values():
|
||||
current_client_time_remaining = self.max_connection_time - (time.time() - start_time)
|
||||
if wait_time is None or current_client_time_remaining < wait_time:
|
||||
wait_time = current_client_time_remaining
|
||||
return wait_time / 60 if wait_time is not None else 0
|
||||
|
||||
return wait_time / 60
|
||||
def is_server_full(self, websocket, options):
|
||||
"""
|
||||
Checks if the server is at its maximum client capacity and sends a wait message to the client if necessary.
|
||||
|
||||
def recv_audio(self, websocket):
|
||||
Args:
|
||||
websocket: The websocket of the client attempting to connect.
|
||||
options: A dictionary of options that may include the client's unique identifier.
|
||||
|
||||
Returns:
|
||||
True if the server is full, False otherwise.
|
||||
"""
|
||||
if len(self.clients) >= self.max_clients:
|
||||
wait_time = self.get_wait_time()
|
||||
response = {"uid": options["uid"], "status": "WAIT", "message": wait_time}
|
||||
websocket.send(json.dumps(response))
|
||||
return True
|
||||
return False
|
||||
|
||||
def is_client_timeout(self, websocket):
|
||||
"""
|
||||
Checks if a client has exceeded the maximum allowed connection time and disconnects them if so, issuing a warning.
|
||||
|
||||
Args:
|
||||
websocket: The websocket associated with the client to check.
|
||||
|
||||
Returns:
|
||||
True if the client's connection time has exceeded the maximum limit, False otherwise.
|
||||
"""
|
||||
elapsed_time = time.time() - self.start_times[websocket]
|
||||
if elapsed_time >= self.max_connection_time:
|
||||
self.clients[websocket].disconnect()
|
||||
logging.warning(f"Client with uid '{self.clients[websocket].client_uid}' disconnected due to overtime.")
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class BackendType(Enum):
|
||||
FASTER_WHISPER = "faster_whisper"
|
||||
TENSORRT = "tensorrt"
|
||||
OPENVINO = "openvino"
|
||||
|
||||
@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
|
||||
|
||||
def is_openvino(self) -> bool:
|
||||
return self == BackendType.OPENVINO
|
||||
|
||||
|
||||
class TranscriptionServer:
|
||||
RATE = 16000
|
||||
|
||||
def __init__(self):
|
||||
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, trt_py_session=False,
|
||||
):
|
||||
client: Optional[ServeClientBase] = None
|
||||
|
||||
if self.backend.is_tensorrt():
|
||||
try:
|
||||
from whisper_live.backend.trt_backend import ServeClientTensorRT
|
||||
client = ServeClientTensorRT(
|
||||
websocket,
|
||||
multilingual=trt_multilingual,
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"],
|
||||
model=whisper_tensorrt_path,
|
||||
single_model=self.single_model,
|
||||
use_py_session=trt_py_session,
|
||||
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||
clip_audio=options.get("clip_audio", False),
|
||||
same_output_threshold=options.get("same_output_threshold", 10),
|
||||
)
|
||||
logging.info("Running TensorRT backend.")
|
||||
except Exception as e:
|
||||
logging.error(f"TensorRT-LLM not supported: {e}")
|
||||
self.client_uid = options["uid"]
|
||||
websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"status": "WARNING",
|
||||
"message": "TensorRT-LLM not supported on Server yet. "
|
||||
"Reverting to available backend: 'faster_whisper'"
|
||||
}))
|
||||
self.backend = BackendType.FASTER_WHISPER
|
||||
|
||||
if self.backend.is_openvino():
|
||||
try:
|
||||
from whisper_live.backend.openvino_backend import ServeClientOpenVINO
|
||||
client = ServeClientOpenVINO(
|
||||
websocket,
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"],
|
||||
model=options["model"],
|
||||
single_model=self.single_model,
|
||||
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||
clip_audio=options.get("clip_audio", False),
|
||||
same_output_threshold=options.get("same_output_threshold", 10),
|
||||
)
|
||||
logging.info("Running OpenVINO backend.")
|
||||
except Exception as e:
|
||||
logging.error(f"OpenVINO not supported: {e}")
|
||||
self.backend = BackendType.FASTER_WHISPER
|
||||
self.client_uid = options["uid"]
|
||||
websocket.send(json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"status": "WARNING",
|
||||
"message": "OpenVINO not supported on Server yet. "
|
||||
"Reverting to available backend: 'faster_whisper'"
|
||||
}))
|
||||
|
||||
try:
|
||||
if self.backend.is_faster_whisper():
|
||||
from whisper_live.backend.faster_whisper_backend import ServeClientFasterWhisper
|
||||
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,
|
||||
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||
clip_audio=options.get("clip_audio", False),
|
||||
same_output_threshold=options.get("same_output_threshold", 10),
|
||||
)
|
||||
|
||||
logging.info("Running faster_whisper backend.")
|
||||
except Exception as e:
|
||||
logging.error(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)
|
||||
|
||||
def get_audio_from_websocket(self, websocket):
|
||||
"""
|
||||
Receives audio buffer from websocket and creates a numpy array out of it.
|
||||
|
||||
Args:
|
||||
websocket: The websocket to receive audio from.
|
||||
|
||||
Returns:
|
||||
A numpy array containing the audio.
|
||||
"""
|
||||
frame_data = websocket.recv()
|
||||
if frame_data == b"END_OF_AUDIO":
|
||||
return False
|
||||
return np.frombuffer(frame_data, dtype=np.float32)
|
||||
|
||||
def handle_new_connection(self, websocket, faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path, trt_multilingual, trt_py_session=False):
|
||||
try:
|
||||
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.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, trt_py_session=trt_py_session)
|
||||
return True
|
||||
except json.JSONDecodeError:
|
||||
logging.error("Failed to decode JSON from client")
|
||||
return False
|
||||
except ConnectionClosed:
|
||||
logging.info("Connection closed by client")
|
||||
return False
|
||||
except Exception as e:
|
||||
logging.error(f"Error during new connection initialization: {str(e)}")
|
||||
return False
|
||||
|
||||
def process_audio_frames(self, websocket):
|
||||
frame_np = self.get_audio_from_websocket(websocket)
|
||||
client = self.client_manager.get_client(websocket)
|
||||
if frame_np is False:
|
||||
if self.backend.is_tensorrt():
|
||||
client.set_eos(True)
|
||||
return False
|
||||
|
||||
if self.backend.is_tensorrt():
|
||||
voice_active = self.voice_activity(websocket, frame_np)
|
||||
if voice_active:
|
||||
self.no_voice_activity_chunks = 0
|
||||
client.set_eos(False)
|
||||
if self.use_vad and not voice_active:
|
||||
return True
|
||||
|
||||
client.add_frames(frame_np)
|
||||
return True
|
||||
|
||||
def recv_audio(self,
|
||||
websocket,
|
||||
backend: BackendType = BackendType.FASTER_WHISPER,
|
||||
faster_whisper_custom_model_path=None,
|
||||
whisper_tensorrt_path=None,
|
||||
trt_multilingual=False,
|
||||
trt_py_session=False):
|
||||
"""
|
||||
Receive audio chunks from a client in an infinite loop.
|
||||
|
||||
|
||||
Continuously receives audio frames from a connected client
|
||||
over a WebSocket connection. It processes the audio frames using a
|
||||
voice activity detection (VAD) model to determine if they contain speech
|
||||
@@ -78,76 +334,42 @@ class TranscriptionServer:
|
||||
|
||||
Args:
|
||||
websocket (WebSocket): The WebSocket connection for the client.
|
||||
|
||||
backend (str): The backend to run the server with.
|
||||
faster_whisper_custom_model_path (str): path to custom faster whisper model.
|
||||
whisper_tensorrt_path (str): Required for tensorrt backend.
|
||||
trt_multilingual(bool): Only used for tensorrt, True if multilingual model.
|
||||
|
||||
Raises:
|
||||
Exception: If there is an error during the audio frame processing.
|
||||
"""
|
||||
logging.info("New client connected")
|
||||
options = websocket.recv()
|
||||
options = json.loads(options)
|
||||
|
||||
if len(self.clients) >= self.max_clients:
|
||||
logging.warning("Client Queue Full. Asking client to wait ...")
|
||||
wait_time = self.get_wait_time()
|
||||
response = {
|
||||
"uid": options["uid"],
|
||||
"status": "WAIT",
|
||||
"message": wait_time,
|
||||
}
|
||||
websocket.send(json.dumps(response))
|
||||
websocket.close()
|
||||
del websocket
|
||||
self.backend = backend
|
||||
if not self.handle_new_connection(websocket, faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path, trt_multilingual, trt_py_session=trt_py_session):
|
||||
return
|
||||
|
||||
client = ServeClient(
|
||||
websocket,
|
||||
multilingual=options["multilingual"],
|
||||
language=options["language"],
|
||||
task=options["task"],
|
||||
client_uid=options["uid"]
|
||||
)
|
||||
|
||||
self.clients[websocket] = client
|
||||
self.clients_start_time[websocket] = time.time()
|
||||
|
||||
while True:
|
||||
try:
|
||||
frame_data = websocket.recv()
|
||||
frame_np = np.frombuffer(frame_data, dtype=np.float32)
|
||||
|
||||
try:
|
||||
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
|
||||
if speech_prob < self.vad_threshold:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
logging.error(e)
|
||||
return
|
||||
|
||||
self.clients[websocket].add_frames(frame_np)
|
||||
|
||||
elapsed_time = time.time() - self.clients_start_time[websocket]
|
||||
if elapsed_time >= self.max_connection_time:
|
||||
self.clients[websocket].disconnect()
|
||||
logging.warning(f"{self.clients[websocket]} Client disconnected due to overtime.")
|
||||
self.clients[websocket].cleanup()
|
||||
self.clients.pop(websocket)
|
||||
self.clients_start_time.pop(websocket)
|
||||
websocket.close()
|
||||
del websocket
|
||||
try:
|
||||
while not self.client_manager.is_client_timeout(websocket):
|
||||
if not self.process_audio_frames(websocket):
|
||||
break
|
||||
except ConnectionClosed:
|
||||
logging.info("Connection closed by client")
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error: {str(e)}")
|
||||
finally:
|
||||
if self.client_manager.get_client(websocket):
|
||||
self.cleanup(websocket)
|
||||
websocket.close()
|
||||
del websocket
|
||||
|
||||
except Exception as e:
|
||||
logging.error(e)
|
||||
self.clients[websocket].cleanup()
|
||||
self.clients.pop(websocket)
|
||||
self.clients_start_time.pop(websocket)
|
||||
logging.info("Connection Closed.")
|
||||
logging.info(self.clients)
|
||||
del websocket
|
||||
break
|
||||
|
||||
def run(self, host, port=9090):
|
||||
def run(self,
|
||||
host,
|
||||
port=9090,
|
||||
backend="tensorrt",
|
||||
faster_whisper_custom_model_path=None,
|
||||
whisper_tensorrt_path=None,
|
||||
trt_multilingual=False,
|
||||
trt_py_session=False,
|
||||
single_model=False):
|
||||
"""
|
||||
Run the transcription server.
|
||||
|
||||
@@ -155,377 +377,70 @@ class TranscriptionServer:
|
||||
host (str): The host address to bind the server.
|
||||
port (int): The port number to bind the server.
|
||||
"""
|
||||
with serve(self.recv_audio, host, port) as 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=BackendType(backend),
|
||||
faster_whisper_custom_model_path=faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path=whisper_tensorrt_path,
|
||||
trt_multilingual=trt_multilingual,
|
||||
trt_py_session=trt_py_session,
|
||||
),
|
||||
host,
|
||||
port
|
||||
) as server:
|
||||
server.serve_forever()
|
||||
|
||||
|
||||
class ServeClient:
|
||||
"""
|
||||
Attributes:
|
||||
RATE (int): The audio sampling rate (constant) set to 16000.
|
||||
SERVER_READY (str): A constant message indicating that the server is ready.
|
||||
DISCONNECT (str): A constant message indicating that the client should disconnect.
|
||||
client_uid (str): A unique identifier for the client.
|
||||
data (bytes): Accumulated audio data.
|
||||
frames (bytes): Accumulated audio frames.
|
||||
language (str): The language for transcription.
|
||||
task (str): The task type, e.g., "transcribe."
|
||||
transcriber (WhisperModel): The Whisper model for speech-to-text.
|
||||
timestamp_offset (float): The offset in audio timestamps.
|
||||
frames_np (numpy.ndarray): NumPy array to store audio frames.
|
||||
frames_offset (float): The offset in audio frames.
|
||||
text (list): List of transcribed text segments.
|
||||
current_out (str): The current incomplete transcription.
|
||||
prev_out (str): The previous incomplete transcription.
|
||||
t_start (float): Timestamp for the start of transcription.
|
||||
exit (bool): A flag to exit the transcription thread.
|
||||
same_output_threshold (int): Threshold for consecutive same output segments.
|
||||
show_prev_out_thresh (int): Threshold for showing previous output segments.
|
||||
add_pause_thresh (int): Threshold for adding a pause (blank) segment.
|
||||
transcript (list): List of transcribed segments.
|
||||
send_last_n_segments (int): Number of last segments to send to the client.
|
||||
wrapper (textwrap.TextWrapper): Text wrapper for formatting text.
|
||||
pick_previous_segments (int): Number of previous segments to include in the output.
|
||||
websocket: The WebSocket connection for the client.
|
||||
"""
|
||||
RATE = 16000
|
||||
SERVER_READY = "SERVER_READY"
|
||||
DISCONNECT = "DISCONNECT"
|
||||
|
||||
def __init__(self, websocket, task="transcribe", device=None, multilingual=False, language=None, client_uid=None):
|
||||
def voice_activity(self, websocket, frame_np):
|
||||
"""
|
||||
Initialize a ServeClient instance.
|
||||
The Whisper model is initialized based on the client's language and device availability.
|
||||
The transcription thread is started upon initialization. A "SERVER_READY" message is sent
|
||||
to the client to indicate that the server is ready.
|
||||
Evaluates the voice activity in a given audio frame and manages the state of voice activity detection.
|
||||
|
||||
This method uses the configured voice activity detection (VAD) model to assess whether the given audio frame
|
||||
contains speech. If the VAD model detects no voice activity for more than three consecutive frames,
|
||||
it sets an end-of-speech (EOS) flag for the associated client. This method aims to efficiently manage
|
||||
speech detection to improve subsequent processing steps.
|
||||
|
||||
Args:
|
||||
websocket (WebSocket): The WebSocket connection for the client.
|
||||
task (str, optional): The task type, e.g., "transcribe." Defaults to "transcribe".
|
||||
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
||||
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.
|
||||
websocket: The websocket associated with the current client. Used to retrieve the client object
|
||||
from the client manager for state management.
|
||||
frame_np (numpy.ndarray): The audio frame to be analyzed. This should be a NumPy array containing
|
||||
the audio data for the current frame.
|
||||
|
||||
"""
|
||||
self.client_uid = client_uid
|
||||
self.data = b""
|
||||
self.frames = b""
|
||||
self.language = language if multilingual else "en"
|
||||
self.task = task
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
self.transcriber = WhisperModel(
|
||||
"small" if multilingual else "small.en",
|
||||
device=device,
|
||||
compute_type="int8" if device=="cpu" else "float16",
|
||||
local_files_only=False,
|
||||
)
|
||||
|
||||
self.timestamp_offset = 0.0
|
||||
self.frames_np = None
|
||||
self.frames_offset = 0.0
|
||||
self.text = []
|
||||
self.current_out = ''
|
||||
self.prev_out = ''
|
||||
self.t_start=None
|
||||
self.exit = False
|
||||
self.same_output_threshold = 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 = []
|
||||
self.send_last_n_segments = 10
|
||||
|
||||
# text formatting
|
||||
self.wrapper = textwrap.TextWrapper(width=50)
|
||||
self.pick_previous_segments = 2
|
||||
|
||||
# threading
|
||||
self.websocket = websocket
|
||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||
self.trans_thread.start()
|
||||
self.websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
"uid": self.client_uid,
|
||||
"message": self.SERVER_READY
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
def fill_output(self, output):
|
||||
"""
|
||||
Format the current incomplete transcription output by combining it with previous complete segments.
|
||||
The resulting transcription is wrapped into two lines, each containing a maximum of 50 characters.
|
||||
|
||||
It ensures that the combined transcription fits within two lines, with a maximum of 50 characters per line.
|
||||
Segments are concatenated in the order they exist in the list of previous segments, with the most
|
||||
recent complete segment first and older segments prepended as needed to maintain the character limit.
|
||||
If a 3-second pause is detected in the previous segments, any text preceding it is discarded to ensure
|
||||
the transcription starts with the most recent complete content. The resulting transcription is returned
|
||||
as a single string.
|
||||
|
||||
Args:
|
||||
output(str): The current incomplete transcription segment.
|
||||
|
||||
Returns:
|
||||
str: A formatted transcription wrapped in two lines.
|
||||
bool: True if voice activity is detected in the current frame, False otherwise. When returning False
|
||||
after detecting no voice activity for more than three consecutive frames, it also triggers the
|
||||
end-of-speech (EOS) flag for the client.
|
||||
"""
|
||||
text = ''
|
||||
pick_prev = min(len(self.text), self.pick_previous_segments)
|
||||
for seg in self.text[-pick_prev:]:
|
||||
# discard everything before a 3 second pause
|
||||
if seg == '':
|
||||
text = ''
|
||||
else:
|
||||
text += seg
|
||||
wrapped = "".join(text + output)
|
||||
return wrapped
|
||||
|
||||
def add_frames(self, frame_np):
|
||||
if not self.vad_detector(frame_np):
|
||||
self.no_voice_activity_chunks += 1
|
||||
if self.no_voice_activity_chunks > 3:
|
||||
client = self.client_manager.get_client(websocket)
|
||||
if not client.eos:
|
||||
client.set_eos(True)
|
||||
time.sleep(0.1) # Sleep 100m; wait some voice activity.
|
||||
return False
|
||||
return True
|
||||
|
||||
def cleanup(self, websocket):
|
||||
"""
|
||||
Add audio frames to the ongoing audio stream buffer.
|
||||
|
||||
This method is responsible for maintaining the audio stream buffer, allowing the continuous addition
|
||||
of audio frames as they are received. It also ensures that the buffer does not exceed a specified size
|
||||
to prevent excessive memory usage.
|
||||
|
||||
If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds
|
||||
of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided
|
||||
audio frame. The audio stream buffer is used for real-time processing of audio data for transcription.
|
||||
Cleans up resources associated with a given client's websocket.
|
||||
|
||||
Args:
|
||||
frame_np (numpy.ndarray): The audio frame data as a NumPy array.
|
||||
|
||||
websocket: The websocket associated with the client to be cleaned up.
|
||||
"""
|
||||
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):]
|
||||
if self.frames_np is None:
|
||||
self.frames_np = frame_np.copy()
|
||||
else:
|
||||
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
|
||||
if self.client_manager.get_client(websocket):
|
||||
self.client_manager.remove_client(websocket)
|
||||
|
||||
def speech_to_text(self):
|
||||
"""
|
||||
Process an audio stream in an infinite loop, continuously transcribing the speech.
|
||||
|
||||
This method continuously receives audio frames, performs real-time transcription, and sends
|
||||
transcribed segments to the client via a WebSocket connection.
|
||||
|
||||
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
|
||||
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
|
||||
are sent to the client in real-time, and a history of segments is maintained to provide context.Pauses in speech
|
||||
(no output from Whisper) are handled by showing the previous output for a set duration. A blank segment is added if
|
||||
there is no speech for a specified duration to indicate a pause.
|
||||
|
||||
Raises:
|
||||
Exception: If there is an issue with audio processing or WebSocket communication.
|
||||
|
||||
"""
|
||||
# detect language
|
||||
if self.language is None:
|
||||
# wait for 30s of audio
|
||||
while self.frames_np is None or self.frames_np.shape[0] < 30*self.RATE:
|
||||
time.sleep(1)
|
||||
input_bytes = self.frames_np[-30*self.RATE:].copy()
|
||||
self.frames_np = None
|
||||
duration = input_bytes.shape[0] / self.RATE
|
||||
|
||||
self.language, lang_prob = self.transcriber.transcribe(
|
||||
input_bytes,
|
||||
initial_prompt=None,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
logging.info(f"Detected language {self.language} with probability {lang_prob}")
|
||||
self.websocket.send(json.dumps(
|
||||
{"uid": self.client_uid, "language": self.language, "language_prob": lang_prob}))
|
||||
|
||||
while True:
|
||||
if self.exit:
|
||||
logging.info("Exiting speech to text thread")
|
||||
break
|
||||
|
||||
if self.frames_np is None:
|
||||
continue
|
||||
|
||||
# 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
|
||||
|
||||
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
|
||||
if duration<1.0:
|
||||
continue
|
||||
try:
|
||||
input_sample = input_bytes.copy()
|
||||
# set previous complete segment as initial prompt
|
||||
if len(self.text) and self.text[-1] != '':
|
||||
initial_prompt = self.text[-1]
|
||||
else:
|
||||
initial_prompt = None
|
||||
|
||||
# whisper transcribe with prompt
|
||||
result = self.transcriber.transcribe(
|
||||
input_sample,
|
||||
initial_prompt=initial_prompt,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
|
||||
if len(result):
|
||||
self.t_start = None
|
||||
last_segment = self.update_segments(result, duration)
|
||||
if len(self.transcript) < self.send_last_n_segments:
|
||||
segments = self.transcript
|
||||
else:
|
||||
segments = self.transcript[-self.send_last_n_segments:]
|
||||
if last_segment is not None:
|
||||
segments = segments + [last_segment]
|
||||
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
else:
|
||||
# show previous output if there is pause i.e. no output from whisper
|
||||
segments = []
|
||||
if self.t_start is None: self.t_start = time.time()
|
||||
if time.time() - self.t_start < self.show_prev_out_thresh:
|
||||
if len(self.transcript) < self.send_last_n_segments:
|
||||
segments = self.transcript
|
||||
else:
|
||||
segments = self.transcript[-self.send_last_n_segments:]
|
||||
|
||||
# add a blank if there is no speech for 3 seconds
|
||||
if len(self.text) and self.text[-1] != '':
|
||||
if time.time() - self.t_start > self.add_pause_thresh:
|
||||
self.text.append('')
|
||||
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
time.sleep(0.01)
|
||||
|
||||
def update_segments(self, segments, duration):
|
||||
"""
|
||||
Processes the segments from whisper. Appends all the segments to the list
|
||||
except for the last segment assuming that it is incomplete.
|
||||
|
||||
Updates the ongoing transcript with transcribed segments, including their start and end times.
|
||||
Complete segments are appended to the transcript in chronological order. Incomplete segments
|
||||
(assumed to be the last one) are processed to identify repeated content. If the same incomplete
|
||||
segment is seen multiple times, it updates the offset and appends the segment to the transcript.
|
||||
A threshold is used to detect repeated content and ensure it is only included once in the transcript.
|
||||
The timestamp offset is updated based on the duration of processed segments. The method returns the
|
||||
last processed segment, allowing it to be sent to the client for real-time updates.
|
||||
|
||||
Args:
|
||||
segments(dict) : dictionary of segments as returned by whisper
|
||||
duration(float): duration of the current chunk
|
||||
|
||||
Returns:
|
||||
dict or None: The last processed segment with its start time, end time, and transcribed text.
|
||||
Returns None if there are no valid segments to process.
|
||||
"""
|
||||
offset = None
|
||||
self.current_out = ''
|
||||
last_segment = None
|
||||
# process complete segments
|
||||
if len(segments) > 1:
|
||||
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)
|
||||
self.transcript.append(
|
||||
{
|
||||
'start': start,
|
||||
'end': end,
|
||||
'text': text_
|
||||
}
|
||||
)
|
||||
|
||||
offset = min(duration, s.end)
|
||||
|
||||
self.current_out += segments[-1].text
|
||||
last_segment = {
|
||||
'start': self.timestamp_offset + segments[-1].start,
|
||||
'end': self.timestamp_offset + min(duration, segments[-1].end),
|
||||
'text': self.current_out
|
||||
}
|
||||
|
||||
# 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 not len(self.text) or self.text[-1].strip().lower()!=self.current_out.strip().lower():
|
||||
self.text.append(self.current_out)
|
||||
self.transcript.append(
|
||||
{
|
||||
'start': self.timestamp_offset,
|
||||
'end': self.timestamp_offset + duration,
|
||||
'text': self.current_out
|
||||
}
|
||||
)
|
||||
self.current_out = ''
|
||||
offset = duration
|
||||
self.same_output_threshold = 0
|
||||
last_segment = None
|
||||
else:
|
||||
self.prev_out = self.current_out
|
||||
|
||||
# update offset
|
||||
if offset is not None:
|
||||
self.timestamp_offset += offset
|
||||
|
||||
return last_segment
|
||||
|
||||
def disconnect(self):
|
||||
"""
|
||||
Notify the client of disconnection and send a disconnect message.
|
||||
|
||||
This method sends a disconnect message to the client via the WebSocket connection to notify them
|
||||
that the transcription service is disconnecting gracefully.
|
||||
|
||||
"""
|
||||
self.websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
"uid": self.client_uid,
|
||||
"message": self.DISCONNECT
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
def cleanup(self):
|
||||
"""
|
||||
Perform cleanup tasks before exiting the transcription service.
|
||||
|
||||
This method performs necessary cleanup tasks, including stopping the transcription thread, marking
|
||||
the exit flag to indicate the transcription thread should exit gracefully, and destroying resources
|
||||
associated with the transcription process.
|
||||
|
||||
"""
|
||||
logging.info("Cleaning up.")
|
||||
self.exit = True
|
||||
self.transcriber.destroy()
|
||||
|
||||
@@ -1,880 +0,0 @@
|
||||
# original https://github.com/guillaumekln/faster-whisper/blob/master/faster_whisper/transcribe.py
|
||||
|
||||
import itertools
|
||||
import logging
|
||||
import os
|
||||
import zlib
|
||||
import logging
|
||||
|
||||
from typing import BinaryIO, Iterable, List, NamedTuple, Optional, Tuple, Union
|
||||
|
||||
import ctranslate2
|
||||
import numpy as np
|
||||
import tokenizers
|
||||
|
||||
from faster_whisper.audio import decode_audio
|
||||
from faster_whisper.feature_extractor import FeatureExtractor
|
||||
from faster_whisper.tokenizer import Tokenizer
|
||||
from faster_whisper.utils import download_model, format_timestamp
|
||||
from faster_whisper.vad import (
|
||||
SpeechTimestampsMap,
|
||||
collect_chunks,
|
||||
get_speech_timestamps,
|
||||
)
|
||||
|
||||
|
||||
# implement logger not available in faster_whisper==0.4.1
|
||||
def get_logger():
|
||||
"""Returns the module logger."""
|
||||
return logging.getLogger("faster_whisper")
|
||||
|
||||
|
||||
class Word(NamedTuple):
|
||||
start: float
|
||||
end: float
|
||||
word: str
|
||||
probability: float
|
||||
|
||||
|
||||
class Segment(NamedTuple):
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
words: Optional[List[Word]]
|
||||
avg_log_prob: float
|
||||
no_speech_prob: float
|
||||
|
||||
|
||||
class AudioInfo(NamedTuple):
|
||||
language: str
|
||||
language_probability: float
|
||||
duration: float
|
||||
|
||||
|
||||
class TranscriptionOptions(NamedTuple):
|
||||
beam_size: int
|
||||
best_of: int
|
||||
patience: float
|
||||
length_penalty: float
|
||||
log_prob_threshold: Optional[float]
|
||||
no_speech_threshold: Optional[float]
|
||||
compression_ratio_threshold: Optional[float]
|
||||
condition_on_previous_text: bool
|
||||
temperatures: List[float]
|
||||
initial_prompt: Optional[str]
|
||||
prefix: Optional[str]
|
||||
suppress_blank: bool
|
||||
suppress_tokens: Optional[List[int]]
|
||||
without_timestamps: bool
|
||||
max_initial_timestamp: float
|
||||
word_timestamps: bool
|
||||
prepend_punctuations: str
|
||||
append_punctuations: str
|
||||
|
||||
|
||||
class WhisperModel:
|
||||
def __init__(
|
||||
self,
|
||||
model_size_or_path: str,
|
||||
device: str = "auto",
|
||||
device_index: Union[int, List[int]] = 0,
|
||||
compute_type: str = "default",
|
||||
cpu_threads: int = 0,
|
||||
num_workers: int = 1,
|
||||
download_root: Optional[str] = None,
|
||||
local_files_only: bool = True,
|
||||
):
|
||||
"""Initializes the Whisper model.
|
||||
|
||||
Args:
|
||||
model_size_or_path: Size of the model to use (tiny, tiny.en, base, base.en,
|
||||
small, small.en, medium, medium.en, large-v1, or large-v2) or a path to a converted
|
||||
model directory. When a size is configured, the converted model is downloaded
|
||||
from the Hugging Face Hub.
|
||||
device: Device to use for computation ("cpu", "cuda", "auto").
|
||||
device_index: Device ID to use.
|
||||
The model can also be loaded on multiple GPUs by passing a list of IDs
|
||||
(e.g. [0, 1, 2, 3]). In that case, multiple transcriptions can run in parallel
|
||||
when transcribe() is called from multiple Python threads (see also num_workers).
|
||||
compute_type: Type to use for computation.
|
||||
See https://opennmt.net/CTranslate2/quantization.html.
|
||||
cpu_threads: Number of threads to use when running on CPU (4 by default).
|
||||
A non zero value overrides the OMP_NUM_THREADS environment variable.
|
||||
num_workers: When transcribe() is called from multiple Python threads,
|
||||
having multiple workers enables true parallelism when running the model
|
||||
(concurrent calls to self.model.generate() will run in parallel).
|
||||
This can improve the global throughput at the cost of increased memory usage.
|
||||
download_root: Directory where the model should be saved. If not set, the model
|
||||
is saved in the standard Hugging Face cache directory.
|
||||
"""
|
||||
self.logger = get_logger()
|
||||
|
||||
if os.path.isdir(model_size_or_path):
|
||||
model_path = model_size_or_path
|
||||
else:
|
||||
model_path = download_model(
|
||||
model_size_or_path,
|
||||
local_files_only=local_files_only,
|
||||
cache_dir=download_root,
|
||||
)
|
||||
|
||||
self.model = ctranslate2.models.Whisper(
|
||||
model_path,
|
||||
device=device,
|
||||
device_index=device_index,
|
||||
compute_type=compute_type,
|
||||
intra_threads=cpu_threads,
|
||||
inter_threads=num_workers,
|
||||
)
|
||||
|
||||
tokenizer_file = os.path.join(model_path, "tokenizer.json")
|
||||
if os.path.isfile(tokenizer_file):
|
||||
self.hf_tokenizer = tokenizers.Tokenizer.from_file(tokenizer_file)
|
||||
else:
|
||||
self.hf_tokenizer = tokenizers.Tokenizer.from_pretrained(
|
||||
"openai/whisper-tiny" + ("" if self.model.is_multilingual else ".en")
|
||||
)
|
||||
|
||||
self.feature_extractor = FeatureExtractor()
|
||||
self.num_samples_per_token = self.feature_extractor.hop_length * 2
|
||||
self.frames_per_second = (
|
||||
self.feature_extractor.sampling_rate // self.feature_extractor.hop_length
|
||||
)
|
||||
self.tokens_per_second = (
|
||||
self.feature_extractor.sampling_rate // self.num_samples_per_token
|
||||
)
|
||||
self.input_stride = 2
|
||||
self.time_precision = 0.02
|
||||
self.max_length = 448
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
audio: Union[str, BinaryIO, np.ndarray],
|
||||
language: Optional[str] = None,
|
||||
task: str = "transcribe",
|
||||
beam_size: int = 5,
|
||||
best_of: int = 5,
|
||||
patience: float = 1,
|
||||
length_penalty: float = 1,
|
||||
temperature: Union[float, List[float], Tuple[float, ...]] = [
|
||||
0.0,
|
||||
0.2,
|
||||
0.4,
|
||||
0.6,
|
||||
0.8,
|
||||
1.0,
|
||||
],
|
||||
compression_ratio_threshold: Optional[float] = 2.4,
|
||||
log_prob_threshold: Optional[float] = -1.0,
|
||||
no_speech_threshold: Optional[float] = 0.6,
|
||||
condition_on_previous_text: bool = True,
|
||||
initial_prompt: Optional[str] = None,
|
||||
prefix: Optional[str] = None,
|
||||
suppress_blank: bool = True,
|
||||
suppress_tokens: Optional[List[int]] = [-1],
|
||||
without_timestamps: bool = False,
|
||||
max_initial_timestamp: float = 1.0,
|
||||
word_timestamps: bool = False,
|
||||
prepend_punctuations: str = "\"'“¿([{-",
|
||||
append_punctuations: str = "\"'.。,,!!??::”)]}、",
|
||||
vad_filter: bool = False,
|
||||
vad_parameters: Optional[dict] = None,
|
||||
) -> Tuple[Iterable[Segment], AudioInfo]:
|
||||
"""Transcribes an input file.
|
||||
|
||||
Arguments:
|
||||
audio: Path to the input file (or a file-like object), or the audio waveform.
|
||||
language: The language spoken in the audio. It should be a language code such
|
||||
as "en" or "fr". If not set, the language will be detected in the first 30 seconds
|
||||
of audio.
|
||||
task: Task to execute (transcribe or translate).
|
||||
beam_size: Beam size to use for decoding.
|
||||
best_of: Number of candidates when sampling with non-zero temperature.
|
||||
patience: Beam search patience factor.
|
||||
length_penalty: Exponential length penalty constant.
|
||||
temperature: Temperature for sampling. It can be a tuple of temperatures,
|
||||
which will be successively used upon failures according to either
|
||||
`compression_ratio_threshold` or `log_prob_threshold`.
|
||||
compression_ratio_threshold: If the gzip compression ratio is above this value,
|
||||
treat as failed.
|
||||
log_prob_threshold: If the average log probability over sampled tokens is
|
||||
below this value, treat as failed.
|
||||
no_speech_threshold: If the no_speech probability is higher than this value AND
|
||||
the average log probability over sampled tokens is below `log_prob_threshold`,
|
||||
consider the segment as silent.
|
||||
condition_on_previous_text: If True, the previous output of the model is provided
|
||||
as a prompt for the next window; disabling may make the text inconsistent across
|
||||
windows, but the model becomes less prone to getting stuck in a failure loop,
|
||||
such as repetition looping or timestamps going out of sync.
|
||||
initial_prompt: Optional text to provide as a prompt for the first window.
|
||||
prefix: Optional text to provide as a prefix for the first window.
|
||||
suppress_blank: Suppress blank outputs at the beginning of the sampling.
|
||||
suppress_tokens: List of token IDs to suppress. -1 will suppress a default set
|
||||
of symbols as defined in the model config.json file.
|
||||
without_timestamps: Only sample text tokens.
|
||||
max_initial_timestamp: The initial timestamp cannot be later than this.
|
||||
word_timestamps: Extract word-level timestamps using the cross-attention pattern
|
||||
and dynamic time warping, and include the timestamps for each word in each segment.
|
||||
prepend_punctuations: If word_timestamps is True, merge these punctuation symbols
|
||||
with the next word
|
||||
append_punctuations: If word_timestamps is True, merge these punctuation symbols
|
||||
with the previous word
|
||||
vad_filter: Enable the voice activity detection (VAD) to filter out parts of the audio
|
||||
without speech. This step is using the Silero VAD model
|
||||
https://github.com/snakers4/silero-vad.
|
||||
vad_parameters: Dictionary of Silero VAD parameters (see available parameters and
|
||||
default values in the function `get_speech_timestamps`).
|
||||
|
||||
Returns:
|
||||
A tuple with:
|
||||
|
||||
- a generator over transcribed segments
|
||||
- an instance of AudioInfo
|
||||
"""
|
||||
sampling_rate = self.feature_extractor.sampling_rate
|
||||
|
||||
if not isinstance(audio, np.ndarray):
|
||||
audio = decode_audio(audio, sampling_rate=sampling_rate)
|
||||
|
||||
duration = audio.shape[0] / sampling_rate
|
||||
|
||||
self.logger.info(
|
||||
"Processing audio with duration %s", format_timestamp(duration)
|
||||
)
|
||||
|
||||
if vad_filter:
|
||||
vad_parameters = {} if vad_parameters is None else vad_parameters
|
||||
speech_chunks = get_speech_timestamps(audio, **vad_parameters)
|
||||
audio = collect_chunks(audio, speech_chunks)
|
||||
|
||||
self.logger.info(
|
||||
"VAD filter removed %s of audio",
|
||||
format_timestamp(duration - (audio.shape[0] / sampling_rate)),
|
||||
)
|
||||
|
||||
if self.logger.isEnabledFor(logging.DEBUG):
|
||||
self.logger.debug(
|
||||
"VAD filter kept the following audio segments: %s",
|
||||
", ".join(
|
||||
"[%s -> %s]"
|
||||
% (
|
||||
format_timestamp(chunk["start"] / sampling_rate),
|
||||
format_timestamp(chunk["end"] / sampling_rate),
|
||||
)
|
||||
for chunk in speech_chunks
|
||||
),
|
||||
)
|
||||
|
||||
else:
|
||||
speech_chunks = None
|
||||
|
||||
features = self.feature_extractor(audio)
|
||||
|
||||
encoder_output = None
|
||||
|
||||
if language is None:
|
||||
if not self.model.is_multilingual:
|
||||
language = "en"
|
||||
language_probability = 1
|
||||
else:
|
||||
segment = features[:, : self.feature_extractor.nb_max_frames]
|
||||
encoder_output = self.encode(segment)
|
||||
results = self.model.detect_language(encoder_output)
|
||||
language_token, language_probability = results[0][0]
|
||||
language = language_token[2:-2]
|
||||
|
||||
self.logger.info(
|
||||
"Detected language '%s' with probability %.2f",
|
||||
language,
|
||||
language_probability,
|
||||
)
|
||||
return language, language_probability
|
||||
else:
|
||||
language_probability = 1
|
||||
|
||||
tokenizer = Tokenizer(
|
||||
self.hf_tokenizer,
|
||||
self.model.is_multilingual,
|
||||
task=task,
|
||||
language=language,
|
||||
)
|
||||
|
||||
options = TranscriptionOptions(
|
||||
beam_size=beam_size,
|
||||
best_of=best_of,
|
||||
patience=patience,
|
||||
length_penalty=length_penalty,
|
||||
log_prob_threshold=log_prob_threshold,
|
||||
no_speech_threshold=no_speech_threshold,
|
||||
compression_ratio_threshold=compression_ratio_threshold,
|
||||
condition_on_previous_text=condition_on_previous_text,
|
||||
temperatures=(
|
||||
temperature if isinstance(temperature, (list, tuple)) else [temperature]
|
||||
),
|
||||
initial_prompt=initial_prompt,
|
||||
prefix=prefix,
|
||||
suppress_blank=suppress_blank,
|
||||
suppress_tokens=get_suppressed_tokens(tokenizer, suppress_tokens),
|
||||
without_timestamps=without_timestamps,
|
||||
max_initial_timestamp=max_initial_timestamp,
|
||||
word_timestamps=word_timestamps,
|
||||
prepend_punctuations=prepend_punctuations,
|
||||
append_punctuations=append_punctuations,
|
||||
)
|
||||
|
||||
segments = self.generate_segments(features, tokenizer, options, encoder_output)
|
||||
|
||||
if speech_chunks:
|
||||
segments = restore_speech_timestamps(segments, speech_chunks, sampling_rate)
|
||||
|
||||
audio_info = AudioInfo(
|
||||
language=language,
|
||||
language_probability=language_probability,
|
||||
duration=duration,
|
||||
)
|
||||
|
||||
return segments
|
||||
|
||||
def generate_segments(
|
||||
self,
|
||||
features: np.ndarray,
|
||||
tokenizer: Tokenizer,
|
||||
options: TranscriptionOptions,
|
||||
encoder_output: Optional[ctranslate2.StorageView] = None,
|
||||
) -> Iterable[Segment]:
|
||||
content_frames = features.shape[-1] - self.feature_extractor.nb_max_frames
|
||||
seek = 0
|
||||
all_tokens = []
|
||||
prompt_reset_since = 0
|
||||
|
||||
if options.initial_prompt is not None:
|
||||
initial_prompt = " " + options.initial_prompt.strip()
|
||||
initial_prompt_tokens = tokenizer.encode(initial_prompt)
|
||||
all_tokens.extend(initial_prompt_tokens)
|
||||
all_segments = []
|
||||
while seek < content_frames:
|
||||
time_offset = seek * self.feature_extractor.time_per_frame
|
||||
segment = features[:, seek : seek + self.feature_extractor.nb_max_frames]
|
||||
segment_size = min(
|
||||
self.feature_extractor.nb_max_frames, content_frames - seek
|
||||
)
|
||||
segment_duration = segment_size * self.feature_extractor.time_per_frame
|
||||
|
||||
if self.logger.isEnabledFor(logging.DEBUG):
|
||||
self.logger.debug(
|
||||
"Processing segment at %s", format_timestamp(time_offset)
|
||||
)
|
||||
|
||||
previous_tokens = all_tokens[prompt_reset_since:]
|
||||
prompt = self.get_prompt(
|
||||
tokenizer,
|
||||
previous_tokens,
|
||||
without_timestamps=options.without_timestamps,
|
||||
prefix=options.prefix if seek == 0 else None,
|
||||
)
|
||||
|
||||
if encoder_output is None:
|
||||
encoder_output = self.encode(segment)
|
||||
|
||||
result, avg_log_prob, temperature = self.generate_with_fallback(
|
||||
encoder_output, prompt, tokenizer, options
|
||||
)
|
||||
|
||||
if options.no_speech_threshold is not None:
|
||||
# no voice activity check
|
||||
should_skip = result.no_speech_prob > options.no_speech_threshold
|
||||
|
||||
if (
|
||||
options.log_prob_threshold is not None
|
||||
and avg_log_prob > options.log_prob_threshold
|
||||
):
|
||||
# don't skip if the logprob is high enough, despite the no_speech_prob
|
||||
should_skip = False
|
||||
|
||||
if should_skip:
|
||||
self.logger.debug(
|
||||
"No speech threshold is met (%f > %f)",
|
||||
result.no_speech_prob,
|
||||
options.no_speech_threshold,
|
||||
)
|
||||
|
||||
# fast-forward to the next segment boundary
|
||||
seek += segment_size
|
||||
continue
|
||||
|
||||
tokens = result.sequences_ids[0]
|
||||
|
||||
previous_seek = seek
|
||||
current_segments = []
|
||||
|
||||
single_timestamp_ending = (
|
||||
len(tokens) >= 2
|
||||
and tokens[-2] < tokenizer.timestamp_begin
|
||||
and tokens[-1] >= tokenizer.timestamp_begin
|
||||
)
|
||||
|
||||
consecutive_timestamps = [
|
||||
i
|
||||
for i in range(len(tokens))
|
||||
if i > 0
|
||||
and tokens[i] >= tokenizer.timestamp_begin
|
||||
and tokens[i - 1] >= tokenizer.timestamp_begin
|
||||
]
|
||||
|
||||
if len(consecutive_timestamps) > 0:
|
||||
slices = list(consecutive_timestamps)
|
||||
if single_timestamp_ending:
|
||||
slices.append(len(tokens))
|
||||
|
||||
last_slice = 0
|
||||
for current_slice in slices:
|
||||
sliced_tokens = tokens[last_slice:current_slice]
|
||||
start_timestamp_position = (
|
||||
sliced_tokens[0] - tokenizer.timestamp_begin
|
||||
)
|
||||
end_timestamp_position = (
|
||||
sliced_tokens[-1] - tokenizer.timestamp_begin
|
||||
)
|
||||
start_time = (
|
||||
time_offset + start_timestamp_position * self.time_precision
|
||||
)
|
||||
end_time = (
|
||||
time_offset + end_timestamp_position * self.time_precision
|
||||
)
|
||||
|
||||
current_segments.append(
|
||||
dict(
|
||||
seek=seek,
|
||||
start=start_time,
|
||||
end=end_time,
|
||||
tokens=sliced_tokens,
|
||||
)
|
||||
)
|
||||
last_slice = current_slice
|
||||
|
||||
if single_timestamp_ending:
|
||||
# single timestamp at the end means no speech after the last timestamp.
|
||||
seek += segment_size
|
||||
else:
|
||||
# otherwise, ignore the unfinished segment and seek to the last timestamp
|
||||
last_timestamp_position = (
|
||||
tokens[last_slice - 1] - tokenizer.timestamp_begin
|
||||
)
|
||||
seek += last_timestamp_position * self.input_stride
|
||||
|
||||
else:
|
||||
duration = segment_duration
|
||||
timestamps = [
|
||||
token for token in tokens if token >= tokenizer.timestamp_begin
|
||||
]
|
||||
if len(timestamps) > 0 and timestamps[-1] != tokenizer.timestamp_begin:
|
||||
last_timestamp_position = timestamps[-1] - tokenizer.timestamp_begin
|
||||
duration = last_timestamp_position * self.time_precision
|
||||
|
||||
current_segments.append(
|
||||
dict(
|
||||
seek=seek,
|
||||
start=time_offset,
|
||||
end=time_offset + duration,
|
||||
tokens=tokens,
|
||||
)
|
||||
)
|
||||
|
||||
seek += segment_size
|
||||
|
||||
if not options.condition_on_previous_text or temperature > 0.5:
|
||||
prompt_reset_since = len(all_tokens)
|
||||
|
||||
if options.word_timestamps:
|
||||
self.add_word_timestamps(
|
||||
current_segments,
|
||||
tokenizer,
|
||||
encoder_output,
|
||||
segment_size,
|
||||
options.prepend_punctuations,
|
||||
options.append_punctuations,
|
||||
)
|
||||
|
||||
word_end_timestamps = [
|
||||
w["end"] for s in current_segments for w in s["words"]
|
||||
]
|
||||
|
||||
if not single_timestamp_ending and len(word_end_timestamps) > 0:
|
||||
seek_shift = round(
|
||||
(word_end_timestamps[-1] - time_offset) * self.frames_per_second
|
||||
)
|
||||
|
||||
if seek_shift > 0:
|
||||
seek = previous_seek + seek_shift
|
||||
|
||||
encoder_output = None
|
||||
|
||||
for segment in current_segments:
|
||||
tokens = segment["tokens"]
|
||||
text = tokenizer.decode(tokens)
|
||||
|
||||
if segment["start"] == segment["end"] or not text.strip():
|
||||
continue
|
||||
|
||||
all_tokens.extend(tokens)
|
||||
|
||||
all_segments.append(Segment(
|
||||
start=segment["start"],
|
||||
end=segment["end"],
|
||||
text=text,
|
||||
words=(
|
||||
[Word(**word) for word in segment["words"]]
|
||||
if options.word_timestamps
|
||||
else None
|
||||
),
|
||||
avg_log_prob=avg_log_prob,
|
||||
no_speech_prob=result.no_speech_prob,
|
||||
))
|
||||
return all_segments
|
||||
|
||||
def encode(self, features: np.ndarray) -> ctranslate2.StorageView:
|
||||
# When the model is running on multiple GPUs, the encoder output should be moved
|
||||
# to the CPU since we don't know which GPU will handle the next job.
|
||||
to_cpu = self.model.device == "cuda" and len(self.model.device_index) > 1
|
||||
|
||||
features = np.expand_dims(features, 0)
|
||||
features = get_ctranslate2_storage(features)
|
||||
|
||||
return self.model.encode(features, to_cpu=to_cpu)
|
||||
|
||||
def generate_with_fallback(
|
||||
self,
|
||||
encoder_output: ctranslate2.StorageView,
|
||||
prompt: List[int],
|
||||
tokenizer: Tokenizer,
|
||||
options: TranscriptionOptions,
|
||||
) -> Tuple[ctranslate2.models.WhisperGenerationResult, float, float]:
|
||||
result = None
|
||||
avg_log_prob = None
|
||||
final_temperature = None
|
||||
|
||||
max_initial_timestamp_index = int(
|
||||
round(options.max_initial_timestamp / self.time_precision)
|
||||
)
|
||||
|
||||
for temperature in options.temperatures:
|
||||
if temperature > 0:
|
||||
kwargs = {
|
||||
"beam_size": 1,
|
||||
"num_hypotheses": options.best_of,
|
||||
"sampling_topk": 0,
|
||||
"sampling_temperature": temperature,
|
||||
}
|
||||
else:
|
||||
kwargs = {
|
||||
"beam_size": options.beam_size,
|
||||
"patience": options.patience,
|
||||
}
|
||||
|
||||
final_temperature = temperature
|
||||
result = self.model.generate(
|
||||
encoder_output,
|
||||
[prompt],
|
||||
length_penalty=options.length_penalty,
|
||||
max_length=self.max_length,
|
||||
return_scores=True,
|
||||
return_no_speech_prob=True,
|
||||
suppress_blank=options.suppress_blank,
|
||||
suppress_tokens=options.suppress_tokens,
|
||||
max_initial_timestamp_index=max_initial_timestamp_index,
|
||||
**kwargs,
|
||||
)[0]
|
||||
|
||||
tokens = result.sequences_ids[0]
|
||||
|
||||
# Recover the average log prob from the returned score.
|
||||
seq_len = len(tokens)
|
||||
cum_log_prob = result.scores[0] * (seq_len**options.length_penalty)
|
||||
avg_log_prob = cum_log_prob / (seq_len + 1)
|
||||
|
||||
text = tokenizer.decode(tokens).strip()
|
||||
compression_ratio = get_compression_ratio(text)
|
||||
|
||||
needs_fallback = False
|
||||
|
||||
if (
|
||||
options.compression_ratio_threshold is not None
|
||||
and compression_ratio > options.compression_ratio_threshold
|
||||
):
|
||||
needs_fallback = True # too repetitive
|
||||
|
||||
self.logger.debug(
|
||||
"Compression ratio threshold is not met with temperature %.1f (%f > %f)",
|
||||
temperature,
|
||||
compression_ratio,
|
||||
options.compression_ratio_threshold,
|
||||
)
|
||||
|
||||
if (
|
||||
options.log_prob_threshold is not None
|
||||
and avg_log_prob < options.log_prob_threshold
|
||||
):
|
||||
needs_fallback = True # average log probability is too low
|
||||
|
||||
self.logger.debug(
|
||||
"Log probability threshold is not met with temperature %.1f (%f < %f)",
|
||||
temperature,
|
||||
avg_log_prob,
|
||||
options.log_prob_threshold,
|
||||
)
|
||||
|
||||
if not needs_fallback:
|
||||
break
|
||||
|
||||
return result, avg_log_prob, final_temperature
|
||||
|
||||
def get_prompt(
|
||||
self,
|
||||
tokenizer: Tokenizer,
|
||||
previous_tokens: List[int],
|
||||
without_timestamps: bool = False,
|
||||
prefix: Optional[str] = None,
|
||||
) -> List[int]:
|
||||
prompt = []
|
||||
|
||||
if previous_tokens:
|
||||
prompt.append(tokenizer.sot_prev)
|
||||
prompt.extend(previous_tokens[-(self.max_length // 2 - 1) :])
|
||||
|
||||
prompt.extend(tokenizer.sot_sequence)
|
||||
|
||||
if without_timestamps:
|
||||
prompt.append(tokenizer.no_timestamps)
|
||||
|
||||
if prefix:
|
||||
prefix_tokens = tokenizer.encode(" " + prefix.strip())
|
||||
if len(prefix_tokens) >= self.max_length // 2:
|
||||
prefix_tokens = prefix_tokens[: self.max_length // 2 - 1]
|
||||
prompt.extend(prefix_tokens)
|
||||
|
||||
return prompt
|
||||
|
||||
def add_word_timestamps(
|
||||
self,
|
||||
segments: List[dict],
|
||||
tokenizer: Tokenizer,
|
||||
encoder_output: ctranslate2.StorageView,
|
||||
num_frames: int,
|
||||
prepend_punctuations: str,
|
||||
append_punctuations: str,
|
||||
):
|
||||
if len(segments) == 0:
|
||||
return
|
||||
|
||||
text_tokens_per_segment = [
|
||||
[token for token in segment["tokens"] if token < tokenizer.eot]
|
||||
for segment in segments
|
||||
]
|
||||
|
||||
text_tokens = list(itertools.chain.from_iterable(text_tokens_per_segment))
|
||||
alignment = self.find_alignment(
|
||||
tokenizer, text_tokens, encoder_output, num_frames
|
||||
)
|
||||
merge_punctuations(alignment, prepend_punctuations, append_punctuations)
|
||||
|
||||
time_offset = (
|
||||
segments[0]["seek"]
|
||||
* self.feature_extractor.hop_length
|
||||
/ self.feature_extractor.sampling_rate
|
||||
)
|
||||
|
||||
word_index = 0
|
||||
|
||||
for segment, text_tokens in zip(segments, text_tokens_per_segment):
|
||||
saved_tokens = 0
|
||||
words = []
|
||||
|
||||
while word_index < len(alignment) and saved_tokens < len(text_tokens):
|
||||
timing = alignment[word_index]
|
||||
|
||||
if timing["word"]:
|
||||
words.append(
|
||||
dict(
|
||||
word=timing["word"],
|
||||
start=round(time_offset + timing["start"], 2),
|
||||
end=round(time_offset + timing["end"], 2),
|
||||
probability=timing["probability"],
|
||||
)
|
||||
)
|
||||
|
||||
saved_tokens += len(timing["tokens"])
|
||||
word_index += 1
|
||||
|
||||
if len(words) > 0:
|
||||
# adjust the segment-level timestamps based on the word-level timestamps
|
||||
segment["start"] = words[0]["start"]
|
||||
segment["end"] = words[-1]["end"]
|
||||
|
||||
segment["words"] = words
|
||||
|
||||
def find_alignment(
|
||||
self,
|
||||
tokenizer: Tokenizer,
|
||||
text_tokens: List[int],
|
||||
encoder_output: ctranslate2.StorageView,
|
||||
num_frames: int,
|
||||
median_filter_width: int = 7,
|
||||
) -> List[dict]:
|
||||
if len(text_tokens) == 0:
|
||||
return []
|
||||
|
||||
result = self.model.align(
|
||||
encoder_output,
|
||||
tokenizer.sot_sequence,
|
||||
[text_tokens],
|
||||
num_frames,
|
||||
median_filter_width=median_filter_width,
|
||||
)[0]
|
||||
|
||||
text_token_probs = result.text_token_probs
|
||||
|
||||
alignments = result.alignments
|
||||
text_indices = np.array([pair[0] for pair in alignments])
|
||||
time_indices = np.array([pair[1] for pair in alignments])
|
||||
|
||||
words, word_tokens = tokenizer.split_to_word_tokens(
|
||||
text_tokens + [tokenizer.eot]
|
||||
)
|
||||
word_boundaries = np.pad(np.cumsum([len(t) for t in word_tokens[:-1]]), (1, 0))
|
||||
|
||||
jumps = np.pad(np.diff(text_indices), (1, 0), constant_values=1).astype(bool)
|
||||
jump_times = time_indices[jumps] / self.tokens_per_second
|
||||
start_times = jump_times[word_boundaries[:-1]]
|
||||
end_times = jump_times[word_boundaries[1:]]
|
||||
word_probabilities = [
|
||||
np.mean(text_token_probs[i:j])
|
||||
for i, j in zip(word_boundaries[:-1], word_boundaries[1:])
|
||||
]
|
||||
|
||||
# hack: ensure the first and second word is not longer than twice the median word duration.
|
||||
# a better segmentation algorithm based on VAD should be able to replace this.
|
||||
word_durations = end_times - start_times
|
||||
word_durations = word_durations[word_durations.nonzero()]
|
||||
if len(word_durations) > 0:
|
||||
median_duration = np.median(word_durations)
|
||||
max_duration = median_duration * 2
|
||||
if len(word_durations) >= 2 and word_durations[1] > max_duration:
|
||||
boundary = max(end_times[2] / 2, end_times[2] - max_duration)
|
||||
end_times[0] = start_times[1] = boundary
|
||||
if (
|
||||
len(word_durations) >= 1
|
||||
and end_times[0] - start_times[0] > max_duration
|
||||
):
|
||||
start_times[0] = max(0, end_times[0] - max_duration)
|
||||
|
||||
return [
|
||||
dict(
|
||||
word=word, tokens=tokens, start=start, end=end, probability=probability
|
||||
)
|
||||
for word, tokens, start, end, probability in zip(
|
||||
words, word_tokens, start_times, end_times, word_probabilities
|
||||
)
|
||||
]
|
||||
|
||||
def destroy(self):
|
||||
del self.model
|
||||
|
||||
|
||||
def restore_speech_timestamps(
|
||||
segments: Iterable[Segment],
|
||||
speech_chunks: List[dict],
|
||||
sampling_rate: int,
|
||||
) -> Iterable[Segment]:
|
||||
ts_map = SpeechTimestampsMap(speech_chunks, sampling_rate)
|
||||
|
||||
for segment in segments:
|
||||
if segment.words:
|
||||
words = []
|
||||
for word in segment.words:
|
||||
# Ensure the word start and end times are resolved to the same chunk.
|
||||
chunk_index = ts_map.get_chunk_index(word.start)
|
||||
word = word._replace(
|
||||
start=ts_map.get_original_time(word.start, chunk_index),
|
||||
end=ts_map.get_original_time(word.end, chunk_index),
|
||||
)
|
||||
words.append(word)
|
||||
|
||||
segment = segment._replace(
|
||||
start=words[0].start,
|
||||
end=words[-1].end,
|
||||
words=words,
|
||||
)
|
||||
|
||||
else:
|
||||
segment = segment._replace(
|
||||
start=ts_map.get_original_time(segment.start),
|
||||
end=ts_map.get_original_time(segment.end),
|
||||
)
|
||||
|
||||
yield segment
|
||||
|
||||
|
||||
def get_ctranslate2_storage(segment: np.ndarray) -> ctranslate2.StorageView:
|
||||
segment = np.ascontiguousarray(segment)
|
||||
segment = ctranslate2.StorageView.from_array(segment)
|
||||
return segment
|
||||
|
||||
|
||||
def get_compression_ratio(text: str) -> float:
|
||||
text_bytes = text.encode("utf-8")
|
||||
return len(text_bytes) / len(zlib.compress(text_bytes))
|
||||
|
||||
|
||||
def get_suppressed_tokens(tokenizer, suppress_tokens):
|
||||
if not suppress_tokens or -1 in suppress_tokens:
|
||||
return suppress_tokens
|
||||
|
||||
suppress_tokens = list(suppress_tokens)
|
||||
|
||||
# Ensure the following special tokens are suppressed when the user does
|
||||
# not use the default set (-1).
|
||||
suppress_tokens.extend(
|
||||
[
|
||||
tokenizer.transcribe,
|
||||
tokenizer.translate,
|
||||
tokenizer.sot,
|
||||
tokenizer.sot_prev,
|
||||
tokenizer.sot_lm,
|
||||
]
|
||||
)
|
||||
|
||||
return sorted(set(suppress_tokens))
|
||||
|
||||
|
||||
def merge_punctuations(alignment: List[dict], prepended: str, appended: str):
|
||||
# merge prepended punctuations
|
||||
i = len(alignment) - 2
|
||||
j = len(alignment) - 1
|
||||
while i >= 0:
|
||||
previous = alignment[i]
|
||||
following = alignment[j]
|
||||
if previous["word"].startswith(" ") and previous["word"].strip() in prepended:
|
||||
# prepend it to the following word
|
||||
following["word"] = previous["word"] + following["word"]
|
||||
following["tokens"] = previous["tokens"] + following["tokens"]
|
||||
previous["word"] = ""
|
||||
previous["tokens"] = []
|
||||
else:
|
||||
j = i
|
||||
i -= 1
|
||||
|
||||
# merge appended punctuations
|
||||
i = 0
|
||||
j = 1
|
||||
while j < len(alignment):
|
||||
previous = alignment[i]
|
||||
following = alignment[j]
|
||||
if not previous["word"].endswith(" ") and following["word"] in appended:
|
||||
# append it to the previous word
|
||||
previous["word"] = previous["word"] + following["word"]
|
||||
previous["tokens"] = previous["tokens"] + following["tokens"]
|
||||
following["word"] = ""
|
||||
following["tokens"] = []
|
||||
else:
|
||||
i = j
|
||||
j += 1
|
||||
@@ -0,0 +1,364 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2022-2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import logging
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from subprocess import CalledProcessError, run
|
||||
from typing import Dict, Iterable, List, Optional, TextIO, Tuple, Union
|
||||
|
||||
import kaldialign
|
||||
import numpy as np
|
||||
import soundfile
|
||||
import av
|
||||
import wave
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from whisper_live.utils import resample
|
||||
|
||||
|
||||
Pathlike = Union[str, Path]
|
||||
|
||||
SAMPLE_RATE = 16000
|
||||
N_FFT = 400
|
||||
HOP_LENGTH = 160
|
||||
CHUNK_LENGTH = 30
|
||||
N_SAMPLES = CHUNK_LENGTH * SAMPLE_RATE # 480000 samples in a 30-second chunk
|
||||
|
||||
|
||||
def load_audio(file: str, sr: int = 16000):
|
||||
"""
|
||||
Open an audio file, resample it, and read as a mono waveform.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
file: str
|
||||
The audio file to open.
|
||||
|
||||
sr: int
|
||||
The sample rate to resample the audio if necessary.
|
||||
|
||||
Returns
|
||||
-------
|
||||
A NumPy array containing the audio waveform, in float32 dtype.
|
||||
"""
|
||||
resampled_file = resample(file, sr)
|
||||
|
||||
with wave.open(resampled_file, "rb") as wav_file:
|
||||
num_frames = wav_file.getnframes()
|
||||
raw_data = wav_file.readframes(num_frames)
|
||||
|
||||
audio_data = np.frombuffer(raw_data, dtype=np.int16)
|
||||
|
||||
audio_data = audio_data.astype(np.float32) / 32768.0
|
||||
|
||||
return audio_data
|
||||
|
||||
|
||||
def load_audio_wav_format(wav_path):
|
||||
# make sure audio in .wav format
|
||||
assert wav_path.endswith(
|
||||
'.wav'), f"Only support .wav format, but got {wav_path}"
|
||||
waveform, sample_rate = soundfile.read(wav_path)
|
||||
assert sample_rate == 16000, f"Only support 16k sample rate, but got {sample_rate}"
|
||||
return waveform, sample_rate
|
||||
|
||||
|
||||
def pad_or_trim(array, length: int = N_SAMPLES, *, axis: int = -1):
|
||||
"""
|
||||
Pad or trim the audio array to N_SAMPLES, as expected by the encoder.
|
||||
"""
|
||||
if torch.is_tensor(array):
|
||||
if array.shape[axis] > length:
|
||||
array = array.index_select(dim=axis,
|
||||
index=torch.arange(length,
|
||||
device=array.device))
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = F.pad(array,
|
||||
[pad for sizes in pad_widths[::-1] for pad in sizes])
|
||||
else:
|
||||
if array.shape[axis] > length:
|
||||
array = array.take(indices=range(length), axis=axis)
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = np.pad(array, pad_widths)
|
||||
|
||||
return array
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def mel_filters(device,
|
||||
n_mels: int,
|
||||
mel_filters_dir: str = None) -> torch.Tensor:
|
||||
"""
|
||||
load the mel filterbank matrix for projecting STFT into a Mel spectrogram.
|
||||
Allows decoupling librosa dependency; saved using:
|
||||
|
||||
np.savez_compressed(
|
||||
"mel_filters.npz",
|
||||
mel_80=librosa.filters.mel(sr=16000, n_fft=400, n_mels=80),
|
||||
)
|
||||
"""
|
||||
assert n_mels in {80, 128}, f"Unsupported n_mels: {n_mels}"
|
||||
if mel_filters_dir is None:
|
||||
mel_filters_path = os.path.join(os.path.dirname(__file__), "assets",
|
||||
"mel_filters.npz")
|
||||
else:
|
||||
mel_filters_path = os.path.join(mel_filters_dir, "mel_filters.npz")
|
||||
with np.load(mel_filters_path) as f:
|
||||
return torch.from_numpy(f[f"mel_{n_mels}"]).to(device)
|
||||
|
||||
|
||||
def log_mel_spectrogram(
|
||||
audio: Union[str, np.ndarray, torch.Tensor],
|
||||
n_mels: int,
|
||||
padding: int = 0,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
return_duration: bool = False,
|
||||
mel_filters_dir: str = None,
|
||||
):
|
||||
"""
|
||||
Compute the log-Mel spectrogram of
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio: Union[str, np.ndarray, torch.Tensor], shape = (*)
|
||||
The path to audio or either a NumPy array or Tensor containing the audio waveform in 16 kHz
|
||||
|
||||
n_mels: int
|
||||
The number of Mel-frequency filters, only 80 and 128 are supported
|
||||
|
||||
padding: int
|
||||
Number of zero samples to pad to the right
|
||||
|
||||
device: Optional[Union[str, torch.device]]
|
||||
If given, the audio tensor is moved to this device before STFT
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor, shape = (80 or 128, n_frames)
|
||||
A Tensor that contains the Mel spectrogram
|
||||
"""
|
||||
if not torch.is_tensor(audio):
|
||||
if isinstance(audio, str):
|
||||
if audio.endswith('.wav'):
|
||||
audio, _ = load_audio_wav_format(audio)
|
||||
else:
|
||||
audio = load_audio(audio)
|
||||
assert isinstance(audio,
|
||||
np.ndarray), f"Unsupported audio type: {type(audio)}"
|
||||
duration = audio.shape[-1] / SAMPLE_RATE
|
||||
audio = pad_or_trim(audio, N_SAMPLES)
|
||||
audio = audio.astype(np.float32)
|
||||
audio = torch.from_numpy(audio)
|
||||
|
||||
if device is not None:
|
||||
audio = audio.to(device)
|
||||
if padding > 0:
|
||||
audio = F.pad(audio, (0, padding))
|
||||
window = torch.hann_window(N_FFT).to(audio.device)
|
||||
stft = torch.stft(audio,
|
||||
N_FFT,
|
||||
HOP_LENGTH,
|
||||
window=window,
|
||||
return_complex=True)
|
||||
magnitudes = stft[..., :-1].abs()**2
|
||||
|
||||
filters = mel_filters(audio.device, n_mels, mel_filters_dir)
|
||||
mel_spec = filters @ magnitudes
|
||||
|
||||
log_spec = torch.clamp(mel_spec, min=1e-10).log10()
|
||||
log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)
|
||||
log_spec = (log_spec + 4.0) / 4.0
|
||||
if return_duration:
|
||||
return log_spec, duration
|
||||
else:
|
||||
return log_spec
|
||||
|
||||
|
||||
def store_transcripts(filename: Pathlike, texts: Iterable[Tuple[str, str,
|
||||
str]]) -> None:
|
||||
"""Save predicted results and reference transcripts to a file.
|
||||
https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py
|
||||
Args:
|
||||
filename:
|
||||
File to save the results to.
|
||||
texts:
|
||||
An iterable of tuples. The first element is the cur_id, the second is
|
||||
the reference transcript and the third element is the predicted result.
|
||||
Returns:
|
||||
Return None.
|
||||
"""
|
||||
with open(filename, "w") as f:
|
||||
for cut_id, ref, hyp in texts:
|
||||
print(f"{cut_id}:\tref={ref}", file=f)
|
||||
print(f"{cut_id}:\thyp={hyp}", file=f)
|
||||
|
||||
|
||||
def write_error_stats( # noqa: C901
|
||||
f: TextIO,
|
||||
test_set_name: str,
|
||||
results: List[Tuple[str, str]],
|
||||
enable_log: bool = True,
|
||||
) -> float:
|
||||
"""Write statistics based on predicted results and reference transcripts.
|
||||
https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py
|
||||
It will write the following to the given file:
|
||||
|
||||
- WER
|
||||
- number of insertions, deletions, substitutions, corrects and total
|
||||
reference words. For example::
|
||||
|
||||
Errors: 23 insertions, 57 deletions, 212 substitutions, over 2606
|
||||
reference words (2337 correct)
|
||||
|
||||
- The difference between the reference transcript and predicted result.
|
||||
An instance is given below::
|
||||
|
||||
THE ASSOCIATION OF (EDISON->ADDISON) ILLUMINATING COMPANIES
|
||||
|
||||
The above example shows that the reference word is `EDISON`,
|
||||
but it is predicted to `ADDISON` (a substitution error).
|
||||
|
||||
Another example is::
|
||||
|
||||
FOR THE FIRST DAY (SIR->*) I THINK
|
||||
|
||||
The reference word `SIR` is missing in the predicted
|
||||
results (a deletion error).
|
||||
results:
|
||||
An iterable of tuples. The first element is the cur_id, the second is
|
||||
the reference transcript and the third element is the predicted result.
|
||||
enable_log:
|
||||
If True, also print detailed WER to the console.
|
||||
Otherwise, it is written only to the given file.
|
||||
Returns:
|
||||
Return None.
|
||||
"""
|
||||
subs: Dict[Tuple[str, str], int] = defaultdict(int)
|
||||
ins: Dict[str, int] = defaultdict(int)
|
||||
dels: Dict[str, int] = defaultdict(int)
|
||||
|
||||
# `words` stores counts per word, as follows:
|
||||
# corr, ref_sub, hyp_sub, ins, dels
|
||||
words: Dict[str, List[int]] = defaultdict(lambda: [0, 0, 0, 0, 0])
|
||||
num_corr = 0
|
||||
ERR = "*"
|
||||
for cut_id, ref, hyp in results:
|
||||
ali = kaldialign.align(ref, hyp, ERR)
|
||||
for ref_word, hyp_word in ali:
|
||||
if ref_word == ERR:
|
||||
ins[hyp_word] += 1
|
||||
words[hyp_word][3] += 1
|
||||
elif hyp_word == ERR:
|
||||
dels[ref_word] += 1
|
||||
words[ref_word][4] += 1
|
||||
elif hyp_word != ref_word:
|
||||
subs[(ref_word, hyp_word)] += 1
|
||||
words[ref_word][1] += 1
|
||||
words[hyp_word][2] += 1
|
||||
else:
|
||||
words[ref_word][0] += 1
|
||||
num_corr += 1
|
||||
ref_len = sum([len(r) for _, r, _ in results])
|
||||
sub_errs = sum(subs.values())
|
||||
ins_errs = sum(ins.values())
|
||||
del_errs = sum(dels.values())
|
||||
tot_errs = sub_errs + ins_errs + del_errs
|
||||
tot_err_rate = "%.2f" % (100.0 * tot_errs / ref_len)
|
||||
|
||||
if enable_log:
|
||||
logging.info(f"[{test_set_name}] %WER {tot_errs / ref_len:.2%} "
|
||||
f"[{tot_errs} / {ref_len}, {ins_errs} ins, "
|
||||
f"{del_errs} del, {sub_errs} sub ]")
|
||||
|
||||
print(f"%WER = {tot_err_rate}", file=f)
|
||||
print(
|
||||
f"Errors: {ins_errs} insertions, {del_errs} deletions, "
|
||||
f"{sub_errs} substitutions, over {ref_len} reference "
|
||||
f"words ({num_corr} correct)",
|
||||
file=f,
|
||||
)
|
||||
print(
|
||||
"Search below for sections starting with PER-UTT DETAILS:, "
|
||||
"SUBSTITUTIONS:, DELETIONS:, INSERTIONS:, PER-WORD STATS:",
|
||||
file=f,
|
||||
)
|
||||
|
||||
print("", file=f)
|
||||
print("PER-UTT DETAILS: corr or (ref->hyp) ", file=f)
|
||||
for cut_id, ref, hyp in results:
|
||||
ali = kaldialign.align(ref, hyp, ERR)
|
||||
combine_successive_errors = True
|
||||
if combine_successive_errors:
|
||||
ali = [[[x], [y]] for x, y in ali]
|
||||
for i in range(len(ali) - 1):
|
||||
if ali[i][0] != ali[i][1] and ali[i + 1][0] != ali[i + 1][1]:
|
||||
ali[i + 1][0] = ali[i][0] + ali[i + 1][0]
|
||||
ali[i + 1][1] = ali[i][1] + ali[i + 1][1]
|
||||
ali[i] = [[], []]
|
||||
ali = [[
|
||||
list(filter(lambda a: a != ERR, x)),
|
||||
list(filter(lambda a: a != ERR, y)),
|
||||
] for x, y in ali]
|
||||
ali = list(filter(lambda x: x != [[], []], ali))
|
||||
ali = [[
|
||||
ERR if x == [] else " ".join(x),
|
||||
ERR if y == [] else " ".join(y),
|
||||
] for x, y in ali]
|
||||
|
||||
print(
|
||||
f"{cut_id}:\t" + " ".join((ref_word if ref_word == hyp_word else
|
||||
f"({ref_word}->{hyp_word})"
|
||||
for ref_word, hyp_word in ali)),
|
||||
file=f,
|
||||
)
|
||||
|
||||
print("", file=f)
|
||||
print("SUBSTITUTIONS: count ref -> hyp", file=f)
|
||||
|
||||
for count, (ref, hyp) in sorted([(v, k) for k, v in subs.items()],
|
||||
reverse=True):
|
||||
print(f"{count} {ref} -> {hyp}", file=f)
|
||||
|
||||
print("", file=f)
|
||||
print("DELETIONS: count ref", file=f)
|
||||
for count, ref in sorted([(v, k) for k, v in dels.items()], reverse=True):
|
||||
print(f"{count} {ref}", file=f)
|
||||
|
||||
print("", file=f)
|
||||
print("INSERTIONS: count hyp", file=f)
|
||||
for count, hyp in sorted([(v, k) for k, v in ins.items()], reverse=True):
|
||||
print(f"{count} {hyp}", file=f)
|
||||
|
||||
print("", file=f)
|
||||
print("PER-WORD STATS: word corr tot_errs count_in_ref count_in_hyp",
|
||||
file=f)
|
||||
for _, word, counts in sorted([(sum(v[1:]), k, v)
|
||||
for k, v in words.items()],
|
||||
reverse=True):
|
||||
(corr, ref_sub, hyp_sub, ins, dels) = counts
|
||||
tot_errs = ref_sub + hyp_sub + ins + dels
|
||||
ref_count = corr + ref_sub + dels
|
||||
hyp_count = corr + hyp_sub + ins
|
||||
|
||||
print(f"{word} {corr} {tot_errs} {ref_count} {hyp_count}", file=f)
|
||||
return float(tot_err_rate)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,23 @@
|
||||
import librosa
|
||||
import os
|
||||
|
||||
import openvino_genai as ov_genai
|
||||
import huggingface_hub as hf_hub
|
||||
|
||||
|
||||
class WhisperOpenVINO(object):
|
||||
def __init__(self, model_id="OpenVINO/whisper-tiny-fp16-ov", device="CPU", language="en", task="transcribe"):
|
||||
model_path = model_id.split('/')[-1]
|
||||
cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "openvino_whisper_models")
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
model_path = os.path.join(cache_dir, model_path)
|
||||
if not os.path.exists(model_path):
|
||||
hf_hub.snapshot_download(model_id, local_dir=model_path)
|
||||
self.model = ov_genai.WhisperPipeline(str(model_path), device=device)
|
||||
self.language = language
|
||||
self.task = task
|
||||
|
||||
def transcribe(self, input_audio):
|
||||
outputs = self.model.generate(input_audio, return_timestamps=True, language=self.language, task=self.task)
|
||||
outputs = [seg for seg in outputs.chunks]
|
||||
return outputs
|
||||
@@ -0,0 +1,479 @@
|
||||
import json
|
||||
import re
|
||||
import math
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
from whisper.tokenizer import get_tokenizer
|
||||
from whisper_live.transcriber.tensorrt_utils import (
|
||||
mel_filters,
|
||||
load_audio_wav_format,
|
||||
pad_or_trim,
|
||||
load_audio
|
||||
)
|
||||
|
||||
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.bindings import GptJsonConfig, KVCacheType
|
||||
from tensorrt_llm.runtime import PYTHON_BINDINGS, ModelConfig, SamplingConfig
|
||||
from tensorrt_llm.runtime.session import Session, TensorInfo
|
||||
if PYTHON_BINDINGS:
|
||||
from tensorrt_llm.runtime import ModelRunnerCpp
|
||||
|
||||
SAMPLE_RATE = 16000
|
||||
N_FFT = 400
|
||||
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):
|
||||
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,
|
||||
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()
|
||||
inputs['input_features'] = mel
|
||||
inputs['input_lengths'] = mel_input_lengths
|
||||
inputs['position_ids'] = position_ids
|
||||
|
||||
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)
|
||||
|
||||
logger.debug(f'output info {output_info}')
|
||||
outputs = {
|
||||
t.name: torch.empty(tuple(t.shape),
|
||||
dtype=trt_dtype_to_torch(t.dtype),
|
||||
device='cuda')
|
||||
for t in output_info
|
||||
}
|
||||
stream = torch.cuda.current_stream()
|
||||
ok = self.session.run(inputs=inputs,
|
||||
outputs=outputs,
|
||||
stream=stream.cuda_stream)
|
||||
assert ok, 'Engine execution failed'
|
||||
stream.synchronize()
|
||||
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 = read_config('decoder', engine_dir)
|
||||
self.decoder_generation_session = self.get_session(
|
||||
engine_dir, runtime_mapping, debug_mode)
|
||||
|
||||
def get_session(self, engine_dir, runtime_mapping, debug_mode=False):
|
||||
serialize_path = engine_dir / 'decoder' / 'rank0.engine'
|
||||
with open(serialize_path, "rb") as f:
|
||||
decoder_engine_buffer = f.read()
|
||||
|
||||
decoder_model_config = ModelConfig(
|
||||
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'],
|
||||
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'],
|
||||
dtype=self.decoder_config['dtype'],
|
||||
has_token_type_embedding=False,
|
||||
)
|
||||
decoder_generation_session = tensorrt_llm.runtime.GenerationSession(
|
||||
decoder_model_config,
|
||||
decoder_engine_buffer,
|
||||
runtime_mapping,
|
||||
debug_mode=debug_mode)
|
||||
|
||||
return decoder_generation_session
|
||||
|
||||
def generate(self,
|
||||
decoder_input_ids,
|
||||
encoder_outputs,
|
||||
encoder_max_input_length,
|
||||
encoder_input_lengths,
|
||||
eot_id,
|
||||
max_new_tokens=40,
|
||||
num_beams=1):
|
||||
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])
|
||||
],
|
||||
dtype=torch.int32,
|
||||
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,
|
||||
num_beams=num_beams)
|
||||
self.decoder_generation_session.setup(
|
||||
decoder_input_lengths.size(0),
|
||||
decoder_max_input_length,
|
||||
max_new_tokens,
|
||||
beam_width=num_beams,
|
||||
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()
|
||||
|
||||
# get the list of int from output_ids tensor
|
||||
output_ids = output_ids.cpu().numpy().tolist()
|
||||
return output_ids
|
||||
|
||||
|
||||
class WhisperTRTLLM(object):
|
||||
|
||||
def __init__(self,
|
||||
engine_dir,
|
||||
assets_dir=None,
|
||||
device=None,
|
||||
is_multilingual=False,
|
||||
language="en",
|
||||
task="transcribe",
|
||||
use_py_session=False,
|
||||
num_beams=1,
|
||||
debug_mode=False,
|
||||
max_output_len=96):
|
||||
world_size = 1
|
||||
runtime_rank = tensorrt_llm.mpi_rank()
|
||||
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.device = device
|
||||
self.tokenizer = get_tokenizer(
|
||||
is_multilingual,
|
||||
num_languages=self.num_languages,
|
||||
language=language,
|
||||
task=task,
|
||||
)
|
||||
|
||||
if use_py_session:
|
||||
self.encoder = WhisperEncoding(engine_dir)
|
||||
self.decoder = WhisperDecoding(engine_dir,
|
||||
runtime_mapping,
|
||||
debug_mode=False)
|
||||
else:
|
||||
json_config = GptJsonConfig.parse_file(engine_dir / 'decoder' /
|
||||
'config.json')
|
||||
assert json_config.model_config.supports_inflight_batching
|
||||
runner_kwargs = dict(engine_dir=engine_dir,
|
||||
is_enc_dec=True,
|
||||
max_batch_size=1,
|
||||
max_input_len=3000,
|
||||
max_output_len=max_output_len,
|
||||
max_beam_width=num_beams,
|
||||
debug_mode=debug_mode,
|
||||
kv_cache_free_gpu_memory_fraction=0.9,
|
||||
cross_kv_cache_fraction=0.5)
|
||||
self.model_runner_cpp = ModelRunnerCpp.from_dir(**runner_kwargs)
|
||||
self.filters = mel_filters(self.device, self.n_mels, assets_dir)
|
||||
self.use_py_session = use_py_session
|
||||
|
||||
def log_mel_spectrogram(
|
||||
self,
|
||||
audio: Union[str, np.ndarray, torch.Tensor],
|
||||
padding: int = 0,
|
||||
return_duration=True
|
||||
):
|
||||
"""
|
||||
Compute the log-Mel spectrogram of
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio: Union[str, np.ndarray, torch.Tensor], shape = (*)
|
||||
The path to audio or either a NumPy array or Tensor containing the audio waveform in 16 kHz
|
||||
|
||||
n_mels: int
|
||||
The number of Mel-frequency filters, only 80 and 128 are supported
|
||||
|
||||
padding: int
|
||||
Number of zero samples to pad to the right
|
||||
|
||||
device: Optional[Union[str, torch.device]]
|
||||
If given, the audio tensor is moved to this device before STFT
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor, shape = (80 or 128, n_frames)
|
||||
A Tensor that contains the Mel spectrogram
|
||||
"""
|
||||
if not torch.is_tensor(audio):
|
||||
if isinstance(audio, str):
|
||||
if audio.endswith('.wav'):
|
||||
audio, _ = load_audio_wav_format(audio)
|
||||
else:
|
||||
audio = load_audio(audio)
|
||||
assert isinstance(audio, np.ndarray), f"Unsupported audio type: {type(audio)}"
|
||||
duration = audio.shape[-1] / SAMPLE_RATE
|
||||
audio = pad_or_trim(audio, N_SAMPLES)
|
||||
audio = audio.astype(np.float32)
|
||||
audio = torch.from_numpy(audio)
|
||||
|
||||
if self.device is not None:
|
||||
audio = audio.to(self.device)
|
||||
if padding > 0:
|
||||
audio = F.pad(audio, (0, padding))
|
||||
window = torch.hann_window(N_FFT).to(audio.device)
|
||||
stft = torch.stft(audio, N_FFT, HOP_LENGTH, window=window, return_complex=True)
|
||||
magnitudes = stft[..., :-1].abs()**2
|
||||
|
||||
mel_spec = self.filters @ magnitudes
|
||||
|
||||
log_spec = torch.clamp(mel_spec, min=1e-10).log10()
|
||||
log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)
|
||||
log_spec = (log_spec + 4.0) / 4.0
|
||||
if return_duration:
|
||||
return log_spec, duration
|
||||
else:
|
||||
return log_spec
|
||||
|
||||
def process_batch(
|
||||
self,
|
||||
mel,
|
||||
mel_input_lengths,
|
||||
text_prefix="<|startoftranscript|><|en|><|transcribe|><|notimestamps|>",
|
||||
num_beams=1,
|
||||
max_new_tokens=96):
|
||||
prompt_id = self.tokenizer.encode(
|
||||
text_prefix, allowed_special=set(self.tokenizer.special_tokens.keys()))
|
||||
|
||||
prompt_id = torch.tensor(prompt_id)
|
||||
batch_size = mel.shape[0]
|
||||
decoder_input_ids = prompt_id.repeat(batch_size, 1)
|
||||
if self.use_py_session:
|
||||
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=max_new_tokens,
|
||||
num_beams=num_beams)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
if isinstance(mel, list):
|
||||
mel = [
|
||||
m.transpose(1, 2).type(
|
||||
str_dtype_to_torch("float16")).squeeze(0)
|
||||
for m in mel
|
||||
]
|
||||
else:
|
||||
mel = mel.transpose(1, 2)
|
||||
outputs = self.model_runner_cpp.generate(
|
||||
batch_input_ids=decoder_input_ids,
|
||||
encoder_input_features=mel,
|
||||
encoder_output_lengths=mel_input_lengths // 2,
|
||||
max_new_tokens=max_new_tokens,
|
||||
end_id=self.tokenizer.eot,
|
||||
pad_id=self.tokenizer.eot,
|
||||
num_beams=num_beams,
|
||||
output_sequence_lengths=True,
|
||||
return_dict=True)
|
||||
torch.cuda.synchronize()
|
||||
output_ids = outputs['output_ids'].cpu().numpy().tolist()
|
||||
texts = []
|
||||
for i in range(len(output_ids)):
|
||||
text = self.tokenizer.decode(output_ids[i][0]).strip()
|
||||
texts.append(text)
|
||||
return texts
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
mel,
|
||||
text_prefix="<|startoftranscript|><|en|><|transcribe|><|notimestamps|>",
|
||||
dtype='float16',
|
||||
batch_size=1,
|
||||
num_beams=1,
|
||||
padding_strategy="max",
|
||||
max_new_tokens=96,
|
||||
):
|
||||
mel = mel.type(str_dtype_to_torch(dtype))
|
||||
mel = mel.unsqueeze(0)
|
||||
# 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,
|
||||
max_new_tokens=max_new_tokens
|
||||
)
|
||||
prediction = predictions[0]
|
||||
|
||||
# remove all special tokens in the prediction
|
||||
prediction = re.sub(r'<\|.*?\|>', '', prediction)
|
||||
return prediction.strip()
|
||||
|
||||
|
||||
def decode_wav_file(
|
||||
model,
|
||||
mel,
|
||||
text_prefix="<|startoftranscript|><|en|><|transcribe|><|notimestamps|>",
|
||||
dtype='float16',
|
||||
batch_size=1,
|
||||
num_beams=1,
|
||||
normalizer=None,
|
||||
mel_filters_dir=None):
|
||||
|
||||
mel = mel.type(str_dtype_to_torch(dtype))
|
||||
mel = mel.unsqueeze(0)
|
||||
# repeat the mel spectrogram to match the batch size
|
||||
mel = mel.repeat(batch_size, 1, 1)
|
||||
predictions = model.process_batch(mel, text_prefix, num_beams)
|
||||
prediction = predictions[0]
|
||||
|
||||
# remove all special tokens in the prediction
|
||||
prediction = re.sub(r'<\|.*?\|>', '', prediction)
|
||||
if normalizer:
|
||||
prediction = normalizer(prediction)
|
||||
|
||||
return prediction.strip()
|
||||
@@ -0,0 +1,82 @@
|
||||
import os
|
||||
import textwrap
|
||||
import scipy
|
||||
import numpy as np
|
||||
import av
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def clear_screen():
|
||||
"""Clears the console screen."""
|
||||
os.system("cls" if os.name == "nt" else "clear")
|
||||
|
||||
|
||||
def print_transcript(text):
|
||||
"""Prints formatted transcript text."""
|
||||
wrapper = textwrap.TextWrapper(width=60)
|
||||
for line in wrapper.wrap(text="".join(text)):
|
||||
print(line)
|
||||
|
||||
|
||||
def format_time(s):
|
||||
"""Convert seconds (float) to SRT time format."""
|
||||
hours = int(s // 3600)
|
||||
minutes = int((s % 3600) // 60)
|
||||
seconds = int(s % 60)
|
||||
milliseconds = int((s - int(s)) * 1000)
|
||||
return f"{hours:02}:{minutes:02}:{seconds:02},{milliseconds:03}"
|
||||
|
||||
|
||||
def create_srt_file(segments, resampled_file):
|
||||
with open(resampled_file, 'w', encoding='utf-8') as srt_file:
|
||||
segment_number = 1
|
||||
for segment in segments:
|
||||
start_time = format_time(float(segment['start']))
|
||||
end_time = format_time(float(segment['end']))
|
||||
text = segment['text']
|
||||
|
||||
srt_file.write(f"{segment_number}\n")
|
||||
srt_file.write(f"{start_time} --> {end_time}\n")
|
||||
srt_file.write(f"{text}\n\n")
|
||||
|
||||
segment_number += 1
|
||||
|
||||
|
||||
def resample(file: str, sr: int = 16000):
|
||||
"""
|
||||
Resample the audio file to 16kHz.
|
||||
|
||||
Args:
|
||||
file (str): The audio file to open
|
||||
sr (int): The sample rate to resample the audio if necessary
|
||||
|
||||
Returns:
|
||||
resampled_file (str): The resampled audio file
|
||||
"""
|
||||
container = av.open(file)
|
||||
stream = next(s for s in container.streams if s.type == 'audio')
|
||||
|
||||
resampler = av.AudioResampler(
|
||||
format='s16',
|
||||
layout='mono',
|
||||
rate=sr,
|
||||
)
|
||||
|
||||
resampled_file = Path(file).stem + "_resampled.wav"
|
||||
output_container = av.open(resampled_file, mode='w')
|
||||
output_stream = output_container.add_stream('pcm_s16le', rate=sr)
|
||||
output_stream.layout = 'mono'
|
||||
|
||||
for frame in container.decode(audio=0):
|
||||
frame.pts = None
|
||||
resampled_frames = resampler.resample(frame)
|
||||
if resampled_frames is not None:
|
||||
for resampled_frame in resampled_frames:
|
||||
for packet in output_stream.encode(resampled_frame):
|
||||
output_container.mux(packet)
|
||||
|
||||
for packet in output_stream.encode(None):
|
||||
output_container.mux(packet)
|
||||
|
||||
output_container.close()
|
||||
return resampled_file
|
||||
+56
-14
@@ -1,16 +1,16 @@
|
||||
# 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():
|
||||
|
||||
def __init__(self, force_onnx_cpu=True):
|
||||
path = self.download()
|
||||
|
||||
opts = onnxruntime.SessionOptions()
|
||||
opts.log_severity_level = 3
|
||||
|
||||
@@ -22,9 +22,12 @@ class VoiceActivityDetection():
|
||||
else:
|
||||
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:
|
||||
@@ -39,22 +42,27 @@ class VoiceActivityDetection():
|
||||
|
||||
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
|
||||
@@ -106,10 +121,37 @@ class VoiceActivityDetection():
|
||||
# Check if the model file already exists
|
||||
if not os.path.exists(model_filename):
|
||||
# If it doesn't exist, download the model using wget
|
||||
print("Downloading VAD ONNX model...")
|
||||
try:
|
||||
subprocess.run(["wget", "-O", model_filename, model_url], check=True)
|
||||
except subprocess.CalledProcessError:
|
||||
print("Failed to download the model using wget.")
|
||||
return model_filename
|
||||
|
||||
|
||||
class VoiceActivityDetector:
|
||||
def __init__(self, threshold=0.5, frame_rate=16000):
|
||||
"""
|
||||
Initializes the VoiceActivityDetector with a voice activity detection model and a threshold.
|
||||
|
||||
Args:
|
||||
threshold (float, optional): The probability threshold for detecting voice activity. Defaults to 0.5.
|
||||
"""
|
||||
self.model = VoiceActivityDetection()
|
||||
self.threshold = threshold
|
||||
self.frame_rate = frame_rate
|
||||
|
||||
def __call__(self, audio_frame):
|
||||
"""
|
||||
Determines if the given audio frame contains speech by comparing the detected speech probability against
|
||||
the threshold.
|
||||
|
||||
Args:
|
||||
audio_frame (np.ndarray): The audio frame to be analyzed for voice activity. It is expected to be a
|
||||
NumPy array of audio samples.
|
||||
|
||||
Returns:
|
||||
bool: True if the speech probability exceeds the threshold, indicating the presence of voice activity;
|
||||
False otherwise.
|
||||
"""
|
||||
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