From 647c576e6a666c0d6f96ad6ee88f116a3af2d0bd Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Thu, 11 Jan 2024 08:17:56 +0000 Subject: [PATCH] update with multilingual option --- run_server.py | 16 +++++++++++++--- whisper_live/server.py | 18 +++++++++++++----- whisper_live/transcriber_tensorrt.py | 11 +++++++---- 3 files changed, 33 insertions(+), 12 deletions(-) diff --git a/run_server.py b/run_server.py index 6353308..6ba1c66 100644 --- a/run_server.py +++ b/run_server.py @@ -3,11 +3,15 @@ from whisper_live.server import TranscriptionServer if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument('--backend', type=str, default='tensorrt', help='Backends from ["tensorrt", "faster_whisper"]') + parser.add_argument('--port', type=int, default=9090, help="Websocket port to run the server on.") + parser.add_argument('--backend', type=str, default='faster_whisper', help='Backends from ["tensorrt", "faster_whisper"]') parser.add_argument('--whisper_tensorrt_path', type=str, - default="/root/TensorRT-LLM/examples/whisper/whisper_small_en", + default=None, help='Whisper TensorRT model path') + parser.add_argument('--trt_multilingual', + action="store_true", + help='Boolean only for TensorRT model. True if multilingual.') args = parser.parse_args() if args.backend == "tensorrt": @@ -15,4 +19,10 @@ if __name__ == "__main__": raise ValueError("Please Provide a valid tensorrt model path") server = TranscriptionServer() - server.run("0.0.0.0", port=6006, backend=args.backend, whisper_tensorrt_path=args.whisper_tensorrt_path) + server.run( + "0.0.0.0", + port=6006, + backend=args.backend, + whisper_tensorrt_path=args.whisper_tensorrt_path, + multilingual=args.trt_multilingual + ) diff --git a/whisper_live/server.py b/whisper_live/server.py index 5561316..5b53f1a 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -68,7 +68,7 @@ class TranscriptionServer: return wait_time / 60 - def recv_audio(self, websocket, backend="tensorrt", whisper_tensorrt_path=None): + def recv_audio(self, websocket, backend="tensorrt", whisper_tensorrt_path=None, multilingual=False): """ Receive audio chunks from a client in an infinite loop. @@ -118,7 +118,7 @@ class TranscriptionServer: self.backend = "tensorrt" client = ServeClientTensorRT( websocket, - multilingual=options["multilingual"], + multilingual=multilingual, language=options["language"], task=options["task"], client_uid=options["uid"], @@ -200,7 +200,7 @@ class TranscriptionServer: del websocket break - def run(self, host, port=9090, backend="tensorrt", whisper_tensorrt_path=None): + def run(self, host, port=9090, backend="tensorrt", whisper_tensorrt_path=None, multilingual=False): """ Run the transcription server. @@ -212,7 +212,8 @@ class TranscriptionServer: functools.partial( self.recv_audio, backend=backend, - whisper_tensorrt_path=whisper_tensorrt_path + whisper_tensorrt_path=whisper_tensorrt_path, + multilingual=multilingual ), host, port @@ -369,7 +370,14 @@ class ServeClientTensorRT(ServeClientBase): self.language = language if multilingual else "en" self.task = task self.eos = False - self.transcriber = WhisperTRTLLM(model_path, False, "assets", device="cuda") + self.transcriber = WhisperTRTLLM( + model_path, + assets_dir="assets", + device="cuda", + is_multilingual=multilingual, + language=self.language, + task=self.task + ) # threading self.trans_thread = threading.Thread(target=self.speech_to_text) diff --git a/whisper_live/transcriber_tensorrt.py b/whisper_live/transcriber_tensorrt.py index dd035ef..8634a8f 100644 --- a/whisper_live/transcriber_tensorrt.py +++ b/whisper_live/transcriber_tensorrt.py @@ -181,7 +181,10 @@ class WhisperTRTLLM(object): engine_dir, debug_mode=False, assets_dir=None, - device=None + device=None, + is_multilingual=False, + language="en", + task="transcribe" ): world_size = 1 runtime_rank = tensorrt_llm.mpi_rank() @@ -198,10 +201,10 @@ class WhisperTRTLLM(object): # tokenizer_dir=assets_dir) self.device = device self.tokenizer = get_tokenizer( - False, + is_multilingual, num_languages=self.encoder.num_languages, - language="en", - task="transcribe", + language=language, + task=task, ) self.filters = mel_filters(self.device, self.encoder.n_mels, assets_dir)