From 4ad64ad1a8d2c69dda85cc2465c7967ae49cf4c9 Mon Sep 17 00:00:00 2001 From: Stephen Tafoya Date: Sat, 13 Jun 2026 02:20:24 -0600 Subject: [PATCH] fix latent bugs --- wyoming_glados/__main__.py | 11 ++++++++--- wyoming_glados/handler.py | 8 +++----- 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/wyoming_glados/__main__.py b/wyoming_glados/__main__.py index 23fbebd..f6e25d7 100644 --- a/wyoming_glados/__main__.py +++ b/wyoming_glados/__main__.py @@ -23,13 +23,18 @@ async def main() -> None: parser.add_argument("--model-dir", type=Path, required=True, help="Directory containing model files (config.json, *.safetensors, style_vectors.npy)") parser.add_argument("--device", default="cpu", - help="Device for PyTorch (cpu, cuda)") + help="Device for PyTorch (cpu, cuda, rocm)") parser.add_argument("--debug", action="store_true", help="Log DEBUG messages") parser.add_argument("--version", action="version", version=__version__) args = parser.parse_args() + device_map = { + "rocm": "hip", + } + device = device_map.get(args.device, args.device) + logging.basicConfig( level=logging.DEBUG if args.debug else logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s", @@ -72,7 +77,7 @@ async def main() -> None: _LOGGER.info("Starting GLaDOS Wyoming TTS server on %s", args.uri) _LOGGER.info("Model directory: %s", model_dir) - _LOGGER.info("Device: %s", args.device) + _LOGGER.info("Device: %s", device) server_task = asyncio.create_task( server.run( @@ -80,7 +85,7 @@ async def main() -> None: GLaDOSEventHandler, wyoming_info, model_dir, - args.device, + device, ) ) ) diff --git a/wyoming_glados/handler.py b/wyoming_glados/handler.py index 67e9301..e55366a 100644 --- a/wyoming_glados/handler.py +++ b/wyoming_glados/handler.py @@ -41,11 +41,9 @@ def _detect_language(text: str) -> Languages: def _load_bert_for_language(language: Languages, device: str) -> None: model_name = _BERT_MODEL_NAMES[language] - if not bert_models.is_model_loaded(language): - _LOGGER.info("Loading BERT model for %s (%s)", language.name, model_name) - bert_models.load_model(language, model_name) - if not bert_models.is_tokenizer_loaded(language): - bert_models.load_tokenizer(language, model_name) + _LOGGER.info("Loading BERT model for %s (%s)", language.name, model_name) + bert_models.load_model(language, model_name) + bert_models.load_tokenizer(language, model_name) bert = bert_models.__loaded_models.get(language) if bert is not None: bert = bert.float()