mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-29 15:15:27 -07:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bcab47f25a |
@@ -193,8 +193,6 @@ class MLXTTSBackend:
|
|||||||
|
|
||||||
logger.info("Generating audio for text: %s", text)
|
logger.info("Generating audio for text: %s", text)
|
||||||
|
|
||||||
model_name = f"qwen-tts-{self._current_model_size}"
|
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
"""Run synchronous generation in thread pool."""
|
"""Run synchronous generation in thread pool."""
|
||||||
# MLX generate() returns a generator yielding GenerationResult objects
|
# MLX generate() returns a generator yielding GenerationResult objects
|
||||||
@@ -220,40 +218,38 @@ class MLXTTSBackend:
|
|||||||
logger.warning("Regenerating without voice prompt.")
|
logger.warning("Regenerating without voice prompt.")
|
||||||
ref_audio = None
|
ref_audio = None
|
||||||
|
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
# lazy tokenizer/config lookups inside mlx_audio don't hang
|
# state. Forcing offline here (previously used to avoid lazy
|
||||||
# when the user is disconnected (issue #462).
|
# mlx_audio lookups hanging when the network drops mid-inference,
|
||||||
with force_offline_if_cached(True, model_name):
|
# issue #462) regressed online users because libraries make
|
||||||
# Check if model supports voice cloning via generate method
|
# legitimate metadata calls during generation.
|
||||||
# MLX API may support ref_audio parameter directly
|
try:
|
||||||
try:
|
if ref_audio:
|
||||||
# Try with voice cloning parameters if supported
|
# Check if generate accepts ref_audio parameter
|
||||||
if ref_audio:
|
import inspect
|
||||||
# Check if generate accepts ref_audio parameter
|
|
||||||
import inspect
|
|
||||||
|
|
||||||
sig = inspect.signature(self.model.generate)
|
sig = inspect.signature(self.model.generate)
|
||||||
if "ref_audio" in sig.parameters:
|
if "ref_audio" in sig.parameters:
|
||||||
# Generate with voice cloning
|
# Generate with voice cloning
|
||||||
for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
|
for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
|
||||||
audio_chunks.append(np.array(result.audio))
|
audio_chunks.append(np.array(result.audio))
|
||||||
sample_rate = result.sample_rate
|
sample_rate = result.sample_rate
|
||||||
else:
|
|
||||||
# Fallback: generate without voice cloning
|
|
||||||
for result in self.model.generate(text, lang_code=lang):
|
|
||||||
audio_chunks.append(np.array(result.audio))
|
|
||||||
sample_rate = result.sample_rate
|
|
||||||
else:
|
else:
|
||||||
# No voice prompt, generate normally
|
# Fallback: generate without voice cloning
|
||||||
for result in self.model.generate(text, lang_code=lang):
|
for result in self.model.generate(text, lang_code=lang):
|
||||||
audio_chunks.append(np.array(result.audio))
|
audio_chunks.append(np.array(result.audio))
|
||||||
sample_rate = result.sample_rate
|
sample_rate = result.sample_rate
|
||||||
except Exception as e:
|
else:
|
||||||
# If voice cloning fails, try without it
|
# No voice prompt, generate normally
|
||||||
logger.warning("Voice cloning failed, generating without voice prompt: %s", e)
|
|
||||||
for result in self.model.generate(text, lang_code=lang):
|
for result in self.model.generate(text, lang_code=lang):
|
||||||
audio_chunks.append(np.array(result.audio))
|
audio_chunks.append(np.array(result.audio))
|
||||||
sample_rate = result.sample_rate
|
sample_rate = result.sample_rate
|
||||||
|
except Exception as e:
|
||||||
|
# If voice cloning fails, try without it
|
||||||
|
logger.warning("Voice cloning failed, generating without voice prompt: %s", e)
|
||||||
|
for result in self.model.generate(text, lang_code=lang):
|
||||||
|
audio_chunks.append(np.array(result.audio))
|
||||||
|
sample_rate = result.sample_rate
|
||||||
|
|
||||||
# Concatenate all chunks
|
# Concatenate all chunks
|
||||||
if audio_chunks:
|
if audio_chunks:
|
||||||
@@ -347,8 +343,6 @@ class MLXSTTBackend:
|
|||||||
"""
|
"""
|
||||||
await self.load_model_async(model_size)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
progress_model_name = f"whisper-{self.model_size}"
|
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
# MLX Whisper transcription using generate method
|
# MLX Whisper transcription using generate method
|
||||||
@@ -357,11 +351,10 @@ class MLXSTTBackend:
|
|||||||
if language:
|
if language:
|
||||||
decode_options["language"] = language
|
decode_options["language"] = language
|
||||||
|
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
# lazy tokenizer/config lookups don't hang when the user is
|
# state — see the comment in MLXTTSBackend.generate for the
|
||||||
# disconnected (issue #462).
|
# regression this revert fixes (issue #462).
|
||||||
with force_offline_if_cached(True, progress_model_name):
|
result = self.model.generate(str(audio_path), **decode_options)
|
||||||
result = self.model.generate(str(audio_path), **decode_options)
|
|
||||||
|
|
||||||
# Extract text from result
|
# Extract text from result
|
||||||
if isinstance(result, str):
|
if isinstance(result, str):
|
||||||
|
|||||||
@@ -172,19 +172,17 @@ class PyTorchTTSBackend:
|
|||||||
# This shouldn't happen in practice, but handle it
|
# This shouldn't happen in practice, but handle it
|
||||||
return {"prompt": cached_prompt}, True
|
return {"prompt": cached_prompt}, True
|
||||||
|
|
||||||
model_name = f"qwen-tts-{self._current_model_size}"
|
|
||||||
|
|
||||||
def _create_prompt_sync():
|
def _create_prompt_sync():
|
||||||
"""Run synchronous voice prompt creation in thread pool."""
|
"""Run synchronous voice prompt creation in thread pool."""
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
# state. Forcing offline here (issue #462) regressed online
|
||||||
# when the user is disconnected (issue #462).
|
# users whose libraries issue legitimate metadata lookups
|
||||||
with force_offline_if_cached(True, model_name):
|
# during voice-prompt creation.
|
||||||
return self.model.create_voice_clone_prompt(
|
return self.model.create_voice_clone_prompt(
|
||||||
ref_audio=str(audio_path),
|
ref_audio=str(audio_path),
|
||||||
ref_text=reference_text,
|
ref_text=reference_text,
|
||||||
x_vector_only_mode=False,
|
x_vector_only_mode=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Run blocking operation in thread pool
|
# Run blocking operation in thread pool
|
||||||
voice_prompt_items = await asyncio.to_thread(_create_prompt_sync)
|
voice_prompt_items = await asyncio.to_thread(_create_prompt_sync)
|
||||||
@@ -227,24 +225,20 @@ class PyTorchTTSBackend:
|
|||||||
# Load model
|
# Load model
|
||||||
await self.load_model_async(None)
|
await self.load_model_async(None)
|
||||||
|
|
||||||
model_name = f"qwen-tts-{self._current_model_size}"
|
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
"""Run synchronous generation in thread pool."""
|
"""Run synchronous generation in thread pool."""
|
||||||
# Set seed if provided
|
# Set seed if provided
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
manual_seed(seed, self.device)
|
manual_seed(seed, self.device)
|
||||||
|
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# See _create_prompt_sync comment — inference runs with the
|
||||||
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
# process's default HF_HUB_OFFLINE state (issue #462).
|
||||||
# when the user is disconnected (issue #462).
|
wavs, sample_rate = self.model.generate_voice_clone(
|
||||||
with force_offline_if_cached(True, model_name):
|
text=text,
|
||||||
wavs, sample_rate = self.model.generate_voice_clone(
|
voice_clone_prompt=voice_prompt,
|
||||||
text=text,
|
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
|
||||||
voice_clone_prompt=voice_prompt,
|
instruct=instruct,
|
||||||
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
|
)
|
||||||
instruct=instruct,
|
|
||||||
)
|
|
||||||
return wavs[0], sample_rate
|
return wavs[0], sample_rate
|
||||||
|
|
||||||
# Run blocking inference in thread pool to avoid blocking event loop
|
# Run blocking inference in thread pool to avoid blocking event loop
|
||||||
@@ -342,46 +336,44 @@ class PyTorchSTTBackend:
|
|||||||
"""
|
"""
|
||||||
await self.load_model_async(model_size)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
progress_model_name = f"whisper-{self.model_size}"
|
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
# Load audio
|
# Load audio
|
||||||
audio, _sr = load_audio(audio_path, sample_rate=16000)
|
audio, _sr = load_audio(audio_path, sample_rate=16000)
|
||||||
|
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
# `get_decoder_prompt_ids` and any lazy tokenizer lookups
|
# state — forcing offline here (issue #462) broke online users
|
||||||
# don't hang when the user is disconnected (issue #462).
|
# whose `get_decoder_prompt_ids` / tokenizer calls issue
|
||||||
with force_offline_if_cached(True, progress_model_name):
|
# legitimate metadata lookups.
|
||||||
# Process audio
|
# Process audio
|
||||||
inputs = self.processor(
|
inputs = self.processor(
|
||||||
audio,
|
audio,
|
||||||
sampling_rate=16000,
|
sampling_rate=16000,
|
||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
inputs = inputs.to(self.device)
|
||||||
|
|
||||||
|
# Generate transcription
|
||||||
|
# If language is provided, force it; otherwise let Whisper auto-detect
|
||||||
|
generate_kwargs = {}
|
||||||
|
if language:
|
||||||
|
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
|
||||||
|
language=language,
|
||||||
|
task="transcribe",
|
||||||
)
|
)
|
||||||
inputs = inputs.to(self.device)
|
generate_kwargs["forced_decoder_ids"] = forced_decoder_ids
|
||||||
|
|
||||||
# Generate transcription
|
with torch.no_grad():
|
||||||
# If language is provided, force it; otherwise let Whisper auto-detect
|
predicted_ids = self.model.generate(
|
||||||
generate_kwargs = {}
|
inputs["input_features"],
|
||||||
if language:
|
**generate_kwargs,
|
||||||
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
|
)
|
||||||
language=language,
|
|
||||||
task="transcribe",
|
|
||||||
)
|
|
||||||
generate_kwargs["forced_decoder_ids"] = forced_decoder_ids
|
|
||||||
|
|
||||||
with torch.no_grad():
|
# Decode
|
||||||
predicted_ids = self.model.generate(
|
transcription = self.processor.batch_decode(
|
||||||
inputs["input_features"],
|
predicted_ids,
|
||||||
**generate_kwargs,
|
skip_special_tokens=True,
|
||||||
)
|
)[0]
|
||||||
|
|
||||||
# Decode
|
|
||||||
transcription = self.processor.batch_decode(
|
|
||||||
predicted_ids,
|
|
||||||
skip_special_tokens=True,
|
|
||||||
)[0]
|
|
||||||
|
|
||||||
return transcription.strip()
|
return transcription.strip()
|
||||||
|
|
||||||
|
|||||||
@@ -186,7 +186,6 @@ class QwenCustomVoiceBackend:
|
|||||||
await self.load_model_async(None)
|
await self.load_model_async(None)
|
||||||
|
|
||||||
speaker = voice_prompt.get("preset_voice_id") or QWEN_CV_DEFAULT_SPEAKER
|
speaker = voice_prompt.get("preset_voice_id") or QWEN_CV_DEFAULT_SPEAKER
|
||||||
model_name = f"qwen-custom-voice-{self._current_model_size}"
|
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
@@ -206,11 +205,11 @@ class QwenCustomVoiceBackend:
|
|||||||
if instruct:
|
if instruct:
|
||||||
kwargs["instruct"] = instruct
|
kwargs["instruct"] = instruct
|
||||||
|
|
||||||
# Model is loaded → weights are on disk. Force offline so
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
# state. Forcing offline here (issue #462) regressed online
|
||||||
# when the user is disconnected (issue #462).
|
# users whose libraries issue legitimate metadata lookups
|
||||||
with force_offline_if_cached(True, model_name):
|
# during generation.
|
||||||
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
||||||
return wavs[0], sample_rate
|
return wavs[0], sample_rate
|
||||||
|
|
||||||
audio, sample_rate = await asyncio.to_thread(_generate_sync)
|
audio, sample_rate = await asyncio.to_thread(_generate_sync)
|
||||||
|
|||||||
Reference in New Issue
Block a user