diff --git a/backend/engines/confucius4/main.py b/backend/engines/confucius4/main.py index 2b188f9b..3983bf6a 100644 --- a/backend/engines/confucius4/main.py +++ b/backend/engines/confucius4/main.py @@ -115,7 +115,9 @@ def _load_model(stdout): import torch from confuciustts.cli.inference import ConfuciusTTS # type: ignore[import-not-found] - device = "cuda" if torch.cuda.is_available() else "cpu" + device = torch.accelerator.current_accelerator().type # 'cuda', 'npu', 'mps', 'xpu', 'cpu' + if device == "mps": + device = "cpu" # ConfuciusTTS is untested on MPS; fall back to CPU for safety _send(stdout, {"op": "progress", "stage": "loading_model", "percent": 50}) _model = ConfuciusTTS(config_path=_config_path(), device=device) diff --git a/backend/engines/dots_tts/main.py b/backend/engines/dots_tts/main.py index 36a490b1..19d055d1 100644 --- a/backend/engines/dots_tts/main.py +++ b/backend/engines/dots_tts/main.py @@ -115,7 +115,7 @@ def _load_runtime(stdout): from dots_tts.runtime import DotsTtsRuntime # type: ignore[import-not-found] repo = os.environ.get("OMNIVOICE_DOTS_TTS_MODEL", _DEFAULT_REPO) - default_precision = "bfloat16" if torch.cuda.is_available() else "float32" + default_precision = "bfloat16" if torch.accelerator.current_accelerator().type != "cpu" else "float32" precision = os.environ.get("OMNIVOICE_DOTS_TTS_PRECISION", default_precision) optimize = os.environ.get("OMNIVOICE_DOTS_TTS_OPTIMIZE", "0") == "1"