Files
voicebox/docs/developer/model-management.mdx
Jamie Pine 7b5e73cfa8 Add .npmrc for bun usage and update dependencies
- Created a new .npmrc file to enforce bun usage.
- Bumped version numbers for multiple packages to 0.1.9 in bun.lock.
- Added react-sound-visualizer dependency to enhance audio visualization features.
- Introduced convert:assets script in package.json for asset optimization.
- Updated CONTRIBUTING.md with instructions for converting assets to web formats.
- Added documentation files for API endpoints and developer guidelines in the docs directory.
2026-01-29 18:56:10 -08:00

342 lines
8.6 KiB
Plaintext

---
title: "Model Management"
description: "How model downloading, loading, and status tracking works in Voicebox"
---
## Overview
Voicebox manages two types of models:
**TTS Models:** Qwen3-TTS for voice cloning (0.6B and 1.7B variants).
**ASR Models:** Whisper for transcription (tiny through large).
Models are downloaded from HuggingFace Hub on first use and cached locally.
## Available Models
### TTS Models
| Model | HuggingFace ID | Size | VRAM |
|-------|----------------|------|------|
| 0.6B | `Qwen/Qwen3-TTS-12Hz-0.6B-Base` | ~1.2GB | ~2GB |
| 1.7B | `Qwen/Qwen3-TTS-12Hz-1.7B-Base` | ~3.4GB | ~6GB |
### Whisper Models
| Model | HuggingFace ID | Size | VRAM |
|-------|----------------|------|------|
| tiny | `openai/whisper-tiny` | ~150MB | ~1GB |
| base | `openai/whisper-base` | ~300MB | ~1GB |
| small | `openai/whisper-small` | ~500MB | ~2GB |
| medium | `openai/whisper-medium` | ~1.5GB | ~5GB |
| large | `openai/whisper-large` | ~3GB | ~10GB |
## Model Storage
Models are cached in the HuggingFace cache directory:
```
~/.cache/huggingface/hub/
├── models--Qwen--Qwen3-TTS-12Hz-1.7B-Base/
├── models--Qwen--Qwen3-TTS-12Hz-0.6B-Base/
├── models--openai--whisper-base/
└── ...
```
## Progress Tracking
### Progress Manager
Tracks download progress across all models:
```python
class ProgressManager:
def __init__(self):
self._progress = {} # model_name -> progress_info
def update_progress(
self,
model_name: str,
current: int,
total: int,
filename: str,
status: str,
):
self._progress[model_name] = {
"current": current,
"total": total,
"filename": filename,
"status": status, # downloading, complete, error
"updated_at": datetime.utcnow(),
}
def get_progress(self, model_name: str) -> Optional[dict]:
return self._progress.get(model_name)
```
### HuggingFace Progress Callback
Hooks into HuggingFace's download system:
```python
class HFProgressTracker:
def __init__(self, callback):
self.callback = callback
@contextmanager
def patch_download(self):
"""Context manager to intercept HF downloads."""
original_download = hf_hub_download
def patched_download(*args, **kwargs):
# Intercept progress
result = original_download(*args, **kwargs)
self.callback(progress_info)
return result
# Apply patch
with patch('huggingface_hub.hf_hub_download', patched_download):
yield
```
### Server-Sent Events (SSE)
Progress is streamed to the frontend:
```python
@app.get("/models/progress/{model_name}")
async def get_model_progress(model_name: str):
async def event_generator():
while True:
progress = progress_manager.get_progress(model_name)
if progress:
yield f"data: {json.dumps(progress)}\n\n"
if progress and progress["status"] in ["complete", "error"]:
break
await asyncio.sleep(0.5)
return StreamingResponse(
event_generator(),
media_type="text/event-stream"
)
```
## Task Manager
Tracks active downloads and generations:
```python
class TaskManager:
def __init__(self):
self._active_downloads = {}
self._active_generations = {}
def start_download(self, model_name: str):
self._active_downloads[model_name] = {
"status": "downloading",
"started_at": datetime.utcnow(),
}
def complete_download(self, model_name: str):
if model_name in self._active_downloads:
del self._active_downloads[model_name]
def get_active_tasks(self) -> dict:
return {
"downloads": list(self._active_downloads.values()),
"generations": list(self._active_generations.values()),
}
```
## Model Status
Check which models are downloaded and loaded:
```python
@app.get("/models/status")
async def get_model_status() -> ModelStatusListResponse:
models = []
# Check TTS models
for size, hf_id in [("1.7B", "Qwen/Qwen3-TTS-12Hz-1.7B-Base"), ...]:
downloaded = is_model_downloaded(hf_id)
loaded = tts_model._current_model_size == size
models.append(ModelStatus(
model_name=f"qwen-tts-{size}",
display_name=f"Qwen3-TTS {size}",
downloaded=downloaded,
size_mb=get_model_size_mb(hf_id),
loaded=loaded,
))
# Check Whisper models
for size in ["tiny", "base", "small", "medium", "large"]:
hf_id = f"openai/whisper-{size}"
downloaded = is_model_downloaded(hf_id)
models.append(ModelStatus(
model_name=f"whisper-{size}",
display_name=f"Whisper {size}",
downloaded=downloaded,
size_mb=get_model_size_mb(hf_id),
loaded=False, # Whisper is loaded on-demand
))
return ModelStatusListResponse(models=models)
```
## Manual Model Operations
### Load Model
```python
@app.post("/models/load")
async def load_model(model_size: str = "1.7B"):
tts_model = get_tts_model()
await tts_model.load_model_async(model_size)
return {"status": "loaded", "model_size": model_size}
```
### Unload Model
```python
@app.post("/models/unload")
async def unload_model():
tts_model = get_tts_model()
tts_model.unload_model()
return {"status": "unloaded"}
```
### Trigger Download
```python
@app.post("/models/download")
async def trigger_model_download(request: ModelDownloadRequest):
# This triggers the download in background
# Progress is tracked via /models/progress/{model_name}
if request.model_name.startswith("qwen-tts"):
size = request.model_name.split("-")[-1]
asyncio.create_task(download_tts_model(size))
elif request.model_name.startswith("whisper"):
size = request.model_name.split("-")[-1]
asyncio.create_task(download_whisper_model(size))
return {"status": "downloading"}
```
### Delete Model
```python
@app.delete("/models/{model_name}")
async def delete_model(model_name: str):
# Find and delete from HuggingFace cache
cache_dir = Path.home() / ".cache" / "huggingface" / "hub"
model_dirs = list(cache_dir.glob(f"models--*--{model_name}*"))
for model_dir in model_dirs:
shutil.rmtree(model_dir)
return {"status": "deleted"}
```
## API Endpoints
| Method | Endpoint | Description |
|--------|----------|-------------|
| GET | `/models/status` | Get status of all models |
| POST | `/models/load` | Load TTS model |
| POST | `/models/unload` | Unload TTS model |
| POST | `/models/download` | Trigger model download |
| GET | `/models/progress/{name}` | Stream download progress (SSE) |
| DELETE | `/models/{name}` | Delete downloaded model |
| GET | `/tasks/active` | Get active downloads/generations |
## Response Schemas
### ModelStatus
```json
{
"model_name": "qwen-tts-1.7B",
"display_name": "Qwen3-TTS 1.7B",
"downloaded": true,
"size_mb": 3400,
"loaded": true
}
```
### ActiveTasksResponse
```json
{
"downloads": [
{
"model_name": "whisper-medium",
"status": "downloading",
"started_at": "2024-01-15T10:30:00Z"
}
],
"generations": [
{
"task_id": "uuid",
"profile_id": "uuid",
"text_preview": "Hello world...",
"started_at": "2024-01-15T10:30:00Z"
}
]
}
```
## Frontend Integration
### Progress Display
```typescript
// Subscribe to download progress via SSE
const eventSource = new EventSource(`/models/progress/${modelName}`);
eventSource.onmessage = (event) => {
const progress = JSON.parse(event.data);
updateProgressBar(progress.current / progress.total);
if (progress.status === 'complete') {
eventSource.close();
}
};
```
### Model Status UI
```typescript
// Fetch model status
const { data: models } = useQuery({
queryKey: ['models', 'status'],
queryFn: () => api.getModelStatus(),
});
// Display download/load buttons based on status
models.map(model => (
<ModelCard
name={model.display_name}
downloaded={model.downloaded}
loaded={model.loaded}
onDownload={() => triggerDownload(model.model_name)}
onLoad={() => loadModel(model.model_name)}
/>
));
```
## Error Handling
| Error | Cause | Solution |
|-------|-------|----------|
| Download failed | Network issue | Retry download |
| OOM on load | Model too large | Use smaller model |
| Model not found | Cache corrupted | Re-download |
| Slow download | HF rate limit | Wait and retry |