update with multilingual option
This commit is contained in:
+13
-3
@@ -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
|
||||
)
|
||||
|
||||
+13
-5
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user