diff --git a/README.md b/README.md index a906db4..36ad28d 100644 --- a/README.md +++ b/README.md @@ -44,10 +44,11 @@ The server supports 3 backends `faster_whisper`, `tensorrt` and `openvino`. If r python3 run_server.py --port 9090 \ --backend faster_whisper -# running with custom model +# running with custom model and cache_dir to save auto-converted ctranslate2 models python3 run_server.py --port 9090 \ --backend faster_whisper \ -fw "/path/to/custom/faster/whisper/model" + -c ~/.cache/whisper-live/ ``` - 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. diff --git a/run_server.py b/run_server.py index db093ef..980b639 100644 --- a/run_server.py +++ b/run_server.py @@ -31,6 +31,10 @@ if __name__ == "__main__": 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('--cache_path', '-c', + type=str, + default="~/.cache/whisper-live/", + help='Path to cache the converted ctranslate2 models.') args = parser.parse_args() if args.backend == "tensorrt": @@ -51,4 +55,5 @@ if __name__ == "__main__": trt_multilingual=args.trt_multilingual, trt_py_session=args.trt_py_session, single_model=not args.no_single_model, + cache_path=args.cache_path ) diff --git a/whisper_live/backend/faster_whisper_backend.py b/whisper_live/backend/faster_whisper_backend.py index c9e01d1..e2384c8 100644 --- a/whisper_live/backend/faster_whisper_backend.py +++ b/whisper_live/backend/faster_whisper_backend.py @@ -31,6 +31,7 @@ class ServeClientFasterWhisper(ServeClientBase): no_speech_thresh=0.45, clip_audio=False, same_output_threshold=10, + cache_path="~/.cache/whisper-live/" ): """ Initialize a ServeClient instance. @@ -61,6 +62,7 @@ class ServeClientFasterWhisper(ServeClientBase): clip_audio, same_output_threshold, ) + 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", @@ -140,7 +142,7 @@ class ServeClientFasterWhisper(ServeClientBase): if ctranslate2.contains_model(local_snapshot): model_to_load = local_snapshot else: - cache_root = os.path.expanduser("~/.cache/whisper-live/whisper-ct2-models/") + 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) diff --git a/whisper_live/server.py b/whisper_live/server.py index e3a6a89..25ce403 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -233,6 +233,7 @@ class TranscriptionServer: no_speech_thresh=options.get("no_speech_thresh", 0.45), clip_audio=options.get("clip_audio", False), same_output_threshold=options.get("same_output_threshold", 10), + cache_path=self.cache_path, ) logging.info("Running faster_whisper backend.") @@ -369,7 +370,8 @@ class TranscriptionServer: whisper_tensorrt_path=None, trt_multilingual=False, trt_py_session=False, - single_model=False): + single_model=False, + cache_path="~/.cache/whisper-live/"): """ Run the transcription server. @@ -377,6 +379,7 @@ class TranscriptionServer: host (str): The host address to bind the server. port (int): The port number to bind the server. """ + self.cache_path = cache_path if faster_whisper_custom_model_path is not None and not os.path.exists(faster_whisper_custom_model_path): raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path.") if whisper_tensorrt_path is not None and not os.path.exists(whisper_tensorrt_path):