485 Commits

Author SHA1 Message Date
cjones 955665401a Override get_segment_end() in ServeClientOpenVINO to match OpenVINO's WhisperDecodedResultChunk class. 2026-07-24 23:16:42 -04:00
cjones 17b1639a9b Get ready to add diarization to openvino backend. 2026-07-23 00:05:02 -04:00
Vineet Suryan 06ec02445d Merge pull request #480 from Kokkini/feature/streaming-transcription-client
Add StreamingTranscriptionClient for streaming from any source
2026-07-17 12:24:02 +02:00
Vineet Suryan daf3da8633 Merge pull request #523 from Kaihui-AMD/faster-whisper-rocm
Add AMD ROCm GPU support for faster_whisper backend
2026-07-15 15:43:46 +02:00
Kaihui-AMD 2e466b765a Add AMD ROCm GPU support for faster_whisper backend
Add ROCm_whisper.md (Docker + native install guide) and docker/Dockerfile.rocm
based on rocm/pytorch:rocm7.2.4 (PyTorch 2.10.0) that installs the official
CTranslate2 v4.8.0 ROCm wheel. The default faster_whisper backend runs on AMD
GPUs out of the box with no code changes.

Tested on Radeon AI PRO R9700 (gfx1201) and Ryzen AI Max+ 395 / Radeon 8060S
(gfx1151) with ROCm 7.2.4.

Addresses #520.
2026-07-15 16:13:37 +08:00
Vineet Suryan dfef376086 Merge pull request #524 from nightcityblade/fix/issue-521
fix: allow NumPy 2 for pyannote audio
2026-07-15 09:53:12 +02:00
nightcityblade 26845cabe9 fix: allow NumPy 2 for pyannote audio 2026-07-15 11:13:12 +08:00
Vineet Suryan c0f101f0d1 Merge pull request #516 from nightcityblade/fix/issue-324
fix: improve subtitle readability
2026-07-14 14:59:11 +02:00
Quang Tran ec1dc7c6aa fix: clean shutdown for StreamingTranscriptionClient 2026-07-08 23:13:20 +07:00
Quang Tran a71c578570 fix: address PR review comments on StreamingTranscriptionClient 2026-07-08 22:49:23 +07:00
Quang Tran f4f1b1d8be feat: support manual audio streaming from any source 2026-07-08 22:46:34 +07:00
nightcityblade ad4d314b79 test: cover subtitle wrapping edge cases 2026-07-07 23:04:01 +08:00
nightcityblade 8c0caf1be0 fix: improve subtitle readability 2026-07-07 23:04:01 +08:00
Vineet Suryan d9459ebf2d Merge pull request #517 from dmaier-ef/fix/idle-transcription-thread-cpu-contention
Fix idle-client busy-wait before first audio frame
2026-07-06 11:37:35 +05:30
David Maier 5b577b34e4 Add configurable timeout for first-frame wait and improve thread-safety 2026-07-03 14:45:09 +02:00
Aaron Boxer 2debc0ee80 Make initial_prompt and vad_parameters accessible from the client
Expose initial_prompt and vad_parameters as flat client parameters
(consistent with hotwords, send_last_n_segments, etc.), send them as
flat handshake keys, and read them server-side with null-safe
options.get(...). Supersedes #283; avoids the options-bag None.get crash.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-26 11:36:26 -04:00
Aaron Boxer 9f7a043d8b ctranslate2: add missing libcublas dep 2026-06-26 10:38:56 -04:00
Alessandro Griseta ee64458194 Use command -v dnf instead 2026-06-26 10:28:41 -04:00
David Maier 056774ea50 Fix idle-client busy-wait before first audio frame 2026-06-26 14:57:03 +02:00
Vineet Suryan e4160d2d06 Merge pull request #513 from SuperCowProducts/clarify-installation-instructions
Make install instructions clearer
2026-06-26 16:03:33 +05:30
nightcityblade 471c3fd6b4 docs: add macOS OpenMP workaround 2026-06-24 11:14:48 -04:00
makaveli10 ac7a9f849c fix(batch): add temperature fallback to prevent decoder runaway
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2026-06-23 10:52:02 -04:00
makaveli10 ecb052c873 Enable single-model mode for stock models when batch_inference is set
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2026-06-23 10:52:02 -04:00
nightcityblade ee3113a507 docs: clarify iOS project setup 2026-06-23 10:50:31 -04:00
Alessandro Griseta 81ff199a70 First venv activate, then pip install 2026-06-18 12:29:43 +02:00
Aaron Boxer d0ad362b20 chrome extension: point to Modal transcription service
- add unit test
2026-06-14 12:18:36 -04:00
Aaron Boxer 44940e2834 Support known speaker hints in REST API 2026-06-14 12:17:44 -04:00
nightcityblade 32c1b18c9f fix: support uint8 websocket audio format 2026-06-04 16:04:38 -04:00
Vineet Suryan 582d5426d6 Bump version v0.9.0 2026-06-02 11:54:07 +05:30
nightcityblade 1814dd9bfa fix: include server API dependencies in package 2026-06-01 19:50:34 -04:00
nightcityblade 0747529910 test: strengthen wheel install smoke test 2026-06-01 19:50:34 -04:00
nightcityblade cef3352c5b test: add clean virtualenv install smoke test 2026-06-01 19:50:34 -04:00
nightcityblade 3516f663e5 fix: support Python 3.13 server installs 2026-06-01 19:24:17 -04:00
nightcityblade d006d4abb3 fix: make caption overlay line count configurable
Fixes collabora/WhisperLive#486
2026-06-01 19:18:15 -04:00
nightcityblade 6242ad7f81 fix: add av to package requirements 2026-06-01 18:58:28 -04:00
Aaron Boxer 5334ea0f7a Add WebSocket authentication via api_key
- When --api_key is set, WebSocket connections require auth too
- Supports Authorization: Bearer <key> header or ?token=<key> query param
- Unauthenticated connections receive HTTP 401 before upgrade
- Uses websockets process_request callback (no resource allocation before auth)
- Added 5 unit tests for WebSocket auth handler
2026-05-26 23:07:50 -04:00
Aaron Boxer b648bcb2a4 Add optional API key auth and rate limiting for REST API
- api_key param: requires 'Authorization: Bearer <key>' header
- rate_limit_rpm param: per-IP sliding-window rate limit (requests/min)
- Both are off by default (backward compatible)
- CLI flags: --api_key, --rate_limit_rpm
- Added 5 unit tests for auth and rate limiting
2026-05-26 23:02:43 -04:00
Aaron Boxer e86d98dd80 Improve REST API unsupported param warnings
- Enumerate each ignored param individually in log message
- Add warnings for 'include' param (was previously silent)
- Added 6 unit tests for REST API param validation
2026-05-26 22:59:12 -04:00
Aaron Boxer c5ec7f4a99 Add reconnect logic to WebSocket client
- New params: max_retries (default 0), retry_delay (default 5s)
- On unexpected close, retries up to max_retries times
- Does not retry on server_error (server rejected connection)
- Extracted _create_websocket() helper for reuse
- Added 4 unit tests for reconnect behavior
2026-05-26 05:45:37 -04:00
Vineet Suryan 52005b94cb Merge pull request #443 from boxerab/configurable-constants
Extract hardcoded buffer constants into class attributes
2026-05-26 13:55:13 +05:30
Vineet Suryan f2f769532b Merge pull request #442 from boxerab/unify-translate-flags
Clarify --translate vs --enable_translation CLI flags
2026-05-26 13:25:23 +05:30
Vineet Suryan cdc661ce28 Merge pull request #434 from nightcityblade/fix/issue-405
docs: fix setup instructions reported in #405
2026-05-26 13:17:37 +05:30
Aaron Boxer 8396763444 fix: set proper mock_info attributes in SSE streaming tests
MagicMock auto-attributes are not JSON serializable. Set language,
language_probability, and duration explicitly. Also exclude metadata
events from segment count assertion.
2026-05-25 09:44:21 -04:00
Aaron Boxer 8ac98dceec feat: add SSE streaming for REST transcription endpoint
- stream=true now returns text/event-stream with per-segment SSE events
- Each segment yields 'data: {json}' followed by 'data: [DONE]'
- Error events streamed as 'data: {"error": ...}'
- Temp files cleaned up in finally block
- 5 new tests in test_server_extended.py (183 total passing)
2026-05-25 09:44:21 -04:00
Aaron Boxer 8bde966c1e docs: remove features from README that moved to Aavaaz
Remove authentication, rate limiting, and auto-reconnect documentation
since these features now live in the Aavaaz project.
2026-05-15 10:45:31 -04:00
Aaron Boxer dc4a707f9a chore: add .gitignore to exclude __pycache__, virtualenvs, and build artifacts 2026-05-15 10:45:31 -04:00
Aaron Boxer c028c4b584 feat: add segment_post_processor hook for external plugins
Add a minimal, non-breaking hook to WhisperLive that allows external
projects to post-process transcription segments before they are sent
to the client.

Changes:
- ServeClientBase: add segment_post_processor attribute (default None)
- ServeClientBase.send_transcription_to_client: apply post_processor
  per-segment with error handling (falls back to original segment)
- TranscriptionServer: add segment_post_processor parameter to run()
  and wire it to each client on creation

This enables downstream projects to plug in custom processing
(e.g. formatting, PII redaction, diarization tagging) without
modifying WhisperLive core code.
2026-05-15 10:45:31 -04:00
Aaron Boxer 4e31f8c61b Add word-level timestamps and confidence scores
- New word_timestamps option (default False) in client handshake
- When enabled, each segment includes 'words' array with per-word
  start/end times and probability scores
- Wired through entire pipeline: client → server → backend → transcribe()
- Words include timestamp_offset for accurate absolute times
- REST API already supported word timestamps; now WebSocket does too
- Added 9 unit tests for word timestamp extraction and formatting
2026-05-15 10:35:27 -04:00
Aaron Boxer 18de3eacc7 Document all new features in README
- Added 'Advanced Features' section with 8 subsections
- Word-level timestamps: WebSocket JSON example
- Custom vocabulary / hotwords: usage and REST API support
- Speaker diarization: setup, pyannote dependency, output format
- Authentication: API key for REST + WebSocket
- Rate limiting: per-IP RPM configuration
- Auto-reconnect: max_retries / retry_delay
- Batch inference: CLI flags
- Raw PCM input: int16 normalization
- Updated table of contents
2026-05-15 09:01:48 -04:00
Aaron Boxer 18b897277f Add real-time speaker diarization support
- New whisper_live/diarization.py: SpeakerDiarizer with online clustering
- Uses pyannote.audio speaker embeddings (optional dependency)
- Cosine similarity threshold for speaker matching (default 0.55)
- Running average embedding update for speaker stability
- Configurable max_speakers limit (default 10)
- Client options: enable_diarization, max_speakers
- Segments include 'speaker' field when diarization is active
- Graceful fallback: logs warning if pyannote not installed
- Added 12 unit tests (mock-based, no GPU required)
2026-05-13 10:50:59 -04:00
Aaron Boxer 3d63e82571 fix: add skip decorator to TestStartMetricsServer for CI without prometheus_client 2026-05-13 10:45:40 -04:00
Aaron Boxer ced4bdb737 feat: add Prometheus metrics instrumentation
- New whisper_live/metrics.py with Counter, Gauge, Histogram metrics
- Track connections (opened/closed/rejected), transcription latency,
  audio processed, segments emitted, REST requests, and errors
- All metric helpers are no-ops when prometheus_client not installed
- --metrics_port CLI flag to expose /metrics endpoint (0 = disabled)
- Metrics integrated into server.py, base.py at key instrumentation points
- 17 new tests in tests/test_metrics.py (178 total passing)
2026-05-13 10:45:40 -04:00
Aaron Boxer 4210697ca6 Add custom vocabulary / hotwords support 2026-05-13 10:32:28 -04:00
Aaron Boxer 445bf26e85 Extract hardcoded buffer constants into class attributes
- MAX_BUFFER_DURATION_S (45): max audio buffer before trimming
- BUFFER_TRIM_DURATION_S (30): duration to discard on trim
- CLIP_THRESHOLD_DURATION_S (25): stale audio clip threshold
- CLIP_TAIL_DURATION_S (5): audio tail to keep after clipping
- All values can now be overridden by subclasses
2026-05-11 20:05:38 -04:00
Aaron Boxer b534f9d249 Clarify --translate vs --enable_translation CLI flags
- --translate: Whisper built-in to-English translation (task=translate)
- --enable_translation: M2M100 any-to-any translation backend
- Added warning when both flags are used simultaneously
- Updated help text to distinguish the two features
2026-05-11 20:05:34 -04:00
Vineet Suryan 9a71a95ca8 Merge pull request #441 from boxerab/fix-clear-screen-shell-injection
Avoid clear screen shell injection by using ANSI escape codes
2026-05-08 17:36:11 +02:00
Vineet Suryan 485b211072 Merge pull request #440 from boxerab/input-validation-server-params
Validate server parameters on startup
2026-05-08 17:35:26 +02:00
Vineet Suryan 19847784ae Merge pull request #439 from boxerab/bounded-transcript-memory
Bound transcript memory and translation queue size
2026-05-08 17:34:41 +02:00
Vineet Suryan 52d94bf1fb Merge pull request #438 from boxerab/thread-safety-client-manager
Add thread safety to client manager with threading lock
2026-04-21 13:29:54 +02:00
Vineet Suryan 1c663d0bba Merge pull request #437 from boxerab/rawpcm
audio: add support for raw pcm input via server flag
2026-04-21 13:17:47 +02:00
Vineet Suryan 298a01f1b0 Merge pull request #436 from boxerab/testing
CI: expand test suite coverage
2026-04-20 17:58:48 +02:00
Aaron Boxer a6147a6745 Replace os.system() in clear_screen() with ANSI escape codes
- Eliminates shell injection risk from os.system('clear'/'cls')
- Uses ANSI escape sequence \033[H\033[2J instead
- Removed unused os import
- Added test verifying ANSI codes are used
2026-04-17 09:31:00 -04:00
Aaron Boxer 18bce1864a Validate server parameters on startup
- max_clients must be >= 1
- max_connection_time must be > 0
- batch_max_size must be >= 1 (when batch enabled)
- batch_window_ms must be >= 0 (when batch enabled)
- Added 5 new tests for parameter validation
2026-04-17 09:30:19 -04:00
Aaron Boxer 9e5e4a9970 Bound transcript memory and translation queue size
- Add MAX_TRANSCRIPT_LENGTH (500) and MAX_TRANSLATION_QUEUE_SIZE (100)
  class constants to ServeClientBase
- Trim transcript and text lists after each update_segments() call
- Create translation queue with maxsize to prevent unbounded growth
- Added tests for _trim_transcript()
2026-04-17 09:29:38 -04:00
Aaron Boxer 81cdbbca95 Add thread safety to ClientManager with threading.Lock
- All ClientManager methods (add_client, get_client, remove_client,
  get_wait_time, is_server_full, is_client_timeout) now protected by
  a threading.Lock
- cleanup() called outside the lock to avoid holding it during I/O
- is_server_full() computes wait time inline under lock instead of
  calling get_wait_time() to avoid nested lock acquisition
- Added concurrent thread safety tests for add/remove and get operations
2026-04-17 09:27:37 -04:00
Aaron Boxer f5340ddf1e audio: add support for raw pcm input via server flag
fixes #
2026-04-17 09:21:06 -04:00
Aaron Boxer b1cd51ac8a CI: expand test suite coverage
these new test cover issues such as thread safety, VAD thresholding,
message routing, error handling etc. that weren't covered by existing
tests. Mocking is used to avoid dependencies on GPU, ONNX etc.
2026-04-17 09:05:34 -04:00
Vineet Suryan e41324bf03 Merge pull request #435 from nightcityblade/fix/issue-327
fix: render transcript text safely in browser extensions
2026-04-16 12:37:33 +02:00
nightcityblade 31efff9330 fix: render transcript text safely in browser extensions 2026-04-15 23:09:39 +08:00
nightcityblade 68a8b57e66 docs: fix setup instructions errors reported in #405
- Clarify that setup.sh installs portaudio system dependency and list
  per-distro package names
- Add missing --gpus all flag to TensorRT Docker run command
- Fix Docker TensorRT example showing multiple --trt_model_path on one
  command (should be separate alternatives)
- Document --trt_py_session flag as workaround for TensorRT C++ session
  crashes (CrossAttentionMask warnings)

Closes #405

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-13 11:09:40 +08:00
Vineet Suryan 6de5c87d2f Bump version v0.8.0 2026-03-17 14:48:07 +05:30
Vineet Suryan 8e09d16ee4 Merge pull request #430 from makaveli10/vineet/fix-run-client
Fix crash when no --files provided; use microphone input instead
2026-03-17 14:46:57 +05:30
Xiaoliang Gao 4943c25ff7 Fix crash when no --files provided; use microphone input instead 2026-03-17 14:36:20 +05:30
Vineet Suryan 710bdffb51 Merge pull request #429 from makaveli10/vineet/update-setup-packages
Expose __version__ in package root and update dependencies in setup.py
2026-03-17 14:19:15 +05:30
makaveli10 5f0010d720 Expose __version__ in package root and update dependencies in setup.py
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2026-03-17 13:34:34 +05:30
Vineet Suryan 6fcae6a30c Merge pull request #425 from nightcityblade/fix/issue-377
feat: make display_segments configurable in Client/TranscriptionClient
2026-03-16 10:34:29 +05:30
Vineet Suryan bc441dea23 Merge pull request #427 from ianwh02/fix/batch-single-vad-none
Fix NoneType crash in _process_single when VAD filters all audio
2026-03-13 17:02:43 +05:30
ianwh02 89466f7b77 Fix NoneType crash in _process_single when VAD filters all audio
When VAD removes all speech from an audio chunk, transcriber.transcribe() returns (None, info). Calling list(None) raises TypeError. The _process_multi path already handles this case; this aligns _process_single to match.
2026-03-11 16:48:23 +00:00
nightcityblade 6ae57c81cd feat: add --n_display_segments CLI arg to run_client.py 2026-03-11 00:05:07 +08:00
Vineet Suryan 5d8629ea0c Merge pull request #422 from ianwh02/feature/batch-inference
Add cross-client GPU batch inference for faster_whisper backend
2026-03-09 23:01:28 +05:30
ianwh02 e7e78a7151 Add unit tests for BatchInferenceWorker 2026-03-09 16:18:36 +00:00
ianwh02 3508b39584 Fix missing batch_config init causing CI test hang 2026-03-09 11:42:15 +00:00
nightcityblade 067573a510 feat: make display_segments configurable in Client/TranscriptionClient
Replace hardcoded [-4:] truncation with a configurable display_segments
parameter (default: 4) in both Client and TranscriptionClient classes.

Fixes #377
2026-03-08 12:19:54 +08:00
ianwh02 e8bd4fd532 Add cross-client GPU batch inference for faster_whisper backend 2026-02-26 00:00:51 +00:00
Marcus Edel f8869906b0 Merge pull request #419 from makaveli10/bump-whisper-version
Bump openai-whisper version to 20250625.
2026-02-20 08:04:31 -05:00
makaveli10 9fa7005511 Bump openai-whisper version to 20250625
Resolves pkg_resources missing during wheel build

Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2026-02-11 12:39:38 +00:00
Vineet Suryan e48d16f923 Merge pull request #398 from AlexStansfield/feature/faster-whisper-1.2.0
feat: update to support faster whisper 1.2.0
2026-02-11 17:57:05 +05:30
Vineet Suryan 98fcc5110b Merge pull request #418 from JenySadadia/enable-timestamps
Enable timestamps for transcripted text
2026-02-11 17:51:58 +05:30
Jeny Sadadia 5e33aa2a7e Enable timestamps for transcripted text
Add `--enable-timestamps` option to `run_client.py`
script to print out transcripted text with timestamps.

Sample output with translation enabled:
```
[0.000 -> 7.440]  And so, my fellow Americans, ask not what your country can do for you.
[7.440 -> 10.300]  Ask what you can do for your country.

TRANSLATION to fr:
[0.000 -> 7.440] Et donc, mes camarades américains, ne demandez pas ce que votre pays peut faire pour vous.
[7.440 -> 10.300] Demandez ce que vous pouvez faire pour votre pays.
```

Signed-off-by: Jeny Sadadia <jeny.sadadia@collabora.com>
2026-02-10 15:43:20 +05:30
Vineet Suryan 6c8142a9d2 Merge pull request #415 from JenySadadia/run-client-docs
README.md: add instructions for running client
2026-02-06 20:03:19 +05:30
Aaron Boxer 29ee640409 api: add support OpenAI REST transcription api 2026-02-05 22:31:59 -05:00
Aaron Boxer b9ae2af8e6 setup.sh: support Fedora 2026-02-05 22:31:59 -05:00
Jeny Sadadia 9251394047 README.md: add instructions for running client
Specify command to run client script.

Signed-off-by: Jeny Sadadia <jeny.sadadia@collabora.com>
2026-02-03 17:11:40 +05:30
Marcus Edel c6ee9a6870 Merge pull request #412 from makaveli10/vineet/fix-faster-whisper-custom-model-loading
Feat: support HuggingFace model IDs for faster_whisper_custom_model_path.
2026-01-13 15:27:43 -05:00
makaveli10 f5256fc62f feat: support HuggingFace model IDs for faster_whisper_custom_model_path
Previously, the server only accepted local file paths for custom Faster Whisper
models. This change allows passing HuggingFace repo IDs which are automatically
downloaded and converted to CTranslate2 format by the backend if not already in
CTranslate2 format.

Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2026-01-13 16:34:04 +05:30
Alex Stansfield c43eb1dd5a update to support faster whisper 1.2.0 2025-10-07 14:20:20 +00:00
Vineet Suryan 3b17bda5f9 Merge pull request #397 from locnnil/patch-1
fix(run_server.py): help text for max_connection_time argument
2025-09-25 21:09:58 +05:30
Lincoln Wallace 95a9b7ef05 fix(run_server.py): help text for max_connection_time argument
The help text for `--max_connection_time` is incorrect. Looks like a copy-paste mistake from `--cache_path`.
2025-09-25 11:14:35 -03:00
Vineet Suryan 5e6be74f6d Merge pull request #391 from makaveli10/integrate_live_translation
Integrate live translation
2025-07-24 12:07:03 +05:30
makaveli10 04db67170b ServeClientTranslation import only when enable_tranlsation is True
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-22 15:47:06 +00:00
makaveli10 5ce401d4c6 Update test_client to expect translation args
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-22 09:06:31 +00:00
makaveli10 39dfd7521f Update requirements
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-22 09:05:54 +00:00
makaveli10 2b8b245fa8 Add translation backend
Translate from any language to any language with alirezamsh/small100
running in a thread and reading from a queue shared with transcription thread.

Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-22 08:54:10 +00:00
makaveli 1ec437e71f Merge pull request #387 from makaveli10/modify_max_client_time_server_only
Change max_clients max_connection_time from server only
2025-07-22 10:13:48 +05:30
makaveli10 bf6251e3b8 Fix max_client, max_connection_time failed tests
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-21 23:09:16 +05:30
makaveli10 8d6ddd4f7b Change max_clients max_connection_time from server only
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-21 22:52:24 +05:30
makaveli 8d785e5681 Merge pull request #384 from klonikar/main
issue 371
2025-07-21 18:44:57 +05:30
Kiran Lonikar 368bcdd81f temporarily delete web_live directory to merge PR 2025-07-17 16:46:51 +05:30
makaveli d9e608f5c8 Merge pull request #385 from makaveli10/add_srt_download_opt_browser_ext
Add srt download opt browser ext
2025-07-16 15:49:40 +05:30
makaveli10 ad0fb23936 Add download srt file option firefox extension
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-14 09:43:52 +05:30
makaveli10 914281f449 Add download srt file option chrome extension
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-07-14 09:43:39 +05:30
Kiran Lonikar 40edd25468 remove commented code 2025-07-12 22:29:30 +05:30
Kiran Lonikar ad11b2b0ef web client which can take microphone input and transcribe the speech 2025-07-09 21:01:54 +05:30
Kiran Lonikar ddd32cc30f adding test client to transcribe audio files 2025-07-07 13:24:13 +05:30
Kiran Lonikar e597c876cf changes to run when audio playback is muted 2025-07-07 13:21:20 +05:30
Kiran Lonikar 9954548075 issue 371
model name is of form namespace/repo_name and not os path.
2025-07-06 16:10:49 +05:30
makaveli 0f21c80ed8 Merge pull request #382 from ParkMazorika/main
Add iOS client for WhisperLive (Audio-Transcription-iOS)
2025-06-30 18:51:31 +05:30
Park hyeon gyu d79e720b34 Delete .DS_Store 2025-06-30 18:50:12 +09:00
바견규 f3acfa2f18 Add iOS client section to main README/ modify iOS README 2025-06-28 07:21:34 +09:00
Park hyeon gyu 2e5aae6585 Update Audio-Transcription-iOS/README.md
Co-authored-by: makaveli <39617050+makaveli10@users.noreply.github.com>
2025-06-23 20:50:39 +09:00
바견규 179b56a260 Add iOS client for WhisperLive (Audio-Transcription-iOS) 2025-06-17 01:19:56 +09:00
makaveli 4ae3825661 Merge pull request #378 from makaveli10/auto_convert_faster_whisper
Auto convert hf custom whisper to ct2(faster-whisper)
2025-06-02 17:52:19 +05:30
Marcus Edel 198a499f96 Merge pull request #381 from makaveli10/add_blog_to_readme
Add hi transcription video; add blog post section.
2025-06-02 07:54:34 -04:00
makaveli10 05002d6ded Add hi transcription video; add blog post section
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-06-02 10:52:48 +00:00
makaveli10 74abf66d48 Make cache path configurable to save auto converted ct2 models
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-06-01 09:04:22 +00:00
Marcus Edel 0520978c0e Merge pull request #379 from adamsz-lume/run-setup-sh-on-mac
Make setup.sh to work on macos.
2025-05-30 13:59:48 -04:00
Adam 12f3bb2012 make setup.sh to work on macos 2025-05-29 17:44:39 +01:00
makaveli10 bff88ed3e7 Auto convert hf custom whisper to ct2(faster-whisper)
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-05-29 09:23:32 +00:00
makaveli 4b46371dac Bump version v0.7.1 2025-05-15 10:25:17 +05:30
makaveli 1f4c918d01 Merge pull request #376 from xXLosKrachosXx/main
Add transcription callback parameter to TranscriptionClient #361
2025-05-15 10:08:59 +05:30
makaveli cd327bab50 Bump version to v0.7.0 2025-05-15 09:47:44 +05:30
Erik b91b3664c2 Add transcription callback parameter to TranscriptionClient #361 2025-05-14 23:03:14 +02:00
Marcus Edel 2375924b45 Merge pull request #375 from makaveli10/update_trt_docs
Update tensorrt_llm docker setup
2025-05-14 11:05:36 -04:00
makaveli ae169245a1 Update tensorrt_llm docker setup
Remove tensorrt build from ci due to space limitations

Signed-off-by: makaveli <vineet.suryan@collabora.com>
2025-05-14 20:29:13 +05:30
makaveli 4ba576fb06 Merge pull request #374 from xXLosKrachosXx/main
Add transcription callback to Client for handling transcription results
2025-05-14 20:19:25 +05:30
makaveli a27ac16d1f Merge pull request #373 from rover0811/main
Add: support for secure WebSocket (WSS) connections
2025-05-13 15:46:37 +05:30
Erik 188b21f1d0 Add transcription callback to Client for handling transcription results 2025-05-12 20:56:04 +02:00
rover0811 d29993048d Fix: Enable support for WebSocket streaming in client.
Added the `use_wss` parameter to allow the client to handle WebSocket-based streaming. This enhances flexibility for real-time transcription scenarios.
2025-05-12 16:42:11 +09:00
rover0811 41d9f683a8 Add: support for secure WebSocket (WSS) connections
Introduce an optional `use_wss` flag to enable secure WebSocket protocol. Updated socket URL generation to dynamically select between `ws` or `wss` based on the flag value. Ensures greater flexibility when connecting to secure servers.
2025-05-12 13:53:41 +09:00
makaveli d9d8d511c7 Merge pull request #367 from giubots/configure-more-params
Add possibility to configure more parameters
2025-05-06 11:29:18 +05:30
giubots 275ed4e45b Merge branch 'main' into configure-more-params 2025-05-02 12:23:45 +02:00
giubots 9cfd8f85b6 test: add new parameters to tests 2025-05-02 11:58:17 +02:00
makaveli 7fb2d356f9 Merge pull request #368 from makaveli10/upgrade_trt_v0_18
Upgrade tensorrt_llm to v0.18.2
2025-04-30 12:19:39 +05:30
makaveli af50fed180 Merge pull request #366 from emmanuel-ferdman/main
Resolve daemon warnings for threading methods
2025-04-30 12:09:03 +05:30
giubots a2271806c3 feat: client sends new parameters to server 2025-04-28 17:21:33 +02:00
giubots 0abf8693ef refactor: include additional parameters
Refactor ServeClientBase and its subclasses to include additional parameters for segment handling and audio clipping.
2025-04-25 13:09:44 +02:00
Emmanuel Ferdman 444a1df740 Resolve daemon warnings for threading methods
Signed-off-by: Emmanuel Ferdman <emmanuelferdman@gmail.com>
2025-04-25 00:30:08 -07:00
makaveli10 47ee035f65 Upgrade tensorrt_llm to v0.18.2
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-04-22 12:33:03 +00:00
makaveli d9cb4ffdd0 Merge pull request #359 from makaveli10/remove_blank_segment
Remove blank segment feature
2025-04-22 17:59:10 +05:30
makaveli10 9b364f267a Remove blank segment feature
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-04-17 08:29:14 +00:00
makaveli 617f587699 Merge pull request #354 from makaveli10/remove_audio_clipping
Remove clip_audio from faster_whisper backend
2025-04-15 18:27:07 +05:30
makaveli10 fb3deb2745 Remove clip_audio from faster_whisper backend
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-04-15 08:00:46 +00:00
makaveli 5e430f8154 Merge pull request #353 from Perseus14/patch-1
Fix typo in setup.py
2025-04-15 13:30:30 +05:30
Rishabh Manoj efb51bf0fa Fix typo in setup.py 2025-04-13 00:36:22 +05:30
makaveli 2abca69c9d Merge pull request #348 from makaveli10/integrate_openvino
Integrate openvino
2025-04-08 23:31:53 +05:30
makaveli a62495b090 Integrate OpenVINO backend
Signed-off-by: makaveli <vineet.suryan@collabora.com>
2025-03-31 12:57:19 +05:30
makaveli10 c1ac71ada0 Refactor 🔨
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-03-24 16:48:33 +05:30
makaveli f5bea0a693 Bump version v0.6.3 2025-02-26 19:41:56 +05:30
Marcus Edel 5b3bef5845 Merge pull request #341 from makaveli10/fix_py312_pypi_install
Fix setup.py onnxruntime version for py312 pypi installation support.
2025-02-24 05:39:41 -05:00
makaveli10 2c761adc32 Fix setup.py onnxruntime version for py312 pypi installation support
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-02-24 10:41:32 +02:00
makaveli 379bd146fc Bump version v0.6.2 2025-02-07 17:07:13 +05:30
makaveli e93c2823b1 Merge pull request #334 from makaveli10/add_option_to_mute_audio_playback
Add option to mute audio playback for file input
2025-02-06 10:54:18 +05:30
makaveli10 87520498e9 Add option to mute audio playback for file input
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-02-05 20:23:41 +05:30
Marcus Edel 23d71fdbce Merge pull request #333 from makaveli10/add_support_py_312
Add support py 312.
2025-02-05 08:40:55 -05:00
makaveli10 ef7c32dc95 Add python 3.12 to test matrix
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-02-03 19:43:57 +05:30
makaveli10 28be23340b Upgrade onnxruntime version to 1.17.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-02-03 19:12:30 +05:30
makaveli ba5aa5aa38 Merge pull request #331 from makaveli10/replace_ffmpeg_with_av_lib
Replace ffmpeg with av lib for resampling, rtsp & hls streams
2025-01-22 22:23:36 +05:30
makaveli10 779baff9c3 Add pynvml missing dep for tensorrt
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-22 05:11:26 -05:00
makaveli10 5aa5826f36 Replace ffmpeg with av lib for resampling, rtsp & hls streams
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-22 05:11:04 -05:00
makaveli 893265bb3f Merge pull request #321 from makaveli10/fix_tensorrt_docker_image
Revert to 12.4.1 base image
2025-01-17 15:55:49 +05:30
makaveli10 5120afbc25 Revert to 12.4.1 base image
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-17 10:19:38 +00:00
makaveli 4baccf75a7 Bump version v0.6.1 2025-01-16 10:43:26 +05:30
makaveli b7acb8c872 Merge pull request #320 from makaveli10/fix_deprecated_package_name
Fix package name
2025-01-16 10:42:44 +05:30
makaveli10 fe7b55efe4 Fix package name
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-16 05:07:13 +00:00
makaveli c1b249ad0d Merge pull request #319 from makaveli10/upgrade_silero_vad_v5
Upgrade silero vad v5
2025-01-13 18:20:13 +05:30
makaveli10 5e4589cfe1 Upgrade silero vad v5.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-13 11:37:58 +00:00
makaveli10 b6b73730fb Fix: typo
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-13 11:35:22 +00:00
makaveli 953a88c7da Merge pull request #318 from makaveli10/fix_skipped_audio_chunk
Fix skipped audio chunk
2025-01-13 11:58:17 +05:30
makaveli10 182b5cbd6d Fix skipped audio chunk by recording the time of the first repition of a segment
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-08 13:58:34 +00:00
makaveli 32ba924d8c Bump version v0.6.0 2025-01-07 18:10:03 +05:30
Marcus Edel 450433b07b Merge pull request #316 from makaveli10/fix_data_incosistency
Add lock to thread shared variables updates/reads.
2025-01-06 10:00:43 -05:00
makaveli10 38bff6a901 Update requirements & versions
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-06 07:05:51 +00:00
makaveli10 c936e5f727 Add lock to thread shared variables updates/reads
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2025-01-06 06:34:02 +00:00
makaveli 18de63c649 Merge pull request #307 from makaveli10/fix_docker_tensorrt
Set docker tesnorrt job timeout to 60 mins
2024-12-18 17:02:16 +05:30
makaveli10 7bcd8b9520 Set docker tesnorrt job timeout to 60 mins
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-12-18 11:19:30 +00:00
Marcus Edel 19c05c8231 Merge pull request #301 from makaveli10/upgrade_tensorrt
Upgrade tensorrt_llm==0.15.0.
2024-12-03 10:04:14 -05:00
makaveli10 49e232bc4d Set segment.completed to False by default
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-12-03 20:09:41 +05:30
makaveli10 30617dfd44 Checkout git tensorrt_llm==v0.15.0 in dockerfile
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-12-03 18:27:11 +05:30
makaveli a55b99c11e Merge pull request #299 from makaveli10/fix-py38-tests
Fix requirements & tests for py38
2024-11-28 15:47:09 +05:30
makaveli10 2725f1aed9 Fix requirements & tests for py38
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-28 15:40:30 +05:30
Marcus Edel 53c31f3570 Merge pull request #297 from makaveli10/support_hf_models
Support loading hf models.
2024-11-27 13:54:03 -05:00
Marcus Edel e65fbcd9fc Merge pull request #298 from makaveli10/upgrade_faster_whisper
Upgrade faster_whisper==1.1.0 official release.
2024-11-27 13:53:39 -05:00
makaveli10 7f0c7a6791 Upgrade faster_whisper==1.1.0 official release
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-26 21:36:28 +05:30
makaveli10 2eff360b9e Support loading hf models
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-26 11:14:18 +05:30
makaveli10 c25a036c02 Upgrade tensorrt_llm to 0.15.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-21 06:13:00 +00:00
Marcus Edel 446fc6e835 Merge pull request #296 from makaveli10/upgrade_faster_whisper
Upgrade faster-whisper 1.1.0rc0.
2024-11-19 08:27:18 -05:00
makaveli10 a1650eaa4f Fix client tests to write srt file
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-19 13:31:55 +05:30
makaveli10 a6523b6b71 Minor fixes for better punctuations
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-19 02:13:11 -05:00
makaveli10 e275d34943 Remove pinned tiktoken version from server requirements
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-18 13:06:50 +05:30
makaveli10 778a9c5903 Upgrade faster-whisper 1.1.0rc0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-11-14 10:07:37 -05:00
Marcus Edel 0e89573798 Merge pull request #292 from makaveli10/fix_srt_file_missing_segments
Fix srt file missing segments.
2024-11-05 08:41:44 -05:00
makaveli10 8d89de22d8 Update tests to incorporate the completed boolean in segments
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-11-05 18:12:09 +05:30
makaveli10 81c57ae40c Send completed bool with each segment
Completed bool represents if the segment is completely processed by the server

Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-11-05 18:11:32 +05:30
Marcus Edel 00f0ff1112 Merge pull request #284 from makaveli10/expose_client_manager_args
Expose client manager args.
2024-10-31 15:40:11 -04:00
makaveli10 8b87a0562d Fix unittest to exposed client manager args
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-10-28 17:01:06 +05:30
makaveli10 617fda2864 Update Readme
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-10-10 08:43:21 -04:00
makaveli10 0d74790c67 Expose ClientManager arguments to be passed from client
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-10-10 08:37:45 -04:00
makaveli10 1322dd3c27 Pin openai-whisper version
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-10-10 08:36:21 -04:00
Marcus Edel be71657397 Merge pull request #276 from makaveli10/fix_tensorrt_docker_deps
Upgrade tensorrt-llm==`0.10.0`.
2024-09-20 12:00:09 -04:00
makaveli10 a317597f01 Update README: fix typo
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-09-16 00:32:16 -04:00
makaveli10 aaa47cfab5 Fix requirements & upgrade tensorrt-llm
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-09-16 00:25:23 -04:00
makaveli bc070d6688 Bump version 0.5.1 2024-09-05 09:34:30 +05:30
Marcus Edel 8e7e329a39 Merge pull request #274 from makaveli10/fallback_to_fp32
Set compute_type based on device capability.
2024-09-03 09:17:31 -04:00
makaveli10 380f07394b Set compute_type based on device capability
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-09-03 00:47:38 -04:00
Marcus Edel 30f78a2cc6 Merge pull request #272 from makaveli10/fix_last_segment_init
Initialize last_segment to None.
2024-08-30 12:39:49 -04:00
makaveli10 01c6bc1ecd Initialize last_segment to None
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-30 07:47:03 -04:00
Marcus Edel bdaed45820 Merge pull request #262 from makaveli10/discard_no_speech_segments
Discard no speech segments.
2024-08-19 10:12:39 -04:00
makaveli10 4870e9fb9e Make text logging optional
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:09:44 -04:00
makaveli10 ccb183b4d8 Pin torch version to 2.3.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:05:44 -04:00
makaveli10 fac62aaccc Fix hallucinations with no_speech_thres
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:05:12 -04:00
makaveli aade67736a Merge pull request #257 from sondt2709/fix-ffmpeg-subprocess-deadlock
Fix deadlock issue in FFmpeg subprocess by ensuring stderr is consumed
2024-07-19 14:49:10 +05:30
Sean Dang abfe830eee Fix deadlock issue in FFmpeg subprocess by ensuring stderr is consumed 2024-07-11 23:30:44 +07:00
makaveli cb392cbb93 Merge pull request #247 from makaveli10/pin_sliero_vad_model_version
Pin silero VAD onnx model version to v4.0
2024-07-09 12:58:27 +05:30
makaveli10 42733da59a Pin numpy version to <2
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-07-02 11:49:30 +05:30
makaveli10 26c517021f Pin silero VAD onnx model version to v4.0
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-07-02 11:01:40 +05:30
makaveli cf721e8b53 Merge pull request #243 from berkaybilik/making_backend_arg_safer
Making backend arg safer
2024-07-02 10:57:58 +05:30
makaveli 5985ec82b6 Merge pull request #236 from t-nil/patch-1
Backslash missing in example
2024-06-30 20:33:07 +05:30
berkaybilik 2f1c934ea2 always use the BackendType enum to reference the backend inside the TranscriptionServer 2024-06-27 00:20:17 +01:00
berkaybilik b220ccb330 fixed reference before assignment error/warning 2024-06-27 00:11:06 +01:00
berkaybilik 5e3906fc7b use enum to validate backend validity in server.run 2024-06-26 23:59:07 +01:00
Florian Meißner a8b9275013 Update README.md 2024-06-15 12:29:36 +02:00
makaveli 815441e8bb Bump version v0.5.0 2024-06-07 11:21:28 +05:30
makaveli 5b9bc2bc0e Merge pull request #223 from peldszus/single-model-mode
Single model mode
2024-06-07 11:10:52 +05:30
Marcus Edel ee132517fa Merge pull request #228 from anshulkharb/patch-1
fix spelling of detection in README.md.
2024-06-05 20:57:44 -04:00
Anshul Kharb 761bb61e87 fix spelling of detection in README.md 2024-06-05 23:14:13 +05:30
Andreas Peldszus 14077315ae Fix argparser option 2024-06-05 10:34:22 +02:00
Andreas Peldszus ab17c4dbc6 Make single model mode the default, update readme 2024-06-05 09:47:52 +02:00
makaveli 5e2421118d Merge pull request #227 from makaveli10/update_tensorrt_llm
Update tensorrt llm to v0.9.0
2024-06-05 08:50:26 +05:30
makaveli d1de2ec3ce Merge pull request #224 from chien-liu/expose-client-srt-location
Expose the srt file location of Transcription client
2024-06-03 21:26:07 +05:30
makaveli10 22a37e7843 Update ci to build and push teensorrt docker image
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-06-03 11:53:27 -04:00
makaveli10 e4579ef291 Dockerfile tensorrt use cuda-runtimee as base image to reduce size
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-06-03 06:44:09 -04:00
makaveli10 f73a146eb9 Update TensorRT backend tensorrt_llm==0.9.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-06-03 05:24:41 -04:00
chien-liu cfba5b3e54 Expose the srt file location of Transcription client 2024-06-01 00:14:59 +02:00
Andreas Peldszus 1ac7a278bb Update README 2024-05-31 15:07:01 +02:00
Andreas Peldszus 3a96f60006 Raise error for invalid model paths 2024-05-31 15:06:54 +02:00
Andreas Peldszus 3c09289dea Add single model mode for custom models
- Use a threadlock around the model in single model mode
2024-05-31 15:06:47 +02:00
makaveli e1a42c22d2 Merge pull request #216 from makaveli10/feature/writing_audio_frames_optional
Make writing audio frames optional
2024-05-29 09:10:43 +05:30
makaveli10 3d043dc906 Remove flake8 warning suppression
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 21:29:38 +05:30
makaveli10 399e9e7efe Fix README typo
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 11:18:35 +05:30
makaveli10 9d2ea75247 Refactor to make record function more readable
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 11:06:49 +05:30
makaveli10 225a98be0c Ignore linting as this file is a copy from faster_whisper
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 11:06:49 +05:30
makaveli10 8f373c3537 Update README
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 11:06:49 +05:30
makaveli10 61d07edabb Make writing output audio file optional when using microphone
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-05-28 11:06:49 +05:30
makaveli 03e30e1fed Merge pull request #212 from dshepelev15/feat/RTSP_support
Add support for RTSP stream
2024-05-28 10:48:13 +05:30
makaveli c0a947a8f6 Merge pull request #215 from makaveli10/fix/omp-num-threads
fix: limit CPU usage for VAD onnxruntime inference session by setting…
2024-05-24 23:42:51 +05:30
makaveli10 819ab35b28 fix: limit CPU usage for VAD onnxruntime inference session by setting OMP_NUM_THREADS
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-05-24 04:58:49 -04:00
dshepelev15 615c9c7aed Add support for RTSP stream 2024-05-17 17:32:09 +03:00
makaveli a9683319e0 Merge pull request #192 from fraic/dev1
Add option: save network stream to local file while transcribing
2024-05-05 17:46:39 +05:30
makaveli 0a2d92c5b8 Merge pull request #206 from peldszus/smaller-dockerimages
Improve cpu and gpu Dockerfiles, resulting in much smaller images
2024-05-02 11:02:16 +05:30
Andreas Peldszus dccfce2a3c Improve cpu and gpu Dockerfiles, resulting in much smaller images 2024-04-26 15:51:11 +02:00
fraic f78fc473c5 Add option: save network stream to local file while transcribing 2024-03-25 19:11:24 +08:00
makaveli 0dfbdb2477 bump version v0.4.1 2024-03-22 12:30:42 +05:30
makaveli e171b9c460 Merge pull request #190 from makaveli10/fix_client_close
Fix client close
2024-03-22 12:29:49 +05:30
makaveli10 66b5dc7c15 Merge remote-tracking branch 'upstream/main' into fix_client_close 2024-03-22 12:13:44 +05:30
makaveli10 0f1d36fc06 fix: microphone client close 2024-03-22 12:13:27 +05:30
makaveli d24c53198c Merge pull request #187 from jsichi/preserve-server-error
Don't clear server_error flag on close.
2024-03-22 12:12:16 +05:30
John Sichi fe1640695c Don't clear server_error flag on close. 2024-03-21 13:28:18 +09:00
makaveli 8d77f0fa5a bump version v0.4.0 2024-03-20 12:04:38 +05:30
makaveli 2e37216282 Merge pull request #174 from jsichi/tee-client
Add support for processing same audio stream via multiple clients running different tasks.
2024-03-17 22:51:26 +05:30
makaveli 7c7a446478 Fix: mock pyaudio for ci to pass the server tests 2024-03-15 12:29:32 +05:30
John Sichi 37d7f2ed66 Merge branch 'main' into tee-client 2024-03-13 20:12:51 +09:00
makaveli 754f22dfae Merge pull request #175 from FlippFuzz/fix-faster-whisper-version-setup
Fix faster whisper version in setup.py
2024-03-11 14:52:46 +05:30
makaveli ebd2dc9568 Merge pull request #173 from FlippFuzz/fix-os-error-no-mic
Handle failure on systems without microphones
2024-03-11 14:50:15 +05:30
FlippFuzz 9b2e17ec4d Fix faster whisper version in setup.py 2024-03-10 20:28:54 +08:00
John Sichi 5b32dc4130 Fix default value for multicast. 2024-03-10 21:02:50 +09:00
John Sichi c0f37c77e9 Remove camelcase 2024-03-10 20:49:17 +09:00
John Sichi 3b15dc76b4 Add support for processing same audio stream via multiple clients with different tasks. 2024-03-10 20:44:15 +09:00
FlippFuzz 4d477e35e7 Handle failure on systems without microphones
Catch the OSError and print a WARN log.
2024-03-10 19:43:53 +08:00
makaveli a17f4041de Merge pull request #163 from makaveli10/upgrade_faster_whisper
Upgrade faster whisper==1.0.1
2024-03-04 22:26:44 +05:30
makaveli10 8a06ba802b update cuda version 12.2.2 gpu dockerfile 2024-03-04 03:24:49 -05:00
makaveli10 02d4566289 upgrade faster whisper 1.0.1 2024-03-04 07:35:43 +00:00
Marcus Edel acd4902bec Merge pull request #161 from makaveli10/fix_docker_workflow
Build & push docker image on every new tag.
2024-02-29 10:42:49 -05:00
makaveli10 a495a49b06 build & push docker image on every new tag 2024-02-29 19:08:32 +05:30
makaveli 9e5ab408cd bump version 0.3.0 2024-02-28 23:39:44 +05:30
Marcus Edel 5e6c26c3a0 Merge pull request #158 from makaveli10/cpu_usage
fix: cpu usage issue.
2024-02-28 09:11:32 -05:00
makaveli10 18b6168807 fix: cpu usage issue 2024-02-28 13:55:37 +05:30
makaveli ec1349360a Merge pull request #157 from makaveli10/trt-multilingual
fix: lanuguage, task prefix in decoder start ids
2024-02-27 18:46:33 +05:30
makaveli10 a41e714801 fix: lanuguage, task prefix in decoder start ids 2024-02-26 23:31:19 -05:00
Marcus Edel 2d16ee552f Merge pull request #156 from makaveli10/fix_docker_image_gpu
Fix docker image gpu.
2024-02-26 09:24:26 -05:00
makaveli10 9699611000 push docker image to ghcr on push to main 2024-02-26 18:50:38 +05:30
makaveli10 ea64d47899 run server with python3 2024-02-26 18:50:18 +05:30
makaveli c067224474 bump version 0.2.1 2024-02-22 11:16:10 +05:30
makaveli e92f53cfd9 Update ci.yml
install wheel
2024-02-22 11:15:34 +05:30
makaveli 308ac1cff7 bump version 0.2.0 2024-02-22 11:03:30 +05:30
makaveli 2fced08705 Merge pull request #149 from makaveli10/docker-ghcr-ci
Docker ghcr ci
2024-02-22 10:14:51 +05:30
makaveli10 8bdaf9249d only run docker image build and push on new version release 2024-02-21 22:56:35 +05:30
makaveli10 b47a56ca6d update readme to use ghcr docker containers 2024-02-21 22:50:56 +05:30
makaveli10 f975bd452e change ghcr owner 2024-02-21 22:49:16 +05:30
makaveli10 cb963c4834 Merge remote-tracking branch 'upstream/main' into docker-ghcr-ci 2024-02-21 22:43:34 +05:30
makaveli 1db94ea96e Merge pull request #147 from makaveli10/vad_option
add VAD a client option
2024-02-21 16:10:54 +05:30
makaveli10 babe5de074 add vad option to firefox extension 2024-02-20 13:10:46 +05:30
makaveli10 99af50208d add vad option in chrome extension 2024-02-20 13:04:24 +05:30
makaveli10 dc22b7da9f add srt_file_path option 2024-02-20 11:53:45 +05:30
makaveli10 2e9f67ba0b Merge remote-tracking branch 'upstream/main' into vad_option 2024-02-20 11:43:56 +05:30
makaveli c919ba3501 Merge pull request #146 from makaveli10/code_formatting
Code formatting
2024-02-20 11:32:27 +05:30
makaveli10 a38fdb494d remove timeout from tests job 2024-02-20 00:28:38 +05:30
makaveli10 fddc244228 Merge remote-tracking branch 'upstream/main' into code_formatting 2024-02-19 22:02:40 +05:30
Marcus Edel 5e1174ff33 Merge pull request #136 from makaveli10/add_tests
Add tests.
2024-02-19 09:03:01 -05:00
Marcus Edel e40414ab1b Merge pull request #144 from makaveli10/update_readme
add whisper live demo video.
2024-02-19 09:02:27 -05:00
makaveli10 17873c66a0 prune docker cache 2024-02-19 06:36:10 -05:00
makaveli10 1147f58225 increase job timeout 2024-02-19 05:13:03 -05:00
makaveli10 0baa1dc0a6 update ci to build and gpu docker image to ghcr 2024-02-19 04:53:39 -05:00
makaveli10 d1de4948ee update base cuda version to 11.8; some dockerfile-gpu fixes 2024-02-19 04:38:41 -05:00
makaveli10 ff871ad485 update dockerfile name 2024-02-16 20:37:59 +05:30
makaveli10 5fe5e0c8ba Merge branch 'vad_option' into develop 2024-02-16 20:32:01 +05:30
makaveli10 b42ced9816 fix: tests for end of speech message while mocking pyaudio 2024-02-16 20:31:38 +05:30
makaveli10 06794470f8 build docker image on pus develop 2024-02-16 19:01:54 +05:30
makaveli10 d530957b2c test docker ci on fork 2024-02-16 18:58:57 +05:30
makaveli10 6cabbe441b update cpu dockerfile with python-slim-buster base image 2024-02-16 18:58:38 +05:30
makaveli10 78da3f6750 Merge branch 'code_formatting' into vad_option 2024-02-16 17:29:37 +05:30
makaveli10 b04cffc458 update readme; remove common content 2024-02-16 17:11:56 +05:30
makaveli10 fd7c5965b3 add whisper live demo video 2024-02-16 13:28:02 +05:30
makaveli10 4471665085 remove test audio from tests 2024-02-15 19:09:57 +05:30
makaveli10 147e97002e clear_screen for updated transcript 2024-02-15 18:56:14 +05:30
makaveli10 e3c7666cf7 update readme with use_vad 2024-02-15 18:13:26 +05:30
makaveli10 8266099ed0 update tensorrt readme 2024-02-15 18:08:02 +05:30
makaveli10 01dc69e068 close when end of audio from client 2024-02-15 18:07:19 +05:30
makaveli10 9bb92b9bb2 use_vad option and send end of audio message 2024-02-15 17:59:12 +05:30
makaveli10 57c4b60e04 remove websocket.path log from exception logging 2024-02-15 15:08:08 +05:30
makaveli10 3cd96367fb make vad an option 2024-02-15 14:59:43 +05:30
makaveli10 c1420cba0d add tests for server exception handling 2024-02-15 12:16:58 +05:30
makaveli10 4db91eed66 update vad tests after refactor 2024-02-15 12:16:39 +05:30
makaveli10 7bcb92c266 create new method for handling a new connection; expcetion handling 2024-02-15 12:16:18 +05:30
makaveli10 170ba22e5b update method docstrings 2024-02-09 16:08:18 +05:30
makaveli10 ac00e28b86 add: VoiceActivityDetector to manage vad 2024-02-09 16:07:43 +05:30
makaveli10 ceb3cc8747 update timeout log 2024-02-09 14:20:47 +05:30
makaveli10 eaec0ead08 add: handle_transcription_output method 2024-02-09 14:19:49 +05:30
makaveli10 9fbff47126 🔨 refactor whisper_live according to flake8 2024-02-09 13:45:13 +05:30
makaveli10 b4abe95fc6 add: code-format job 2024-02-09 13:44:17 +05:30
makaveli10 14974af951 update ci to run tests 2024-02-08 14:07:44 +05:30
makaveli10 bc474b4a76 update on_close; on_error 2024-02-08 14:06:06 +05:30
makaveli10 9ccf940f51 remove debug import excpetion tensorrt llm 2024-02-08 14:05:46 +05:30
makaveli10 9a9972007e remove debug stats 2024-02-08 14:04:52 +05:30
makaveli10 b2ad6478f5 update requirement for tests 2024-02-08 14:04:32 +05:30
makaveli10 490efdeacc mv test audio to assets 2024-02-08 14:04:16 +05:30
makaveli10 cb570d28ce add vad tests 2024-02-08 14:03:57 +05:30
makaveli10 4e5e086c38 add server tests 2024-02-08 14:03:39 +05:30
makaveli10 cf78d5d608 add client tests 2024-02-08 14:03:13 +05:30
makaveli 9d29b08cea Merge pull request #135 from collabora/revert-134-test_pypi_upload
Revert "Test pypi upload"
2024-02-08 12:23:36 +05:30
makaveli 6071cc1cc5 Revert "Test pypi upload" 2024-02-08 12:23:15 +05:30
makaveli f98e309663 Merge pull request #134 from makaveli10/test_pypi_upload
Test pypi upload
2024-02-08 12:23:07 +05:30
makaveli10 30b00d6c89 upload to testpypi 2024-02-08 12:18:50 +05:30
makaveli10 da2992bcaf add tests to ci.yml 2024-02-08 11:42:39 +05:30
makaveli10 16c5ed8ce9 add more python versions 2024-02-08 11:26:17 +05:30
makaveli10 e14fefb671 cache req 2024-02-08 11:13:12 +05:30
makaveli10 98399707a3 add pyaudio mock 2024-02-08 11:12:57 +05:30
makaveli10 567ceb1246 add pyaudio mock; refactor 🔨 2024-02-08 11:12:40 +05:30
makaveli10 28ea8a20f1 update tests ci 2024-02-07 23:54:02 +05:30
makaveli10 acf6dfe5b7 update python version 2024-02-07 23:47:54 +05:30
makaveli10 84a97f5fdd remove whisper_live from patch to mock websocket 2024-02-07 23:47:29 +05:30
makaveli10 ca2634bbb6 add tests workflow 2024-02-07 23:33:49 +05:30
makaveli10 444ce63440 add unit tests 2024-02-07 23:33:24 +05:30
makaveli10 8db063ee33 update log level to info 2024-02-07 23:31:58 +05:30
makaveli10 92cbc37e9c remove debug stats vad 2024-02-07 23:31:12 +05:30
makaveli10 24fd835356 update requirements 2024-02-07 23:30:34 +05:30
makaveli10 5409d14bcb move audio files to assets 2024-02-07 23:30:16 +05:30
makaveli10 d6edf8e847 update on_error; on_close 2024-02-07 23:29:22 +05:30
makaveli10 cc3ed74c0e remove unused imports 2024-02-07 23:26:42 +05:30
makaveli10 20a8a8ad3d update log level to warning 2024-02-07 23:26:13 +05:30
makaveli10 07387abbc0 silence WhisperTRTLLM import warning 2024-02-07 23:25:50 +05:30
makaveli10 4ecc59783e bump version v0.1.0 2024-02-05 22:37:11 +05:30
makaveli ec9074d712 Merge pull request #128 from lightwastak3n/firefox_remove_multilingual
Firefox remove multilingual
2024-02-05 21:42:10 +05:30
Marcus Edel 3a25db4cb9 Merge pull request #127 from makaveli10/update_setup_requirements
Update required packages for setup.
2024-02-05 08:11:01 -05:00
makaveli 4924ec0adb remove empty lines 2024-02-05 18:40:03 +05:30
makaveli f35abc7f81 Merge branch 'collabora:main' into update_setup_requirements 2024-02-05 18:26:38 +05:30
Sasa Trivic f383121ec3 Remove multilingual option description from the extension readme 2024-02-04 18:03:41 +01:00
Sasa Trivic 17d62272cf Update README.md
Remove multilingual from README
2024-02-04 17:52:21 +01:00
Sasa Trivic 91e1b75bfc Merge branch 'collabora:main' into firefox_remove_multilingual 2024-02-04 17:36:57 +01:00
Sasa Trivic 7aad2ae721 Remove multilingual from Firefox. Sort languages, disable all inputs when capturing. Move both transcripts to the bottom center. 2024-02-04 17:35:20 +01:00
makaveli10 6a1b82f953 update required packages for setup 2024-02-04 16:02:02 +05:30
makaveli dc84839873 Merge pull request #126 from makaveli10/fix_typo_multilingual
fix: typo; remove multilingual debug stat
2024-02-04 14:38:54 +05:30
makaveli10 56d19f5469 fix: typo; remove multilingual debug stat 2024-02-04 14:34:16 +05:30
makaveli 08fa183ba4 Merge pull request #123 from lightwastak3n/main
Chrome extension update; remove multililingual option faster-whisper
2024-02-03 19:24:45 +05:30
Marcus Edel d89b27b8aa Merge pull request #124 from Stinosko/patch-1
Add scipy to server.txt.
2024-02-02 15:31:35 -05:00
Stinosko ce68cc6c87 Add scipy to server.txt
The server script uses scipy but is not installed with the current server requirements file.
2024-02-02 21:07:57 +01:00
Sasa Trivic e697574870 Remove multilingual from client. Remove multilingual from faster whisper backend. Disable task dropdown when capturing in chrome extension. 2024-02-02 13:56:26 +01:00
Sasa Trivic b098b52a4d Merge branch 'main' of github.com:lightwastak3n/WhisperLive 2024-02-01 17:37:18 +01:00
Sasa Trivic ad5543b03e Merge remote-tracking branch 'upstream/main' 2024-02-01 17:31:35 +01:00
Marcus Edel 60455b1583 Merge pull request #121 from makaveli10/save_transcript
Save transcript.
2024-02-01 11:11:05 -05:00
Sasa Trivic cb458fc207 Merge pull request #1 from lightwastak3n/extension_rewrite
Extension rewrite
2024-02-01 15:16:47 +01:00
Sasa Trivic 32ed089a76 Change faster whisper to work with new extension 2024-02-01 14:54:29 +01:00
Sasa Trivic 5b28ddefbd Chrome extension - QOL. Remove multilingual part. 2024-02-01 14:44:07 +01:00
Sasa Trivic e1f531eccf Remove duplicate assignment 2024-02-01 13:56:09 +01:00
Sasa Trivic 8200207530 Center transcription div 2024-02-01 13:16:23 +01:00
makaveli10 08575a03c2 write srt file only for faster_whisper backend 2024-02-01 14:22:45 +05:30
makaveli10 f590446865 Merge remote-tracking branch 'upstream/main' into save_transcript 2024-02-01 12:14:09 +05:30
Sasa Trivic 36d137888e Merge branch 'collabora:main' into main 2024-01-31 18:03:29 +01:00
Marcus Edel e64bc9f3d6 Merge pull request #116 from makaveli10/tensorrt_model_warmup
Tensorrt model warmup.
2024-01-31 11:49:40 -05:00
Sasa Trivic 7cc945aded self.client_uid accessed without being defined 2024-01-31 16:45:31 +01:00
makaveli 2c8a25d355 Merge pull request #119 from gchust/main
fix: keyError: 'model' in server, when using browser extension
2024-01-31 18:40:27 +05:30
makaveli10 f4027de343 add: save_transcript to srt file 2024-01-31 17:37:30 +05:30
makaveli10 d1754d2c46 fix: model_size, no_speech, segment timings 2024-01-31 17:37:03 +05:30
gchust 0e6b1c0632 fix: keyError: 'model' in server, when using browser extension 2024-01-31 11:33:48 +00:00
makaveli d6b51ccd7d Update TensorRT_whisper.md
typo: setup.sh file path
2024-01-29 12:33:12 +05:30
makaveli 2f3c1cd172 Update TensorRT_whisper.md
fix: typo in bash script name to build tensorrt engine
2024-01-29 12:27:51 +05:30
makaveli 703263b375 wamrup tensorrt engine 2024-01-29 12:17:27 +05:30
makaveli 30d2cffb93 load audio for warmup 2024-01-29 12:15:39 +05:30
makaveli 4d94c6b38b Update server.txt
install ffmpeg-python to load test file for model warmup
2024-01-29 12:10:18 +05:30
makaveli d5a0f5859e Update TensorRT_whisper.md
ffmpeg is needed for model warmup
2024-01-29 12:06:56 +05:30
makaveli 025873d2ca Update TensorRT server requirements 2024-01-29 11:57:03 +05:30
makaveli 8c36768f7f Merge pull request #112 from lightwastak3n/main
Readme: Fix transcribe examples
2024-01-26 00:27:43 +05:30
Sasa Trivic ce13e7b622 Fix transcribe examples 2024-01-25 18:44:23 +01:00
makaveli 3498787ccd Merge pull request #104 from makaveli10/tensorrt_backend
Tensorrt backend
2024-01-24 16:42:18 +05:30
makaveli 5cd59b1e4c Update README.md
Co-authored-by: Marcus Edel <marcus.edel@fu-berlin.de>
2024-01-24 10:09:22 +05:30
makaveli10 bd543295f3 update readme 2024-01-22 11:54:21 +00:00
makaveli10 8e2642283a update tensorrt readme 2024-01-22 11:48:23 +00:00
makaveli10 3bf5b47947 fix: server; remove debug stats 2024-01-22 11:47:01 +00:00
makaveli10 634dae835b add numba to req(trt-llm) 2024-01-22 11:45:11 +00:00
makaveli10 969a5aa9e5 remove trt_llm install script 2024-01-22 11:44:48 +00:00
makaveli10 44a2e20c68 remove trt-llm dockerfile 2024-01-22 11:44:21 +00:00
makaveli10 986823dbef Merge remote-tracking branch 'upstream/main' into tensorrt_backend 2024-01-19 10:45:01 -05:00
makaveli10 1e2faa3f2b fix: tensorrt llm idocker setup & docs 2024-01-19 10:40:59 -05:00
makaveli10 e3084b34cb update tensorrt docker & readme 2024-01-19 07:30:50 -05:00
makaveli10 b955e63dc1 update READM 2024-01-19 12:02:18 +00:00
makaveli10 f25ff1785a increase chunk size from 64ms to 256ms 2024-01-19 12:02:05 +00:00
makaveli10 867ff522ae add tensorrt installation & whisper conversion script 2024-01-19 11:59:16 +00:00
makaveli10 75001ae6b7 updatetensorrt-llm dockerfile 2024-01-19 11:58:40 +00:00
makaveli10 6f1d13f25b update requirements 2024-01-19 11:57:46 +00:00
makaveli10 7a9dc6db40 add tensorrt readme 2024-01-19 11:53:47 +00:00
makaveli10 735d6c7763 merge with main 2024-01-19 11:43:21 +00:00
Marcus Edel 0942dc2cfd Merge pull request #102 from makaveli10/change_model_size_param_name
Server to control custom model usage.
2024-01-18 11:07:18 -05:00
makaveli10 881fd55776 run server with custom model from args 2024-01-18 15:02:51 +08:00
makaveli 0c01d7b1e5 Merge pull request #98 from makaveli10/change_model_size_param_name
Change model size param name
2024-01-15 21:08:26 +05:30
makaveli10 c810369324 revert the default port of chrom/firefox extension to 9090 2024-01-15 23:30:05 +08:00
makaveli10 71d0fe69c6 add option to use custom model 2024-01-15 23:28:02 +08:00
makaveli10 67232fffd5 install whl 2024-01-12 11:39:05 +00:00
makaveli10 076aebf3b6 bump version v0.0.11 2024-01-12 15:21:27 +05:30
makaveli10 4cf9d95f73 merge main 2024-01-12 08:17:46 +00:00
makaveli10 389bb5ae37 add docker setup for tensorrt-llm; update readme 2024-01-12 08:15:53 +00:00
Marcus Edel a7eedc5d84 Merge pull request #94 from makaveli10/fix_error_messages
Fix: error messages.
2024-01-11 09:52:07 -05:00
makaveli d91330d790 Merge branch 'collabora:main' into fix_error_messages 2024-01-11 18:18:58 +05:30
makaveli 783d147316 Merge pull request #96 from makaveli10/fix_key_error
fix: key error
2024-01-11 16:29:07 +05:30
makaveli10 058c93e55e fix: key error 2024-01-11 18:55:43 +08:00
makaveli10 f06b9bc827 remove torch req 2024-01-11 08:18:51 +00:00
makaveli10 3c202bf836 update README; add TensorRT doc 2024-01-11 08:18:25 +00:00
makaveli10 647c576e6a update with multilingual option 2024-01-11 08:17:56 +00:00
makaveli10 71a062b726 update dockerfiles 2024-01-10 14:30:57 +00:00
makaveli10 a26f990586 update readme to new setup.sh path 2024-01-10 14:22:51 +00:00
makaveli10 ddb1e0947f move setup.sh to scripts 2024-01-10 14:22:20 +00:00
makaveli10 6dff4fbdd3 add tensorrt_llm installation script 2024-01-10 14:21:39 +00:00
makaveli10 244ca9e6ba remove duplicate code 2024-01-10 14:14:48 +00:00
makaveli10 0f9e93d203 add: tensorrt backend to server 2024-01-09 18:10:17 +00:00
makaveli10 fd86340f30 add: tensorrt backend 2024-01-09 18:09:50 +00:00
makaveli10 2300eedc8b fix: error messages 2024-01-09 13:35:38 +08:00
makaveli cafcb04fbc Merge pull request #92 from hcljsq/main
feat: set initial_prompt and vad_parameters in the first message
2024-01-09 10:00:48 +05:30
Chen Hua 72ead71eeb feat: set initial_prompt and vad_parameters in the first message 2024-01-08 15:33:55 +08:00
makaveli 7b2f5cff72 Merge pull request #89 from hcljsq/main
format segment timestamps
2024-01-03 21:27:06 +05:30
Chen Hua 32c6a565d7 refactor(server): add format_segment helper to standardize timestamp output 2024-01-03 12:25:25 +08:00
makaveli e30286c046 Merge pull request #83 from Chronoz/fix_exception_on_overflow
fix exception on overflow
2024-01-02 19:01:36 +05:30
makaveli10 7c0b32b85e bump version 0.0.10 2024-01-01 18:12:51 +05:30
makaveli 01665a54c1 Merge pull request #82 from hcljsq/main
Add the `large-v3` model
2024-01-01 18:09:41 +05:30
华晨 02793a93f8 feat: Update transcriber to support large-v3 model with 128 mel filters 2024-01-01 20:22:37 +08:00
makaveli db2e0bbcdd Merge pull request #81 from k0hacuu/main
README Spelling correction
2024-01-01 13:53:38 +05:30
Chronoz e92ddd291a fix exception on overflow 2023-12-31 20:59:33 +07:00
华晨 5918b5ed42 fix: update faster-whisper 2023-12-31 15:45:41 +08:00
华晨 71d207a607 feat: add large-v3 model 2023-12-30 23:19:30 +08:00
Jonny Yang 6ee4cd09f2 Update README.md 2023-12-25 13:12:24 +00:00
Marcus Edel 5de4de4b84 Merge pull request #76 from makaveli10/model_size_option
Model size option.
2023-12-20 09:10:57 -05:00
makaveli10 e006722da7 add model size option to client 2023-12-20 18:06:02 +05:30
makaveli10 a52dc0cbf8 update chrome/firefox plugin readme 2023-12-15 00:09:55 +05:30
makaveli10 048ab0a8f4 remove debugging script 2023-12-15 02:26:01 +08:00
makaveli10 261bb9e961 Merge branch 'model_size_option' of github.com:makaveli10/whisper-live into model_size_option 2023-12-15 02:23:20 +08:00
makaveli10 091f6179d4 add sample audio for testing 2023-12-15 02:23:13 +08:00
makaveli10 1e1349cd80 update firefox plugin with model size dropdown 2023-12-14 23:52:51 +05:30
makaveli10 14beb4f942 update chrome plugin with model size dropdown 2023-12-14 23:52:25 +05:30
makaveli10 7ffcad64ba update readme 2023-12-15 02:21:35 +08:00
makaveli10 402fceb9f3 remove emptyline 2023-12-15 02:21:09 +08:00
makaveli10 09b18e8ab8 add model size option from server 2023-12-14 22:08:12 +08:00
makaveli da72d03073 Merge pull request #73 from jhormigo/hls_support
Support for HLS transcription
2023-12-13 00:05:11 +05:30
Jesús Hormigo a1a8d5f92a Added a HLS stream sample URL 2023-12-12 13:46:52 +01:00
Jesús Hormigo b6dee4e46e Using ffmpeg-python package instead of requiring having ffmpeg installed in system 2023-12-10 19:34:27 +01:00
Jesús Hormigo f3cd20fbf3 Support for HLS transcription (resolves #62) 2023-12-09 21:24:55 +01:00
makaveli10 da86c18205 bump version: 0.0.9 2023-12-06 15:09:50 +05:30
makaveli 8097e9b44a Merge pull request #70 from ethanzrd/main
Update `faster_whisper` version setup.py
2023-12-06 15:04:06 +05:30
Ethan Zerad 222852ff33 Update setup.py
Update faster-whisper version to match requirements/server.txt
2023-12-02 11:07:41 +02:00
82 changed files with 13886 additions and 2442 deletions
+208 -36
View File
@@ -1,4 +1,4 @@
name: CI
name: Test & Build CI/CD
on:
push:
@@ -7,46 +7,218 @@ 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 }}
venv-install-smoke-test:
runs-on: ubuntu-22.04
steps:
- uses: actions/checkout@v2
- name: Set up Python 3.12
uses: actions/setup-python@v2
with:
python-version: '3.12'
- name: Install system dependencies
run: sudo apt-get update && sudo apt-get install -y portaudio19-dev
- name: Build package artifacts
run: |
python -m pip install --upgrade pip
python -m pip install build
python -m build --sdist --wheel
- name: Verify install in a clean virtualenv
run: |
python -m venv smoke-test-venv
source smoke-test-venv/bin/activate
pip install dist/*.whl
python -c "import whisper_live.client; import whisper_live.server"
build-and-push-docker-cpu:
needs: [run-tests, check-code-format, venv-install-smoke-test]
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, venv-install-smoke-test]
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 }}
+25
View File
@@ -0,0 +1,25 @@
__pycache__/
*.pyc
*.pyo
*.egg-info/
dist/
build/
*.egg
.eggs/
whisper_env/
venv/
.venv/
env/
.env
*.so
*.o
.pytest_cache/
.mypy_cache/
.ruff_cache/
output*.srt
transcript*.srt
translation*.srt
*.wav
docs/site/
Audio-Transcription-Chrome/node_modules/
Audio-Transcription-Chrome/package-lock.json
+2 -1
View File
@@ -26,9 +26,10 @@ 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.
- **Download SRT file at Stop Capture**: Select if you want to download the srt file for the session at stop capture.
- **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).
@@ -0,0 +1,242 @@
'use strict';
// ---------------------------------------------------------------------------
// Chrome API mock — defined before any extension script is loaded
// ---------------------------------------------------------------------------
const storageData = {};
global.chrome = {
storage: {
local: {
get: jest.fn((keys, cb) => {
if (typeof keys === 'string') {
cb({ [keys]: storageData[keys] });
} else if (Array.isArray(keys)) {
const result = {};
keys.forEach(k => { result[k] = storageData[k]; });
cb(result);
} else {
const result = {};
Object.keys(keys).forEach(k => {
result[k] = k in storageData ? storageData[k] : keys[k];
});
cb(result);
}
}),
set: jest.fn((obj, cb) => {
Object.assign(storageData, obj);
if (cb) cb();
}),
},
},
runtime: {
sendMessage: jest.fn(),
onMessage: { addListener: jest.fn() },
getURL: jest.fn(path => `chrome-extension://fake-id/${path}`),
id: 'fake-extension-id',
},
tabs: {
query: jest.fn(),
get: jest.fn(),
create: jest.fn(),
remove: jest.fn(),
sendMessage: jest.fn(),
},
tabCapture: { capture: jest.fn() },
scripting: { executeScript: jest.fn() },
};
// Flush all pending microtasks and one macrotask round so async click
// handlers (which await at least one Promise inside) complete fully.
const flushPromises = () => new Promise(resolve => setTimeout(resolve, 0));
function resetStorage() {
Object.keys(storageData).forEach(k => delete storageData[k]);
}
function buildPopupDOM() {
document.body.innerHTML = `
<div id="startCapture"></div>
<div id="stopCapture"></div>
<input type="checkbox" id="useServerCheckbox">
<input type="checkbox" id="useVadCheckbox">
<input type="checkbox" id="saveCaptionsCheckbox">
<select id="languageDropdown">
<option value="" selected></option>
<option value="en">English</option>
</select>
<select id="taskDropdown">
<option value="transcribe" selected>Transcribe</option>
</select>
<select id="modelSizeDropdown">
<option value="small" selected>Small</option>
<option value="large-v3">Large-v3</option>
</select>
<select id="captionLinesDropdown">
<option value="3" selected>3</option>
</select>
`;
}
function loadPopup() {
jest.resetModules();
require('../popup.js');
document.dispatchEvent(new Event('DOMContentLoaded'));
}
function clickStart(tabId = 1) {
chrome.tabs.query.mockImplementation((_, cb) => cb([{ id: tabId }]));
document.getElementById('startCapture').click();
}
// ---------------------------------------------------------------------------
// 1. WebSocket URL construction — pure logic extracted from options.js
// ---------------------------------------------------------------------------
describe('WebSocket URL construction', () => {
function buildWsUrl(host, port) {
return port
? `ws://${host}:${port}/`
: `wss://${host}/ws`;
}
test('local server: ws:// with port and trailing slash', () => {
expect(buildWsUrl('localhost', '9090')).toBe('ws://localhost:9090/');
});
test('Modal service (empty port): wss:// with /ws path', () => {
expect(buildWsUrl('boxerab--aavaaz-live-livetranscriber-web.modal.run', '')).toBe(
'wss://boxerab--aavaaz-live-livetranscriber-web.modal.run/ws'
);
});
test('custom host+port stays on ws://', () => {
expect(buildWsUrl('my-server.example.com', '7090')).toBe(
'ws://my-server.example.com:7090/'
);
});
});
// ---------------------------------------------------------------------------
// 2. popup.js — host/port selection based on checkbox
// ---------------------------------------------------------------------------
describe('popup.js host/port selection', () => {
beforeEach(() => {
resetStorage();
buildPopupDOM();
jest.clearAllMocks();
});
test('checkbox unchecked → localhost:9090', async () => {
document.getElementById('useServerCheckbox').checked = false;
loadPopup();
clickStart();
await flushPromises();
const call = chrome.runtime.sendMessage.mock.calls.find(
c => c[0] && c[0].action === 'startCapture'
);
expect(call).toBeDefined();
expect(call[0].host).toBe('localhost');
expect(call[0].port).toBe('9090');
});
test('checkbox checked → Modal host with empty port', async () => {
document.getElementById('useServerCheckbox').checked = true;
loadPopup();
clickStart();
await flushPromises();
const call = chrome.runtime.sendMessage.mock.calls.find(
c => c[0] && c[0].action === 'startCapture'
);
expect(call).toBeDefined();
expect(call[0].host).toBe('boxerab--aavaaz-live-livetranscriber-web.modal.run');
expect(call[0].port).toBe('');
});
test('startCapture message carries language, task, and modelSize', async () => {
document.getElementById('languageDropdown').value = 'en';
document.getElementById('taskDropdown').value = 'transcribe';
document.getElementById('modelSizeDropdown').value = 'large-v3';
loadPopup();
clickStart();
await flushPromises();
const call = chrome.runtime.sendMessage.mock.calls.find(
c => c[0] && c[0].action === 'startCapture'
);
expect(call).toBeDefined();
expect(call[0].task).toBe('transcribe');
expect(call[0].modelSize).toBe('large-v3');
});
});
// ---------------------------------------------------------------------------
// 3. popup.js — button state management
// ---------------------------------------------------------------------------
describe('popup.js button state', () => {
beforeEach(() => {
resetStorage();
buildPopupDOM();
jest.clearAllMocks();
});
test('start button enabled and stop button disabled on load (not capturing)', () => {
storageData.capturingState = { isCapturing: false };
loadPopup();
expect(document.getElementById('startCapture').disabled).toBeFalsy();
expect(document.getElementById('stopCapture').disabled).toBe(true);
});
test('stop button enabled on load when already capturing', () => {
storageData.capturingState = { isCapturing: true };
loadPopup();
expect(document.getElementById('stopCapture').disabled).toBe(false);
});
test('clicking stop sends stopCapture action to runtime', () => {
storageData.capturingState = { isCapturing: true };
loadPopup();
// toggleCaptureButtons(true) disables startCapture and enables stopCapture
document.getElementById('stopCapture').click();
expect(chrome.runtime.sendMessage).toHaveBeenCalledWith(
expect.objectContaining({ action: 'stopCapture' }),
expect.any(Function)
);
});
});
// ---------------------------------------------------------------------------
// 4. popup.js — storage state restoration on open
// ---------------------------------------------------------------------------
describe('popup.js storage restoration', () => {
beforeEach(() => {
resetStorage();
buildPopupDOM();
jest.clearAllMocks();
});
test('restores useServerCheckbox from storage', () => {
storageData.useServerState = true;
loadPopup();
expect(document.getElementById('useServerCheckbox').checked).toBe(true);
});
test('restores language selection from storage', () => {
storageData.selectedLanguage = 'en';
loadPopup();
expect(document.getElementById('languageDropdown').value).toBe('en');
});
test('restores model size from storage', () => {
storageData.selectedModelSize = 'large-v3';
loadPopup();
expect(document.getElementById('modelSizeDropdown').value).toBe('large-v3');
});
});
@@ -0,0 +1,77 @@
class AudioPreProcessor extends AudioWorkletProcessor {
constructor() {
super();
this.sampleRate = sampleRate || 48000;
this.targetSampleRate = 16000;
this.inputSamplesNeeded = this.sampleRate * 0.5; // 0.5s
this.inputBuffer = new Float32Array(this.inputSamplesNeeded);
this.inputWriteOffset = 0;
this.processCount = 0;
this.audioDetectedCount = 0;
}
process(inputs, outputs) {
this.processCount++;
const input = inputs[0];
const output = outputs[0];
if (!input || input.length === 0) {
return true;
}
for (let channel = 0; channel < Math.min(input.length, output.length); channel++) {
if (input[channel] && output[channel]) {
output[channel].set(input[channel]);
}
}
let monoInput;
if (input.length === 1) {
monoInput = input[0];
} else if (input.length >= 2) {
monoInput = new Float32Array(input[0].length);
for (let i = 0; i < input[0].length; i++) {
monoInput[i] = (input[0][i] + (input[1] ? input[1][i] : 0)) * 0.5;
}
} else {
return true;
}
if (!monoInput || monoInput.length === 0) {
return true;
}
let inputOffset = 0;
while (inputOffset < monoInput.length) {
const remainingBuffer = this.inputSamplesNeeded - this.inputWriteOffset;
const toCopy = Math.min(remainingBuffer, monoInput.length - inputOffset);
this.inputBuffer.set(monoInput.subarray(inputOffset, inputOffset + toCopy), this.inputWriteOffset);
this.inputWriteOffset += toCopy;
inputOffset += toCopy;
if (this.inputWriteOffset === this.inputSamplesNeeded) {
const downsampled = this.downsampleTo16kHz(this.inputBuffer);
this.port.postMessage(downsampled);
this.inputWriteOffset = 0;
}
}
return true;
}
downsampleTo16kHz(inputBuffer) {
const ratio = this.sampleRate / this.targetSampleRate;
const length = Math.floor(inputBuffer.length / ratio);
const result = new Float32Array(length);
for (let i = 0; i < length; i++) {
const idx = Math.floor(i * ratio);
result[i] = inputBuffer[idx];
}
return result;
}
}
registerProcessor('audiopreprocessor', AudioPreProcessor);
+8 -5
View File
@@ -156,7 +156,10 @@ 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,
saveCaptions: options.saveCaptions,
},
});
} else {
@@ -172,14 +175,14 @@ async function startCapture(options) {
* Stops the capture process and performs cleanup.
* @returns {Promise<void>} - A Promise that resolves when the capture process is stopped successfully.
*/
async function stopCapture() {
async function stopCapture(options) {
const optionTabId = await getLocalStorageValue("optionTabId");
const currentTabId = await getLocalStorageValue("currentTabId");
if (optionTabId) {
res = await sendMessageToTab(currentTabId, {
type: "STOP",
data: { currentTabId: currentTabId },
data: { currentTabId: currentTabId, saveCaptions: options.saveCaptions },
});
await removeChromeTab(optionTabId);
}
@@ -194,7 +197,7 @@ chrome.runtime.onMessage.addListener(async (message) => {
if (message.action === "startCapture") {
startCapture(message);
} else if (message.action === "stopCapture") {
stopCapture();
stopCapture(message);
} else if (message.action === "updateSelectedLanguage") {
const detectedLanguage = message.detectedLanguage;
chrome.runtime.sendMessage({ action: "updateSelectedLanguage", detectedLanguage });
@@ -202,7 +205,7 @@ chrome.runtime.onMessage.addListener(async (message) => {
} else if (message.action === "toggleCaptureButtons") {
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false });
chrome.storage.local.set({ capturingState: { isCapturing: false } })
stopCapture();
stopCapture({saveCaptions: message.saveCaptions});
}
});
+137 -58
View File
@@ -1,10 +1,46 @@
var elem_container = null;
var elem_text = null;
var segments = [];
var text_segments = [];
var captionLineCount = 3;
var allSegments = [];
var lastIncompleteSegment = null;
function formatTime(seconds) {
const date = new Date(seconds * 1000);
const hh = String(date.getUTCHours()).padStart(2, '0');
const mm = String(date.getUTCMinutes()).padStart(2, '0');
const ss = String(date.getUTCSeconds()).padStart(2, '0');
const mmm = String(date.getUTCMilliseconds()).padStart(3, '0');
return `${hh}:${mm}:${ss},${mmm}`;
}
function generateSRT() {
return allSegments
.map((seg, i) => {
const start = formatTime(seg.start);
const end = formatTime(seg.end);
const text = seg.text.trim().replace(/[\r\n]+/g, ' ');
return `${i + 1}\n${start} --> ${end}\n${text}`;
})
.join('\n\n');
}
function downloadSRT() {
console.log("downloadSRT called");
console.log("Total segments for SRT:", allSegments.length);
const srtBlob = new Blob([generateSRT()], { type: 'text/srt;charset=utf-8' });
const url = URL.createObjectURL(srtBlob);
const a = document.createElement('a');
a.href = url;
a.download = 'captions.srt';
a.style.display = 'none';
document.body.appendChild(a);
a.click();
URL.revokeObjectURL(url);
document.body.removeChild(a);
}
function initPopupElement() {
if (document.getElementById('popupElement')) {
@@ -32,7 +68,7 @@ function initPopupElement() {
closePopupButton.style.cursor = 'pointer';
closePopupButton.addEventListener('click', async () => {
popupContainer.style.display = 'none';
await browser.runtime.sendMessage({ action: 'toggleCaptureButtons', data: false });
await chrome.runtime.sendMessage({ action: 'toggleCaptureButtons', data: false });
});
buttonContainer.appendChild(closePopupButton);
popupContainer.appendChild(buttonContainer);
@@ -52,22 +88,23 @@ function showPopup(customText) {
}
function init_element() {
function init_element(lines = 3) {
captionLineCount = Math.min(Math.max(parseInt(lines, 10) || 3, 1), 8);
if (document.getElementById('transcription')) {
return;
}
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:0 24px;font-family:Arial,Helvetica,sans-serif;font-size:22px;font-weight:600;line-height:30px;position:fixed;top:85%;left:50%;transform:translate(-50%,-50%);width:min(80vw,900px);min-height:' + (captionLineCount * 30) + 'px;z-index:2147483647;color:white;text-align:center;letter-spacing:0.01em;text-shadow:0 0 2px #000,0 2px 4px rgba(0,0,0,0.95);cursor:move;';
for (var i = 0; i < 4; i++) {
for (var i = 0; i <= captionLineCount; i++) {
elem_text = document.createElement('span');
elem_text.style.cssText = 'position: absolute;padding-left:16px;padding-right:16px;';
elem_text.style.cssText = 'position:absolute;left:50%;transform:translateX(-50%);max-width:100%;padding:2px 14px;background:rgba(0,0,0,0.72);border-radius:6px;box-decoration-break:clone;-webkit-box-decoration-break:clone;';
elem_text.id = "t" + i;
elem_container.appendChild(elem_text);
if (i == 3) {
if (i == captionLineCount) {
elem_text.style.top = "-1000px"
}
}
@@ -129,7 +166,7 @@ function get_lines(elem, line_height) {
var divHeight = elem.offsetHeight;
var lines = divHeight / line_height;
var original_text = elem.innerHTML;
var original_text = elem.textContent;
var words = original_text.split(' ');
var segments = [];
@@ -139,7 +176,7 @@ function get_lines(elem, line_height) {
for (var i = 0; i < words.length; i++)
{
segment += words[i] + ' ';
elem.innerHTML = segment;
elem.textContent = segment;
divHeight = elem.offsetHeight;
if ((divHeight / line_height) > current_lines) {
@@ -153,7 +190,7 @@ function get_lines(elem, line_height) {
var line_segment = segment.substring(segment_len, segment.length - 1)
segments.push(line_segment);
elem.innerHTML = original_text;
elem.textContent = original_text;
return segments;
@@ -161,7 +198,7 @@ function get_lines(elem, line_height) {
function remove_element() {
var elem = document.getElementById('transcription')
for (var i = 0; i < 4; i++) {
for (var i = 0; i <= captionLineCount; i++) {
document.getElementById("t" + i).remove();
}
elem.remove()
@@ -169,8 +206,26 @@ function remove_element() {
chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
const { type, data } = request;
if (type === "STOP") {
const saveCaptions = data.saveCaptions;
const captionLines = data.captionLines || captionLineCount;
if (type === "STOP") {
if (saveCaptions === true) {
// If there is a last incomplete segment, push it to allSegments
if (lastIncompleteSegment && lastIncompleteSegment.text && lastIncompleteSegment.text.trim() !== "") {
// Apply same Python logic: check if transcript is empty OR start >= last end
if (allSegments.length === 0 || parseFloat(lastIncompleteSegment.start) >= parseFloat(allSegments[allSegments.length - 1].end)) {
allSegments.push({
start: lastIncompleteSegment.start,
end: lastIncompleteSegment.end,
text: lastIncompleteSegment.text
});
console.log("Added final incomplete segment");
}
}
downloadSRT();
}
remove_element();
sendResponse({data: "STOPPED"});
return true;
@@ -182,55 +237,79 @@ chrome.runtime.onMessage.addListener((request, sender, sendResponse) => {
return true;
}
init_element();
init_element(captionLines);
message = JSON.parse(data);
message = message["segments"];
var text = '';
for (var i = 0; i < message.length; i++) {
text += message[i].text + ' ';
}
text = text.replace(/(\r\n|\n|\r)/gm, "");
var elem = document.getElementById('t3');
elem.innerHTML = text;
var line_height_style = getStyle('t3', 'line-height');
var line_height = parseInt(line_height_style.substring(0, line_height_style.length - 2));
var divHeight = elem.offsetHeight;
var lines = divHeight / line_height;
text_segments = [];
text_segments = get_lines(elem, line_height);
elem.innerHTML = '';
if (text_segments.length > 2) {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = text_segments[text_segments.length - 3 + i];
try {
const message = JSON.parse(data.data);
const segments = message["segments"];
if (saveCaptions === true) {
segments.forEach(seg => {
if (seg.completed === true &&
(allSegments.length === 0 || parseFloat(seg.start) >= parseFloat(allSegments[allSegments.length - 1].end))) {
allSegments.push({
start: seg.start,
end: seg.end,
text: seg.text
});
lastIncompleteSegment = null;
} else if (seg.completed !== true) {
lastIncompleteSegment = seg;
}
});
}
} else {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = '';
var text = '';
for (var i = 0; i < segments.length; i++) {
text += segments[i].text + ' ';
}
}
text = text.replace(/(\r\n|\n|\r)/gm, "");
var elem = document.getElementById('t' + captionLineCount);
if (elem) {
elem.textContent = text;
if (text_segments.length <= 2) {
for (var i = 0; i < text_segments.length; i++) {
document.getElementById('t' + i).innerHTML = text_segments[i];
}
} else {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = text_segments[text_segments.length - 3 + i];
}
}
var line_height_style = getStyle('t' + captionLineCount, 'line-height');
var line_height = parseInt(line_height_style.substring(0, line_height_style.length - 2));
var divHeight = elem.offsetHeight;
var lines = divHeight / line_height;
for (var i = 1; i < 3; i++)
{
var parent_elem = document.getElementById('t' + (i - 1));
var elem = document.getElementById('t' + i);
elem.style.top = parent_elem.offsetHeight + parent_elem.offsetTop + 'px';
text_segments = [];
text_segments = get_lines(elem, line_height);
elem.textContent = '';
if (text_segments.length > captionLineCount - 1) {
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = text_segments[text_segments.length - captionLineCount + i];
}
} else {
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = '';
}
}
if (text_segments.length <= captionLineCount - 1) {
for (var i = 0; i < text_segments.length; i++) {
document.getElementById('t' + i).textContent = text_segments[i];
}
} else {
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = text_segments[text_segments.length - captionLineCount + i];
}
}
for (var i = 1; i < captionLineCount; i++)
{
var parent_elem = document.getElementById('t' + (i - 1));
var elem = document.getElementById('t' + i);
if (parent_elem && elem) {
elem.style.top = parent_elem.offsetHeight + parent_elem.offsetTop + 'px';
}
}
}
} catch (error) {
console.error("Error processing message:", error);
}
sendResponse({});
+8 -2
View File
@@ -1,14 +1,20 @@
{
{
"manifest_version": 3,
"name": "Audio Transcription",
"version": "1.0.0",
"description": "This extension captures the audio on the current tab, sends it to a server for transcription and shows the transcription in Real-time.",
"options_page": "options.html",
"background": {
"service_worker": "background.js"
},
"web_accessible_resources": [
{
"resources": ["audiopreprocessor.js"],
"matches": ["<all_urls>"]
}
],
"permissions": [
"storage",
"activeTab",
+98 -68
View File
@@ -31,41 +31,6 @@ function sendMessageToTab(tabId, data) {
});
}
/**
* Resamples the audio data to a target sample rate of 16kHz.
* @param {Array|ArrayBuffer|TypedArray} audioData - The input audio data.
* @param {number} [origSampleRate=44100] - The original sample rate of the audio data.
* @returns {Float32Array} The resampled audio data at 16kHz.
*/
function resampleTo16kHZ(audioData, origSampleRate = 44100) {
// Convert the audio data to a Float32Array
const data = new Float32Array(audioData);
// Calculate the desired length of the resampled data
const targetLength = Math.round(data.length * (16000 / origSampleRate));
// Create a new Float32Array for the resampled data
const resampledData = new Float32Array(targetLength);
// Calculate the spring factor and initialize the first and last values
const springFactor = (data.length - 1) / (targetLength - 1);
resampledData[0] = data[0];
resampledData[targetLength - 1] = data[data.length - 1];
// Resample the audio data
for (let i = 1; i < targetLength - 1; i++) {
const index = i * springFactor;
const leftIndex = Math.floor(index).toFixed();
const rightIndex = Math.ceil(index).toFixed();
const fraction = index - leftIndex;
resampledData[i] = data[leftIndex] + (data[rightIndex] - data[leftIndex]) * fraction;
}
// Return the resampled data
return resampledData;
}
function generateUUID() {
let dt = new Date().getTime();
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, function(c) {
@@ -76,33 +41,109 @@ function generateUUID() {
return uuid;
}
// Global variables for audio processing
let audioContext = null;
let preNode = null;
let socket = null;
let isServerReady = false;
let currentStream = null;
let currentOptions = null;
// AudioWorklet URL - make sure this path matches your manifest.json
const WORKLET_URL = chrome.runtime.getURL('audiopreprocessor.js');
async function initAudioWorklet(stream) {
audioContext = new AudioContext();
if (audioContext.state === 'suspended') {
await audioContext.resume();
}
try {
await audioContext.audioWorklet.addModule(WORKLET_URL);
preNode = new AudioWorkletNode(audioContext, 'audiopreprocessor');
const mediaStream = audioContext.createMediaStreamSource(stream);
mediaStream.connect(preNode);
preNode.connect(audioContext.destination);
preNode.port.onmessage = (event) => {
const data = event.data;
const audio16k = data; // Float32Array @ 16 kHz
if (socket && socket.readyState === WebSocket.OPEN && isServerReady) {
socket.send(audio16k);
}
};
// Test if we can hear audio (this will help verify the audio path)
} catch (error) {
console.error("Error initializing AudioWorklet:", error);
throw error;
}
}
function cleanupAudio() {
if (preNode) {
preNode.port.onmessage = null;
preNode.disconnect();
preNode = null;
}
if (audioContext) {
audioContext.close();
audioContext = null;
}
if (currentStream) {
currentStream.getTracks().forEach(track => {
track.stop();
console.log("Stopped track:", track.kind);
});
currentStream = null;
}
}
/**
* Starts recording audio from the captured tab.
* @param {Object} option - The options object containing the currentTabId.
*/
async function startRecord(option) {
currentOptions = option;
const stream = await captureTabAudio();
const uuid = generateUUID();
if (stream) {
// call when the stream inactive
currentStream = stream;
stream.oninactive = () => {
cleanupAudio();
window.close();
};
const socket = new WebSocket(`ws://${option.host}:${option.port}/`);
let isServerReady = false;
let language = option.language;
if (language === null && !option.multilingual) {
language = 'en';
try {
await initAudioWorklet(stream);
} catch (error) {
console.error("Failed to initialize AudioWorklet:", error);
return;
}
socket.onopen = function(e) {
const wsUrl = option.port
? `ws://${option.host}:${option.port}/`
: `wss://${option.host}/ws`;
socket = new WebSocket(wsUrl);
isServerReady = false;
let language = option.language;
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
})
);
};
@@ -131,7 +172,6 @@ async function startRecord(option) {
language = data["language"];
// send message to popup.js to update dropdown
// console.log(language);
chrome.runtime.sendMessage({
action: "updateSelectedLanguage",
detectedLanguage: language,
@@ -141,43 +181,33 @@ async function startRecord(option) {
}
if (data["message"] === "DISCONNECT"){
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false })
chrome.runtime.sendMessage({ action: "toggleCaptureButtons", data: false, saveCaptions: option.saveCaptions });
return;
}
res = await sendMessageToTab(option.currentTabId, {
const res = await sendMessageToTab(option.currentTabId, {
type: "transcript",
data: event.data,
data: {
data: event.data,
saveCaptions: option.saveCaptions,
},
});
};
const audioDataCache = [];
const context = new AudioContext();
const mediaStream = context.createMediaStreamSource(stream);
const recorder = context.createScriptProcessor(4096, 1, 1);
recorder.onaudioprocess = async (event) => {
if (!context || !isServerReady) return;
const inputData = event.inputBuffer.getChannelData(0);
const audioData16kHz = resampleTo16kHZ(inputData, context.sampleRate);
audioDataCache.push(inputData);
socket.send(audioData16kHz);
socket.onclose = () => {
cleanupAudio();
};
socket.onerror = (error) => {
cleanupAudio();
};
// Prevent page mute
mediaStream.connect(recorder);
recorder.connect(context.destination);
mediaStream.connect(context.destination);
// }
} else {
window.close();
}
}
/**
* Listener for incoming messages from the extension's background script.
* @param {Object} request - The message request object.
+19
View File
@@ -0,0 +1,19 @@
{
"name": "audio-transcription-chrome",
"version": "1.0.0",
"description": "Audio Transcription is a Chrome extension that allows users to capture any audio playing on the current tab and transcribe it using OpenAI-whisper in real time. Users will have the option to do voice activity detection as well to not send audio to server when there is no speech.",
"main": "audiopreprocessor.js",
"scripts": {
"test": "jest"
},
"jest": {
"testEnvironment": "jsdom"
},
"keywords": [],
"author": "",
"license": "ISC",
"devDependencies": {
"jest": "^30.4.2",
"jest-environment-jsdom": "^30.4.1"
}
}
+126 -97
View File
@@ -16,120 +16,149 @@
<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="checkbox-container">
<input type="checkbox" id="saveCaptionsCheckbox">
<label for="saveCaptions">Download SRT file at Stop Capture</label>
</div>
<div class="dropdown-container">
<label for="captionLinesDropdown">Caption Lines:</label>
<select id="captionLinesDropdown">
<option value="3" selected>3 lines</option>
<option value="5">5 lines</option>
<option value="8">8 lines</option>
</select>
</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>
+66 -24
View File
@@ -4,11 +4,16 @@ 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 saveCaptionsCheckbox = document.getElementById("saveCaptionsCheckbox");
const languageDropdown = document.getElementById('languageDropdown');
const taskDropdown = document.getElementById('taskDropdown');
const modelSizeDropdown = document.getElementById('modelSizeDropdown');
const captionLinesDropdown = document.getElementById('captionLinesDropdown');
let selectedLanguage = null;
let selectedTask = taskDropdown.value;
let selectedModelSize = modelSizeDropdown.value;
let selectedCaptionLines = captionLinesDropdown.value;
// Add click event listeners to the buttons
startButton.addEventListener("click", startCapture);
@@ -30,11 +35,15 @@ 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;
}
});
chrome.storage.local.get("saveCaptionsState", ({ saveCaptionsState }) => {
if (saveCaptionsState !== undefined) {
saveCaptionsCheckbox.checked = saveCaptionsState;
}
});
@@ -52,6 +61,20 @@ document.addEventListener("DOMContentLoaded", function () {
}
});
chrome.storage.local.get("selectedModelSize", ({ selectedModelSize: storedModelSize }) => {
if (storedModelSize !== undefined) {
modelSizeDropdown.value = storedModelSize;
selectedModelSize = storedModelSize;
}
});
chrome.storage.local.get("selectedCaptionLines", ({ selectedCaptionLines: storedCaptionLines }) => {
if (storedCaptionLines !== undefined) {
captionLinesDropdown.value = storedCaptionLines;
selectedCaptionLines = storedCaptionLines;
}
});
// Function to handle the start capture button click event
async function startCapture() {
// Ignore click if the button is disabled
@@ -67,8 +90,8 @@ document.addEventListener("DOMContentLoaded", function () {
let port = "9090";
const useCollaboraServer = useServerCheckbox.checked;
if (useCollaboraServer){
host = "transcription.kurg.org"
port = "7090"
host = "boxerab--aavaaz-live-livetranscriber-web.modal.run"
port = ""
}
chrome.runtime.sendMessage(
@@ -77,9 +100,12 @@ 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,
saveCaptions: saveCaptionsCheckbox.checked,
captionLines: Number(selectedCaptionLines),
}, () => {
// Update capturing state in storage and toggle the buttons
chrome.storage.local.set({ capturingState: { isCapturing: true } }, () => {
@@ -97,7 +123,11 @@ document.addEventListener("DOMContentLoaded", function () {
}
// Send a message to the background script to stop capturing
chrome.runtime.sendMessage({ action: "stopCapture" }, () => {
chrome.runtime.sendMessage(
{
action: "stopCapture",
saveCaptions: saveCaptionsCheckbox.checked,
}, () => {
// Update capturing state in storage and toggle the buttons
chrome.storage.local.set({ capturingState: { isCapturing: false } }, () => {
toggleCaptureButtons(false);
@@ -118,9 +148,13 @@ 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;
saveCaptionsCheckbox.disabled = isCapturing;
modelSizeDropdown.disabled = isCapturing;
languageDropdown.disabled = isCapturing;
taskDropdown.disabled = isCapturing;
captionLinesDropdown.disabled = isCapturing;
startButton.classList.toggle("disabled", isCapturing);
stopButton.classList.toggle("disabled", !isCapturing);
}
@@ -131,16 +165,14 @@ 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 });
});
saveCaptionsCheckbox.addEventListener("change", () => {
const saveCaptionsState = saveCaptionsCheckbox.checked;
chrome.storage.local.set({ saveCaptionsState });
});
languageDropdown.addEventListener('change', function() {
@@ -157,6 +189,16 @@ document.addEventListener("DOMContentLoaded", function () {
chrome.storage.local.set({ selectedTask });
});
modelSizeDropdown.addEventListener('change', function() {
selectedModelSize = modelSizeDropdown.value;
chrome.storage.local.set({ selectedModelSize });
});
captionLinesDropdown.addEventListener('change', function() {
selectedCaptionLines = captionLinesDropdown.value;
chrome.storage.local.set({ selectedCaptionLines });
});
chrome.runtime.onMessage.addListener(async (request, sender, sendResponse) => {
if (request.action === "updateSelectedLanguage") {
const detectedLanguage = request.detectedLanguage;
+2 -1
View File
@@ -24,9 +24,10 @@ 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.
- **Download SRT file at Stop Capture**: Select if you want to download the srt file for the session at stop capture.
- **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).
@@ -0,0 +1,70 @@
class AudioPreProcessor extends AudioWorkletProcessor {
constructor() {
super();
this.sampleRate = sampleRate || 48000;
this.targetSampleRate = 16000;
this.inputSamplesNeeded = this.sampleRate * 0.5;
this.inputBuffer = new Float32Array(this.inputSamplesNeeded);
this.inputWriteOffset = 0;
}
process(inputs, outputs) {
const input = inputs[0];
const output = outputs[0];
if (!input || input.length === 0) {
return true;
}
for (let channel = 0; channel < Math.min(input.length, output.length); channel++) {
if (input[channel] && output[channel]) {
output[channel].set(input[channel]);
}
}
let monoInput;
if (input.length === 1) {
monoInput = input[0];
} else if (input.length > 1) {
monoInput = new Float32Array(input[0].length);
for (let channel = 0; channel < input.length; channel++) {
monoInput.set(input[channel], 0);
}
}
if (!monoInput) {
return true;
}
let inputOffset = 0;
while (inputOffset < monoInput.length) {
const remainingBuffer = this.inputSamplesNeeded - this.inputWriteOffset;
const toCopy = Math.min(remainingBuffer, monoInput.length - inputOffset);
this.inputBuffer.set(monoInput.subarray(inputOffset, inputOffset + toCopy), this.inputWriteOffset);
this.inputWriteOffset += toCopy;
inputOffset += toCopy;
if (this.inputWriteOffset === this.inputSamplesNeeded) {
const downsampled = this.downsampleTo16kHz(this.inputBuffer);
this.port.postMessage(downsampled);
this.inputWriteOffset = 0;
}
}
return true;
}
downsampleTo16kHz(inputBuffer) {
const ratio = this.sampleRate / this.targetSampleRate;
const length = Math.floor(inputBuffer.length / ratio);
const result = new Float32Array(length);
for (let i = 0; i < length; i++) {
const idx = Math.floor(i * ratio);
result[i] = inputBuffer[idx];
}
return result;
}
}
registerProcessor('audiopreprocessor', AudioPreProcessor);
+192 -146
View File
@@ -1,151 +1,168 @@
let socket = null;
let isCapturing = false;
let mediaStream = null;
let audioContext = null;
let scriptProcessor = null;
let language = null;
let isPaused = false;
let preNode = null;
let allSegments = [];
let lastIncompleteSegment = null;
const mediaElements = document.querySelectorAll('video, audio');
mediaElements.forEach((mediaElement) => {
mediaElement.addEventListener('play', handlePlaybackStateChange);
mediaElement.addEventListener('pause', handlePlaybackStateChange);
});
function handlePlaybackStateChange(event) {
isPaused = event.target.paused;
function formatTime(seconds) {
const date = new Date(seconds * 1000);
const hh = String(date.getUTCHours()).padStart(2, '0');
const mm = String(date.getUTCMinutes()).padStart(2, '0');
const ss = String(date.getUTCSeconds()).padStart(2, '0');
const mmm = String(date.getUTCMilliseconds()).padStart(3, '0');
return `${hh}:${mm}:${ss},${mmm}`;
}
function generateSRT() {
return allSegments
.map((seg, i) => {
const start = formatTime(seg.start);
const end = formatTime(seg.end);
const text = seg.text.trim().replace(/[\r\n]+/g, ' ');
return `${i + 1}\n${start} --> ${end}\n${text}`;
})
.join('\n\n');
}
function downloadSRT() {
const srtBlob = new Blob([generateSRT()], { type: 'text/srt;charset=utf-8' });
const url = URL.createObjectURL(srtBlob);
const a = document.createElement('a');
a.href = url;
a.download = 'captions.srt';
a.style.display = 'none';
document.body.appendChild(a);
a.click();
URL.revokeObjectURL(url);
document.body.removeChild(a);
}
function generateUUID() {
let dt = new Date().getTime();
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, function(c) {
return 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, c => {
const r = (dt + Math.random() * 16) % 16 | 0;
dt = Math.floor(dt / 16);
return (c === 'x' ? r : (r & 0x3 | 0x8)).toString(16);
});
return uuid;
}
/**
* Resamples the audio data to a target sample rate of 16kHz.
* @param {Array|ArrayBuffer|TypedArray} audioData - The input audio data.
* @param {number} [origSampleRate=44100] - The original sample rate of the audio data.
* @returns {Float32Array} The resampled audio data at 16kHz.
*/
function resampleTo16kHZ(audioData, origSampleRate = 44100) {
// Convert the audio data to a Float32Array
const data = new Float32Array(audioData);
document.querySelectorAll('video, audio').forEach(el => {
el.addEventListener('play', () => { isPaused = false; });
el.addEventListener('pause', () => { isPaused = true; });
});
// Calculate the desired length of the resampled data
const targetLength = Math.round(data.length * (16000 / origSampleRate));
// Create a new Float32Array for the resampled data
const resampledData = new Float32Array(targetLength);
function setupMessageHandler() {
if (preNode) {
preNode.port.onmessage = e => {
const audio16k = e.data;
if (isCapturing && socket && socket.readyState === WebSocket.OPEN && !isPaused) {
socket.send(audio16k);
}
};
}
}
// Calculate the spring factor and initialize the first and last values
const springFactor = (data.length - 1) / (targetLength - 1);
resampledData[0] = data[0];
resampledData[targetLength - 1] = data[data.length - 1];
// Resample the audio data
for (let i = 1; i < targetLength - 1; i++) {
const index = i * springFactor;
const leftIndex = Math.floor(index).toFixed();
const rightIndex = Math.ceil(index).toFixed();
const fraction = index - leftIndex;
resampledData[i] = data[leftIndex] + (data[rightIndex] - data[leftIndex]) * fraction;
const WORKLET_URL = browser.runtime.getURL('audiopreprocessor.js');
async function initAudioWorklet() {
if (audioContext && preNode) {
setupMessageHandler();
return;
}
audioContext = new AudioContext();
await audioContext.audioWorklet.addModule(WORKLET_URL);
preNode = new AudioWorkletNode(audioContext, 'audiopreprocessor');
document.querySelectorAll('audio, video').forEach(el => {
let src;
try {
src = audioContext.createMediaElementSource(el);
} catch(e) {
console.warn('Could not create MediaElementSource for', el, e);
return;
}
src.connect(preNode);
src.connect(audioContext.destination);
});
preNode.connect(audioContext.destination);
setupMessageHandler();
}
async function startRecording(data) {
if (!audioContext) {
await initAudioWorklet();
}
// Return the resampled data
return resampledData;
}
const uid = generateUUID();
socket = new WebSocket(`ws://${data.host}:${data.port}/`);
language = data.language;
function startRecording(data) {
socket = new WebSocket(`ws://${data.host}:${data.port}/`);
language = data.language;
if (language === null && !data.useMultilingual) {
language = 'en';
socket.onopen = () => {
socket.send(JSON.stringify({
uid,
language: data.language,
task: data.task,
model: data.modelSize,
use_vad: data.useVad
}));
};
let serverReady = false;
socket.onmessage = async event => {
const msg = JSON.parse(event.data);
if (msg.uid !== uid) return;
if (msg.status === 'WAIT') {
await browser.runtime.sendMessage({ action: 'showPopup', data: msg.message });
return;
}
if (!serverReady && msg.message === 'SERVER_READY') {
serverReady = true;
return;
}
if (!language && msg.language) {
language = msg.language;
await browser.runtime.sendMessage({ action: 'updateSelectedLanguage', data: language });
return;
}
if (msg.message === 'DISCONNECT') {
await browser.runtime.sendMessage({ action: 'toggleCaptureButtons' });
return;
}
if (msg.segments) {
await browser.runtime.sendMessage({ action: 'transcript', data: {data: event.data, saveCaption: data.saveCaption} });
}
};
const uuid = generateUUID();
socket.onopen = function(e) {
socket.send(
JSON.stringify({
uid: uuid,
multilingual: data.useMultilingual,
language: data.language,
task: data.task
})
);
};
let isServerReady = false;
socket.onmessage = async (event) => {
const data = JSON.parse(event.data);
if (data["uid"] !== uuid)
return;
if (data["status"] === "WAIT"){
await browser.runtime.sendMessage({ action: "showPopup", data: data["message"] })
return;
}
if (!isServerReady && data["message"] === "SERVER_READY"){
isServerReady = true;
return;
}
if (language === null ){
language = data["language"];
await browser.runtime.sendMessage({ action: "updateSelectedLanguage", data: language })
return
}
if (data["message"] === "DISCONNECT"){
await browser.runtime.sendMessage({ action: "toggleCaptureButtons", data: false })
return
}
await browser.runtime.sendMessage({ action: "transcript", data: event.data })
.catch(function(error) {
console.error("Error sending message:", error);
});
};
// Access the audio stream from the current tab
navigator.mediaDevices.getUserMedia({ audio: true })
.then(function(stream) {
// Create a new MediaRecorder instance
const audioDataCache = [];
audioContext = new AudioContext();
mediaStream = audioContext.createMediaStreamSource(stream);
recorder = audioContext.createScriptProcessor(4096, 1, 1);
recorder.onaudioprocess = async (event) => {
if (!audioContext || !isCapturing || !isServerReady || isPaused) return;
const inputData = event.inputBuffer.getChannelData(0);
const audioData16kHz = resampleTo16kHZ(inputData, audioContext.sampleRate);
audioDataCache.push(inputData);
socket.send(audioData16kHz);
};
// Prevent page mute
mediaStream.connect(recorder);
recorder.connect(audioContext.destination);
})
isCapturing = true;
}
function stopRecording() {
isCapturing = false;
if (socket) {
socket.close();
socket = null;
}
remove_element();
}
var elem_container = null;
var elem_text = null;
var segments = [];
var text_segments = [];
var captionLineCount = 3;
function initPopupElement() {
if (document.getElementById('popupElement')) {
@@ -193,22 +210,23 @@ function showPopup(customText) {
}
function init_element() {
function init_element(lines = 3) {
captionLineCount = Math.min(Math.max(parseInt(lines, 10) || 3, 1), 8);
if (document.getElementById('transcription')) {
return;
}
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:0 24px;font-family:Arial,Helvetica,sans-serif;font-size:22px;font-weight:600;line-height:30px;position:fixed;top:85%;left:50%;transform:translate(-50%,-50%);width:min(80vw,900px);min-height:' + (captionLineCount * 30) + 'px;z-index:2147483647;color:white;text-align:center;letter-spacing:0.01em;text-shadow:0 0 2px #000,0 2px 4px rgba(0,0,0,0.95);cursor:move;';
for (var i = 0; i < 4; i++) {
for (var i = 0; i <= captionLineCount; i++) {
elem_text = document.createElement('span');
elem_text.style.cssText = 'position: absolute;padding-left:16px;padding-right:16px;';
elem_text.style.cssText = 'position:absolute;left:50%;transform:translateX(-50%);max-width:100%;padding:2px 14px;background:rgba(0,0,0,0.72);border-radius:6px;box-decoration-break:clone;-webkit-box-decoration-break:clone;';
elem_text.id = "t" + i;
elem_container.appendChild(elem_text);
if (i == 3) {
if (i == captionLineCount) {
elem_text.style.top = "-1000px"
}
}
@@ -270,7 +288,7 @@ function get_lines(elem, line_height) {
var divHeight = elem.offsetHeight;
var lines = divHeight / line_height;
var original_text = elem.innerHTML;
var original_text = elem.textContent;
var words = original_text.split(' ');
var segments = [];
@@ -280,7 +298,7 @@ function get_lines(elem, line_height) {
for (var i = 0; i < words.length; i++)
{
segment += words[i] + ' ';
elem.innerHTML = segment;
elem.textContent = segment;
divHeight = elem.offsetHeight;
if ((divHeight / line_height) > current_lines) {
@@ -294,7 +312,7 @@ function get_lines(elem, line_height) {
var line_segment = segment.substring(segment_len, segment.length - 1)
segments.push(line_segment);
elem.innerHTML = original_text;
elem.textContent = original_text;
return segments;
@@ -302,7 +320,7 @@ function get_lines(elem, line_height) {
function remove_element() {
var elem = document.getElementById('transcription')
for (var i = 0; i < 4; i++) {
for (var i = 0; i <= captionLineCount; i++) {
document.getElementById("t" + i).remove();
}
elem.remove()
@@ -310,6 +328,9 @@ function remove_element() {
browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
const { action, data } = request;
const saveCaption = data.saveCaption || false;
const captionLines = data.captionLines || captionLineCount;
if (action === "startCapture") {
isCapturing = true;
startRecording(data);
@@ -320,12 +341,20 @@ browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
socket.close();
socket = null;
}
if (audioContext) {
audioContext.close();
audioContext = null;
mediaStream = null;
recorder = null;
if (saveCaption === true) {
if (lastIncompleteSegment && lastIncompleteSegment.text && lastIncompleteSegment.text.trim() !== "") {
if (allSegments.length === 0 || parseFloat(lastIncompleteSegment.start) >= parseFloat(allSegments[allSegments.length - 1].end)) {
allSegments.push({
start: lastIncompleteSegment.start,
end: lastIncompleteSegment.end,
text: lastIncompleteSegment.text
});
}
}
downloadSRT();
}
remove_element();
@@ -338,9 +367,26 @@ browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
} else if (action === "show_transcript"){
if (!isCapturing) return;
init_element();
message = JSON.parse(data);
init_element(captionLines);
message = JSON.parse(data.data);
message = message["segments"];
if (saveCaption === true) {
message.forEach(seg => {
if (seg.completed === true &&
(allSegments.length === 0 || parseFloat(seg.start) >= parseFloat(allSegments[allSegments.length - 1].end))) {
allSegments.push({
start: seg.start,
end: seg.end,
text: seg.text
});
lastIncompleteSegment = null;
} else if (seg.completed !== true) {
lastIncompleteSegment = seg;
}
});
}
var text = '';
for (var i = 0; i < message.length; i++) {
@@ -348,10 +394,10 @@ browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
}
text = text.replace(/(\r\n|\n|\r)/gm, "");
var elem = document.getElementById('t3');
elem.innerHTML = text;
var elem = document.getElementById('t' + captionLineCount);
elem.textContent = text;
var line_height_style = getStyle('t3', 'line-height');
var line_height_style = getStyle('t' + captionLineCount, 'line-height');
var line_height = parseInt(line_height_style.substring(0, line_height_style.length - 2));
var divHeight = elem.offsetHeight;
var lines = divHeight / line_height;
@@ -359,29 +405,29 @@ browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
text_segments = [];
text_segments = get_lines(elem, line_height);
elem.innerHTML = '';
elem.textContent = '';
if (text_segments.length > 2) {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = text_segments[text_segments.length - 3 + i];
if (text_segments.length > captionLineCount - 1) {
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = text_segments[text_segments.length - captionLineCount + i];
}
} else {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = '';
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = '';
}
}
if (text_segments.length <= 2) {
if (text_segments.length <= captionLineCount - 1) {
for (var i = 0; i < text_segments.length; i++) {
document.getElementById('t' + i).innerHTML = text_segments[i];
document.getElementById('t' + i).textContent = text_segments[i];
}
} else {
for (var i = 0; i < 3; i++) {
document.getElementById('t' + i).innerHTML = text_segments[text_segments.length - 3 + i];
for (var i = 0; i < captionLineCount; i++) {
document.getElementById('t' + i).textContent = text_segments[text_segments.length - captionLineCount + i];
}
}
for (var i = 1; i < 3; i++)
for (var i = 1; i < captionLineCount; i++)
{
var parent_elem = document.getElementById('t' + (i - 1));
var elem = document.getElementById('t' + i);
@@ -8,6 +8,9 @@
"activeTab",
"<all_urls>"
],
"web_accessible_resources": [
"audiopreprocessor.js"
],
"background": {
"scripts": ["background.js"],
"persistent": false
+127 -99
View File
@@ -15,114 +15,126 @@
<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>
<div class="checkbox-container">
<input type="checkbox" id="saveCaptionCheckbox">
<label for="saveCaption">Download SRT file at Stop Capture</label>
</div>
<textarea id="waitTextBox" style="display: none;"></textarea>
<div class="dropdown-container">
<label for="captionLinesDropdown">Caption Lines:</label>
<select id="captionLinesDropdown">
<option value="3" selected>3 lines</option>
<option value="5">5 lines</option>
<option value="8">8 lines</option>
</select>
</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">
@@ -133,5 +145,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>
+60 -21
View File
@@ -3,11 +3,17 @@ 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 saveCaptionCheckbox = document.getElementById("saveCaptionCheckbox");
const languageDropdown = document.getElementById('languageDropdown');
const taskDropdown = document.getElementById('taskDropdown');
const modelSizeDropdown = document.getElementById('modelSizeDropdown');
const captionLinesDropdown = document.getElementById('captionLinesDropdown');
let selectedLanguage = null;
let selectedTask = taskDropdown.value;
let selectedModelSize = modelSizeDropdown.value;
let selectedCaptionLines = captionLinesDropdown.value;
browser.storage.local.get("capturingState")
.then(function(result) {
@@ -32,11 +38,15 @@ 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;
}
});
browser.storage.local.get("saveCaptionState", ({ saveCaptionState }) => {
if (saveCaptionState !== undefined) {
saveCaptionCheckbox.checked = saveCaptionState;
}
});
@@ -54,6 +64,20 @@ document.addEventListener("DOMContentLoaded", function() {
}
});
browser.storage.local.get("selectedModelSize", ({ selectedModelSize: storedModelSize }) => {
if (storedModelSize !== undefined) {
modelSizeDropdown.value = storedModelSize;
selectedModelSize = storedModelSize;
}
});
browser.storage.local.get("selectedCaptionLines", ({ selectedCaptionLines: storedCaptionLines }) => {
if (storedCaptionLines !== undefined) {
captionLinesDropdown.value = storedCaptionLines;
selectedCaptionLines = storedCaptionLines;
}
});
startButton.addEventListener("click", function() {
let host = "localhost";
let port = "9090";
@@ -73,9 +97,12 @@ document.addEventListener("DOMContentLoaded", function() {
data: {
host: host,
port: port,
useMultilingual: useMultilingualCheckbox.checked,
language: selectedLanguage,
task: selectedTask
task: selectedTask,
modelSize: selectedModelSize,
useVad: useVadCheckbox.checked,
saveCaption: saveCaptionCheckbox.checked,
captionLines: Number(selectedCaptionLines),
}
});
toggleCaptureButtons(true);
@@ -92,7 +119,7 @@ document.addEventListener("DOMContentLoaded", function() {
stopButton.addEventListener("click", function() {
browser.tabs.query({ active: true, currentWindow: true })
.then(function(tabs) {
browser.tabs.sendMessage(tabs[0].id, { action: "stopCapture" })
browser.tabs.sendMessage(tabs[0].id, { action: "stopCapture", data: {saveCaption: saveCaptionCheckbox.checked, } })
.then(function(response) {
toggleCaptureButtons(false);
browser.storage.local.set({ capturingState: { isCapturing: false } })
@@ -114,8 +141,12 @@ document.addEventListener("DOMContentLoaded", function() {
startButton.disabled = isCapturing;
stopButton.disabled = !isCapturing;
useServerCheckbox.disabled = isCapturing;
useMultilingualCheckbox.disabled = isCapturing;
useVadCheckbox.disabled = isCapturing;
saveCaptionCheckbox.disabled = isCapturing;
modelSizeDropdown.disabled = isCapturing;
languageDropdown.disabled = isCapturing;
taskDropdown.disabled = isCapturing;
captionLinesDropdown.disabled = isCapturing;
startButton.classList.toggle("disabled", isCapturing);
stopButton.classList.toggle("disabled", !isCapturing);
}
@@ -126,16 +157,14 @@ 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 });
});
saveCaptionCheckbox.addEventListener("change", () => {
const saveCaptionState = saveCaptionCheckbox.checked;
browser.storage.local.set({ saveCaptionState });
});
languageDropdown.addEventListener('change', function() {
@@ -152,6 +181,16 @@ document.addEventListener("DOMContentLoaded", function() {
browser.storage.local.set({ selectedTask });
});
modelSizeDropdown.addEventListener('change', function() {
selectedModelSize = modelSizeDropdown.value;
browser.storage.local.set({ selectedModelSize });
});
captionLinesDropdown.addEventListener('change', function() {
selectedCaptionLines = captionLinesDropdown.value;
browser.storage.local.set({ selectedCaptionLines });
});
browser.runtime.onMessage.addListener((request, sender, sendResponse) => {
if (request.action === "updateSelectedLanguage") {
const detectedLanguage = request.data;
+1 -1
View File
@@ -108,4 +108,4 @@ label {
.dropdown-container {
padding: 10px;
}
}
+229
View File
@@ -0,0 +1,229 @@
// AudioStream.swift
// Lecture2Quiz
//
// Created by ParkMazorika on 4/27/25.
//
import AVFoundation
/// Streams audio input to a WebSocket after converting and normalizing.
class AudioStreamer {
private let engine = AVAudioEngine()
private let inputNode: AVAudioInputNode
private var inputFormat: AVAudioFormat?
private var isPaused: Bool = false
private var audioWebSocket: AudioWebSocket?
private var partialBuffer = Data()
private var isStreaming: Bool = false
private var bufferSize: AVAudioFrameCount = 1600 // ~100ms of audio
private var sampleRate: Double = 16000
private var channels: UInt32 = 1
private var converter: AVAudioConverter?
init(webSocket: AudioWebSocket) {
self.inputNode = engine.inputNode
self.audioWebSocket = webSocket
let inputFormat = inputNode.outputFormat(forBus: 0)
print("Input format: \(inputFormat)")
let outputFormat = AVAudioFormat(
commonFormat: .pcmFormatInt16,
sampleRate: 16000,
channels: 1,
interleaved: true
)!
self.converter = AVAudioConverter(from: inputFormat, to: outputFormat)
self.inputFormat = outputFormat
}
/// Configures the audio session for recording.
func configureAudioSession() {
let session = AVAudioSession.sharedInstance()
do {
try session.setCategory(.playAndRecord, mode: .default, options: [.allowBluetooth, .defaultToSpeaker])
try session.setPreferredSampleRate(48000)
try session.setPreferredInputNumberOfChannels(1)
try session.setMode(.videoChat)
try session.setActive(true, options: .notifyOthersOnDeactivation)
sampleRate = session.sampleRate
channels = UInt32(session.inputNumberOfChannels)
print("Sample rate: \(sampleRate)")
print("Input channels: \(channels)")
} catch {
print("Failed to configure audio session: \(error.localizedDescription)")
}
}
/// Starts capturing and streaming audio data.
func startStreaming() {
guard !isStreaming else {
print("Already streaming.")
return
}
configureAudioSession()
let format = AVAudioFormat(
commonFormat: .pcmFormatFloat32,
sampleRate: 48000,
channels: channels,
interleaved: true
)
guard let hardwareFormat = format else {
print("Failed to create audio format.")
return
}
self.inputFormat = hardwareFormat
inputNode.installTap(onBus: 0, bufferSize: bufferSize, format: hardwareFormat) { [weak self] buffer, _ in
self?.processAudioBuffer(buffer)
}
do {
try engine.start()
isStreaming = true
print("AVAudioEngine started.")
} catch {
print("Failed to start AVAudioEngine: \(error.localizedDescription)")
}
}
/// Converts and sends the audio buffer to the server via WebSocket.
func processAudioBuffer(_ buffer: AVAudioPCMBuffer) {
guard let converter = self.converter else {
print("Audio converter is nil.")
return
}
if let floatChannelData = buffer.floatChannelData {
let frameLength = Int(buffer.frameLength)
let channelData = Array(UnsafeBufferPointer(start: floatChannelData.pointee, count: frameLength))
let rms = sqrt(channelData.map { $0 * $0 }.reduce(0, +) / Float(frameLength))
print("Audio RMS: \(rms)")
if rms < 0.001 {
print("Warning: Input volume is too low.")
}
}
let outputFormat = AVAudioFormat(
commonFormat: .pcmFormatInt16,
sampleRate: 16000,
channels: 1,
interleaved: true
)!
guard let newBuffer = AVAudioPCMBuffer(pcmFormat: outputFormat, frameCapacity: 1600) else {
print("Failed to allocate PCM buffer.")
return
}
let inputBlock: AVAudioConverterInputBlock = { _, outStatus in
outStatus.pointee = .haveData
return buffer
}
var error: NSError?
converter.convert(to: newBuffer, error: &error, withInputFrom: inputBlock)
if let error = error {
print("Audio conversion failed: \(error.localizedDescription)")
return
}
print("Converted buffer frameLength: \(newBuffer.frameLength), sampleRate: \(newBuffer.format.sampleRate)")
if let audioData = convertToFloat32BytesLikePython(newBuffer) {
var completeData = partialBuffer + audioData
let chunkSize = 4096
while completeData.count >= chunkSize {
let chunk = completeData.prefix(chunkSize)
audioWebSocket?.sendDataToServer(chunk)
print("Sent 4096 bytes of audio.")
completeData.removeFirst(chunkSize)
}
partialBuffer = completeData
}
}
/// Converts the audio buffer to Float32 Data with RMS normalization and soft clipping.
func convertToFloat32BytesLikePython(_ buffer: AVAudioPCMBuffer) -> Data? {
guard let int16ChannelData = buffer.int16ChannelData else {
print("int16ChannelData is nil.")
return nil
}
let frameLength = Int(buffer.frameLength)
let channelPointer = int16ChannelData.pointee
var floatArray = [Float32](repeating: 0, count: frameLength)
for i in 0..<frameLength {
let int16Value = channelPointer[i]
floatArray[i] = Float32(Int16(littleEndian: int16Value)) / 32768.0
}
let rms = sqrt(floatArray.map { $0 * $0 }.reduce(0, +) / Float(frameLength))
let targetRMS: Float32 = 0.25
let gain = targetRMS / max(rms, 0.00001)
print("Original RMS: \(rms), applied gain: \(gain)")
for i in 0..<frameLength {
let scaled = floatArray[i] * gain
let clipped = tanh(scaled * 3.0)
floatArray[i] = clipped
}
let floatData = Data(bytes: floatArray, count: frameLength * MemoryLayout<Float32>.size)
if let minVal = floatArray.min(), let maxVal = floatArray.max() {
print("Float32 value range after normalization: \(minVal)...\(maxVal)")
}
print("Converted to Float32 data: \(floatData.count) bytes")
return floatData
}
/// Pauses audio streaming by removing the input tap.
func pauseStreaming() {
guard !isPaused else { return }
inputNode.removeTap(onBus: 0)
isPaused = true
print("Audio streaming paused.")
}
/// Resumes audio streaming by reinstalling the input tap.
func resumeStreaming() {
guard isPaused else { return }
guard let inputFormat = inputFormat else {
print("inputFormat is nil.")
return
}
inputNode.installTap(onBus: 0, bufferSize: bufferSize, format: inputFormat) { [weak self] buffer, _ in
self?.processAudioBuffer(buffer)
}
isPaused = false
print("Audio streaming resumed.")
}
/// Stops the AVAudioEngine and resets streaming state.
func stopStreaming() {
guard isStreaming else {
print("Already stopped.")
return
}
inputNode.removeTap(onBus: 0)
engine.stop()
isStreaming = false
print("AVAudioEngine stopped.")
}
}
@@ -0,0 +1,256 @@
//
// RecordingViewModel.swift
// Lecture2Quiz
//
// Created by ParkMazorika on 4/27/25.
//
import Foundation
/// WebSocket client that connects to a transcription server and handles streaming, JSON messages, and retries.
class AudioWebSocket: NSObject, URLSessionWebSocketDelegate {
private var webSocketTask: URLSessionWebSocketTask?
private var urlSession: URLSession!
private let host: String
private let port: Int
private var retryCount = 0
private let maxRetries = 3
private var uid: String
private let modelSize: String
private var pingTimer: Timer?
private var processedTexts = Set<String>()
var onServerReady: (() -> Void)?
var onTranscriptionReceived: ((String) -> Void)?
init(host: String, port: Int, modelSize: String = "medium") {
self.host = host
self.port = port
self.uid = UUID().uuidString
self.modelSize = modelSize
super.init()
self.urlSession = URLSession(
configuration: .default,
delegate: self,
delegateQueue: .main
)
connect()
}
/// Establishes a WebSocket connection with the configured server.
private func connect() {
guard retryCount <= maxRetries else {
print("Maximum reconnect attempts exceeded.")
return
}
let socketURL = port == 443 || port == 80
? "wss://\(host)"
: "wss://\(host):\(port)"
guard let url = URL(string: socketURL) else {
print("Invalid URL: \(socketURL)")
return
}
webSocketTask = urlSession.webSocketTask(with: url)
webSocketTask?.resume()
print("Attempting WebSocket connection: \(socketURL)")
listen()
sendInitialJSON()
startPing()
}
/// Sends the initial JSON payload to identify and configure the session.
private func sendInitialJSON() {
let jsonPayload: [String: Any] = [
"uid": uid,
"language": "en",
"task": "transcribe",
"model": modelSize,
"use_vad": true,
"max_clients": 4,
"max_connection_time": 600
]
do {
let jsonData = try JSONSerialization.data(withJSONObject: jsonPayload, options: [])
let jsonString = String(data: jsonData, encoding: .utf8) ?? ""
print("Sending config JSON: \(jsonString)")
webSocketTask?.send(.string(jsonString)) { [weak self] error in
if let error = error {
print("Failed to send config JSON: \(error.localizedDescription)")
self?.reconnect()
} else {
print("Config JSON sent successfully.")
}
}
} catch {
print("JSON serialization error: \(error.localizedDescription)")
}
}
/// Sends audio data to the server.
func sendDataToServer(_ data: Data) {
guard isConnected else {
print("Not connected - skipping data send.")
reconnect()
return
}
webSocketTask?.send(.data(data)) { [weak self] error in
if let error = error {
print("Failed to send audio data: \(error.localizedDescription)")
self?.reconnect()
} else {
print("Sent audio data: \(data.count) bytes")
}
}
}
/// Returns true if the WebSocket is currently connected.
internal var isConnected: Bool {
webSocketTask?.state == .running
}
/// Attempts reconnection with exponential backoff.
private func reconnect() {
retryCount += 1
stopPing()
let delay = min(5.0, pow(2.0, Double(retryCount)))
DispatchQueue.global().asyncAfter(deadline: .now() + delay) { [weak self] in
print("Reconnecting... (\(self?.retryCount ?? 0)/\(self?.maxRetries ?? 0))")
self?.connect()
}
}
/// Starts listening for incoming messages from the server.
private func listen() {
webSocketTask?.receive { [weak self] result in
switch result {
case .success(let message):
self?.handleMessage(message)
self?.listen()
case .failure(let error):
print("Receive error: \(error.localizedDescription)")
self?.reconnect()
}
}
}
/// Handles incoming WebSocket messages (text or binary).
private func handleMessage(_ message: URLSessionWebSocketTask.Message) {
switch message {
case .data(let data):
print("Received binary data: \(data.count) bytes")
case .string(let text):
print("Received text message: \(text)")
guard let data = text.data(using: .utf8) else { return }
do {
if let json = try JSONSerialization.jsonObject(with: data) as? [String: Any] {
if let status = json["status"] as? String {
handleStatusMessage(status: status, message: json["message"] as? String)
return
}
if let message = json["message"] as? String, message == "SERVER_READY" {
print("Server is ready.")
onServerReady?()
return
}
if let segments = json["segments"] as? [[String: Any]] {
let wrapped = ["segments": segments]
let segmentData = try JSONSerialization.data(withJSONObject: wrapped, options: [])
let segmentString = String(data: segmentData, encoding: .utf8)!
onTranscriptionReceived?(segmentString)
print("Transcription segments forwarded.")
}
}
} catch {
print("JSON parsing error: \(error.localizedDescription)")
}
@unknown default:
print("Unknown message type received.")
}
}
/// Handles status message JSON from the server.
private func handleStatusMessage(status: String, message: String?) {
switch status {
case "WAIT":
print("Waiting: \(message ?? "")")
case "ERROR":
print("Error: \(message ?? "")")
case "WARNING":
print("Warning: \(message ?? "")")
default:
print("\(status): \(message ?? "")")
}
}
/// Sends the "END_OF_AUDIO" signal to the server.
func sendEndOfAudio() {
guard isConnected else {
print("Not connected - skipping END_OF_AUDIO.")
return
}
webSocketTask?.send(.string("END_OF_AUDIO")) { error in
if let error = error {
print("Failed to send END_OF_AUDIO: \(error.localizedDescription)")
} else {
print("END_OF_AUDIO sent.")
}
}
}
/// Gracefully closes the WebSocket connection.
func closeConnection() {
stopPing()
webSocketTask?.cancel(with: .normalClosure, reason: nil)
retryCount = maxRetries
print("WebSocket closed.")
}
/// Starts periodic ping to keep the WebSocket alive.
private func startPing() {
stopPing()
pingTimer = Timer.scheduledTimer(withTimeInterval: 15.0, repeats: true) { [weak self] _ in
self?.webSocketTask?.sendPing { error in
if let error = error {
print("Ping failed: \(error.localizedDescription)")
} else {
print("Ping sent successfully.")
}
}
}
RunLoop.main.add(pingTimer!, forMode: .common)
}
/// Stops the periodic ping timer.
private func stopPing() {
pingTimer?.invalidate()
pingTimer = nil
}
/// Called when the WebSocket is closed by the server.
func urlSession(_ session: URLSession,
webSocketTask: URLSessionWebSocketTask,
didCloseWith closeCode: URLSessionWebSocketTask.CloseCode,
reason: Data?) {
let reasonString = String(data: reason ?? Data(), encoding: .utf8) ?? "No reason"
print("WebSocket closed - code: \(closeCode.rawValue), reason: \(reasonString)")
stopPing()
reconnect()
}
}
+99
View File
@@ -0,0 +1,99 @@
//
// ContentView.swift
// WhisperLive_iOS_Client
//
// Created by ParkMazorika on 6/17/25.
//
import SwiftUI
/// A standalone view for recording and real-time transcription display.
struct RecordingView: View {
var onDismiss: () -> Void
@StateObject private var recordingViewModel = AudioViewModel()
@State private var showSubmitView = false
var body: some View {
VStack(spacing: 0) {
// Stop button (only visible when recording)
HStack {
Spacer()
if recordingViewModel.isRecording {
Button("Stop Recording") {
recordingViewModel.stopRecording()
recordingViewModel.finalizeTranscription()
showSubmitView = true
}
.font(.headline)
.padding()
.foregroundColor(.gray)
}
}
// Transcription display
ScrollView {
VStack(spacing: 8) {
ForEach(recordingViewModel.transcriptionList.indices, id: \.self) { index in
Text(recordingViewModel.transcriptionList[index])
.padding()
.frame(maxWidth: .infinity, alignment: .leading)
.background(Color.gray.opacity(0.1))
.cornerRadius(8)
.font(.system(size: 14, weight: .semibold))
}
}
.padding(.horizontal)
}
Divider().padding(.top, 8)
// Timer and Record/Pause/Resume button
VStack(spacing: 16) {
Text(recordingViewModel.timeLabel)
.font(.system(size: 40))
Button(action: {
if recordingViewModel.isRecording {
recordingViewModel.isPaused
? recordingViewModel.resumeRecording()
: recordingViewModel.pauseRecording()
} else {
recordingViewModel.startRecording()
}
}) {
Image(systemName: recordingViewModel.isRecording
? (recordingViewModel.isPaused ? "play.circle.fill" : "pause.circle.fill")
: "mic.circle.fill")
.font(.system(size: 50))
.foregroundStyle(.black)
}
}
.padding(.bottom, 40)
}
.padding(.top)
.background(Color(.systemBackground))
.overlay(
Group {
if recordingViewModel.isLoading {
ZStack {
Color.black.opacity(0.4).ignoresSafeArea()
ProgressView("Processing...")
.padding()
.background(Color.white)
.cornerRadius(10)
}
}
}
)
.sheet(isPresented: $showSubmitView) {
//anotherView
}
}
}
#Preview("Recording View") {
RecordingView {
// Dummy dismiss handler
print("RecordingView dismissed")
}
}
+119
View File
@@ -0,0 +1,119 @@
# Audio-Transcription-iOS
This is an iOS client for [WhisperLive](https://github.com/collabora/WhisperLive), a real-time speech-to-text server based on OpenAI Whisper.
The app streams microphone audio to a WhisperLive server via WebSocket and displays live transcription results in real time.
> ⚠️ This client is designed to work specifically with the [WhisperLive Python WebSocket server](https://github.com/collabora/WhisperLive?tab=readme-ov-file#running-the-server).
> Make sure the server is running and reachable from your iOS device.
## Features
- Real-time microphone capture with AVAudioEngine
- Streaming to WhisperLive backend using WebSocket
- Displays transcription as segments arrive
- Start / Pause / Resume / Stop recording with SwiftUI interface
- Final transcription view on stop
## Requirements
- iOS 15.0+
- Swift 5.8+
- AVFoundation (for microphone)
- Working WhisperLive WebSocket server
## Getting Started
This directory contains the Swift source files for the iOS client, but it does not include a generated `.xcodeproj` or `.xcodeworkspace`. Create a new Xcode project and add these files to it.
1. Clone the repository (your fork):
```bash
git clone https://github.com/yourusername/whisperlive.git
cd whisperlive/Audio-Transcription-iOS
```
2. In Xcode, choose **File ▸ New ▸ Project…**, then create an iOS **App** project with SwiftUI.
3. Add the Swift files from this directory to the new app target:
- `AudioStream.swift`
- `AudioWebSocket.swift`
- `ContentView.swift`
- `RecordingViewModel.swift`
- `WhisperLive_iOS_ClientApp.swift`
4. Use `WhisperLive-iOS-Client-Info.plist` as a reference for your app's `Info.plist`, or add the microphone usage description manually:
```xml
<key>NSMicrophoneUsageDescription</key>
<string>This app requires microphone access for transcription.</string>
```
5. Run the app on a physical device (recommended)
## Running on a Physical Device (with Free Apple ID)
You can run this app on a real iPhone without a paid Apple Developer account. Follow these steps:
### 1. Register a Free Apple ID in Xcode
1. Open Xcode ▸ Settings… (or Preferences) ▸ **Accounts**
2. Click the **+** button ▸ Select **Apple ID**
3. Sign in with your Apple ID (a free one is fine)
4. A "Personal Team" will be created automatically
> ✅ You can deploy up to 3 apps on a physical device using a free Apple ID with a 7-day provisioning profile.
---
### 2. Set Up Signing in Your Project
1. In Xcode, select your **project** in the Project Navigator
2. Go to **TARGETS ▸ YourAppName ▸ Signing & Capabilities**
3. Set **Team** to your Personal Team
4. Set a unique **Bundle Identifier** (e.g., `com.yourname.whisperlive`)
5. Make sure **Automatically manage signing** is checked
6. If a red warning appears, click **"Resolve Issues"**
---
### 3. Connect and Trust Your iPhone
1. Connect your iPhone via USB
2. When prompted, tap **“Trust This Computer”** on your iPhone
3. Make sure your iPhone appears in Xcode's device list
---
### 4. Enable Developer Mode on iPhone
1. Press the **Build (▶︎)** button in Xcode
2. Your iPhone will ask to enable **Developer Mode**
3. On iPhone, go to:
**Settings ▸ Privacy & Security ▸ Developer Mode**
4. Enable it and restart the device if required
---
Now you can run and debug the app on your real device!
## Folder Structure
```
Audio-Transcription-iOS/
├── AudioStream.swift
├── AudioWebSocket.swift
├── ContentView.swift
├── RecordingViewModel.swift
├── WhisperLive-iOS-Client-Info.plist
├── WhisperLive_iOS_ClientApp.swift
├── README.md
```
## License
MIT
This iOS client is provided as an open-source example to complement WhisperLive's real-time transcription ecosystem.
@@ -0,0 +1,174 @@
//
// RecordingViewModel.swift
// Lecture2Quiz
//
// Created by ParkMazorika on 4/27/25.
//
import AVFoundation
import Combine
/// Represents a segment of transcribed audio with start/end timestamps and completion flag.
struct TranscriptionSegment: Identifiable, Equatable {
var id = UUID()
var start: Double
var end: Double
var text: String
var completed: Bool
}
/// ViewModel responsible for managing audio recording and transcription logic.
class AudioViewModel: ObservableObject {
@Published var isRecording = false // Indicates if recording is active
@Published var isPaused = false // Indicates if recording is currently paused
@Published var timeLabel = "00:00" // Timer label formatted as mm:ss
@Published var transcriptionList: [String] = [] // Live transcription output
@Published var isLoading = false // True while waiting for server response
@Published var finalScript: String = "" // Final script from completed segments
private var timer: Timer?
private var elapsedTime: Int = 0
private var audioStreamer: AudioStreamer? // Handles audio capture and streaming
private var audioWebSocket: AudioWebSocket? // Manages WebSocket communication
private var segments: [TranscriptionSegment] = [] // Stores all transcription segments
init() {}
/// Starts audio recording and initializes WebSocket + AVAudioEngine.
func startRecording() {
let audioAPIUrl = "your server url"
audioWebSocket = AudioWebSocket(host: audioAPIUrl, port: 443)
audioStreamer = AudioStreamer(webSocket: audioWebSocket!)
isLoading = true
// Handle server transcription message
audioWebSocket?.onTranscriptionReceived = { [weak self] text in
self?.handleRawTranscriptionJSON(text)
}
// When server sends SERVER_READY
audioWebSocket?.onServerReady = { [weak self] in
guard let self = self else { return }
DispatchQueue.main.async {
self.isLoading = false
self.isRecording = true
self.isPaused = false
self.timeLabel = "00:00"
self.elapsedTime = 0
self.startTimer()
self.audioStreamer?.startStreaming()
}
}
}
/// Pauses the recording and stops the timer.
func pauseRecording() {
isPaused = true
audioStreamer?.pauseStreaming()
timer?.invalidate()
}
/// Resumes recording and restarts the timer.
func resumeRecording() {
isPaused = false
audioStreamer?.resumeStreaming()
startTimer()
}
/// Stops recording and finalizes connection to server.
func stopRecording() {
isRecording = false
isPaused = false
timer?.invalidate()
audioStreamer?.stopStreaming()
audioWebSocket?.sendEndOfAudio()
audioWebSocket?.onTranscriptionReceived = nil
audioWebSocket?.closeConnection()
}
/// Starts the recording timer (1-second interval).
private func startTimer() {
timer = Timer.scheduledTimer(withTimeInterval: 1.0, repeats: true) { _ in
self.elapsedTime += 1
let minutes = self.elapsedTime / 60
let seconds = self.elapsedTime % 60
self.timeLabel = String(format: "%02d:%02d", minutes, seconds)
}
}
/// Finalizes the transcription by joining all completed segments into one string.
func finalizeTranscription() {
isLoading = false
let completedText = segments
.filter { $0.completed }
.map { $0.text.trimmingCharacters(in: .whitespaces) }
.joined(separator: " ")
finalScript = completedText
print("Final transcript:\n\(finalScript)")
}
/// Handles incoming JSON from the server and updates UI state.
/// Supports both full JSON and raw string cases.
func handleRawTranscriptionJSON(_ jsonString: String) {
let trimmed = jsonString.trimmingCharacters(in: .whitespacesAndNewlines)
guard let data = trimmed.data(using: .utf8) else { return }
if trimmed.hasPrefix("{") {
// Parse JSON containing segment list
do {
if let dict = try JSONSerialization.jsonObject(with: data) as? [String: Any],
let segmentDicts = dict["segments"] as? [[String: Any]] {
for item in segmentDicts {
guard let startStr = item["start"] as? String,
let endStr = item["end"] as? String,
let text = item["text"] as? String,
let completed = item["completed"] as? Bool,
let start = Double(startStr),
let end = Double(endStr) else { continue }
let newSegment = TranscriptionSegment(start: start, end: end, text: text, completed: completed)
// Overwrite if already exists, else append
if let index = self.segments.firstIndex(where: { $0.start == start }) {
self.segments[index] = newSegment
} else {
self.segments.append(newSegment)
}
}
// Update the UI
DispatchQueue.main.async {
let completedTexts = self.segments
.filter { $0.completed }
.sorted(by: { $0.start < $1.start })
.map { $0.text.trimmingCharacters(in: .whitespaces) }
let pendingText = self.segments
.filter { !$0.completed }
.sorted(by: { $0.start < $1.start })
.map { $0.text.trimmingCharacters(in: .whitespaces) }
.last ?? ""
self.transcriptionList = completedTexts + (pendingText.isEmpty ? [] : [pendingText])
self.finalScript = self.transcriptionList.joined(separator: " ")
}
}
} catch {
print("JSON parsing error: \(error)")
}
} else {
// Handle raw text line
DispatchQueue.main.async {
if self.transcriptionList.last != trimmed {
self.transcriptionList.append(trimmed)
self.finalScript = self.transcriptionList.joined(separator: " ")
}
}
}
}
}
@@ -0,0 +1,8 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
<plist version="1.0">
<dict>
<key>NSMicrophoneUsageDescription</key>
<string>This app requires microphone access for voice transcription.</string>
</dict>
</plist>
@@ -0,0 +1,20 @@
//
// WhisperLive_iOS_ClientApp.swift
// WhisperLive_iOS_Client
//
// Created by on 6/17/25.
//
import SwiftUI
@main
struct WhisperLive_iOS_ClientApp: App {
var body: some Scene {
WindowGroup {
RecordingView {
// Handle dismiss action here, or leave it empty for now
print("RecordingView dismissed")
}
}
}
}
+374 -49
View File
@@ -1,14 +1,49 @@
# 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>
<a href="https://www.youtube.com/watch?v=0f5oiG4oPWQ"><img
src="https://img.youtube.com/vi/0f5oiG4oPWQ/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)
- [Advanced Features](#advanced-features)
- [Word-Level Timestamps](#word-level-timestamps)
- [Custom Vocabulary / Hotwords](#custom-vocabulary--hotwords)
- [Speaker Diarization](#speaker-diarization)
- [Batch Inference](#batch-inference)
- [Raw PCM Input](#raw-pcm-input)
- [Streaming Client (Manual Audio Chunking)](#streaming-client-manual-audio-chunking)
- [Browser Extensions](#browser-extensions)
- [Whisper Live Server in Docker](#whisper-live-server-in-docker)
- [Troubleshooting](#troubleshooting)
- [Future Work](#future-work)
- [Blog Posts](#blog-posts)
- [Contact](#contact)
- [Citations](#citations)
## Installation
- Install PyAudio and ffmpeg
- Install PortAudio (required system dependency for microphone input via PyAudio)
```bash
bash setup.sh
bash scripts/setup.sh
```
On Debian/Ubuntu this installs `portaudio19-dev`, on Fedora `portaudio-devel`, on macOS it uses Homebrew (`portaudio`).
- Install 3.12 venv (on Fedora `sudo dnf install -y python3.12 python3.12-pip`)
```bash
python3.12 -m venv whisper_env
source whisper_env/bin/activate
```
- Install whisper-live from pip
@@ -16,69 +51,360 @@ Unlike traditional speech recognition systems that rely on continuous audio stre
pip install whisper-live
```
### OpenAI REST interface
#### Server
```bash
python3 run_server.py --port 9090 --backend faster_whisper --max_clients 4 --max_connection_time 600 --enable_rest --cors-origins="http://localhost:8080,http://127.0.0.1:8080"
```
#### Client
```bash
python3 client_openai.py $AUDIO_FILE
```
### 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 \
--max_clients 4 \
--max_connection_time 600
# running with custom model and cache_dir to save auto-converted ctranslate2 models
python3 run_server.py --port 9090 \
--backend faster_whisper \
--max_clients 4 \
--max_connection_time 600 \
-fw "/path/to/custom/faster/whisper/model" \
-c ~/.cache/whisper-live/
```
- 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 \
--max_clients 4 \
--max_connection_time 600
- 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 \
--max_clients 4 \
--max_connection_time 600
```
This would start the websocket server on port ```9090```.
> **Note:** The TensorRT backend uses a C++ session by default. If you experience issues (e.g. repeated `CrossAttentionMask` warnings or crashes), add the `--trt_py_session` flag to use the Python session instead.
- Use `--max_clients` option to restrict the number of clients the server should allow. Defaults to 4.
- Use `--max_connection_time` options to limit connection time for a client in seconds. Defaults to 600.
- 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#).
### Chrome Extension
- Refer to [Audio-Transcription-Chrome](https://github.com/collabora/whisper-live/tree/main/Audio-Transcription-Chrome#readme) to use Chrome extension.
```
python3 run_server.py -p 9090 -b openvino
```
### Setting up AMD ROCm for faster_whisper backend
- Please follow [ROCm_whisper readme](https://github.com/collabora/WhisperLive/blob/main/ROCm_whisper.md) for setup of AMD ROCm GPU support with the CTranslate2 ROCm wheel.
#### 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
Use the below command to run the client:
```bash
python3 run_client.py --files <audio-file-name>
```
This will connect to the localhost server running on port 9090 by default. Use flags `--server` and `--port` to use different configurations. The above command will transcribe audio file provided with `--files` flag.
Here are the details of client instance implemented in `run_client.py` script:
- `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`.
- `mute_audio_playback`: Whether to mute audio playback when transcribing an audio file. Defaults to False.
- `enable_translation`: Start translation thread on the server (from any to any).
- `target_language`: Server translation thread's target translation language.
```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
mute_audio_playback=False, # Only used for file input, False by Default
enable_translation=True,
target_language="hi",
initial_prompt=None, # Add context for the model, e.g. 'Jane Doe context'
)
```
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")
```
## Advanced Features
#### Word-Level Timestamps
Enable per-word timing and confidence scores in transcription segments:
```python
client = TranscriptionClient(
"localhost", 9090,
word_timestamps=True,
)
```
When enabled, each segment in the WebSocket response includes a `words` array:
```json
{
"segments": [{
"start": "0.000", "end": "2.500", "text": "Hello world",
"words": [
{"word": "Hello", "start": "0.000", "end": "0.800", "probability": 0.95},
{"word": " world", "start": "0.900", "end": "2.500", "probability": 0.88}
]
}]
}
```
#### Custom Vocabulary / Hotwords
Boost recognition of specific terms (product names, acronyms, domain jargon):
```python
client = TranscriptionClient(
"localhost", 9090,
hotwords="WhisperLive,TensorRT,OpenVINO",
)
```
The `hotwords` parameter is a comma-separated string passed directly to faster-whisper's keyword boosting. Also available in the REST API via the `hotwords` form field.
#### Speaker Diarization
Real-time speaker identification using pyannote.audio embeddings (optional dependency):
```bash
pip install pyannote.audio
```
```python
client = TranscriptionClient(
"localhost", 9090,
enable_diarization=True,
max_speakers=4,
)
```
When enabled, completed segments include a `speaker` field:
```json
{"start": "0.000", "end": "2.500", "text": "Hello", "speaker": "SPEAKER_00", "completed": true}
```
Diarization uses online cosine-similarity clustering of speaker embeddings. If `pyannote.audio` is not installed, the server logs a warning and continues without diarization.
The OpenAI-compatible REST endpoint also accepts `known_speaker_names` and uploaded `known_speaker_references` multipart fields. When speaker fields are supplied with `response_format="verbose_json"`, segments include a `speaker` field.
#### Batch Inference
Batch multiple client sessions into single GPU calls for higher throughput:
```bash
python3 run_server.py --port 9090 --backend faster_whisper \
--batch_inference --batch_max_size 8 --batch_window_ms 50
```
#### Raw PCM Input
Accept raw PCM int16 audio from clients (useful for embedded devices):
```bash
python3 run_server.py --port 9090 --backend faster_whisper --raw_pcm_input
```
Audio is automatically normalized to float32 range [-1.0, 1.0]. Clients can also set `audio_format` in the initial websocket options to `float32` (default), `int16`, or `uint8`.
## Streaming Client (manual audio streaming from any source)
`StreamingTranscriptionClient` lets you push raw PCM audio bytes from any source — a live microphone capture loop, a network stream, an audio pipeline — and receive transcripts via callbacks as speech is detected. Unlike `TranscriptionClient`, it does not manage audio capture internally; you control when and how audio is fed.
A runnable example that reads from an audio file and streams the chunks is at [`examples/manual_audio_chunking.py`](examples/manual_audio_chunking.py):
```bash
python examples/manual_audio_chunking.py --file assets/jfk.flac
```
Example usage:
```python
from whisper_live.client import StreamingTranscriptionClient
client = StreamingTranscriptionClient(
"localhost", 9090,
lang="en",
model="small",
on_session_started=lambda: print("Server ready"),
on_partial_transcript=lambda text, segs: print(f"{text}", end="\r"),
on_committed_transcript=lambda text, segs: print(f"{text}"),
on_error=lambda e: print(f"Error: {e}"),
on_close=lambda: print("Closed"),
)
with client:
for chunk in my_audio_source: # any cadence, any chunk size
client.send(chunk, pcm_format="int16")
```
Audio must be **mono, 16 kHz PCM**. Two formats are accepted:
| `pcm_format` | Description |
|---|---|
| `"int16"` (default for raw microphone data) | 16-bit signed integers, normalized internally |
| `"float32"` | 32-bit floats in `[-1, 1]`, passed through directly |
NumPy arrays can be sent with `send_array()`:
```python
import numpy as np
samples = np.frombuffer(raw_bytes, dtype=np.int16)
client.send_array(samples)
```
**Callbacks**
| Callback | Signature | When fired |
|---|---|---|
| `on_session_started` | `() -> None` | Server handshake complete, ready to receive audio |
| `on_partial_transcript` | `(text, segments) -> None` | In-progress segment updated |
| `on_committed_transcript` | `(text, segments) -> None` | Segment finalized |
| `on_translation` | `(text, segments) -> None` | Translated segment ready (requires `enable_translation=True`) |
| `on_error` | `(error) -> None` | WebSocket error |
| `on_close` | `() -> None` | Connection closed |
## 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
## iOS Client
Use WhisperLive on iOS with our native iOS client.
Refer to [`ios-client`](https://github.com/collabora/WhisperLive/tree/main/Audio-Transcription-iOS) and [`ios-client/README.md`](https://github.com/collabora/WhisperLive/blob/main/Audio-Transcription-iOS/README.md) for setup and usage instructions.
### Firefox Extension
- Refer to [Audio-Transcription-Firefox](https://github.com/collabora/whisper-live/tree/main/Audio-Transcription-Firefox#readme) to use Mozilla Firefox extension.
## 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 (pick one engine)
python3 run_server.py --port 9090 \
--backend tensorrt \
--trt_model_path "/app/TensorRT-LLM-examples/whisper/whisper_small_en_float16"
# or int8 / int4:
# --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
```
- AMD ROCm (faster-whisper on AMD GPU via CTranslate2 ROCm wheel)
```bash
docker build -f docker/Dockerfile.rocm -t whisperlive-rocm .
docker run --rm -it --device=/dev/kfd --device=/dev/dri \
--group-add "$(getent group video | cut -d: -f3)" \
--group-add "$(getent group render | cut -d: -f3)" \
-p 9090:9090 whisperlive-rocm
```
- CPU
- Faster-whisper
```bash
docker run -it -p 9090:9090 ghcr.io/collabora/whisperlive-cpu:latest
```
## Troubleshooting
#### macOS OpenMP runtime conflict
On macOS, especially on Intel Macs, `faster_whisper`/`ctranslate2` can conflict with OpenMP runtimes loaded by other Python packages. If the server aborts with a duplicate OpenMP runtime error, run the server with `KMP_DUPLICATE_LIB_OK=TRUE`:
```bash
docker build . -t whisper-live -f docker/Dockerfile.cpu
docker run -it -p 9090:9090 whisper-live:latest
KMP_DUPLICATE_LIB_OK=TRUE python3 run_server.py --port 9090 \
--backend faster_whisper \
--max_clients 4 \
--max_connection_time 600 \
--no_single_model
```
**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.
This workaround is intended for local development and testing. For production deployments, prefer using a clean environment that loads only one OpenMP runtime.
## Future Work
- [ ] Add translation to other languages on top of transcription.
- [ ] TensorRT backend for Whisper.
- [x] Add translation to other languages on top of transcription.
## Blog Posts
- [Transforming speech technology with WhisperLive](https://www.collabora.com/news-and-blog/blog/2024/05/28/transforming-speech-technology-with-whisperlive/)
- [WhisperFusion: Ultra-low latency conversations with an AI chatbot](https://www.collabora.com/news-and-blog/news-and-events/whisperfusion-ultra-low-latency-conversations-with-an-ai-chatbot.html) powered by WhisperLive
- [Breaking language barriers 2.0: Moving closer towards fully reliable, production-ready Hindi ASR](https://www.collabora.com/news-and-blog/news-and-events/breaking-language-barriers-20-moving-closer-production-ready-hindi-asr.html) which is used in WhisperLive for hindi.
## Contact
We are available to help you with both Open Source and proprietary AI projects. You can reach us via the Collabora website or [vineet.suryan@collabora.com](mailto:vineet.suryan@collabora.com) and [marcus.edel@collabora.com](mailto:marcus.edel@collabora.com).
## Citations
```bibtex
@article{Whisper
@@ -98,6 +424,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}
}
+66
View File
@@ -0,0 +1,66 @@
# WhisperLive-ROCm
Run WhisperLive's `faster_whisper` backend on AMD GPUs using the official [CTranslate2 ROCm wheel](https://github.com/OpenNMT/CTranslate2/releases). Tested on Radeon AI PRO R9700 (gfx1201/RDNA4) and Ryzen AI Max+ 395 / Radeon 8060S (gfx1151/Strix Halo).
## Docker Installation (recommended)
- Install [docker](https://docs.docker.com/engine/install/)
- Build and run the WhisperLive ROCm image:
```bash
docker build -f docker/Dockerfile.rocm -t whisperlive-rocm .
docker run --rm -it \
--device=/dev/kfd --device=/dev/dri \
--group-add "$(getent group video | cut -d: -f3)" \
--group-add "$(getent group render | cut -d: -f3)" \
-p 9090:9090 whisperlive-rocm
```
## Native Installation
### Prerequisites
- AMD GPU with ROCm support (see [supported GPUs](https://rocm.docs.amd.com/en/latest/compatibility/compatibility-matrix.html))
- ROCm 7.2+ installed ([installation guide](https://rocm.docs.amd.com/en/latest/deploy/linux/quick_start.html))
- User in `video` and `render` groups (`sudo usermod -aG video,render $USER`, re-login)
- Python 3.12
### Verify ROCm is working
```bash
rocminfo | grep -E 'Name:|gfx'
# Should show your GPU, e.g. "Name: gfx1151" or "Name: gfx1201"
```
### Install CTranslate2 ROCm wheel
The default `pip install ctranslate2` installs a CUDA-only wheel. Replace it with the official ROCm wheel from the [CTranslate2 releases page](https://github.com/OpenNMT/CTranslate2/releases):
```bash
# Download the ROCm wheels archive (v4.8.0)
curl -LO https://github.com/OpenNMT/CTranslate2/releases/download/v4.8.0/rocm-python-wheels-Linux.zip
# Extract the Python 3.12 wheel
unzip -j rocm-python-wheels-Linux.zip 'temp-linux/ctranslate2-*-cp312-*manylinux*x86_64.whl'
# Install (replaces any existing ctranslate2)
pip install ctranslate2-*-cp312-*.whl
```
### Install WhisperLive server requirements
```bash
pip install -r requirements/server.txt
```
### Verify GPU is visible to CTranslate2
```bash
python -c "import ctranslate2; print('devices:', ctranslate2.get_cuda_device_count())"
```
Expected output: `devices: 1` (CTranslate2 uses the name "cuda" even on ROCm).
If you see `devices: 0`, check:
- Your user is in `video` and `render` groups (re-login after adding)
- `/dev/kfd` exists and is accessible
- The ROCm wheel was installed (not the default PyPI CUDA-only one)
## Run WhisperLive Server with ROCm
```bash
python3 run_server.py --port 9090 --backend faster_whisper
```
The server automatically uses the AMD GPU when the CTranslate2 ROCm wheel is installed. For multi-GPU systems, use `HIP_VISIBLE_DEVICES=N` to select a specific GPU.
+47
View File
@@ -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
```
BIN
View File
Binary file not shown.
+25
View File
@@ -0,0 +1,25 @@
import sys
from whisper_live.client import TranscriptionClient
if len(sys.argv) < 2:
print("Usage: python transcribe_file.py <path_to_audio_file>")
sys.exit(1)
audio_file = sys.argv[1]
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
mute_audio_playback=False, # Only used for file input, False by Default
enable_translation=True,
target_language="hi",
)
# Transcribe the offline audio file
client(audio_file)
+38
View File
@@ -0,0 +1,38 @@
import sys
import requests
if len(sys.argv) < 2:
print("Usage: python transcribe_file.py <path_to_audio_file>")
sys.exit(1)
audio_file = sys.argv[1]
# Configuration
host = "localhost"
port = 8000 # Default REST port; change if you used --rest_port
url = f"http://{host}:{port}/v1/audio/transcriptions"
model = "small" # Or "whisper-1" (mapped to small internally)
language = "en" # Or "hi" for Hindi
response_format = "json" # Options: "json", "text", "verbose_json", "srt", "vtt"
# Prepare the request
files = {"file": open(audio_file, "rb")}
data = {
"model": model,
"language": language,
"response_format": response_format,
# Optional: Add "prompt" for style guidance, "temperature" (0-1), etc.
}
# Send the request
response = requests.post(url, files=files, data=data)
if response.status_code == 200:
if response_format == "json" or response_format == "verbose_json":
result = response.json()
print("Transcript:", result.get("text", "No text found"))
# If you need translation, post-process here (e.g., using another API like Google Translate)
else:
print("Transcript:", response.text)
else:
print("Error:", response.status_code, response.json().get("error", "Unknown error"))
+12 -32
View File
@@ -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
View File
@@ -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"]
+19
View File
@@ -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"]
+44
View File
@@ -0,0 +1,44 @@
# docker/Dockerfile.rocm
#
# WhisperLive faster_whisper backend on AMD ROCm GPUs.
# Uses the official CTranslate2 ROCm wheel (ships kernels for gfx803 through
# gfx1201 including Strix Halo gfx1151 and RDNA4 gfx1200/1201).
#
# Build:
# docker build -f docker/Dockerfile.rocm -t whisperlive-rocm .
#
# Run (expose the WebSocket port; add --enable_rest --rest_port 8000 -p 8000:8000 for REST):
# docker run --rm -it \
# --device=/dev/kfd --device=/dev/dri \
# --group-add "$(getent group video | cut -d: -f3)" \
# --group-add "$(getent group render | cut -d: -f3)" \
# -p 9090:9090 whisperlive-rocm
FROM rocm/pytorch:rocm7.2.4_ubuntu24.04_py3.12_pytorch_release_2.10.0
ARG DEBIAN_FRONTEND=noninteractive
ARG CT2_WHEEL_URL=https://github.com/OpenNMT/CTranslate2/releases/download/v4.8.0/rocm-python-wheels-Linux.zip
RUN apt-get update -qq && \
apt-get install -y --no-install-recommends curl unzip portaudio19-dev && \
apt-get clean && rm -rf /var/lib/apt/lists/*
WORKDIR /app
# Install the CTranslate2 ROCm wheel (official release artifact).
# This replaces any CUDA-only ctranslate2 and enables GPU on AMD.
RUN curl -sL "${CT2_WHEEL_URL}" -o /tmp/ct2-rocm.zip && \
unzip -j /tmp/ct2-rocm.zip 'temp-linux/ctranslate2-*-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl' -d /tmp && \
pip install --no-cache-dir --force-reinstall /tmp/ctranslate2-*-cp312-*.whl && \
rm -f /tmp/ct2-rocm.zip /tmp/ctranslate2-*.whl
# Install server requirements
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
EXPOSE 9090
CMD ["python", "run_server.py", "--port", "9090", "--backend", "faster_whisper"]
+30
View File
@@ -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 .
+78
View File
@@ -0,0 +1,78 @@
"""
Manual audio chunking example for WhisperLive.
Streams an audio file to a running WhisperLive server in real-time sized chunks,
printing partial transcripts when speech is detected and committed transcripts
when each segment is finalized.
Usage:
python examples/manual_audio_chunking.py --file assets/jfk.flac
"""
import argparse
import os
import sys
import time
import wave
try:
from whisper_live.client import StreamingTranscriptionClient
from whisper_live.utils import resample
except ImportError: # just in case whisper_live isn't installed.
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
print("[INFO] whisper_live not installed or the current version does not have StreamingTranscriptionClient. Will attempt to import from local source.")
from whisper_live.client import StreamingTranscriptionClient
from whisper_live.utils import resample
SAMPLE_RATE = 16000
def stream_audio_file(path: str, client: StreamingTranscriptionClient, chunk_ms: int = 50) -> None:
"""Read an audio file, resample to 16 kHz mono if needed, and pace chunks in real time."""
resampled_path = resample(path)
try:
with wave.open(resampled_path, "rb") as wf:
frames_per_chunk = SAMPLE_RATE * chunk_ms // 1000
chunk_duration = frames_per_chunk / SAMPLE_RATE
while chunk := wf.readframes(frames_per_chunk):
client.send(chunk, pcm_format="int16")
time.sleep(chunk_duration)
finally:
os.remove(resampled_path)
def main():
parser = argparse.ArgumentParser(description="Stream an audio file to WhisperLive.")
parser.add_argument("--file", "-f", required=True, help="Audio file to transcribe (any format supported by ffmpeg).")
parser.add_argument("--server", "-s", default="localhost")
parser.add_argument("--port", "-p", type=int, default=9090)
parser.add_argument("--model", "-m", default="small")
parser.add_argument("--lang", "-l", default="en")
parser.add_argument("--chunk_ms", type=int, default=50, help="Chunk size in ms.")
args = parser.parse_args()
client = StreamingTranscriptionClient(
args.server, args.port,
lang=args.lang,
model=args.model,
on_session_started=lambda: print("[INFO] Server ready.\n"),
on_partial_transcript=lambda text, _: print(f"\r{text:<80}", end="", flush=True),
on_committed_transcript=lambda text, _: print(f"\r{text:<80}"),
on_error=lambda e: print(f"\n[ERROR] {e}"),
on_close=lambda: print("\n[INFO] Connection closed."),
)
with client:
print(f"[INFO] Streaming {args.file} in {args.chunk_ms} ms chunks.")
stream_audio_file(args.file, client, chunk_ms=args.chunk_ms)
print("\n[INFO] Final transcript:")
for seg in client.transcript:
print(f" [{float(seg['start']):.2f}s → {float(seg['end']):.2f}s] {seg['text'].strip()}")
seg = client.last_partial
if seg:
print(f" [{float(seg['start']):.2f}s → {float(seg['end']):.2f}s] {seg['text'].strip()} (partial)")
if __name__ == "__main__":
main()
+5
View File
@@ -0,0 +1,5 @@
[pytest]
testpaths = tests
python_files = test_*.py
python_classes = Test*
python_functions = test_*
+1 -1
View File
@@ -1,4 +1,4 @@
PyAudio
ffmpeg-python
av
scipy
websocket-client
+28 -6
View File
@@ -1,7 +1,29 @@
PyAudio
faster-whisper==0.9.0
--extra-index-url https://download.pytorch.org/whl/cu111
torch==1.10.1
torchaudio==0.10.1
faster-whisper==1.2.0
websockets
onnxruntime==1.16.0
onnxruntime>=1.17.0,<1.20.0; python_version < "3.10"
onnxruntime>=1.20.0,<2; python_version >= "3.10"
numba
kaldialign
soundfile
scipy
av
jiwer
evaluate
numpy>=1.26.4,<2.5
openai-whisper==20250625
pyannote.audio
tokenizers==0.20.3
transformers[torch]
sentencepiece
# openvino
librosa
openvino
openvino-genai
openvino-tokenizers
optimum
optimum-intel
fastapi
uvicorn
python-multipart
+105
View File
@@ -0,0 +1,105 @@
from pathlib import Path
import sys
from whisper_live.client import TranscriptionClient
import argparse
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('--server', '-s',
type=str,
default='localhost',
help='hostname or ip address of server')
parser.add_argument('--files', '-f',
type=str,
nargs='+',
help='Files to transcribe, separated by spaces. '
'If not provided, will use microphone input.')
parser.add_argument('--output_file', '-o',
type=str,
default='./output_recording.wav',
help='output recording filename, only used for microphone input.')
parser.add_argument('--model', '-m',
type=str,
default='small',
help='Model to use for transcription, e.g., "tiny, small.en, large-v3".')
parser.add_argument('--lang', '-l',
type=str,
default='en',
help='Language code for transcription, e.g., "en" for English.')
parser.add_argument('--translate', '-t',
action='store_true',
help='Use Whisper built-in translation to English (sets task=translate). '
'For any-to-any translation, use --enable_translation instead.')
parser.add_argument('--mute_audio_playback', '-a',
action='store_true',
help='Mute audio playback during transcription.')
parser.add_argument('--save_output_recording', '-r',
action='store_true',
help='Save the output recording, only used for microphone input.')
parser.add_argument('--enable_translation',
action='store_true',
help='Enable any-to-any translation via M2M100 model (separate from Whisper --translate).')
parser.add_argument('--target_language', '-tl',
type=str,
default='fr',
help='Target language for translation, e.g., "fr" for French.')
parser.add_argument('--enable_timestamps',
action='store_true',
help='Show transcription with timestamps')
parser.add_argument('--n_display_segments',
type=int,
default=4,
help='Number of transcript segments to display in terminal (default: 4).')
args = parser.parse_args()
if args.translate and args.enable_translation:
print("[WARN]: Both --translate and --enable_translation are set. "
"--translate uses Whisper's built-in to-English translation, "
"while --enable_translation uses M2M100 for any-to-any. "
"Both will be active.")
client = TranscriptionClient(
args.server,
args.port,
lang=args.lang,
translate=args.translate,
model=args.model, # also support hf_model => `Systran/faster-whisper-small`
use_vad=True,
save_output_recording=args.save_output_recording, # Only used for microphone input, False by Default
output_recording_filename=args.output_file, # Only used for microphone input
mute_audio_playback=args.mute_audio_playback, # Only used for file input, False by Default
enable_translation=args.enable_translation, # Enable translation of the transcription output
target_language=args.target_language, # Target language for translation, e.g., "fr
enable_timestamps=args.enable_timestamps,
display_segments=args.n_display_segments,
)
if args.files is None:
client()
sys.exit(0)
# Validate audio files
valid_files = []
for file_path in args.files:
path = Path(file_path)
if path.exists() and path.is_file():
valid_files.append(str(path))
else:
print(f"Warning: File not found: {file_path}")
if not valid_files:
print("Error: No valid audio files found!")
sys.exit(1)
print(f"Found {len(valid_files)} audio file(s) to stream:")
for file_path in valid_files:
print(f" - {file_path}")
for f in valid_files:
client(f)
+142 -2
View File
@@ -1,5 +1,145 @@
from whisper_live.server import TranscriptionServer
import argparse
import os
import threading
import logging
from fastapi import FastAPI
from fastapi import UploadFile, Form
import uvicorn
import tempfile
import shutil
import json
from starlette.responses import PlainTextResponse, JSONResponse
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.')
parser.add_argument('--max_clients',
type=int,
default=4,
help='Maximum clients supported by the server.')
parser.add_argument('--max_connection_time',
type=int,
default=300,
help='The maximum duration (in seconds) a client can stay connected. Defaults to 300 seconds (5 minutes)')
parser.add_argument('--cache_path', '-c',
type=str,
default="~/.cache/whisper-live/",
help='Path to cache the converted ctranslate2 models.')
parser.add_argument(
"--rest_port", type=int, default=8000, help="Port for the REST API server."
)
parser.add_argument(
"--enable_rest",
action="store_true",
help="Enable the OpenAI-compatible REST API endpoint.",
)
parser.add_argument(
'--cors-origins',
type=str,
default=None,
help="Comma-separated list of allowed CORS origins (e.g., 'http://localhost:3000,http://example.com'). Defaults to localhost/127.0.0.1 on the WebSocket port."
)
parser.add_argument(
'--batch_inference',
action='store_true',
help='Enable batched GPU inference for concurrent sessions. '
'Batches multiple sessions into a single GPU call for higher throughput.'
)
parser.add_argument(
'--batch_max_size',
type=int,
default=8,
help='Maximum batch size for batched inference (default: 8).'
)
parser.add_argument(
'--batch_window_ms',
type=int,
default=50,
help='Maximum time in ms to wait for batch to fill (default: 50).'
)
parser.add_argument(
'--raw_pcm_input',
action='store_true',
help='Expect raw PCM int16 audio from clients instead of float32. '
'Audio will be normalized to float32 range [-1.0, 1.0].'
)
parser.add_argument(
'--metrics_port',
type=int,
default=0,
help='Port for Prometheus /metrics endpoint. 0 = disabled (default). Requires prometheus_client.'
)
parser.add_argument(
'--api_key',
type=str,
default=None,
help='Optional API key for authenticating REST API and WebSocket connections. '
'Clients must send "Authorization: Bearer <key>" header or "?token=<key>" query parameter.'
)
parser.add_argument(
'--rate_limit_rpm',
type=int,
default=0,
help='Maximum REST API requests per minute per client IP. 0 = unlimited (default).'
)
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,
max_clients=args.max_clients,
max_connection_time=args.max_connection_time,
cache_path=args.cache_path,
rest_port=args.rest_port,
enable_rest=args.enable_rest,
cors_origins=args.cors_origins,
batch_enabled=args.batch_inference,
batch_max_size=args.batch_max_size,
batch_window_ms=args.batch_window_ms,
raw_pcm_input=args.raw_pcm_input,
metrics_port=args.metrics_port,
api_key=args.api_key,
rate_limit_rpm=args.rate_limit_rpm,
)
+120
View File
@@ -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"
+32
View File
@@ -0,0 +1,32 @@
#!/bin/bash
# Detect the operating system
if [[ "$OSTYPE" == "darwin"* ]]; then
# macOS
echo "Detected macOS, using Homebrew for installation"
# Check if Homebrew is installed
if ! command -v brew &> /dev/null; then
echo "Homebrew not found. Please install Homebrew first: https://brew.sh/"
exit 1
fi
# Install packages using Homebrew
brew install portaudio wget
elif [[ "$OSTYPE" == "linux-gnu"* ]]; then
# Linux
if [[ -f /etc/os-release ]]; then
source /etc/os-release
fi
if [[ "$(command -v dnf)" ]]; then
echo "Detected Fedora, using dnf for installation"
dnf install -y portaudio-devel wget
else
echo "Detected Linux (assuming Debian/Ubuntu), using apt-get for installation"
apt-get install -y portaudio19-dev wget
fi
else
echo "Unsupported operating system: $OSTYPE"
exit 1
fi
+66 -36
View File
@@ -10,45 +10,75 @@ 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.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"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",
"av",
"faster-whisper==1.2.0",
"torch",
"torchaudio",
"websockets",
"onnxruntime",
"ffmpeg-python",
"onnxruntime>=1.17.0,<1.20.0; python_version < '3.10'",
"onnxruntime>=1.20.0,<2; python_version >= '3.10'",
"scipy",
"websocket-client",
],
python_requires=">=3.8"
)
"numba",
"openai-whisper==20250625",
"kaldialign",
"soundfile",
"tokenizers==0.20.3",
"librosa",
"numpy>=1.26.4,<2.5",
"openvino",
"openvino-genai",
"openvino-tokenizers",
"optimum",
"optimum-intel",
"fastapi",
"uvicorn",
"python-multipart",
# CTranslate2 (faster-whisper's backend) is hard-linked against
# libcublas.so.12 / libcudnn.so.9 but doesn't declare the matching
# wheels as runtime deps. torch >=2.12 also dropped the cu12
# wheels in favor of cu13, so users no longer get cu12 transitively.
# Without these wheels GPU inference dies at first transcription:
# ERROR: Library libcublas.so.12 is not found or cannot be loaded
# Skip only for CPU-only inference.
"nvidia-cublas-cu12; sys_platform == 'linux'",
"nvidia-cudnn-cu12; sys_platform == 'linux'",
],
python_requires=">=3.9"
)
-3
View File
@@ -1,3 +0,0 @@
#! /bin/bash
apt-get install portaudio19-dev ffmpeg wget -y
View File
+644
View File
@@ -0,0 +1,644 @@
import json
import queue
import threading
import time
import unittest
from unittest.mock import MagicMock, patch
import numpy as np
from whisper_live.backend.base import ServeClientBase
class ConcreteServeClient(ServeClientBase):
"""Concrete subclass for testing the abstract base class."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.language = "en"
def transcribe_audio(self, input_sample):
return None
def handle_transcription_output(self, result, duration):
pass
class WaitTrackingEvent:
"""Threading event that records when wait() is entered."""
def __init__(self):
self._event = threading.Event()
self.wait_started = threading.Event()
def wait(self, timeout=None):
self.wait_started.set()
return self._event.wait(timeout)
def set(self):
self._event.set()
def __getattr__(self, name):
return getattr(self._event, name)
class TestServeClientBaseInit(unittest.TestCase):
def test_default_values(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="test-uid", websocket=ws)
self.assertEqual(client.client_uid, "test-uid")
self.assertEqual(client.send_last_n_segments, 10)
self.assertAlmostEqual(client.no_speech_thresh, 0.45)
self.assertFalse(client.clip_audio)
self.assertEqual(client.same_output_threshold, 10)
self.assertIsNone(client.frames_np)
self.assertAlmostEqual(client.timestamp_offset, 0.0)
self.assertFalse(client.exit)
self.assertEqual(client.transcript, [])
def test_custom_values(self):
ws = MagicMock()
q = queue.Queue()
client = ConcreteServeClient(
client_uid="uid2",
websocket=ws,
send_last_n_segments=5,
no_speech_thresh=0.6,
clip_audio=True,
same_output_threshold=20,
translation_queue=q,
)
self.assertEqual(client.send_last_n_segments, 5)
self.assertAlmostEqual(client.no_speech_thresh, 0.6)
self.assertTrue(client.clip_audio)
self.assertEqual(client.same_output_threshold, 20)
self.assertIs(client.translation_queue, q)
class TestAddFrames(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
def test_first_frame_initializes_buffer(self):
frame = np.array([0.1, 0.2, 0.3], dtype=np.float32)
self.client.add_frames(frame)
np.testing.assert_array_equal(self.client.frames_np, frame)
def test_subsequent_frames_concatenated(self):
frame1 = np.array([0.1, 0.2], dtype=np.float32)
frame2 = np.array([0.3, 0.4], dtype=np.float32)
self.client.add_frames(frame1)
self.client.add_frames(frame2)
expected = np.array([0.1, 0.2, 0.3, 0.4], dtype=np.float32)
np.testing.assert_array_equal(self.client.frames_np, expected)
def test_buffer_trimmed_at_45_seconds(self):
# 45 seconds + 1 sample at 16kHz = 720001 samples
self.client.frames_np = np.zeros(45 * 16000 + 1, dtype=np.float32)
self.client.add_frames(np.array([1.0], dtype=np.float32))
# after trimming 30s, buffer should be ~15s + 1 original + 1 new
expected_len = (45 * 16000 + 1) - (30 * 16000) + 1
self.assertEqual(self.client.frames_np.shape[0], expected_len)
self.assertAlmostEqual(self.client.frames_offset, 30.0)
def test_timestamp_offset_updated_on_trim(self):
self.client.frames_np = np.zeros(45 * 16000 + 1, dtype=np.float32)
self.client.timestamp_offset = 5.0 # behind frames_offset after trim
self.client.add_frames(np.array([1.0], dtype=np.float32))
# timestamp_offset should be bumped to at least frames_offset
self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset)
class TestAddFramesThreadSafety(unittest.TestCase):
def test_concurrent_add_frames(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="test", websocket=ws)
errors = []
def add_many():
try:
for _ in range(100):
client.add_frames(np.random.randn(160).astype(np.float32))
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=add_many) for _ in range(4)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
self.assertIsNotNone(client.frames_np)
def test_exception_releases_lock_without_signaling_frames_ready(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="test", websocket=ws)
client.frames_np = np.array([0.1], dtype=np.float32)
with patch("whisper_live.backend.base.np.concatenate", side_effect=RuntimeError("boom")):
with self.assertRaisesRegex(RuntimeError, "boom"):
client.add_frames(np.array([0.2], dtype=np.float32))
self.assertFalse(client.lock.locked())
self.assertFalse(client.frames_ready.is_set())
class TestGetAudioChunkForProcessing(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
def test_empty_buffer_returns_empty(self):
self.client.frames_np = np.array([], dtype=np.float32)
chunk, duration = self.client.get_audio_chunk_for_processing()
self.assertEqual(duration, 0.0)
self.assertEqual(chunk.shape[0], 0)
def test_full_buffer_no_offset(self):
audio = np.random.randn(16000).astype(np.float32) # 1 second
self.client.frames_np = audio
chunk, duration = self.client.get_audio_chunk_for_processing()
self.assertAlmostEqual(duration, 1.0)
np.testing.assert_array_equal(chunk, audio)
def test_with_offset(self):
audio = np.random.randn(32000).astype(np.float32) # 2 seconds
self.client.frames_np = audio
self.client.timestamp_offset = 1.0 # skip first second
chunk, duration = self.client.get_audio_chunk_for_processing()
self.assertAlmostEqual(duration, 1.0)
self.assertEqual(chunk.shape[0], 16000)
class TestClipAudioIfNoValidSegment(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(
client_uid="test", websocket=self.ws, clip_audio=True
)
def test_clips_when_chunk_exceeds_25s(self):
# 30 seconds of audio with no valid segments
self.client.frames_np = np.zeros(30 * 16000, dtype=np.float32)
self.client.timestamp_offset = 0.0
self.client.frames_offset = 0.0
self.client.clip_audio_if_no_valid_segment()
# offset should have advanced to leave ~5s of remaining audio
expected_offset = (30 * 16000 / 16000) - 5
self.assertAlmostEqual(self.client.timestamp_offset, expected_offset, places=1)
def test_no_clip_when_short(self):
self.client.frames_np = np.zeros(10 * 16000, dtype=np.float32)
self.client.timestamp_offset = 0.0
self.client.frames_offset = 0.0
self.client.clip_audio_if_no_valid_segment()
self.assertAlmostEqual(self.client.timestamp_offset, 0.0)
class TestPrepareSegments(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(
client_uid="test", websocket=self.ws, send_last_n_segments=3
)
def test_empty_transcript_no_last(self):
segments = self.client.prepare_segments()
self.assertEqual(segments, [])
def test_empty_transcript_with_last(self):
last = {"start": "0.000", "end": "1.000", "text": "hello", "completed": False}
segments = self.client.prepare_segments(last_segment=last)
self.assertEqual(len(segments), 1)
self.assertEqual(segments[0]["text"], "hello")
def test_fewer_than_n_segments(self):
self.client.transcript = [
{"start": "0.000", "end": "1.000", "text": "a", "completed": True},
{"start": "1.000", "end": "2.000", "text": "b", "completed": True},
]
segments = self.client.prepare_segments()
self.assertEqual(len(segments), 2)
def test_more_than_n_segments_truncated(self):
self.client.transcript = [
{"start": f"{i}.000", "end": f"{i+1}.000", "text": f"seg{i}", "completed": True}
for i in range(10)
]
segments = self.client.prepare_segments()
self.assertEqual(len(segments), 3)
self.assertEqual(segments[0]["text"], "seg7")
def test_last_segment_appended(self):
self.client.transcript = [
{"start": "0.000", "end": "1.000", "text": "a", "completed": True},
]
last = {"start": "1.000", "end": "2.000", "text": "in progress", "completed": False}
segments = self.client.prepare_segments(last_segment=last)
self.assertEqual(len(segments), 2)
self.assertEqual(segments[-1]["text"], "in progress")
class TestFormatSegment(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
def test_format(self):
seg = self.client.format_segment(1.234, 5.678, "hello world", completed=True)
self.assertEqual(seg["start"], "1.234")
self.assertEqual(seg["end"], "5.678")
self.assertEqual(seg["text"], "hello world")
self.assertTrue(seg["completed"])
def test_format_not_completed(self):
seg = self.client.format_segment(0.0, 1.0, "text")
self.assertFalse(seg["completed"])
class TestSendTranscriptionToClient(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test-uid", websocket=self.ws)
def test_sends_json(self):
segments = [{"start": "0.000", "end": "1.000", "text": "hi", "completed": True}]
self.client.send_transcription_to_client(segments)
self.ws.send.assert_called_once()
sent = json.loads(self.ws.send.call_args[0][0])
self.assertEqual(sent["uid"], "test-uid")
self.assertEqual(len(sent["segments"]), 1)
def test_send_failure_logged_not_raised(self):
self.ws.send.side_effect = ConnectionError("broken pipe")
# should not raise
self.client.send_transcription_to_client([])
class TestDisconnect(unittest.TestCase):
def test_sends_disconnect_message(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="uid1", websocket=ws)
client.disconnect()
sent = json.loads(ws.send.call_args[0][0])
self.assertEqual(sent["uid"], "uid1")
self.assertEqual(sent["message"], "DISCONNECT")
class TestCleanup(unittest.TestCase):
def test_sets_exit_flag(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="uid1", websocket=ws)
self.assertFalse(client.exit)
client.cleanup()
self.assertTrue(client.exit)
def _supports_thread_time():
thread_time = getattr(time, "thread_time", None)
if thread_time is None:
return False
try:
thread_time()
except NotImplementedError:
return False
return True
class TestSpeechToTextWaitingBehavior(unittest.TestCase):
"""Tests the first-frame wait behavior in speech_to_text()."""
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
self.client.frames_ready = WaitTrackingEvent()
self.transcribe_called = threading.Event()
self.thread_started = threading.Event()
self.cpu_used = None
self.thread = None
def tearDown(self):
if self.thread is not None and self.thread.is_alive():
self.client.exit = True
# release wait() directly so a broken cleanup() cannot hang the test process
self.client.frames_ready.set()
self.thread.join(timeout=1.0)
def _start_speech_thread(self, target=None):
self.thread = threading.Thread(target=target or self.client.speech_to_text)
self.thread.start()
return self.thread
def _join_speech_thread(self):
self.thread.join(timeout=1.0)
return not self.thread.is_alive()
def _transcribe_once(self, input_sample):
# mark the first processing step after wait and stop the loop
self.transcribe_called.set()
self.client.exit = True
return []
def _measure_waiting_cpu(self):
# measure CPU consumed by speech_to_text loop while it waits for the first frame
self.thread_started.set()
start_cpu = time.thread_time()
self.client.speech_to_text()
self.cpu_used = time.thread_time() - start_cpu
def test_waits_for_first_frame_before_transcribing(self):
self.client.transcribe_audio = MagicMock(side_effect=self._transcribe_once)
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.assertFalse(self.transcribe_called.is_set())
self.client.add_frames(np.zeros(self.client.RATE, dtype=np.float32))
self.assertTrue(self.transcribe_called.wait(timeout=1.0))
self.assertTrue(self._join_speech_thread())
def test_cleanup_unblocks_waiting_thread_without_audio(self):
self.client.transcribe_audio = MagicMock()
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.client.cleanup()
self.assertTrue(self._join_speech_thread())
self.assertTrue(self.client.exit)
self.client.transcribe_audio.assert_not_called()
def test_exit_flag_unblocks_waiting_thread_without_signal(self):
self.client.transcribe_audio = MagicMock()
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.client.exit = True
self.assertTrue(self._join_speech_thread())
self.client.transcribe_audio.assert_not_called()
@unittest.skipUnless(_supports_thread_time(), "time.thread_time() not supported")
def test_waiting_for_first_frame_uses_negligible_thread_cpu(self):
self._start_speech_thread(target=self._measure_waiting_cpu)
self.assertTrue(self.thread_started.wait(timeout=1.0))
time.sleep(0.25)
self.client.cleanup()
self.assertTrue(self._join_speech_thread())
self.assertIsNotNone(self.cpu_used)
self.assertLess(self.cpu_used, 0.1)
class TestTrimTranscript(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
def test_transcript_trimmed_when_over_max(self):
self.client.transcript = [
{"start": f"{i}.000", "end": f"{i+1}.000", "text": f"seg{i}", "completed": True}
for i in range(self.client.MAX_TRANSCRIPT_LENGTH + 100)
]
self.client._trim_transcript()
self.assertEqual(len(self.client.transcript), self.client.MAX_TRANSCRIPT_LENGTH)
self.assertEqual(self.client.transcript[0]["text"], "seg100")
def test_transcript_not_trimmed_when_under_max(self):
self.client.transcript = [
{"start": "0.000", "end": "1.000", "text": "a", "completed": True}
]
self.client._trim_transcript()
self.assertEqual(len(self.client.transcript), 1)
def test_text_list_trimmed(self):
self.client.text = ["word"] * (self.client.MAX_TRANSCRIPT_LENGTH + 50)
self.client._trim_transcript()
self.assertEqual(len(self.client.text), self.client.MAX_TRANSCRIPT_LENGTH)
class TestUpdateSegments(unittest.TestCase):
"""Tests for the core update_segments() logic."""
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(
client_uid="test",
websocket=self.ws,
no_speech_thresh=0.45,
same_output_threshold=3,
)
self.client.frames_np = np.zeros(16000 * 5, dtype=np.float32)
def _make_segment(self, start, end, text, no_speech_prob=0.0):
seg = MagicMock()
seg.start = start
seg.end = end
seg.text = text
seg.no_speech_prob = no_speech_prob
return seg
def test_single_segment_becomes_last(self):
segs = [self._make_segment(0.0, 1.0, " hello")]
last = self.client.update_segments(segs, duration=2.0)
self.assertIsNotNone(last)
self.assertIn("hello", last["text"])
self.assertFalse(last["completed"])
self.assertEqual(len(self.client.transcript), 0)
def test_multiple_segments_completes_all_but_last(self):
segs = [
self._make_segment(0.0, 1.0, " first"),
self._make_segment(1.0, 2.0, " second"),
]
last = self.client.update_segments(segs, duration=3.0)
self.assertEqual(len(self.client.transcript), 1)
self.assertTrue(self.client.transcript[0]["completed"])
self.assertIn("first", self.client.transcript[0]["text"])
self.assertIsNotNone(last)
self.assertIn("second", last["text"])
def test_high_no_speech_prob_skipped(self):
segs = [
self._make_segment(0.0, 1.0, " noise", no_speech_prob=0.9),
self._make_segment(1.0, 2.0, " also noise", no_speech_prob=0.9),
]
last = self.client.update_segments(segs, duration=3.0)
self.assertEqual(len(self.client.transcript), 0)
self.assertIsNone(last)
def test_segment_with_start_gte_end_skipped(self):
segs = [
self._make_segment(1.0, 0.5, " backwards"),
self._make_segment(1.5, 2.0, " normal"),
]
last = self.client.update_segments(segs, duration=3.0)
self.assertEqual(len(self.client.transcript), 0)
self.assertIsNotNone(last)
def test_repeated_output_triggers_completion(self):
seg = self._make_segment(0.0, 1.0, " repeated")
for _ in range(self.client.same_output_threshold + 2):
last = self.client.update_segments([seg], duration=2.0)
# after enough repeats, should be added to transcript
self.assertTrue(len(self.client.transcript) >= 1)
def test_translation_queue_receives_completed(self):
q = queue.Queue()
self.client.translation_queue = q
segs = [
self._make_segment(0.0, 1.0, " first"),
self._make_segment(1.0, 2.0, " second"),
]
self.client.update_segments(segs, duration=3.0)
self.assertFalse(q.empty())
item = q.get_nowait()
self.assertIn("first", item["text"])
def test_timestamp_offset_advances(self):
segs = [
self._make_segment(0.0, 1.0, " first"),
self._make_segment(1.0, 2.0, " second"),
]
self.client.update_segments(segs, duration=3.0)
self.assertGreater(self.client.timestamp_offset, 0.0)
class TestGetSegmentHelpers(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
def test_get_segment_no_speech_prob_attr(self):
seg = MagicMock()
seg.no_speech_prob = 0.3
self.assertAlmostEqual(self.client.get_segment_no_speech_prob(seg), 0.3)
def test_get_segment_no_speech_prob_fallback(self):
seg = MagicMock(spec=[]) # no attributes
self.assertEqual(self.client.get_segment_no_speech_prob(seg), 0)
def test_get_segment_start_uses_start(self):
seg = MagicMock()
seg.start = 1.5
self.assertAlmostEqual(self.client.get_segment_start(seg), 1.5)
def test_get_segment_end_uses_end(self):
seg = MagicMock()
seg.end = 3.0
self.assertAlmostEqual(self.client.get_segment_end(seg), 3.0)
def test_get_segment_start_fallback_to_start_ts(self):
seg = MagicMock(spec=["start_ts"])
seg.start_ts = 2.0
self.assertAlmostEqual(self.client.get_segment_start(seg), 2.0)
class TestWordTimestamps(unittest.TestCase):
"""Tests for word-level timestamp extraction."""
def _make_client(self, word_timestamps=False):
ws = MagicMock()
return ConcreteServeClient(
client_uid="wt-uid", websocket=ws, word_timestamps=word_timestamps
)
def _make_word(self, word, start, end, prob):
w = MagicMock()
w.word = word
w.start = start
w.end = end
w.probability = prob
return w
def _make_segment(self, text, start, end, no_speech_prob=0.0, words=None):
seg = MagicMock()
seg.text = text
seg.start = start
seg.end = end
seg.no_speech_prob = no_speech_prob
seg.words = words
return seg
def test_word_timestamps_disabled_by_default(self):
client = self._make_client()
self.assertFalse(client.word_timestamps)
def test_word_timestamps_enabled(self):
client = self._make_client(word_timestamps=True)
self.assertTrue(client.word_timestamps)
def test_extract_words_when_disabled(self):
client = self._make_client(word_timestamps=False)
seg = self._make_segment("hello", 0.0, 1.0, words=[self._make_word("hello", 0.0, 0.5, 0.99)])
result = client._extract_words(seg, 0.0)
self.assertIsNone(result)
def test_extract_words_when_enabled(self):
client = self._make_client(word_timestamps=True)
words = [
self._make_word("hello", 0.0, 0.3, 0.95),
self._make_word("world", 0.4, 0.8, 0.88),
]
seg = self._make_segment("hello world", 0.0, 1.0, words=words)
result = client._extract_words(seg, 10.0)
self.assertEqual(len(result), 2)
self.assertEqual(result[0]["word"], "hello")
self.assertEqual(result[0]["start"], "10.000")
self.assertEqual(result[0]["end"], "10.300")
self.assertEqual(result[0]["probability"], 0.95)
self.assertEqual(result[1]["word"], "world")
self.assertEqual(result[1]["start"], "10.400")
def test_extract_words_no_words_on_segment(self):
client = self._make_client(word_timestamps=True)
seg = self._make_segment("hello", 0.0, 1.0, words=None)
result = client._extract_words(seg, 0.0)
self.assertIsNone(result)
def test_format_segment_without_words(self):
client = self._make_client()
seg = client.format_segment(0.0, 1.0, "hello")
self.assertNotIn("words", seg)
def test_format_segment_with_words(self):
client = self._make_client(word_timestamps=True)
words = [{"word": "hello", "start": "0.000", "end": "0.500", "probability": 0.95}]
seg = client.format_segment(0.0, 1.0, "hello", words=words)
self.assertIn("words", seg)
self.assertEqual(len(seg["words"]), 1)
self.assertEqual(seg["words"][0]["word"], "hello")
def test_update_segments_includes_words(self):
client = self._make_client(word_timestamps=True)
words1 = [self._make_word("hello", 0.0, 0.5, 0.9)]
words2 = [self._make_word("world", 0.6, 1.0, 0.85)]
segments = [
self._make_segment(" hello", 0.0, 0.5, words=words1),
self._make_segment(" world", 0.6, 1.0, words=words2),
]
last = client.update_segments(segments, 2.0)
# First segment should be completed (in transcript) with words
self.assertTrue(len(client.transcript) > 0)
self.assertIn("words", client.transcript[-1])
# Last segment should be in-progress with words
self.assertIsNotNone(last)
self.assertIn("words", last)
def test_update_segments_no_words_when_disabled(self):
client = self._make_client(word_timestamps=False)
words1 = [self._make_word("hello", 0.0, 0.5, 0.9)]
words2 = [self._make_word("world", 0.6, 1.0, 0.85)]
segments = [
self._make_segment(" hello", 0.0, 0.5, words=words1),
self._make_segment(" world", 0.6, 1.0, words=words2),
]
last = client.update_segments(segments, 2.0)
self.assertTrue(len(client.transcript) > 0)
self.assertNotIn("words", client.transcript[-1])
self.assertNotIn("words", last)
if __name__ == "__main__":
unittest.main()
+163
View File
@@ -0,0 +1,163 @@
import time
import unittest
from unittest import mock
from unittest.mock import MagicMock
import numpy as np
from whisper_live.batch_inference import BatchInferenceWorker, BatchRequest
class TestBatchInferenceWorker(unittest.TestCase):
def setUp(self):
self.mock_transcriber = MagicMock()
self.worker = BatchInferenceWorker(
transcriber=self.mock_transcriber,
max_batch_size=8,
batch_window_ms=200,
)
self.worker.start()
def tearDown(self):
self.worker.stop()
def _make_audio(self, duration_s=1.0):
return np.random.randn(int(16000 * duration_s)).astype(np.float32)
def test_single_request_uses_transcribe(self):
"""Single request should fall back to transcriber.transcribe()."""
fake_segment = MagicMock()
fake_info = MagicMock()
self.mock_transcriber.transcribe.return_value = ([fake_segment], fake_info)
req = BatchRequest(audio=self._make_audio(), language="en", use_vad=False)
self.worker.submit(req)
req.future.wait(timeout=5)
self.assertTrue(req.future.is_set())
self.assertIsNone(req.error)
self.assertEqual(req.result, [fake_segment])
self.assertEqual(req.info, fake_info)
self.mock_transcriber.transcribe.assert_called_once()
@mock.patch('whisper_live.batch_inference.get_suppressed_tokens', return_value=[-1])
@mock.patch('whisper_live.batch_inference.Tokenizer')
def test_multiple_requests_batched(self, mock_tokenizer_cls, mock_suppress):
"""Multiple concurrent requests should go through the batched GPU path."""
# Mock tokenizer
mock_tok = MagicMock()
mock_tok.decode.return_value = "hello world"
mock_tokenizer_cls.return_value = mock_tok
# Mock feature extractor
self.mock_transcriber.feature_extractor.return_value = np.zeros(
(80, 3000), dtype=np.float32
)
self.mock_transcriber.feature_extractor.sampling_rate = 16000
# Mock encode
self.mock_transcriber.encode.return_value = np.zeros(
(3, 1500, 512), dtype=np.float32
)
# Mock model.generate — one result per item
gen_result = MagicMock()
gen_result.sequences_ids = [[50257, 50362, 1234, 50256]]
gen_result.scores = [np.float32(-1.0)]
gen_result.no_speech_prob = 0.1
self.mock_transcriber.model.generate.return_value = [gen_result] * 3
# Mock remaining model attributes
self.mock_transcriber.model.is_multilingual = False
self.mock_transcriber.max_length = 448
self.mock_transcriber.frames_per_second = 50
self.mock_transcriber.get_prompt.return_value = [50258]
self.mock_transcriber._split_segments_by_timestamps.return_value = (
[{"start": 0.0, "end": 1.0, "tokens": [1234], "seek": 0}],
None,
None,
)
requests = [
BatchRequest(audio=self._make_audio(), language="en", use_vad=False)
for _ in range(3)
]
for req in requests:
self.worker.submit(req)
for req in requests:
req.future.wait(timeout=5)
for req in requests:
self.assertTrue(req.future.is_set())
self.assertIsNone(req.error)
self.assertIsNotNone(req.result)
# Verify the batched encode path was used (not transcribe)
self.mock_transcriber.encode.assert_called()
self.mock_transcriber.transcribe.assert_not_called()
def test_error_propagation(self):
"""Transcriber errors should propagate to the request without crashing the worker."""
self.mock_transcriber.transcribe.side_effect = RuntimeError("GPU OOM")
req = BatchRequest(audio=self._make_audio(), language="en", use_vad=False)
self.worker.submit(req)
req.future.wait(timeout=5)
self.assertTrue(req.future.is_set())
self.assertIsInstance(req.error, RuntimeError)
self.assertIn("GPU OOM", str(req.error))
# Worker should still be alive — submit another request
self.mock_transcriber.transcribe.side_effect = None
self.mock_transcriber.transcribe.return_value = ([MagicMock()], MagicMock())
req2 = BatchRequest(audio=self._make_audio(), language="en", use_vad=False)
self.worker.submit(req2)
req2.future.wait(timeout=5)
self.assertIsNone(req2.error)
self.assertIsNotNone(req2.result)
def test_worker_stop(self):
"""Worker thread should exit cleanly when stop() is called."""
self.assertTrue(self.worker._thread.is_alive())
self.worker.stop()
self.assertFalse(self.worker._thread.is_alive())
def test_batch_respects_max_size(self):
"""Batches should not exceed max_batch_size."""
self.worker.stop() # Stop the default worker
observed_batch_sizes = []
original_process = BatchInferenceWorker._process_batch
def tracking_process(self_inner, batch):
observed_batch_sizes.append(len(batch))
original_process(self_inner, batch)
self.worker = BatchInferenceWorker(
transcriber=self.mock_transcriber,
max_batch_size=2,
batch_window_ms=100,
)
self.mock_transcriber.transcribe.return_value = ([MagicMock()], MagicMock())
with mock.patch.object(
BatchInferenceWorker, '_process_batch', tracking_process
):
self.worker.start()
requests = [
BatchRequest(audio=self._make_audio(), language="en", use_vad=False)
for _ in range(4)
]
for req in requests:
self.worker.submit(req)
for req in requests:
req.future.wait(timeout=5)
for size in observed_batch_sizes:
self.assertLessEqual(size, 2)
self.assertTrue(all(req.future.is_set() for req in requests))
+194
View File
@@ -0,0 +1,194 @@
import json
import os
import scipy
import websocket
import copy
import unittest
from io import StringIO
from unittest.mock import patch, MagicMock
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
from whisper_live.utils import print_transcript, 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):
self.client.on_open(self.mock_ws_app)
self.mock_ws_app.send.assert_called_once()
sent_message = json.loads(self.mock_ws_app.send.call_args[0][0])
self.assertEqual(sent_message["uid"], self.client.uid)
self.assertEqual(sent_message["language"], self.client.language)
self.assertEqual(sent_message["task"], self.client.task)
self.assertEqual(sent_message["model"], self.client.model)
self.assertTrue(sent_message["use_vad"])
self.assertEqual(sent_message["send_last_n_segments"], 10)
self.assertAlmostEqual(sent_message["no_speech_thresh"], 0.45)
self.assertFalse(sent_message["clip_audio"])
self.assertEqual(sent_message["same_output_threshold"], 10)
self.assertFalse(sent_message["enable_translation"])
self.assertEqual(sent_message["target_language"], "fr")
self.assertIsNone(sent_message["hotwords"])
self.assertFalse(sent_message["enable_diarization"])
self.assertEqual(sent_message["max_speakers"], 10)
self.assertFalse(sent_message["word_timestamps"])
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 TestPrintTranscript(unittest.TestCase):
@patch("whisper_live.utils.shutil.get_terminal_size")
@patch("sys.stdout", new_callable=StringIO)
def test_print_transcript_respects_narrow_terminals(self, mock_stdout, mock_terminal_size):
mock_terminal_size.return_value = os.terminal_size((20, 20))
print_transcript(["This transcript should still wrap cleanly on a narrow terminal."])
output_lines = [line for line in mock_stdout.getvalue().splitlines() if line.strip()]
self.assertGreater(len(output_lines), 1)
self.assertTrue(all(len(line) <= 20 for line in output_lines))
@patch("whisper_live.utils.shutil.get_terminal_size")
@patch("sys.stdout", new_callable=StringIO)
def test_print_transcript_indents_timestamp_continuations(self, mock_stdout, mock_terminal_size):
mock_terminal_size.return_value = os.terminal_size((32, 20))
print_transcript(
[{"start": "00:00", "end": "00:05", "text": "This line should wrap and keep its timestamp indentation."}],
timestamps=True,
)
output_lines = [line.rstrip() for line in mock_stdout.getvalue().splitlines() if line.strip()]
self.assertGreater(len(output_lines), 1)
self.assertTrue(output_lines[1].startswith(" " * len("[00:00 -> 00:05] ")))
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())
+305
View File
@@ -0,0 +1,305 @@
import json
import time
import unittest
from unittest.mock import patch, MagicMock, PropertyMock
from whisper_live.client import Client, TranscriptionTeeClient
class TestClientStatusMessages(unittest.TestCase):
"""Tests for Client.handle_status_messages() and on_message() branches."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def setUp(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
self.client = Client(host="localhost", port=9090, lang="en")
def tearDown(self):
self.client.close_websocket()
def test_wait_status(self):
msg = {"uid": self.client.uid, "status": "WAIT", "message": 5.0}
self.client.handle_status_messages(msg)
self.assertTrue(self.client.waiting)
def test_error_status(self):
msg = {"uid": self.client.uid, "status": "ERROR", "message": "model not found"}
self.client.handle_status_messages(msg)
self.assertTrue(self.client.server_error)
def test_warning_status_no_side_effects(self):
msg = {"uid": self.client.uid, "status": "WARNING", "message": "fallback backend"}
self.client.handle_status_messages(msg)
self.assertFalse(self.client.server_error)
self.assertFalse(self.client.waiting)
def test_on_message_wrong_uid_ignored(self):
msg = json.dumps({"uid": "wrong-uid", "segments": [{"start": 0, "end": 1, "text": "hi", "completed": True}]})
self.client.on_message(MagicMock(), msg)
self.assertEqual(len(self.client.transcript), 0)
def test_on_message_disconnect(self):
self.client.recording = True
msg = json.dumps({"uid": self.client.uid, "message": "DISCONNECT"})
self.client.on_message(MagicMock(), msg)
self.assertFalse(self.client.recording)
def test_on_message_server_ready(self):
msg = json.dumps({
"uid": self.client.uid,
"message": "SERVER_READY",
"backend": "faster_whisper",
})
self.client.on_message(MagicMock(), msg)
self.assertTrue(self.client.recording)
self.assertEqual(self.client.server_backend, "faster_whisper")
def test_on_message_language_detection(self):
msg = json.dumps({
"uid": self.client.uid,
"language": "fr",
"language_prob": 0.95,
})
self.client.on_message(MagicMock(), msg)
self.assertEqual(self.client.language, "fr")
class TestClientTranslationFlow(unittest.TestCase):
"""Tests for the translation-related client functionality."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def setUp(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
self.client = Client(
host="localhost",
port=9090,
lang="en",
enable_translation=True,
target_language="es",
)
# simulate SERVER_READY so server_backend is set
ready_msg = json.dumps({
"uid": self.client.uid,
"message": "SERVER_READY",
"backend": "faster_whisper",
})
self.client.on_message(MagicMock(), ready_msg)
def tearDown(self):
self.client.close_websocket()
def test_on_open_includes_translation_fields(self):
mock_ws = MagicMock()
self.client.on_open(mock_ws)
sent = json.loads(mock_ws.send.call_args[0][0])
self.assertTrue(sent["enable_translation"])
self.assertEqual(sent["target_language"], "es")
def test_translated_segments_processed(self):
msg = json.dumps({
"uid": self.client.uid,
"translated_segments": [
{"start": "0.000", "end": "1.000", "text": "Hola mundo", "completed": True},
],
})
self.client.on_message(MagicMock(), msg)
self.assertEqual(len(self.client.translated_transcript), 1)
self.assertEqual(self.client.translated_transcript[0]["text"], "Hola mundo")
def test_translation_callback_invoked(self):
callback = MagicMock()
self.client.translation_callback = callback
msg = json.dumps({
"uid": self.client.uid,
"translated_segments": [
{"start": "0.000", "end": "1.000", "text": "Hola", "completed": True},
],
})
self.client.on_message(MagicMock(), msg)
callback.assert_called_once()
def test_translation_callback_exception_handled(self):
callback = MagicMock(side_effect=RuntimeError("callback broke"))
self.client.translation_callback = callback
msg = json.dumps({
"uid": self.client.uid,
"translated_segments": [
{"start": "0.000", "end": "1.000", "text": "Hola", "completed": True},
],
})
# should not raise
self.client.on_message(MagicMock(), msg)
class TestClientTranscriptionCallback(unittest.TestCase):
"""Tests for the transcription callback feature."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def setUp(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
self.callback = MagicMock()
self.client = Client(
host="localhost",
port=9090,
lang="en",
transcription_callback=self.callback,
)
ready_msg = json.dumps({
"uid": self.client.uid,
"message": "SERVER_READY",
"backend": "faster_whisper",
})
self.client.on_message(MagicMock(), ready_msg)
def tearDown(self):
self.client.close_websocket()
def test_callback_receives_text_and_segments(self):
msg = json.dumps({
"uid": self.client.uid,
"segments": [
{"start": "0.000", "end": "1.000", "text": "Hello", "completed": True},
],
})
self.client.on_message(MagicMock(), msg)
self.callback.assert_called_once()
text_arg, segments_arg = self.callback.call_args[0]
self.assertIn("Hello", text_arg)
self.assertIsInstance(segments_arg, list)
def test_callback_exception_does_not_crash(self):
self.callback.side_effect = ValueError("boom")
msg = json.dumps({
"uid": self.client.uid,
"segments": [
{"start": "0.000", "end": "1.000", "text": "Test", "completed": True},
],
})
# should not raise
self.client.on_message(MagicMock(), msg)
class TestClientSrtWriting(unittest.TestCase):
"""Tests for Client.write_srt_file() edge cases."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def setUp(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
self.client = Client(host="localhost", port=9090, lang="en")
self.client.server_backend = "faster_whisper"
def tearDown(self):
self.client.close_websocket()
import os
for f in ["test_out.srt"]:
if os.path.exists(f):
os.remove(f)
def test_write_srt_empty_transcript_with_last_segment(self):
self.client.transcript = []
self.client.last_segment = {"start": "0.000", "end": "1.000", "text": "final"}
self.client.write_srt_file("test_out.srt")
self.assertEqual(len(self.client.transcript), 1)
self.assertEqual(self.client.transcript[0]["text"], "final")
def test_write_srt_appends_last_segment_if_different(self):
self.client.transcript = [{"start": "0.000", "end": "1.000", "text": "first"}]
self.client.last_segment = {"start": "1.000", "end": "2.000", "text": "second"}
self.client.write_srt_file("test_out.srt")
self.assertEqual(len(self.client.transcript), 2)
def test_write_srt_no_duplicate_last_segment(self):
self.client.transcript = [{"start": "0.000", "end": "1.000", "text": "same"}]
self.client.last_segment = {"start": "0.000", "end": "1.000", "text": "same"}
self.client.write_srt_file("test_out.srt")
self.assertEqual(len(self.client.transcript), 1)
class TestWaitBeforeDisconnect(unittest.TestCase):
"""Tests for Client.wait_before_disconnect()."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def setUp(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
self.client = Client(host="localhost", port=9090, lang="en")
def tearDown(self):
self.client.close_websocket()
def test_raises_if_no_response(self):
self.client.last_response_received = None
with self.assertRaises(AssertionError):
self.client.wait_before_disconnect()
def test_returns_immediately_if_timeout_elapsed(self):
self.client.last_response_received = time.time() - 100
self.client.disconnect_if_no_response_for = 15
start = time.time()
self.client.wait_before_disconnect()
elapsed = time.time() - start
self.assertLess(elapsed, 1.0)
class TestTeeClientEdgeCases(unittest.TestCase):
"""Edge cases for TranscriptionTeeClient."""
def test_empty_clients_raises(self):
with self.assertRaises(Exception):
TranscriptionTeeClient([])
class TestClientReconnect(unittest.TestCase):
"""Tests for reconnection logic."""
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def test_reconnect_on_close(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0)
initial_socket = client.client_socket
client.on_close(MagicMock(), 1006, "abnormal closure")
self.assertEqual(client._retry_count, 1)
# A new websocket should have been created
self.assertIsNotNone(client.client_socket)
client.close_websocket()
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def test_no_reconnect_on_server_error(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0)
client.server_error = True
client.on_close(MagicMock(), 1000, "normal")
self.assertEqual(client._retry_count, 0)
client.close_websocket()
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def test_no_reconnect_when_max_retries_zero(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
client = Client(host="localhost", port=9090, lang="en", max_retries=0, retry_delay=0)
client.on_close(MagicMock(), 1006, "abnormal closure")
self.assertEqual(client._retry_count, 0)
client.close_websocket()
@patch("whisper_live.client.websocket.WebSocketApp")
@patch("whisper_live.client.pyaudio.PyAudio")
def test_stops_after_max_retries(self, mock_pyaudio, mock_websocket):
mock_pyaudio.return_value.open.return_value = MagicMock()
client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0)
client.on_close(MagicMock(), 1006, "closed")
client.on_close(MagicMock(), 1006, "closed")
self.assertEqual(client._retry_count, 2)
# third close should NOT retry
client.on_close(MagicMock(), 1006, "closed")
self.assertEqual(client._retry_count, 2)
client.close_websocket()
if __name__ == "__main__":
unittest.main()
+201
View File
@@ -0,0 +1,201 @@
import unittest
from unittest.mock import MagicMock, patch
import numpy as np
class TestSpeakerDiarizer(unittest.TestCase):
"""Tests for SpeakerDiarizer with mocked embedding model."""
def _make_diarizer(self, **kwargs):
from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer(**kwargs)
# Mock the embedding model to return deterministic embeddings
d._model = MagicMock()
return d
def _set_embedding(self, diarizer, embedding):
"""Configure mock model to return a specific embedding."""
emb = np.array(embedding, dtype=np.float32)
emb = emb / np.linalg.norm(emb)
diarizer._model.return_value = emb
def test_first_speaker_creates_new(self):
d = self._make_diarizer()
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32) # 1 second of audio
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "SPEAKER_00")
self.assertEqual(len(d.speakers), 1)
def test_same_speaker_matches(self):
d = self._make_diarizer(similarity_threshold=0.8)
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
d.identify_speaker(audio) # SPEAKER_00
# Same embedding should match
self._set_embedding(d, [0.99, 0.01, 0.0])
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "SPEAKER_00")
self.assertEqual(len(d.speakers), 1)
def test_different_speaker_creates_new(self):
d = self._make_diarizer(similarity_threshold=0.8)
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
d.identify_speaker(audio) # SPEAKER_00
# Very different embedding
self._set_embedding(d, [0.0, 1.0, 0.0])
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "SPEAKER_01")
self.assertEqual(len(d.speakers), 2)
def test_max_speakers_limit(self):
d = self._make_diarizer(similarity_threshold=0.95, max_speakers=2)
audio = np.zeros(16000, dtype=np.float32)
self._set_embedding(d, [1.0, 0.0, 0.0])
d.identify_speaker(audio) # SPEAKER_00
self._set_embedding(d, [0.0, 1.0, 0.0])
d.identify_speaker(audio) # SPEAKER_01
# Third distinct speaker should be assigned to closest existing
self._set_embedding(d, [0.0, 0.0, 1.0])
speaker = d.identify_speaker(audio)
self.assertIn(speaker, ["SPEAKER_00", "SPEAKER_01"])
self.assertEqual(len(d.speakers), 2)
def test_short_audio_returns_none(self):
d = self._make_diarizer()
# Less than 0.3 seconds
audio = np.zeros(3000, dtype=np.float32)
speaker = d.identify_speaker(audio)
self.assertIsNone(speaker)
def test_reset_clears_state(self):
d = self._make_diarizer()
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
d.identify_speaker(audio)
self.assertEqual(len(d.speakers), 1)
d.reset()
self.assertEqual(len(d.speakers), 0)
self.assertEqual(d._speaker_count, 0)
def test_enroll_speaker_uses_known_name(self):
d = self._make_diarizer(similarity_threshold=0.8)
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
self.assertTrue(d.enroll_speaker("Alice", audio))
self._set_embedding(d, [0.99, 0.01, 0.0])
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_speaker_names_label_new_speakers(self):
d = self._make_diarizer(speaker_names=["Alice"])
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_import_error_without_pyannote(self):
from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer()
with patch.dict("sys.modules", {"pyannote": None, "pyannote.audio": None}):
with self.assertRaises(ImportError):
d._load_model()
class TestDiarizationInBase(unittest.TestCase):
"""Test diarization integration in ServeClientBase."""
def _make_client(self, diarization=None):
from whisper_live.backend.base import ServeClientBase
class ConcreteClient(ServeClientBase):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.language = "en"
def transcribe_audio(self, input_sample):
return None
def handle_transcription_output(self, result, duration):
pass
ws = MagicMock()
return ConcreteClient(
client_uid="test-uid", websocket=ws, diarization=diarization
)
def test_no_diarization_by_default(self):
client = self._make_client()
self.assertIsNone(client.diarization)
def test_format_segment_with_speaker(self):
client = self._make_client()
seg = client.format_segment(0.0, 1.0, "hello", speaker="SPEAKER_00")
self.assertEqual(seg["speaker"], "SPEAKER_00")
def test_format_segment_without_speaker(self):
client = self._make_client()
seg = client.format_segment(0.0, 1.0, "hello")
self.assertNotIn("speaker", seg)
def test_identify_speaker_disabled(self):
client = self._make_client(diarization=None)
seg = MagicMock()
seg.start = 0.0
seg.end = 1.0
result = client._identify_speaker(seg)
self.assertIsNone(result)
def test_identify_speaker_calls_diarizer(self):
mock_diarizer = MagicMock()
mock_diarizer.identify_speaker.return_value = "SPEAKER_01"
client = self._make_client(diarization=mock_diarizer)
# Set up audio buffer
client.frames_np = np.zeros(48000, dtype=np.float32)
client.frames_offset = 0.0
client.timestamp_offset = 0.0
seg = MagicMock()
seg.start = 0.5
seg.end = 1.5
result = client._identify_speaker(seg)
self.assertEqual(result, "SPEAKER_01")
mock_diarizer.identify_speaker.assert_called_once()
class TestRestDiarizationHelpers(unittest.TestCase):
def test_normalize_form_list_accepts_repeated_or_comma_separated_values(self):
from whisper_live.server import TranscriptionServer
self.assertEqual(
TranscriptionServer._normalize_form_list(["Alice,Bob", "Carol"]),
["Alice", "Bob", "Carol"],
)
def test_speaker_labels_for_segments(self):
from whisper_live.server import TranscriptionServer
segment = MagicMock()
segment.start = 0.0
segment.end = 1.0
diarizer = MagicMock()
diarizer.identify_speaker.return_value = "Alice"
speakers = TranscriptionServer._speaker_labels_for_segments(
[segment],
np.zeros(16000, dtype=np.float32),
diarizer,
)
self.assertEqual(speakers, {0: "Alice"})
diarizer.identify_speaker.assert_called_once()
if __name__ == "__main__":
unittest.main()
+138
View File
@@ -0,0 +1,138 @@
import unittest
from unittest.mock import patch, MagicMock
from whisper_live import metrics as wl_metrics
_skip_no_prometheus = unittest.skipUnless(
wl_metrics.is_available(), "prometheus_client not installed"
)
class TestMetricsAvailability(unittest.TestCase):
def test_is_available_returns_bool(self):
self.assertIsInstance(wl_metrics.is_available(), bool)
@_skip_no_prometheus
class TestTrackConnectionOpened(unittest.TestCase):
def test_increments_total_and_active(self):
total_before = wl_metrics.CONNECTIONS_TOTAL._value.get()
active_before = wl_metrics.CONNECTIONS_ACTIVE._value.get()
wl_metrics.track_connection_opened()
self.assertEqual(wl_metrics.CONNECTIONS_TOTAL._value.get(), total_before + 1)
self.assertEqual(wl_metrics.CONNECTIONS_ACTIVE._value.get(), active_before + 1)
@_skip_no_prometheus
class TestTrackConnectionClosed(unittest.TestCase):
def test_decrements_active(self):
wl_metrics.track_connection_opened()
active_before = wl_metrics.CONNECTIONS_ACTIVE._value.get()
wl_metrics.track_connection_closed()
self.assertEqual(wl_metrics.CONNECTIONS_ACTIVE._value.get(), active_before - 1)
@_skip_no_prometheus
class TestTrackConnectionRejected(unittest.TestCase):
def test_rejected_full(self):
before = wl_metrics.CONNECTIONS_REJECTED.labels(reason="full")._value.get()
wl_metrics.track_connection_rejected(reason="full")
self.assertEqual(wl_metrics.CONNECTIONS_REJECTED.labels(reason="full")._value.get(), before + 1)
def test_rejected_auth(self):
before = wl_metrics.CONNECTIONS_REJECTED.labels(reason="auth")._value.get()
wl_metrics.track_connection_rejected(reason="auth")
self.assertEqual(wl_metrics.CONNECTIONS_REJECTED.labels(reason="auth")._value.get(), before + 1)
@_skip_no_prometheus
class TestTrackTranscriptionLatency(unittest.TestCase):
def test_observe_records_value(self):
count_before = wl_metrics.TRANSCRIPTION_LATENCY._sum.get()
wl_metrics.track_transcription_latency(0.5)
self.assertAlmostEqual(wl_metrics.TRANSCRIPTION_LATENCY._sum.get(), count_before + 0.5, places=3)
@_skip_no_prometheus
class TestTrackAudioProcessed(unittest.TestCase):
def test_increments_by_duration(self):
before = wl_metrics.AUDIO_PROCESSED._value.get()
wl_metrics.track_audio_processed(3.5)
self.assertAlmostEqual(wl_metrics.AUDIO_PROCESSED._value.get(), before + 3.5, places=3)
@_skip_no_prometheus
class TestTrackSegmentEmitted(unittest.TestCase):
def test_completed_true(self):
before = wl_metrics.SEGMENTS_EMITTED.labels(completed="true")._value.get()
wl_metrics.track_segment_emitted(completed=True)
self.assertEqual(wl_metrics.SEGMENTS_EMITTED.labels(completed="true")._value.get(), before + 1)
def test_completed_false(self):
before = wl_metrics.SEGMENTS_EMITTED.labels(completed="false")._value.get()
wl_metrics.track_segment_emitted(completed=False)
self.assertEqual(wl_metrics.SEGMENTS_EMITTED.labels(completed="false")._value.get(), before + 1)
@_skip_no_prometheus
class TestTrackRestRequest(unittest.TestCase):
def test_tracks_200(self):
before = wl_metrics.REST_REQUESTS.labels(endpoint="transcriptions", status="200")._value.get()
wl_metrics.track_rest_request(endpoint="transcriptions", status=200)
self.assertEqual(wl_metrics.REST_REQUESTS.labels(endpoint="transcriptions", status="200")._value.get(), before + 1)
def test_tracks_500(self):
before = wl_metrics.REST_REQUESTS.labels(endpoint="transcriptions", status="500")._value.get()
wl_metrics.track_rest_request(endpoint="transcriptions", status=500)
self.assertEqual(wl_metrics.REST_REQUESTS.labels(endpoint="transcriptions", status="500")._value.get(), before + 1)
@_skip_no_prometheus
class TestTrackError(unittest.TestCase):
def test_tracks_transcription_error(self):
before = wl_metrics.ERRORS.labels(type="transcription")._value.get()
wl_metrics.track_error("transcription")
self.assertEqual(wl_metrics.ERRORS.labels(type="transcription")._value.get(), before + 1)
def test_tracks_rest_error(self):
before = wl_metrics.ERRORS.labels(type="rest_transcription")._value.get()
wl_metrics.track_error("rest_transcription")
self.assertEqual(wl_metrics.ERRORS.labels(type="rest_transcription")._value.get(), before + 1)
@_skip_no_prometheus
class TestStartMetricsServer(unittest.TestCase):
@patch("whisper_live.metrics.start_http_server")
def test_starts_on_given_port(self, mock_start):
wl_metrics.start_metrics_server(9999)
mock_start.assert_called_once_with(9999)
@patch("whisper_live.metrics.start_http_server", side_effect=OSError("port in use"))
def test_logs_error_on_failure(self, mock_start):
with self.assertLogs(level="ERROR") as cm:
wl_metrics.start_metrics_server(9999)
self.assertTrue(any("Failed to start" in msg for msg in cm.output))
class TestNoOpWhenUnavailable(unittest.TestCase):
"""Verify helper functions are no-ops when _AVAILABLE is False."""
def test_all_helpers_are_noop(self):
original = wl_metrics._AVAILABLE
try:
wl_metrics._AVAILABLE = False
# None of these should raise
wl_metrics.track_connection_opened()
wl_metrics.track_connection_closed()
wl_metrics.track_connection_rejected("full")
wl_metrics.track_transcription_latency(1.0)
wl_metrics.track_audio_processed(1.0)
wl_metrics.track_segment_emitted()
wl_metrics.track_rest_request()
wl_metrics.track_error()
finally:
wl_metrics._AVAILABLE = original
if __name__ == "__main__":
unittest.main()
+150
View File
@@ -0,0 +1,150 @@
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()
self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600)
self.server.cache_path = "~/.cache/whisper-live/"
@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))
+700
View File
@@ -0,0 +1,700 @@
import json
import time
import threading
import collections
import unittest
from unittest import mock
from unittest.mock import MagicMock, patch
from whisper_live.server import TranscriptionServer, BackendType, ClientManager
class TestClientManagerAddRemove(unittest.TestCase):
def setUp(self):
self.cm = ClientManager(max_clients=2, max_connection_time=60)
def test_add_and_get_client(self):
ws = MagicMock()
client = MagicMock()
self.cm.add_client(ws, client)
self.assertIs(self.cm.get_client(ws), client)
def test_get_nonexistent_client(self):
ws = MagicMock()
self.assertFalse(self.cm.get_client(ws))
def test_remove_client_calls_cleanup(self):
ws = MagicMock()
client = MagicMock()
self.cm.add_client(ws, client)
self.cm.remove_client(ws)
client.cleanup.assert_called_once()
self.assertNotIn(ws, self.cm.clients)
self.assertNotIn(ws, self.cm.start_times)
def test_remove_nonexistent_client_no_error(self):
ws = MagicMock()
self.cm.remove_client(ws) # should not raise
class TestClientManagerThreadSafety(unittest.TestCase):
def test_concurrent_add_remove(self):
cm = ClientManager(max_clients=100, max_connection_time=600)
errors = []
def add_clients(start_idx):
try:
for i in range(50):
ws = MagicMock(name=f"ws-{start_idx}-{i}")
client = MagicMock(name=f"client-{start_idx}-{i}")
cm.add_client(ws, client)
except Exception as e:
errors.append(e)
def remove_clients():
try:
for _ in range(25):
with cm.lock:
if cm.clients:
ws = next(iter(cm.clients))
else:
continue
cm.remove_client(ws)
except Exception as e:
errors.append(e)
threads = [
threading.Thread(target=add_clients, args=(0,)),
threading.Thread(target=add_clients, args=(1,)),
threading.Thread(target=remove_clients),
threading.Thread(target=remove_clients),
]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
def test_concurrent_get_client(self):
cm = ClientManager(max_clients=100, max_connection_time=600)
ws = MagicMock()
client = MagicMock()
cm.add_client(ws, client)
errors = []
results = []
def get_many():
try:
for _ in range(100):
results.append(cm.get_client(ws))
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=get_many) for _ in range(4)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
self.assertTrue(all(r is client for r in results))
class TestClientManagerServerFull(unittest.TestCase):
def setUp(self):
self.cm = ClientManager(max_clients=1, max_connection_time=60)
def test_not_full_returns_false(self):
ws = MagicMock()
options = {"uid": "test"}
self.assertFalse(self.cm.is_server_full(ws, options))
def test_full_sends_wait_and_returns_true(self):
ws1 = MagicMock()
self.cm.add_client(ws1, MagicMock())
ws2 = MagicMock()
options = {"uid": "new-client"}
self.assertTrue(self.cm.is_server_full(ws2, options))
ws2.send.assert_called_once()
sent = json.loads(ws2.send.call_args[0][0])
self.assertEqual(sent["status"], "WAIT")
self.assertEqual(sent["uid"], "new-client")
class TestClientManagerTimeout(unittest.TestCase):
def setUp(self):
self.cm = ClientManager(max_clients=4, max_connection_time=10)
def test_not_timed_out(self):
ws = MagicMock()
client = MagicMock()
self.cm.add_client(ws, client)
self.assertFalse(self.cm.is_client_timeout(ws))
def test_timed_out(self):
ws = MagicMock()
client = MagicMock()
self.cm.add_client(ws, client)
self.cm.start_times[ws] = time.time() - 20
self.assertTrue(self.cm.is_client_timeout(ws))
client.disconnect.assert_called_once()
class TestClientManagerGetWaitTime(unittest.TestCase):
def test_no_clients_returns_zero(self):
cm = ClientManager(max_clients=4, max_connection_time=600)
self.assertEqual(cm.get_wait_time(), 0)
def test_single_client_wait_time(self):
cm = ClientManager(max_clients=4, max_connection_time=600)
ws = MagicMock()
cm.add_client(ws, MagicMock())
cm.start_times[ws] = time.time() - 300
wait = cm.get_wait_time()
self.assertAlmostEqual(wait, 5.0, places=0)
def test_multiple_clients_returns_minimum(self):
cm = ClientManager(max_clients=4, max_connection_time=600)
ws1, ws2 = MagicMock(), MagicMock()
cm.add_client(ws1, MagicMock())
cm.add_client(ws2, MagicMock())
cm.start_times[ws1] = time.time() - 100
cm.start_times[ws2] = time.time() - 500
wait = cm.get_wait_time()
# ws2 has 100s remaining = ~1.67 minutes
self.assertAlmostEqual(wait, 100 / 60, places=0)
class TestBackendType(unittest.TestCase):
def test_valid_types(self):
valid = BackendType.valid_types()
self.assertIn("faster_whisper", valid)
self.assertIn("tensorrt", valid)
self.assertIn("openvino", valid)
def test_is_valid(self):
self.assertTrue(BackendType.is_valid("faster_whisper"))
self.assertFalse(BackendType.is_valid("nonexistent"))
def test_type_checks(self):
self.assertTrue(BackendType.FASTER_WHISPER.is_faster_whisper())
self.assertFalse(BackendType.FASTER_WHISPER.is_tensorrt())
self.assertTrue(BackendType.TENSORRT.is_tensorrt())
self.assertTrue(BackendType.OPENVINO.is_openvino())
def test_enum_from_string(self):
bt = BackendType("faster_whisper")
self.assertEqual(bt, BackendType.FASTER_WHISPER)
def test_invalid_enum_raises(self):
with self.assertRaises(ValueError):
BackendType("invalid_backend")
class TestTranscriptionServerInit(unittest.TestCase):
def test_defaults(self):
server = TranscriptionServer()
self.assertIsNone(server.client_manager)
self.assertTrue(server.use_vad)
self.assertFalse(server.single_model)
self.assertIsNone(server.batch_config)
def test_run_invalid_backend_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, backend="nonexistent")
def test_run_invalid_trt_path_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(
host="localhost",
port=9090,
backend="tensorrt",
whisper_tensorrt_path="/nonexistent/path",
)
def test_run_max_clients_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_clients=0)
def test_run_max_clients_negative_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_clients=-1)
def test_run_max_connection_time_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_connection_time=0)
def test_run_batch_max_size_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, batch_enabled=True, batch_max_size=0)
def test_run_batch_window_ms_negative_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, batch_enabled=True, batch_window_ms=-1)
class TestTranscriptionServerGetAudio(unittest.TestCase):
def setUp(self):
self.server = TranscriptionServer()
def test_end_of_audio_returns_false(self):
ws = MagicMock()
ws.recv.return_value = b"END_OF_AUDIO"
result = self.server.get_audio_from_websocket(ws)
self.assertFalse(result)
def test_valid_audio_returns_numpy(self):
import numpy as np
ws = MagicMock()
audio = np.array([0.1, 0.2, 0.3], dtype=np.float32)
ws.recv.return_value = audio.tobytes()
result = self.server.get_audio_from_websocket(ws)
np.testing.assert_array_almost_equal(result, audio)
def test_raw_pcm_input_normalizes_int16(self):
import numpy as np
self.server.raw_pcm_input = True
ws = MagicMock()
pcm = np.array([0, 16384, -16384, 32767], dtype=np.int16)
ws.recv.return_value = pcm.tobytes()
result = self.server.get_audio_from_websocket(ws)
expected = pcm.astype(np.float32) / 32768.0
np.testing.assert_array_almost_equal(result, expected)
self.assertTrue(result.dtype == np.float32)
self.assertTrue(np.all(result >= -1.0))
self.assertTrue(np.all(result <= 1.0))
def test_uint8_audio_format_normalizes_unsigned_pcm(self):
import numpy as np
ws = MagicMock()
self.server.audio_formats[ws] = "uint8"
pcm = np.array([0, 128, 255], dtype=np.uint8)
ws.recv.return_value = pcm.tobytes()
result = self.server.get_audio_from_websocket(ws)
expected = (pcm.astype(np.float32) - 128.0) / 128.0
np.testing.assert_array_almost_equal(result, expected)
def test_raw_pcm_input_off_reads_float32(self):
import numpy as np
self.server.raw_pcm_input = False
ws = MagicMock()
audio = np.array([0.5, -0.5], dtype=np.float32)
ws.recv.return_value = audio.tobytes()
result = self.server.get_audio_from_websocket(ws)
np.testing.assert_array_almost_equal(result, audio)
class TestTranscriptionServerHandleNewConnection(unittest.TestCase):
def setUp(self):
self.server = TranscriptionServer()
self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600)
self.server.cache_path = "~/.cache/whisper-live/"
self.server.backend = BackendType.FASTER_WHISPER
@mock.patch("websockets.WebSocketCommonProtocol")
def test_invalid_json_returns_false(self, mock_ws):
mock_ws.recv.return_value = "not valid json {{"
result = self.server.handle_new_connection(mock_ws, None, None, False)
self.assertFalse(result)
@mock.patch("websockets.WebSocketCommonProtocol")
def test_server_full_returns_false(self, mock_ws):
# Fill server
for i in range(4):
self.server.client_manager.add_client(MagicMock(), MagicMock())
mock_ws.recv.return_value = json.dumps({
"uid": "test",
"language": "en",
"task": "transcribe",
"model": "tiny.en",
})
result = self.server.handle_new_connection(mock_ws, None, None, False)
self.assertFalse(result)
class TestTranscriptionServerCleanup(unittest.TestCase):
def setUp(self):
self.server = TranscriptionServer()
self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600)
def test_cleanup_removes_client(self):
ws = MagicMock()
client = MagicMock()
self.server.client_manager.add_client(ws, client)
self.cleanup_server = self.server
self.server.cleanup(ws)
self.assertNotIn(ws, self.server.client_manager.clients)
client.cleanup.assert_called_once()
class TestStreamTranscription(unittest.TestCase):
"""Tests for the SSE streaming endpoint (stream=true)."""
def _make_app(self):
"""Create a FastAPI app with the transcribe endpoint that has streaming support."""
from fastapi import FastAPI, UploadFile, Form
app = FastAPI()
server = TranscriptionServer()
@app.post("/v1/audio/transcriptions")
async def transcribe(
file: UploadFile,
stream: bool = Form(default=False),
language: str = Form(default=None),
response_format: str = Form(default="json"),
):
if stream:
return server._stream_transcription(
file, language, None, 0.0, None, None
)
return {"text": "non-streamed"}
return app
@patch("whisper_live.server.WhisperModel")
def test_stream_returns_sse_content_type(self, mock_model_cls):
mock_seg = MagicMock()
mock_seg.id = 0
mock_seg.start = 0.0
mock_seg.end = 1.0
mock_seg.text = " hello "
mock_seg.words = []
mock_info = MagicMock()
mock_info.language = "en"
mock_info.language_probability = 0.98
mock_info.duration = 1.0
mock_model = MagicMock()
mock_model.transcribe.return_value = (iter([mock_seg]), mock_info)
mock_model_cls.return_value = mock_model
import io
from fastapi.testclient import TestClient
app = self._make_app()
client = TestClient(app)
resp = client.post(
"/v1/audio/transcriptions",
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
data={"stream": "true"},
)
self.assertEqual(resp.status_code, 200)
self.assertIn("text/event-stream", resp.headers.get("content-type", ""))
@patch("whisper_live.server.WhisperModel")
def test_stream_yields_segment_and_done(self, mock_model_cls):
mock_seg = MagicMock()
mock_seg.id = 0
mock_seg.start = 0.0
mock_seg.end = 1.5
mock_seg.text = " hello world "
mock_seg.words = []
mock_info = MagicMock()
mock_info.language = "en"
mock_info.language_probability = 0.95
mock_info.duration = 1.5
mock_model = MagicMock()
mock_model.transcribe.return_value = (iter([mock_seg]), mock_info)
mock_model_cls.return_value = mock_model
import io
from fastapi.testclient import TestClient
app = self._make_app()
client = TestClient(app)
resp = client.post(
"/v1/audio/transcriptions",
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
data={"stream": "true"},
)
body = resp.text
self.assertIn('"text": "hello world"', body)
self.assertIn("[DONE]", body)
@patch("whisper_live.server.WhisperModel")
def test_stream_multiple_segments(self, mock_model_cls):
segs = []
for i in range(3):
s = MagicMock()
s.id = i
s.start = float(i)
s.end = float(i + 1)
s.text = f" segment {i} "
s.words = []
segs.append(s)
mock_info = MagicMock()
mock_info.language = "en"
mock_info.language_probability = 0.99
mock_info.duration = 3.0
mock_model = MagicMock()
mock_model.transcribe.return_value = (iter(segs), mock_info)
mock_model_cls.return_value = mock_model
import io
from fastapi.testclient import TestClient
app = self._make_app()
client = TestClient(app)
resp = client.post(
"/v1/audio/transcriptions",
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
data={"stream": "true"},
)
body = resp.text
events = [line for line in body.split("\n") if line.startswith("data: ") and "[DONE]" not in line and '"type": "metadata"' not in line]
self.assertEqual(len(events), 3)
for i, event in enumerate(events):
data = json.loads(event.removeprefix("data: "))
self.assertEqual(data["text"], f"segment {i}")
@patch("whisper_live.server.WhisperModel", side_effect=RuntimeError("model error"))
def test_stream_error_yields_error_event(self, mock_model_cls):
import io
from fastapi.testclient import TestClient
app = self._make_app()
client = TestClient(app)
resp = client.post(
"/v1/audio/transcriptions",
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
data={"stream": "true"},
)
body = resp.text
self.assertIn('"error"', body)
self.assertIn("model error", body)
def test_non_stream_still_works(self):
import io
from fastapi.testclient import TestClient
app = self._make_app()
client = TestClient(app)
resp = client.post(
"/v1/audio/transcriptions",
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
data={"stream": "false"},
)
self.assertEqual(resp.status_code, 200)
self.assertEqual(resp.json()["text"], "non-streamed")
class TestRESTAPIParamWarnings(unittest.TestCase):
"""Test that unsupported OpenAI-compatible REST params produce warnings."""
@classmethod
def setUpClass(cls):
"""Build a FastAPI test app by extracting the endpoint definition."""
import logging
from fastapi import FastAPI, UploadFile, Form, File
from fastapi.testclient import TestClient
from typing import Optional, List
app = FastAPI()
@app.post("/v1/audio/transcriptions")
async def transcribe(
file: UploadFile,
model: str = Form(default="whisper-1"),
language: Optional[str] = Form(default=None),
prompt: Optional[str] = Form(default=None),
response_format: str = Form(default="json"),
temperature: float = Form(default=0.0),
timestamp_granularities: Optional[List[str]] = Form(default=None),
chunking_strategy: Optional[str] = Form(default=None),
include: Optional[List[str]] = Form(default=None),
known_speaker_names: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[UploadFile]] = File(default=None),
stream: bool = Form(default=False),
):
ignored_params = []
if chunking_strategy:
ignored_params.append(f"chunking_strategy='{chunking_strategy}'")
if include:
ignored_params.append(f"include={include}")
if ignored_params:
logging.warning(f"Unsupported OpenAI params ignored: {', '.join(ignored_params)}")
# Return a JSON response with the ignored list for testing
return {"text": "test", "ignored": ignored_params}
cls.test_client = TestClient(app)
def _post(self, **extra_fields):
import io
data = {**extra_fields}
files = {"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}
return self.test_client.post("/v1/audio/transcriptions", data=data, files=files)
def test_no_warnings_when_no_extra_params(self):
resp = self._post()
self.assertEqual(resp.status_code, 200)
self.assertEqual(resp.json()["ignored"], [])
def test_chunking_strategy_warning(self):
resp = self._post(chunking_strategy="auto")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertTrue(any("chunking_strategy" in p for p in ignored))
def test_include_warning(self):
resp = self._post(include="logprobs")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertTrue(any("include" in p for p in ignored))
def test_known_speaker_names_supported(self):
resp = self._post(known_speaker_names="alice")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertFalse(any("known_speaker_names" in p for p in ignored))
def test_multiple_ignored_params(self):
resp = self._post(chunking_strategy="auto", known_speaker_names="bob")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertEqual(len(ignored), 1)
class TestAPIKeyAuth(unittest.TestCase):
"""Test optional API key authentication middleware."""
@classmethod
def setUpClass(cls):
from fastapi import FastAPI, Request
from fastapi.testclient import TestClient
from fastapi.responses import JSONResponse as JSONR
app = FastAPI()
@app.middleware("http")
async def _check_api_key(request: Request, call_next):
auth = request.headers.get("Authorization", "")
if auth != "Bearer test-secret":
return JSONR({"error": "Invalid or missing API key"}, status_code=401)
return await call_next(request)
@app.get("/ping")
async def ping():
return {"status": "ok"}
cls.test_client = TestClient(app)
def test_missing_key_returns_401(self):
resp = self.test_client.get("/ping")
self.assertEqual(resp.status_code, 401)
def test_wrong_key_returns_401(self):
resp = self.test_client.get("/ping", headers={"Authorization": "Bearer wrong"})
self.assertEqual(resp.status_code, 401)
def test_correct_key_returns_200(self):
resp = self.test_client.get("/ping", headers={"Authorization": "Bearer test-secret"})
self.assertEqual(resp.status_code, 200)
self.assertEqual(resp.json()["status"], "ok")
class TestRateLimiting(unittest.TestCase):
"""Test per-IP rate limiting middleware."""
def _make_app(self, rpm_limit=3):
from fastapi import FastAPI, Request
from fastapi.testclient import TestClient
from fastapi.responses import JSONResponse as JSONR
_rate_lock = threading.Lock()
_rate_buckets: dict = {}
app = FastAPI()
@app.middleware("http")
async def _rate_limit(request: Request, call_next):
client_ip = request.client.host if request.client else "unknown"
now = time.time()
with _rate_lock:
bucket = _rate_buckets.setdefault(client_ip, collections.deque())
while bucket and bucket[0] < now - 60:
bucket.popleft()
if len(bucket) >= rpm_limit:
return JSONR({"error": "Rate limit exceeded"}, status_code=429)
bucket.append(now)
return await call_next(request)
@app.get("/ping")
async def ping():
return {"status": "ok"}
return TestClient(app)
def test_within_limit_succeeds(self):
client = self._make_app(rpm_limit=3)
for _ in range(3):
resp = client.get("/ping")
self.assertEqual(resp.status_code, 200)
def test_exceeding_limit_returns_429(self):
client = self._make_app(rpm_limit=3)
for _ in range(3):
client.get("/ping")
resp = client.get("/ping")
self.assertEqual(resp.status_code, 429)
self.assertIn("Rate limit", resp.json()["error"])
class TestWebSocketAuth(unittest.TestCase):
"""Tests for the WebSocket process_request auth callback."""
def _make_auth_handler(self, api_key):
"""Build the same auth function the server creates."""
def _ws_auth(path, request_headers):
auth = request_headers.get("Authorization", "")
token_param = None
if "?" in path:
from urllib.parse import urlparse, parse_qs
parsed = urlparse(path)
token_param = parse_qs(parsed.query).get("token", [None])[0]
if auth == f"Bearer {api_key}" or token_param == api_key:
return None
return (401, [("Content-Type", "text/plain")], b"Unauthorized\n")
return _ws_auth
def test_valid_bearer_token(self):
handler = self._make_auth_handler("my-secret")
result = handler("/", {"Authorization": "Bearer my-secret"})
self.assertIsNone(result)
def test_invalid_bearer_token(self):
handler = self._make_auth_handler("my-secret")
result = handler("/", {"Authorization": "Bearer wrong"})
self.assertEqual(result[0], 401)
def test_missing_auth_header(self):
handler = self._make_auth_handler("my-secret")
result = handler("/", {})
self.assertEqual(result[0], 401)
def test_valid_query_token(self):
handler = self._make_auth_handler("my-secret")
result = handler("/?token=my-secret", {})
self.assertIsNone(result)
def test_invalid_query_token(self):
handler = self._make_auth_handler("my-secret")
result = handler("/?token=wrong", {})
self.assertEqual(result[0], 401)
if __name__ == "__main__":
unittest.main()
+210
View File
@@ -0,0 +1,210 @@
import json
import time
import unittest
from unittest.mock import patch, MagicMock
import numpy as np
from whisper_live.client import Client, StreamingTranscriptionClient
class StreamingClientTestCase(unittest.TestCase):
@patch('whisper_live.client.websocket.WebSocketApp')
def setUp(self, mock_websocket):
self.mock_websocket = mock_websocket
self.mock_ws_app = mock_websocket.return_value
self.mock_ws_app.send = MagicMock()
self.committed = []
self.partials = []
self.session_started = []
self.client = StreamingTranscriptionClient(
host='localhost',
port=9090,
lang="en",
on_session_started=lambda: self.session_started.append(True),
on_committed_transcript=lambda text, segs: self.committed.append((text, segs)),
on_partial_transcript=lambda text, segs: self.partials.append((text, segs)),
)
self._inner = self.client._client
def tearDown(self):
self._inner.close_websocket()
self.mock_websocket.stop()
def _server_ready(self, backend="faster_whisper"):
self._inner.on_message(self.mock_ws_app, json.dumps({
"uid": self._inner.uid,
"message": "SERVER_READY",
"backend": backend,
}))
def _send_segments(self, segments):
self._inner.on_message(self.mock_ws_app, json.dumps({
"uid": self._inner.uid,
"segments": segments,
}))
class TestPcmFormatConversion(StreamingClientTestCase):
def test_int16_is_normalized_to_float32(self):
self._server_ready()
raw = np.array([0, 16384, -32768], dtype=np.int16).tobytes()
with patch.object(self._inner, 'send_packet_to_server') as mock_send:
self.client.send(raw, pcm_format="int16")
sent = np.frombuffer(mock_send.call_args[0][0], dtype=np.float32)
np.testing.assert_allclose(sent, [0.0, 0.5, -1.0], atol=1e-4)
def test_float32_passes_through(self):
self._server_ready()
raw = np.array([0.1, -0.2], dtype=np.float32).tobytes()
with patch.object(self._inner, 'send_packet_to_server') as mock_send:
self.client.send(raw, pcm_format="float32")
self.assertEqual(mock_send.call_args[0][0], raw)
def test_default_format_is_int16(self):
self._server_ready()
raw = np.array([32767], dtype=np.int16).tobytes()
with patch.object(self._inner, 'send_packet_to_server') as mock_send:
self.client.send(raw)
sent = np.frombuffer(mock_send.call_args[0][0], dtype=np.float32)
self.assertAlmostEqual(float(sent[0]), 32767 / 32768.0, places=4)
def test_unsupported_format_raises(self):
self._server_ready()
with self.assertRaises(ValueError):
self.client.send(b"\x00\x00", pcm_format="int8")
def test_send_after_close_raises(self):
self.client._closed = True
with self.assertRaises(RuntimeError):
self.client.send(b"\x00\x00", pcm_format="int16")
def test_send_array_normalizes_integers(self):
with patch.object(self._inner, 'send_packet_to_server') as mock_send:
self.client.send_array(np.array([0, 16384, -32768], dtype=np.int16))
sent = np.frombuffer(mock_send.call_args[0][0], dtype=np.float32)
np.testing.assert_allclose(sent, [0.0, 0.5, -1.0], atol=1e-4)
class TestTranscriptDispatch(StreamingClientTestCase):
def test_partial_then_committed(self):
self._server_ready()
self._send_segments([{"start": 0, "end": 1, "text": "hello", "completed": False}])
self.assertEqual(len(self.partials), 1)
self.assertEqual(self.partials[0][0], "hello")
self.assertEqual(len(self.committed), 0)
self._send_segments([{"start": 0, "end": 1, "text": "hello world", "completed": True}])
self.assertEqual(len(self.committed), 1)
self.assertEqual(self.committed[0][0], "hello world")
self.assertEqual(len(self.client.transcript), 1)
def test_committed_deduplicated(self):
self._server_ready()
seg = {"start": 0, "end": 1, "text": "hi", "completed": True}
self._send_segments([seg])
self._send_segments([seg])
self.assertEqual(len(self.committed), 1)
self.assertEqual(len(self.client.transcript), 1)
def test_committed_backend_agnostic(self):
"""Committed dispatch must work for non-faster_whisper backends."""
self._server_ready(backend="tensorrt")
self._send_segments([{"start": 0, "end": 1, "text": "trt seg", "completed": True}])
self.assertEqual(len(self.committed), 1)
self.assertEqual(len(self.client.transcript), 1)
def test_last_partial_alias(self):
self._server_ready()
self._send_segments([{"start": 0, "end": 1, "text": "pending", "completed": False}])
self.assertIsNotNone(self.client.last_partial)
self.assertIs(self.client.last_partial, self.client.last_segment)
class TestConnectLifecycle(StreamingClientTestCase):
def test_connect_returns_after_ready(self):
self._server_ready()
self.assertIs(self.client.connect(), self.client)
self.assertEqual(len(self.session_started), 1)
def test_connect_times_out(self):
self.client._ready_timeout = 0.1
with self.assertRaises(TimeoutError):
self.client.connect()
def test_connect_raises_on_server_error(self):
self._inner.on_message(self.mock_ws_app, json.dumps({
"uid": self._inner.uid,
"status": "ERROR",
"message": "boom",
}))
with self.assertRaises(RuntimeError):
self.client.connect()
def test_connect_raises_when_server_full(self):
self._inner.on_message(self.mock_ws_app, json.dumps({
"uid": self._inner.uid,
"status": "WAIT",
"message": 5,
}))
with self.assertRaises(RuntimeError):
self.client.connect()
def test_close_sends_end_of_audio(self):
self._server_ready()
self._inner.recording = False # pretend server already closed
with patch.object(self._inner, 'send_packet_to_server') as mock_send, \
patch.object(self._inner, 'close_websocket') as mock_close:
self.client.close()
mock_send.assert_called_once_with(Client.END_OF_AUDIO.encode("utf-8"))
mock_close.assert_called_once()
def test_close_waits_for_server_then_times_out(self):
self._server_ready()
self.assertTrue(self._inner.recording) # server still "processing"
start = time.time()
with patch.object(self._inner, 'send_packet_to_server'), \
patch.object(self._inner, 'close_websocket') as mock_close:
self.client.close(timeout=0.2)
self.assertGreaterEqual(time.time() - start, 0.2)
mock_close.assert_called_once()
def test_close_returns_early_when_server_closes(self):
self._server_ready()
def close_soon(_msg):
self._inner.recording = False
with patch.object(self._inner, 'send_packet_to_server', side_effect=close_soon), \
patch.object(self._inner, 'close_websocket') as mock_close:
start = time.time()
self.client.close(timeout=10.0)
self.assertLess(time.time() - start, 1.0)
mock_close.assert_called_once()
class TestErrorHandling(StreamingClientTestCase):
def test_close_frame_not_reported_as_error(self):
"""A normal CLOSE control frame (opcode 8) must not fire on_error."""
self._server_ready()
errors = []
self.client._client._on_error_hook = errors.append
close_frame = MagicMock()
close_frame.opcode = 8
self._inner.on_error(self.mock_ws_app, close_frame)
self.assertEqual(errors, [])
self.assertFalse(self._inner.server_error)
def test_real_error_still_reported(self):
self._server_ready()
errors = []
self.client._client._on_error_hook = errors.append
self._inner.on_error(self.mock_ws_app, RuntimeError("boom"))
self.assertEqual(len(errors), 1)
self.assertTrue(self._inner.server_error)
if __name__ == '__main__':
unittest.main()
+140
View File
@@ -0,0 +1,140 @@
import os
import tempfile
import unittest
from io import StringIO
from unittest.mock import patch
from whisper_live.utils import format_time, create_srt_file, print_transcript, clear_screen
class TestFormatTime(unittest.TestCase):
def test_zero(self):
self.assertEqual(format_time(0), "00:00:00,000")
def test_seconds_only(self):
self.assertEqual(format_time(5.0), "00:00:05,000")
def test_fractional_seconds(self):
self.assertEqual(format_time(1.5), "00:00:01,500")
def test_minutes(self):
self.assertEqual(format_time(65.0), "00:01:05,000")
def test_hours(self):
self.assertEqual(format_time(3661.123), "01:01:01,123")
def test_millisecond_precision(self):
self.assertEqual(format_time(0.001), "00:00:00,001")
def test_large_value(self):
# float precision: int((86399.999 - 86399) * 1000) may be 998 or 999
result = format_time(86399.999)
self.assertIn(result, ("23:59:59,998", "23:59:59,999"))
def test_rounding_edge(self):
result = format_time(0.9999)
# 0.9999 -> int(s%60)=0, milliseconds=int(0.9999*1000)=999
self.assertEqual(result, "00:00:00,999")
class TestCreateSrtFile(unittest.TestCase):
def test_single_segment(self):
segments = [{"start": "0.000", "end": "1.500", "text": "Hello world"}]
with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f:
path = f.name
try:
create_srt_file(segments, path)
with open(path, "r", encoding="utf-8") as f:
content = f.read()
self.assertIn("1\n", content)
self.assertIn("00:00:00,000 --> 00:00:01,500", content)
self.assertIn("Hello world", content)
finally:
os.remove(path)
def test_multiple_segments(self):
segments = [
{"start": "0.000", "end": "1.000", "text": "First"},
{"start": "1.000", "end": "2.500", "text": "Second"},
{"start": "2.500", "end": "4.000", "text": "Third"},
]
with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f:
path = f.name
try:
create_srt_file(segments, path)
with open(path, "r", encoding="utf-8") as f:
content = f.read()
self.assertIn("1\n", content)
self.assertIn("2\n", content)
self.assertIn("3\n", content)
self.assertIn("First", content)
self.assertIn("Third", content)
finally:
os.remove(path)
def test_empty_segments(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f:
path = f.name
try:
create_srt_file([], path)
with open(path, "r", encoding="utf-8") as f:
content = f.read()
self.assertEqual(content, "")
finally:
os.remove(path)
def test_unicode_text(self):
segments = [{"start": "0.000", "end": "1.000", "text": "日本語テスト"}]
with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f:
path = f.name
try:
create_srt_file(segments, path)
with open(path, "r", encoding="utf-8") as f:
content = f.read()
self.assertIn("日本語テスト", content)
finally:
os.remove(path)
class TestPrintTranscript(unittest.TestCase):
@patch("sys.stdout", new_callable=StringIO)
def test_clear_screen_uses_ansi(self, mock_stdout):
clear_screen()
output = mock_stdout.getvalue()
self.assertIn("\033[H\033[2J", output)
@patch("sys.stdout", new_callable=StringIO)
def test_print_plain_text(self, mock_stdout):
text = ["Hello", " world"]
print_transcript(text)
output = mock_stdout.getvalue()
self.assertIn("Hello world", output)
@patch("sys.stdout", new_callable=StringIO)
def test_print_with_timestamps(self, mock_stdout):
text = [
{"start": 0.0, "end": 1.0, "text": "Hello"},
{"start": 1.0, "end": 2.0, "text": "world"},
]
print_transcript(text, timestamps=True)
output = mock_stdout.getvalue()
self.assertIn("[0.0 -> 1.0]", output)
self.assertIn("Hello", output)
@patch("sys.stdout", new_callable=StringIO)
def test_print_translated(self, mock_stdout):
text = ["Bonjour", "le monde"]
print_transcript(text, translated=True)
output = mock_stdout.getvalue()
self.assertIn("Bonjour le monde", output)
@patch("sys.stdout", new_callable=StringIO)
def test_print_empty(self, mock_stdout):
print_transcript([])
output = mock_stdout.getvalue()
# empty text joined is empty string, should not crash
self.assertEqual(output.strip(), "")
if __name__ == "__main__":
unittest.main()
+26
View File
@@ -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.")
+131
View File
@@ -0,0 +1,131 @@
import unittest
from unittest.mock import patch, MagicMock
import numpy as np
import torch
from whisper_live.vad import VoiceActivityDetection, VoiceActivityDetector
class TestVoiceActivityDetectionValidation(unittest.TestCase):
"""Tests for VoiceActivityDetection input validation without requiring the ONNX model."""
@patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None)
def setUp(self):
self.vad = VoiceActivityDetection()
self.vad.sample_rates = [8000, 16000]
def test_1d_input_unsqueezed(self):
x = torch.randn(512)
x_out, sr_out = self.vad._validate_input(x, 16000)
self.assertEqual(x_out.dim(), 2)
self.assertEqual(sr_out, 16000)
def test_3d_input_raises(self):
x = torch.randn(1, 1, 512)
with self.assertRaises(ValueError):
self.vad._validate_input(x, 16000)
def test_unsupported_sample_rate_raises(self):
x = torch.randn(1, 512)
with self.assertRaises(ValueError):
self.vad._validate_input(x, 44100)
def test_too_short_audio_raises(self):
x = torch.randn(1, 1)
with self.assertRaises(ValueError):
self.vad._validate_input(x, 16000)
def test_downsample_multiple_of_16k(self):
x = torch.randn(1, 512 * 3)
x_out, sr_out = self.vad._validate_input(x, 48000)
self.assertEqual(sr_out, 16000)
self.assertEqual(x_out.shape[1], 512)
class TestVoiceActivityDetectionStateReset(unittest.TestCase):
"""Tests for VoiceActivityDetection.reset_states()."""
@patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None)
def setUp(self):
self.vad = VoiceActivityDetection()
def test_reset_creates_correct_shapes(self):
self.vad.reset_states(batch_size=4)
self.assertEqual(self.vad._state.shape, (2, 4, 128))
self.assertEqual(self.vad._context.shape[0], 0)
self.assertEqual(self.vad._last_sr, 0)
self.assertEqual(self.vad._last_batch_size, 0)
def test_reset_default_batch_size(self):
self.vad.reset_states()
self.assertEqual(self.vad._state.shape, (2, 1, 128))
class TestVoiceActivityDetectionDownload(unittest.TestCase):
"""Tests for the model download function."""
@patch("os.path.exists", return_value=True)
def test_skips_download_if_exists(self, mock_exists):
path = VoiceActivityDetection.download()
self.assertTrue(path.endswith("silero_vad.onnx"))
@patch("os.path.exists", return_value=False)
@patch("subprocess.run")
@patch("os.makedirs")
def test_downloads_if_missing(self, mock_makedirs, mock_run, mock_exists):
path = VoiceActivityDetection.download()
mock_run.assert_called_once()
self.assertIn("silero_vad.onnx", path)
@patch("os.path.exists", return_value=False)
@patch("subprocess.run", side_effect=Exception("wget not found"))
@patch("os.makedirs")
def test_handles_download_failure(self, mock_makedirs, mock_run, mock_exists):
# should not raise, just prints an error
with self.assertRaises(Exception):
VoiceActivityDetection.download()
class TestVoiceActivityDetectorThreshold(unittest.TestCase):
"""Tests for VoiceActivityDetector threshold behavior."""
@patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None)
def test_above_threshold_returns_true(self):
detector = VoiceActivityDetector.__new__(VoiceActivityDetector)
detector.model = VoiceActivityDetection()
detector.threshold = 0.5
detector.frame_rate = 16000
mock_probs = torch.tensor([[0.9, 0.8, 0.7]])
with patch.object(detector.model, "audio_forward", return_value=mock_probs):
result = detector(np.random.randn(16000).astype(np.float32))
self.assertTrue(result)
@patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None)
def test_below_threshold_returns_false(self):
detector = VoiceActivityDetector.__new__(VoiceActivityDetector)
detector.model = VoiceActivityDetection()
detector.threshold = 0.5
detector.frame_rate = 16000
mock_probs = torch.tensor([[0.1, 0.2, 0.3]])
with patch.object(detector.model, "audio_forward", return_value=mock_probs):
result = detector(np.random.randn(16000).astype(np.float32))
self.assertFalse(result)
@patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None)
def test_custom_threshold(self):
detector = VoiceActivityDetector.__new__(VoiceActivityDetector)
detector.model = VoiceActivityDetection()
detector.threshold = 0.95
detector.frame_rate = 16000
mock_probs = torch.tensor([[0.9]])
with patch.object(detector.model, "audio_forward", return_value=mock_probs):
result = detector(np.random.randn(16000).astype(np.float32))
self.assertFalse(result)
if __name__ == "__main__":
unittest.main()
+3
View File
@@ -0,0 +1,3 @@
from whisper_live.__version__ import __version__
__all__ = ['__version__']
+1 -1
View File
@@ -1 +1 @@
__version__="0.0.8"
__version__ = "0.9.0"
View File
+490
View File
@@ -0,0 +1,490 @@
import json
import logging
import threading
import time
import queue
import numpy as np
from whisper_live import metrics as wl_metrics
class ServeClientBase(object):
RATE = 16000
SERVER_READY = "SERVER_READY"
DISCONNECT = "DISCONNECT"
MAX_BUFFER_DURATION_S = 45
"""Maximum audio buffer duration in seconds before trimming."""
BUFFER_TRIM_DURATION_S = 30
"""Duration in seconds to trim from the buffer when it exceeds MAX_BUFFER_DURATION_S."""
CLIP_THRESHOLD_DURATION_S = 25
"""Duration threshold in seconds for clipping audio with no valid segments."""
CLIP_TAIL_DURATION_S = 5
"""Duration in seconds of audio to keep after clipping."""
FIRST_FRAME_WAIT_TIMEOUT_S = 0.1
"""Interval in seconds for re-checking exit while waiting for the first audio frame."""
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."""
MAX_TRANSCRIPT_LENGTH = 500
MAX_TRANSLATION_QUEUE_SIZE = 100
def __init__(
self,
client_uid,
websocket,
send_last_n_segments=10,
no_speech_thresh=0.45,
clip_audio=False,
same_output_threshold=10,
translation_queue=None,
diarization=None,
word_timestamps=False,
):
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.diarization = diarization
self.word_timestamps = word_timestamps
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
self.translation_queue = translation_queue
# Optional post-processing callable for segments.
# If set, called with a segment dict and must return a segment dict.
# Allows external projects to plug in custom post-processing
# (e.g. PII redaction, formatting, diarization) without modifying
# WhisperLive's core code.
self.segment_post_processor = None
# threading
self.lock = threading.Lock()
self.frames_ready = threading.Event()
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. The loop blocks until the first
audio frame arrives when a client is connected but still idle.
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:
while self.frames_np is None and not self.exit:
self.frames_ready.wait(timeout=self.FIRST_FRAME_WAIT_TIMEOUT_S)
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()
t0 = time.time()
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
wl_metrics.track_transcription_latency(time.time() - t0)
wl_metrics.track_audio_processed(duration)
self.handle_transcription_output(result, duration)
except Exception as e:
logging.error(f"[ERROR]: Failed to transcribe audio chunk: {e}")
wl_metrics.track_error("transcription")
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, speaker=None, words=None):
"""
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.
speaker (str, optional): Speaker label from diarization.
words (list, optional): Word-level timestamps and probabilities.
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.
"""
seg = {
'start': "{:.3f}".format(start),
'end': "{:.3f}".format(end),
'text': text,
'completed': completed,
}
if speaker is not None:
seg['speaker'] = speaker
if words is not None:
seg['words'] = words
return seg
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. When the first frame arrives, it also wakes the transcription
thread so processing can begin.
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.
"""
with self.lock:
if self.frames_np is not None and self.frames_np.shape[0] > self.MAX_BUFFER_DURATION_S*self.RATE:
self.frames_offset += float(self.BUFFER_TRIM_DURATION_S)
self.frames_np = self.frames_np[int(self.BUFFER_TRIM_DURATION_S*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.frames_ready.set()
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] > self.CLIP_THRESHOLD_DURATION_S * self.RATE:
duration = self.frames_np.shape[0] / self.RATE
self.timestamp_offset = self.frames_offset + duration - self.CLIP_TAIL_DURATION_S
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.
If a ``segment_post_processor`` callable is set, each segment is passed through it
before sending. The callable receives a segment dict and must return a segment dict.
Returns:
segments (list): A list of transcription segments to be sent to the client.
"""
if self.segment_post_processor is not None:
processed = []
for seg in segments:
try:
result = self.segment_post_processor(seg)
processed.append(result if result is not None else seg)
except Exception as e:
logging.error(f"[ERROR]: segment_post_processor failed: {e}")
processed.append(seg)
segments = processed
try:
self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments,
})
)
for seg in segments:
wl_metrics.track_segment_emitted(completed=seg.get("completed", False))
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
self.frames_ready.set()
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 _identify_speaker(self, segment):
"""Run diarization on a segment's audio slice if diarization is enabled.
Returns:
str or None: Speaker label, or None if diarization is disabled or audio unavailable.
"""
if self.diarization is None or self.frames_np is None:
return None
try:
seg_start = self.get_segment_start(segment)
seg_end = self.get_segment_end(segment)
start_sample = int(seg_start * self.RATE)
end_sample = int(seg_end * self.RATE)
samples_offset = max(0, int((self.timestamp_offset - self.frames_offset) * self.RATE))
audio_slice = self.frames_np[samples_offset + start_sample:samples_offset + end_sample]
if len(audio_slice) < self.RATE * 0.3:
return None
return self.diarization.identify_speaker(audio_slice, self.RATE)
except Exception as e:
logging.error(f"Diarization error: {e}")
return None
def _extract_words(self, segment, time_offset):
"""Extracts word-level timestamps from a segment if word_timestamps is enabled."""
if not self.word_timestamps:
return None
words = getattr(segment, "words", None)
if not words:
return None
return [
{
"word": w.word,
"start": "{:.3f}".format(time_offset + w.start),
"end": "{:.3f}".format(time_offset + w.end),
"probability": round(w.probability, 4),
}
for w in words
]
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
speaker = self._identify_speaker(s)
words = self._extract_words(s, self.timestamp_offset)
completed_segment = self.format_segment(start, end, text_, completed=True, speaker=speaker, words=words)
self.transcript.append(completed_segment)
if self.translation_queue:
try:
self.translation_queue.put(completed_segment.copy(), timeout=0.1)
except queue.Full:
logging.warning("Translation queue is full, skipping segment")
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
words = self._extract_words(segments[-1], self.timestamp_offset)
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,
words=words
)
# 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:
completed_segment = self.format_segment(
self.timestamp_offset,
self.timestamp_offset + min(duration, self.end_time_for_same_output),
self.current_out,
completed=True
)
self.transcript.append(completed_segment)
if self.translation_queue:
try:
self.translation_queue.put(completed_segment.copy(), timeout=0.1)
except queue.Full:
logging.warning("Translation queue is full, skipping segment")
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
self._trim_transcript()
return last_segment
def _trim_transcript(self):
"""Trims transcript and text lists to prevent unbounded memory growth."""
if len(self.transcript) > self.MAX_TRANSCRIPT_LENGTH:
self.transcript = self.transcript[-self.MAX_TRANSCRIPT_LENGTH:]
if len(self.text) > self.MAX_TRANSCRIPT_LENGTH:
self.text = self.text[-self.MAX_TRANSCRIPT_LENGTH:]
@@ -0,0 +1,267 @@
import os
import json
import logging
import threading
import time
import torch
import ctranslate2
from huggingface_hub import snapshot_download
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()
BATCH_WORKER = None
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=7,
cache_path="~/.cache/whisper-live/",
translation_queue=None,
hotwords=None,
diarization=None,
word_timestamps=False,
):
"""
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,
translation_queue,
diarization,
word_timestamps,
)
self.cache_path = cache_path
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 {"threshold": 0.5}
self.hotwords = hotwords
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. If model is a huggingface model_id
then it is automatically converted to ctranslate2(faster_whisper) format.
"""
model_ref = self.model_size_or_path
if model_ref in self.model_sizes:
model_to_load = model_ref
else:
logging.info(f"Model not in model_sizes")
if os.path.isdir(model_ref) and ctranslate2.contains_model(model_ref):
model_to_load = model_ref
else:
local_snapshot = snapshot_download(
repo_id = model_ref,
repo_type = "model",
)
if ctranslate2.contains_model(local_snapshot):
model_to_load = local_snapshot
else:
cache_root = os.path.expanduser(os.path.join(self.cache_path, "whisper-ct2-models/"))
os.makedirs(cache_root, exist_ok=True)
safe_name = model_ref.replace("/", "--")
ct2_dir = os.path.join(cache_root, safe_name)
if not ctranslate2.contains_model(ct2_dir):
logging.info(f"Converting '{model_ref}' to CTranslate2 @ {ct2_dir}")
ct2_converter = ctranslate2.converters.TransformersConverter(
local_snapshot,
copy_files=["tokenizer.json", "preprocessor_config.json"]
)
ct2_converter.convert(
output_dir=ct2_dir,
quantization=self.compute_type,
force=False, # skip if already up-to-date
)
model_to_load = ct2_dir
logging.info(f"Loading model: {model_to_load}")
self.transcriber = WhisperModel(
model_to_load,
device=device,
compute_type=self.compute_type,
local_files_only=False,
)
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.
"""
# Batch inference path: submit to central queue and wait
if ServeClientFasterWhisper.BATCH_WORKER is not None:
from whisper_live.batch_inference import BatchRequest
request = BatchRequest(
audio=input_sample,
language=self.language,
task=self.task,
initial_prompt=self.initial_prompt,
use_vad=self.use_vad,
vad_parameters=self.vad_parameters if self.use_vad else None,
word_timestamps=self.word_timestamps,
client_uid=self.client_uid,
)
ServeClientFasterWhisper.BATCH_WORKER.submit(request)
request.future.wait(timeout=30)
if request.error:
raise request.error
if self.language is None and request.info is not None:
self.set_language(request.info)
return request.result
# Original lock-based path (backward compatible)
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,
hotwords=self.hotwords,
word_timestamps=self.word_timestamps)
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)
+165
View File
@@ -0,0 +1,165 @@
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,
diarization=None,
):
"""
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,
None, # translation_queue — OpenVINO backend does not support translation
diarization, # speaker diarization — passed through to base class
)
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 get_segment_end(self, segment):
"""
Override base class implementation to handle OpenVINO's end timestamp sentinel value.
WhisperDecodedResultChunk.end_ts is -1.0 when the model did not predict an ending
timestamp (e.g. audio cut off mid-word). A negative end_ts causes a negative array
index in _identify_speaker(), producing an empty audio slice and silently disabling
diarization. Fall back to start_ts + 1.0 second in that case.
"""
end = getattr(segment, "end_ts", -1.0)
if end < 0:
return getattr(segment, "start_ts", 0) + 1.0
return end
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,365 @@
# Copyright (c) 2022 Idiap Research Institute, http://www.idiap.ch/
# Written by Alireza Mohammadshahi <alireza.mohammadshahi@idiap.ch>
# This is a modified version of https://github.com/huggingface/transformers/blob/main/src/transformers/models/m2m_100/tokenization_m2m_100.py
# which owns by Fariseq Authors and The HuggingFace Inc. team.
#
#
# 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.
"""Tokenization classes for SMALL100."""
import json
import os
from pathlib import Path
from shutil import copyfile
from typing import Any, Dict, List, Optional, Tuple, Union
import sentencepiece
from transformers.tokenization_utils import BatchEncoding, PreTrainedTokenizer
from transformers.utils import logging
logger = logging.get_logger(__name__)
SPIECE_UNDERLINE = ""
VOCAB_FILES_NAMES = {
"vocab_file": "vocab.json",
"spm_file": "sentencepiece.bpe.model",
"tokenizer_config_file": "tokenizer_config.json",
}
PRETRAINED_VOCAB_FILES_MAP = {
"vocab_file": {
"alirezamsh/small100": "https://huggingface.co/alirezamsh/small100/resolve/main/vocab.json",
},
"spm_file": {
"alirezamsh/small100": "https://huggingface.co/alirezamsh/small100/resolve/main/sentencepiece.bpe.model",
},
"tokenizer_config_file": {
"alirezamsh/small100": "https://huggingface.co/alirezamsh/small100/resolve/main/tokenizer_config.json",
},
}
PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {
"alirezamsh/small100": 1024,
}
# fmt: off
FAIRSEQ_LANGUAGE_CODES = {
"m2m100": ["af", "am", "ar", "ast", "az", "ba", "be", "bg", "bn", "br", "bs", "ca", "ceb", "cs", "cy", "da", "de", "el", "en", "es", "et", "fa", "ff", "fi", "fr", "fy", "ga", "gd", "gl", "gu", "ha", "he", "hi", "hr", "ht", "hu", "hy", "id", "ig", "ilo", "is", "it", "ja", "jv", "ka", "kk", "km", "kn", "ko", "lb", "lg", "ln", "lo", "lt", "lv", "mg", "mk", "ml", "mn", "mr", "ms", "my", "ne", "nl", "no", "ns", "oc", "or", "pa", "pl", "ps", "pt", "ro", "ru", "sd", "si", "sk", "sl", "so", "sq", "sr", "ss", "su", "sv", "sw", "ta", "th", "tl", "tn", "tr", "uk", "ur", "uz", "vi", "wo", "xh", "yi", "yo", "zh", "zu"]
}
# fmt: on
class SMALL100Tokenizer(PreTrainedTokenizer):
"""
Construct an SMALL100 tokenizer. Based on [SentencePiece](https://github.com/google/sentencepiece).
This tokenizer inherits from [`PreTrainedTokenizer`] which contains most of the main methods. Users should refer to
this superclass for more information regarding those methods.
Args:
vocab_file (`str`):
Path to the vocabulary file.
spm_file (`str`):
Path to [SentencePiece](https://github.com/google/sentencepiece) file (generally has a .spm extension) that
contains the vocabulary.
tgt_lang (`str`, *optional*):
A string representing the target language.
eos_token (`str`, *optional*, defaults to `"</s>"`):
The end of sequence token.
sep_token (`str`, *optional*, defaults to `"</s>"`):
The separator token, which is used when building a sequence from multiple sequences, e.g. two sequences for
sequence classification or for a text and a question for question answering. It is also used as the last
token of a sequence built with special tokens.
unk_token (`str`, *optional*, defaults to `"<unk>"`):
The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this
token instead.
pad_token (`str`, *optional*, defaults to `"<pad>"`):
The token used for padding, for example when batching sequences of different lengths.
language_codes (`str`, *optional*):
What language codes to use. Should be `"m2m100"`.
sp_model_kwargs (`dict`, *optional*):
Will be passed to the `SentencePieceProcessor.__init__()` method. The [Python wrapper for
SentencePiece](https://github.com/google/sentencepiece/tree/master/python) can be used, among other things,
to set:
- `enable_sampling`: Enable subword regularization.
- `nbest_size`: Sampling parameters for unigram. Invalid for BPE-Dropout.
- `nbest_size = {0,1}`: No sampling is performed.
- `nbest_size > 1`: samples from the nbest_size results.
- `nbest_size < 0`: assuming that nbest_size is infinite and samples from the all hypothesis (lattice)
using forward-filtering-and-backward-sampling algorithm.
- `alpha`: Smoothing parameter for unigram sampling, and dropout probability of merge operations for
BPE-dropout.
Examples:
```python
>>> from tokenization_small100 import SMALL100Tokenizer
>>> tokenizer = SMALL100Tokenizer.from_pretrained("alirezamsh/small100", tgt_lang="ro")
>>> src_text = " UN Chief Says There Is No Military Solution in Syria"
>>> tgt_text = "Şeful ONU declară că nu există o soluţie militară în Siria"
>>> model_inputs = tokenizer(src_text, text_target=tgt_text, return_tensors="pt")
>>> model(**model_inputs) # should work
```"""
vocab_files_names = VOCAB_FILES_NAMES
max_model_input_sizes = PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES
pretrained_vocab_files_map = PRETRAINED_VOCAB_FILES_MAP
model_input_names = ["input_ids", "attention_mask"]
prefix_tokens: List[int] = []
suffix_tokens: List[int] = []
def __init__(
self,
vocab_file,
spm_file,
tgt_lang=None,
bos_token="<s>",
eos_token="</s>",
sep_token="</s>",
pad_token="<pad>",
unk_token="<unk>",
language_codes="m2m100",
sp_model_kwargs: Optional[Dict[str, Any]] = None,
num_madeup_words=8,
**kwargs,
) -> None:
self.sp_model_kwargs = {} if sp_model_kwargs is None else sp_model_kwargs
self.language_codes = language_codes
fairseq_language_code = FAIRSEQ_LANGUAGE_CODES[language_codes]
self.lang_code_to_token = {lang_code: f"__{lang_code}__" for lang_code in fairseq_language_code}
kwargs["additional_special_tokens"] = kwargs.get("additional_special_tokens", [])
kwargs["additional_special_tokens"] += [
self.get_lang_token(lang_code)
for lang_code in fairseq_language_code
if self.get_lang_token(lang_code) not in kwargs["additional_special_tokens"]
]
self.vocab_file = vocab_file
self.encoder = load_json(vocab_file)
self.decoder = {v: k for k, v in self.encoder.items()}
self.spm_file = spm_file
self.sp_model = load_spm(spm_file, self.sp_model_kwargs)
self.encoder_size = len(self.encoder)
self.lang_token_to_id = {
self.get_lang_token(lang_code): self.encoder_size + i for i, lang_code in enumerate(fairseq_language_code)
}
self.lang_code_to_id = {lang_code: self.encoder_size + i for i, lang_code in enumerate(fairseq_language_code)}
self.id_to_lang_token = {v: k for k, v in self.lang_token_to_id.items()}
self._tgt_lang = tgt_lang if tgt_lang is not None else "en"
self.cur_lang_id = self.get_lang_id(self._tgt_lang)
self.num_madeup_words = num_madeup_words
super().__init__(
tgt_lang=tgt_lang,
bos_token=bos_token,
eos_token=eos_token,
sep_token=sep_token,
unk_token=unk_token,
pad_token=pad_token,
language_codes=language_codes,
sp_model_kwargs=self.sp_model_kwargs,
num_madeup_words=num_madeup_words,
**kwargs,
)
self.set_lang_special_tokens(self._tgt_lang)
@property
def vocab_size(self) -> int:
return len(self.encoder) + len(self.lang_token_to_id) + self.num_madeup_words
@property
def tgt_lang(self) -> str:
return self._tgt_lang
@tgt_lang.setter
def tgt_lang(self, new_tgt_lang: str) -> None:
self._tgt_lang = new_tgt_lang
self.set_lang_special_tokens(self._tgt_lang)
def _tokenize(self, text: str) -> List[str]:
return self.sp_model.encode(text, out_type=str)
def _convert_token_to_id(self, token):
if token in self.lang_token_to_id:
return self.lang_token_to_id[token]
return self.encoder.get(token, self.encoder[self.unk_token])
def _convert_id_to_token(self, index: int) -> str:
"""Converts an index (integer) in a token (str) using the decoder."""
if index in self.id_to_lang_token:
return self.id_to_lang_token[index]
return self.decoder.get(index, self.unk_token)
def convert_tokens_to_string(self, tokens: List[str]) -> str:
"""Converts a sequence of tokens (strings for sub-words) in a single string."""
return self.sp_model.decode(tokens)
def get_special_tokens_mask(
self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None, already_has_special_tokens: bool = False
) -> List[int]:
"""
Retrieve sequence ids from a token list that has no special tokens added. This method is called when adding
special tokens using the tokenizer `prepare_for_model` method.
Args:
token_ids_0 (`List[int]`):
List of IDs.
token_ids_1 (`List[int]`, *optional*):
Optional second list of IDs for sequence pairs.
already_has_special_tokens (`bool`, *optional*, defaults to `False`):
Whether or not the token list is already formatted with special tokens for the model.
Returns:
`List[int]`: A list of integers in the range [0, 1]: 1 for a special token, 0 for a sequence token.
"""
if already_has_special_tokens:
return super().get_special_tokens_mask(
token_ids_0=token_ids_0, token_ids_1=token_ids_1, already_has_special_tokens=True
)
prefix_ones = [1] * len(self.prefix_tokens)
suffix_ones = [1] * len(self.suffix_tokens)
if token_ids_1 is None:
return prefix_ones + ([0] * len(token_ids_0)) + suffix_ones
return prefix_ones + ([0] * len(token_ids_0)) + ([0] * len(token_ids_1)) + suffix_ones
def build_inputs_with_special_tokens(
self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None
) -> List[int]:
"""
Build model inputs from a sequence or a pair of sequence for sequence classification tasks by concatenating and
adding special tokens. An MBART sequence has the following format, where `X` represents the sequence:
- `input_ids` (for encoder) `X [eos, src_lang_code]`
- `decoder_input_ids`: (for decoder) `X [eos, tgt_lang_code]`
BOS is never used. Pairs of sequences are not the expected use case, but they will be handled without a
separator.
Args:
token_ids_0 (`List[int]`):
List of IDs to which the special tokens will be added.
token_ids_1 (`List[int]`, *optional*):
Optional second list of IDs for sequence pairs.
Returns:
`List[int]`: List of [input IDs](../glossary#input-ids) with the appropriate special tokens.
"""
if token_ids_1 is None:
if self.prefix_tokens is None:
return token_ids_0 + self.suffix_tokens
else:
return self.prefix_tokens + token_ids_0 + self.suffix_tokens
# We don't expect to process pairs, but leave the pair logic for API consistency
if self.prefix_tokens is None:
return token_ids_0 + token_ids_1 + self.suffix_tokens
else:
return self.prefix_tokens + token_ids_0 + token_ids_1 + self.suffix_tokens
def get_vocab(self) -> Dict:
vocab = {self.convert_ids_to_tokens(i): i for i in range(self.vocab_size)}
vocab.update(self.added_tokens_encoder)
return vocab
def __getstate__(self) -> Dict:
state = self.__dict__.copy()
state["sp_model"] = None
return state
def __setstate__(self, d: Dict) -> None:
self.__dict__ = d
# for backward compatibility
if not hasattr(self, "sp_model_kwargs"):
self.sp_model_kwargs = {}
self.sp_model = load_spm(self.spm_file, self.sp_model_kwargs)
def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> Tuple[str]:
save_dir = Path(save_directory)
if not save_dir.is_dir():
raise OSError(f"{save_directory} should be a directory")
vocab_save_path = save_dir / (
(filename_prefix + "-" if filename_prefix else "") + self.vocab_files_names["vocab_file"]
)
spm_save_path = save_dir / (
(filename_prefix + "-" if filename_prefix else "") + self.vocab_files_names["spm_file"]
)
save_json(self.encoder, vocab_save_path)
if os.path.abspath(self.spm_file) != os.path.abspath(spm_save_path) and os.path.isfile(self.spm_file):
copyfile(self.spm_file, spm_save_path)
elif not os.path.isfile(self.spm_file):
with open(spm_save_path, "wb") as fi:
content_spiece_model = self.sp_model.serialized_model_proto()
fi.write(content_spiece_model)
return (str(vocab_save_path), str(spm_save_path))
def prepare_seq2seq_batch(
self,
src_texts: List[str],
tgt_texts: Optional[List[str]] = None,
tgt_lang: str = "ro",
**kwargs,
) -> BatchEncoding:
self.tgt_lang = tgt_lang
self.set_lang_special_tokens(self.tgt_lang)
return super().prepare_seq2seq_batch(src_texts, tgt_texts, **kwargs)
def _build_translation_inputs(self, raw_inputs, tgt_lang: Optional[str], **extra_kwargs):
"""Used by translation pipeline, to prepare inputs for the generate function"""
if tgt_lang is None:
raise ValueError("Translation requires a `tgt_lang` for this model")
self.tgt_lang = tgt_lang
inputs = self(raw_inputs, add_special_tokens=True, **extra_kwargs)
return inputs
def _switch_to_input_mode(self):
self.set_lang_special_tokens(self.tgt_lang)
def _switch_to_target_mode(self):
self.prefix_tokens = None
self.suffix_tokens = [self.eos_token_id]
def set_lang_special_tokens(self, src_lang: str) -> None:
"""Reset the special tokens to the tgt lang setting. No prefix and suffix=[eos, tgt_lang_code]."""
lang_token = self.get_lang_token(src_lang)
self.cur_lang_id = self.lang_token_to_id[lang_token]
self.prefix_tokens = [self.cur_lang_id]
self.suffix_tokens = [self.eos_token_id]
def get_lang_token(self, lang: str) -> str:
return self.lang_code_to_token[lang]
def get_lang_id(self, lang: str) -> int:
lang_token = self.get_lang_token(lang)
return self.lang_token_to_id[lang_token]
def load_spm(path: str, sp_model_kwargs: Dict[str, Any]) -> sentencepiece.SentencePieceProcessor:
spm = sentencepiece.SentencePieceProcessor(**sp_model_kwargs)
spm.Load(str(path))
return spm
def load_json(path: str) -> Union[Dict, List]:
with open(path, "r") as f:
return json.load(f)
def save_json(data, path: str) -> None:
with open(path, "w") as f:
json.dump(data, f, indent=2)
+218
View File
@@ -0,0 +1,218 @@
import json
import logging
import threading
import time
import queue
from typing import Dict, Any, Optional
import torch
import threading
from transformers import M2M100ForConditionalGeneration
from whisper_live.backend.tokenization_small100 import SMALL100Tokenizer
from whisper_live.backend.base import ServeClientBase
class ServeClientTranslation(ServeClientBase):
"""
Handles translation of completed transcription segments in a separate thread.
Reads from a queue populated by the transcription backend and sends translated
segments back to the client via WebSocket.
"""
def __init__(
self,
client_uid,
websocket,
translation_queue,
target_language="fr",
send_last_n_segments=10,
model_name="alirezamsh/small100"
):
"""
Initialize the translation client.
Args:
client_uid (str): Unique identifier for the client
websocket: WebSocket connection to the client
translation_queue (queue.Queue): Queue containing completed segments to translate
target_language (str): Target language code (default: "fr" for French)
send_last_n_segments (int): Number of recent translated segments to send
model_name (str): Translation model name to use
"""
super().__init__(client_uid, websocket, send_last_n_segments)
self.translation_queue = translation_queue
self.target_language = target_language
self.model_name = model_name
self.translated_segments = []
self.translation_model = None
self.tokenizer = None
self.device = None
self.model_loaded = False
self.load_translation_model()
def load_translation_model(self):
"""Load the translation model and tokenizer."""
try:
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
logging.info(f"Loading translation model on device: {self.device}")
self.translation_model = M2M100ForConditionalGeneration.from_pretrained(
self.model_name
).to(self.device)
self.tokenizer = SMALL100Tokenizer.from_pretrained(self.model_name)
self.tokenizer.tgt_lang = self.target_language
self.model_loaded = True
logging.info(f"Translation model loaded successfully. Target language: {self.target_language}")
except Exception as e:
logging.error(f"Failed to load translation model: {e}")
self.translation_model = None
self.tokenizer = None
self.model_loaded = False
def translate_text(self, text: str) -> str:
"""
Translate a single text segment.
Args:
text (str): Text to translate
Returns:
str: Translated text or original text if translation fails
"""
if not self.model_loaded or not text.strip():
return text
try:
# Encode input and move to device
encoded_input = self.tokenizer(text, return_tensors="pt").to(self.device)
# Generate translation
with torch.no_grad():
generated_tokens = self.translation_model.generate(**encoded_input)
# Decode output
output = self.tokenizer.batch_decode(generated_tokens, skip_special_tokens=True)
return output[0] if output else text
except Exception as e:
logging.error(f"Translation failed for text '{text}': {e}")
return text
def process_translation_queue(self):
"""
Process segments from the translation queue.
Continuously reads from the queue until None is received (exit signal).
"""
logging.info(f"Starting translation processing for client {self.client_uid}")
while not self.exit:
try:
# Get segment from queue with timeout
segment = self.translation_queue.get(timeout=1.0)
# Check for exit signal
if segment is None:
logging.info(f"Received exit signal for translation client {self.client_uid}")
break
# Only translate completed segments
if not segment.get("completed", False):
self.translation_queue.task_done()
continue
# Translate the segment
original_text = segment.get("text", "")
translated_text = self.translate_text(original_text)
# Create translated segment
translated_segment = {
"start": segment["start"],
"end": segment["end"],
"text": translated_text,
"completed": segment.get("completed", False),
"target_language": self.target_language
}
self.translated_segments.append(translated_segment)
segments_to_send = self.prepare_translated_segments()
self.send_translation_to_client(segments_to_send)
self.translation_queue.task_done()
except queue.Empty:
continue
except Exception as e:
logging.error(f"Error processing translation queue: {e}")
continue
logging.info(f"Translation processing ended for client {self.client_uid}")
def prepare_translated_segments(self):
"""
Prepare the last n translated segments to send to client.
Returns:
list: List of recent translated segments
"""
if len(self.translated_segments) >= self.send_last_n_segments:
return self.translated_segments[-self.send_last_n_segments:]
return self.translated_segments[:]
def send_translation_to_client(self, translated_segments):
"""
Send translated segments to the client via WebSocket.
Args:
translated_segments (list): List of translated segments to send
"""
try:
self.websocket.send(
json.dumps({
"uid": self.client_uid,
"translated_segments": translated_segments,
})
)
except Exception as e:
logging.error(f"[ERROR]: Sending translation data to client: {e}")
def speech_to_text(self):
"""
Override parent method to handle translation processing.
This method will be called when the translation thread starts.
"""
self.process_translation_queue()
def set_target_language(self, language: str):
"""
Change the target language for translation.
Args:
language (str): New target language code
"""
self.target_language = language
if self.tokenizer:
self.tokenizer.tgt_lang = language
logging.info(f"Target language changed to: {language}")
def cleanup(self):
"""Clean up translation resources."""
logging.info(f"Cleaning up translation resources for client {self.client_uid}")
self.exit = True
try:
self.translation_queue.put(None, timeout=1.0)
except:
pass
self.translated_segments.clear()
if self.translation_model:
del self.translation_model
self.translation_model = None
if self.tokenizer:
del self.tokenizer
self.tokenizer = None
if self.device and self.device.type == 'cuda':
torch.cuda.empty_cache()
+210
View File
@@ -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}")
+450
View File
@@ -0,0 +1,450 @@
"""
Batch inference scheduler for WhisperLive.
Replaces the per-session SINGLE_MODEL_LOCK with a queue-based batch system.
Multiple sessions submit audio to a central queue; a single dedicated thread
collects pending requests and runs them as a GPU batch via CTranslate2's
batched encode() + generate() API.
For batch_size=1, falls back to standard transcriber.transcribe() for
identical behavior to the non-batched path.
Usage:
Enable via ``--batch_inference`` CLI flag. The batch worker is lazily
started after the first client connects and the shared model is loaded.
Thread safety:
- ``queue.Queue`` is stdlib thread-safe.
- Each ``BatchRequest.future`` (``threading.Event``) is written by the
batch worker BEFORE ``.set()``, read by the session thread AFTER
``.wait()`` — no data race.
- Only the batch worker thread touches the GPU model — zero lock
contention between session threads.
"""
import logging
import queue
import threading
import time
from dataclasses import dataclass, field
from math import ceil
from typing import Any, Dict, List, Optional
import numpy as np
from faster_whisper.audio import pad_or_trim
from faster_whisper.tokenizer import Tokenizer
from faster_whisper.vad import (
VadOptions,
collect_chunks,
get_speech_timestamps,
)
from whisper_live.transcriber.transcriber_faster_whisper import (
Segment,
TranscriptionInfo,
get_compression_ratio,
get_suppressed_tokens,
)
@dataclass
class BatchRequest:
"""A single inference request submitted by a session thread.
The session thread creates this, calls ``BatchInferenceWorker.submit()``,
then blocks on ``future.wait()``. The batch worker fills ``result``
and/or ``error``, then signals ``future.set()``.
Attributes:
audio: Raw audio samples (float32, 16 kHz mono).
language: ISO language code or None for auto-detection.
task: ``"transcribe"`` or ``"translate"``.
initial_prompt: Optional prompt for Whisper conditioning.
use_vad: Whether to apply Voice Activity Detection.
vad_parameters: Parameters forwarded to ``VadOptions``.
future: Event signaled when the result is ready.
result: List of ``Segment`` objects (filled by worker).
info: ``TranscriptionInfo`` metadata (filled by worker).
error: Exception instance if processing failed.
"""
audio: np.ndarray
language: Optional[str] = None
task: str = "transcribe"
initial_prompt: Optional[str] = None
use_vad: bool = True
vad_parameters: Optional[Dict] = None
word_timestamps: bool = False
client_uid: Optional[str] = None
# Signaling
future: threading.Event = field(default_factory=threading.Event)
# Results (filled by batch worker)
result: Optional[Any] = None
info: Optional[Any] = None
error: Optional[Exception] = None
class BatchInferenceWorker:
"""Central batch inference scheduler for the faster_whisper backend.
Owns a single daemon thread that is the **only** thread touching the GPU
model. Per-session transcription threads submit ``BatchRequest`` objects
and block on ``future.wait()`` instead of competing for
``SINGLE_MODEL_LOCK``.
The worker loop:
1. Blocks until the first request arrives from the queue.
2. Waits up to ``batch_window_ms`` for additional requests (up to
``max_batch_size``).
3. Processes the collected batch:
- **batch_size == 1**: delegates to ``transcriber.transcribe()`` for
identical behavior to the non-batched path.
- **batch_size > 1**: runs a custom batched GPU path using
CTranslate2's ``encode()`` + ``generate()`` APIs.
Args:
transcriber: The shared ``WhisperModel`` instance.
max_batch_size: Maximum number of requests per batch.
batch_window_ms: Maximum time (ms) to wait for the batch to fill
after the first request arrives.
"""
def __init__(
self,
transcriber,
max_batch_size: int = 8,
batch_window_ms: int = 50,
):
self.transcriber = transcriber
self.max_batch_size = max_batch_size
self.batch_window_ms = batch_window_ms
self._queue: queue.Queue = queue.Queue()
self._stop_event = threading.Event()
self._thread: Optional[threading.Thread] = None
def start(self):
"""Start the background batch worker thread."""
self._thread = threading.Thread(target=self._worker_loop, daemon=True)
self._thread.start()
logging.info(
f"[BatchInference] Started (max_batch={self.max_batch_size}, "
f"window={self.batch_window_ms}ms)"
)
def stop(self):
"""Signal the worker to stop and wait for it to finish."""
self._stop_event.set()
if self._thread:
self._thread.join(timeout=5)
def submit(self, request: BatchRequest):
"""Submit an inference request to the batch queue.
Args:
request: The ``BatchRequest`` to enqueue. The caller should
then call ``request.future.wait()`` to block until the
result is ready.
"""
self._queue.put(request)
# -------------------------------------------------------------------------
# Worker loop
# -------------------------------------------------------------------------
def _worker_loop(self):
"""Main loop: collect requests into batches and process them."""
while not self._stop_event.is_set():
batch: List[BatchRequest] = []
# Block until first request arrives
try:
first = self._queue.get(timeout=0.5)
batch.append(first)
except queue.Empty:
continue
# Collect more requests within the batch window
deadline = time.monotonic() + (self.batch_window_ms / 1000.0)
while len(batch) < self.max_batch_size:
remaining = deadline - time.monotonic()
if remaining <= 0:
break
try:
item = self._queue.get(timeout=remaining)
batch.append(item)
except queue.Empty:
break
# Process the collected batch
try:
self._process_batch(batch)
except Exception as e:
logging.error(f"[BatchInference] Batch processing error: {e}")
for req in batch:
if not req.future.is_set():
req.error = e
req.future.set()
# -------------------------------------------------------------------------
# Batch processing
# -------------------------------------------------------------------------
def _process_batch(self, batch: List[BatchRequest]):
"""Dispatch to single or multi-item processing."""
if len(batch) == 1:
self._process_single(batch[0])
return
logging.info(f"[BatchInference] Processing batch of {len(batch)}")
self._process_multi(batch)
def _process_single(self, req: BatchRequest):
"""Process a single request using standard ``transcriber.transcribe()``.
This path is used when only one request is available in the batch
window, ensuring identical behavior to the non-batched code path.
"""
try:
result, info = self.transcriber.transcribe(
req.audio,
language=req.language,
task=req.task,
initial_prompt=req.initial_prompt,
vad_filter=req.use_vad,
vad_parameters=req.vad_parameters if req.use_vad else None,
)
# Materialize the generator into a list
req.result = list(result) if result is not None else []
req.info = info
except Exception as e:
req.error = e
finally:
req.future.set()
def _process_multi(self, batch: List[BatchRequest]):
"""Batched GPU path: encode + generate for multiple sessions at once.
Pipeline:
1. Per-item CPU preprocessing (VAD filtering + mel feature extraction)
2. Batch GPU encode — single ``transcriber.encode()`` call
3. Per-item prompt construction (handles different languages/tasks)
4. Batch GPU generate — single ``transcriber.model.generate()`` call
5. Per-item segment parsing and result dispatch
"""
# Step 1: Per-item CPU preprocessing (VAD + feature extraction)
preprocessed = []
for req in batch:
try:
audio = req.audio
speech_chunks = None
if req.use_vad:
vad_params = req.vad_parameters or {}
vad_opts = VadOptions(**vad_params) if isinstance(vad_params, dict) else vad_params
speech_chunks = get_speech_timestamps(audio, vad_opts)
if speech_chunks:
audio_chunks, _ = collect_chunks(audio, speech_chunks)
audio = np.concatenate(audio_chunks, axis=0) if audio_chunks else audio
if audio.shape[0] == 0:
# No speech detected — return empty result immediately
req.result = []
req.info = self._make_info(req, 0.0, 0.0)
req.future.set()
continue
duration = audio.shape[0] / self.transcriber.feature_extractor.sampling_rate
features = self.transcriber.feature_extractor(audio)
features = pad_or_trim(features) # -> [n_mels, 3000]
preprocessed.append((req, features, audio, duration, speech_chunks))
except Exception as e:
req.error = e
req.future.set()
if not preprocessed:
return
try:
# Step 2: Batch GPU encode
feature_batch = np.stack([p[1] for p in preprocessed]) # [B, n_mels, 3000]
encoder_output = self.transcriber.encode(feature_batch)
# Step 3: Build per-item prompts (handles different languages/tasks)
tokenizers_list = []
prompts = []
resolved_languages = []
for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed):
lang = req.language
# If language unknown, detect from encoder output
if lang is None:
try:
lang_results = self.transcriber.model.detect_language(encoder_output)
if lang_results and len(lang_results) > i:
detected = lang_results[i]
if detected:
lang = detected[0][0].strip("<|>")
except Exception:
lang = "en" # fallback
resolved_languages.append(lang or "en")
tokenizer = Tokenizer(
self.transcriber.hf_tokenizer,
self.transcriber.model.is_multilingual,
task=req.task,
language=lang or "en",
)
previous_tokens = []
if req.initial_prompt:
previous_tokens = tokenizer.encode(" " + req.initial_prompt.strip())
prompt = self.transcriber.get_prompt(
tokenizer,
previous_tokens=previous_tokens,
without_timestamps=False,
)
tokenizers_list.append(tokenizer)
prompts.append(prompt)
# Step 4: Batch GPU generate with per-item temperature fallback.
# Mirrors faster_whisper.transcribe()'s fallback loop. Items that
# pass quality thresholds at lower temperature keep their result;
# only failed items are re-decoded at the next temperature.
suppress_tokens = get_suppressed_tokens(tokenizers_list[0], [-1])
temperatures = [0.0, 0.2, 0.4, 0.6, 0.8, 1.0]
comp_thresh = 2.4
logprob_thresh = -1.0
no_speech_thresh = 0.6
n = len(preprocessed)
final_results = [None] * n # tuples of (gen_result, avg_logprob, used_temp)
pending_indices = list(range(n))
for temp in temperatures:
if not pending_indices:
break
if len(pending_indices) == n:
sub_encoder = encoder_output
else:
# Re-encode features for just the pending items to get
# an encoder_output of the right batch dimension.
sub_feature_batch = np.stack(
[preprocessed[i][1] for i in pending_indices]
)
sub_encoder = self.transcriber.encode(sub_feature_batch)
sub_prompts = [prompts[i] for i in pending_indices]
gen_kwargs = dict(
beam_size=5 if temp == 0.0 else 1,
patience=1,
length_penalty=1,
max_length=self.transcriber.max_length,
suppress_blank=True,
suppress_tokens=suppress_tokens,
return_scores=True,
return_no_speech_prob=True,
sampling_temperature=temp,
repetition_penalty=1,
no_repeat_ngram_size=0,
)
batch_results = self.transcriber.model.generate(
sub_encoder, sub_prompts, **gen_kwargs
)
next_pending = []
for j, idx in enumerate(pending_indices):
gen_result = batch_results[j]
tokens = gen_result.sequences_ids[0]
seq_len = len(tokens)
cum_logprob = gen_result.scores[0] * seq_len
avg_logprob = cum_logprob / (seq_len + 1) if seq_len > 0 else 0.0
raw_text = tokenizers_list[idx].decode(tokens).strip()
comp_ratio = get_compression_ratio(raw_text) if raw_text else 0.0
bad = (
comp_ratio > comp_thresh
or avg_logprob < logprob_thresh
)
# High no_speech + low logprob -> treat as silence, accept empty.
is_silence = (
gen_result.no_speech_prob > no_speech_thresh
and avg_logprob < logprob_thresh
)
if not bad or is_silence or temp == temperatures[-1]:
final_results[idx] = (gen_result, avg_logprob, temp)
else:
next_pending.append(idx)
pending_indices = next_pending
# Step 5: Per-item segment parsing and result dispatch
for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed):
try:
tokenizer = tokenizers_list[i]
gen_result, avg_logprob, used_temp = final_results[i]
tokens = gen_result.sequences_ids[0]
segment_size = int(ceil(duration) * self.transcriber.frames_per_second)
subsegments, _, _ = self.transcriber._split_segments_by_timestamps(
tokenizer=tokenizer,
tokens=tokens,
time_offset=0,
segment_size=segment_size,
segment_duration=duration,
seek=0,
)
segments = []
for seg_idx, subseg in enumerate(subsegments):
text = tokenizer.decode(subseg["tokens"]).strip()
if not text:
continue
segments.append(Segment(
id=seg_idx,
seek=subseg.get("seek", 0),
start=subseg["start"],
end=subseg["end"],
text=text,
tokens=subseg["tokens"],
avg_logprob=avg_logprob,
compression_ratio=get_compression_ratio(text),
no_speech_prob=gen_result.no_speech_prob,
words=None,
temperature=used_temp,
))
req.result = segments
req.info = self._make_info(
req, duration, duration,
language=resolved_languages[i],
)
except Exception as e:
req.error = e
finally:
req.future.set()
except Exception as e:
logging.error(f"[BatchInference] GPU batch error: {e}")
for req, *_ in preprocessed:
if not req.future.is_set():
req.error = e
req.future.set()
def _make_info(self, req, duration, duration_after_vad, language=None):
"""Build a ``TranscriptionInfo`` for the given request."""
return TranscriptionInfo(
language=language or req.language or "en",
language_probability=1.0,
duration=duration,
duration_after_vad=duration_after_vad,
all_language_probs=None,
transcription_options=None,
vad_options=None,
)
+916 -265
View File
File diff suppressed because it is too large Load Diff
+181
View File
@@ -0,0 +1,181 @@
"""
Optional speaker diarization module for WhisperLive.
Uses speaker embeddings and online clustering to assign speaker labels
to transcription segments in real-time. Requires pyannote.audio as an
optional dependency.
Install: pip install pyannote.audio
"""
import logging
import numpy as np
def load_audio(file_path, sample_rate=16000):
"""Load an audio file as mono float32 PCM at the requested sample rate."""
import av
container = av.open(file_path)
resampler = av.AudioResampler(format="flt", layout="mono", rate=sample_rate)
chunks = []
try:
for frame in container.decode(audio=0):
for resampled_frame in resampler.resample(frame):
chunks.append(
resampled_frame.to_ndarray().reshape(-1).astype(np.float32)
)
finally:
container.close()
if not chunks:
return np.array([], dtype=np.float32)
return np.concatenate(chunks)
class SpeakerDiarizer:
"""Real-time speaker diarization using speaker embeddings and online clustering.
Each completed transcription segment's audio is passed through a speaker
embedding model. The embedding is compared against known speakers using
cosine similarity. If no match exceeds the threshold, a new speaker is
created.
Args:
similarity_threshold (float): Minimum cosine similarity to match an
existing speaker. Lower values merge speakers more aggressively.
Default 0.55.
max_speakers (int): Maximum number of distinct speakers to track.
Once reached, new segments are assigned to the closest existing
speaker. Default 10.
embedding_model (str): The pyannote embedding model to use.
Default "pyannote/wespeaker-voxceleb-resnet34-LM".
hf_token (str or None): HuggingFace token for gated model access.
"""
def __init__(
self,
similarity_threshold=0.55,
max_speakers=10,
embedding_model="pyannote/wespeaker-voxceleb-resnet34-LM",
hf_token=None,
speaker_names=None,
):
self.similarity_threshold = similarity_threshold
self.max_speakers = max_speakers
self.speaker_names = list(speaker_names or [])
self.speakers = {} # speaker_id -> embedding (averaged)
self._speaker_count = 0
self._model = None
self._embedding_model_name = embedding_model
self._hf_token = hf_token
def _next_speaker_id(self):
if self._speaker_count < len(self.speaker_names):
return self.speaker_names[self._speaker_count]
return f"SPEAKER_{self._speaker_count:02d}"
def _load_model(self):
"""Lazy-load the embedding model on first use."""
if self._model is not None:
return
try:
from pyannote.audio import Model, Inference
import torch
model = Model.from_pretrained(
self._embedding_model_name,
use_auth_token=self._hf_token,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
self._model = Inference(model, window="whole", device=torch.device(device))
logging.info(f"Speaker embedding model loaded on {device}")
except ImportError:
raise ImportError(
"pyannote.audio is required for speaker diarization. "
"Install it with: pip install pyannote.audio"
)
def _compute_embedding(self, audio_np, sample_rate=16000):
"""Compute a speaker embedding from an audio numpy array.
Args:
audio_np (np.ndarray): 1-D float32 audio samples.
sample_rate (int): Sample rate of the audio.
Returns:
np.ndarray: Speaker embedding vector, or None if audio is too short.
"""
self._load_model()
if len(audio_np) < sample_rate * 0.3:
return None
waveform = {
"waveform": __import__("torch").tensor(audio_np).unsqueeze(0),
"sample_rate": sample_rate,
}
embedding = self._model(waveform)
return embedding / np.linalg.norm(embedding)
@staticmethod
def _cosine_similarity(a, b):
"""Compute cosine similarity between two vectors."""
return float(np.dot(a, b))
def identify_speaker(self, audio_np, sample_rate=16000):
"""Identify or create a speaker from an audio segment.
Args:
audio_np (np.ndarray): 1-D float32 audio for the segment.
sample_rate (int): Sample rate. Default 16000.
Returns:
str or None: Speaker label (e.g. "SPEAKER_00"), or None if
the audio is too short to embed.
"""
embedding = self._compute_embedding(audio_np, sample_rate)
if embedding is None:
return None
best_speaker = None
best_sim = -1.0
for speaker_id, stored_emb in self.speakers.items():
sim = self._cosine_similarity(embedding, stored_emb)
if sim > best_sim:
best_sim = sim
best_speaker = speaker_id
if best_sim >= self.similarity_threshold:
# Update running average for the matched speaker
self.speakers[best_speaker] = (
self.speakers[best_speaker] * 0.9 + embedding * 0.1
)
# Re-normalize
self.speakers[best_speaker] /= np.linalg.norm(self.speakers[best_speaker])
return best_speaker
if len(self.speakers) >= self.max_speakers:
# Assign to closest speaker
return (
best_speaker if best_speaker else f"SPEAKER_{self._speaker_count:02d}"
)
# Create a new speaker
speaker_id = self._next_speaker_id()
self._speaker_count += 1
self.speakers[speaker_id] = embedding
return speaker_id
def enroll_speaker(self, speaker_name, audio_np, sample_rate=16000):
"""Enroll a known speaker from reference audio."""
embedding = self._compute_embedding(audio_np, sample_rate)
if embedding is None:
return False
self.speakers[speaker_name] = embedding
return True
def reset(self):
"""Reset all speaker state."""
self.speakers.clear()
self._speaker_count = 0
+122
View File
@@ -0,0 +1,122 @@
"""
Prometheus metrics for WhisperLive server.
Exposes a /metrics HTTP endpoint on a configurable port for Prometheus scraping.
All metrics are optional — the server works fine without prometheus_client installed.
"""
import logging
import threading
try:
from prometheus_client import (
Counter,
Gauge,
Histogram,
start_http_server,
)
CONNECTIONS_TOTAL = Counter(
"whisperlive_connections_total",
"Total WebSocket connections accepted",
)
CONNECTIONS_ACTIVE = Gauge(
"whisperlive_connections_active",
"Currently active WebSocket connections",
)
CONNECTIONS_REJECTED = Counter(
"whisperlive_connections_rejected_total",
"Connections rejected (server full or auth failure)",
["reason"],
)
TRANSCRIPTION_LATENCY = Histogram(
"whisperlive_transcription_latency_seconds",
"Time to transcribe a single audio chunk",
buckets=(0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0),
)
AUDIO_PROCESSED = Counter(
"whisperlive_audio_processed_seconds_total",
"Total seconds of audio processed",
)
SEGMENTS_EMITTED = Counter(
"whisperlive_segments_emitted_total",
"Total transcription segments sent to clients",
["completed"],
)
REST_REQUESTS = Counter(
"whisperlive_rest_requests_total",
"Total REST API requests",
["endpoint", "status"],
)
ERRORS = Counter(
"whisperlive_errors_total",
"Total errors by type",
["type"],
)
_AVAILABLE = True
except ImportError:
_AVAILABLE = False
def is_available():
"""Check if prometheus_client is installed."""
return _AVAILABLE
def start_metrics_server(port=9091):
"""Start the Prometheus metrics HTTP server on the given port.
Args:
port (int): Port to serve /metrics on. Default 9091.
"""
if not _AVAILABLE:
logging.warning("prometheus_client not installed; metrics endpoint disabled")
return
try:
start_http_server(port)
logging.info(f"Prometheus metrics available at http://0.0.0.0:{port}/metrics")
except Exception as e:
logging.error(f"Failed to start metrics server: {e}")
def track_connection_opened():
if _AVAILABLE:
CONNECTIONS_TOTAL.inc()
CONNECTIONS_ACTIVE.inc()
def track_connection_closed():
if _AVAILABLE:
CONNECTIONS_ACTIVE.dec()
def track_connection_rejected(reason="full"):
if _AVAILABLE:
CONNECTIONS_REJECTED.labels(reason=reason).inc()
def track_transcription_latency(seconds):
if _AVAILABLE:
TRANSCRIPTION_LATENCY.observe(seconds)
def track_audio_processed(seconds):
if _AVAILABLE:
AUDIO_PROCESSED.inc(seconds)
def track_segment_emitted(completed=True):
if _AVAILABLE:
SEGMENTS_EMITTED.labels(completed=str(completed).lower()).inc()
def track_rest_request(endpoint="/v1/audio/transcriptions", status="200"):
if _AVAILABLE:
REST_REQUESTS.labels(endpoint=endpoint, status=str(status)).inc()
def track_error(error_type="transcription"):
if _AVAILABLE:
ERRORS.labels(type=error_type).inc()
+867 -432
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+364
View File
@@ -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()
+99
View File
@@ -0,0 +1,99 @@
import os
import shutil
import textwrap
import scipy
import numpy as np
import av
from pathlib import Path
def clear_screen():
"""Clears the console screen."""
print("\033[H\033[2J", end="", flush=True)
def print_transcript(text, translated=False, timestamps=False):
"""Prints formatted transcript text in a subtitle-like block."""
terminal_width = shutil.get_terminal_size((80, 20)).columns
wrap_width = max(10, min(80, terminal_width - 8))
if timestamps:
lines = []
for t in text:
prefix = f'[{t["start"]} -> {t["end"]}] '
wrapper = textwrap.TextWrapper(
width=wrap_width,
subsequent_indent=" " * len(prefix),
)
lines.extend(wrapper.wrap(f'{prefix}{t["text"]}'))
else:
wrapper = textwrap.TextWrapper(width=wrap_width)
transcript = " ".join(text) if translated else "".join(text)
lines = wrapper.wrap(text=transcript)
for line in lines[-3:]:
print(line.center(terminal_width))
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
+157
View File
@@ -0,0 +1,157 @@
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
opts.inter_op_num_threads = 1
opts.intra_op_num_threads = 1
if force_onnx_cpu and 'CPUExecutionProvider' in onnxruntime.get_available_providers():
self.session = onnxruntime.InferenceSession(path, providers=['CPUExecutionProvider'], sess_options=opts)
else:
self.session = onnxruntime.InferenceSession(path, providers=['CUDAExecutionProvider'], sess_options=opts)
self.reset_states()
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:
x = x.unsqueeze(0)
if x.dim() > 2:
raise ValueError(f"Too many dimensions for input audio chunk {x.dim()}")
if sr != 16000 and (sr % 16000 == 0):
step = sr // 16000
x = x[:,::step]
sr = 16000
if sr not in self.sample_rates:
raise ValueError(f"Supported sampling rates: {self.sample_rates} (or multiply of 16000)")
if sr / x.shape[1] > 31.25:
raise ValueError("Input audio chunk is too short")
return x, sr
def reset_states(self, batch_size=1):
self._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)
if (self._last_sr) and (self._last_sr != sr):
self.reset_states(batch_size)
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(), 'state': self._state.numpy(), 'sr': np.array(sr, dtype='int64')}
ort_outs = self.session.run(None, ort_inputs)
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.from_numpy(out)
return out
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)
for i in range(0, x.shape[1], num_samples):
wavs_batch = x[:, i:i+num_samples]
out_chunk = self.__call__(wavs_batch, sr)
outs.append(out_chunk)
stacked = torch.cat(outs, dim=1)
return stacked.cpu()
@staticmethod
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
os.makedirs(target_dir, exist_ok=True)
# Define the target file path
model_filename = os.path.join(target_dir, "silero_vad.onnx")
# Check if the model file already exists
if not os.path.exists(model_filename):
# If it doesn't exist, download the model using wget
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()