mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-12 22:33:54 -04:00
feat(backends): add Whisper-Medusa transcription
Add a dedicated Python gRPC backend for aiola Whisper-Medusa checkpoints, including mono 16 kHz preprocessing, bounded clip validation, CPU/CUDA 12 images, backend gallery metadata, and user documentation. Assisted-by: Codex:gpt-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
a64fc865ba
commit
cec62c3dd3
16 files changed
+499
-1
No files matched your search
@@ -0,0 +1,14 @@
|
||||
.DEFAULT_GOAL := install
|
||||
|
||||
.PHONY: install
|
||||
install:
|
||||
bash install.sh
|
||||
|
||||
.PHONY: clean
|
||||
clean:
|
||||
$(RM) backend_pb2_grpc.py backend_pb2.py
|
||||
rm -rf venv __pycache__
|
||||
|
||||
.PHONY: test
|
||||
test:
|
||||
python3 -m unittest test_unit.py
|
||||
Executable
+167
@@ -0,0 +1,167 @@
|
||||
#!/usr/bin/env python3
|
||||
"""LocalAI gRPC backend for aiola Whisper-Medusa speech recognition."""
|
||||
|
||||
import argparse
|
||||
from concurrent import futures
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import time
|
||||
|
||||
import backend_pb2
|
||||
import backend_pb2_grpc
|
||||
import grpc
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "common"))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "common"))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
|
||||
|
||||
SAMPLE_RATE = 16000
|
||||
MAX_DURATION_SECONDS = 30
|
||||
MAX_WORKERS = int(os.environ.get("PYTHON_GRPC_MAX_WORKERS", "1"))
|
||||
_ONE_DAY_IN_SECONDS = 60 * 60 * 24
|
||||
|
||||
|
||||
def _parse_options(raw_options):
|
||||
options = {}
|
||||
for option in raw_options:
|
||||
if ":" not in option:
|
||||
continue
|
||||
key, value = option.split(":", 1)
|
||||
try:
|
||||
value = int(value)
|
||||
except ValueError:
|
||||
try:
|
||||
value = float(value)
|
||||
except ValueError:
|
||||
pass
|
||||
options[key] = value
|
||||
return options
|
||||
|
||||
|
||||
def _prepare_audio(path, torchaudio):
|
||||
waveform, sample_rate = torchaudio.load(path)
|
||||
if waveform.shape[0] > 1:
|
||||
waveform = waveform.mean(dim=0, keepdim=True)
|
||||
if sample_rate != SAMPLE_RATE:
|
||||
waveform = torchaudio.transforms.Resample(sample_rate, SAMPLE_RATE)(waveform)
|
||||
sample_rate = SAMPLE_RATE
|
||||
return waveform, sample_rate
|
||||
|
||||
|
||||
class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
def __init__(self):
|
||||
self.model = None
|
||||
self.processor = None
|
||||
self.device = None
|
||||
self.options = {}
|
||||
|
||||
def Health(self, request, context):
|
||||
return backend_pb2.Reply(message=b"OK")
|
||||
|
||||
def LoadModel(self, request, context):
|
||||
try:
|
||||
import torch
|
||||
from transformers import WhisperProcessor
|
||||
from whisper_medusa import WhisperMedusaModel
|
||||
|
||||
if request.CUDA and not torch.cuda.is_available():
|
||||
return backend_pb2.Result(success=False, message="CUDA is not available")
|
||||
|
||||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
self.device = torch.device("mps")
|
||||
|
||||
self.options = _parse_options(request.Options)
|
||||
model_path, local_only = resolve_model_reference(
|
||||
request, "aiola/whisper-medusa-linear-libri"
|
||||
)
|
||||
self.model = WhisperMedusaModel.from_pretrained(
|
||||
model_path, local_files_only=local_only
|
||||
).to(self.device)
|
||||
self.model.eval()
|
||||
self.processor = WhisperProcessor.from_pretrained(
|
||||
model_path, local_files_only=local_only
|
||||
)
|
||||
except Exception as err:
|
||||
print(f"Whisper-Medusa model load failed: {err}", file=sys.stderr)
|
||||
return backend_pb2.Result(success=False, message=str(err))
|
||||
|
||||
return backend_pb2.Result(success=True, message="Model loaded successfully")
|
||||
|
||||
def AudioTranscription(self, request, context):
|
||||
if self.model is None or self.processor is None:
|
||||
return backend_pb2.TranscriptResult(segments=[], text="")
|
||||
|
||||
try:
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
waveform, sample_rate = _prepare_audio(request.dst, torchaudio)
|
||||
duration = waveform.shape[-1] / sample_rate
|
||||
if duration > MAX_DURATION_SECONDS:
|
||||
raise ValueError(
|
||||
f"Whisper-Medusa supports audio clips up to {MAX_DURATION_SECONDS} seconds"
|
||||
)
|
||||
|
||||
language = request.language or str(self.options.get("language", "en"))
|
||||
regulation_start = int(self.options.get("regulation_start", 140))
|
||||
regulation_factor = float(self.options.get("regulation_factor", 1.01))
|
||||
features = self.processor(
|
||||
waveform.squeeze(), return_tensors="pt", sampling_rate=sample_rate
|
||||
).input_features.to(self.device)
|
||||
with torch.inference_mode():
|
||||
output = self.model.generate(
|
||||
features,
|
||||
language=language,
|
||||
exponential_decay_length_penalty=(
|
||||
regulation_start,
|
||||
regulation_factor,
|
||||
),
|
||||
)
|
||||
text = self.processor.decode(output[0], skip_special_tokens=True).strip()
|
||||
segment = backend_pb2.TranscriptSegment(
|
||||
id=0,
|
||||
start=0,
|
||||
end=int(duration * 1_000_000_000),
|
||||
text=text,
|
||||
)
|
||||
return backend_pb2.TranscriptResult(segments=[segment], text=text)
|
||||
except Exception as err:
|
||||
print(f"Whisper-Medusa transcription failed: {err}", file=sys.stderr)
|
||||
return backend_pb2.TranscriptResult(segments=[], text="")
|
||||
|
||||
|
||||
def serve(address):
|
||||
server = grpc.server(
|
||||
futures.ThreadPoolExecutor(max_workers=MAX_WORKERS),
|
||||
options=[
|
||||
("grpc.max_send_message_length", 50 * 1024 * 1024),
|
||||
("grpc.max_receive_message_length", 50 * 1024 * 1024),
|
||||
],
|
||||
interceptors=get_auth_interceptors(),
|
||||
)
|
||||
backend_pb2_grpc.add_BackendServicer_to_server(BackendServicer(), server)
|
||||
server.add_insecure_port(address)
|
||||
server.start()
|
||||
print(f"Server started. Listening on: {address}", file=sys.stderr)
|
||||
|
||||
def stop_server(_signal, _frame):
|
||||
server.stop(0)
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, stop_server)
|
||||
signal.signal(signal.SIGTERM, stop_server)
|
||||
try:
|
||||
while True:
|
||||
time.sleep(_ONE_DAY_IN_SECONDS)
|
||||
except KeyboardInterrupt:
|
||||
server.stop(0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run the Whisper-Medusa backend")
|
||||
parser.add_argument("--addr", default="localhost:50051")
|
||||
serve(parser.parse_args().addr)
|
||||
Executable
+12
@@ -0,0 +1,12 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
backend_dir=$(dirname "$0")
|
||||
if [ -d "$backend_dir/common" ]; then
|
||||
source "$backend_dir/common/libbackend.sh"
|
||||
else
|
||||
source "$backend_dir/../common/libbackend.sh"
|
||||
fi
|
||||
|
||||
PYTHON_VERSION="3.11"
|
||||
installRequirements
|
||||
Executable
+3
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
python3 -m grpc_tools.protoc -I../../ --python_out=. --grpc_python_out=. ../../backend.proto
|
||||
@@ -0,0 +1,3 @@
|
||||
--extra-index-url https://download.pytorch.org/whl/cpu
|
||||
torch==2.2.2
|
||||
torchaudio==2.2.2
|
||||
@@ -0,0 +1,3 @@
|
||||
--extra-index-url https://download.pytorch.org/whl/cu121
|
||||
torch==2.2.2
|
||||
torchaudio==2.2.2
|
||||
@@ -0,0 +1,5 @@
|
||||
grpcio==1.71.0
|
||||
protobuf
|
||||
grpcio-tools
|
||||
transformers==4.49.0
|
||||
git+https://github.com/aiola-lab/whisper-medusa.git@19819c37ab15db6e68826e406614a2c86fbb946e
|
||||
Executable
+9
@@ -0,0 +1,9 @@
|
||||
#!/bin/bash
|
||||
backend_dir=$(dirname "$0")
|
||||
if [ -d "$backend_dir/common" ]; then
|
||||
source "$backend_dir/common/libbackend.sh"
|
||||
else
|
||||
source "$backend_dir/../common/libbackend.sh"
|
||||
fi
|
||||
|
||||
startBackend "$@"
|
||||
Executable
+3
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
python3 -m unittest test_unit.py
|
||||
@@ -0,0 +1,89 @@
|
||||
import importlib.util
|
||||
import pathlib
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
|
||||
|
||||
class _Message:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
|
||||
backend_pb2 = types.ModuleType("backend_pb2")
|
||||
for name in ("Reply", "Result", "TranscriptResult", "TranscriptSegment"):
|
||||
setattr(backend_pb2, name, _Message)
|
||||
sys.modules["backend_pb2"] = backend_pb2
|
||||
|
||||
backend_pb2_grpc = types.ModuleType("backend_pb2_grpc")
|
||||
backend_pb2_grpc.BackendServicer = object
|
||||
backend_pb2_grpc.add_BackendServicer_to_server = lambda *args: None
|
||||
sys.modules["backend_pb2_grpc"] = backend_pb2_grpc
|
||||
|
||||
grpc = types.ModuleType("grpc")
|
||||
grpc.server = lambda *args, **kwargs: None
|
||||
sys.modules["grpc"] = grpc
|
||||
|
||||
grpc_auth = types.ModuleType("grpc_auth")
|
||||
grpc_auth.get_auth_interceptors = lambda: []
|
||||
sys.modules["grpc_auth"] = grpc_auth
|
||||
|
||||
model_utils = types.ModuleType("model_utils")
|
||||
model_utils.resolve_model_reference = lambda request, default: (request.Model or default, False)
|
||||
sys.modules["model_utils"] = model_utils
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"whisper_medusa_backend", pathlib.Path(__file__).with_name("backend.py")
|
||||
)
|
||||
backend = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(backend)
|
||||
|
||||
|
||||
class FakeTensor:
|
||||
def __init__(self, channels=1, samples=16000):
|
||||
self.channels = channels
|
||||
self.samples = samples
|
||||
self.shape = (channels, samples)
|
||||
self.mean_calls = []
|
||||
|
||||
def mean(self, dim, keepdim):
|
||||
self.mean_calls.append((dim, keepdim))
|
||||
return FakeTensor(1, self.samples)
|
||||
|
||||
def squeeze(self):
|
||||
return self
|
||||
|
||||
def to(self, device):
|
||||
return self
|
||||
|
||||
|
||||
class FakeTorchaudio:
|
||||
class transforms:
|
||||
class Resample:
|
||||
def __init__(self, source, target):
|
||||
self.source = source
|
||||
self.target = target
|
||||
|
||||
def __call__(self, waveform):
|
||||
return FakeTensor(waveform.channels, waveform.samples * self.target // self.source)
|
||||
|
||||
@staticmethod
|
||||
def load(path):
|
||||
return FakeTensor(2, 8000), 8000
|
||||
|
||||
|
||||
class BackendHelpersTest(unittest.TestCase):
|
||||
def test_parse_options_converts_supported_scalar_types(self):
|
||||
self.assertEqual(
|
||||
backend._parse_options(["regulation_start:120", "regulation_factor:1.05", "ignored"]),
|
||||
{"regulation_start": 120, "regulation_factor": 1.05},
|
||||
)
|
||||
|
||||
def test_prepare_audio_mixes_to_mono_and_resamples_to_16khz(self):
|
||||
waveform, sample_rate = backend._prepare_audio("clip.wav", FakeTorchaudio)
|
||||
self.assertEqual(sample_rate, 16000)
|
||||
self.assertEqual(waveform.shape, (1, 16000))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in new issue
Block a user