mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
Unloading a TTS/Whisper/LLM model on the MLX backend only dropped the
Python reference (`del self.model`). MLX keeps freed array buffers in
its own allocator pool for reuse instead of returning them to the OS,
so the process's memory footprint never actually shrank after unload
on Apple Silicon (the default backend there) until the process exited.
Add empty_mlx_cache() (backend/backends/base.py), wrapping
mx.clear_cache(), and call it from the three MLX unload_model()
implementations: MLXTTSBackend, MLXSTTBackend, MLXQwenLLMBackend.
Separately, the voice-clone prompt cache (backend/utils/cache.py) is a
process-lifetime dict populated by create_voice_prompt() across every
TTS engine, but nothing ever cleared it on model unload — only the
unrelated /tasks/clear-cache endpoint touched it. Add
clear_voice_prompt_memory_cache() (memory only, disk cache untouched
so a later generation still reloads the prompt instead of recomputing
it) and wire it into every TTS unload path (services/tts.py and the
qwen_custom_voice / generic branches of unload_model_by_config).
Whisper and the LLM backends never produce voice prompts, so their
unload paths are left alone.
Testing:
- New unit tests: backend/tests/test_mlx_unload_clears_cache.py,
backend/tests/test_voice_prompt_cache_unload.py (8 tests, all pass).
- Verified end-to-end on Apple Silicon against real cached models
(Qwen TTS 1.7B, Whisper Turbo, Qwen3 0.6B): loaded each via the
running app, unloaded via the real /models/{name}/unload endpoint,
and confirmed via mx.get_cache_memory()/get_active_memory() that the
MLX allocator's cache drops to 0 on every cycle. Ran a real
voice-clone generation end to end and confirmed the in-memory prompt
cache goes from 1 entry to 0 on unload while the on-disk .prompt
file is left intact.
320 lines
11 KiB
Python
320 lines
11 KiB
Python
"""
|
|
Qwen3 LLM backend implementations.
|
|
|
|
Provides MLX (Apple Silicon, 4-bit community quants) and PyTorch
|
|
(transformers AutoModelForCausalLM) paths that share the same
|
|
`LLMBackend` protocol and model-load progress plumbing as the TTS
|
|
and STT engines.
|
|
"""
|
|
|
|
import asyncio
|
|
import logging
|
|
from typing import Optional
|
|
|
|
from . import LLMBackend, DEFAULT_LLM_MAX_TOKENS, DEFAULT_LLM_TEMPERATURE
|
|
from .base import (
|
|
is_model_cached,
|
|
get_torch_device,
|
|
empty_device_cache,
|
|
empty_mlx_cache,
|
|
manual_seed,
|
|
model_load_progress,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
PYTORCH_HF_REPOS = {
|
|
"0.6B": "Qwen/Qwen3-0.6B",
|
|
"1.7B": "Qwen/Qwen3-1.7B",
|
|
"4B": "Qwen/Qwen3-4B",
|
|
}
|
|
|
|
MLX_HF_REPOS = {
|
|
"0.6B": "mlx-community/Qwen3-0.6B-4bit",
|
|
"1.7B": "mlx-community/Qwen3-1.7B-4bit",
|
|
"4B": "mlx-community/Qwen3-4B-4bit",
|
|
}
|
|
|
|
|
|
def _progress_name(model_size: str) -> str:
|
|
return f"qwen3-{model_size.lower()}"
|
|
|
|
|
|
def _build_messages(
|
|
prompt: str,
|
|
system: Optional[str],
|
|
examples: Optional[list[tuple[str, str]]] = None,
|
|
) -> list[dict]:
|
|
messages: list[dict] = []
|
|
if system:
|
|
messages.append({"role": "system", "content": system})
|
|
if examples:
|
|
for user_text, assistant_text in examples:
|
|
messages.append({"role": "user", "content": user_text})
|
|
messages.append({"role": "assistant", "content": assistant_text})
|
|
messages.append({"role": "user", "content": prompt})
|
|
return messages
|
|
|
|
|
|
class PyTorchQwenLLMBackend:
|
|
"""Qwen3 LLM backend using HuggingFace transformers."""
|
|
|
|
def __init__(self, model_size: str = "0.6B"):
|
|
self.model = None
|
|
self.tokenizer = None
|
|
self.model_size = model_size
|
|
self._current_model_size: Optional[str] = None
|
|
self.device = self._get_device()
|
|
|
|
def _get_device(self) -> str:
|
|
return get_torch_device(allow_xpu=True, allow_directml=True, allow_mps=True)
|
|
|
|
def is_loaded(self) -> bool:
|
|
return self.model is not None
|
|
|
|
def _get_model_path(self, model_size: str) -> str:
|
|
if model_size not in PYTORCH_HF_REPOS:
|
|
raise ValueError(f"Unknown Qwen3 size: {model_size}")
|
|
return PYTORCH_HF_REPOS[model_size]
|
|
|
|
def _is_model_cached(self, model_size: str) -> bool:
|
|
return is_model_cached(self._get_model_path(model_size))
|
|
|
|
async def load_model(self, model_size: Optional[str] = None) -> None:
|
|
if model_size is None:
|
|
model_size = self.model_size
|
|
|
|
if self.model is not None and self._current_model_size == model_size:
|
|
return
|
|
|
|
if self.model is not None and self._current_model_size != model_size:
|
|
self.unload_model()
|
|
|
|
await asyncio.to_thread(self._load_model_sync, model_size)
|
|
|
|
def _load_model_sync(self, model_size: str) -> None:
|
|
import torch
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
|
|
progress_model_name = _progress_name(model_size)
|
|
is_cached = self._is_model_cached(model_size)
|
|
repo = self._get_model_path(model_size)
|
|
|
|
with model_load_progress(progress_model_name, is_cached):
|
|
logger.info("Loading Qwen3 %s on %s...", model_size, self.device)
|
|
# Loads run with the process's default HF_HUB_OFFLINE state.
|
|
# Forcing offline for cached models flips process-global state
|
|
# and silently switches every concurrent download/load on other
|
|
# threads to offline mode (issue #841) — the same regression
|
|
# removed app-wide in #524/#530.
|
|
self.tokenizer = AutoTokenizer.from_pretrained(repo)
|
|
dtype = torch.float16 if self.device in ("cuda", "mps") else torch.float32
|
|
self.model = AutoModelForCausalLM.from_pretrained(
|
|
repo,
|
|
dtype=dtype,
|
|
)
|
|
self.model.to(self.device)
|
|
self.model.eval()
|
|
|
|
self._current_model_size = model_size
|
|
self.model_size = model_size
|
|
logger.info("Qwen3 %s loaded successfully", model_size)
|
|
|
|
def unload_model(self) -> None:
|
|
if self.model is None:
|
|
return
|
|
del self.model
|
|
del self.tokenizer
|
|
self.model = None
|
|
self.tokenizer = None
|
|
self._current_model_size = None
|
|
empty_device_cache(self.device)
|
|
logger.info("Qwen3 unloaded")
|
|
|
|
async def generate(
|
|
self,
|
|
prompt: str,
|
|
system: Optional[str] = None,
|
|
max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
|
|
temperature: float = DEFAULT_LLM_TEMPERATURE,
|
|
model_size: Optional[str] = None,
|
|
examples: Optional[list[tuple[str, str]]] = None,
|
|
) -> str:
|
|
await self.load_model(model_size)
|
|
return await asyncio.to_thread(
|
|
self._generate_sync, prompt, system, max_tokens, temperature, examples
|
|
)
|
|
|
|
def _generate_sync(
|
|
self,
|
|
prompt: str,
|
|
system: Optional[str],
|
|
max_tokens: int,
|
|
temperature: float,
|
|
examples: Optional[list[tuple[str, str]]] = None,
|
|
) -> str:
|
|
import torch
|
|
|
|
messages = _build_messages(prompt, system, examples)
|
|
text = self.tokenizer.apply_chat_template(
|
|
messages,
|
|
tokenize=False,
|
|
add_generation_prompt=True,
|
|
enable_thinking=False,
|
|
)
|
|
inputs = self.tokenizer(text, return_tensors="pt").to(self.device)
|
|
|
|
do_sample = temperature > 0
|
|
generate_kwargs = {
|
|
"max_new_tokens": max_tokens,
|
|
"do_sample": do_sample,
|
|
"pad_token_id": self.tokenizer.eos_token_id,
|
|
}
|
|
if do_sample:
|
|
generate_kwargs["temperature"] = temperature
|
|
generate_kwargs["top_p"] = 0.9
|
|
|
|
with torch.no_grad():
|
|
output_ids = self.model.generate(**inputs, **generate_kwargs)
|
|
|
|
input_len = inputs["input_ids"].shape[1]
|
|
new_tokens = output_ids[0, input_len:]
|
|
return self.tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
|
|
|
|
|
|
class MLXQwenLLMBackend:
|
|
"""Qwen3 LLM backend using mlx-lm (Apple Silicon)."""
|
|
|
|
def __init__(self, model_size: str = "0.6B"):
|
|
self.model = None
|
|
self.tokenizer = None
|
|
self.model_size = model_size
|
|
self._current_model_size: Optional[str] = None
|
|
# Same role as MLXTTSBackend._op_lock: keeps two coroutines from
|
|
# racing their reload decisions around one load-then-generate.
|
|
self._op_lock = asyncio.Lock()
|
|
|
|
def is_loaded(self) -> bool:
|
|
return self.model is not None
|
|
|
|
def _get_model_path(self, model_size: str) -> str:
|
|
if model_size not in MLX_HF_REPOS:
|
|
raise ValueError(f"Unknown Qwen3 size: {model_size}")
|
|
return MLX_HF_REPOS[model_size]
|
|
|
|
def _is_model_cached(self, model_size: str) -> bool:
|
|
return is_model_cached(
|
|
self._get_model_path(model_size),
|
|
weight_extensions=(".safetensors", ".bin", ".npz"),
|
|
)
|
|
|
|
async def load_model(self, model_size: Optional[str] = None) -> None:
|
|
if model_size is None:
|
|
model_size = self.model_size
|
|
|
|
if self.model is not None and self._current_model_size == model_size:
|
|
return
|
|
|
|
# Routed through the same dedicated MLX thread as TTS/STT — MLX's
|
|
# Metal stream is thread-local, so load and generate must run on
|
|
# one thread (see mlx_backend._run_on_mlx_thread / issue #699).
|
|
from .mlx_backend import _run_on_mlx_thread
|
|
|
|
async with self._op_lock:
|
|
await _run_on_mlx_thread(self._reload_sync, model_size)
|
|
|
|
def _reload_sync(self, model_size: str) -> None:
|
|
"""Unload a mismatched model and load the requested one, in one MLX-thread op."""
|
|
if self.model is not None and self._current_model_size != model_size:
|
|
self.unload_model()
|
|
self._load_model_sync(model_size)
|
|
|
|
def _load_model_sync(self, model_size: str) -> None:
|
|
from mlx_lm import load as mlx_load
|
|
|
|
progress_model_name = _progress_name(model_size)
|
|
is_cached = self._is_model_cached(model_size)
|
|
repo = self._get_model_path(model_size)
|
|
|
|
with model_load_progress(progress_model_name, is_cached):
|
|
logger.info("Loading Qwen3 %s via MLX...", model_size)
|
|
# See the PyTorch loader comment — no offline forcing (issue #841).
|
|
loaded = mlx_load(repo)
|
|
|
|
# mlx_lm.load returns (model, tokenizer) by default and
|
|
# (model, tokenizer, config) when return_config=True.
|
|
self.model = loaded[0]
|
|
self.tokenizer = loaded[1]
|
|
|
|
self._current_model_size = model_size
|
|
self.model_size = model_size
|
|
logger.info("Qwen3 %s (MLX) loaded successfully", model_size)
|
|
|
|
def unload_model(self) -> None:
|
|
if self.model is None:
|
|
return
|
|
del self.model
|
|
del self.tokenizer
|
|
self.model = None
|
|
self.tokenizer = None
|
|
self._current_model_size = None
|
|
empty_mlx_cache()
|
|
logger.info("Qwen3 (MLX) unloaded")
|
|
|
|
async def generate(
|
|
self,
|
|
prompt: str,
|
|
system: Optional[str] = None,
|
|
max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
|
|
temperature: float = DEFAULT_LLM_TEMPERATURE,
|
|
model_size: Optional[str] = None,
|
|
examples: Optional[list[tuple[str, str]]] = None,
|
|
) -> str:
|
|
from .mlx_backend import _run_on_mlx_thread
|
|
|
|
resolved_size = model_size if model_size is not None else self.model_size
|
|
|
|
def _reload_and_generate_sync() -> str:
|
|
# One worker submission for load + generate, as in
|
|
# MLXTTSBackend._reload_and_generate_sync: no gap between the two
|
|
# where another request's load_model (different size) or an
|
|
# unload can swap or null out self.model / self.tokenizer.
|
|
if self.model is None or self._current_model_size != resolved_size:
|
|
self._reload_sync(resolved_size)
|
|
return self._generate_sync(prompt, system, max_tokens, temperature, examples)
|
|
|
|
async with self._op_lock:
|
|
return await _run_on_mlx_thread(_reload_and_generate_sync)
|
|
|
|
def _generate_sync(
|
|
self,
|
|
prompt: str,
|
|
system: Optional[str],
|
|
max_tokens: int,
|
|
temperature: float,
|
|
examples: Optional[list[tuple[str, str]]] = None,
|
|
) -> str:
|
|
from mlx_lm import generate as mlx_generate
|
|
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,
|
|
)
|
|
|
|
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()
|