mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
fix(mlx): release the model binding inside the sync bodies so a failing op's traceback cannot pin it
Review follow-up: the post-op drain ran in a finally, but on the error path the exception's traceback kept the _generate_sync/_transcribe_sync frame (and its model local) alive, so mx.clear_cache() had nothing to return. Each sync body now drops its model (and tokenizer) binding in its own finally before the exception propagates.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
fb67ad436c
commit
63899fd865
@@ -241,70 +241,76 @@ class MLXTTSBackend:
|
||||
# thread mid-generation cannot turn a later self.model read into
|
||||
# None (the fallback path below runs seconds into a request).
|
||||
model = self.model
|
||||
# MLX generate() returns a generator yielding GenerationResult objects
|
||||
audio_chunks = []
|
||||
sample_rate = 24000
|
||||
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
|
||||
|
||||
# Set seed if provided (MLX uses numpy random)
|
||||
if seed is not None:
|
||||
import mlx.core as mx
|
||||
|
||||
np.random.seed(seed)
|
||||
mx.random.seed(seed)
|
||||
|
||||
# Extract voice prompt info
|
||||
ref_audio = voice_prompt.get("ref_audio") or voice_prompt.get("ref_audio_path")
|
||||
ref_text = voice_prompt.get("ref_text", "")
|
||||
|
||||
# Validate that the audio file exists
|
||||
if ref_audio and not Path(ref_audio).exists():
|
||||
logger.warning("Audio file not found: %s", ref_audio)
|
||||
logger.warning("This may be due to a cached voice prompt referencing a deleted temp file.")
|
||||
logger.warning("Regenerating without voice prompt.")
|
||||
ref_audio = None
|
||||
|
||||
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||
# state. Forcing offline here (previously used to avoid lazy
|
||||
# mlx_audio lookups hanging when the network drops mid-inference,
|
||||
# issue #462) regressed online users because libraries make
|
||||
# legitimate metadata calls during generation.
|
||||
try:
|
||||
if ref_audio:
|
||||
# Check if generate accepts ref_audio parameter
|
||||
import inspect
|
||||
# MLX generate() returns a generator yielding GenerationResult objects
|
||||
audio_chunks = []
|
||||
sample_rate = 24000
|
||||
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
|
||||
|
||||
sig = inspect.signature(model.generate)
|
||||
if "ref_audio" in sig.parameters:
|
||||
# Generate with voice cloning
|
||||
for result in model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
sample_rate = result.sample_rate
|
||||
# Set seed if provided (MLX uses numpy random)
|
||||
if seed is not None:
|
||||
import mlx.core as mx
|
||||
|
||||
np.random.seed(seed)
|
||||
mx.random.seed(seed)
|
||||
|
||||
# Extract voice prompt info
|
||||
ref_audio = voice_prompt.get("ref_audio") or voice_prompt.get("ref_audio_path")
|
||||
ref_text = voice_prompt.get("ref_text", "")
|
||||
|
||||
# Validate that the audio file exists
|
||||
if ref_audio and not Path(ref_audio).exists():
|
||||
logger.warning("Audio file not found: %s", ref_audio)
|
||||
logger.warning("This may be due to a cached voice prompt referencing a deleted temp file.")
|
||||
logger.warning("Regenerating without voice prompt.")
|
||||
ref_audio = None
|
||||
|
||||
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||
# state. Forcing offline here (previously used to avoid lazy
|
||||
# mlx_audio lookups hanging when the network drops mid-inference,
|
||||
# issue #462) regressed online users because libraries make
|
||||
# legitimate metadata calls during generation.
|
||||
try:
|
||||
if ref_audio:
|
||||
# Check if generate accepts ref_audio parameter
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(model.generate)
|
||||
if "ref_audio" in sig.parameters:
|
||||
# Generate with voice cloning
|
||||
for result in model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
sample_rate = result.sample_rate
|
||||
else:
|
||||
# Fallback: generate without voice cloning
|
||||
for result in model.generate(text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
sample_rate = result.sample_rate
|
||||
else:
|
||||
# Fallback: generate without voice cloning
|
||||
# No voice prompt, generate normally
|
||||
for result in model.generate(text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
sample_rate = result.sample_rate
|
||||
else:
|
||||
# No voice prompt, generate normally
|
||||
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 model.generate(text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
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 model.generate(text, lang_code=lang):
|
||||
audio_chunks.append(np.array(result.audio))
|
||||
sample_rate = result.sample_rate
|
||||
|
||||
# Concatenate all chunks
|
||||
if audio_chunks:
|
||||
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
|
||||
else:
|
||||
# Fallback: empty audio
|
||||
audio = np.array([], dtype=np.float32)
|
||||
# Concatenate all chunks
|
||||
if audio_chunks:
|
||||
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
|
||||
else:
|
||||
# Fallback: empty audio
|
||||
audio = np.array([], dtype=np.float32)
|
||||
|
||||
return audio, sample_rate
|
||||
return audio, sample_rate
|
||||
finally:
|
||||
# Drop the local binding here, inside the frame a propagating
|
||||
# traceback would keep alive, so the caller's drain can
|
||||
# actually return the model's buffers to MLX.
|
||||
del model
|
||||
|
||||
def _reload_and_generate_sync():
|
||||
"""Ensure the configured model is loaded, then generate — as ONE
|
||||
@@ -324,9 +330,10 @@ class MLXTTSBackend:
|
||||
finally:
|
||||
# An unload_model() that landed while we were generating only
|
||||
# dropped the backend's reference; the model's buffers were
|
||||
# kept alive by _generate_sync's local and have just been
|
||||
# returned to MLX's pool (on success or failure). Drain it now,
|
||||
# or they stay resident until the next load/unload cycle.
|
||||
# kept alive by _generate_sync's local, which that function
|
||||
# drops in its own finally (so a propagating traceback cannot
|
||||
# pin it), and have just been returned to MLX's pool. Drain it
|
||||
# now, or they stay resident until the next load/unload cycle.
|
||||
if self.model is None:
|
||||
empty_mlx_cache()
|
||||
|
||||
@@ -423,24 +430,30 @@ class MLXSTTBackend:
|
||||
# MLX Whisper transcription using generate method
|
||||
# The generate method accepts audio path directly
|
||||
model = self.model
|
||||
decode_options = {}
|
||||
if language:
|
||||
decode_options["language"] = language
|
||||
try:
|
||||
decode_options = {}
|
||||
if language:
|
||||
decode_options["language"] = language
|
||||
|
||||
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||
# state — see the comment in MLXTTSBackend.generate for the
|
||||
# regression this revert fixes (issue #462).
|
||||
result = model.generate(str(audio_path), **decode_options)
|
||||
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||
# state — see the comment in MLXTTSBackend.generate for the
|
||||
# regression this revert fixes (issue #462).
|
||||
result = model.generate(str(audio_path), **decode_options)
|
||||
|
||||
# Extract text from result
|
||||
if isinstance(result, str):
|
||||
return result.strip()
|
||||
elif isinstance(result, dict):
|
||||
return result.get("text", "").strip()
|
||||
elif hasattr(result, "text"):
|
||||
return result.text.strip()
|
||||
else:
|
||||
return str(result).strip()
|
||||
# Extract text from result
|
||||
if isinstance(result, str):
|
||||
return result.strip()
|
||||
elif isinstance(result, dict):
|
||||
return result.get("text", "").strip()
|
||||
elif hasattr(result, "text"):
|
||||
return result.text.strip()
|
||||
else:
|
||||
return str(result).strip()
|
||||
finally:
|
||||
# Drop the local binding here, inside the frame a propagating
|
||||
# traceback would keep alive, so the caller's drain can
|
||||
# actually return the model's buffers to MLX.
|
||||
del model
|
||||
|
||||
def _reload_and_transcribe_sync():
|
||||
"""Ensure the requested model is loaded, then transcribe — as ONE
|
||||
|
||||
@@ -305,21 +305,27 @@ class MLXQwenLLMBackend:
|
||||
from mlx_lm.sample_utils import make_sampler
|
||||
|
||||
model, tokenizer = self.model, self.tokenizer
|
||||
messages = _build_messages(prompt, system, examples)
|
||||
chat_prompt = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
try:
|
||||
messages = _build_messages(prompt, system, examples)
|
||||
chat_prompt = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
|
||||
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
|
||||
text = mlx_generate(
|
||||
model,
|
||||
tokenizer,
|
||||
prompt=chat_prompt,
|
||||
max_tokens=max_tokens,
|
||||
sampler=sampler,
|
||||
verbose=False,
|
||||
)
|
||||
return text.strip()
|
||||
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
|
||||
text = mlx_generate(
|
||||
model,
|
||||
tokenizer,
|
||||
prompt=chat_prompt,
|
||||
max_tokens=max_tokens,
|
||||
sampler=sampler,
|
||||
verbose=False,
|
||||
)
|
||||
return text.strip()
|
||||
finally:
|
||||
# Drop the local binding here, inside the frame a propagating
|
||||
# traceback would keep alive, so the caller's drain can
|
||||
# actually return the model's buffers to MLX.
|
||||
del model, tokenizer
|
||||
|
||||
Reference in New Issue
Block a user