whisper remote preload

frdel committed Dec 3, 2024 at 19:41 UTC 357909c16a66c0e7ec78a2a993a9a4e54dd67bf9
4 files changed +68 -29
preload.py
+13 -2
@@ -1,7 +1,18 @@
1 -from python.helpers import runtime, whisper
1 +import asyncio
2 +from python.helpers import runtime, whisper, settings
3
4 print("Running preload...")
5 runtime.initialize()
6
7 +
8 +async def preload():
9 + set = settings.get_settings()
10 +
11 + # async tasks to preload
12 + tasks = [whisper.preload(set["stt_model_size"])]
13 +
14 + return asyncio.gather(*tasks, return_exceptions=True)
15 +
16 +
17 # preload transcription model
7 -whisper.preload()
\ No newline at end of file
18 +asyncio.run(preload())
python/api/transcribe.py
+9 -2
@@ -1,10 +1,17 @@
1 from python.helpers.api import ApiHandler
2 from flask import Request, Response
3
4 -from python.helpers import runtime, whisper
4 +from python.helpers import runtime, settings, whisper
5
6 class Transcribe(ApiHandler):
7 async def process(self, input: dict, request: Request) -> dict | Response:
8 audio = input.get("audio")
9 - result = await whisper.transcribe(audio) # type: ignore
9 + ctxid = input.get("ctxid", "")
10 +
11 + context = self.get_context(ctxid)
12 + if await whisper.is_downloading():
13 + context.log.log(type="info", content="Whisper model is currently being downloaded, please wait...")
14 +
15 + set = settings.get_settings()
16 + result = await whisper.transcribe(set["stt_model_size"], audio) # type: ignore
17 return result
python/helpers/settings.py
+5 -4
@@ -1,3 +1,4 @@
1 +import asyncio
2 import json
3 import os
4 import re
@@ -5,12 +6,13 @@ import subprocess
6 from typing import Any, Literal, TypedDict
7
8 import models
8 -from python.helpers import runtime, whisper
9 +from python.helpers import runtime, whisper, defer
10 from . import files, dotenv
11 from models import get_model, ModelProvider, ModelType
12 from langchain_core.language_models.chat_models import BaseChatModel
13 from langchain_core.embeddings import Embeddings
14
15 +
16 class Settings(TypedDict):
17 chat_model_provider: str
18 chat_model_name: str
@@ -743,8 +745,8 @@ def _apply_settings():
745 agent.config = ctx.config
746 agent = agent.get_data(agent.DATA_NAME_SUBORDINATE)
747
746 - # reload whisper model if necessary
747 - whisper.preload()
748 + # reload whisper model if necessary
749 + task = defer.DeferredTask(whisper.preload, _settings["stt_model_size"])
750
751
752 def _env_to_dict(data: str):
@@ -776,7 +778,6 @@ def set_root_password(password: str):
778 raise Exception("root password can only be set in dockerized environments")
779 subprocess.run(f"echo 'root:{password}' | chpasswd", shell=True, check=True)
780 dotenv.save_dotenv_value(dotenv.KEY_ROOT_PASSWORD, password)
779 -
781
782
783 def get_runtime_config(set: Settings):
python/helpers/whisper.py
+41 -21
@@ -1,39 +1,59 @@
1 -# Import the necessary libraries
1 import base64
2 import warnings
3 import whisper
4 import tempfile
5 +import asyncio
6 from python.helpers import runtime, rfc, settings
7
8 -# suppress FutureWarning from torch.load
9 -warnings.filterwarnings('ignore', category=FutureWarning)
8 +# Suppress FutureWarning from torch.load
9 +warnings.filterwarnings("ignore", category=FutureWarning)
10
11 -model = None
12 -model_name = ""
11 +_model = None
12 +_model_name = ""
13 +is_updating_model = False # Tracks whether the model is currently updating
14
14 -def preload():
15 - global model, model_name
16 - set = settings.get_settings()
17 - if not model or model_name != set["stt_model_size"]:
18 - model = whisper.load_model(set["stt_model_size"])
19 - model_name = set["stt_model_size"]
20 - return model
15 +async def preload(model_name:str):
16 + try:
17 + return await runtime.call_development_function(_preload, model_name)
18 + except Exception as e:
19 + if not runtime.is_development():
20 + raise e
21 +
22 +async def _preload(model_name:str):
23 + global _model, _model_name, is_updating_model
24
22 -async def transcribe(audio_bytes_b64: str):
23 - return await runtime.call_development_function(_transcribe, audio_bytes_b64)
25 + while is_updating_model:
26 + await asyncio.sleep(0.1)
27
25 -def _transcribe(audio_bytes_b64: str):
26 - global model
27 - if model is None:
28 - model = preload()
28 + try:
29 + is_updating_model = True
30 + if not _model or _model_name != model_name:
31 + print(f"Loading Whisper model: {model_name}")
32 + _model = whisper.load_model(model_name)
33 + _model_name = model_name
34 + finally:
35 + is_updating_model = False
36
37 +async def is_downloading():
38 + return await runtime.call_development_function(_is_downloading)
39 +
40 +def _is_downloading():
41 + return is_updating_model
42 +
43 +async def transcribe(model_name:str, audio_bytes_b64: str):
44 + return await runtime.call_development_function(_transcribe, model_name, audio_bytes_b64)
45 +
46 +
47 +async def _transcribe(model_name:str, audio_bytes_b64: str):
48 + await _preload(model_name)
49 +
50 # Decode audio bytes if encoded as a base64 string
51 audio_bytes = base64.b64decode(audio_bytes_b64)
52
33 - #create temp audio file
53 + # Create temp audio file
54 with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as audio_file:
55 audio_file.write(audio_bytes)
56
57 # Transcribe the audio file
38 - result = model.transcribe(audio_file.name, fp16=False )
39 - return result
\ No newline at end of file
58 + result = _model.transcribe(audio_file.name, fp16=False) # type: ignore
59 + return result