fixed a lot of things and actually got it to work
This commit is contained in:
@@ -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