mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
refactor(generation): simplify post-generation cache cleanup
Initialise tts_model before the try so the finally can read its device without a locals() probe, import empty_device_cache alongside the other backend imports instead of inside a bare try/except, and flatten the MPS branch in empty_device_cache.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
38389db82b
commit
17fd1ddd1b
@@ -219,8 +219,7 @@ def empty_device_cache(device: str) -> None:
|
||||
torch.cuda.empty_cache()
|
||||
elif device == "xpu" and hasattr(torch, "xpu"):
|
||||
torch.xpu.empty_cache()
|
||||
elif device == "mps" and hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
if hasattr(torch.mps, "empty_cache"):
|
||||
elif device == "mps" and torch.backends.mps.is_available() and hasattr(torch.mps, "empty_cache"):
|
||||
torch.mps.empty_cache()
|
||||
|
||||
|
||||
|
||||
@@ -54,11 +54,13 @@ async def run_generation(
|
||||
get_tts_backend_for_engine,
|
||||
load_engine_model,
|
||||
)
|
||||
from ..backends.base import empty_device_cache
|
||||
from ..utils.chunked_tts import generate_chunked
|
||||
from ..utils.audio import has_tts_runaway, normalize_audio, save_audio, trim_tts_output
|
||||
|
||||
task_manager = get_task_manager()
|
||||
bg_db = next(get_db())
|
||||
tts_model = None
|
||||
|
||||
try:
|
||||
tts_model = get_tts_backend_for_engine(engine)
|
||||
@@ -156,12 +158,7 @@ async def run_generation(
|
||||
finally:
|
||||
task_manager.complete_generation(generation_id)
|
||||
bg_db.close()
|
||||
try:
|
||||
from ..backends.base import empty_device_cache
|
||||
device = getattr(tts_model, "device", "cpu") if "tts_model" in locals() else "cpu"
|
||||
empty_device_cache(device)
|
||||
except Exception:
|
||||
pass
|
||||
empty_device_cache(getattr(tts_model, "device", "cpu"))
|
||||
|
||||
|
||||
def _notify_speak_end(generation_id: str, *, status: str) -> None:
|
||||
@@ -286,11 +283,13 @@ async def generate_audio_sync(
|
||||
get_tts_backend_for_engine,
|
||||
load_engine_model,
|
||||
)
|
||||
from ..backends.base import empty_device_cache
|
||||
from ..utils.chunked_tts import generate_chunked
|
||||
from ..utils.audio import has_tts_runaway, normalize_audio, trim_tts_output
|
||||
from . import tts
|
||||
|
||||
bg_db = next(get_db())
|
||||
tts_model = None
|
||||
try:
|
||||
tts_model = get_tts_backend_for_engine(engine)
|
||||
await load_engine_model(engine, model_size)
|
||||
@@ -329,12 +328,7 @@ async def generate_audio_sync(
|
||||
|
||||
return tts.audio_to_wav_bytes(audio, sample_rate)
|
||||
finally:
|
||||
try:
|
||||
from ..backends.base import empty_device_cache
|
||||
device = getattr(tts_model, "device", "cpu") if "tts_model" in locals() else "cpu"
|
||||
empty_device_cache(device)
|
||||
except Exception:
|
||||
pass
|
||||
empty_device_cache(getattr(tts_model, "device", "cpu"))
|
||||
|
||||
|
||||
def _save_regenerate(
|
||||
|
||||
Reference in New Issue
Block a user