Add: Test for ONNX inference code (without PyTorch dependency)

This commit is contained in:
tsukumi
2024-09-23 10:04:02 +09:00
parent 5a82105df3
commit eb5d245903
6 changed files with 104 additions and 22 deletions

View File

@@ -8,11 +8,12 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
"""
from __future__ import annotations
import gc
import time
from typing import Optional, Union, cast
from typing import TYPE_CHECKING, Optional, Union, cast
import torch
from transformers import (
AutoModelForMaskedLM,
AutoTokenizer,
@@ -27,6 +28,10 @@ from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
from style_bert_vits2.logging import logger
if TYPE_CHECKING:
import torch
# 各言語ごとのロード済みの BERT モデルを格納する辞書
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
@@ -208,6 +213,8 @@ def unload_model(language: Languages) -> None:
language (Languages): アンロードする BERT モデルの言語
"""
import torch
if language in __loaded_models:
del __loaded_models[language]
gc.collect()
@@ -224,6 +231,8 @@ def unload_tokenizer(language: Languages) -> None:
language (Languages): アンロードする BERT トークナイザーの言語
"""
import torch
if language in __loaded_tokenizers:
del __loaded_tokenizers[language]
gc.collect()

View File

@@ -154,6 +154,16 @@ class TTSModel:
f"Model loaded successfully from {self.model_path} to {self.onnx_session.get_providers()[0]} ({time.time() - start_time:.2f}s)"
)
def unload(self) -> None:
"""
音声合成モデルをデバイスからアンロードする。
"""
if self.net_g is not None:
self.net_g = None
if self.onnx_session is not None:
self.onnx_session = None
def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
"""
スタイルベクトルを取得する。