Merge pull request #116 from makaveli10/tensorrt_model_warmup

Tensorrt model warmup.
This commit is contained in:
Marcus Edel
2024-01-31 11:49:40 -05:00
committed by GitHub
4 changed files with 17 additions and 5 deletions
+3 -2
View File
@@ -37,10 +37,10 @@ python -c "import torch; import tensorrt; import tensorrt_llm"
- We build `small.en` and `small` multilingual TensorRT engine. The script logs the path of the directory with Whisper TensorRT engine. We need the model_path to run the server. - We build `small.en` and `small` multilingual TensorRT engine. The script logs the path of the directory with Whisper TensorRT engine. We need the model_path to run the server.
```bash ```bash
# convert small.en # convert small.en
bash build_whisper_tensorrt /root/TensorRT-LLM-examples small.en bash scripts/build_whisper_tensorrt.sh /root/TensorRT-LLM-examples small.en
# convert small multilingual model # convert small multilingual model
bash build_whisper_tensorrt /root/TensorRT-LLM-examples small bash scripts/build_whisper_tensorrt.sh /root/TensorRT-LLM-examples small
``` ```
## Run WhisperLive Server with TensorRT Backend ## Run WhisperLive Server with TensorRT Backend
@@ -48,6 +48,7 @@ bash build_whisper_tensorrt /root/TensorRT-LLM-examples small
cd /home/WhisperLive cd /home/WhisperLive
# Install requirements # Install requirements
bash scripts/setup.sh
pip install -r requirements/server.txt pip install -r requirements/server.txt
# Required to create mel spectogram # Required to create mel spectogram
+4
View File
@@ -3,3 +3,7 @@ torch
websockets websockets
onnxruntime==1.16.0 onnxruntime==1.16.0
numba numba
openai-whisper
kaldialign
soundfile
ffmpeg-python
+7
View File
@@ -397,6 +397,7 @@ class ServeClientTensorRT(ServeClientBase):
language=self.language, language=self.language,
task=self.task task=self.task
) )
self.warmup()
# threading # threading
self.trans_thread = threading.Thread(target=self.speech_to_text) self.trans_thread = threading.Thread(target=self.speech_to_text)
@@ -411,6 +412,12 @@ class ServeClientTensorRT(ServeClientBase):
) )
) )
def warmup(self, warmup_steps=10):
logging.info("[INFO:] Warming up TensorRT engine..")
mel, duration = self.transcriber.log_mel_spectrogram("tests/jfk.flac")
for i in range(warmup_steps):
last_segment = self.transcriber.transcribe(mel)
def set_eos(self, eos): def set_eos(self, eos):
self.lock.acquire() self.lock.acquire()
self.eos = eos self.eos = eos
+1 -1
View File
@@ -11,7 +11,7 @@ import numpy as np
from whisper.tokenizer import get_tokenizer from whisper.tokenizer import get_tokenizer
from whisper_live.tensorrt_utils import (mel_filters, store_transcripts, from whisper_live.tensorrt_utils import (mel_filters, store_transcripts,
write_error_stats, load_audio_wav_format, write_error_stats, load_audio_wav_format,
pad_or_trim) pad_or_trim, load_audio)
import tensorrt_llm import tensorrt_llm
import tensorrt_llm.logger as logger import tensorrt_llm.logger as logger