Files
freedakgmail 6116b1f3c6
Update API Documentation / build-api-docs (push) Has been cancelled
Initial commit: FunASR Speech Recognition Toolkit
Add complete FunASR codebase including models, runtime, and documentation.
2026-07-09 22:38:58 +08:00

433 lines
18 KiB
Python

#!/usr/bin/env python3
"""Fun-ASR-Nano vLLM Inference Server.
Unified server with three interfaces:
- HTTP REST: POST /asr (file upload)
- WebSocket: ws://host:port/ws (streaming audio)
- OpenAI API: POST /v1/audio/transcriptions (Whisper-compatible)
All endpoints share the same vLLM engine + dynamic VAD + SPK + timestamps.
Usage:
CUDA_VISIBLE_DEVICES=0 python serve_vllm.py --port 8000
CUDA_VISIBLE_DEVICES=0 python serve_vllm.py --port 8000 --model FunAudioLLM/Fun-ASR-Nano-2512
"""
import asyncio
import argparse
import io
import json
import logging
import os
import re
import time
import tempfile
import numpy as np
import soundfile as sf
import torch
import warnings
warnings.filterwarnings('ignore')
logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s')
logger = logging.getLogger(__name__)
def truncate_repetition(text, min_repeat_len=3, max_repeats=3):
"""Detect and truncate repetitive patterns in ASR output."""
if not text or len(text) < 20:
return text
n = len(text)
for length in range(min_repeat_len, min(n // max_repeats, 30)):
for start in range(n - length * max_repeats):
chunk = text[start:start + length]
if text[start:start + length * max_repeats] == chunk * max_repeats:
return text[:start + length]
return text
try:
from fastapi import FastAPI, File, UploadFile, Form, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
import uvicorn
except ImportError:
raise ImportError("pip install fastapi uvicorn python-multipart")
from funasr.models.fun_asr_nano.inference_vllm import FunASRNanoVLLM
from funasr.models.fsmn_vad_streaming.dynamic_vad import DynamicStreamingVAD
from funasr import AutoModel
# ============================================================
# Global state
# ============================================================
_engine = None
_vad_model = None
_spk_model = None
_args = None
def prepare_audio_for_inference(audio_data, sr, target_sr=16000):
"""Return mono float32 audio at target_sr for ASR inference."""
audio_data = np.asarray(audio_data)
if audio_data.ndim > 1:
channel_axis = -1 if audio_data.shape[-1] <= audio_data.shape[0] else 0
audio_data = audio_data.mean(axis=channel_axis)
if sr != target_sr:
import librosa
audio_data = librosa.resample(audio_data, orig_sr=sr, target_sr=target_sr)
sr = target_sr
return audio_data.astype(np.float32), sr
def load_engine(args):
global _engine, _vad_model, _spk_model, _args
_args = args
if _engine is None:
logger.info(f"Loading vLLM engine: {args.model}")
_engine = FunASRNanoVLLM.from_pretrained(
model=args.model, hub=args.hub, device=args.device, dtype=args.dtype,
max_model_len=args.max_model_len,
gpu_memory_utilization=args.gpu_memory_utilization,
)
logger.info(f"Loading VAD: {args.vad_model}")
_vad_model = AutoModel(model=args.vad_model, device=args.device, disable_update=True)
if args.spk_model:
logger.info(f"Loading SPK: {args.spk_model}")
_spk_model = AutoModel(model=args.spk_model, device=args.device, disable_update=True)
else:
logger.info("SPK disabled")
logger.info("All models ready!")
def process_audio(audio_data, sr=16000, language=None, hotwords=None,
use_vad=True, use_spk=False, use_timestamp=True):
"""Core processing: VAD segment → vLLM ASR → timestamps → SPK."""
audio_data, sr = prepare_audio_for_inference(audio_data, sr)
# VAD segmentation
if use_vad and len(audio_data) > sr * 1:
vad_res = _vad_model.generate(input=audio_data, fs=sr)
segments = vad_res[0]["value"]
else:
segments = [[0, int(len(audio_data) * 1000 / sr)]]
if not segments:
return {"text": "", "segments": [], "duration": len(audio_data) / sr}
# Extract segment audio
seg_audios = []
seg_times = []
for seg in segments:
s0 = int(seg[0] * sr / 1000)
s1 = int(seg[1] * sr / 1000)
seg_audio = audio_data[s0:s1]
if len(seg_audio) > sr * 0.3:
seg_audios.append(seg_audio)
seg_times.append((seg[0], seg[1]))
if not seg_audios:
return {"text": "", "segments": [], "duration": len(audio_data) / sr}
# vLLM batch ASR
gen_kwargs = {"max_new_tokens": 500}
if language:
gen_kwargs["language"] = language
if hotwords:
gen_kwargs["hotwords"] = hotwords
results = _engine.generate(inputs=seg_audios, **gen_kwargs)
# Build segments with timestamps
output_segments = []
full_text_parts = []
for i, (r, (start_ms, end_ms)) in enumerate(zip(results, seg_times)):
r["text"] = truncate_repetition(r["text"])
seg_info = {
"text": r["text"],
"start": start_ms / 1000,
"end": end_ms / 1000,
}
if use_timestamp and "timestamps" in r:
# Offset timestamps by segment start
offset = start_ms / 1000
seg_info["words"] = [
{"word": ts["token"], "start": ts["start_time"] + offset, "end": ts["end_time"] + offset}
for ts in r["timestamps"]
]
output_segments.append(seg_info)
full_text_parts.append(r["text"])
# SPK diarization
if use_spk and _spk_model is not None:
from funasr.models.campplus.utils import sv_chunk, postprocess, distribute_spk
from funasr.models.campplus.cluster_backend import ClusterBackend
vad_segs = [[st, et, audio_data[int(st*sr):int(et*sr)]]
for st, et in [(s["start"], s["end"]) for s in output_segments]]
chunks = sv_chunk(vad_segs)
if chunks:
speech_list = [ch[2] for ch in chunks]
spk_res = _spk_model.generate(input=speech_list, cache={}, is_final=True)
embs = torch.cat([r["spk_embedding"] for r in spk_res], dim=0)
cluster = ClusterBackend(merge_thr=0.78).to(_args.device)
labels = cluster(embs.cpu(), oracle_num=None)
if not isinstance(labels, np.ndarray):
labels = np.array(labels)
all_sorted = sorted(chunks, key=lambda x: x[0])
sv_output = postprocess(all_sorted, None, labels, embs.cpu())
sentences = [{"text": s["text"], "start": int(s["start"]*1000), "end": int(s["end"]*1000)}
for s in output_segments]
distribute_spk(sentences, sv_output)
for i, s in enumerate(sentences):
output_segments[i]["speaker"] = f"SPK{s.get('spk', 0)}"
return {
"text": " ".join(full_text_parts),
"segments": output_segments,
"duration": len(audio_data) / sr,
}
def build_openai_verbose_json(result, language=None):
"""Build OpenAI-compatible verbose_json while preserving FunASR extensions."""
segments = []
for i, seg in enumerate(result["segments"]):
item = {
"id": i,
"start": seg["start"],
"end": seg["end"],
"text": seg["text"],
"words": seg.get("words", []),
}
if "speaker" in seg:
item["speaker"] = seg["speaker"]
segments.append(item)
return {
"task": "transcribe",
"language": language or "zh",
"duration": result["duration"],
"text": result["text"],
"segments": segments,
}
# ============================================================
# FastAPI App
# ============================================================
app = FastAPI(title="Fun-ASR-Nano vLLM Server", version="1.0")
@app.on_event("startup")
async def startup():
load_engine(_args)
# --- HTTP REST: POST /asr ---
@app.post("/asr")
async def asr_endpoint(
file: UploadFile = File(...),
language: str = Form(default=None),
hotwords: str = Form(default=""),
spk: bool = Form(default=False),
timestamp: bool = Form(default=True),
):
"""ASR with file upload. Returns text + segments + timestamps + speaker."""
content = await file.read()
audio_data, sr = sf.read(io.BytesIO(content))
hw_list = [w.strip() for w in hotwords.split(",") if w.strip()] if hotwords else None
t0 = time.perf_counter()
result = process_audio(audio_data, sr=sr, language=language,
hotwords=hw_list, use_spk=spk, use_timestamp=timestamp)
t1 = time.perf_counter()
result["processing_time"] = round(t1 - t0, 3)
result["rtf"] = round((t1 - t0) / result["duration"], 4) if result["duration"] > 0 else 0
return JSONResponse(content=result)
# --- OpenAI API: POST /v1/audio/transcriptions ---
@app.post("/v1/audio/transcriptions")
async def openai_transcriptions(
file: UploadFile = File(...),
model: str = Form(default="fun-asr-nano"),
language: str = Form(default=None),
response_format: str = Form(default="json"),
timestamp_granularities: str = Form(default="word"),
spk: bool = Form(default=False),
):
"""OpenAI Whisper-compatible transcription API (extended with spk support)."""
content = await file.read()
audio_data, sr = sf.read(io.BytesIO(content))
use_ts = "word" in timestamp_granularities or "segment" in timestamp_granularities
result = process_audio(audio_data, sr=sr, language=language, use_spk=spk, use_timestamp=use_ts)
if response_format == "text":
return JSONResponse(content=result["text"])
elif response_format == "verbose_json":
return JSONResponse(content=build_openai_verbose_json(result, language=language))
else:
return JSONResponse(content={"text": result["text"]})
# --- WebSocket: ws://host:port/ws ---
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
"""Streaming WebSocket ASR with dynamic VAD + SPK."""
await websocket.accept()
logger.info(f"WebSocket connected: {websocket.client}")
vad = DynamicStreamingVAD(_vad_model)
audio_buffer = np.array([], dtype=np.float32)
locked_sentences = []
language = None
hotwords = None
use_spk = False
is_active = False
try:
while True:
message = await websocket.receive()
if "text" in message:
cmd = message["text"].strip()
if cmd.upper() == "START":
vad.reset()
audio_buffer = np.array([], dtype=np.float32)
locked_sentences = []
is_active = True
await websocket.send_json({"event": "started"})
elif cmd.upper().startswith("LANGUAGE:"):
language = cmd[9:].strip() or None
await websocket.send_json({"event": "language_set", "language": language})
elif cmd.upper().startswith("HOTWORDS:"):
hotwords = [w.strip() for w in cmd[9:].split(",") if w.strip()]
await websocket.send_json({"event": "hotwords_set", "hotwords": hotwords})
elif cmd.upper().startswith("SPK:"):
use_spk = cmd[4:].strip().lower() in ("true", "1", "on", "yes")
await websocket.send_json({"event": "spk_set", "spk": use_spk})
elif cmd.upper() == "STOP":
if is_active and len(audio_buffer) > 0:
# Final: process remaining audio
final_segs = vad.finalize()
for seg in final_segs:
seg_audio = audio_buffer[int(seg[0]*16):int(seg[1]*16)]
if len(seg_audio) > 8000:
gen_kw = {"max_new_tokens": 500}
if language: gen_kw["language"] = language
if hotwords: gen_kw["hotwords"] = hotwords
res = _engine.generate(inputs=[seg_audio], **gen_kw)
if res[0]["text"].strip():
locked_sentences.append({
"text": res[0]["text"], "start": seg[0], "end": seg[1]
})
# Handle ongoing speech
if vad.is_speaking:
end_ms = int(len(audio_buffer) * 1000 / 16000)
start_ms = int(vad.current_speech_start) if hasattr(vad, 'current_speech_start') and vad.current_speech_start else 0
seg_audio = audio_buffer[int(start_ms*16):]
if len(seg_audio) > 8000:
gen_kw = {"max_new_tokens": 500}
if language: gen_kw["language"] = language
if hotwords: gen_kw["hotwords"] = hotwords
res = _engine.generate(inputs=[seg_audio], **gen_kw)
if res[0]["text"].strip():
locked_sentences.append({
"text": res[0]["text"], "start": start_ms, "end": end_ms
})
# SPK: run full clustering on all sentences (only if enabled)
if use_spk and locked_sentences and _spk_model is not None:
try:
from funasr.models.campplus.utils import sv_chunk, postprocess, distribute_spk
from funasr.models.campplus.cluster_backend import ClusterBackend
vad_segs = [[s["start"]/1000, s["end"]/1000,
audio_buffer[int(s["start"]*16):int(s["end"]*16)]]
for s in locked_sentences]
chunks = sv_chunk(vad_segs)
if chunks:
speech_list = [ch[2] for ch in chunks]
spk_res = _spk_model.generate(input=speech_list, cache={}, is_final=True)
import torch as _torch
embs = _torch.cat([r["spk_embedding"] for r in spk_res], dim=0)
cluster = ClusterBackend(merge_thr=0.78).to(_args.device)
labels = cluster(embs.cpu(), oracle_num=None)
if not isinstance(labels, np.ndarray):
labels = np.array(labels)
all_sorted = sorted(chunks, key=lambda x: x[0])
sv_output = postprocess(all_sorted, None, labels, embs.cpu())
spk_sents = [{"text": s["text"], "start": int(s["start"]), "end": int(s["end"])}
for s in locked_sentences]
distribute_spk(spk_sents, sv_output)
for i, ss in enumerate(spk_sents):
locked_sentences[i]["spk"] = ss.get("spk", 0)
except Exception as e:
logger.warning(f"SPK failed: {e}")
await websocket.send_json({
"sentences": locked_sentences,
"is_final": True,
"duration_ms": int(len(audio_buffer) * 1000 / 16000),
})
is_active = False
await websocket.send_json({"event": "stopped"})
elif "bytes" in message and is_active:
pcm = np.frombuffer(message["bytes"], dtype=np.int16).astype(np.float32) / 32768.0
audio_buffer = np.concatenate([audio_buffer, pcm])
# Feed VAD
new_confirmed = vad.feed(torch.from_numpy(pcm).float())
for seg in new_confirmed:
seg_audio = audio_buffer[int(seg[0]*16):int(seg[1]*16)]
if len(seg_audio) > 8000:
gen_kw = {"max_new_tokens": 500}
if language: gen_kw["language"] = language
if hotwords: gen_kw["hotwords"] = hotwords
res = _engine.generate(inputs=[seg_audio], **gen_kw)
if res[0]["text"].strip():
locked_sentences.append({
"text": res[0]["text"], "start": seg[0], "end": seg[1]
})
# Send partial update
await websocket.send_json({
"sentences": locked_sentences,
"is_final": False,
"duration_ms": int(len(audio_buffer) * 1000 / 16000),
})
except WebSocketDisconnect:
logger.info("WebSocket disconnected")
except Exception as e:
logger.error(f"WebSocket error: {e}", exc_info=True)
# ============================================================
# Main
# ============================================================
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Fun-ASR-Nano vLLM Server")
parser.add_argument("--port", type=int, default=8000)
parser.add_argument("--host", type=str, default="0.0.0.0")
parser.add_argument("--model", type=str, default="FunAudioLLM/Fun-ASR-Nano-2512")
parser.add_argument("--hub", type=str, default="ms")
parser.add_argument("--device", type=str, default="cuda:0")
parser.add_argument("--dtype", type=str, default="bf16")
parser.add_argument("--max-model-len", type=int, default=4096)
parser.add_argument("--gpu-memory-utilization", type=float, default=0.5)
parser.add_argument("--vad-model", type=str, default="fsmn-vad", help="VAD model name or local path")
parser.add_argument("--spk-model", type=str, default="iic/speech_eres2netv2_sv_zh-cn_16k-common", help="Speaker model name or local path (set empty to disable)")
_args = parser.parse_args()
load_engine(_args)
uvicorn.run(app, host=_args.host, port=_args.port)