api: add support OpenAI REST transcription api

This commit is contained in:
Aaron Boxer
2026-02-05 10:42:54 -05:00
committed by Aaron Boxer
parent b9ae2af8e6
commit 29ee640409
5 changed files with 218 additions and 7 deletions
+16
View File
@@ -45,6 +45,22 @@ source whisper_env/bin/activate
``` ```
### OpenAI REST interface
#### Server
```bash
python3 run_server.py --port 9090 --backend faster_whisper --max_clients 4 --max_connection_time 600 --enable_rest --cors-origins="http://localhost:8080,http://127.0.0.1:8080"
```
#### Client
```bash
python3 client_openai.py $AUDIO_FILE
```
### Setting up NVIDIA/TensorRT-LLM for TensorRT backend ### Setting up NVIDIA/TensorRT-LLM for TensorRT backend
- Please follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) for setup of [NVIDIA/TensorRT-LLM](https://github.com/NVIDIA/TensorRT-LLM) and for building Whisper-TensorRT engine. - Please follow [TensorRT_whisper readme](https://github.com/collabora/WhisperLive/blob/main/TensorRT_whisper.md) for setup of [NVIDIA/TensorRT-LLM](https://github.com/NVIDIA/TensorRT-LLM) and for building Whisper-TensorRT engine.
+38
View File
@@ -0,0 +1,38 @@
import sys
import requests
if len(sys.argv) < 2:
print("Usage: python transcribe_file.py <path_to_audio_file>")
sys.exit(1)
audio_file = sys.argv[1]
# Configuration
host = "localhost"
port = 8000 # Default REST port; change if you used --rest_port
url = f"http://{host}:{port}/v1/audio/transcriptions"
model = "small" # Or "whisper-1" (mapped to small internally)
language = "en" # Or "hi" for Hindi
response_format = "json" # Options: "json", "text", "verbose_json", "srt", "vtt"
# Prepare the request
files = {"file": open(audio_file, "rb")}
data = {
"model": model,
"language": language,
"response_format": response_format,
# Optional: Add "prompt" for style guidance, "temperature" (0-1), etc.
}
# Send the request
response = requests.post(url, files=files, data=data)
if response.status_code == 200:
if response_format == "json" or response_format == "verbose_json":
result = response.json()
print("Transcript:", result.get("text", "No text found"))
# If you need translation, post-process here (e.g., using another API like Google Translate)
else:
print("Transcript:", response.text)
else:
print("Error:", response.status_code, response.json().get("error", "Unknown error"))
+4
View File
@@ -21,3 +21,7 @@ openvino-genai
openvino-tokenizers openvino-tokenizers
optimum optimum
optimum-intel optimum-intel
fastapi
uvicorn
python-multipart
+27 -1
View File
@@ -1,5 +1,14 @@
import argparse import argparse
import os import os
import threading
import logging
from fastapi import FastAPI
from fastapi import UploadFile, Form
import uvicorn
import tempfile
import shutil
import json
from starlette.responses import PlainTextResponse, JSONResponse
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
@@ -43,6 +52,20 @@ if __name__ == "__main__":
type=str, type=str,
default="~/.cache/whisper-live/", default="~/.cache/whisper-live/",
help='Path to cache the converted ctranslate2 models.') help='Path to cache the converted ctranslate2 models.')
parser.add_argument(
"--rest_port", type=int, default=8000, help="Port for the REST API server."
)
parser.add_argument(
"--enable_rest",
action="store_true",
help="Enable the OpenAI-compatible REST API endpoint.",
)
parser.add_argument(
'--cors-origins',
type=str,
default=None,
help="Comma-separated list of allowed CORS origins (e.g., 'http://localhost:3000,http://example.com'). Defaults to localhost/127.0.0.1 on the WebSocket port."
)
args = parser.parse_args() args = parser.parse_args()
if args.backend == "tensorrt": if args.backend == "tensorrt":
@@ -65,5 +88,8 @@ if __name__ == "__main__":
single_model=not args.no_single_model, single_model=not args.no_single_model,
max_clients=args.max_clients, max_clients=args.max_clients,
max_connection_time=args.max_connection_time, max_connection_time=args.max_connection_time,
cache_path=args.cache_path cache_path=args.cache_path,
rest_port=args.rest_port,
enable_rest=args.enable_rest,
cors_origins=args.cors_origins,
) )
+130 -3
View File
@@ -5,9 +5,18 @@ import queue
import json import json
import functools import functools
import logging import logging
import shutil
import tempfile
from typing import Optional, List
from fastapi import FastAPI, UploadFile, Form
from fastapi.middleware.cors import CORSMiddleware
from starlette.responses import PlainTextResponse, JSONResponse
import uvicorn
from faster_whisper import WhisperModel
import torch
from enum import Enum from enum import Enum
from typing import List, Optional from typing import List, Optional
import numpy as np import numpy as np
from websockets.sync.server import serve from websockets.sync.server import serve
from websockets.exceptions import ConnectionClosed from websockets.exceptions import ConnectionClosed
@@ -403,7 +412,10 @@ class TranscriptionServer:
single_model=False, single_model=False,
max_clients=4, max_clients=4,
max_connection_time=600, max_connection_time=600,
cache_path="~/.cache/whisper-live/"): cache_path="~/.cache/whisper-live/",
rest_port=8000,
enable_rest=False,
cors_origins: Optional[str] = None):
""" """
Run the transcription server. Run the transcription server.
@@ -427,6 +439,122 @@ class TranscriptionServer:
logging.info("Single model mode currently only works with custom models.") logging.info("Single model mode currently only works with custom models.")
if not BackendType.is_valid(backend): if not BackendType.is_valid(backend):
raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}") raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}")
# New OpenAI-compatible REST API (toggleable via enable_rest boolean)
if enable_rest:
app = FastAPI(title="WhisperLive OpenAI-Compatible API")
origins = [o.strip() for o in cors_origins.split(',')] if cors_origins else []
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"], # Allows all methods (GET, POST, etc.)
allow_headers=["*"], # Allows all headers
)
@app.post("/v1/audio/transcriptions")
async def transcribe(
file: UploadFile,
model: str = Form(default="whisper-1"),
language: Optional[str] = Form(default=None),
prompt: Optional[str] = Form(default=None),
response_format: str = Form(default="json"),
temperature: float = Form(default=0.0),
timestamp_granularities: Optional[List[str]] = Form(default=None),
# Stubs for unsupported OpenAI params
chunking_strategy: Optional[str] = Form(default=None),
include: Optional[List[str]] = Form(default=None),
known_speaker_names: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[str]] = Form(default=None),
stream: bool = Form(default=False)
):
if stream:
return JSONResponse({"error": "Streaming not supported in this backend."}, status_code=400)
if chunking_strategy or known_speaker_names or known_speaker_references:
logging.warning("Diarization/chunking params ignored; not supported.")
supported_formats = ["json", "text", "srt", "verbose_json", "vtt"]
if response_format not in supported_formats:
return JSONResponse({"error": f"Unsupported response_format. Supported: {supported_formats}"}, status_code=400)
if model != "whisper-1":
logging.warning(f"Model '{model}' requested; using 'small' as fallback.")
model_name = faster_whisper_custom_model_path or "small"
try:
suffix = os.path.splitext(file.filename)[1] or ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
shutil.copyfileobj(file.file, tmp)
tmp_path = tmp.name
device = "cuda" if torch.cuda.is_available() else "cpu"
compute_type = "float16" if device == "cuda" else "int8"
transcriber = WhisperModel(model_name, device=device, compute_type=compute_type)
segments, info = transcriber.transcribe(
tmp_path,
language=language,
initial_prompt=prompt,
temperature=temperature,
vad_filter=False,
word_timestamps=(timestamp_granularities and "word" in timestamp_granularities)
)
text = " ".join([s.text.strip() for s in segments])
os.unlink(tmp_path)
if response_format == "text":
return PlainTextResponse(text)
elif response_format == "json":
return {"text": text}
elif response_format == "verbose_json":
verbose = {
"task": "transcribe",
"language": info.language,
"duration": info.duration,
"text": text,
"segments": []
}
for seg in segments:
seg_dict = {
"id": seg.id,
"seek": seg.seek,
"start": seg.start,
"end": seg.end,
"text": seg.text.strip(),
"tokens": seg.tokens,
"temperature": seg.temperature,
"avg_logprob": seg.avg_logprob,
"compression_ratio": seg.compression_ratio,
"no_speech_prob": seg.no_speech_prob
}
if timestamp_granularities and "word" in timestamp_granularities:
seg_dict["words"] = [{"word": w.word, "start": w.start, "end": w.end, "probability": w.probability} for w in seg.words]
verbose["segments"].append(seg_dict)
return verbose
elif response_format in ["srt", "vtt"]:
output = []
for i, seg in enumerate(segments, 1):
start = f"{int(seg.start // 3600):02}:{int((seg.start % 3600) // 60):02}:{seg.start % 60:06.3f}"
end = f"{int(seg.end // 3600):02}:{int((seg.end % 3600) // 60):02}:{seg.end % 60:06.3f}"
if response_format == "srt":
output.append(f"{i}\n{start.replace('.', ',')} --> {end.replace('.', ',')}\n{seg.text.strip()}\n")
else: # vtt
output.append(f"{start} --> {end}\n{seg.text.strip()}\n")
return PlainTextResponse("\n".join(output))
except Exception as e:
return JSONResponse({"error": str(e)}, status_code=500)
threading.Thread(
target=uvicorn.run,
args=(app,),
kwargs={"host": "0.0.0.0", "port": rest_port, "log_level": "info"},
daemon=True
).start()
logging.info(f"✅ OpenAI-Compatible API started on http://0.0.0.0:{rest_port}")
# Original WebSocket server (always supported)
with serve( with serve(
functools.partial( functools.partial(
self.recv_audio, self.recv_audio,
@@ -487,4 +615,3 @@ class TranscriptionServer:
if hasattr(client, 'translation_thread') and client.translation_thread: if hasattr(client, 'translation_thread') and client.translation_thread:
client.translation_thread.join(timeout=2.0) client.translation_thread.join(timeout=2.0)
self.client_manager.remove_client(websocket) self.client_manager.remove_client(websocket)