更新
This commit is contained in:
@@ -0,0 +1,238 @@
|
||||
import argparse
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, File, Form, UploadFile
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
logging.getLogger("matplotlib").setLevel(logging.WARNING)
|
||||
ROOT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.append(os.path.join(ROOT_DIR, "../../.."))
|
||||
sys.path.append(os.path.join(ROOT_DIR, "../../../third_party/Matcha-TTS"))
|
||||
|
||||
from cosyvoice.cli.cosyvoice import AutoModel
|
||||
|
||||
app = FastAPI(title="CosyVoice 3 streaming API")
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
cosyvoice = None
|
||||
model_name = ""
|
||||
fp16_enabled = False
|
||||
speaker_cache_lock = threading.Lock()
|
||||
cancel_events = {}
|
||||
cancel_events_lock = threading.Lock()
|
||||
|
||||
|
||||
def persist_upload(upload: UploadFile) -> str:
|
||||
suffix = os.path.splitext(upload.filename or "")[1] or ".wav"
|
||||
with tempfile.NamedTemporaryFile(prefix="cosyvoice-prompt-", suffix=suffix, delete=False) as target:
|
||||
upload.file.seek(0)
|
||||
shutil.copyfileobj(upload.file, target)
|
||||
return target.name
|
||||
|
||||
|
||||
def speaker_cache_id(prompt_text: str, prompt_wav: str) -> str:
|
||||
digest = hashlib.sha256(prompt_text.encode("utf-8"))
|
||||
with open(prompt_wav, "rb") as source:
|
||||
for chunk in iter(lambda: source.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return "voice-" + digest.hexdigest()[:24]
|
||||
|
||||
|
||||
def ensure_cached_speaker(prompt_text: str, prompt_wav: str) -> str:
|
||||
cache_id = speaker_cache_id(prompt_text, prompt_wav)
|
||||
with speaker_cache_lock:
|
||||
if cache_id not in cosyvoice.frontend.spk2info:
|
||||
cosyvoice.add_zero_shot_spk(prompt_text, prompt_wav, cache_id)
|
||||
logging.info("cached zero-shot speaker %s", cache_id)
|
||||
return cache_id
|
||||
|
||||
|
||||
def register_cancel_event(request_id: str):
|
||||
request_id = request_id.strip()
|
||||
if not request_id:
|
||||
return None
|
||||
|
||||
event = threading.Event()
|
||||
with cancel_events_lock:
|
||||
previous = cancel_events.get(request_id)
|
||||
if previous is not None:
|
||||
previous.set()
|
||||
cancel_events[request_id] = event
|
||||
return event
|
||||
|
||||
|
||||
def release_cancel_event(request_id: str, event):
|
||||
if not request_id or event is None:
|
||||
return
|
||||
with cancel_events_lock:
|
||||
if cancel_events.get(request_id) is event:
|
||||
cancel_events.pop(request_id, None)
|
||||
|
||||
|
||||
def pcm_stream(model_output, cleanup_path: str = "", request_id: str = "", cancel_event=None):
|
||||
iterator = iter(model_output)
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
try:
|
||||
item = next(iterator)
|
||||
except StopIteration:
|
||||
break
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
audio = item["tts_speech"].detach().cpu().numpy()
|
||||
yield (audio * (2**15)).astype(np.int16).tobytes()
|
||||
finally:
|
||||
close = getattr(iterator, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
release_cancel_event(request_id, cancel_event)
|
||||
if cleanup_path:
|
||||
try:
|
||||
os.remove(cleanup_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
def response(model_output, cleanup_path: str = "", request_id: str = ""):
|
||||
request_id = request_id.strip()[:128]
|
||||
cancel_event = register_cancel_event(request_id)
|
||||
return StreamingResponse(
|
||||
pcm_stream(model_output, cleanup_path, request_id, cancel_event),
|
||||
media_type="application/octet-stream",
|
||||
headers={
|
||||
"X-Sample-Rate": str(cosyvoice.sample_rate),
|
||||
"X-Audio-Format": "pcm_s16le",
|
||||
"X-Accel-Buffering": "no",
|
||||
"Cache-Control": "no-store, no-transform",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
def health():
|
||||
return {
|
||||
"status": "ok",
|
||||
"model": model_name,
|
||||
"sample_rate": cosyvoice.sample_rate,
|
||||
"streaming": True,
|
||||
"fp16": fp16_enabled,
|
||||
"cached_speakers": len(cosyvoice.frontend.spk2info),
|
||||
"active_streams": len(cancel_events),
|
||||
}
|
||||
|
||||
|
||||
@app.post("/cancel/{request_id}")
|
||||
def cancel(request_id: str):
|
||||
with cancel_events_lock:
|
||||
event = cancel_events.get(request_id)
|
||||
if event is not None:
|
||||
event.set()
|
||||
return {"cancelled": event is not None, "request_id": request_id}
|
||||
|
||||
|
||||
@app.get("/inference_sft")
|
||||
@app.post("/inference_sft")
|
||||
def inference_sft(tts_text: str = Form(), spk_id: str = Form(), request_id: str = Form("")):
|
||||
return response(cosyvoice.inference_sft(tts_text, spk_id, stream=True), request_id=request_id)
|
||||
|
||||
|
||||
@app.get("/inference_zero_shot")
|
||||
@app.post("/inference_zero_shot")
|
||||
def inference_zero_shot(
|
||||
tts_text: str = Form(),
|
||||
prompt_text: str = Form(),
|
||||
prompt_wav: UploadFile = File(),
|
||||
request_id: str = Form(""),
|
||||
):
|
||||
prompt_path = persist_upload(prompt_wav)
|
||||
cache_id = ensure_cached_speaker(prompt_text, prompt_path)
|
||||
return response(
|
||||
cosyvoice.inference_zero_shot(
|
||||
tts_text,
|
||||
prompt_text,
|
||||
prompt_path,
|
||||
zero_shot_spk_id=cache_id,
|
||||
stream=True,
|
||||
),
|
||||
prompt_path,
|
||||
request_id,
|
||||
)
|
||||
|
||||
|
||||
@app.get("/inference_cross_lingual")
|
||||
@app.post("/inference_cross_lingual")
|
||||
def inference_cross_lingual(
|
||||
tts_text: str = Form(),
|
||||
prompt_wav: UploadFile = File(),
|
||||
request_id: str = Form(""),
|
||||
):
|
||||
prompt_path = persist_upload(prompt_wav)
|
||||
return response(
|
||||
cosyvoice.inference_cross_lingual(tts_text, prompt_path, stream=True),
|
||||
prompt_path,
|
||||
request_id,
|
||||
)
|
||||
|
||||
|
||||
@app.get("/inference_instruct")
|
||||
@app.post("/inference_instruct")
|
||||
def inference_instruct(
|
||||
tts_text: str = Form(),
|
||||
spk_id: str = Form(),
|
||||
instruct_text: str = Form(),
|
||||
request_id: str = Form(""),
|
||||
):
|
||||
return response(
|
||||
cosyvoice.inference_instruct(tts_text, spk_id, instruct_text, stream=True),
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
|
||||
@app.get("/inference_instruct2")
|
||||
@app.post("/inference_instruct2")
|
||||
def inference_instruct2(
|
||||
tts_text: str = Form(),
|
||||
instruct_text: str = Form(),
|
||||
prompt_wav: UploadFile = File(),
|
||||
request_id: str = Form(""),
|
||||
):
|
||||
prompt_path = persist_upload(prompt_wav)
|
||||
return response(
|
||||
cosyvoice.inference_instruct2(tts_text, instruct_text, prompt_path, stream=True),
|
||||
prompt_path,
|
||||
request_id,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--port", type=int, default=50000)
|
||||
parser.add_argument("--model_dir", type=str, default="FunAudioLLM/Fun-CosyVoice3-0.5B-2512")
|
||||
parser.add_argument("--fp16", action="store_true", help="Run the PyTorch model in FP16 on CUDA")
|
||||
args = parser.parse_args()
|
||||
model_name = args.model_dir
|
||||
fp16_enabled = bool(args.fp16)
|
||||
cosyvoice = AutoModel(model_dir=args.model_dir, fp16=fp16_enabled)
|
||||
|
||||
default_prompt_wav = os.path.join(ROOT_DIR, "../../../asset/zero_shot_prompt.wav")
|
||||
default_prompt_text = "You are a helpful assistant.<|endofprompt|>希望你以后能够做的比我还好呦。"
|
||||
if os.path.isfile(default_prompt_wav):
|
||||
ensure_cached_speaker(default_prompt_text, default_prompt_wav)
|
||||
|
||||
uvicorn.run(app, host="0.0.0.0", port=args.port)
|
||||
Reference in New Issue
Block a user