Refactor App and Sidebar components to support macOS, add TitleBarDragRegion for improved window dragging, and enhance model management with delete functionality and download progress tracking. Update audio player to handle audio resets more effectively and improve generation form with model download notifications.

This commit is contained in:
Jamie Pine
2026-01-25 23:25:21 -08:00
parent 090b1f6dde
commit b2659e6a6d
19 changed files with 919 additions and 154 deletions
+85 -3
View File
@@ -314,9 +314,9 @@ async def generate_speech(
# Generate audio
tts_model = tts.get_tts_model()
# Load the requested model size if different from current
# Load the requested model size if different from current (async to not block)
model_size = data.model_size or "1.7B"
tts_model.load_model(model_size)
await tts_model.load_model_async(model_size)
audio, sample_rate = await tts_model.generate(
data.text,
voice_prompt,
@@ -519,7 +519,7 @@ async def load_model(model_size: str = "1.7B"):
"""Manually load TTS model."""
try:
tts_model = tts.get_tts_model()
tts_model.load_model(model_size)
await tts_model.load_model_async(model_size)
return {"message": f"Model {model_size} loaded successfully"}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@@ -784,6 +784,88 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
raise HTTPException(status_code=500, detail=str(e))
@app.delete("/models/{model_name}")
async def delete_model(model_name: str):
"""Delete a downloaded model from the HuggingFace cache."""
import shutil
import os
# Map model names to HuggingFace repo IDs
model_configs = {
"qwen-tts-1.7B": {
"hf_repo_id": "Qwen/Qwen3-TTS-12Hz-1.7B-Base",
"model_size": "1.7B",
"model_type": "tts",
},
"qwen-tts-0.6B": {
"hf_repo_id": "Qwen/Qwen3-TTS-12Hz-0.6B-Base",
"model_size": "0.6B",
"model_type": "tts",
},
"whisper-base": {
"hf_repo_id": "openai/whisper-base",
"model_size": "base",
"model_type": "whisper",
},
"whisper-small": {
"hf_repo_id": "openai/whisper-small",
"model_size": "small",
"model_type": "whisper",
},
"whisper-medium": {
"hf_repo_id": "openai/whisper-medium",
"model_size": "medium",
"model_type": "whisper",
},
"whisper-large": {
"hf_repo_id": "openai/whisper-large",
"model_size": "large",
"model_type": "whisper",
},
}
if model_name not in model_configs:
raise HTTPException(status_code=400, detail=f"Unknown model: {model_name}")
config = model_configs[model_name]
hf_repo_id = config["hf_repo_id"]
try:
# Check if model is loaded and unload it first
if config["model_type"] == "tts":
tts_model = tts.get_tts_model()
if tts_model.is_loaded() and tts_model.model_size == config["model_size"]:
tts.unload_tts_model()
elif config["model_type"] == "whisper":
whisper_model = transcribe.get_whisper_model()
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
transcribe.unload_whisper_model()
# Find and delete the cache directory
cache_dir = os.path.expanduser("~/.cache/huggingface/hub")
repo_cache_dir = Path(cache_dir) / ("models--" + hf_repo_id.replace("/", "--"))
# Check if the cache directory exists
if not repo_cache_dir.exists():
raise HTTPException(status_code=404, detail=f"Model {model_name} not found in cache")
# Delete the entire cache directory for this model
try:
shutil.rmtree(repo_cache_dir)
except OSError as e:
raise HTTPException(
status_code=500,
detail=f"Failed to delete model cache directory: {str(e)}"
)
return {"message": f"Model {model_name} deleted successfully"}
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to delete model: {str(e)}")
# ============================================
# STARTUP & SHUTDOWN
# ============================================
+102 -81
View File
@@ -3,6 +3,7 @@ Whisper ASR module for transcription.
"""
from typing import Optional, List, Dict
import asyncio
import torch
import numpy as np
from pathlib import Path
@@ -79,6 +80,22 @@ class WhisperModel:
progress_manager.mark_error(f"whisper-{model_size}", str(e))
raise
async def load_model_async(self, model_size: Optional[str] = None):
"""
Async version of load_model that runs in thread pool.
This prevents blocking the event loop during model loading.
"""
if model_size is None:
model_size = self.model_size
# If already loaded with correct size, return immediately
if self.model is not None and self.model_size == model_size:
return
# Run the blocking load operation in a thread pool
await asyncio.to_thread(self.load_model, model_size)
def unload_model(self):
"""Unload the model to free memory."""
if self.model is not None:
@@ -107,44 +124,49 @@ class WhisperModel:
Returns:
Transcribed text
"""
self.load_model()
await self.load_model_async()
from .utils.audio import load_audio
# Load audio
audio, sr = load_audio(audio_path, sample_rate=16000)
# Process audio
inputs = self.processor(
audio,
sampling_rate=16000,
return_tensors="pt",
)
inputs = inputs.to(self.device)
# Set language if provided
forced_decoder_ids = None
if language:
lang_code = "en" if language == "en" else "zh"
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
language=lang_code,
task="transcribe",
def _transcribe_sync():
"""Run synchronous transcription in thread pool."""
# Load audio
audio, sr = load_audio(audio_path, sample_rate=16000)
# Process audio
inputs = self.processor(
audio,
sampling_rate=16000,
return_tensors="pt",
)
inputs = inputs.to(self.device)
# Set language if provided
forced_decoder_ids = None
if language:
lang_code = "en" if language == "en" else "zh"
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
language=lang_code,
task="transcribe",
)
# Generate transcription
with torch.no_grad():
predicted_ids = self.model.generate(
inputs["input_features"],
forced_decoder_ids=forced_decoder_ids,
)
# Decode
transcription = self.processor.batch_decode(
predicted_ids,
skip_special_tokens=True,
)[0]
return transcription.strip()
# Generate transcription
with torch.no_grad():
predicted_ids = self.model.generate(
inputs["input_features"],
forced_decoder_ids=forced_decoder_ids,
)
# Decode
transcription = self.processor.batch_decode(
predicted_ids,
skip_special_tokens=True,
)[0]
return transcription.strip()
# Run blocking transcription in thread pool
return await asyncio.to_thread(_transcribe_sync)
async def transcribe_with_timestamps(
self,
@@ -161,59 +183,58 @@ class WhisperModel:
Returns:
List of word segments with timestamps
"""
self.load_model()
await self.load_model_async()
from .utils.audio import load_audio
# Load audio
audio, sr = load_audio(audio_path, sample_rate=16000)
# Process audio
inputs = self.processor(
audio,
sampling_rate=16000,
return_tensors="pt",
)
inputs = inputs.to(self.device)
# Set language if provided
forced_decoder_ids = None
if language:
lang_code = "en" if language == "en" else "zh"
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
language=lang_code,
task="transcribe",
def _transcribe_timestamps_sync():
"""Run synchronous transcription with timestamps in thread pool."""
# Load audio
audio, sr = load_audio(audio_path, sample_rate=16000)
# Process audio
inputs = self.processor(
audio,
sampling_rate=16000,
return_tensors="pt",
)
inputs = inputs.to(self.device)
# Set language if provided
forced_decoder_ids = None
if language:
lang_code = "en" if language == "en" else "zh"
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
language=lang_code,
task="transcribe",
)
# Generate with timestamps
with torch.no_grad():
predicted_ids = self.model.generate(
inputs["input_features"],
forced_decoder_ids=forced_decoder_ids,
return_timestamps=True,
)
# Parse timestamps (simplified - would need more robust parsing)
# For now, return basic transcription
# TODO: Implement proper timestamp parsing
transcription = self.processor.batch_decode(
predicted_ids,
skip_special_tokens=True,
)[0]
return [
{
"text": transcription,
"start": 0.0,
"end": len(audio) / sr,
}
]
# Generate with timestamps
with torch.no_grad():
predicted_ids = self.model.generate(
inputs["input_features"],
forced_decoder_ids=forced_decoder_ids,
return_timestamps=True,
)
# Decode with timestamps
result = self.processor.batch_decode(
predicted_ids,
skip_special_tokens=False,
)[0]
# Parse timestamps (simplified - would need more robust parsing)
# For now, return basic transcription
# TODO: Implement proper timestamp parsing
transcription = self.processor.batch_decode(
predicted_ids,
skip_special_tokens=True,
)[0]
return [
{
"text": transcription,
"start": 0.0,
"end": len(audio) / sr,
}
]
# Run blocking transcription in thread pool
return await asyncio.to_thread(_transcribe_timestamps_sync)
# Global model instance
+55 -20
View File
@@ -3,6 +3,7 @@ TTS inference module using Qwen3-TTS.
"""
from typing import Optional, List, Tuple
import asyncio
import torch
import numpy as np
import io
@@ -110,6 +111,15 @@ class TTSModel:
if model_path.startswith("Qwen/"):
print(f"Loading TTS model {model_size} on {self.device}...")
# Initialize progress state to show download has started
progress_manager.update_progress(
model_name=model_name,
current=0,
total=1, # Set to 1 initially, will be updated by callback
filename="",
status="downloading",
)
# Set up progress callback
progress_callback = create_hf_progress_callback(model_name, progress_manager)
tracker = HFProgressTracker(progress_callback)
@@ -151,6 +161,22 @@ class TTSModel:
progress_manager.mark_error(f"qwen-tts-{model_size}", str(e))
raise
async def load_model_async(self, model_size: Optional[str] = None):
"""
Async version of load_model that runs in thread pool.
This prevents blocking the event loop during model loading.
"""
if model_size is None:
model_size = self.model_size
# If already loaded with correct size, return immediately
if self.model is not None and self._current_model_size == model_size:
return
# Run the blocking load operation in a thread pool
await asyncio.to_thread(self.load_model, model_size)
def unload_model(self):
"""Unload the model to free memory."""
if self.model is not None:
@@ -180,7 +206,7 @@ class TTSModel:
Returns:
Tuple of (voice_prompt_dict, was_cached)
"""
self.load_model()
await self.load_model_async()
# Check cache if enabled
if use_cache:
@@ -189,12 +215,16 @@ class TTSModel:
if cached_prompt is not None:
return cached_prompt, True
# Create new voice prompt
voice_prompt_items = self.model.create_voice_clone_prompt(
ref_audio=str(audio_path),
ref_text=reference_text,
x_vector_only_mode=False,
)
def _create_prompt_sync():
"""Run synchronous voice prompt creation in thread pool."""
return self.model.create_voice_clone_prompt(
ref_audio=str(audio_path),
ref_text=reference_text,
x_vector_only_mode=False,
)
# Run blocking operation in thread pool
voice_prompt_items = await asyncio.to_thread(_create_prompt_sync)
# Cache if enabled
if use_cache:
@@ -256,22 +286,27 @@ class TTSModel:
Returns:
Tuple of (audio_array, sample_rate)
"""
self.load_model()
# Load model (already handles async via to_thread if needed)
await self.load_model_async()
# Set seed if provided
if seed is not None:
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
def _generate_sync():
"""Run synchronous generation in thread pool."""
# Set seed if provided
if seed is not None:
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
# Generate audio
wavs, sample_rate = self.model.generate_voice_clone(
text=text,
voice_clone_prompt=voice_prompt,
instruct=instruct,
)
# Generate audio - this is the blocking operation
wavs, sample_rate = self.model.generate_voice_clone(
text=text,
voice_clone_prompt=voice_prompt,
instruct=instruct,
)
return wavs[0], sample_rate
audio = wavs[0] # Get first result
# Run blocking inference in thread pool to avoid blocking event loop
audio, sample_rate = await asyncio.to_thread(_generate_sync)
return audio, sample_rate
+83 -11
View File
@@ -8,14 +8,18 @@ import threading
class HFProgressTracker:
"""Tracks HuggingFace Hub download progress by intercepting hf_hub_download."""
"""Tracks HuggingFace Hub download progress by intercepting hf_hub_download and snapshot_download."""
def __init__(self, progress_callback: Optional[Callable] = None):
self.progress_callback = progress_callback
self._original_hf_hub_download = None
self._original_snapshot_download = None
self._lock = threading.Lock()
self._total_downloaded = 0
self._total_size = 0
self._file_sizes = {} # Track sizes of individual files
self._file_downloaded = {} # Track downloaded bytes per file
self._current_filename = ""
def _tracked_hf_hub_download(self, *args, **kwargs):
"""Wrapper for hf_hub_download with progress tracking."""
@@ -24,12 +28,53 @@ class HFProgressTracker:
# Get original callback if present
original_resume_callback = kwargs.get("resume_download", None)
# Extract filename if available
filename = kwargs.get("filename", "")
if not filename and len(args) > 1:
filename = args[1] if isinstance(args[1], str) else ""
with self._lock:
self._current_filename = filename
def combined_callback(downloaded: int, total: int):
"""Combined callback that tracks progress."""
# Update totals
# Update per-file tracking
with self._lock:
# Estimate: assume each file contributes equally
# This is a simplification - in reality we'd track per-file
if filename:
self._file_sizes[filename] = total
self._file_downloaded[filename] = downloaded
# Calculate totals across all files
self._total_size = sum(self._file_sizes.values())
self._total_downloaded = sum(self._file_downloaded.values())
# Call original callback if present
if original_resume_callback:
original_resume_callback(downloaded, total)
# Call our progress callback
if self.progress_callback:
with self._lock:
# Pass filename for better progress display
self.progress_callback(self._total_downloaded, self._total_size, filename)
# Replace callback
kwargs["resume_download"] = combined_callback
# Call original download
return self._original_hf_hub_download(*args, **kwargs)
def _tracked_snapshot_download(self, *args, **kwargs):
"""Wrapper for snapshot_download with progress tracking."""
import huggingface_hub
# snapshot_download also uses resume_download callback
original_resume_callback = kwargs.get("resume_download", None)
def combined_callback(downloaded: int, total: int):
"""Combined callback that tracks progress."""
with self._lock:
# For snapshot_download, we track overall progress
if total > 0:
self._total_size = max(self._total_size, total)
self._total_downloaded = downloaded
@@ -41,53 +86,80 @@ class HFProgressTracker:
# Call our progress callback
if self.progress_callback:
with self._lock:
self.progress_callback(self._total_downloaded, self._total_size)
self.progress_callback(self._total_downloaded, self._total_size, "")
# Replace callback
kwargs["resume_download"] = combined_callback
# Call original download
return self._original_hf_hub_download(*args, **kwargs)
return self._original_snapshot_download(*args, **kwargs)
def _tracked_tqdm_update(self, n=1):
"""Track tqdm updates for progress."""
if self._original_tqdm:
# Get current tqdm instance
import tqdm
# Try to get progress info from tqdm
# This is a fallback if hf_hub_download callback doesn't work
pass
@contextmanager
def patch_download(self):
"""Context manager to patch hf_hub_download for progress tracking."""
"""Context manager to patch hf_hub_download and snapshot_download for progress tracking."""
try:
import huggingface_hub
self._original_hf_hub_download = huggingface_hub.hf_hub_download
# Also patch snapshot_download if available (used by from_pretrained)
try:
self._original_snapshot_download = huggingface_hub.snapshot_download
except AttributeError:
self._original_snapshot_download = None
# Reset totals
with self._lock:
self._total_downloaded = 0
self._total_size = 0
self._file_sizes = {}
self._file_downloaded = {}
self._current_filename = ""
# Patch the function
# Patch the functions
huggingface_hub.hf_hub_download = self._tracked_hf_hub_download
if self._original_snapshot_download:
huggingface_hub.snapshot_download = self._tracked_snapshot_download
yield
except ImportError:
# If huggingface_hub not available, just yield without patching
yield
finally:
# Restore original
# Restore original functions
if self._original_hf_hub_download:
try:
import huggingface_hub
huggingface_hub.hf_hub_download = self._original_hf_hub_download
except ImportError:
pass
if self._original_snapshot_download:
try:
import huggingface_hub
huggingface_hub.snapshot_download = self._original_snapshot_download
except (ImportError, AttributeError):
pass
def create_hf_progress_callback(model_name: str, progress_manager):
"""Create a progress callback for HuggingFace downloads."""
def callback(downloaded: int, total: int):
def callback(downloaded: int, total: int, filename: str = ""):
"""Progress callback."""
if total > 0:
progress_manager.update_progress(
model_name=model_name,
current=downloaded,
total=total,
filename="",
filename=filename or "",
status="downloading",
)
return callback