mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-29 07:05:14 -07:00
Refactor MLX and PyTorch Backend Model Loading
- Updated hidden imports in build_binary.py to replace 'mlx_audio.asr' with 'mlx_audio.stt'. - Enhanced model loading logic in MLX and PyTorch backends to ensure proper progress tracking during model downloads. - Improved error handling and context management for progress tracking in both backends. - Bumped version to 0.1.10 in Cargo.lock to reflect recent changes.
This commit is contained in:
@@ -341,15 +341,28 @@ class MLXSTTBackend:
|
|||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
try:
|
try:
|
||||||
from mlx_audio.asr import load
|
# IMPORTANT: Set up progress tracking BEFORE importing mlx_audio
|
||||||
|
# This ensures tqdm is patched before any HuggingFace Hub imports
|
||||||
# MLX Whisper model naming
|
|
||||||
model_name = f"mlx-community/whisper-{model_size}"
|
|
||||||
|
|
||||||
# Set up progress tracking
|
|
||||||
progress_manager = get_progress_manager()
|
progress_manager = get_progress_manager()
|
||||||
progress_model_name = f"whisper-{model_size}"
|
progress_model_name = f"whisper-{model_size}"
|
||||||
|
|
||||||
|
# Set up progress callback and tracker
|
||||||
|
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback)
|
||||||
|
|
||||||
|
# Patch tqdm BEFORE importing mlx_audio
|
||||||
|
# This is critical because mlx_audio imports huggingface_hub which imports tqdm
|
||||||
|
print("[DEBUG] Starting tqdm patch BEFORE mlx_audio import")
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
print("[DEBUG] tqdm patched, now importing mlx_audio")
|
||||||
|
|
||||||
|
# NOW import mlx_audio - it will use our patched tqdm
|
||||||
|
from mlx_audio.stt import load
|
||||||
|
|
||||||
|
# MLX Whisper uses the standard OpenAI models
|
||||||
|
model_name = f"openai/whisper-{model_size}"
|
||||||
|
|
||||||
# Start tracking download task
|
# Start tracking download task
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
task_manager.start_download(progress_model_name)
|
task_manager.start_download(progress_model_name)
|
||||||
@@ -365,13 +378,12 @@ class MLXSTTBackend:
|
|||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set up progress callback
|
# Load the model (tqdm is already patched from above)
|
||||||
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
try:
|
||||||
tracker = HFProgressTracker(progress_callback)
|
|
||||||
|
|
||||||
# Use progress tracker during download
|
|
||||||
with tracker.patch_download():
|
|
||||||
self.model = load(model_name)
|
self.model = load(model_name)
|
||||||
|
finally:
|
||||||
|
# Exit the patch context
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
self.model_size = model_size
|
self.model_size = model_size
|
||||||
|
|
||||||
|
|||||||
@@ -85,15 +85,25 @@ class PyTorchTTSBackend:
|
|||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
try:
|
try:
|
||||||
|
# IMPORTANT: Set up progress tracking BEFORE importing qwen_tts
|
||||||
|
# This ensures tqdm is patched before any HuggingFace Hub imports
|
||||||
|
progress_manager = get_progress_manager()
|
||||||
|
model_name = f"qwen-tts-{model_size}"
|
||||||
|
|
||||||
|
# Set up progress callback and tracker
|
||||||
|
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback)
|
||||||
|
|
||||||
|
# Patch tqdm BEFORE importing qwen_tts
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
|
||||||
|
# NOW import qwen_tts - it will use our patched tqdm
|
||||||
from qwen_tts import Qwen3TTSModel
|
from qwen_tts import Qwen3TTSModel
|
||||||
|
|
||||||
# Get model path (local or HuggingFace Hub ID)
|
# Get model path (local or HuggingFace Hub ID)
|
||||||
model_path = self._get_model_path(model_size)
|
model_path = self._get_model_path(model_size)
|
||||||
|
|
||||||
# Set up progress tracking
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
|
|
||||||
print(f"Loading TTS model {model_size} on {self.device}...")
|
print(f"Loading TTS model {model_size} on {self.device}...")
|
||||||
|
|
||||||
# Start tracking download task
|
# Start tracking download task
|
||||||
@@ -109,18 +119,16 @@ class PyTorchTTSBackend:
|
|||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set up progress callback
|
# Load the model (tqdm is already patched from above)
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
try:
|
||||||
tracker = HFProgressTracker(progress_callback)
|
|
||||||
|
|
||||||
# Use progress tracker during download
|
|
||||||
with tracker.patch_download():
|
|
||||||
# Load the model - downloads will happen automatically with progress tracking
|
|
||||||
self.model = Qwen3TTSModel.from_pretrained(
|
self.model = Qwen3TTSModel.from_pretrained(
|
||||||
model_path,
|
model_path,
|
||||||
device_map=self.device,
|
device_map=self.device,
|
||||||
torch_dtype=torch.float32 if self.device == "cpu" else torch.bfloat16,
|
torch_dtype=torch.float32 if self.device == "cpu" else torch.bfloat16,
|
||||||
)
|
)
|
||||||
|
finally:
|
||||||
|
# Exit the patch context
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
# Mark as complete
|
# Mark as complete
|
||||||
progress_manager.mark_complete(model_name)
|
progress_manager.mark_complete(model_name)
|
||||||
@@ -318,36 +326,57 @@ class PyTorchSTTBackend:
|
|||||||
Args:
|
Args:
|
||||||
model_size: Model size (tiny, base, small, medium, large)
|
model_size: Model size (tiny, base, small, medium, large)
|
||||||
"""
|
"""
|
||||||
|
print(f"[DEBUG] load_model_async called with size: {model_size}")
|
||||||
if model_size is None:
|
if model_size is None:
|
||||||
model_size = self.model_size
|
model_size = self.model_size
|
||||||
|
|
||||||
|
print(f"[DEBUG] Model already loaded? {self.model is not None}, current size: {self.model_size}, requested: {model_size}")
|
||||||
if self.model is not None and self.model_size == model_size:
|
if self.model is not None and self.model_size == model_size:
|
||||||
|
print(f"[DEBUG] Early return - model already loaded")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
print(f"[DEBUG] Calling asyncio.to_thread for _load_model_sync")
|
||||||
# Run blocking load in thread pool
|
# Run blocking load in thread pool
|
||||||
await asyncio.to_thread(self._load_model_sync, model_size)
|
await asyncio.to_thread(self._load_model_sync, model_size)
|
||||||
|
print(f"[DEBUG] asyncio.to_thread completed")
|
||||||
|
|
||||||
# Alias for compatibility
|
# Alias for compatibility
|
||||||
load_model = load_model_async
|
load_model = load_model_async
|
||||||
|
|
||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
|
print(f"[DEBUG] _load_model_sync called for Whisper {model_size}")
|
||||||
try:
|
try:
|
||||||
|
# IMPORTANT: Set up progress tracking BEFORE importing transformers
|
||||||
|
# This ensures tqdm is patched before any HuggingFace Hub imports
|
||||||
|
progress_manager = get_progress_manager()
|
||||||
|
progress_model_name = f"whisper-{model_size}"
|
||||||
|
|
||||||
|
# Set up progress callback and tracker
|
||||||
|
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback)
|
||||||
|
|
||||||
|
# Patch tqdm BEFORE importing transformers
|
||||||
|
print("[DEBUG] Starting tqdm patch BEFORE transformers import")
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
print("[DEBUG] tqdm patched, now importing transformers")
|
||||||
|
|
||||||
|
# NOW import transformers - it will use our patched tqdm
|
||||||
from transformers import WhisperProcessor, WhisperForConditionalGeneration
|
from transformers import WhisperProcessor, WhisperForConditionalGeneration
|
||||||
|
|
||||||
model_name = f"openai/whisper-{model_size}"
|
model_name = f"openai/whisper-{model_size}"
|
||||||
|
print(f"[DEBUG] Model name: {model_name}")
|
||||||
# Set up progress tracking
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
|
||||||
|
|
||||||
# Start tracking download task
|
# Start tracking download task
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
task_manager.start_download(progress_model_name)
|
task_manager.start_download(progress_model_name)
|
||||||
|
print(f"[DEBUG] Task manager started download")
|
||||||
|
|
||||||
print(f"Loading Whisper model {model_size} on {self.device}...")
|
print(f"Loading Whisper model {model_size} on {self.device}...")
|
||||||
|
|
||||||
# Initialize progress state to show download has started
|
# Initialize progress state to show download has started
|
||||||
|
print(f"[DEBUG] Calling update_progress...")
|
||||||
progress_manager.update_progress(
|
progress_manager.update_progress(
|
||||||
model_name=progress_model_name,
|
model_name=progress_model_name,
|
||||||
current=0,
|
current=0,
|
||||||
@@ -355,15 +384,15 @@ class PyTorchSTTBackend:
|
|||||||
filename="",
|
filename="",
|
||||||
status="downloading",
|
status="downloading",
|
||||||
)
|
)
|
||||||
|
print(f"[DEBUG] update_progress called, listeners: {len(progress_manager._listeners.get(progress_model_name, []))}")
|
||||||
|
|
||||||
# Set up progress callback
|
# Load models (tqdm is already patched from above)
|
||||||
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
try:
|
||||||
tracker = HFProgressTracker(progress_callback)
|
|
||||||
|
|
||||||
# Use progress tracker during download
|
|
||||||
with tracker.patch_download():
|
|
||||||
self.processor = WhisperProcessor.from_pretrained(model_name)
|
self.processor = WhisperProcessor.from_pretrained(model_name)
|
||||||
self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
|
self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
|
||||||
|
finally:
|
||||||
|
# Exit the patch context
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
self.model.to(self.device)
|
self.model.to(self.device)
|
||||||
self.model_size = model_size
|
self.model_size = model_size
|
||||||
|
|||||||
@@ -80,7 +80,7 @@ def build_server():
|
|||||||
'--hidden-import', 'mlx.nn',
|
'--hidden-import', 'mlx.nn',
|
||||||
'--hidden-import', 'mlx_audio',
|
'--hidden-import', 'mlx_audio',
|
||||||
'--hidden-import', 'mlx_audio.tts',
|
'--hidden-import', 'mlx_audio.tts',
|
||||||
'--hidden-import', 'mlx_audio.asr',
|
'--hidden-import', 'mlx_audio.stt',
|
||||||
'--collect-submodules', 'mlx',
|
'--collect-submodules', 'mlx',
|
||||||
'--collect-submodules', 'mlx_audio',
|
'--collect-submodules', 'mlx_audio',
|
||||||
# Collect MLX data files including Metal shader libraries (.metallib)
|
# Collect MLX data files including Metal shader libraries (.metallib)
|
||||||
|
|||||||
+5
-1
@@ -1393,7 +1393,11 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
|||||||
async def download_in_background():
|
async def download_in_background():
|
||||||
"""Download model in background without blocking the HTTP request."""
|
"""Download model in background without blocking the HTTP request."""
|
||||||
try:
|
try:
|
||||||
await asyncio.to_thread(config["load_func"])
|
# Call the load function (which may be async)
|
||||||
|
result = config["load_func"]()
|
||||||
|
# If it's a coroutine, await it
|
||||||
|
if asyncio.iscoroutine(result):
|
||||||
|
await result
|
||||||
task_manager.complete_download(request.model_name)
|
task_manager.complete_download(request.model_name)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
task_manager.error_download(request.model_name, str(e))
|
task_manager.error_download(request.model_name, str(e))
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ class HFProgressTracker:
|
|||||||
"""A tqdm subclass that reports progress to our tracker."""
|
"""A tqdm subclass that reports progress to our tracker."""
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
def __init__(self, *args, **kwargs):
|
||||||
|
print(f"[DEBUG TrackedTqdm] __init__ called with desc: {kwargs.get('desc', '')}")
|
||||||
# Extract filename from desc before passing to parent
|
# Extract filename from desc before passing to parent
|
||||||
desc = kwargs.get("desc", "")
|
desc = kwargs.get("desc", "")
|
||||||
if not desc and args:
|
if not desc and args:
|
||||||
@@ -79,6 +80,7 @@ class HFProgressTracker:
|
|||||||
}
|
}
|
||||||
|
|
||||||
def update(self, n=1):
|
def update(self, n=1):
|
||||||
|
print(f"[DEBUG TrackedTqdm] update called with n={n}")
|
||||||
result = super().update(n)
|
result = super().update(n)
|
||||||
|
|
||||||
# Report progress
|
# Report progress
|
||||||
@@ -118,11 +120,13 @@ class HFProgressTracker:
|
|||||||
@contextmanager
|
@contextmanager
|
||||||
def patch_download(self):
|
def patch_download(self):
|
||||||
"""Context manager to patch tqdm for progress tracking."""
|
"""Context manager to patch tqdm for progress tracking."""
|
||||||
|
print("[DEBUG HFProgressTracker] patch_download called")
|
||||||
try:
|
try:
|
||||||
import tqdm as tqdm_module
|
import tqdm as tqdm_module
|
||||||
|
|
||||||
# Store original tqdm class
|
# Store original tqdm class
|
||||||
self._original_tqdm_class = tqdm_module.tqdm
|
self._original_tqdm_class = tqdm_module.tqdm
|
||||||
|
print(f"[DEBUG HFProgressTracker] Original tqdm class: {self._original_tqdm_class}")
|
||||||
|
|
||||||
# Reset totals
|
# Reset totals
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -135,18 +139,22 @@ class HFProgressTracker:
|
|||||||
|
|
||||||
# Create our tracked tqdm class
|
# Create our tracked tqdm class
|
||||||
tracked_tqdm = self._create_tracked_tqdm_class()
|
tracked_tqdm = self._create_tracked_tqdm_class()
|
||||||
|
print(f"[DEBUG HFProgressTracker] Created TrackedTqdm class: {tracked_tqdm}")
|
||||||
|
|
||||||
# Patch tqdm.tqdm
|
# Patch tqdm.tqdm
|
||||||
tqdm_module.tqdm = tracked_tqdm
|
tqdm_module.tqdm = tracked_tqdm
|
||||||
|
print(f"[DEBUG HFProgressTracker] Patched tqdm.tqdm")
|
||||||
|
|
||||||
# Also patch tqdm.auto.tqdm if it exists (used by huggingface_hub)
|
# Also patch tqdm.auto.tqdm if it exists (used by huggingface_hub)
|
||||||
self._original_tqdm_auto = None
|
self._original_tqdm_auto = None
|
||||||
if hasattr(tqdm_module, "auto") and hasattr(tqdm_module.auto, "tqdm"):
|
if hasattr(tqdm_module, "auto") and hasattr(tqdm_module.auto, "tqdm"):
|
||||||
self._original_tqdm_auto = tqdm_module.auto.tqdm
|
self._original_tqdm_auto = tqdm_module.auto.tqdm
|
||||||
tqdm_module.auto.tqdm = tracked_tqdm
|
tqdm_module.auto.tqdm = tracked_tqdm
|
||||||
|
print(f"[DEBUG HFProgressTracker] Patched tqdm.auto.tqdm")
|
||||||
|
|
||||||
# Patch in sys.modules to catch already-imported references
|
# Patch in sys.modules to catch already-imported references
|
||||||
self._patched_modules = {}
|
self._patched_modules = {}
|
||||||
|
patched_count = 0
|
||||||
for module_name in list(sys.modules.keys()):
|
for module_name in list(sys.modules.keys()):
|
||||||
if "huggingface" in module_name or module_name.startswith("tqdm"):
|
if "huggingface" in module_name or module_name.startswith("tqdm"):
|
||||||
try:
|
try:
|
||||||
@@ -159,8 +167,11 @@ class HFProgressTracker:
|
|||||||
):
|
):
|
||||||
self._patched_modules[module_name] = attr
|
self._patched_modules[module_name] = attr
|
||||||
setattr(module, "tqdm", tracked_tqdm)
|
setattr(module, "tqdm", tracked_tqdm)
|
||||||
|
patched_count += 1
|
||||||
|
print(f"[DEBUG HFProgressTracker] Patched {module_name}.tqdm")
|
||||||
except (AttributeError, TypeError):
|
except (AttributeError, TypeError):
|
||||||
pass
|
pass
|
||||||
|
print(f"[DEBUG HFProgressTracker] Patched {patched_count} modules in sys.modules")
|
||||||
|
|
||||||
yield
|
yield
|
||||||
|
|
||||||
|
|||||||
@@ -90,15 +90,25 @@ class ProgressManager:
|
|||||||
"timestamp": datetime.now().isoformat(),
|
"timestamp": datetime.now().isoformat(),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
print(f"[DEBUG] update_progress called: {model_name}, {progress_pct:.1f}%")
|
||||||
|
|
||||||
# Thread-safe update of progress dict
|
# Thread-safe update of progress dict
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._progress[model_name] = progress_data
|
self._progress[model_name] = progress_data
|
||||||
|
|
||||||
# Notify all listeners (thread-safe)
|
# Notify all listeners (thread-safe)
|
||||||
listener_count = len(self._listeners.get(model_name, []))
|
listener_count = len(self._listeners.get(model_name, []))
|
||||||
|
print(f"[DEBUG] Listener count for {model_name}: {listener_count}")
|
||||||
|
print(f"[DEBUG] All listeners: {list(self._listeners.keys())}")
|
||||||
|
print(f"[DEBUG] Main loop set: {self._main_loop is not None}")
|
||||||
|
if self._main_loop:
|
||||||
|
print(f"[DEBUG] Main loop running: {self._main_loop.is_running()}")
|
||||||
|
|
||||||
if listener_count > 0:
|
if listener_count > 0:
|
||||||
logger.debug(f"Notifying {listener_count} listeners for {model_name}: {progress_pct:.1f}% ({filename})")
|
logger.debug(f"Notifying {listener_count} listeners for {model_name}: {progress_pct:.1f}% ({filename})")
|
||||||
|
print(f"[DEBUG] About to notify listeners...")
|
||||||
self._notify_listeners_threadsafe(model_name, progress_data)
|
self._notify_listeners_threadsafe(model_name, progress_data)
|
||||||
|
print(f"[DEBUG] Notified listeners")
|
||||||
else:
|
else:
|
||||||
logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%")
|
logger.debug(f"No listeners for {model_name}, progress update stored: {progress_pct:.1f}%")
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from PyInstaller.utils.hooks import collect_submodules
|
|||||||
from PyInstaller.utils.hooks import copy_metadata
|
from PyInstaller.utils.hooks import copy_metadata
|
||||||
|
|
||||||
datas = []
|
datas = []
|
||||||
hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.profiles', 'backend.history', 'backend.tts', 'backend.transcribe', 'backend.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.utils.validation', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'librosa', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'pkg_resources.extern', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.asr']
|
hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.profiles', 'backend.history', 'backend.tts', 'backend.transcribe', 'backend.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.utils.validation', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'librosa', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'pkg_resources.extern', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.stt']
|
||||||
datas += collect_data_files('qwen_tts')
|
datas += collect_data_files('qwen_tts')
|
||||||
datas += collect_data_files('mlx')
|
datas += collect_data_files('mlx')
|
||||||
datas += collect_data_files('mlx_audio')
|
datas += collect_data_files('mlx_audio')
|
||||||
|
|||||||
Generated
+1
-1
@@ -5041,7 +5041,7 @@ checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "voicebox"
|
name = "voicebox"
|
||||||
version = "0.1.9"
|
version = "0.1.10"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"base64 0.22.1",
|
"base64 0.22.1",
|
||||||
"core-foundation-sys",
|
"core-foundation-sys",
|
||||||
|
|||||||
Binary file not shown.
Reference in New Issue
Block a user