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