SenseVoice_AgenticRAG / python /SenseVoice_server.py
H022329's picture
Upload folder using huggingface_hub
748b06f verified
Raw
History Blame Contribute Delete
12.6 kB
#!/usr/bin/env python3
"""
OpenAI-compatible ASR server for SenseVoice (AX650 NPU).
Wraps SenseVoiceAx with standard OpenAI /v1/audio/transcriptions endpoint.
Usage:
python openai_server.py --port 8087
# or
uvicorn openai_server:app --host 0.0.0.0 --port 8087
"""
import argparse
import logging
import os
import sys
import time
import uuid
from pathlib import Path
from typing import Optional
import librosa
import numpy as np
# Ensure SenseVoiceAx can be imported
SCRIPT_DIR = Path(__file__).resolve().parent
sys.path.insert(0, str(SCRIPT_DIR))
from SenseVoiceAx import SenseVoiceAx
# ── Logging ────────────────────────────────────────────────────────────
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("openai_asr")
# ── FastAPI ────────────────────────────────────────────────────────────
from fastapi import FastAPI, File, Form, HTTPException, Request, UploadFile
from fastapi.responses import JSONResponse, PlainTextResponse
app = FastAPI(
title="SenseVoice ASR API (OpenAI-compatible)",
version="1.0.0",
description="OpenAI-compatible /v1/audio/transcriptions backed by SenseVoice on AX650 NPU",
)
# ── Globals ────────────────────────────────────────────────────────────
asr_model: Optional["SenseVoiceAx"] = None
_model_config: dict = {}
# ── Model Loading ──────────────────────────────────────────────────────
def build_model(
model_root: str,
max_seq_len: int = 256,
beam_size: int = 3,
streaming: bool = False,
) -> SenseVoiceAx:
"""Create a SenseVoiceAx instance from a model directory."""
model_root = Path(model_root)
if streaming:
model_path = model_root / "streaming_sensevoice.axmodel"
max_seq_len = 26
else:
model_path = model_root / "sensevoice.axmodel"
if not model_path.exists():
# also try the nested path
alt_path = model_root / "sensevoice" / model_path.name
if alt_path.exists():
model_path = alt_path
if not model_path.exists():
raise FileNotFoundError(
f"Model file not found: {model_path}. "
f"Please check model_root={model_root}"
)
cmvn_file = model_root / "am.mvn"
bpe_model = model_root / "chn_jpn_yue_eng_ko_spectok.bpe.model"
token_file = model_root / "tokens.txt"
for f, name in [
(cmvn_file, "am.mvn"),
(token_file, "tokens.txt"),
(bpe_model, "bpe.model"),
]:
if not f.exists():
raise FileNotFoundError(f"Required file {name} not found at {f}")
logger.info(f"Loading SenseVoice from: {model_path}")
logger.info(f" cmvn: {cmvn_file}")
logger.info(f" tokens: {token_file}")
model = SenseVoiceAx(
str(model_path),
str(cmvn_file),
str(token_file),
str(bpe_model),
max_seq_len=max_seq_len,
beam_size=beam_size,
streaming=streaming,
)
logger.info(f"Model loaded. Language options: {model.language_options}")
return model
# ── Lifespan ────────────────────────────────────────────────────────────
from contextlib import asynccontextmanager
@asynccontextmanager
async def lifespan(app: FastAPI):
global asr_model, _model_config
model_root = os.getenv("SENSEVOICE_MODEL_ROOT", "../sensevoice_ax650")
max_seq_len = int(os.getenv("SENSEVOICE_MAX_SEQ_LEN", "256"))
beam_size = int(os.getenv("SENSEVOICE_BEAM_SIZE", "3"))
streaming = os.getenv("SENSEVOICE_STREAMING", "false").lower() == "true"
_model_config = {
"model_root": str(Path(model_root).resolve()),
"max_seq_len": max_seq_len,
"beam_size": beam_size,
"streaming": streaming,
}
print("=" * 60)
print(f"Loading ASR model from: {_model_config['model_root']}")
print("=" * 60)
asr_model = build_model(model_root, max_seq_len, beam_size, streaming)
print("=" * 60)
print("Model ready.")
print("=" * 60)
yield
asr_model = None
app.router.lifespan_context = lifespan
# ── Middleware ──────────────────────────────────────────────────────────
@app.middleware("http")
async def request_logger(request: Request, call_next):
start = time.perf_counter()
response = await call_next(request)
elapsed = time.perf_counter() - start
logger.info(
f'{request.client.host} "{request.method} {request.url.path}" '
f"{response.status_code} ({elapsed:.3f}s)"
)
return response
# ── Health ─────────────────────────────────────────────────────────────
@app.get("/health")
async def health():
return {
"status": "ok" if asr_model is not None else "model_not_loaded",
"model_config": _model_config,
}
# ── Models ─────────────────────────────────────────────────────────────
@app.get("/v1/models")
async def list_models():
return {
"object": "list",
"data": [
{
"id": "sensevoice",
"object": "model",
"created": 0,
"owned_by": "AXERA-TECH",
"type": "audio",
"languages": list(SenseVoiceAx.lid_dict.keys())
if asr_model
else [],
}
],
}
# ── Audio Transcription (OpenAI-compatible) ────────────────────────────
@app.post("/v1/audio/transcriptions")
async def create_transcription(
file: UploadFile = File(...),
model: str = Form("sensevoice"),
language: Optional[str] = Form("auto"),
response_format: str = Form("json"),
):
"""
Transcribe audio to text β€” OpenAI-compatible.
Parameters:
- file: audio file (wav, mp3, flac, m4a, etc.)
- model: model name (ignored, always uses sensevoice)
- language: zh/en/yue/ja/ko/auto
- response_format: json | text | verbose_json
"""
trace_id = f"trc_{uuid.uuid4().hex}"
# ── Validate ───────────────────────────────────────────────────
if asr_model is None:
raise HTTPException(status_code=503, detail="ASR model not loaded")
supported_formats = ("json", "text", "verbose_json")
if response_format not in supported_formats:
return JSONResponse(
{
"error": {
"message": f"response_format '{response_format}' not supported. "
f"Use: {', '.join(supported_formats)}",
"type": "invalid_request_error",
"trace_id": trace_id,
}
},
status_code=400,
)
if response_format in ("srt", "vtt"):
return JSONResponse(
{
"error": {
"message": f"{response_format} not implemented yet",
"type": "invalid_request_error",
"trace_id": trace_id,
}
},
status_code=400,
)
lang = language or "auto"
valid_langs = asr_model.language_options
if lang not in valid_langs:
logger.warning(
f"Language '{lang}' not in {valid_langs}, falling back to 'auto'"
)
lang = "auto"
# ── Read audio ─────────────────────────────────────────────────
try:
audio_bytes = await file.read()
except Exception as e:
raise HTTPException(status_code=400, detail=f"Failed to read file: {e}")
if not audio_bytes:
raise HTTPException(status_code=400, detail="Empty audio file")
# ── Save temp file ─────────────────────────────────────────────
filename = file.filename or "audio.wav"
suffix = Path(filename).suffix or ".wav"
tmp_path = Path(f"/tmp/sensevoice_{uuid.uuid4().hex}{suffix}")
tmp_path.write_bytes(audio_bytes)
try:
# ── Load and resample to 16kHz ─────────────────────────────
waveform, sr = librosa.load(str(tmp_path), sr=16000)
duration_s = len(waveform) / 16000
# ── Inference ──────────────────────────────────────────────
t0 = time.perf_counter()
text = asr_model.infer_waveform(waveform, lang)
elapsed = time.perf_counter() - t0
logger.info(
f"ASR done | lang={lang} | duration={duration_s:.2f}s | "
f"latency={elapsed:.2f}s | text={text[:100]}"
)
# ── Format response ────────────────────────────────────────
if response_format == "text":
return PlainTextResponse(text, headers={"X-Trace-Id": trace_id})
if response_format == "verbose_json":
return JSONResponse(
{
"task": "transcribe",
"language": lang,
"duration": round(duration_s, 3),
"text": text,
"segments": [],
"trace_id": trace_id,
"provider": "sensevoice_ax650",
"model": model,
"processing_ms": round(elapsed * 1000),
}
)
# default: "json"
return JSONResponse({"text": text})
except Exception as e:
logger.exception("ASR inference failed")
return JSONResponse(
{
"error": {
"message": f"ASR inference failed: {e}",
"type": "server_error",
"trace_id": trace_id,
}
},
status_code=500,
)
finally:
tmp_path.unlink(missing_ok=True)
# ── Main ───────────────────────────────────────────────────────────────
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="SenseVoice OpenAI ASR Server")
parser.add_argument(
"--model-root",
default=os.getenv("SENSEVOICE_MODEL_ROOT", "../sensevoice_ax650"),
help="Path to model directory (default: ../sensevoice_ax650)",
)
parser.add_argument(
"--max-seq-len",
type=int,
default=int(os.getenv("SENSEVOICE_MAX_SEQ_LEN", "256")),
help="Max sequence length (default: 256)",
)
parser.add_argument(
"--beam-size",
type=int,
default=int(os.getenv("SENSEVOICE_BEAM_SIZE", "3")),
help="Beam size for decoding (default: 3)",
)
parser.add_argument("--streaming", action="store_true", help="Use streaming model")
parser.add_argument("--host", default="0.0.0.0", help="Listen host")
parser.add_argument("--port", type=int, default=8087, help="Listen port")
args = parser.parse_args()
# Set env vars for lifespan handler
os.environ["SENSEVOICE_MODEL_ROOT"] = args.model_root
os.environ["SENSEVOICE_MAX_SEQ_LEN"] = str(args.max_seq_len)
os.environ["SENSEVOICE_BEAM_SIZE"] = str(args.beam_size)
os.environ["SENSEVOICE_STREAMING"] = "true" if args.streaming else "false"
import uvicorn
uvicorn.run(app, host=args.host, port=args.port)