mirror of
https://github.com/debpalash/VoiceStudio.git
synced 2026-10-02 01:26:35 +08:00
fix(clone): find the installed Whisper the Model Catalogue leaves at a pinned commit (#2442)
snapshot_download(repo, local_files_only=True) only resolves the main ref, which catalogue installs never write, so cloning reported an installed speech-to-text model as missing. The model fallback now also tries every cached snapshot of the configured PyTorch Whisper, large-v3-turbo and large-v3, skipping truncated ones. Adds end-to-end resolution tests (short, long, supplied, no ASR). Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 5.5
parent
245ba3b8f2
commit
939ec45662
@@ -94,11 +94,9 @@ metadata and the backend fallback mirror it.
|
||||
- Voice Clone explains when the selected model cannot clone instead of silently ignoring the reference (#2419)
|
||||
- GitHub Star count refreshes from the repository API, and macOS tray icons keep the intended menu-bar size (#2419)
|
||||
- MLX-Audio OuteTTS generates again with a reference clip or its default voice (#2419)
|
||||
- A reference over 20 s with no speech-to-text model installed now says the clip is too long and to trim it to 3-10 s, instead of asking for a transcript it would refuse (#2442) — thanks @drakeo338!
|
||||
- MLX Qwen3-TTS receives the selected language correctly; MeloTTS explains missing text resources without downloading during generation (#2419)
|
||||
- The language picker offers only the languages each MLX-Audio model supports (Kokoro, CSM, Qwen3-TTS, Dia, Chatterbox, MeloTTS, OuteTTS), per their model cards, instead of every language (#977)
|
||||
|
||||
- The call agent holds back reasoning from chat templates that prefill the opening tag instead of speaking it to the caller (#2428) — thanks @swadhinbiswas!
|
||||
- The Enter that confirms Korean, Japanese or Chinese input no longer also submits project renames, language search, pronunciation, worker or MCP fields (#2338) — thanks @HEOJUNFO!
|
||||
- English text normalization speaks a dollar amount followed by a period or comma ("It costs $5.") instead of leaving the digits (#2390) — thanks @kevin9327!
|
||||
- Voice Design sends descriptions unchanged to free-text engines and restores the original design when reopening a take (#2401) — thanks @CauaMatheus and @dominikj-cf!
|
||||
|
||||
@@ -61,6 +61,10 @@ The env var overrides the persisted UI choice.
|
||||
Whisper Turbo CT2 build; cloning does not require a second Transformers copy.
|
||||
The model-level Whisper fallback
|
||||
also requires cached weights; it never downloads another ASR during cloning.
|
||||
It reuses any installed snapshot of the configured PyTorch Whisper model
|
||||
(`OMNIVOICE_PYTORCH_ASR_MODEL`), `openai/whisper-large-v3-turbo` or
|
||||
`openai/whisper-large-v3`, including the pinned-commit layout Model Catalogue
|
||||
leaves in the Hugging Face cache.
|
||||
If no recognizer is installed, supply a matching reference transcript or
|
||||
explicitly install and select a speech-to-text model in Model Catalogue. A clip with a supplied transcript is limited
|
||||
to 20 seconds so the two stay aligned; trim both to the same passage. Without
|
||||
@@ -71,7 +75,9 @@ The env var overrides the persisted UI choice.
|
||||
whole-clip transcription and ignores a saved profile transcript, so long
|
||||
saved voices use this selection too. OmniVoice's own Whisper snapshot is only
|
||||
a fallback when no catalogue recognizer can transcribe a window, and it is
|
||||
never downloaded during cloning. A transcript typed on the request for
|
||||
never downloaded during cloning. When neither can run, the error says the clip
|
||||
is too long and to trim it to 3–10 seconds, since a transcript cannot rescue it.
|
||||
A transcript typed on the request for
|
||||
such a clip is rejected with `[clone_ref_too_long]`.
|
||||
Only that 15-second window is sent to the model. Clips longer than 75 seconds
|
||||
must be trimmed first. If no spoken words are detected, trim to
|
||||
|
||||
@@ -309,6 +309,104 @@ def _resolve_snapshot_dir(checkpoint) -> str:
|
||||
_DEFAULT_ASR_MODEL = "openai/whisper-large-v3-turbo"
|
||||
|
||||
|
||||
def _reference_asr_repos() -> List[str]:
|
||||
"""Whisper checkpoints the implicit cloning fallback may reuse, best first.
|
||||
|
||||
The configured PyTorch Whisper model leads; ``openai/whisper-large-v3`` is
|
||||
the checkpoint Model Catalogue lists for that engine, so a user who
|
||||
installed it there is not told that no speech-to-text model exists.
|
||||
"""
|
||||
configured = os.environ.get("OMNIVOICE_PYTORCH_ASR_MODEL", "").strip()
|
||||
repos: List[str] = []
|
||||
for repo in (configured, _DEFAULT_ASR_MODEL, "openai/whisper-large-v3"):
|
||||
if repo and repo not in repos:
|
||||
repos.append(repo)
|
||||
return repos
|
||||
|
||||
|
||||
def _hub_cache_roots() -> List[str]:
|
||||
"""Directories that directly contain ``models--*`` folders."""
|
||||
from huggingface_hub import constants
|
||||
|
||||
candidates = [
|
||||
os.environ.get("HF_HUB_CACHE"),
|
||||
os.environ.get("HUGGINGFACE_HUB_CACHE"),
|
||||
getattr(constants, "HF_HUB_CACHE", None),
|
||||
]
|
||||
home = os.environ.get("HF_HOME")
|
||||
if home:
|
||||
candidates += [os.path.join(home, "hub"), home]
|
||||
roots: List[str] = []
|
||||
for root in candidates:
|
||||
if root and root not in roots:
|
||||
roots.append(root)
|
||||
return roots
|
||||
|
||||
|
||||
def _has_asr_weights(snapshot: str) -> bool:
|
||||
try:
|
||||
names = os.listdir(snapshot)
|
||||
except OSError:
|
||||
return False
|
||||
return "config.json" in names and any(
|
||||
name.endswith((".safetensors", ".bin")) for name in names
|
||||
)
|
||||
|
||||
|
||||
def _find_cached_reference_asr() -> Optional[str]:
|
||||
"""Local snapshot directory of an installed Whisper checkpoint, or None.
|
||||
|
||||
``snapshot_download(repo, local_files_only=True)`` only resolves the
|
||||
``main`` ref. Model Catalogue installs every checkpoint at a pinned commit
|
||||
and never writes that ref, so the plain lookup reported an installed model
|
||||
as missing and cloning asked for a speech-to-text model the user already
|
||||
had. Each cache root is therefore also asked for its concrete snapshot
|
||||
revisions. Nothing here touches the network.
|
||||
"""
|
||||
from huggingface_hub import snapshot_download
|
||||
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||
|
||||
configured = os.environ.get("OMNIVOICE_PYTORCH_ASR_MODEL", "").strip()
|
||||
if configured and os.path.isdir(configured) and _has_asr_weights(configured):
|
||||
return configured
|
||||
for repo in _reference_asr_repos():
|
||||
if os.path.isdir(repo):
|
||||
continue
|
||||
try:
|
||||
return snapshot_download(repo, local_files_only=True)
|
||||
except (LocalEntryNotFoundError, ValueError):
|
||||
pass
|
||||
for root in _hub_cache_roots():
|
||||
snapshots = os.path.join(
|
||||
root, "models--" + repo.replace("/", "--"), "snapshots"
|
||||
)
|
||||
try:
|
||||
revisions = sorted(
|
||||
(
|
||||
name
|
||||
for name in os.listdir(snapshots)
|
||||
if os.path.isdir(os.path.join(snapshots, name))
|
||||
),
|
||||
key=lambda name: os.path.getmtime(os.path.join(snapshots, name)),
|
||||
reverse=True,
|
||||
)
|
||||
except OSError:
|
||||
continue
|
||||
for revision in revisions:
|
||||
try:
|
||||
found = snapshot_download(
|
||||
repo,
|
||||
revision=revision,
|
||||
local_files_only=True,
|
||||
cache_dir=root,
|
||||
)
|
||||
except (LocalEntryNotFoundError, ValueError):
|
||||
continue
|
||||
if _has_asr_weights(found):
|
||||
return found
|
||||
return None
|
||||
|
||||
|
||||
class OmniVoice(PreTrainedModel):
|
||||
_supports_flex_attn = True
|
||||
_supports_flash_attn_2 = True
|
||||
@@ -478,17 +576,13 @@ class OmniVoice(PreTrainedModel):
|
||||
|
||||
def _load_cached_reference_asr(self):
|
||||
"""Implicit cloning fallback may reuse local weights, never download them."""
|
||||
from huggingface_hub import snapshot_download
|
||||
from huggingface_hub.errors import LocalEntryNotFoundError
|
||||
|
||||
try:
|
||||
snapshot = snapshot_download(_DEFAULT_ASR_MODEL, local_files_only=True)
|
||||
except LocalEntryNotFoundError as exc:
|
||||
snapshot = _find_cached_reference_asr()
|
||||
if snapshot is None:
|
||||
raise _ReferenceAsrNotInstalled(
|
||||
"Automatic reference transcription needs an installed speech-to-text "
|
||||
"model. Provide a matching reference transcript, or install and select "
|
||||
"a speech-to-text model in Model Catalogue, then try again."
|
||||
) from exc
|
||||
)
|
||||
# Pass a directory rather than the repo ID so transformers cannot make
|
||||
# metadata requests or download missing assets from a partial snapshot.
|
||||
self.load_asr_model(model_name=snapshot)
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
"""#2442: cloning must find the speech-to-text model the user already installed.
|
||||
|
||||
Model Catalogue installs every checkpoint at a pinned commit and never writes
|
||||
the ``main`` ref, so ``snapshot_download(repo, local_files_only=True)`` — the
|
||||
only lookup OmniVoice's implicit-ASR fallback made — reported an installed
|
||||
Whisper as missing. Voice cloning then failed with "needs an installed
|
||||
speech-to-text model" although the Clone page's transcribe button (which asks
|
||||
the selected engine) worked, and a long reference that a supplied transcript
|
||||
cannot align got the same misleading advice.
|
||||
|
||||
One resolution order, simulated end to end here:
|
||||
supplied/stored transcript -> installed engine (transcribe_reference)
|
||||
-> the model's own cached Whisper (any installed snapshot, pinned or not)
|
||||
-> a typed error that names the real cause.
|
||||
"""
|
||||
import importlib
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import soundfile as sf
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("OMNIVOICE_MODEL", "test")
|
||||
os.environ.setdefault("OMNIVOICE_DISABLE_FILE_LOG", "1")
|
||||
|
||||
SR = 24_000
|
||||
PINNED = "06f233fe06e710322aca913c1bc4249a0d71fce1"
|
||||
|
||||
|
||||
def _tts():
|
||||
return importlib.import_module("services.tts_backend")
|
||||
|
||||
|
||||
def _wav(path, seconds, value=0.1):
|
||||
sf.write(path, torch.full((int(seconds * SR),), value).numpy(), SR)
|
||||
return str(path)
|
||||
|
||||
|
||||
class _Tokenizer:
|
||||
config = SimpleNamespace(hop_length=320)
|
||||
device = "cpu"
|
||||
|
||||
def encode(self, audio):
|
||||
return SimpleNamespace(audio_codes=torch.zeros((1, 1, 1), dtype=torch.long))
|
||||
|
||||
|
||||
def _model():
|
||||
from omnivoice.models.omnivoice import OmniVoice
|
||||
|
||||
model = OmniVoice.__new__(OmniVoice)
|
||||
model.sampling_rate = SR
|
||||
model.audio_tokenizer = _Tokenizer()
|
||||
model._asr_pipe = None
|
||||
model.transcribe = lambda _audio: "Words from the cached Whisper."
|
||||
model.loaded = []
|
||||
|
||||
def load(*, model_name):
|
||||
model.loaded.append(model_name)
|
||||
model._asr_pipe = object()
|
||||
|
||||
model.load_asr_model = load
|
||||
return model
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def hf_cache(tmp_path, monkeypatch):
|
||||
"""An empty, isolated Hugging Face cache; nothing from the host leaks in."""
|
||||
from huggingface_hub import constants
|
||||
|
||||
cache = tmp_path / "hf"
|
||||
cache.mkdir()
|
||||
monkeypatch.setattr(constants, "HF_HUB_CACHE", str(cache))
|
||||
for name in ("HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE", "HF_HOME",
|
||||
"OMNIVOICE_PYTORCH_ASR_MODEL"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
return cache
|
||||
|
||||
|
||||
def _install_pinned(cache, repo="openai/whisper-large-v3", complete=True):
|
||||
"""The layout Model Catalogue leaves behind: a snapshot, no refs/main."""
|
||||
snapshot = cache / ("models--" + repo.replace("/", "--")) / "snapshots" / PINNED
|
||||
snapshot.mkdir(parents=True)
|
||||
(snapshot / "config.json").write_text("{}")
|
||||
if complete:
|
||||
(snapshot / "model.safetensors").write_bytes(b"0" * 16)
|
||||
return snapshot
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def no_prompt_disk_cache(monkeypatch):
|
||||
tts = _tts()
|
||||
monkeypatch.setattr(tts, "_prompt_disk_dir", lambda: None)
|
||||
tts.clear_clone_prompt_cache()
|
||||
yield
|
||||
tts.clear_clone_prompt_cache()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("seconds", [4, 25])
|
||||
def test_catalogue_installed_whisper_without_a_main_ref_is_reused(
|
||||
tmp_path, hf_cache, monkeypatch, seconds
|
||||
):
|
||||
"""Short or long, an installed pinned snapshot is found and loaded."""
|
||||
snapshot = _install_pinned(hf_cache)
|
||||
# No recognizer from the registry produced words (e.g. the selected engine
|
||||
# is the PyTorch pipeline, which transcribe_reference leaves to the model).
|
||||
import services.asr_backend as ab
|
||||
|
||||
monkeypatch.setattr(ab, "transcribe_reference", lambda *_a, **_k: None)
|
||||
model = _model()
|
||||
|
||||
prompt = model.create_voice_clone_prompt(
|
||||
_wav(tmp_path / "ref.wav", seconds), None, preprocess_prompt=False
|
||||
)
|
||||
|
||||
assert model.loaded == [str(snapshot)]
|
||||
assert prompt.ref_text.startswith("Words from the cached Whisper")
|
||||
|
||||
|
||||
def test_configured_pytorch_model_is_preferred(tmp_path, hf_cache, monkeypatch):
|
||||
other = _install_pinned(hf_cache, "openai/whisper-large-v3")
|
||||
chosen = _install_pinned(hf_cache, "openai/whisper-small")
|
||||
monkeypatch.setenv("OMNIVOICE_PYTORCH_ASR_MODEL", "openai/whisper-small")
|
||||
model = _model()
|
||||
|
||||
model.create_voice_clone_prompt(
|
||||
_wav(tmp_path / "ref.wav", 3), None, preprocess_prompt=False
|
||||
)
|
||||
|
||||
assert model.loaded == [str(chosen)]
|
||||
assert str(other) not in model.loaded
|
||||
|
||||
|
||||
def test_partial_snapshot_is_not_loaded(tmp_path, hf_cache):
|
||||
"""A config-only snapshot is a truncated download, not an installed model."""
|
||||
_install_pinned(hf_cache, complete=False)
|
||||
model = _model()
|
||||
|
||||
with pytest.raises(ValueError, match="installed speech-to-text model"):
|
||||
model.create_voice_clone_prompt(
|
||||
_wav(tmp_path / "ref.wav", 3), None, preprocess_prompt=False
|
||||
)
|
||||
assert model.loaded == []
|
||||
|
||||
|
||||
def test_short_reference_without_any_asr_asks_for_a_transcript(tmp_path, hf_cache):
|
||||
model = _model()
|
||||
|
||||
with pytest.raises(ValueError) as caught:
|
||||
model.create_voice_clone_prompt(
|
||||
_wav(tmp_path / "ref.wav", 5), None, preprocess_prompt=False
|
||||
)
|
||||
|
||||
message = str(caught.value)
|
||||
assert "installed speech-to-text model" in message
|
||||
assert "reference transcript" in message
|
||||
assert "too long" not in message
|
||||
|
||||
|
||||
def test_long_reference_without_any_asr_names_its_length(tmp_path, hf_cache):
|
||||
"""A transcript cannot rescue a 35 s clip, so the advice is to trim it."""
|
||||
model = _model()
|
||||
|
||||
with pytest.raises(ValueError) as caught:
|
||||
model.create_voice_clone_prompt(
|
||||
_wav(tmp_path / "long.wav", 35), None, preprocess_prompt=False
|
||||
)
|
||||
|
||||
message = str(caught.value)
|
||||
assert "[clone_ref_too_long]" in message
|
||||
assert "35.0" in message and "3-10 second" in message
|
||||
|
||||
|
||||
def test_installed_engine_transcript_needs_no_model_whisper(
|
||||
tmp_path, hf_cache, monkeypatch, no_prompt_disk_cache
|
||||
):
|
||||
"""ASR installed through the catalogue, no Whisper in the model's cache:
|
||||
the engine's transcript is used and the model fallback is never consulted."""
|
||||
import services.asr_backend as ab
|
||||
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
ab, "transcribe_reference", lambda path, **_k: calls.append(path) or "Engine words."
|
||||
)
|
||||
model = _model()
|
||||
model._load_cached_reference_asr = lambda: (_ for _ in ()).throw(
|
||||
AssertionError("model Whisper consulted")
|
||||
)
|
||||
|
||||
prompt = _tts()._get_clone_prompt(model, _wav(tmp_path / "ref.wav", 6), None)
|
||||
|
||||
assert prompt is not None and prompt.ref_text.startswith("Engine words")
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
def test_supplied_transcript_never_looks_for_asr(
|
||||
tmp_path, hf_cache, monkeypatch, no_prompt_disk_cache
|
||||
):
|
||||
import services.asr_backend as ab
|
||||
|
||||
monkeypatch.setattr(
|
||||
ab, "transcribe_reference",
|
||||
lambda *_a, **_k: (_ for _ in ()).throw(AssertionError("ASR not needed")),
|
||||
)
|
||||
model = _model()
|
||||
model._load_cached_reference_asr = lambda: (_ for _ in ()).throw(
|
||||
AssertionError("model Whisper consulted")
|
||||
)
|
||||
|
||||
prompt = _tts()._get_clone_prompt(
|
||||
model, _wav(tmp_path / "ref.wav", 6), "Typed words."
|
||||
)
|
||||
|
||||
assert prompt is not None and prompt.ref_text.startswith("Typed words")
|
||||
|
||||
|
||||
def test_stored_transcript_on_a_long_clip_still_clones_from_catalogue_whisper(
|
||||
tmp_path, hf_cache, monkeypatch, no_prompt_disk_cache
|
||||
):
|
||||
"""The reported flow: a 35 s voice saved with its transcript, Whisper only
|
||||
in the catalogue layout, no engine transcript for the windows."""
|
||||
snapshot = _install_pinned(hf_cache)
|
||||
import services.asr_backend as ab
|
||||
|
||||
monkeypatch.setattr(ab, "transcribe_reference", lambda *_a, **_k: None)
|
||||
model = _model()
|
||||
original = _wav(tmp_path / "long.wav", 35)
|
||||
|
||||
prompt = _tts()._get_clone_prompt(model, original, "The whole clip transcript.")
|
||||
|
||||
assert prompt is not None
|
||||
assert model.loaded == [str(snapshot)]
|
||||
assert prompt.ref_text.startswith("Words from the cached Whisper")
|
||||
@@ -35,7 +35,10 @@ def test_missing_implicit_asr_never_calls_network_capable_loader(monkeypatch, se
|
||||
preprocess_prompt=False,
|
||||
)
|
||||
loader.assert_not_called()
|
||||
lookup.assert_called_once_with("openai/whisper-large-v3-turbo", local_files_only=True)
|
||||
# The default checkpoint is asked for first; the other reusable Whisper
|
||||
# checkpoints are tried before the model concludes nothing is installed.
|
||||
assert lookup.call_args_list[0].args == ("openai/whisper-large-v3-turbo",)
|
||||
assert all(call.kwargs["local_files_only"] is True for call in lookup.call_args_list)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("seconds", [1, 21])
|
||||
|
||||
Reference in New Issue
Block a user