mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-02 16:45:15 -07:00
feat: add per-model unload endpoint and UI button
- POST /models/{model_name}/unload — unloads a specific model from
memory without deleting from disk, supports all engine types
- Frontend: Unload button in model detail dialog when model is loaded
- Delete button remains disabled while loaded (unload first)
This commit is contained in:
+63
-1
@@ -1479,7 +1479,7 @@ async def load_model(model_size: str = "1.7B"):
|
||||
|
||||
@app.post("/models/unload")
|
||||
async def unload_model():
|
||||
"""Unload TTS model to free memory."""
|
||||
"""Unload the default Qwen TTS model to free memory."""
|
||||
try:
|
||||
tts.unload_tts_model()
|
||||
return {"message": "Model unloaded successfully"}
|
||||
@@ -1487,6 +1487,68 @@ async def unload_model():
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.post("/models/{model_name}/unload")
|
||||
async def unload_model_by_name(model_name: str):
|
||||
"""Unload a specific model from memory without deleting it from disk."""
|
||||
# Map of model_name -> (model_type, model_size)
|
||||
model_types = {
|
||||
"qwen-tts-1.7B": ("tts", "1.7B"),
|
||||
"qwen-tts-0.6B": ("tts", "0.6B"),
|
||||
"luxtts": ("luxtts", "default"),
|
||||
"chatterbox-tts": ("chatterbox", "default"),
|
||||
"chatterbox-turbo": ("chatterbox_turbo", "default"),
|
||||
"whisper-base": ("whisper", "base"),
|
||||
"whisper-small": ("whisper", "small"),
|
||||
"whisper-medium": ("whisper", "medium"),
|
||||
"whisper-large": ("whisper", "large"),
|
||||
"whisper-turbo": ("whisper", "turbo"),
|
||||
}
|
||||
|
||||
if model_name not in model_types:
|
||||
raise HTTPException(status_code=400, detail=f"Unknown model: {model_name}")
|
||||
|
||||
model_type, model_size = model_types[model_name]
|
||||
|
||||
try:
|
||||
if model_type == "tts":
|
||||
tts_model = tts.get_tts_model()
|
||||
if tts_model.is_loaded() and tts_model.model_size == model_size:
|
||||
tts.unload_tts_model()
|
||||
else:
|
||||
return {"message": f"Model {model_name} is not loaded"}
|
||||
elif model_type == "luxtts":
|
||||
from .backends import get_tts_backend_for_engine
|
||||
backend = get_tts_backend_for_engine("luxtts")
|
||||
if backend.is_loaded():
|
||||
backend.unload_model()
|
||||
else:
|
||||
return {"message": f"Model {model_name} is not loaded"}
|
||||
elif model_type == "chatterbox":
|
||||
from .backends import get_tts_backend_for_engine
|
||||
backend = get_tts_backend_for_engine("chatterbox")
|
||||
if backend.is_loaded():
|
||||
backend.unload_model()
|
||||
else:
|
||||
return {"message": f"Model {model_name} is not loaded"}
|
||||
elif model_type == "chatterbox_turbo":
|
||||
from .backends import get_tts_backend_for_engine
|
||||
backend = get_tts_backend_for_engine("chatterbox_turbo")
|
||||
if backend.is_loaded():
|
||||
backend.unload_model()
|
||||
else:
|
||||
return {"message": f"Model {model_name} is not loaded"}
|
||||
elif model_type == "whisper":
|
||||
whisper_model = transcribe.get_whisper_model()
|
||||
if whisper_model.is_loaded() and whisper_model.model_size == model_size:
|
||||
transcribe.unload_whisper_model()
|
||||
else:
|
||||
return {"message": f"Model {model_name} is not loaded"}
|
||||
|
||||
return {"message": f"Model {model_name} unloaded successfully"}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.get("/models/progress/{model_name}")
|
||||
async def get_model_progress(model_name: str):
|
||||
"""Get model download progress via Server-Sent Events."""
|
||||
|
||||
Reference in New Issue
Block a user