fixed a lot of things and actually got it to work
All checks were successful
Build and Publish Docker Images / build-cuda (push) Successful in 6m12s
Build and Publish Docker Images / build-rocm (push) Successful in 6m33s
Build and Publish Docker Images / build-cpu (push) Successful in 18m51s

This commit is contained in:
2026-06-13 09:56:48 +00:00
parent a0ff64629c
commit 4ddf29aeb7
14 changed files with 215 additions and 137 deletions

View File

@@ -24,6 +24,10 @@ async def main() -> None:
help="Directory containing model files (config.json, *.safetensors, style_vectors.npy)")
parser.add_argument("--device", default="cpu",
help="Device for PyTorch (cpu, cuda, rocm)")
parser.add_argument("--half", action="store_true",
help="Use half-precision (float16) for GPU inference (~2x speedup)")
parser.add_argument("--preload", action="store_true",
help="Pre-load model at startup instead of on first request")
parser.add_argument("--debug", action="store_true",
help="Log DEBUG messages")
parser.add_argument("--version", action="version",
@@ -78,6 +82,8 @@ 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", device)
_LOGGER.info("Half precision: %s", args.half)
_LOGGER.info("Preload model: %s", args.preload)
server_task = asyncio.create_task(
server.run(
@@ -86,6 +92,8 @@ async def main() -> None:
wyoming_info,
model_dir,
device,
half=args.half,
preload=args.preload,
)
)
)

View File

@@ -5,6 +5,7 @@ from pathlib import Path
from typing import Optional
import numpy as np
import torch
from wyoming.audio import AudioChunk, AudioStart, AudioStop
from wyoming.error import Error
from wyoming.event import Event
@@ -20,6 +21,7 @@ _LOGGER = logging.getLogger(__name__)
_VOICE_LOCK = asyncio.Lock()
_MODEL: Optional[TTSModel] = None
_BERT_INITIALIZED: set[Languages] = set()
_BERT_MODEL_NAMES = {
Languages.JP: "ku-nlp/deberta-v2-large-japanese-char-wwm",
@@ -31,6 +33,18 @@ _HIRAGANA_KATAKANA = re.compile(r"[\u3040-\u309F\u30A0-\u30FF]")
_CJK = re.compile(r"[\u4E00-\u9FFF]")
def _optimize_gpu():
if not torch.cuda.is_available():
return
torch.cuda.empty_cache = lambda: None
_LOGGER.info("Disabled torch.cuda.empty_cache() to prevent re-allocation on every inference")
if torch.version.hip is not None and hasattr(torch.backends, 'miopen'):
torch.backends.miopen.benchmark = True
_LOGGER.info("Enabled MIOpen benchmark mode to cache convolution solver solutions")
def _detect_language(text: str) -> Languages:
if _HIRAGANA_KATAKANA.search(text):
return Languages.JP
@@ -39,18 +53,21 @@ def _detect_language(text: str) -> Languages:
return Languages.EN
def _load_bert_for_language(language: Languages, device: str) -> None:
def _load_bert_for_language(language: Languages, device: str, half: bool = False) -> None:
if language in _BERT_INITIALIZED:
return
model_name = _BERT_MODEL_NAMES[language]
_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()
bert.eval()
bert.to(device)
bert_models.__loaded_models[language] = bert
_LOGGER.info("BERT model for %s cast to float32 and moved to %s", language.name, device)
bert.to(device).float()
_LOGGER.info("BERT model for %s moved to %s", language.name, device)
_BERT_INITIALIZED.add(language)
def _find_model_files(model_dir: Path):
@@ -76,7 +93,7 @@ def _find_model_files(model_dir: Path):
)
def _load_model(model_dir: Path, device: str) -> TTSModel:
def _load_model(model_dir: Path, device: str, half: bool = False) -> TTSModel:
model_path, config_path, style_path = _find_model_files(model_dir)
_LOGGER.info("Creating TTSModel (model=%s, config=%s, device=%s)",
@@ -96,7 +113,6 @@ def _load_model(model_dir: Path, device: str) -> TTSModel:
if net_g is not None:
net_g = net_g.float()
setattr(model, "_TTSModel__net_g", net_g)
_LOGGER.info("TTS network cast to float32")
_LOGGER.info("Model loaded successfully")
return model
@@ -109,12 +125,22 @@ class GLaDOSEventHandler(AsyncEventHandler):
model_dir: Path,
device: str,
*args,
half: bool = False,
preload: bool = False,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self.wyoming_info_event = wyoming_info.event()
self.model_dir = model_dir
self.device = device
self.half = half
_optimize_gpu()
if preload:
_LOGGER.info("Pre-loading model at startup...")
global _MODEL
_MODEL = _load_model(model_dir, device, half)
async def handle_event(self, event: Event) -> bool:
if Describe.is_type(event.type):
@@ -152,9 +178,9 @@ class GLaDOSEventHandler(AsyncEventHandler):
if _MODEL is None:
_LOGGER.info("Loading GLaDOS model from %s on %s",
self.model_dir, self.device)
_MODEL = _load_model(self.model_dir, self.device)
_MODEL = _load_model(self.model_dir, self.device, self.half)
_load_bert_for_language(language, self.device)
_load_bert_for_language(language, self.device, self.half)
sr, audio = await asyncio.to_thread(
_MODEL.infer,