Files
VoiceStudio/tests/test_clone_prompt_cache.py
T
Palash Debnath 4a0d18f510 perf(omnivoice): cache voice-clone prompt embeddings (#427) (#473)
Every cloned generation re-encoded the reference audio from scratch — a fixed
per-request latency that compounds on batch / long-form / dataset workloads that
reuse one saved voice across many calls.

The OmniVoice model already exposes the fast path (create_voice_clone_prompt →
VoiceClonePrompt, generate(voice_clone_prompt=)); the Studio backend just wasn't
using it. OmniVoiceBackend.generate now:
- builds a VoiceClonePrompt once per reference and caches it (bounded LRU, max 8,
  keyed by ref path + mtime + ref_text; thread-safe — generation runs in a GPU
  thread pool), then passes voice_clone_prompt= to skip the re-encode;
- falls back to the inline ref_audio/ref_text path on ANY cache miss or error,
  so output is identical either way (the model documents the two as equivalent)
  — this is purely a latency optimization, never a behaviour change;
- the design/instruct path (no ref_audio) is untouched.
- unload() clears the cache so a flush / engine-switch frees the prompt tensors.

tests/test_clone_prompt_cache.py: 6 cases (encode-once-then-hit, ref_text +
mtime invalidation, LRU eviction at the cap, encode-failure → None fallback,
clear). 6 passed.

Closes #427.
2026-06-14 22:08:35 +05:30

89 lines
2.5 KiB
Python

"""Voice-clone prompt cache (#427) — the bounded-LRU reference-encode cache.
Pure cache logic: the model is a stub whose create_voice_clone_prompt counts
calls, so we assert the reference is encoded ONCE per (path, mtime, ref_text)
and that misses/errors fall back cleanly. (No real model / torch math here.)
"""
from __future__ import annotations
import pytest
from services import tts_backend as tb
class _StubModel:
def __init__(self, *, fail=False):
self.calls = 0
self.fail = fail
def create_voice_clone_prompt(self, ref_audio, ref_text=None):
self.calls += 1
if self.fail:
raise RuntimeError("encode boom")
return f"PROMPT::{ref_audio}::{ref_text}"
@pytest.fixture(autouse=True)
def _clear_cache():
tb.clear_clone_prompt_cache()
yield
tb.clear_clone_prompt_cache()
def _wav(tmp_path, name="ref.wav", data=b"\x00" * 100):
p = tmp_path / name
p.write_bytes(data)
return str(p)
def test_encodes_once_then_hits_cache(tmp_path):
m = _StubModel()
ref = _wav(tmp_path)
first = tb._get_clone_prompt(m, ref, "hello")
second = tb._get_clone_prompt(m, ref, "hello")
assert first == second
assert m.calls == 1 # second call was a cache hit — no re-encode
def test_different_ref_text_re_encodes(tmp_path):
m = _StubModel()
ref = _wav(tmp_path)
tb._get_clone_prompt(m, ref, "hello")
tb._get_clone_prompt(m, ref, "different")
assert m.calls == 2
def test_mtime_change_invalidates(tmp_path):
m = _StubModel()
ref = _wav(tmp_path)
tb._get_clone_prompt(m, ref, "hi")
# Rewrite with a different mtime → key changes → re-encode.
import os
os.utime(ref, (1, 1))
tb._get_clone_prompt(m, ref, "hi")
assert m.calls == 2
def test_lru_eviction_bounds_cache(tmp_path):
m = _StubModel()
# Fill past the cap with distinct refs.
for i in range(tb._PROMPT_CACHE_MAX + 3):
tb._get_clone_prompt(m, _wav(tmp_path, f"r{i}.wav"), "t")
assert len(tb._prompt_cache) == tb._PROMPT_CACHE_MAX
assert m.calls == tb._PROMPT_CACHE_MAX + 3
def test_encode_failure_returns_none_and_does_not_cache(tmp_path):
m = _StubModel(fail=True)
ref = _wav(tmp_path)
assert tb._get_clone_prompt(m, ref, "x") is None # caller falls back to inline ref
assert len(tb._prompt_cache) == 0
def test_clear_empties_cache(tmp_path):
m = _StubModel()
tb._get_clone_prompt(m, _wav(tmp_path), "x")
assert len(tb._prompt_cache) == 1
tb.clear_clone_prompt_cache()
assert len(tb._prompt_cache) == 0