#!/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)