simulstreaming warmup is done for each instance of online, not for the backend
This commit is contained in:
parent
87b9ed6ecd
commit
728e1f1290
4 changed files with 80 additions and 78 deletions
|
|
@ -1,10 +1,10 @@
|
||||||
try:
|
try:
|
||||||
from whisperlivekit.whisper_streaming_custom.whisper_online import backend_factory, warmup_asr
|
from whisperlivekit.whisper_streaming_custom.whisper_online import backend_factory
|
||||||
from whisperlivekit.whisper_streaming_custom.online_asr import VACOnlineASRProcessor, OnlineASRProcessor
|
from whisperlivekit.whisper_streaming_custom.online_asr import VACOnlineASRProcessor, OnlineASRProcessor
|
||||||
except ImportError:
|
except ImportError:
|
||||||
from .whisper_streaming_custom.whisper_online import backend_factory, warmup_asr
|
from .whisper_streaming_custom.whisper_online import backend_factory, warmup_asr
|
||||||
from .whisper_streaming_custom.online_asr import VACOnlineASRProcessor, OnlineASRProcessor
|
from .whisper_streaming_custom.online_asr import VACOnlineASRProcessor, OnlineASRProcessor
|
||||||
|
from whisperlivekit.warmup import warmup_asr, warmup_online
|
||||||
from argparse import Namespace
|
from argparse import Namespace
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|
@ -130,6 +130,7 @@ def online_factory(args, asr, tokenizer, logfile=sys.stderr):
|
||||||
asr,
|
asr,
|
||||||
logfile=logfile,
|
logfile=logfile,
|
||||||
)
|
)
|
||||||
|
warmup_online(online, args.warmup_file)
|
||||||
elif args.vac:
|
elif args.vac:
|
||||||
online = VACOnlineASRProcessor(
|
online = VACOnlineASRProcessor(
|
||||||
args.min_chunk_size,
|
args.min_chunk_size,
|
||||||
|
|
|
||||||
|
|
@ -26,6 +26,7 @@ class SimulStreamingOnlineProcessor:
|
||||||
self,
|
self,
|
||||||
asr,
|
asr,
|
||||||
logfile=sys.stderr,
|
logfile=sys.stderr,
|
||||||
|
warmup_file=None
|
||||||
):
|
):
|
||||||
self.asr = asr
|
self.asr = asr
|
||||||
self.logfile = logfile
|
self.logfile = logfile
|
||||||
|
|
@ -121,6 +122,19 @@ class SimulStreamingOnlineProcessor:
|
||||||
logger.exception(f"SimulStreaming processing error: {e}")
|
logger.exception(f"SimulStreaming processing error: {e}")
|
||||||
return [], self.end
|
return [], self.end
|
||||||
|
|
||||||
|
def warmup(self, audio, init_prompt=""):
|
||||||
|
"""Warmup the SimulStreaming model."""
|
||||||
|
try:
|
||||||
|
if isinstance(audio, np.ndarray):
|
||||||
|
audio = torch.from_numpy(audio).float()
|
||||||
|
self.model.insert_audio(audio)
|
||||||
|
self.model.infer(True)
|
||||||
|
self.model.refresh_segment(complete=True)
|
||||||
|
logger.info("SimulStreaming model warmed up successfully")
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"SimulStreaming warmup failed: {e}")
|
||||||
|
|
||||||
|
|
||||||
class SimulStreamingASR():
|
class SimulStreamingASR():
|
||||||
"""SimulStreaming backend with AlignAtt policy."""
|
"""SimulStreaming backend with AlignAtt policy."""
|
||||||
sep = ""
|
sep = ""
|
||||||
|
|
@ -203,15 +217,3 @@ class SimulStreamingASR():
|
||||||
num_languages=self.model.model.num_languages,
|
num_languages=self.model.model.num_languages,
|
||||||
task="translate"
|
task="translate"
|
||||||
)
|
)
|
||||||
|
|
||||||
# def warmup(self, audio, init_prompt=""):
|
|
||||||
# """Warmup the SimulStreaming model."""
|
|
||||||
# try:
|
|
||||||
# if isinstance(audio, np.ndarray):
|
|
||||||
# audio = torch.from_numpy(audio).float()
|
|
||||||
# self.model.insert_audio(audio)
|
|
||||||
# self.model.infer(True)
|
|
||||||
# self.model.refresh_segment(complete=True)
|
|
||||||
# logger.info("SimulStreaming model warmed up successfully")
|
|
||||||
# except Exception as e:
|
|
||||||
# logger.exception(f"SimulStreaming warmup failed: {e}")
|
|
||||||
|
|
|
||||||
62
whisperlivekit/warmup.py
Normal file
62
whisperlivekit/warmup.py
Normal file
|
|
@ -0,0 +1,62 @@
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def load_file(warmup_file=None, timeout=5):
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import librosa
|
||||||
|
|
||||||
|
if warmup_file is None:
|
||||||
|
# Download JFK sample if not already present
|
||||||
|
jfk_url = "https://github.com/ggerganov/whisper.cpp/raw/master/samples/jfk.wav"
|
||||||
|
temp_dir = tempfile.gettempdir()
|
||||||
|
warmup_file = os.path.join(temp_dir, "whisper_warmup_jfk.wav")
|
||||||
|
|
||||||
|
if not os.path.exists(warmup_file):
|
||||||
|
logger.debug(f"Downloading warmup file from {jfk_url}")
|
||||||
|
print(f"Downloading warmup file from {jfk_url}")
|
||||||
|
import time
|
||||||
|
import urllib.request
|
||||||
|
import urllib.error
|
||||||
|
import socket
|
||||||
|
|
||||||
|
original_timeout = socket.getdefaulttimeout()
|
||||||
|
socket.setdefaulttimeout(timeout)
|
||||||
|
|
||||||
|
start_time = time.time()
|
||||||
|
try:
|
||||||
|
urllib.request.urlretrieve(jfk_url, warmup_file)
|
||||||
|
logger.debug(f"Download successful in {time.time() - start_time:.2f}s")
|
||||||
|
except (urllib.error.URLError, socket.timeout) as e:
|
||||||
|
logger.warning(f"Download failed: {e}. Proceeding without warmup.")
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
socket.setdefaulttimeout(original_timeout)
|
||||||
|
elif not warmup_file:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if not warmup_file or not os.path.exists(warmup_file) or os.path.getsize(warmup_file) == 0:
|
||||||
|
logger.warning(f"Warmup file {warmup_file} invalid or missing.")
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
audio, sr = librosa.load(warmup_file, sr=16000)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to load audio file: {e}")
|
||||||
|
return False
|
||||||
|
return audio
|
||||||
|
|
||||||
|
def warmup_asr(asr, warmup_file=None, timeout=5):
|
||||||
|
"""
|
||||||
|
Warmup the ASR model by transcribing a short audio file.
|
||||||
|
"""
|
||||||
|
audio = load_file(warmup_file=None, timeout=5)
|
||||||
|
asr.warmup(audio)
|
||||||
|
logger.info("ASR model is warmed up")
|
||||||
|
|
||||||
|
def warmup_online(online, warmup_file=None, timeout=5):
|
||||||
|
audio = load_file(warmup_file=None, timeout=5)
|
||||||
|
online.warmup(audio)
|
||||||
|
logger.warning("ASR is warmed up")
|
||||||
|
|
@ -107,67 +107,4 @@ def backend_factory(args):
|
||||||
tokenizer = create_tokenizer(tgt_language)
|
tokenizer = create_tokenizer(tgt_language)
|
||||||
else:
|
else:
|
||||||
tokenizer = None
|
tokenizer = None
|
||||||
return asr, tokenizer
|
return asr, tokenizer
|
||||||
|
|
||||||
def warmup_asr(asr, warmup_file=None, timeout=5):
|
|
||||||
"""
|
|
||||||
Warmup the ASR model by transcribing a short audio file.
|
|
||||||
"""
|
|
||||||
import os
|
|
||||||
import tempfile
|
|
||||||
|
|
||||||
is_simulstreaming = hasattr(asr, 'warmup') and callable(getattr(asr, 'warmup'))
|
|
||||||
|
|
||||||
if warmup_file is None:
|
|
||||||
# Download JFK sample if not already present
|
|
||||||
jfk_url = "https://github.com/ggerganov/whisper.cpp/raw/master/samples/jfk.wav"
|
|
||||||
temp_dir = tempfile.gettempdir()
|
|
||||||
warmup_file = os.path.join(temp_dir, "whisper_warmup_jfk.wav")
|
|
||||||
|
|
||||||
if not os.path.exists(warmup_file):
|
|
||||||
logger.debug(f"Downloading warmup file from {jfk_url}")
|
|
||||||
print(f"Downloading warmup file from {jfk_url}")
|
|
||||||
import time
|
|
||||||
import urllib.request
|
|
||||||
import urllib.error
|
|
||||||
import socket
|
|
||||||
|
|
||||||
original_timeout = socket.getdefaulttimeout()
|
|
||||||
socket.setdefaulttimeout(timeout)
|
|
||||||
|
|
||||||
start_time = time.time()
|
|
||||||
try:
|
|
||||||
urllib.request.urlretrieve(jfk_url, warmup_file)
|
|
||||||
logger.debug(f"Download successful in {time.time() - start_time:.2f}s")
|
|
||||||
except (urllib.error.URLError, socket.timeout) as e:
|
|
||||||
logger.warning(f"Download failed: {e}. Proceeding without warmup.")
|
|
||||||
return False
|
|
||||||
finally:
|
|
||||||
socket.setdefaulttimeout(original_timeout)
|
|
||||||
elif not warmup_file:
|
|
||||||
return False
|
|
||||||
|
|
||||||
if not warmup_file or not os.path.exists(warmup_file) or os.path.getsize(warmup_file) == 0:
|
|
||||||
logger.warning(f"Warmup file {warmup_file} invalid or missing.")
|
|
||||||
return False
|
|
||||||
|
|
||||||
print(f"Warming up {'SimulStreaming' if is_simulstreaming else 'Whisper'} with {warmup_file}")
|
|
||||||
try:
|
|
||||||
import librosa
|
|
||||||
audio, sr = librosa.load(warmup_file, sr=16000)
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Failed to load audio file: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
try:
|
|
||||||
if is_simulstreaming:
|
|
||||||
asr.warmup(audio)
|
|
||||||
else:
|
|
||||||
asr.transcribe(audio)
|
|
||||||
|
|
||||||
logger.info(f"{'SimulStreaming' if is_simulstreaming else 'Whisper'} is warmed up")
|
|
||||||
return True
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Warmup failed: {e}")
|
|
||||||
return False
|
|
||||||
Loading…
Reference in a new issue