56 lines
2.7 KiB
Diff
56 lines
2.7 KiB
Diff
diff --git a/examples/openai_server.py b/examples/openai_server.py
|
|
index 61047ea..2f95bbd 100644
|
|
--- a/examples/openai_server.py
|
|
+++ b/examples/openai_server.py
|
|
@@ -169,6 +169,16 @@ def resolve_voice(voice_name: str) -> dict:
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
+def _load_voice_clone_prompt(voice_cfg: dict):
|
|
+ spk_emb_path = voice_cfg.get("speaker_embeddings") or voice_cfg.get("speaker embeddings")
|
|
+ if spk_emb_path and os.path.isfile(spk_emb_path):
|
|
+ try:
|
|
+ return torch.load(spk_emb_path, map_location="cpu", weights_only=False)
|
|
+ except Exception as e:
|
|
+ logger.error("Failed to load speaker embeddings from %s: %s", spk_emb_path, e)
|
|
+ return None
|
|
+
|
|
+
|
|
async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, None]:
|
|
"""
|
|
Run generate_voice_clone_streaming in a background thread and yield
|
|
@@ -183,11 +193,15 @@ async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, No
|
|
for chunk, _sr, _timing in tts_model.generate_voice_clone_streaming(
|
|
text=text,
|
|
language=voice_cfg.get("language", "Auto"),
|
|
- ref_audio=voice_cfg["ref_audio"],
|
|
+ ref_audio=voice_cfg.get("ref_audio"),
|
|
ref_text=voice_cfg.get("ref_text", ""),
|
|
chunk_size=voice_cfg.get("chunk_size", 12),
|
|
instruct=voice_cfg.get("instruct"),
|
|
- non_streaming_mode=False,
|
|
+ voice_clone_prompt=_load_voice_clone_prompt(voice_cfg),
|
|
+ non_streaming_mode=True,
|
|
+ temperature=voice_cfg.get("temperature", 0.8),
|
|
+ top_k=voice_cfg.get("top_k", 50),
|
|
+ top_p=voice_cfg.get("top_p", 0.9),
|
|
):
|
|
q.put(chunk)
|
|
except Exception as exc:
|
|
@@ -249,9 +263,14 @@ async def create_speech(req: SpeechRequest):
|
|
return tts_model.generate_voice_clone(
|
|
text=req.input,
|
|
language=voice_cfg.get("language", "Auto"),
|
|
- ref_audio=voice_cfg["ref_audio"],
|
|
+ ref_audio=voice_cfg.get("ref_audio"),
|
|
ref_text=voice_cfg.get("ref_text", ""),
|
|
instruct=voice_cfg.get("instruct"),
|
|
+ voice_clone_prompt=_load_voice_clone_prompt(voice_cfg),
|
|
+ non_streaming_mode=True,
|
|
+ temperature=voice_cfg.get("temperature", 0.8),
|
|
+ top_k=voice_cfg.get("top_k", 50),
|
|
+ top_p=voice_cfg.get("top_p", 0.9),
|
|
)
|
|
|
|
audio_arrays, sr = await loop.run_in_executor(None, _generate)
|