fixed a lot of things and actually got it to work
This commit is contained in:
@@ -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,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user