feat: precompute and load speaker embeddings (.pt) from voices.json

This commit is contained in:
mARTin-B78 2026-06-20 13:04:55 +02:00
parent 789398c5df
commit 8f09b7e797
4 changed files with 128 additions and 21 deletions

View File

@ -340,6 +340,14 @@ The first request after container startup can be slower because CUDA graph captu
## Changelog ## Changelog
### v6.3 — 2026-06-20
**Feature: Precomputed Speaker Embeddings (.pt files)**
- Implemented a way to precompute and store speaker embeddings to avoid recalculating the prompt on the server for every single generation.
- `generate_voices.py` now automatically adds `"speaker_embeddings": ""` (or the path to a `.pt` file if it exists) to `voices.json`.
- `openai_server.py` parses `"speaker_embeddings"` and loads the `.pt` file directly into the model's `voice_clone_prompt`, speeding up TTFA.
- Added a `config/extract_embeddings.py` utility script to generate `.pt` files from existing `voices.json` configurations.
### v6.2 — 2026-05-30 ### v6.2 — 2026-05-30
**Fix: voice drift on streaming port 8023 + per-voice temperature persists across restarts** **Fix: voice drift on streaming port 8023 + per-voice temperature persists across restarts**

View File

@ -0,0 +1,81 @@
import os
import json
import torch
import sys
# Append the app directory to import faster_qwen3_tts
sys.path.append("/app/examples")
sys.path.insert(0, "/app")
from faster_qwen3_tts.model import FasterQwen3TTS
def main():
config_file = "/config/voices.json"
if not os.path.exists(config_file):
print(f"Error: {config_file} not found.")
return
with open(config_file, "r") as f:
voices = json.load(f)
# Need to load the model
# Read model path from QWEN_TTS_MODEL or default
model_path = os.environ.get("QWEN_TTS_MODEL", "Qwen/Qwen3-TTS-12Hz-1.7B-Base")
print(f"Loading model {model_path} for embedding extraction...")
tts_model = FasterQwen3TTS.from_pretrained(
model_path,
device="cuda",
dtype=torch.bfloat16,
)
print("Model loaded successfully.")
updates_made = False
for voice_id, entry in voices.items():
spk_emb_path = entry.get("speaker_embeddings") or entry.get("speaker embeddings")
# If there's already a valid path and it exists, skip
if spk_emb_path and os.path.exists(spk_emb_path):
print(f"Skipping {voice_id}, embedding already exists at {spk_emb_path}")
continue
ref_audio = entry.get("ref_audio")
ref_text = entry.get("ref_text", "")
if not ref_audio or not os.path.exists(ref_audio):
print(f"Skipping {voice_id}, ref_audio not found: {ref_audio}")
continue
print(f"Extracting embeddings for {voice_id}...")
try:
prompt_items = tts_model.model.create_voice_clone_prompt(
ref_audio, [ref_text]
)
vcp = tts_model.model._prompt_items_to_voice_clone_prompt(prompt_items)
# Save the .pt file in /config/speakers
pt_path = f"/config/speakers/{voice_id}.pt"
torch.save(vcp, pt_path)
print(f"Saved {pt_path}")
# Update voices.json entry
entry["speaker_embeddings"] = pt_path
# Remove the legacy "speaker embeddings" with space if it exists
if "speaker embeddings" in entry:
del entry["speaker embeddings"]
updates_made = True
except Exception as e:
print(f"Error extracting embedding for {voice_id}: {e}")
if updates_made:
with open(config_file, "w") as f:
json.dump(voices, f, indent=2, ensure_ascii=False)
print("Updated voices.json with new speaker_embeddings paths.")
else:
print("No new embeddings extracted.")
if __name__ == "__main__":
main()

View File

@ -96,6 +96,12 @@ for scan_dir in SCAN_DIRS:
"chunk_size": 4, "chunk_size": 4,
} }
pt_path = os.path.join("/config/speakers", f"{voice_id}.pt")
if os.path.exists(pt_path):
entry["speaker_embeddings"] = pt_path
else:
entry["speaker_embeddings"] = ""
ref_txt = os.path.join(root, f"{base_name}.reference.txt") ref_txt = os.path.join(root, f"{base_name}.reference.txt")
txt = os.path.join(root, f"{base_name}.txt") txt = os.path.join(root, f"{base_name}.txt")
if os.path.exists(ref_txt): if os.path.exists(ref_txt):

View File

@ -1,12 +1,35 @@
diff --git a/examples/openai_server.py b/examples/openai_server.py diff --git a/examples/openai_server.py b/examples/openai_server.py
index 2199e14..cb44644 100644 index 61047ea..2f95bbd 100644
--- a/examples/openai_server.py --- a/examples/openai_server.py
+++ b/examples/openai_server.py +++ b/examples/openai_server.py
@@ -185,7 +185,10 @@ async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, No @@ -169,6 +169,16 @@ def resolve_voice(voice_name: str) -> dict:
ref_audio=voice_cfg["ref_audio"], # ---------------------------------------------------------------------------
+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", ""), ref_text=voice_cfg.get("ref_text", ""),
chunk_size=voice_cfg.get("chunk_size", 12), chunk_size=voice_cfg.get("chunk_size", 12),
instruct=voice_cfg.get("instruct"),
- non_streaming_mode=False, - non_streaming_mode=False,
+ voice_clone_prompt=_load_voice_clone_prompt(voice_cfg),
+ non_streaming_mode=True, + non_streaming_mode=True,
+ temperature=voice_cfg.get("temperature", 0.8), + temperature=voice_cfg.get("temperature", 0.8),
+ top_k=voice_cfg.get("top_k", 50), + top_k=voice_cfg.get("top_k", 50),
@ -14,10 +37,15 @@ index 2199e14..cb44644 100644
): ):
q.put(chunk) q.put(chunk)
except Exception as exc: except Exception as exc:
@@ -249,6 +252,10 @@ async def create_speech(req: SpeechRequest): @@ -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"), language=voice_cfg.get("language", "Auto"),
ref_audio=voice_cfg["ref_audio"], - ref_audio=voice_cfg["ref_audio"],
+ ref_audio=voice_cfg.get("ref_audio"),
ref_text=voice_cfg.get("ref_text", ""), 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, + non_streaming_mode=True,
+ temperature=voice_cfg.get("temperature", 0.8), + temperature=voice_cfg.get("temperature", 0.8),
+ top_k=voice_cfg.get("top_k", 50), + top_k=voice_cfg.get("top_k", 50),
@ -25,19 +53,3 @@ index 2199e14..cb44644 100644
) )
audio_arrays, sr = await loop.run_in_executor(None, _generate) audio_arrays, sr = await loop.run_in_executor(None, _generate)
@@ -306,6 +313,7 @@ def _parse_args():
p.add_argument("--host", default="0.0.0.0", help="Bind host (default: 0.0.0.0)")
p.add_argument("--port", type=int, default=8000, help="Bind port (default: 8000)")
p.add_argument("--device", default="cuda", help="Torch device (default: cuda)")
+ p.add_argument("--max-seq-len", type=int, default=4096, help="Max sequence length for CUDA graph static cache (default: 4096)")
return p.parse_args()
@@ -344,6 +352,7 @@ def main():
args.model,
device=args.device,
dtype=torch.bfloat16,
+ max_seq_len=args.max_seq_len,
)
SAMPLE_RATE = tts_model.sample_rate
logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)