fix latent bugs
This commit is contained in:
@@ -23,13 +23,18 @@ async def main() -> None:
|
|||||||
parser.add_argument("--model-dir", type=Path, required=True,
|
parser.add_argument("--model-dir", type=Path, required=True,
|
||||||
help="Directory containing model files (config.json, *.safetensors, style_vectors.npy)")
|
help="Directory containing model files (config.json, *.safetensors, style_vectors.npy)")
|
||||||
parser.add_argument("--device", default="cpu",
|
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",
|
parser.add_argument("--debug", action="store_true",
|
||||||
help="Log DEBUG messages")
|
help="Log DEBUG messages")
|
||||||
parser.add_argument("--version", action="version",
|
parser.add_argument("--version", action="version",
|
||||||
version=__version__)
|
version=__version__)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
device_map = {
|
||||||
|
"rocm": "hip",
|
||||||
|
}
|
||||||
|
device = device_map.get(args.device, args.device)
|
||||||
|
|
||||||
logging.basicConfig(
|
logging.basicConfig(
|
||||||
level=logging.DEBUG if args.debug else logging.INFO,
|
level=logging.DEBUG if args.debug else logging.INFO,
|
||||||
format="%(asctime)s %(levelname)s %(name)s %(message)s",
|
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("Starting GLaDOS Wyoming TTS server on %s", args.uri)
|
||||||
_LOGGER.info("Model directory: %s", model_dir)
|
_LOGGER.info("Model directory: %s", model_dir)
|
||||||
_LOGGER.info("Device: %s", args.device)
|
_LOGGER.info("Device: %s", device)
|
||||||
|
|
||||||
server_task = asyncio.create_task(
|
server_task = asyncio.create_task(
|
||||||
server.run(
|
server.run(
|
||||||
@@ -80,7 +85,7 @@ async def main() -> None:
|
|||||||
GLaDOSEventHandler,
|
GLaDOSEventHandler,
|
||||||
wyoming_info,
|
wyoming_info,
|
||||||
model_dir,
|
model_dir,
|
||||||
args.device,
|
device,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -41,11 +41,9 @@ def _detect_language(text: str) -> Languages:
|
|||||||
|
|
||||||
def _load_bert_for_language(language: Languages, device: str) -> None:
|
def _load_bert_for_language(language: Languages, device: str) -> None:
|
||||||
model_name = _BERT_MODEL_NAMES[language]
|
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)
|
||||||
_LOGGER.info("Loading BERT model for %s (%s)", language.name, model_name)
|
bert_models.load_model(language, model_name)
|
||||||
bert_models.load_model(language, model_name)
|
bert_models.load_tokenizer(language, model_name)
|
||||||
if not bert_models.is_tokenizer_loaded(language):
|
|
||||||
bert_models.load_tokenizer(language, model_name)
|
|
||||||
bert = bert_models.__loaded_models.get(language)
|
bert = bert_models.__loaded_models.get(language)
|
||||||
if bert is not None:
|
if bert is not None:
|
||||||
bert = bert.float()
|
bert = bert.float()
|
||||||
|
|||||||
Reference in New Issue
Block a user