Merge pull request #165 from tsukumijima/master
ONNX への変換と ONNXRuntime による推論サポートを追加
This commit is contained in:
5
.gitignore
vendored
5
.gitignore
vendored
@@ -2,9 +2,11 @@ __pycache__/
|
|||||||
venv/
|
venv/
|
||||||
.venv/
|
.venv/
|
||||||
dist/
|
dist/
|
||||||
.coverage
|
.coverage*
|
||||||
.ipynb_checkpoints/
|
.ipynb_checkpoints/
|
||||||
.ruff_cache/
|
.ruff_cache/
|
||||||
|
.DS_Store
|
||||||
|
._*
|
||||||
|
|
||||||
/*.yml
|
/*.yml
|
||||||
!/default_config.yml
|
!/default_config.yml
|
||||||
@@ -13,6 +15,7 @@ dist/
|
|||||||
/bert/*/*.model
|
/bert/*/*.model
|
||||||
/bert/*/*.safetensors
|
/bert/*/*.safetensors
|
||||||
/bert/*/*.msgpack
|
/bert/*/*.msgpack
|
||||||
|
/bert/*/*.onnx
|
||||||
|
|
||||||
/configs/paths.yml
|
/configs/paths.yml
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ chcp 65001 > NUL
|
|||||||
|
|
||||||
pushd %~dp0
|
pushd %~dp0
|
||||||
echo Running gradio_tabs/dataset.py...
|
echo Running gradio_tabs/dataset.py...
|
||||||
venv\Scripts\python gradio_tabs/dataset.py
|
venv\Scripts\python -m gradio_tabs.dataset
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ chcp 65001 > NUL
|
|||||||
|
|
||||||
pushd %~dp0
|
pushd %~dp0
|
||||||
echo Running gradio_tabs/inference.py...
|
echo Running gradio_tabs/inference.py...
|
||||||
venv\Scripts\python gradio_tabs/inference.py
|
venv\Scripts\python -m gradio_tabs.inference
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ chcp 65001 > NUL
|
|||||||
|
|
||||||
pushd %~dp0
|
pushd %~dp0
|
||||||
echo Running gradio_tabs/merge.py...
|
echo Running gradio_tabs/merge.py...
|
||||||
venv\Scripts\python gradio_tabs/merge.py
|
venv\Scripts\python -m gradio_tabs.merge
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,8 @@ You can install via `pip install style-bert-vits2` (inference only), see [librar
|
|||||||
- [Zennの解説記事](https://zenn.dev/litagin/articles/034819a5256ff4)
|
- [Zennの解説記事](https://zenn.dev/litagin/articles/034819a5256ff4)
|
||||||
|
|
||||||
- [**リリースページ**](https://github.com/litagin02/Style-Bert-VITS2/releases/)、[更新履歴](/docs/CHANGELOG.md)
|
- [**リリースページ**](https://github.com/litagin02/Style-Bert-VITS2/releases/)、[更新履歴](/docs/CHANGELOG.md)
|
||||||
- 2024-06-16: Ver 2.6.0 (モデルの差分マージ・加重マージ・ヌルモデルマージの追加)
|
- 2024-09-09: Ver 2.6.1: Google colabでうまく学習できない等のバグ修正のみ
|
||||||
|
- 2024-06-16: Ver 2.6.0 (モデルの差分マージ・加重マージ・ヌルモデルマージの追加、使い道については[この記事](https://zenn.dev/litagin/articles/1297b1dc7bdc79)参照)
|
||||||
- 2024-06-14: Ver 2.5.1 (利用規約をお願いへ変更したのみ)
|
- 2024-06-14: Ver 2.5.1 (利用規約をお願いへ変更したのみ)
|
||||||
- 2024-06-02: Ver 2.5.0 (**[利用規約](/docs/TERMS_OF_USE.md)の追加**、フォルダ分けからのスタイル生成、小春音アミ・あみたろモデルの追加、インストールの高速化等)
|
- 2024-06-02: Ver 2.5.0 (**[利用規約](/docs/TERMS_OF_USE.md)の追加**、フォルダ分けからのスタイル生成、小春音アミ・あみたろモデルの追加、インストールの高速化等)
|
||||||
- 2024-03-16: ver 2.4.1 (**batファイルによるインストール方法の変更**)
|
- 2024-03-16: ver 2.4.1 (**batファイルによるインストール方法の変更**)
|
||||||
@@ -78,9 +79,9 @@ powershell -c "irm https://astral.sh/uv/install.ps1 | iex"
|
|||||||
git clone https://github.com/litagin02/Style-Bert-VITS2.git
|
git clone https://github.com/litagin02/Style-Bert-VITS2.git
|
||||||
cd Style-Bert-VITS2
|
cd Style-Bert-VITS2
|
||||||
uv venv venv
|
uv venv venv
|
||||||
uv pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu118
|
|
||||||
uv pip install -r requirements.txt
|
|
||||||
venv\Scripts\activate
|
venv\Scripts\activate
|
||||||
|
uv pip install "torch<2.4" "torchaudio<2.4" --index-url https://download.pytorch.org/whl/cu118
|
||||||
|
uv pip install -r requirements.txt
|
||||||
python initialize.py # 必要なモデルとデフォルトTTSモデルをダウンロード
|
python initialize.py # 必要なモデルとデフォルトTTSモデルをダウンロード
|
||||||
```
|
```
|
||||||
最後を忘れずに。
|
最後を忘れずに。
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ chcp 65001 > NUL
|
|||||||
|
|
||||||
pushd %~dp0
|
pushd %~dp0
|
||||||
echo Running gradio_tabs/style_vectors.py...
|
echo Running gradio_tabs/style_vectors.py...
|
||||||
venv\Scripts\python gradio_tabs/style_vectors.py
|
venv\Scripts\python -m gradio_tabs.style_vectors
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ chcp 65001 > NUL
|
|||||||
|
|
||||||
pushd %~dp0
|
pushd %~dp0
|
||||||
echo Running gradio_tabs/train.py...
|
echo Running gradio_tabs/train.py...
|
||||||
venv\Scripts\python gradio_tabs/train.py
|
venv\Scripts\python -m gradio_tabs.train
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
|||||||
5
app.py
5
app.py
@@ -14,6 +14,7 @@ from style_bert_vits2.constants import GRADIO_THEME, VERSION
|
|||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
||||||
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
from style_bert_vits2.tts_model import TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModelHolder
|
||||||
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
||||||
@@ -40,7 +41,9 @@ if device == "cuda" and not torch.cuda.is_available():
|
|||||||
# download_default_models()
|
# download_default_models()
|
||||||
|
|
||||||
path_config = get_path_config()
|
path_config = get_path_config()
|
||||||
model_holder = TTSModelHolder(Path(path_config.assets_root), device)
|
model_holder = TTSModelHolder(
|
||||||
|
Path(path_config.assets_root), device, torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
|
|
||||||
with gr.Blocks(theme=GRADIO_THEME) as app:
|
with gr.Blocks(theme=GRADIO_THEME) as app:
|
||||||
gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})")
|
gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})")
|
||||||
|
|||||||
@@ -3,12 +3,24 @@
|
|||||||
"repo_id": "ku-nlp/deberta-v2-large-japanese-char-wwm",
|
"repo_id": "ku-nlp/deberta-v2-large-japanese-char-wwm",
|
||||||
"files": ["pytorch_model.bin"]
|
"files": ["pytorch_model.bin"]
|
||||||
},
|
},
|
||||||
|
"deberta-v2-large-japanese-char-wwm-onnx": {
|
||||||
|
"repo_id": "tsukumijima/deberta-v2-large-japanese-char-wwm-onnx",
|
||||||
|
"files": ["model.onnx"]
|
||||||
|
},
|
||||||
"chinese-roberta-wwm-ext-large": {
|
"chinese-roberta-wwm-ext-large": {
|
||||||
"repo_id": "hfl/chinese-roberta-wwm-ext-large",
|
"repo_id": "hfl/chinese-roberta-wwm-ext-large",
|
||||||
"files": ["pytorch_model.bin"]
|
"files": ["pytorch_model.bin"]
|
||||||
},
|
},
|
||||||
|
"chinese-roberta-wwm-ext-large-onnx": {
|
||||||
|
"repo_id": "tsukumijima/chinese-roberta-wwm-ext-large-onnx",
|
||||||
|
"files": ["model.onnx"]
|
||||||
|
},
|
||||||
"deberta-v3-large": {
|
"deberta-v3-large": {
|
||||||
"repo_id": "microsoft/deberta-v3-large",
|
"repo_id": "microsoft/deberta-v3-large",
|
||||||
"files": ["spm.model", "pytorch_model.bin"]
|
"files": ["spm.model", "pytorch_model.bin"]
|
||||||
|
},
|
||||||
|
"deberta-v3-large-onnx": {
|
||||||
|
"repo_id": "tsukumijima/deberta-v3-large-onnx",
|
||||||
|
"files": ["model.onnx"]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
{}
|
||||||
28
bert/chinese-roberta-wwm-ext-large-onnx/config.json
Normal file
28
bert/chinese-roberta-wwm-ext-large-onnx/config.json
Normal file
@@ -0,0 +1,28 @@
|
|||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"BertForMaskedLM"
|
||||||
|
],
|
||||||
|
"attention_probs_dropout_prob": 0.1,
|
||||||
|
"bos_token_id": 0,
|
||||||
|
"directionality": "bidi",
|
||||||
|
"eos_token_id": 2,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_dropout_prob": 0.1,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"layer_norm_eps": 1e-12,
|
||||||
|
"max_position_embeddings": 512,
|
||||||
|
"model_type": "bert",
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"output_past": true,
|
||||||
|
"pad_token_id": 0,
|
||||||
|
"pooler_fc_size": 768,
|
||||||
|
"pooler_num_attention_heads": 12,
|
||||||
|
"pooler_num_fc_layers": 3,
|
||||||
|
"pooler_size_per_head": 128,
|
||||||
|
"pooler_type": "first_token_transform",
|
||||||
|
"type_vocab_size": 2,
|
||||||
|
"vocab_size": 21128
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
{"unk_token": "[UNK]", "sep_token": "[SEP]", "pad_token": "[PAD]", "cls_token": "[CLS]", "mask_token": "[MASK]"}
|
||||||
21232
bert/chinese-roberta-wwm-ext-large-onnx/tokenizer.json
Normal file
21232
bert/chinese-roberta-wwm-ext-large-onnx/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1 @@
|
|||||||
|
{"init_inputs": []}
|
||||||
21128
bert/chinese-roberta-wwm-ext-large-onnx/vocab.txt
Normal file
21128
bert/chinese-roberta-wwm-ext-large-onnx/vocab.txt
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,9 +0,0 @@
|
|||||||
*.bin.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,57 +0,0 @@
|
|||||||
---
|
|
||||||
language:
|
|
||||||
- zh
|
|
||||||
tags:
|
|
||||||
- bert
|
|
||||||
license: "apache-2.0"
|
|
||||||
---
|
|
||||||
|
|
||||||
# Please use 'Bert' related functions to load this model!
|
|
||||||
|
|
||||||
## Chinese BERT with Whole Word Masking
|
|
||||||
For further accelerating Chinese natural language processing, we provide **Chinese pre-trained BERT with Whole Word Masking**.
|
|
||||||
|
|
||||||
**[Pre-Training with Whole Word Masking for Chinese BERT](https://arxiv.org/abs/1906.08101)**
|
|
||||||
Yiming Cui, Wanxiang Che, Ting Liu, Bing Qin, Ziqing Yang, Shijin Wang, Guoping Hu
|
|
||||||
|
|
||||||
This repository is developed based on:https://github.com/google-research/bert
|
|
||||||
|
|
||||||
You may also interested in,
|
|
||||||
- Chinese BERT series: https://github.com/ymcui/Chinese-BERT-wwm
|
|
||||||
- Chinese MacBERT: https://github.com/ymcui/MacBERT
|
|
||||||
- Chinese ELECTRA: https://github.com/ymcui/Chinese-ELECTRA
|
|
||||||
- Chinese XLNet: https://github.com/ymcui/Chinese-XLNet
|
|
||||||
- Knowledge Distillation Toolkit - TextBrewer: https://github.com/airaria/TextBrewer
|
|
||||||
|
|
||||||
More resources by HFL: https://github.com/ymcui/HFL-Anthology
|
|
||||||
|
|
||||||
## Citation
|
|
||||||
If you find the technical report or resource is useful, please cite the following technical report in your paper.
|
|
||||||
- Primary: https://arxiv.org/abs/2004.13922
|
|
||||||
```
|
|
||||||
@inproceedings{cui-etal-2020-revisiting,
|
|
||||||
title = "Revisiting Pre-Trained Models for {C}hinese Natural Language Processing",
|
|
||||||
author = "Cui, Yiming and
|
|
||||||
Che, Wanxiang and
|
|
||||||
Liu, Ting and
|
|
||||||
Qin, Bing and
|
|
||||||
Wang, Shijin and
|
|
||||||
Hu, Guoping",
|
|
||||||
booktitle = "Proceedings of the 2020 Conference on Empirical Methods in Natural Language Processing: Findings",
|
|
||||||
month = nov,
|
|
||||||
year = "2020",
|
|
||||||
address = "Online",
|
|
||||||
publisher = "Association for Computational Linguistics",
|
|
||||||
url = "https://www.aclweb.org/anthology/2020.findings-emnlp.58",
|
|
||||||
pages = "657--668",
|
|
||||||
}
|
|
||||||
```
|
|
||||||
- Secondary: https://arxiv.org/abs/1906.08101
|
|
||||||
```
|
|
||||||
@article{chinese-bert-wwm,
|
|
||||||
title={Pre-Training with Whole Word Masking for Chinese BERT},
|
|
||||||
author={Cui, Yiming and Che, Wanxiang and Liu, Ting and Qin, Bing and Yang, Ziqing and Wang, Shijin and Hu, Guoping},
|
|
||||||
journal={arXiv preprint arXiv:1906.08101},
|
|
||||||
year={2019}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
File diff suppressed because one or more lines are too long
37
bert/deberta-v2-large-japanese-char-wwm-onnx/config.json
Normal file
37
bert/deberta-v2-large-japanese-char-wwm-onnx/config.json
Normal file
@@ -0,0 +1,37 @@
|
|||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"DebertaV2ForMaskedLM"
|
||||||
|
],
|
||||||
|
"attention_head_size": 64,
|
||||||
|
"attention_probs_dropout_prob": 0.1,
|
||||||
|
"conv_act": "gelu",
|
||||||
|
"conv_kernel_size": 3,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_dropout_prob": 0.1,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"layer_norm_eps": 1e-07,
|
||||||
|
"max_position_embeddings": 512,
|
||||||
|
"max_relative_positions": -1,
|
||||||
|
"model_type": "deberta-v2",
|
||||||
|
"norm_rel_ebd": "layer_norm",
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"pad_token_id": 0,
|
||||||
|
"pooler_dropout": 0,
|
||||||
|
"pooler_hidden_act": "gelu",
|
||||||
|
"pooler_hidden_size": 1024,
|
||||||
|
"pos_att_type": [
|
||||||
|
"p2c",
|
||||||
|
"c2p"
|
||||||
|
],
|
||||||
|
"position_biased_input": false,
|
||||||
|
"position_buckets": 256,
|
||||||
|
"relative_attention": true,
|
||||||
|
"share_att_key": true,
|
||||||
|
"torch_dtype": "float16",
|
||||||
|
"transformers_version": "4.25.1",
|
||||||
|
"type_vocab_size": 0,
|
||||||
|
"vocab_size": 22012
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
{
|
||||||
|
"cls_token": "[CLS]",
|
||||||
|
"mask_token": "[MASK]",
|
||||||
|
"pad_token": "[PAD]",
|
||||||
|
"sep_token": "[SEP]",
|
||||||
|
"unk_token": "[UNK]"
|
||||||
|
}
|
||||||
22116
bert/deberta-v2-large-japanese-char-wwm-onnx/tokenizer.json
Normal file
22116
bert/deberta-v2-large-japanese-char-wwm-onnx/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,19 @@
|
|||||||
|
{
|
||||||
|
"cls_token": "[CLS]",
|
||||||
|
"do_lower_case": false,
|
||||||
|
"do_subword_tokenize": true,
|
||||||
|
"do_word_tokenize": true,
|
||||||
|
"jumanpp_kwargs": null,
|
||||||
|
"mask_token": "[MASK]",
|
||||||
|
"mecab_kwargs": null,
|
||||||
|
"model_max_length": 1000000000000000019884624838656,
|
||||||
|
"never_split": null,
|
||||||
|
"pad_token": "[PAD]",
|
||||||
|
"sep_token": "[SEP]",
|
||||||
|
"special_tokens_map_file": null,
|
||||||
|
"subword_tokenizer_type": "character",
|
||||||
|
"sudachi_kwargs": null,
|
||||||
|
"tokenizer_class": "BertJapaneseTokenizer",
|
||||||
|
"unk_token": "[UNK]",
|
||||||
|
"word_tokenizer_type": "basic"
|
||||||
|
}
|
||||||
22012
bert/deberta-v2-large-japanese-char-wwm-onnx/vocab.txt
Normal file
22012
bert/deberta-v2-large-japanese-char-wwm-onnx/vocab.txt
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,34 +0,0 @@
|
|||||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.model filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.npy filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.npz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pickle filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pkl filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
|
||||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.wasm filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zst filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,89 +0,0 @@
|
|||||||
---
|
|
||||||
language: ja
|
|
||||||
license: cc-by-sa-4.0
|
|
||||||
library_name: transformers
|
|
||||||
tags:
|
|
||||||
- deberta
|
|
||||||
- deberta-v2
|
|
||||||
- fill-mask
|
|
||||||
- character
|
|
||||||
- wwm
|
|
||||||
datasets:
|
|
||||||
- wikipedia
|
|
||||||
- cc100
|
|
||||||
- oscar
|
|
||||||
metrics:
|
|
||||||
- accuracy
|
|
||||||
mask_token: "[MASK]"
|
|
||||||
widget:
|
|
||||||
- text: "京都大学で自然言語処理を[MASK][MASK]する。"
|
|
||||||
---
|
|
||||||
|
|
||||||
# Model Card for Japanese character-level DeBERTa V2 large
|
|
||||||
|
|
||||||
## Model description
|
|
||||||
|
|
||||||
This is a Japanese DeBERTa V2 large model pre-trained on Japanese Wikipedia, the Japanese portion of CC-100, and the Japanese portion of OSCAR.
|
|
||||||
This model is trained with character-level tokenization and whole word masking.
|
|
||||||
|
|
||||||
## How to use
|
|
||||||
|
|
||||||
You can use this model for masked language modeling as follows:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from transformers import AutoTokenizer, AutoModelForMaskedLM
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained('ku-nlp/deberta-v2-large-japanese-char-wwm')
|
|
||||||
model = AutoModelForMaskedLM.from_pretrained('ku-nlp/deberta-v2-large-japanese-char-wwm')
|
|
||||||
|
|
||||||
sentence = '京都大学で自然言語処理を[MASK][MASK]する。'
|
|
||||||
encoding = tokenizer(sentence, return_tensors='pt')
|
|
||||||
...
|
|
||||||
```
|
|
||||||
|
|
||||||
You can also fine-tune this model on downstream tasks.
|
|
||||||
|
|
||||||
## Tokenization
|
|
||||||
|
|
||||||
There is no need to tokenize texts in advance, and you can give raw texts to the tokenizer.
|
|
||||||
The texts are tokenized into character-level tokens by [sentencepiece](https://github.com/google/sentencepiece).
|
|
||||||
|
|
||||||
## Training data
|
|
||||||
|
|
||||||
We used the following corpora for pre-training:
|
|
||||||
|
|
||||||
- Japanese Wikipedia (as of 20221020, 3.2GB, 27M sentences, 1.3M documents)
|
|
||||||
- Japanese portion of CC-100 (85GB, 619M sentences, 66M documents)
|
|
||||||
- Japanese portion of OSCAR (54GB, 326M sentences, 25M documents)
|
|
||||||
|
|
||||||
Note that we filtered out documents annotated with "header", "footer", or "noisy" tags in OSCAR.
|
|
||||||
Also note that Japanese Wikipedia was duplicated 10 times to make the total size of the corpus comparable to that of CC-100 and OSCAR. As a result, the total size of the training data is 171GB.
|
|
||||||
|
|
||||||
## Training procedure
|
|
||||||
|
|
||||||
We first segmented texts in the corpora into words using [Juman++ 2.0.0-rc3](https://github.com/ku-nlp/jumanpp/releases/tag/v2.0.0-rc3) for whole word masking.
|
|
||||||
Then, we built a sentencepiece model with 22,012 tokens including all characters that appear in the training corpus.
|
|
||||||
|
|
||||||
We tokenized raw corpora into character-level subwords using the sentencepiece model and trained the Japanese DeBERTa model using [transformers](https://github.com/huggingface/transformers) library.
|
|
||||||
The training took 26 days using 16 NVIDIA A100-SXM4-40GB GPUs.
|
|
||||||
|
|
||||||
The following hyperparameters were used during pre-training:
|
|
||||||
|
|
||||||
- learning_rate: 1e-4
|
|
||||||
- per_device_train_batch_size: 26
|
|
||||||
- distributed_type: multi-GPU
|
|
||||||
- num_devices: 16
|
|
||||||
- gradient_accumulation_steps: 8
|
|
||||||
- total_train_batch_size: 3,328
|
|
||||||
- max_seq_length: 512
|
|
||||||
- optimizer: Adam with betas=(0.9,0.999) and epsilon=1e-06
|
|
||||||
- lr_scheduler_type: linear schedule with warmup (lr = 0 at 300k steps)
|
|
||||||
- training_steps: 260,000
|
|
||||||
- warmup_steps: 10,000
|
|
||||||
|
|
||||||
The accuracy of the trained model on the masked language modeling task was 0.795.
|
|
||||||
The evaluation set consists of 5,000 randomly sampled documents from each of the training corpora.
|
|
||||||
|
|
||||||
## Acknowledgments
|
|
||||||
|
|
||||||
This work was supported by Joint Usage/Research Center for Interdisciplinary Large-scale Information Infrastructures (JHPCN) through General Collaboration Project no. jh221004, "Developing a Platform for Constructing and Sharing of Large-Scale Japanese Language Models".
|
|
||||||
For training models, we used the mdx: a platform for the data-driven future.
|
|
||||||
22116
bert/deberta-v2-large-japanese-char-wwm/tokenizer.json
Normal file
22116
bert/deberta-v2-large-japanese-char-wwm/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
@@ -11168,7 +11168,7 @@ $
|
|||||||
🪴
|
🪴
|
||||||
🫑
|
🫑
|
||||||
𤢖
|
𤢖
|
||||||
|
|
||||||
ǀ
|
ǀ
|
||||||
ǚ
|
ǚ
|
||||||
ɂ
|
ɂ
|
||||||
|
|||||||
22
bert/deberta-v3-large-onnx/config.json
Normal file
22
bert/deberta-v3-large-onnx/config.json
Normal file
@@ -0,0 +1,22 @@
|
|||||||
|
{
|
||||||
|
"model_type": "deberta-v2",
|
||||||
|
"attention_probs_dropout_prob": 0.1,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_dropout_prob": 0.1,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"max_position_embeddings": 512,
|
||||||
|
"relative_attention": true,
|
||||||
|
"position_buckets": 256,
|
||||||
|
"norm_rel_ebd": "layer_norm",
|
||||||
|
"share_att_key": true,
|
||||||
|
"pos_att_type": "p2c|c2p",
|
||||||
|
"layer_norm_eps": 1e-7,
|
||||||
|
"max_relative_positions": -1,
|
||||||
|
"position_biased_input": false,
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"type_vocab_size": 0,
|
||||||
|
"vocab_size": 128100
|
||||||
|
}
|
||||||
22
bert/deberta-v3-large-onnx/generator_config.json
Normal file
22
bert/deberta-v3-large-onnx/generator_config.json
Normal file
@@ -0,0 +1,22 @@
|
|||||||
|
{
|
||||||
|
"model_type": "deberta-v2",
|
||||||
|
"attention_probs_dropout_prob": 0.1,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_dropout_prob": 0.1,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"max_position_embeddings": 512,
|
||||||
|
"relative_attention": true,
|
||||||
|
"position_buckets": 256,
|
||||||
|
"norm_rel_ebd": "layer_norm",
|
||||||
|
"share_att_key": true,
|
||||||
|
"pos_att_type": "p2c|c2p",
|
||||||
|
"layer_norm_eps": 1e-7,
|
||||||
|
"max_relative_positions": -1,
|
||||||
|
"position_biased_input": false,
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 12,
|
||||||
|
"type_vocab_size": 0,
|
||||||
|
"vocab_size": 128100
|
||||||
|
}
|
||||||
128105
bert/deberta-v3-large-onnx/tokenizer.json
Normal file
128105
bert/deberta-v3-large-onnx/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
4
bert/deberta-v3-large-onnx/tokenizer_config.json
Normal file
4
bert/deberta-v3-large-onnx/tokenizer_config.json
Normal file
@@ -0,0 +1,4 @@
|
|||||||
|
{
|
||||||
|
"do_lower_case": false,
|
||||||
|
"vocab_type": "spm"
|
||||||
|
}
|
||||||
27
bert/deberta-v3-large/.gitattributes
vendored
27
bert/deberta-v3-large/.gitattributes
vendored
@@ -1,27 +0,0 @@
|
|||||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.model filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
|
||||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zstandard filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,93 +0,0 @@
|
|||||||
---
|
|
||||||
language: en
|
|
||||||
tags:
|
|
||||||
- deberta
|
|
||||||
- deberta-v3
|
|
||||||
- fill-mask
|
|
||||||
thumbnail: https://huggingface.co/front/thumbnails/microsoft.png
|
|
||||||
license: mit
|
|
||||||
---
|
|
||||||
|
|
||||||
## DeBERTaV3: Improving DeBERTa using ELECTRA-Style Pre-Training with Gradient-Disentangled Embedding Sharing
|
|
||||||
|
|
||||||
[DeBERTa](https://arxiv.org/abs/2006.03654) improves the BERT and RoBERTa models using disentangled attention and enhanced mask decoder. With those two improvements, DeBERTa out perform RoBERTa on a majority of NLU tasks with 80GB training data.
|
|
||||||
|
|
||||||
In [DeBERTa V3](https://arxiv.org/abs/2111.09543), we further improved the efficiency of DeBERTa using ELECTRA-Style pre-training with Gradient Disentangled Embedding Sharing. Compared to DeBERTa, our V3 version significantly improves the model performance on downstream tasks. You can find more technique details about the new model from our [paper](https://arxiv.org/abs/2111.09543).
|
|
||||||
|
|
||||||
Please check the [official repository](https://github.com/microsoft/DeBERTa) for more implementation details and updates.
|
|
||||||
|
|
||||||
The DeBERTa V3 large model comes with 24 layers and a hidden size of 1024. It has 304M backbone parameters with a vocabulary containing 128K tokens which introduces 131M parameters in the Embedding layer. This model was trained using the 160GB data as DeBERTa V2.
|
|
||||||
|
|
||||||
|
|
||||||
#### Fine-tuning on NLU tasks
|
|
||||||
|
|
||||||
We present the dev results on SQuAD 2.0 and MNLI tasks.
|
|
||||||
|
|
||||||
| Model |Vocabulary(K)|Backbone #Params(M)| SQuAD 2.0(F1/EM) | MNLI-m/mm(ACC)|
|
|
||||||
|-------------------|----------|-------------------|-----------|----------|
|
|
||||||
| RoBERTa-large |50 |304 | 89.4/86.5 | 90.2 |
|
|
||||||
| XLNet-large |32 |- | 90.6/87.9 | 90.8 |
|
|
||||||
| DeBERTa-large |50 |- | 90.7/88.0 | 91.3 |
|
|
||||||
| **DeBERTa-v3-large**|128|304 | **91.5/89.0**| **91.8/91.9**|
|
|
||||||
|
|
||||||
|
|
||||||
#### Fine-tuning with HF transformers
|
|
||||||
|
|
||||||
```bash
|
|
||||||
#!/bin/bash
|
|
||||||
|
|
||||||
cd transformers/examples/pytorch/text-classification/
|
|
||||||
|
|
||||||
pip install datasets
|
|
||||||
export TASK_NAME=mnli
|
|
||||||
|
|
||||||
output_dir="ds_results"
|
|
||||||
|
|
||||||
num_gpus=8
|
|
||||||
|
|
||||||
batch_size=8
|
|
||||||
|
|
||||||
python -m torch.distributed.launch --nproc_per_node=${num_gpus} \
|
|
||||||
run_glue.py \
|
|
||||||
--model_name_or_path microsoft/deberta-v3-large \
|
|
||||||
--task_name $TASK_NAME \
|
|
||||||
--do_train \
|
|
||||||
--do_eval \
|
|
||||||
--evaluation_strategy steps \
|
|
||||||
--max_seq_length 256 \
|
|
||||||
--warmup_steps 50 \
|
|
||||||
--per_device_train_batch_size ${batch_size} \
|
|
||||||
--learning_rate 6e-6 \
|
|
||||||
--num_train_epochs 2 \
|
|
||||||
--output_dir $output_dir \
|
|
||||||
--overwrite_output_dir \
|
|
||||||
--logging_steps 1000 \
|
|
||||||
--logging_dir $output_dir
|
|
||||||
|
|
||||||
```
|
|
||||||
|
|
||||||
### Citation
|
|
||||||
|
|
||||||
If you find DeBERTa useful for your work, please cite the following papers:
|
|
||||||
|
|
||||||
``` latex
|
|
||||||
@misc{he2021debertav3,
|
|
||||||
title={DeBERTaV3: Improving DeBERTa using ELECTRA-Style Pre-Training with Gradient-Disentangled Embedding Sharing},
|
|
||||||
author={Pengcheng He and Jianfeng Gao and Weizhu Chen},
|
|
||||||
year={2021},
|
|
||||||
eprint={2111.09543},
|
|
||||||
archivePrefix={arXiv},
|
|
||||||
primaryClass={cs.CL}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
``` latex
|
|
||||||
@inproceedings{
|
|
||||||
he2021deberta,
|
|
||||||
title={DEBERTA: DECODING-ENHANCED BERT WITH DISENTANGLED ATTENTION},
|
|
||||||
author={Pengcheng He and Xiaodong Liu and Jianfeng Gao and Weizhu Chen},
|
|
||||||
booktitle={International Conference on Learning Representations},
|
|
||||||
year={2021},
|
|
||||||
url={https://openreview.net/forum?id=XPZIaotutsD}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
128105
bert/deberta-v3-large/tokenizer.json
Normal file
128105
bert/deberta-v3-large/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@
|
|||||||
"id": "F7aJhsgLAWvO"
|
"id": "F7aJhsgLAWvO"
|
||||||
},
|
},
|
||||||
"source": [
|
"source": [
|
||||||
"# Style-Bert-VITS2 (ver 2.6.0) のGoogle Colabでの学習\n",
|
"# Style-Bert-VITS2 (ver 2.6.1) のGoogle Colabでの学習\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
|||||||
@@ -69,5 +69,5 @@
|
|||||||
"use_spectral_norm": false,
|
"use_spectral_norm": false,
|
||||||
"gin_channels": 256
|
"gin_channels": 256
|
||||||
},
|
},
|
||||||
"version": "2.6.0"
|
"version": "2.6.1"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -76,5 +76,5 @@
|
|||||||
"initial_channel": 64
|
"initial_channel": 64
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"version": "2.6.0-JP-Extra"
|
"version": "2.6.1-JP-Extra"
|
||||||
}
|
}
|
||||||
|
|||||||
119
convert_bert_onnx.py
Normal file
119
convert_bert_onnx.py
Normal file
@@ -0,0 +1,119 @@
|
|||||||
|
# Usage: .venv/bin/python convert_bert_onnx.py --language JP
|
||||||
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_deberta.py
|
||||||
|
|
||||||
|
import time
|
||||||
|
from argparse import ArgumentParser
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import onnx
|
||||||
|
import torch
|
||||||
|
from onnxsim import model_info, simplify
|
||||||
|
from rich import print
|
||||||
|
from rich.rule import Rule
|
||||||
|
from rich.style import Style
|
||||||
|
from torch import nn
|
||||||
|
from transformers.convert_slow_tokenizer import BertConverter
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
start_time = time.time()
|
||||||
|
parser = ArgumentParser()
|
||||||
|
parser.add_argument(
|
||||||
|
"--language",
|
||||||
|
default=Languages.JP,
|
||||||
|
help="Language of the BERT model to be converted",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# モデルの入出力先ファイルパスを取得
|
||||||
|
language = Languages(args.language)
|
||||||
|
pretrained_model_name_or_path = DEFAULT_BERT_MODEL_PATHS[language]
|
||||||
|
onnx_temp_model_path = Path(pretrained_model_name_or_path) / f"model_temp.onnx"
|
||||||
|
onnx_optimized_model_path = Path(pretrained_model_name_or_path) / f"model.onnx"
|
||||||
|
tokenizer_json_path = Path(pretrained_model_name_or_path) / "tokenizer.json"
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold cyan]Language:[/bold cyan] {language.name}")
|
||||||
|
print(f"[bold cyan]Pretrained model:[/bold cyan] {pretrained_model_name_or_path}")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
|
||||||
|
# トークナイザーを Fast Tokenizer 用形式に変換して保存
|
||||||
|
tokenizer = bert_models.load_tokenizer(language)
|
||||||
|
converter = BertConverter(tokenizer)
|
||||||
|
converter.converted().save(str(tokenizer_json_path))
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold green]Tokenizer JSON saved to {tokenizer_json_path}[/bold green]")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
|
||||||
|
class ONNXBert(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super(ONNXBert, self).__init__()
|
||||||
|
self.model = bert_models.load_model(language)
|
||||||
|
|
||||||
|
def forward(self, input_ids, token_type_ids, attention_mask):
|
||||||
|
inputs = {
|
||||||
|
"input_ids": input_ids,
|
||||||
|
"token_type_ids": token_type_ids,
|
||||||
|
"attention_mask": attention_mask,
|
||||||
|
}
|
||||||
|
res = self.model(**inputs, output_hidden_states=True)
|
||||||
|
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
return res
|
||||||
|
|
||||||
|
# ONNX 変換用の BERT モデルをロード
|
||||||
|
model = ONNXBert()
|
||||||
|
inputs = tokenizer("今日はいい天気ですね", return_tensors="pt")
|
||||||
|
|
||||||
|
# モデルを ONNX に変換
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold cyan]Exporting ONNX model...[/bold cyan]")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
export_start_time = time.time()
|
||||||
|
torch.onnx.export(
|
||||||
|
model=model,
|
||||||
|
args=(
|
||||||
|
inputs["input_ids"],
|
||||||
|
inputs["token_type_ids"],
|
||||||
|
inputs["attention_mask"],
|
||||||
|
),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
verbose=False,
|
||||||
|
input_names=[
|
||||||
|
"input_ids",
|
||||||
|
"token_type_ids",
|
||||||
|
"attention_mask",
|
||||||
|
],
|
||||||
|
output_names=["output"],
|
||||||
|
dynamic_axes={
|
||||||
|
"input_ids": {0: "batch_size", 1: "sequence_length"},
|
||||||
|
"token_type_ids": {0: "batch_size", 1: "sequence_length"},
|
||||||
|
"attention_mask": {0: "batch_size", 1: "sequence_length"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ONNX モデルを最適化
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold cyan]Optimizing ONNX model...[/bold cyan]")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
optimize_start_time = time.time()
|
||||||
|
onnx_model = onnx.load(onnx_temp_model_path)
|
||||||
|
simplified_onnx_model, check = simplify(onnx_model)
|
||||||
|
onnx.save(simplified_onnx_model, onnx_optimized_model_path)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold]Total Time: {time.time() - start_time:.2f}s / "
|
||||||
|
f"Size: {onnx_temp_model_path.stat().st_size / 1000 / 1000:.2f}MB -> "
|
||||||
|
f"{onnx_optimized_model_path.stat().st_size / 1000 / 1000:.2f}MB[/bold]"
|
||||||
|
)
|
||||||
|
onnx_temp_model_path.unlink()
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print("[bold cyan]Optimized model info:[/bold cyan]")
|
||||||
|
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
331
convert_onnx.py
Normal file
331
convert_onnx.py
Normal file
@@ -0,0 +1,331 @@
|
|||||||
|
# Usage: .venv/bin/python convert_onnx.py --model model_assets/koharune-ami/koharune-ami.safetensors
|
||||||
|
# Usage: .venv/bin/python convert_onnx.py --model model_assets/ (All models in the directory will be converted)
|
||||||
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py
|
||||||
|
|
||||||
|
import time
|
||||||
|
from argparse import ArgumentParser
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import cast
|
||||||
|
|
||||||
|
import onnx
|
||||||
|
import torch
|
||||||
|
from onnxsim import model_info, simplify
|
||||||
|
from rich import print
|
||||||
|
from rich.rule import Rule
|
||||||
|
from rich.style import Style
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import (
|
||||||
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
DEFAULT_STYLE,
|
||||||
|
DEFAULT_STYLE_WEIGHT,
|
||||||
|
Languages,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.models.infer import get_text
|
||||||
|
from style_bert_vits2.models.models import SynthesizerTrn
|
||||||
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.tts_model import TTSModel
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
start_time = time.time()
|
||||||
|
parser = ArgumentParser()
|
||||||
|
parser.add_argument(
|
||||||
|
"--model", required=True, help="Path to the model file or directory"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--force-convert",
|
||||||
|
action="store_true",
|
||||||
|
help="Already converted models will be overwritten",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# --model に指定されたパスがディレクトリの時、配下にある全ての .safetensors ファイルを対象に変換する
|
||||||
|
model_paths: list[Path] = []
|
||||||
|
if Path(args.model).is_dir():
|
||||||
|
for path in Path(args.model).glob("**/*.safetensors"):
|
||||||
|
# . から始まるファイルは除外
|
||||||
|
if not path.name.startswith("."):
|
||||||
|
model_paths.append(path)
|
||||||
|
else:
|
||||||
|
model_paths.append(Path(args.model))
|
||||||
|
|
||||||
|
for model_path in model_paths:
|
||||||
|
|
||||||
|
# モデルの入出力先ファイルパスを取得
|
||||||
|
onnx_temp_model_path = model_path.parent / f"{model_path.stem}_temp.onnx"
|
||||||
|
onnx_optimized_model_path = model_path.parent / f"{model_path.stem}.onnx"
|
||||||
|
config_path = model_path.parent / "config.json"
|
||||||
|
style_vec_path = model_path.parent / "style_vectors.npy"
|
||||||
|
assert model_path.exists(), "Model file does not exist"
|
||||||
|
assert config_path.exists(), "Config file does not exist"
|
||||||
|
assert style_vec_path.exists(), "Style vector file does not exist"
|
||||||
|
assert model_path.suffix != ".onnx", "Model file is already ONNX"
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold cyan]Model file:[/bold cyan] {model_path}")
|
||||||
|
print(f"[bold cyan]Config file:[/bold cyan] {config_path}")
|
||||||
|
print(f"[bold cyan]Style vector file:[/bold cyan] {style_vec_path}")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
|
||||||
|
# すでに ONNX モデルが存在する場合、--force-convert オプションが指定されていない場合はスキップ
|
||||||
|
if onnx_optimized_model_path.exists() and not args.force_convert:
|
||||||
|
print(
|
||||||
|
f"[bold yellow]ONNX model already exists: {onnx_optimized_model_path}[/bold yellow]"
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
"[bold]If you want to overwrite it, use the --force-convert option.[/bold]"
|
||||||
|
)
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
continue
|
||||||
|
|
||||||
|
# PyTorch モデルを読み込む
|
||||||
|
device = "cpu"
|
||||||
|
tts_model = TTSModel(
|
||||||
|
model_path=model_path,
|
||||||
|
config_path=config_path,
|
||||||
|
style_vec_path=style_vec_path,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
tts_model.load()
|
||||||
|
style_id = tts_model.style2id[DEFAULT_STYLE]
|
||||||
|
assert tts_model.net_g is not None, "Model is not loaded"
|
||||||
|
|
||||||
|
# 音声合成に必要な BERT 特徴量・音素列・アクセント列・言語 ID を取得
|
||||||
|
# JP-Extra モデルアーキテクチャの場合、bert (中国語の BERT 特徴量) や en_bert (英語の BERT 特徴量) は
|
||||||
|
# torch.zeros() で適当に埋められており、推論には ja_bert (日本語の BERT 特徴量) のみが使用される
|
||||||
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
|
"今日はいい天気ですね。",
|
||||||
|
Languages.JP,
|
||||||
|
tts_model.hyper_parameters,
|
||||||
|
device,
|
||||||
|
assist_text=None,
|
||||||
|
assist_text_weight=DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
given_phone=None,
|
||||||
|
given_tone=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# スタイルベクトルを取得
|
||||||
|
style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT)
|
||||||
|
|
||||||
|
# モデルの入力を作成
|
||||||
|
x_tst = phones.to(device).unsqueeze(0)
|
||||||
|
tones = tones.to(device).unsqueeze(0)
|
||||||
|
lang_ids = lang_ids.to(device).unsqueeze(0)
|
||||||
|
bert = bert.to(device).unsqueeze(0)
|
||||||
|
ja_bert = ja_bert.to(device).unsqueeze(0)
|
||||||
|
en_bert = en_bert.to(device).unsqueeze(0)
|
||||||
|
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
|
||||||
|
style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0)
|
||||||
|
sid = 0
|
||||||
|
sid_tensor = torch.LongTensor([sid]).to(device)
|
||||||
|
length_scale = torch.tensor(1.0)
|
||||||
|
sdp_ratio = torch.tensor(0.0)
|
||||||
|
noise_scale = torch.tensor(0.667)
|
||||||
|
noise_scale_w = torch.tensor(0.8)
|
||||||
|
|
||||||
|
# JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
||||||
|
if isinstance(tts_model.net_g, SynthesizerTrnJPExtra):
|
||||||
|
|
||||||
|
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
||||||
|
def forward_jp_extra(
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
sid: torch.Tensor,
|
||||||
|
tone: torch.Tensor,
|
||||||
|
language: torch.Tensor,
|
||||||
|
bert: torch.Tensor,
|
||||||
|
style_vec: torch.Tensor,
|
||||||
|
length_scale: float = 1.0,
|
||||||
|
sdp_ratio: float = 0.0,
|
||||||
|
noise_scale: float = 0.667,
|
||||||
|
noise_scale_w: float = 0.8,
|
||||||
|
) -> tuple[
|
||||||
|
torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]
|
||||||
|
]:
|
||||||
|
return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer(
|
||||||
|
x,
|
||||||
|
x_lengths,
|
||||||
|
sid,
|
||||||
|
tone,
|
||||||
|
language,
|
||||||
|
bert,
|
||||||
|
style_vec,
|
||||||
|
length_scale=length_scale,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise_scale,
|
||||||
|
noise_scale_w=noise_scale_w,
|
||||||
|
)
|
||||||
|
|
||||||
|
tts_model.net_g.forward = forward_jp_extra # type: ignore
|
||||||
|
|
||||||
|
# モデルを ONNX に変換
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(
|
||||||
|
f"[bold cyan]Exporting ONNX model... (Architecture: JP-Extra)[/bold cyan]"
|
||||||
|
)
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
export_start_time = time.time()
|
||||||
|
torch.onnx.export(
|
||||||
|
model=tts_model.net_g,
|
||||||
|
args=(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
ja_bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
length_scale,
|
||||||
|
sdp_ratio,
|
||||||
|
noise_scale,
|
||||||
|
noise_scale_w,
|
||||||
|
),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
verbose=False,
|
||||||
|
input_names=[
|
||||||
|
"x_tst",
|
||||||
|
"x_tst_lengths",
|
||||||
|
"sid",
|
||||||
|
"tones",
|
||||||
|
"language",
|
||||||
|
"bert",
|
||||||
|
"style_vec",
|
||||||
|
"length_scale",
|
||||||
|
"sdp_ratio",
|
||||||
|
"noise_scale",
|
||||||
|
"noise_scale_w",
|
||||||
|
],
|
||||||
|
output_names=["output"],
|
||||||
|
dynamic_axes={
|
||||||
|
"x_tst": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"x_tst_lengths": {0: "batch_size"},
|
||||||
|
"sid": {0: "batch_size"},
|
||||||
|
"tones": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"language": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"bert": {0: "batch_size", 2: "x_tst_max_length"},
|
||||||
|
"style_vec": {0: "batch_size"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 非 JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
||||||
|
else:
|
||||||
|
|
||||||
|
# SynthesizerTrn の forward メソッドをオーバーライド
|
||||||
|
def forward_non_jp_extra(
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
sid: torch.Tensor,
|
||||||
|
tone: torch.Tensor,
|
||||||
|
language: torch.Tensor,
|
||||||
|
bert: torch.Tensor,
|
||||||
|
ja_bert: torch.Tensor,
|
||||||
|
en_bert: torch.Tensor,
|
||||||
|
style_vec: torch.Tensor,
|
||||||
|
length_scale: float = 1.0,
|
||||||
|
sdp_ratio: float = 0.0,
|
||||||
|
noise_scale: float = 0.667,
|
||||||
|
noise_scale_w: float = 0.8,
|
||||||
|
) -> tuple[
|
||||||
|
torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]
|
||||||
|
]:
|
||||||
|
return cast(SynthesizerTrn, tts_model.net_g).infer(
|
||||||
|
x,
|
||||||
|
x_lengths,
|
||||||
|
sid,
|
||||||
|
tone,
|
||||||
|
language,
|
||||||
|
bert,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec,
|
||||||
|
length_scale=length_scale,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise_scale,
|
||||||
|
noise_scale_w=noise_scale_w,
|
||||||
|
)
|
||||||
|
|
||||||
|
tts_model.net_g.forward = forward_non_jp_extra # type: ignore
|
||||||
|
|
||||||
|
# モデルを ONNX に変換
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(
|
||||||
|
f"[bold cyan]Exporting ONNX model... (Architecture: Non-JP-Extra)[/bold cyan]"
|
||||||
|
)
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
export_start_time = time.time()
|
||||||
|
torch.onnx.export(
|
||||||
|
model=tts_model.net_g,
|
||||||
|
args=(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
length_scale,
|
||||||
|
sdp_ratio,
|
||||||
|
noise_scale,
|
||||||
|
noise_scale_w,
|
||||||
|
),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
verbose=False,
|
||||||
|
input_names=[
|
||||||
|
"x_tst",
|
||||||
|
"x_tst_lengths",
|
||||||
|
"sid",
|
||||||
|
"tones",
|
||||||
|
"language",
|
||||||
|
"bert",
|
||||||
|
"ja_bert",
|
||||||
|
"en_bert",
|
||||||
|
"style_vec",
|
||||||
|
"length_scale",
|
||||||
|
"sdp_ratio",
|
||||||
|
"noise_scale",
|
||||||
|
"noise_scale_w",
|
||||||
|
],
|
||||||
|
output_names=["output"],
|
||||||
|
dynamic_axes={
|
||||||
|
"x_tst": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"x_tst_lengths": {0: "batch_size"},
|
||||||
|
"sid": {0: "batch_size"},
|
||||||
|
"tones": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"language": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"bert": {0: "batch_size", 2: "x_tst_max_length"},
|
||||||
|
"ja_bert": {0: "batch_size", 2: "x_tst_max_length"},
|
||||||
|
"en_bert": {0: "batch_size", 2: "x_tst_max_length"},
|
||||||
|
"style_vec": {0: "batch_size"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ONNX モデルを最適化
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print(f"[bold cyan]Optimizing ONNX model...[/bold cyan]")
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
optimize_start_time = time.time()
|
||||||
|
onnx_model = onnx.load(onnx_temp_model_path)
|
||||||
|
simplified_onnx_model, check = simplify(onnx_model)
|
||||||
|
onnx.save(simplified_onnx_model, onnx_optimized_model_path)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold]Total Time: {time.time() - start_time:.2f}s / "
|
||||||
|
f"Size: {onnx_temp_model_path.stat().st_size / 1000 / 1000:.2f}MB -> "
|
||||||
|
f"{onnx_optimized_model_path.stat().st_size / 1000 / 1000:.2f}MB[/bold]"
|
||||||
|
)
|
||||||
|
onnx_temp_model_path.unlink()
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
print("[bold cyan]Optimized model info:[/bold cyan]")
|
||||||
|
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
||||||
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
@@ -1,5 +1,10 @@
|
|||||||
# Changelog
|
# Changelog
|
||||||
|
|
||||||
|
## v2.6.1 (2024-09-09)
|
||||||
|
|
||||||
|
- Google colabで、torchのバージョン由来でエラーが発生する不具合の修正(たぶん)
|
||||||
|
- WebUIからのスタイル作成での、サブフォルダによるスタイル分けでエラーが発生していた点の修正
|
||||||
|
|
||||||
## v2.6.0 (2024-06-16)
|
## v2.6.0 (2024-06-16)
|
||||||
|
|
||||||
### 新機能
|
### 新機能
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import datetime
|
import datetime
|
||||||
import json
|
import json
|
||||||
from typing import Any, Optional, Union
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
|
|
||||||
@@ -18,11 +19,12 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.models.infer import InvalidToneError
|
from style_bert_vits2.nlp import InvalidToneError
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
from style_bert_vits2.tts_model import TTSModelHolder
|
from style_bert_vits2.tts_model import NullModelParam, TTSModelHolder
|
||||||
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
# pyopenjtalk_worker を起動
|
# pyopenjtalk_worker を起動
|
||||||
@@ -217,16 +219,16 @@ def change_null_model_row(
|
|||||||
null_voice_pitch_weights: float,
|
null_voice_pitch_weights: float,
|
||||||
null_speech_style_weights: float,
|
null_speech_style_weights: float,
|
||||||
null_tempo_weights: float,
|
null_tempo_weights: float,
|
||||||
null_models: dict[int, dict[str, Any]],
|
null_models: dict[int, NullModelParam],
|
||||||
):
|
):
|
||||||
null_models[null_model_index] = {
|
null_models[null_model_index] = NullModelParam(
|
||||||
"name": null_model_name,
|
name=null_model_name,
|
||||||
"path": null_model_path,
|
path=Path(null_model_path),
|
||||||
"weight": null_voice_weights,
|
weight=null_voice_weights,
|
||||||
"pitch": null_voice_pitch_weights,
|
pitch=null_voice_pitch_weights,
|
||||||
"style": null_speech_style_weights,
|
style=null_speech_style_weights,
|
||||||
"tempo": null_tempo_weights,
|
tempo=null_tempo_weights,
|
||||||
}
|
)
|
||||||
if len(null_models) > null_models_frame:
|
if len(null_models) > null_models_frame:
|
||||||
keys_to_keep = list(range(null_models_frame))
|
keys_to_keep = list(range(null_models_frame))
|
||||||
result = {k: null_models[k] for k in keys_to_keep}
|
result = {k: null_models[k] for k in keys_to_keep}
|
||||||
@@ -258,7 +260,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
speaker,
|
speaker,
|
||||||
pitch_scale,
|
pitch_scale,
|
||||||
intonation_scale,
|
intonation_scale,
|
||||||
null_models: dict[int, dict[str, Union[str, float]]],
|
null_models: dict[int, NullModelParam],
|
||||||
force_reload_model: bool,
|
force_reload_model: bool,
|
||||||
):
|
):
|
||||||
model_holder.get_model(model_name, model_path)
|
model_holder.get_model(model_name, model_path)
|
||||||
@@ -736,6 +738,8 @@ if __name__ == "__main__":
|
|||||||
path_config = get_path_config()
|
path_config = get_path_config()
|
||||||
assets_root = path_config.assets_root
|
assets_root = path_config.assets_root
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
model_holder = TTSModelHolder(assets_root, device)
|
model_holder = TTSModelHolder(
|
||||||
|
assets_root, device, torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
app = create_inference_app(model_holder)
|
app = create_inference_app(model_holder)
|
||||||
app.launch(inbrowser=True)
|
app.launch(inbrowser=True)
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from config import get_path_config
|
|||||||
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
|
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
||||||
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
voice_keys = ["dec"]
|
voice_keys = ["dec"]
|
||||||
@@ -105,7 +106,7 @@ def merge_style_usual(
|
|||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -162,7 +163,7 @@ def merge_style_add_diff(
|
|||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -223,7 +224,7 @@ def merge_style_weighted_sum(
|
|||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -274,7 +275,7 @@ def merge_style_add_null(
|
|||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -370,7 +371,7 @@ def merge_models_usual(
|
|||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
neutral_vector_a = style_vectors_a[0]
|
neutral_vector_a = style_vectors_a[0]
|
||||||
@@ -454,7 +455,7 @@ def merge_models_add_diff(
|
|||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
@@ -531,7 +532,7 @@ def merge_models_weighted_sum(
|
|||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
@@ -609,7 +610,7 @@ def merge_models_add_null(
|
|||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
if new_config["data"]["n_speakers"] == 1:
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
new_config["data"]["spk2id"] = { output_name : 0}
|
new_config["data"]["spk2id"] = {output_name: 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
@@ -1018,6 +1019,14 @@ def method_change(x: str):
|
|||||||
|
|
||||||
|
|
||||||
def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
||||||
|
# ONNX モデルが混じらないよう、渡された TTSModelHolder のインスタンス変数を使って TTSModelHolder を作り直す
|
||||||
|
model_holder = TTSModelHolder(
|
||||||
|
model_holder.root_dir,
|
||||||
|
model_holder.device,
|
||||||
|
model_holder.onnx_providers,
|
||||||
|
ignore_onnx=True,
|
||||||
|
)
|
||||||
|
|
||||||
model_names = model_holder.model_names
|
model_names = model_holder.model_names
|
||||||
if len(model_names) == 0:
|
if len(model_names) == 0:
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -1540,8 +1549,9 @@ def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
model_holder = TTSModelHolder(
|
model_holder = TTSModelHolder(
|
||||||
assets_root, device="cuda" if torch.cuda.is_available() else "cpu"
|
assets_root, device, torch_device_to_onnx_providers(device), ignore_onnx=True
|
||||||
)
|
)
|
||||||
app = create_merge_app(model_holder)
|
app = create_merge_app(model_holder)
|
||||||
app.launch(inbrowser=True)
|
app.launch(inbrowser=True)
|
||||||
|
|||||||
@@ -336,7 +336,12 @@ def save_style_vectors_by_dirs(model_name: str, audio_dir_str: str):
|
|||||||
if style_vector_path.exists():
|
if style_vector_path.exists():
|
||||||
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
|
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
|
||||||
shutil.copy(style_vector_path, f"{style_vector_path}.bak")
|
shutil.copy(style_vector_path, f"{style_vector_path}.bak")
|
||||||
save_styles_by_dirs(audio_dir, result_dir)
|
save_styles_by_dirs(
|
||||||
|
wav_dir=audio_dir,
|
||||||
|
output_dir=result_dir,
|
||||||
|
config_path=config_path,
|
||||||
|
config_output_path=config_path,
|
||||||
|
)
|
||||||
return f"成功!\n{result_dir}にスタイルベクトルを保存しました。"
|
return f"成功!\n{result_dir}にスタイルベクトルを保存しました。"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -27,18 +27,25 @@ dependencies = [
|
|||||||
"g2p_en",
|
"g2p_en",
|
||||||
"jieba",
|
"jieba",
|
||||||
"loguru",
|
"loguru",
|
||||||
|
"nltk<=3.8.1",
|
||||||
"num2words",
|
"num2words",
|
||||||
"numba",
|
"numba",
|
||||||
"numpy",
|
"numpy<2",
|
||||||
|
"onnxruntime",
|
||||||
"pydantic>=2.0",
|
"pydantic>=2.0",
|
||||||
"pyopenjtalk-dict",
|
"pyopenjtalk-dict",
|
||||||
"pypinyin",
|
"pypinyin",
|
||||||
"pyworld-prebuilt",
|
"pyworld-prebuilt",
|
||||||
"safetensors",
|
"safetensors",
|
||||||
"torch>=2.1",
|
|
||||||
"transformers",
|
"transformers",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
torch = [
|
||||||
|
"accelerate",
|
||||||
|
"torch>=2.1",
|
||||||
|
]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Documentation = "https://github.com/litagin02/Style-Bert-VITS2#readme"
|
Documentation = "https://github.com/litagin02/Style-Bert-VITS2#readme"
|
||||||
Issues = "https://github.com/litagin02/Style-Bert-VITS2/issues"
|
Issues = "https://github.com/litagin02/Style-Bert-VITS2/issues"
|
||||||
@@ -64,35 +71,71 @@ exclude = [".git", ".gitignore", ".gitattributes"]
|
|||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
packages = ["style_bert_vits2"]
|
packages = ["style_bert_vits2"]
|
||||||
|
|
||||||
|
# for PyTorch inference
|
||||||
[tool.hatch.envs.test]
|
[tool.hatch.envs.test]
|
||||||
dependencies = ["coverage[toml]>=6.5", "pytest"]
|
dependencies = [
|
||||||
|
"coverage[toml]>=6.5",
|
||||||
|
"pytest",
|
||||||
|
"scipy",
|
||||||
|
]
|
||||||
|
features = ["torch"]
|
||||||
[tool.hatch.envs.test.scripts]
|
[tool.hatch.envs.test.scripts]
|
||||||
# Usage: `hatch run test:test`
|
# Usage: `hatch run test:test`
|
||||||
test = "pytest {args:tests}"
|
test = "pytest -s tests/test_main.py::test_synthesize_cpu"
|
||||||
|
# Usage: `hatch run test:test-cuda`
|
||||||
|
test-cuda = "pytest -s tests/test_main.py::test_synthesize_cuda"
|
||||||
# Usage: `hatch run test:coverage`
|
# Usage: `hatch run test:coverage`
|
||||||
test-cov = "coverage run -m pytest {args:tests}"
|
test-cov = "coverage run -m pytest -s tests/test_main.py::test_synthesize_cpu"
|
||||||
# Usage: `hatch run test:cov-report`
|
# Usage: `hatch run test:cov-report`
|
||||||
cov-report = ["- coverage combine", "coverage report"]
|
cov-report = ["- coverage combine", "coverage report"]
|
||||||
# Usage: `hatch run test:cov`
|
# Usage: `hatch run test:cov`
|
||||||
cov = ["test-cov", "cov-report"]
|
cov = ["test-cov", "cov-report"]
|
||||||
|
[[tool.hatch.envs.test.matrix]]
|
||||||
|
python = ["3.9", "3.10", "3.11"]
|
||||||
|
|
||||||
|
# for ONNX inference (without PyTorch dependency)
|
||||||
|
[tool.hatch.envs.test-onnx]
|
||||||
|
dependencies = [
|
||||||
|
"coverage[toml]>=6.5",
|
||||||
|
"pytest",
|
||||||
|
"scipy",
|
||||||
|
"onnxruntime-directml; sys_platform == 'win32'",
|
||||||
|
"onnxruntime-gpu; sys_platform != 'darwin'",
|
||||||
|
]
|
||||||
|
[tool.hatch.envs.test-onnx.scripts]
|
||||||
|
# Usage: `hatch run test-onnx:test`
|
||||||
|
test = "pytest -s tests/test_main.py::test_synthesize_onnx_cpu"
|
||||||
|
# Usage: `hatch run test-onnx:test-cuda`
|
||||||
|
test-cuda = "pytest -s tests/test_main.py::test_synthesize_onnx_cuda"
|
||||||
|
# Usage: `hatch run test-onnx:test-directml`
|
||||||
|
test-directml = "pytest -s tests/test_main.py::test_synthesize_onnx_directml"
|
||||||
|
# Usage: `hatch run test-onnx:test-coreml`
|
||||||
|
test-coreml = "pytest -s tests/test_main.py::test_synthesize_onnx_coreml"
|
||||||
|
# Usage: `hatch run test-onnx:coverage`
|
||||||
|
test-cov = "coverage run -m pytest -s tests/test_main.py::test_synthesize_onnx_cpu"
|
||||||
|
# Usage: `hatch run test-onnx:cov-report`
|
||||||
|
cov-report = ["- coverage combine", "coverage report"]
|
||||||
|
# Usage: `hatch run test-onnx:cov`
|
||||||
|
cov = ["test-cov", "cov-report"]
|
||||||
|
[[tool.hatch.envs.test-onnx.matrix]]
|
||||||
|
python = ["3.9", "3.10", "3.11"]
|
||||||
|
|
||||||
[tool.hatch.envs.style]
|
[tool.hatch.envs.style]
|
||||||
detached = true
|
detached = true
|
||||||
dependencies = ["black[jupyter]", "isort"]
|
dependencies = ["black[jupyter]", "isort"]
|
||||||
[tool.hatch.envs.style.scripts]
|
[tool.hatch.envs.style.scripts]
|
||||||
|
# Usage: `hatch run style:check`
|
||||||
check = [
|
check = [
|
||||||
"black --check --diff .",
|
"black --check --diff .",
|
||||||
"isort --check-only --diff --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
"isort --check-only --diff --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
]
|
]
|
||||||
|
# Usage: `hatch run style:fmt`
|
||||||
fmt = [
|
fmt = [
|
||||||
"black .",
|
"black .",
|
||||||
"isort --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
"isort --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
"check",
|
"check",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[tool.hatch.envs.test.matrix]]
|
|
||||||
python = ["3.9", "3.10", "3.11"]
|
|
||||||
|
|
||||||
[tool.coverage.run]
|
[tool.coverage.run]
|
||||||
source_pkgs = ["style_bert_vits2", "tests"]
|
source_pkgs = ["style_bert_vits2", "tests"]
|
||||||
branch = true
|
branch = true
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
g2p_en
|
g2p_en
|
||||||
@@ -5,9 +6,11 @@ gradio>=4.32
|
|||||||
jieba
|
jieba
|
||||||
librosa==0.9.2
|
librosa==0.9.2
|
||||||
loguru
|
loguru
|
||||||
|
nltk<=3.8.1
|
||||||
num2words
|
num2words
|
||||||
numpy<2
|
numpy<2
|
||||||
onnxruntime
|
onnxruntime
|
||||||
|
onnxruntime-gpu
|
||||||
pyannote.audio>=3.1.0
|
pyannote.audio>=3.1.0
|
||||||
pyloudnorm
|
pyloudnorm
|
||||||
pyopenjtalk-dict
|
pyopenjtalk-dict
|
||||||
@@ -15,5 +18,6 @@ pypinyin
|
|||||||
pyworld-prebuilt
|
pyworld-prebuilt
|
||||||
torch
|
torch
|
||||||
torchaudio
|
torchaudio
|
||||||
|
torchvision
|
||||||
transformers
|
transformers
|
||||||
umap-learn
|
umap-learn
|
||||||
|
|||||||
@@ -1,14 +1,20 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
# faster-whisper==0.10.1
|
# faster-whisper==0.10.1
|
||||||
g2p_en
|
g2p_en
|
||||||
GPUtil
|
GPUtil
|
||||||
gradio
|
gradio>=4.32
|
||||||
jieba
|
jieba
|
||||||
# librosa==0.9.2
|
# librosa==0.9.2
|
||||||
loguru
|
loguru
|
||||||
|
nltk<=3.8.1
|
||||||
num2words
|
num2words
|
||||||
numpy<2
|
numpy<2
|
||||||
|
onnxruntime
|
||||||
|
onnxruntime-directml; sys_platform == 'win32'
|
||||||
|
onnxruntime-gpu; sys_platform != 'darwin'
|
||||||
|
# onnxsim
|
||||||
# protobuf==4.25
|
# protobuf==4.25
|
||||||
psutil
|
psutil
|
||||||
# punctuators
|
# punctuators
|
||||||
@@ -20,5 +26,6 @@ pyworld-prebuilt
|
|||||||
# stable_ts
|
# stable_ts
|
||||||
# tensorboard
|
# tensorboard
|
||||||
torch
|
torch
|
||||||
|
torchaudio
|
||||||
transformers
|
transformers
|
||||||
umap-learn
|
umap-learn
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
faster-whisper==0.10.1
|
faster-whisper==0.10.1
|
||||||
@@ -7,8 +8,13 @@ gradio>=4.32
|
|||||||
jieba
|
jieba
|
||||||
librosa==0.9.2
|
librosa==0.9.2
|
||||||
loguru
|
loguru
|
||||||
|
nltk<=3.8.1
|
||||||
num2words
|
num2words
|
||||||
numpy<2
|
numpy<2
|
||||||
|
onnxruntime
|
||||||
|
onnxruntime-directml; sys_platform == 'win32'
|
||||||
|
onnxruntime-gpu; sys_platform != 'darwin'
|
||||||
|
onnxsim
|
||||||
protobuf==4.25
|
protobuf==4.25
|
||||||
psutil
|
psutil
|
||||||
punctuators
|
punctuators
|
||||||
|
|||||||
@@ -110,8 +110,8 @@ if !errorlevel! neq 0 ( pause & popd & exit /b !errorlevel! )
|
|||||||
echo --------------------------------------------------
|
echo --------------------------------------------------
|
||||||
echo Installing PyTorch...
|
echo Installing PyTorch...
|
||||||
echo --------------------------------------------------
|
echo --------------------------------------------------
|
||||||
echo Executing: uv pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu118
|
echo Executing: uv pip install "torch<2.4" "torchaudio<2.4" --index-url https://download.pytorch.org/whl/cu118
|
||||||
uv pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu118
|
uv pip install "torch<2.4" "torchaudio<2.4" --index-url https://download.pytorch.org/whl/cu118
|
||||||
if !errorlevel! neq 0 ( pause & popd & exit /b !errorlevel! )
|
if !errorlevel! neq 0 ( pause & popd & exit /b !errorlevel! )
|
||||||
|
|
||||||
echo --------------------------------------------------
|
echo --------------------------------------------------
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
@@ -53,6 +53,7 @@ from style_bert_vits2.nlp.japanese.user_dict import (
|
|||||||
update_dict,
|
update_dict,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo
|
from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo
|
||||||
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
# ---フロントエンド部分に関する処理---
|
# ---フロントエンド部分に関する処理---
|
||||||
@@ -156,12 +157,6 @@ pyopenjtalk.initialize_worker()
|
|||||||
# pyopenjtalk の辞書を更新
|
# pyopenjtalk の辞書を更新
|
||||||
update_dict()
|
update_dict()
|
||||||
|
|
||||||
# 事前に BERT モデル/トークナイザーをロードしておく
|
|
||||||
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
|
||||||
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
|
||||||
bert_models.load_model(Languages.JP)
|
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
|
||||||
|
|
||||||
|
|
||||||
class AudioResponse(Response):
|
class AudioResponse(Response):
|
||||||
media_type = "audio/wav"
|
media_type = "audio/wav"
|
||||||
@@ -184,6 +179,7 @@ parser.add_argument("--line_length", type=int, default=None)
|
|||||||
parser.add_argument("--line_count", type=int, default=None)
|
parser.add_argument("--line_count", type=int, default=None)
|
||||||
# parser.add_argument("--skip_default_models", action="store_true")
|
# parser.add_argument("--skip_default_models", action="store_true")
|
||||||
parser.add_argument("--skip_static_files", action="store_true")
|
parser.add_argument("--skip_static_files", action="store_true")
|
||||||
|
parser.add_argument("--preload_onnx_bert", action="store_true")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
device = args.device
|
device = args.device
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
@@ -194,7 +190,19 @@ port = int(args.port)
|
|||||||
# download_default_models()
|
# download_default_models()
|
||||||
skip_static_files = bool(args.skip_static_files)
|
skip_static_files = bool(args.skip_static_files)
|
||||||
|
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
# 事前に BERT モデル/トークナイザーをロードしておく
|
||||||
|
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
||||||
|
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
||||||
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
|
# VRAM 節約のため、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||||
|
if args.preload_onnx_bert:
|
||||||
|
onnx_bert_models.load_model(
|
||||||
|
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
|
model_holder = TTSModelHolder(model_dir, device, torch_device_to_onnx_providers(device))
|
||||||
if len(model_holder.model_names) == 0:
|
if len(model_holder.model_names) == 0:
|
||||||
logger.error(f"Models not found in {model_dir}.")
|
logger.error(f"Models not found in {model_dir}.")
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|||||||
@@ -34,10 +34,11 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
||||||
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
config = get_config()
|
config = get_config()
|
||||||
@@ -51,15 +52,6 @@ pyopenjtalk.initialize_worker()
|
|||||||
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
update_dict()
|
update_dict()
|
||||||
|
|
||||||
# 事前に BERT モデル/トークナイザーをロードしておく
|
|
||||||
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
|
||||||
bert_models.load_model(Languages.JP)
|
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
|
||||||
bert_models.load_model(Languages.EN)
|
|
||||||
bert_models.load_tokenizer(Languages.EN)
|
|
||||||
bert_models.load_model(Languages.ZH)
|
|
||||||
bert_models.load_tokenizer(Languages.ZH)
|
|
||||||
|
|
||||||
|
|
||||||
def raise_validation_error(msg: str, param: str):
|
def raise_validation_error(msg: str, param: str):
|
||||||
logger.warning(f"Validation error: {msg}")
|
logger.warning(f"Validation error: {msg}")
|
||||||
@@ -97,6 +89,7 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--dir", "-d", type=str, help="Model directory", default=config.assets_root
|
"--dir", "-d", type=str, help="Model directory", default=config.assets_root
|
||||||
)
|
)
|
||||||
|
parser.add_argument("--preload_onnx_bert", action="store_true")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.cpu:
|
if args.cpu:
|
||||||
@@ -104,8 +97,22 @@ if __name__ == "__main__":
|
|||||||
else:
|
else:
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
|
# 事前に BERT モデル/トークナイザーをロードしておく
|
||||||
|
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
||||||
|
## 英語や中国語で音声合成するユースケースは限られていることから、VRAM 節約のため日本語の BERT モデル/トークナイザーのみロードする
|
||||||
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
|
# VRAM 節約のため、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||||
|
if args.preload_onnx_bert:
|
||||||
|
onnx_bert_models.load_model(
|
||||||
|
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_dir = Path(args.dir)
|
model_dir = Path(args.dir)
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(
|
||||||
|
model_dir, device, torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
if len(model_holder.model_names) == 0:
|
if len(model_holder.model_names) == 0:
|
||||||
logger.error(f"Models not found in {model_dir}.")
|
logger.error(f"Models not found in {model_dir}.")
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
@@ -141,6 +148,10 @@ if __name__ == "__main__":
|
|||||||
request: Request,
|
request: Request,
|
||||||
text: str = Query(..., min_length=1, max_length=limit, description="セリフ"),
|
text: str = Query(..., min_length=1, max_length=limit, description="セリフ"),
|
||||||
encoding: str = Query(None, description="textをURLデコードする(ex, `utf-8`)"),
|
encoding: str = Query(None, description="textをURLデコードする(ex, `utf-8`)"),
|
||||||
|
model_name: str = Query(
|
||||||
|
None,
|
||||||
|
description="モデル名(model_idより優先)。model_assets内のディレクトリ名を指定",
|
||||||
|
),
|
||||||
model_id: int = Query(
|
model_id: int = Query(
|
||||||
0, description="モデルID。`GET /models/info`のkeyの値を指定ください"
|
0, description="モデルID。`GET /models/info`のkeyの値を指定ください"
|
||||||
),
|
),
|
||||||
@@ -198,6 +209,24 @@ if __name__ == "__main__":
|
|||||||
): # /models/refresh があるためQuery(le)で表現不可
|
): # /models/refresh があるためQuery(le)で表現不可
|
||||||
raise_validation_error(f"model_id={model_id} not found", "model_id")
|
raise_validation_error(f"model_id={model_id} not found", "model_id")
|
||||||
|
|
||||||
|
if model_name:
|
||||||
|
# load_models() の 処理内容が i の正当性を担保していることに注意
|
||||||
|
model_ids = [
|
||||||
|
i
|
||||||
|
for i, x in enumerate(model_holder.models_info)
|
||||||
|
if x.name == model_name
|
||||||
|
]
|
||||||
|
if not model_ids:
|
||||||
|
raise_validation_error(
|
||||||
|
f"model_name={model_name} not found", "model_name"
|
||||||
|
)
|
||||||
|
# 今の実装ではディレクトリ名が重複することは無いはずだが...
|
||||||
|
if len(model_ids) > 1:
|
||||||
|
raise_validation_error(
|
||||||
|
f"model_name={model_name} is ambiguous", "model_name"
|
||||||
|
)
|
||||||
|
model_id = model_ids[0]
|
||||||
|
|
||||||
model = loaded_models[model_id]
|
model = loaded_models[model_id]
|
||||||
if speaker_name is None:
|
if speaker_name is None:
|
||||||
if speaker_id not in model.id2spk.keys():
|
if speaker_id not in model.id2spk.keys():
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from style_bert_vits2.utils.strenum import StrEnum
|
|||||||
|
|
||||||
|
|
||||||
# Style-Bert-VITS2 のバージョン
|
# Style-Bert-VITS2 のバージョン
|
||||||
VERSION = "2.6.0"
|
VERSION = "2.6.1"
|
||||||
|
|
||||||
# Style-Bert-VITS2 のベースディレクトリ
|
# Style-Bert-VITS2 のベースディレクトリ
|
||||||
BASE_DIR = Path(__file__).parent.parent
|
BASE_DIR = Path(__file__).parent.parent
|
||||||
@@ -18,13 +18,20 @@ class Languages(StrEnum):
|
|||||||
ZH = "ZH"
|
ZH = "ZH"
|
||||||
|
|
||||||
|
|
||||||
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
# 言語ごとのデフォルトの BERT モデルのパス
|
||||||
DEFAULT_BERT_TOKENIZER_PATHS = {
|
DEFAULT_BERT_MODEL_PATHS = {
|
||||||
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
||||||
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
||||||
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 言語ごとのデフォルトの BERT モデル (ONNX 版) のパス
|
||||||
|
DEFAULT_ONNX_BERT_MODEL_PATHS = {
|
||||||
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm-onnx",
|
||||||
|
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large-onnx",
|
||||||
|
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large-onnx",
|
||||||
|
}
|
||||||
|
|
||||||
# デフォルトのユーザー辞書ディレクトリ
|
# デフォルトのユーザー辞書ディレクトリ
|
||||||
## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される
|
## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される
|
||||||
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
||||||
|
|||||||
@@ -38,7 +38,8 @@ class HyperParametersTrain(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
class HyperParametersData(BaseModel):
|
class HyperParametersData(BaseModel):
|
||||||
use_jp_extra: bool = True
|
# use_jp_extra フィールドが存在しない旧モデルとの互換性のために False をデフォルト値とする
|
||||||
|
use_jp_extra: bool = False
|
||||||
training_files: str = "Data/Dummy/train.list"
|
training_files: str = "Data/Dummy/train.list"
|
||||||
validation_files: str = "Data/Dummy/val.list"
|
validation_files: str = "Data/Dummy/val.list"
|
||||||
max_wav_value: float = 32768.0
|
max_wav_value: float = 32768.0
|
||||||
|
|||||||
@@ -12,14 +12,16 @@ from style_bert_vits2.models.models_jp_extra import (
|
|||||||
SynthesizerTrn as SynthesizerTrnJPExtra,
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.nlp import (
|
from style_bert_vits2.nlp import (
|
||||||
clean_text,
|
clean_text_with_given_phone_tone,
|
||||||
cleaned_text_to_sequence,
|
cleaned_text_to_sequence,
|
||||||
extract_bert_feature,
|
extract_bert_feature,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.nlp.symbols import SYMBOLS
|
from style_bert_vits2.nlp.symbols import SYMBOLS
|
||||||
|
|
||||||
|
|
||||||
def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
|
def get_net_g(
|
||||||
|
model_path: str, version: str, device: str, hps: HyperParameters
|
||||||
|
) -> Union[SynthesizerTrn, SynthesizerTrnJPExtra]:
|
||||||
if version.endswith("JP-Extra"):
|
if version.endswith("JP-Extra"):
|
||||||
logger.info("Using JP-Extra model")
|
logger.info("Using JP-Extra model")
|
||||||
net_g = SynthesizerTrnJPExtra(
|
net_g = SynthesizerTrnJPExtra(
|
||||||
@@ -86,10 +88,10 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
|
|||||||
_ = net_g.eval()
|
_ = net_g.eval()
|
||||||
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
||||||
_ = utils.checkpoints.load_checkpoint(
|
_ = utils.checkpoints.load_checkpoint(
|
||||||
model_path, net_g, None, skip_optimizer=True
|
model_path, net_g, None, skip_optimizer=True, device=device
|
||||||
)
|
)
|
||||||
elif model_path.endswith(".safetensors"):
|
elif model_path.endswith(".safetensors"):
|
||||||
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
|
_ = utils.safetensors.load_safetensors(model_path, net_g, True, device=device)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown model format: {model_path}")
|
raise ValueError(f"Unknown model format: {model_path}")
|
||||||
return net_g
|
return net_g
|
||||||
@@ -104,55 +106,19 @@ def get_text(
|
|||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
given_phone: Optional[list[str]] = None,
|
given_phone: Optional[list[str]] = None,
|
||||||
given_tone: Optional[list[int]] = None,
|
given_tone: Optional[list[int]] = None,
|
||||||
):
|
) -> tuple[
|
||||||
|
torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor
|
||||||
|
]:
|
||||||
use_jp_extra = hps.version.endswith("JP-Extra")
|
use_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
# 推論時のみ呼び出されるので、raise_yomi_error は False に設定
|
norm_text, phone, tone, word2ph = clean_text_with_given_phone_tone(
|
||||||
norm_text, phone, tone, word2ph = clean_text(
|
|
||||||
text,
|
text,
|
||||||
language_str,
|
language_str,
|
||||||
|
given_phone=given_phone,
|
||||||
|
given_tone=given_tone,
|
||||||
use_jp_extra=use_jp_extra,
|
use_jp_extra=use_jp_extra,
|
||||||
|
# 推論時のみ呼び出されるので、raise_yomi_error は False に設定
|
||||||
raise_yomi_error=False,
|
raise_yomi_error=False,
|
||||||
)
|
)
|
||||||
# phone と tone の両方が与えられた場合はそれを使う
|
|
||||||
if given_phone is not None and given_tone is not None:
|
|
||||||
# 指定された phone と指定された tone 両方の長さが一致していなければならない
|
|
||||||
if len(given_phone) != len(given_tone):
|
|
||||||
raise InvalidPhoneError(
|
|
||||||
f"Length of given_phone ({len(given_phone)}) != length of given_tone ({len(given_tone)})"
|
|
||||||
)
|
|
||||||
# 与えられた音素数と pyopenjtalk で生成した読みの音素数が一致しない
|
|
||||||
if len(given_phone) != sum(word2ph):
|
|
||||||
# 日本語の場合、len(given_phone) と sum(word2ph) が一致するように word2ph を適切に調整する
|
|
||||||
# 他の言語は word2ph の調整方法が思いつかないのでエラー
|
|
||||||
if language_str == Languages.JP:
|
|
||||||
from style_bert_vits2.nlp.japanese.g2p import adjust_word2ph
|
|
||||||
|
|
||||||
word2ph = adjust_word2ph(word2ph, phone, given_phone)
|
|
||||||
# 上記処理により word2ph の合計が given_phone の長さと一致するはず
|
|
||||||
# それでも一致しない場合、大半は読み上げテキストと given_phone が著しく乖離していて調整し切れなかったことを意味する
|
|
||||||
if len(given_phone) != sum(word2ph):
|
|
||||||
raise InvalidPhoneError(
|
|
||||||
f"Length of given_phone ({len(given_phone)}) != sum of word2ph ({sum(word2ph)})"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raise InvalidPhoneError(
|
|
||||||
f"Length of given_phone ({len(given_phone)}) != sum of word2ph ({sum(word2ph)})"
|
|
||||||
)
|
|
||||||
phone = given_phone
|
|
||||||
# 生成あるいは指定された phone と指定された tone 両方の長さが一致していなければならない
|
|
||||||
if len(phone) != len(given_tone):
|
|
||||||
raise InvalidToneError(
|
|
||||||
f"Length of phone ({len(phone)}) != length of given_tone ({len(given_tone)})"
|
|
||||||
)
|
|
||||||
tone = given_tone
|
|
||||||
# tone だけが与えられた場合は clean_text() で生成した phone と合わせて使う
|
|
||||||
elif given_tone is not None:
|
|
||||||
# 生成した phone と指定された tone 両方の長さが一致していなければならない
|
|
||||||
if len(phone) != len(given_tone):
|
|
||||||
raise InvalidToneError(
|
|
||||||
f"Length of phone ({len(phone)}) != length of given_tone ({len(given_tone)})"
|
|
||||||
)
|
|
||||||
tone = given_tone
|
|
||||||
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
||||||
|
|
||||||
if hps.data.add_blank:
|
if hps.data.add_blank:
|
||||||
@@ -216,7 +182,7 @@ def infer(
|
|||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
given_phone: Optional[list[str]] = None,
|
given_phone: Optional[list[str]] = None,
|
||||||
given_tone: Optional[list[int]] = None,
|
given_tone: Optional[list[int]] = None,
|
||||||
):
|
) -> NDArray[Any]:
|
||||||
is_jp_extra = hps.version.endswith("JP-Extra")
|
is_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
text,
|
text,
|
||||||
@@ -242,6 +208,7 @@ def infer(
|
|||||||
bert = bert[:, :-2]
|
bert = bert[:, :-2]
|
||||||
ja_bert = ja_bert[:, :-2]
|
ja_bert = ja_bert[:, :-2]
|
||||||
en_bert = en_bert[:, :-2]
|
en_bert = en_bert[:, :-2]
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
x_tst = phones.to(device).unsqueeze(0)
|
x_tst = phones.to(device).unsqueeze(0)
|
||||||
tones = tones.to(device).unsqueeze(0)
|
tones = tones.to(device).unsqueeze(0)
|
||||||
@@ -253,6 +220,7 @@ def infer(
|
|||||||
style_vec_tensor = torch.from_numpy(style_vec).to(device).unsqueeze(0)
|
style_vec_tensor = torch.from_numpy(style_vec).to(device).unsqueeze(0)
|
||||||
del phones
|
del phones
|
||||||
sid_tensor = torch.LongTensor([sid]).to(device)
|
sid_tensor = torch.LongTensor([sid]).to(device)
|
||||||
|
|
||||||
if is_jp_extra:
|
if is_jp_extra:
|
||||||
output = cast(SynthesizerTrnJPExtra, net_g).infer(
|
output = cast(SynthesizerTrnJPExtra, net_g).infer(
|
||||||
x_tst,
|
x_tst,
|
||||||
@@ -262,10 +230,10 @@ def infer(
|
|||||||
lang_ids,
|
lang_ids,
|
||||||
ja_bert,
|
ja_bert,
|
||||||
style_vec=style_vec_tensor,
|
style_vec=style_vec_tensor,
|
||||||
|
length_scale=length_scale,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise_scale,
|
noise_scale=noise_scale,
|
||||||
noise_scale_w=noise_scale_w,
|
noise_scale_w=noise_scale_w,
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
output = cast(SynthesizerTrn, net_g).infer(
|
output = cast(SynthesizerTrn, net_g).infer(
|
||||||
@@ -278,12 +246,14 @@ def infer(
|
|||||||
ja_bert,
|
ja_bert,
|
||||||
en_bert,
|
en_bert,
|
||||||
style_vec=style_vec_tensor,
|
style_vec=style_vec_tensor,
|
||||||
|
length_scale=length_scale,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise_scale,
|
noise_scale=noise_scale,
|
||||||
noise_scale_w=noise_scale_w,
|
noise_scale_w=noise_scale_w,
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
audio = output[0][0, 0].data.cpu().float().numpy()
|
audio = output[0][0, 0].data.cpu().float().numpy()
|
||||||
|
|
||||||
del (
|
del (
|
||||||
x_tst,
|
x_tst,
|
||||||
tones,
|
tones,
|
||||||
@@ -297,12 +267,5 @@ def infer(
|
|||||||
) # , emo
|
) # , emo
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
return audio
|
return audio
|
||||||
|
|
||||||
|
|
||||||
class InvalidPhoneError(ValueError):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class InvalidToneError(ValueError):
|
|
||||||
pass
|
|
||||||
|
|||||||
244
style_bert_vits2/models/infer_onnx.py
Normal file
244
style_bert_vits2/models/infer_onnx.py
Normal file
@@ -0,0 +1,244 @@
|
|||||||
|
from typing import Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import onnxruntime
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
|
from style_bert_vits2.nlp import (
|
||||||
|
clean_text_with_given_phone_tone,
|
||||||
|
cleaned_text_to_sequence,
|
||||||
|
extract_bert_feature_onnx,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def __intersperse(lst: list[Any], item: Any) -> list[Any]:
|
||||||
|
"""
|
||||||
|
リストの要素の間に特定のアイテムを挿入する
|
||||||
|
style_bert_vits2.models.commons.intersperse と同一実装
|
||||||
|
style_bert_vits2.models.commons モジュールは PyTorch に依存しているため、ONNX 推論時は import できない
|
||||||
|
|
||||||
|
Args:
|
||||||
|
lst (list[Any]): 元のリスト
|
||||||
|
item (Any): 挿入するアイテム
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[Any]: 新しいリスト
|
||||||
|
"""
|
||||||
|
result = [item] * (len(lst) * 2 + 1)
|
||||||
|
result[1::2] = lst
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def get_text_onnx(
|
||||||
|
text: str,
|
||||||
|
language_str: Languages,
|
||||||
|
hps: HyperParameters,
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
given_phone: Optional[list[str]] = None,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
) -> tuple[
|
||||||
|
NDArray[Any], NDArray[Any], NDArray[Any], NDArray[Any], NDArray[Any], NDArray[Any]
|
||||||
|
]:
|
||||||
|
use_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
|
norm_text, phone, tone, word2ph = clean_text_with_given_phone_tone(
|
||||||
|
text,
|
||||||
|
language_str,
|
||||||
|
given_phone=given_phone,
|
||||||
|
given_tone=given_tone,
|
||||||
|
use_jp_extra=use_jp_extra,
|
||||||
|
# 推論時のみ呼び出されるので、raise_yomi_error は False に設定
|
||||||
|
raise_yomi_error=False,
|
||||||
|
)
|
||||||
|
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
||||||
|
|
||||||
|
if hps.data.add_blank:
|
||||||
|
phone = __intersperse(phone, 0)
|
||||||
|
tone = __intersperse(tone, 0)
|
||||||
|
language = __intersperse(language, 0)
|
||||||
|
for i in range(len(word2ph)):
|
||||||
|
word2ph[i] = word2ph[i] * 2
|
||||||
|
word2ph[0] += 1
|
||||||
|
bert_ori = extract_bert_feature_onnx(
|
||||||
|
norm_text,
|
||||||
|
word2ph,
|
||||||
|
language_str,
|
||||||
|
onnx_providers,
|
||||||
|
assist_text,
|
||||||
|
assist_text_weight,
|
||||||
|
)
|
||||||
|
del word2ph
|
||||||
|
assert bert_ori.shape[-1] == len(phone), phone
|
||||||
|
|
||||||
|
if language_str == Languages.ZH:
|
||||||
|
bert = bert_ori
|
||||||
|
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
en_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
elif language_str == Languages.JP:
|
||||||
|
bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
ja_bert = bert_ori
|
||||||
|
en_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
elif language_str == Languages.EN:
|
||||||
|
bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
|
en_bert = bert_ori
|
||||||
|
else:
|
||||||
|
raise ValueError("language_str should be ZH, JP or EN")
|
||||||
|
|
||||||
|
assert bert.shape[-1] == len(
|
||||||
|
phone
|
||||||
|
), f"Bert seq len {bert.shape[-1]} != {len(phone)}"
|
||||||
|
|
||||||
|
phone = np.array(phone, dtype=np.int64)
|
||||||
|
tone = np.array(tone, dtype=np.int64)
|
||||||
|
language = np.array(language, dtype=np.int64)
|
||||||
|
return bert, ja_bert, en_bert, phone, tone, language
|
||||||
|
|
||||||
|
|
||||||
|
def infer_onnx(
|
||||||
|
text: str,
|
||||||
|
style_vec: NDArray[Any],
|
||||||
|
sdp_ratio: float,
|
||||||
|
noise_scale: float,
|
||||||
|
noise_scale_w: float,
|
||||||
|
length_scale: float,
|
||||||
|
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
|
||||||
|
language: Languages,
|
||||||
|
hps: HyperParameters,
|
||||||
|
onnx_session: onnxruntime.InferenceSession,
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
skip_start: bool = False,
|
||||||
|
skip_end: bool = False,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
given_phone: Optional[list[str]] = None,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
is_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text_onnx(
|
||||||
|
text,
|
||||||
|
language,
|
||||||
|
hps,
|
||||||
|
onnx_providers=onnx_providers,
|
||||||
|
assist_text=assist_text,
|
||||||
|
assist_text_weight=assist_text_weight,
|
||||||
|
given_phone=given_phone,
|
||||||
|
given_tone=given_tone,
|
||||||
|
)
|
||||||
|
if skip_start:
|
||||||
|
phones = phones[3:]
|
||||||
|
tones = tones[3:]
|
||||||
|
lang_ids = lang_ids[3:]
|
||||||
|
bert = bert[:, 3:]
|
||||||
|
ja_bert = ja_bert[:, 3:]
|
||||||
|
en_bert = en_bert[:, 3:]
|
||||||
|
if skip_end:
|
||||||
|
phones = phones[:-2]
|
||||||
|
tones = tones[:-2]
|
||||||
|
lang_ids = lang_ids[:-2]
|
||||||
|
bert = bert[:, :-2]
|
||||||
|
ja_bert = ja_bert[:, :-2]
|
||||||
|
en_bert = en_bert[:, :-2]
|
||||||
|
|
||||||
|
x_tst = np.expand_dims(phones, axis=0)
|
||||||
|
tones = np.expand_dims(tones, axis=0)
|
||||||
|
lang_ids = np.expand_dims(lang_ids, axis=0)
|
||||||
|
bert = np.expand_dims(bert, axis=0)
|
||||||
|
ja_bert = np.expand_dims(ja_bert, axis=0)
|
||||||
|
en_bert = np.expand_dims(en_bert, axis=0)
|
||||||
|
x_tst_lengths = np.array([phones.shape[0]], dtype=np.int64)
|
||||||
|
style_vec_tensor = np.expand_dims(style_vec, axis=0)
|
||||||
|
del phones
|
||||||
|
sid_tensor = np.array([sid], dtype=np.int64)
|
||||||
|
|
||||||
|
input_names = [input.name for input in onnx_session.get_inputs()]
|
||||||
|
output_name = onnx_session.get_outputs()[0].name
|
||||||
|
if is_jp_extra:
|
||||||
|
input_tensor = [
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
ja_bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
np.array(length_scale, dtype=np.float32),
|
||||||
|
np.array(sdp_ratio, dtype=np.float32),
|
||||||
|
np.array(noise_scale, dtype=np.float32),
|
||||||
|
np.array(noise_scale_w, dtype=np.float32),
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
input_tensor = [
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
np.array(length_scale, dtype=np.float32),
|
||||||
|
np.array(sdp_ratio, dtype=np.float32),
|
||||||
|
np.array(noise_scale, dtype=np.float32),
|
||||||
|
np.array(noise_scale_w, dtype=np.float32),
|
||||||
|
]
|
||||||
|
|
||||||
|
# 入力テンソルを転送する GPU デバイスを取得
|
||||||
|
## 本来は device_type="dml" もサポートされているはずだが、手元環境だと常に謎の RuntimeError が発生するため当面無効化している
|
||||||
|
first_provider = onnx_session.get_providers()[0]
|
||||||
|
if first_provider == "CUDAExecutionProvider":
|
||||||
|
device_type = "cuda"
|
||||||
|
# elif first_provider == "DmlExecutionProvider":
|
||||||
|
# device_type = "dml"
|
||||||
|
else:
|
||||||
|
device_type = "cpu"
|
||||||
|
|
||||||
|
# 入力テンソルを転送する GPU デバイスの ID を取得
|
||||||
|
## ExecutionProvider に指定したオプションの中から device_id を取得し、入力テンソルの転送先として指定する
|
||||||
|
## InferenceSession で利用するデバイス ID と入力テンソルの転送先デバイス ID は一致している必要がある
|
||||||
|
## 本来は ExecutionProvider に指定したオプションは InferenceSession.get_provider_options() で取得できるはずだが、
|
||||||
|
## 手元環境だと DmlExecutionProvider のみ常に空の辞書が返されるため、当面 onnx_providers から直接オプションを取り出している
|
||||||
|
device_id = 0
|
||||||
|
onnx_providers_dict: dict[str, dict[str, Any]] = {}
|
||||||
|
for provider in onnx_providers:
|
||||||
|
if isinstance(provider, tuple):
|
||||||
|
provider_name, options = provider
|
||||||
|
onnx_providers_dict[provider_name] = options
|
||||||
|
else:
|
||||||
|
onnx_providers_dict[provider] = {}
|
||||||
|
first_provider_options = onnx_providers_dict[first_provider]
|
||||||
|
if "device_id" in first_provider_options:
|
||||||
|
device_id = int(first_provider_options["device_id"])
|
||||||
|
|
||||||
|
# GPU メモリに入力テンソルを割り当て
|
||||||
|
io_binding = onnx_session.io_binding()
|
||||||
|
for name, value in zip(input_names, input_tensor):
|
||||||
|
gpu_tensor = onnxruntime.OrtValue.ortvalue_from_numpy(
|
||||||
|
value, device_type, device_id
|
||||||
|
)
|
||||||
|
io_binding.bind_ortvalue_input(name, gpu_tensor)
|
||||||
|
|
||||||
|
# 推論の実行
|
||||||
|
io_binding.bind_output(output_name, device_type)
|
||||||
|
onnx_session.run_with_iobinding(io_binding)
|
||||||
|
output = io_binding.get_outputs()
|
||||||
|
|
||||||
|
audio = output[0].numpy()[0, 0]
|
||||||
|
|
||||||
|
del (
|
||||||
|
x_tst,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec,
|
||||||
|
) # , emo
|
||||||
|
|
||||||
|
return audio
|
||||||
@@ -15,6 +15,7 @@ def load_checkpoint(
|
|||||||
optimizer: Optional[torch.optim.Optimizer] = None,
|
optimizer: Optional[torch.optim.Optimizer] = None,
|
||||||
skip_optimizer: bool = False,
|
skip_optimizer: bool = False,
|
||||||
for_infer: bool = False,
|
for_infer: bool = False,
|
||||||
|
device: Union[str, torch.device] = "cpu",
|
||||||
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
|
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
|
||||||
"""
|
"""
|
||||||
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
|
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
|
||||||
@@ -31,7 +32,7 @@ def load_checkpoint(
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
assert os.path.isfile(checkpoint_path)
|
assert os.path.isfile(checkpoint_path)
|
||||||
checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")
|
checkpoint_dict = torch.load(checkpoint_path, map_location=device)
|
||||||
iteration = checkpoint_dict["iteration"]
|
iteration = checkpoint_dict["iteration"]
|
||||||
learning_rate = checkpoint_dict["learning_rate"]
|
learning_rate = checkpoint_dict["learning_rate"]
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ def load_safetensors(
|
|||||||
checkpoint_path: Union[str, Path],
|
checkpoint_path: Union[str, Path],
|
||||||
model: torch.nn.Module,
|
model: torch.nn.Module,
|
||||||
for_infer: bool = False,
|
for_infer: bool = False,
|
||||||
|
device: Union[str, torch.device] = "cpu",
|
||||||
) -> tuple[torch.nn.Module, Optional[int]]:
|
) -> tuple[torch.nn.Module, Optional[int]]:
|
||||||
"""
|
"""
|
||||||
指定されたパスから safetensors モデルを読み込み、モデルとイテレーションを返す。
|
指定されたパスから safetensors モデルを読み込み、モデルとイテレーションを返す。
|
||||||
@@ -27,7 +28,7 @@ def load_safetensors(
|
|||||||
|
|
||||||
tensors: dict[str, Any] = {}
|
tensors: dict[str, Any] = {}
|
||||||
iteration: Optional[int] = None
|
iteration: Optional[int] = None
|
||||||
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as f: # type: ignore
|
with safe_open(str(checkpoint_path), framework="pt", device=device) as f: # type: ignore
|
||||||
for key in f.keys():
|
for key in f.keys():
|
||||||
if key == "iteration":
|
if key == "iteration":
|
||||||
iteration = f.get_tensor(key).item()
|
iteration = f.get_tensor(key).item()
|
||||||
|
|||||||
@@ -1,4 +1,8 @@
|
|||||||
from typing import TYPE_CHECKING, Optional
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
from style_bert_vits2.constants import Languages
|
from style_bert_vits2.constants import Languages
|
||||||
from style_bert_vits2.nlp.symbols import (
|
from style_bert_vits2.nlp.symbols import (
|
||||||
@@ -24,9 +28,9 @@ def extract_bert_feature(
|
|||||||
device: str,
|
device: str,
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> "torch.Tensor":
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
テキストから BERT の特徴量を抽出する
|
テキストから BERT の特徴量を抽出する (PyTorch 推論)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text (str): テキスト
|
text (str): テキスト
|
||||||
@@ -52,6 +56,47 @@ def extract_bert_feature(
|
|||||||
return extract_bert_feature(text, word2ph, device, assist_text, assist_text_weight)
|
return extract_bert_feature(text, word2ph, device, assist_text, assist_text_weight)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature_onnx(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
language: Languages,
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
テキストから BERT の特徴量を抽出する (ONNX 推論)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): テキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
language (Languages): テキストの言語
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language == Languages.JP:
|
||||||
|
from style_bert_vits2.nlp.japanese.bert_feature import extract_bert_feature_onnx
|
||||||
|
elif language == Languages.EN:
|
||||||
|
from style_bert_vits2.nlp.english.bert_feature import extract_bert_feature_onnx
|
||||||
|
elif language == Languages.ZH:
|
||||||
|
from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature_onnx
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Language {language} not supported")
|
||||||
|
|
||||||
|
return extract_bert_feature_onnx(
|
||||||
|
text,
|
||||||
|
word2ph,
|
||||||
|
onnx_providers,
|
||||||
|
assist_text,
|
||||||
|
assist_text_weight,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def clean_text(
|
def clean_text(
|
||||||
text: str,
|
text: str,
|
||||||
language: Languages,
|
language: Languages,
|
||||||
@@ -96,6 +141,87 @@ def clean_text(
|
|||||||
return norm_text, phones, tones, word2ph
|
return norm_text, phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
|
def clean_text_with_given_phone_tone(
|
||||||
|
text: str,
|
||||||
|
language: Languages,
|
||||||
|
given_phone: Optional[list[str]] = None,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
use_jp_extra: bool = True,
|
||||||
|
raise_yomi_error: bool = False,
|
||||||
|
) -> tuple[str, list[str], list[int], list[int]]:
|
||||||
|
"""
|
||||||
|
テキストをクリーニングし、音素に変換する
|
||||||
|
変換時、given_phone や given_tone が与えられた場合はそれを調整して使う
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): クリーニングするテキスト
|
||||||
|
language (Languages): テキストの言語
|
||||||
|
given_phone (Optional[list[int]], optional): 読み上げテキストの読みを表す音素列。指定する場合は given_tone も別途指定が必要. Defaults to None.
|
||||||
|
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
||||||
|
use_jp_extra (bool, optional): テキストが日本語の場合に JP-Extra モデルを利用するかどうか。Defaults to True.
|
||||||
|
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[str, list[str], list[int], list[int]]: クリーニングされたテキストと、音素・アクセント・元のテキストの各文字に音素が何個割り当てられるかのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 与えられたテキストをクリーニング
|
||||||
|
norm_text, phone, tone, word2ph = clean_text(
|
||||||
|
text,
|
||||||
|
language,
|
||||||
|
use_jp_extra=use_jp_extra,
|
||||||
|
raise_yomi_error=raise_yomi_error,
|
||||||
|
)
|
||||||
|
|
||||||
|
# phone と tone の両方が与えられた場合はそれを使う
|
||||||
|
if given_phone is not None and given_tone is not None:
|
||||||
|
# 指定された phone と指定された tone 両方の長さが一致していなければならない
|
||||||
|
if len(given_phone) != len(given_tone):
|
||||||
|
raise InvalidPhoneError(
|
||||||
|
f"Length of given_phone ({len(given_phone)}) != length of given_tone ({len(given_tone)})"
|
||||||
|
)
|
||||||
|
# 与えられた音素数と pyopenjtalk で生成した読みの音素数が一致しない
|
||||||
|
if len(given_phone) != sum(word2ph):
|
||||||
|
# 日本語の場合、len(given_phone) と sum(word2ph) が一致するように word2ph を適切に調整する
|
||||||
|
# 他の言語は word2ph の調整方法が思いつかないのでエラー
|
||||||
|
if language == Languages.JP:
|
||||||
|
from style_bert_vits2.nlp.japanese.g2p import adjust_word2ph
|
||||||
|
|
||||||
|
# use_jp_extra でない場合は given_phone 内の「N」を「n」に変換
|
||||||
|
if not use_jp_extra:
|
||||||
|
given_phone = [p if p != "N" else "n" for p in given_phone]
|
||||||
|
# clean_text() から取得した word2ph を調整結果で上書き
|
||||||
|
word2ph = adjust_word2ph(word2ph, phone, given_phone)
|
||||||
|
# 上記処理により word2ph の合計が given_phone の長さと一致するはず
|
||||||
|
# それでも一致しない場合、大半は読み上げテキストと given_phone が著しく乖離していて調整し切れなかったことを意味する
|
||||||
|
if len(given_phone) != sum(word2ph):
|
||||||
|
raise InvalidPhoneError(
|
||||||
|
f"Length of given_phone ({len(given_phone)}) != sum of word2ph ({sum(word2ph)})"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise InvalidPhoneError(
|
||||||
|
f"Length of given_phone ({len(given_phone)}) != sum of word2ph ({sum(word2ph)})"
|
||||||
|
)
|
||||||
|
phone = given_phone
|
||||||
|
# 生成あるいは指定された phone と指定された tone 両方の長さが一致していなければならない
|
||||||
|
if len(phone) != len(given_tone):
|
||||||
|
raise InvalidToneError(
|
||||||
|
f"Length of phone ({len(phone)}) != length of given_tone ({len(given_tone)})"
|
||||||
|
)
|
||||||
|
tone = given_tone
|
||||||
|
|
||||||
|
# tone だけが与えられた場合は clean_text() で生成した phone と合わせて使う
|
||||||
|
elif given_tone is not None:
|
||||||
|
# 生成した phone と指定された tone 両方の長さが一致していなければならない
|
||||||
|
if len(phone) != len(given_tone):
|
||||||
|
raise InvalidToneError(
|
||||||
|
f"Length of phone ({len(phone)}) != length of given_tone ({len(given_tone)})"
|
||||||
|
)
|
||||||
|
tone = given_tone
|
||||||
|
|
||||||
|
return norm_text, phone, tone, word2ph
|
||||||
|
|
||||||
|
|
||||||
def cleaned_text_to_sequence(
|
def cleaned_text_to_sequence(
|
||||||
cleaned_phones: list[str], tones: list[int], language: Languages
|
cleaned_phones: list[str], tones: list[int], language: Languages
|
||||||
) -> tuple[list[int], list[int], list[int]]:
|
) -> tuple[list[int], list[int], list[int]]:
|
||||||
@@ -118,3 +244,11 @@ def cleaned_text_to_sequence(
|
|||||||
lang_ids = [lang_id for i in phones]
|
lang_ids = [lang_id for i in phones]
|
||||||
|
|
||||||
return phones, tones, lang_ids
|
return phones, tones, lang_ids
|
||||||
|
|
||||||
|
|
||||||
|
class InvalidPhoneError(ValueError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class InvalidToneError(ValueError):
|
||||||
|
pass
|
||||||
|
|||||||
@@ -8,22 +8,29 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル
|
|||||||
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
|
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import gc
|
from __future__ import annotations
|
||||||
from typing import Optional, Union, cast
|
|
||||||
|
import gc
|
||||||
|
import time
|
||||||
|
from typing import TYPE_CHECKING, Optional, Union, cast
|
||||||
|
|
||||||
import torch
|
|
||||||
from transformers import (
|
from transformers import (
|
||||||
AutoModelForMaskedLM,
|
AutoModelForMaskedLM,
|
||||||
AutoTokenizer,
|
AutoTokenizer,
|
||||||
DebertaV2Model,
|
DebertaV2Model,
|
||||||
DebertaV2Tokenizer,
|
DebertaV2TokenizerFast,
|
||||||
PreTrainedModel,
|
PreTrainedModel,
|
||||||
PreTrainedTokenizer,
|
PreTrainedTokenizer,
|
||||||
PreTrainedTokenizerFast,
|
PreTrainedTokenizerFast,
|
||||||
)
|
)
|
||||||
|
|
||||||
from style_bert_vits2.constants import DEFAULT_BERT_TOKENIZER_PATHS, Languages
|
from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.nlp import onnx_bert_models
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
||||||
@@ -31,13 +38,17 @@ __loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
|||||||
|
|
||||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
__loaded_tokenizers: dict[
|
__loaded_tokenizers: dict[
|
||||||
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
|
Languages,
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast],
|
||||||
] = {}
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
def load_model(
|
def load_model(
|
||||||
language: Languages,
|
language: Languages,
|
||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
device_map: Optional[
|
||||||
|
Union[str, dict[str, Union[int, str, torch.device]], int, torch.device]
|
||||||
|
] = None,
|
||||||
cache_dir: Optional[str] = None,
|
cache_dir: Optional[str] = None,
|
||||||
revision: str = "main",
|
revision: str = "main",
|
||||||
) -> Union[PreTrainedModel, DebertaV2Model]:
|
) -> Union[PreTrainedModel, DebertaV2Model]:
|
||||||
@@ -46,6 +57,7 @@ def load_model(
|
|||||||
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
||||||
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
device_map は既に指定された言語の BERT モデルがロードされている場合は効果がない。
|
||||||
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている。
|
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている。
|
||||||
@@ -57,6 +69,9 @@ def load_model(
|
|||||||
Args:
|
Args:
|
||||||
language (Languages): ロードする学習済みモデルの対象言語
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
device_map (Optional[str]): accelerate を使用して高速にデバイスにモデルをロードするためのデバイスマップ。
|
||||||
|
指定しない場合は通常のモデルロード処理になる (デフォルト: None)
|
||||||
|
ref: https://huggingface.co/docs/accelerate/usage_guides/big_modeling
|
||||||
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
@@ -70,30 +85,35 @@ def load_model(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
assert DEFAULT_BERT_MODEL_PATHS[language].exists(), \
|
||||||
language
|
f"The default {language.name} BERT model does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip
|
||||||
].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
|
||||||
|
|
||||||
# BERT モデルをロードし、辞書に格納して返す
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Model でロードする必要がある
|
## 英語のみ DebertaV2Model でロードする必要がある
|
||||||
|
start_time = time.time()
|
||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
model = cast(
|
__loaded_models[language] = cast(
|
||||||
DebertaV2Model,
|
DebertaV2Model,
|
||||||
DebertaV2Model.from_pretrained(
|
DebertaV2Model.from_pretrained(
|
||||||
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
pretrained_model_name_or_path,
|
||||||
|
device_map=device_map,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
model = AutoModelForMaskedLM.from_pretrained(
|
__loaded_models[language] = AutoModelForMaskedLM.from_pretrained(
|
||||||
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
pretrained_model_name_or_path,
|
||||||
|
device_map=device_map,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
)
|
)
|
||||||
__loaded_models[language] = model
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
f"Loaded the {language.name} BERT model from {pretrained_model_name_or_path} ({time.time() - start_time:.2f}s)"
|
||||||
)
|
)
|
||||||
|
|
||||||
return model
|
return __loaded_models[language]
|
||||||
|
|
||||||
|
|
||||||
def load_tokenizer(
|
def load_tokenizer(
|
||||||
@@ -101,9 +121,9 @@ def load_tokenizer(
|
|||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
cache_dir: Optional[str] = None,
|
cache_dir: Optional[str] = None,
|
||||||
revision: str = "main",
|
revision: str = "main",
|
||||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]:
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]:
|
||||||
"""
|
"""
|
||||||
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
指定された言語の BERT トークナイザーをロードし、ロード済みの BERT トークナイザーを返す。
|
||||||
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
||||||
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
@@ -131,31 +151,78 @@ def load_tokenizer(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
# ライブラリ利用時、特例的にこの状況で ONNX 版 BERT トークナイザーがロードされている場合はそのまま返す
|
||||||
language
|
## ONNX 版 BERT トークナイザー単独で g2p 処理を行うために必要 (各言語の g2p.py はこの関数に依存している)
|
||||||
].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
## 設計的には微妙だがこの方が差異を吸収できて手っ取り早い
|
||||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
if DEFAULT_BERT_MODEL_PATHS[language].exists() is False and onnx_bert_models.is_tokenizer_loaded(language): # fmt: skip
|
||||||
|
return onnx_bert_models.load_tokenizer(language)
|
||||||
|
assert DEFAULT_BERT_MODEL_PATHS[language].exists(), \
|
||||||
|
f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip
|
||||||
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
## 英語のみ DebertaV2TokenizerFast でロードする必要がある
|
||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
tokenizer = DebertaV2Tokenizer.from_pretrained(
|
__loaded_tokenizers[language] = DebertaV2TokenizerFast.from_pretrained(
|
||||||
pretrained_model_name_or_path,
|
pretrained_model_name_or_path,
|
||||||
cache_dir=cache_dir,
|
cache_dir=cache_dir,
|
||||||
revision=revision,
|
revision=revision,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
tokenizer = AutoTokenizer.from_pretrained(
|
__loaded_tokenizers[language] = AutoTokenizer.from_pretrained(
|
||||||
pretrained_model_name_or_path,
|
pretrained_model_name_or_path,
|
||||||
cache_dir=cache_dir,
|
cache_dir=cache_dir,
|
||||||
revision=revision,
|
revision=revision,
|
||||||
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
)
|
)
|
||||||
__loaded_tokenizers[language] = tokenizer
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
f"Loaded the {language.name} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
)
|
)
|
||||||
|
|
||||||
return tokenizer
|
return __loaded_tokenizers[language]
|
||||||
|
|
||||||
|
|
||||||
|
def transfer_model(language: Languages, device: str) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルを、指定されたデバイスに移動する。
|
||||||
|
モデルのロード後に推論デバイスを変更したい場合に利用する。
|
||||||
|
既に指定されたデバイスにモデルがロードされている場合は何も行われない。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): モデルを移動する言語
|
||||||
|
device (str): モデルを移動するデバイス
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language not in __loaded_models:
|
||||||
|
raise ValueError(f"BERT model for {language.name} is not loaded.")
|
||||||
|
|
||||||
|
# 既に指定されたデバイスにモデルがロードされている場合は何もしない
|
||||||
|
# ex: current_device="cuda:0", device="cuda" → 何もしない
|
||||||
|
# ex: current_device="cuda:0", device="cpu" → モデルを CPU に移動
|
||||||
|
current_device = str(__loaded_models[language].device)
|
||||||
|
if current_device.startswith(device):
|
||||||
|
return
|
||||||
|
|
||||||
|
__loaded_models[language].to(device) # type: ignore
|
||||||
|
logger.info(
|
||||||
|
f"Transferred the {language.name} BERT model from {current_device} to {device}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_model_loaded(language: Languages) -> bool:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルがロード済みかどうかを返す。
|
||||||
|
"""
|
||||||
|
|
||||||
|
return language in __loaded_models
|
||||||
|
|
||||||
|
|
||||||
|
def is_tokenizer_loaded(language: Languages) -> bool:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT トークナイザーがロード済みかどうかを返す。
|
||||||
|
"""
|
||||||
|
|
||||||
|
return language in __loaded_tokenizers
|
||||||
|
|
||||||
|
|
||||||
def unload_model(language: Languages) -> None:
|
def unload_model(language: Languages) -> None:
|
||||||
@@ -166,12 +233,14 @@ def unload_model(language: Languages) -> None:
|
|||||||
language (Languages): アンロードする BERT モデルの言語
|
language (Languages): アンロードする BERT モデルの言語
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
if language in __loaded_models:
|
if language in __loaded_models:
|
||||||
del __loaded_models[language]
|
del __loaded_models[language]
|
||||||
gc.collect()
|
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
logger.info(f"Unloaded the {language} BERT model")
|
gc.collect()
|
||||||
|
logger.info(f"Unloaded the {language.name} BERT model")
|
||||||
|
|
||||||
|
|
||||||
def unload_tokenizer(language: Languages) -> None:
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
@@ -182,12 +251,12 @@ def unload_tokenizer(language: Languages) -> None:
|
|||||||
language (Languages): アンロードする BERT トークナイザーの言語
|
language (Languages): アンロードする BERT トークナイザーの言語
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
if language in __loaded_tokenizers:
|
if language in __loaded_tokenizers:
|
||||||
del __loaded_tokenizers[language]
|
del __loaded_tokenizers[language]
|
||||||
gc.collect()
|
gc.collect()
|
||||||
if torch.cuda.is_available():
|
logger.info(f"Unloaded the {language.name} BERT tokenizer")
|
||||||
torch.cuda.empty_cache()
|
|
||||||
logger.info(f"Unloaded the {language} BERT tokenizer")
|
|
||||||
|
|
||||||
|
|
||||||
def unload_all_models() -> None:
|
def unload_all_models() -> None:
|
||||||
|
|||||||
@@ -1,9 +1,16 @@
|
|||||||
from typing import Optional
|
from __future__ import annotations
|
||||||
|
|
||||||
import torch
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
from style_bert_vits2.constants import Languages
|
from style_bert_vits2.constants import Languages
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
def extract_bert_feature(
|
def extract_bert_feature(
|
||||||
@@ -14,7 +21,7 @@ def extract_bert_feature(
|
|||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
中国語のテキストから BERT の特徴量を抽出する
|
中国語のテキストから BERT の特徴量を抽出する (PyTorch 推論)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text (str): 中国語のテキスト
|
text (str): 中国語のテキスト
|
||||||
@@ -27,9 +34,12 @@ def extract_bert_feature(
|
|||||||
torch.Tensor: BERT の特徴量
|
torch.Tensor: BERT の特徴量
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.ZH).to(device) # type: ignore
|
model = bert_models.load_model(Languages.ZH, device_map=device)
|
||||||
|
bert_models.transfer_model(Languages.ZH, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
@@ -66,6 +76,76 @@ def extract_bert_feature(
|
|||||||
return phone_level_feature.T
|
return phone_level_feature.T
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature_onnx(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
中国語のテキストから BERT の特徴量を抽出する (ONNX 推論)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 中国語のテキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
tokenizer = onnx_bert_models.load_tokenizer(Languages.ZH)
|
||||||
|
inputs = tokenizer(text, return_tensors="np")
|
||||||
|
|
||||||
|
session = onnx_bert_models.load_model(
|
||||||
|
language=Languages.ZH,
|
||||||
|
onnx_providers=onnx_providers,
|
||||||
|
)
|
||||||
|
output_name = session.get_outputs()[0].name
|
||||||
|
res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"token_type_ids": inputs["token_type_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="np")
|
||||||
|
style_res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"token_type_ids": style_inputs["token_type_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
style_res_mean = np.mean(style_res, axis=0)
|
||||||
|
|
||||||
|
assert len(word2ph) == len(text) + 2
|
||||||
|
word2phone = word2ph
|
||||||
|
phone_level_feature = []
|
||||||
|
for i in range(len(word2phone)):
|
||||||
|
if assist_text:
|
||||||
|
assert style_res_mean is not None
|
||||||
|
repeat_feature = (
|
||||||
|
np.tile(res[i], (word2phone[i], 1)) * (1 - assist_text_weight)
|
||||||
|
+ np.tile(style_res_mean, (word2phone[i], 1)) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = np.tile(res[i], (word2phone[i], 1))
|
||||||
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
|
phone_level_feature = np.concatenate(phone_level_feature, axis=0)
|
||||||
|
|
||||||
|
return phone_level_feature.T
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
word_level_feature = torch.rand(38, 1024) # 12个词,每个词1024维特征
|
word_level_feature = torch.rand(38, 1024) # 12个词,每个词1024维特征
|
||||||
word2phone = [
|
word2phone = [
|
||||||
|
|||||||
@@ -1,9 +1,16 @@
|
|||||||
from typing import Optional
|
from __future__ import annotations
|
||||||
|
|
||||||
import torch
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
from style_bert_vits2.constants import Languages
|
from style_bert_vits2.constants import Languages
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
def extract_bert_feature(
|
def extract_bert_feature(
|
||||||
@@ -14,7 +21,7 @@ def extract_bert_feature(
|
|||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
英語のテキストから BERT の特徴量を抽出する
|
英語のテキストから BERT の特徴量を抽出する (PyTorch 推論)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text (str): 英語のテキスト
|
text (str): 英語のテキスト
|
||||||
@@ -27,9 +34,12 @@ def extract_bert_feature(
|
|||||||
torch.Tensor: BERT の特徴量
|
torch.Tensor: BERT の特徴量
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.EN).to(device) # type: ignore
|
model = bert_models.load_model(Languages.EN, device_map=device)
|
||||||
|
bert_models.transfer_model(Languages.EN, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
@@ -64,3 +74,71 @@ def extract_bert_feature(
|
|||||||
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||||
|
|
||||||
return phone_level_feature.T
|
return phone_level_feature.T
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature_onnx(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
英語のテキストから BERT の特徴量を抽出する (ONNX 推論)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 英語のテキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
tokenizer = onnx_bert_models.load_tokenizer(Languages.EN)
|
||||||
|
inputs = tokenizer(text, return_tensors="np")
|
||||||
|
|
||||||
|
session = onnx_bert_models.load_model(
|
||||||
|
language=Languages.EN,
|
||||||
|
onnx_providers=onnx_providers,
|
||||||
|
)
|
||||||
|
output_name = session.get_outputs()[0].name
|
||||||
|
res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="np")
|
||||||
|
style_res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
style_res_mean = np.mean(style_res, axis=0)
|
||||||
|
|
||||||
|
assert len(word2ph) == res.shape[0], (text, res.shape[0], len(word2ph))
|
||||||
|
word2phone = word2ph
|
||||||
|
phone_level_feature = []
|
||||||
|
for i in range(len(word2phone)):
|
||||||
|
if assist_text:
|
||||||
|
assert style_res_mean is not None
|
||||||
|
repeat_feature = (
|
||||||
|
np.tile(res[i], (word2phone[i], 1)) * (1 - assist_text_weight)
|
||||||
|
+ np.tile(style_res_mean, (word2phone[i], 1)) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = np.tile(res[i], (word2phone[i], 1))
|
||||||
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
|
phone_level_feature = np.concatenate(phone_level_feature, axis=0)
|
||||||
|
|
||||||
|
return phone_level_feature.T
|
||||||
|
|||||||
@@ -1,12 +1,19 @@
|
|||||||
from typing import Optional
|
from __future__ import annotations
|
||||||
|
|
||||||
import torch
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
from style_bert_vits2.constants import Languages
|
from style_bert_vits2.constants import Languages
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
def extract_bert_feature(
|
def extract_bert_feature(
|
||||||
text: str,
|
text: str,
|
||||||
word2ph: list[int],
|
word2ph: list[int],
|
||||||
@@ -15,7 +22,7 @@ def extract_bert_feature(
|
|||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
日本語のテキストから BERT の特徴量を抽出する
|
日本語のテキストから BERT の特徴量を抽出する (PyTorch 推論)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text (str): 日本語のテキスト
|
text (str): 日本語のテキスト
|
||||||
@@ -28,6 +35,8 @@ def extract_bert_feature(
|
|||||||
torch.Tensor: BERT の特徴量
|
torch.Tensor: BERT の特徴量
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
# 各単語が何文字かを作る `word2ph` を使う必要があるので、読めない文字は必ず無視する
|
# 各単語が何文字かを作る `word2ph` を使う必要があるので、読めない文字は必ず無視する
|
||||||
# でないと `word2ph` の結果とテキストの文字数結果が整合性が取れない
|
# でないと `word2ph` の結果とテキストの文字数結果が整合性が取れない
|
||||||
text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0])
|
text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0])
|
||||||
@@ -36,7 +45,8 @@ def extract_bert_feature(
|
|||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.JP).to(device) # type: ignore
|
model = bert_models.load_model(Languages.JP, device_map=device)
|
||||||
|
bert_models.transfer_model(Languages.JP, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
@@ -71,3 +81,77 @@ def extract_bert_feature(
|
|||||||
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||||
|
|
||||||
return phone_level_feature.T
|
return phone_level_feature.T
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature_onnx(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
日本語のテキストから BERT の特徴量を抽出する (ONNX 推論)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 日本語のテキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 各単語が何文字かを作る `word2ph` を使う必要があるので、読めない文字は必ず無視する
|
||||||
|
# でないと `word2ph` の結果とテキストの文字数結果が整合性が取れない
|
||||||
|
text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0])
|
||||||
|
if assist_text:
|
||||||
|
assist_text = "".join(text_to_sep_kata(assist_text, raise_yomi_error=False)[0])
|
||||||
|
|
||||||
|
tokenizer = onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
inputs = tokenizer(text, return_tensors="np")
|
||||||
|
|
||||||
|
session = onnx_bert_models.load_model(
|
||||||
|
language=Languages.JP,
|
||||||
|
onnx_providers=onnx_providers,
|
||||||
|
)
|
||||||
|
output_name = session.get_outputs()[0].name
|
||||||
|
res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="np")
|
||||||
|
style_res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
|
||||||
|
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
style_res_mean = np.mean(style_res, axis=0)
|
||||||
|
|
||||||
|
assert len(word2ph) == len(text) + 2, text
|
||||||
|
word2phone = word2ph
|
||||||
|
phone_level_feature = []
|
||||||
|
for i in range(len(word2phone)):
|
||||||
|
if assist_text:
|
||||||
|
assert style_res_mean is not None
|
||||||
|
repeat_feature = (
|
||||||
|
np.tile(res[i], (word2phone[i], 1)) * (1 - assist_text_weight)
|
||||||
|
+ np.tile(style_res_mean, (word2phone[i], 1)) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = np.tile(res[i], (word2phone[i], 1))
|
||||||
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
|
phone_level_feature = np.concatenate(phone_level_feature, axis=0)
|
||||||
|
|
||||||
|
return phone_level_feature.T
|
||||||
|
|||||||
@@ -289,7 +289,7 @@ def adjust_word2ph(
|
|||||||
current_generated_index = 0
|
current_generated_index = 0
|
||||||
|
|
||||||
# word2ph の要素数 (=正規化された読み上げテキストの文字数) を維持しながら、差分情報を使って word2ph を修正
|
# word2ph の要素数 (=正規化された読み上げテキストの文字数) を維持しながら、差分情報を使って word2ph を修正
|
||||||
## 音素数が generated_phone と given_phone で異なる場合にこの align_word2ph() が呼び出される
|
## 音素数が generated_phone と given_phone で異なる場合にこの adjust_word2ph() が呼び出される
|
||||||
## word2ph は正規化された読み上げテキストの文字数に対応しているので、要素数はそのまま given_phone で増減した音素数に合わせて各要素の値を増減する
|
## word2ph は正規化された読み上げテキストの文字数に対応しているので、要素数はそのまま given_phone で増減した音素数に合わせて各要素の値を増減する
|
||||||
for word2ph_element_index, word2ph_element in enumerate(word2ph):
|
for word2ph_element_index, word2ph_element in enumerate(word2ph):
|
||||||
# ここの word2ph_element は、正規化された読み上げテキストの各文字に割り当てられる音素の数を示す
|
# ここの word2ph_element は、正規化された読み上げテキストの各文字に割り当てられる音素の数を示す
|
||||||
@@ -717,5 +717,5 @@ class YomiError(Exception):
|
|||||||
"""
|
"""
|
||||||
OpenJTalk で、読みが正しく取得できない箇所があるときに発生する例外。
|
OpenJTalk で、読みが正しく取得できない箇所があるときに発生する例外。
|
||||||
基本的に「学習の前処理のテキスト処理時」には発生させ、そうでない場合は、
|
基本的に「学習の前処理のテキスト処理時」には発生させ、そうでない場合は、
|
||||||
ignore_yomi_error=True にしておいて、この例外を発生させないようにする。
|
raise_yomi_error=False にしておいて、この例外を発生させないようにする。
|
||||||
"""
|
"""
|
||||||
|
|||||||
237
style_bert_vits2/nlp/onnx_bert_models.py
Normal file
237
style_bert_vits2/nlp/onnx_bert_models.py
Normal file
@@ -0,0 +1,237 @@
|
|||||||
|
"""
|
||||||
|
Style-Bert-VITS2 の ONNX 推論に必要な各言語ごとの ONNX 版 BERT モデルをロード/取得するためのモジュール。
|
||||||
|
このモジュールは style_bert_vits2.nlp.bert_models での実装を ONNX 推論向けに変更したもの。
|
||||||
|
|
||||||
|
オリジナルの Bert-VITS2 では各言語ごとの BERT モデルが初回インポート時にハードコードされたパスから「暗黙的に」ロードされているが、
|
||||||
|
場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。
|
||||||
|
|
||||||
|
そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。
|
||||||
|
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import onnxruntime
|
||||||
|
from huggingface_hub import hf_hub_download
|
||||||
|
from transformers import (
|
||||||
|
AutoTokenizer,
|
||||||
|
DebertaV2TokenizerFast,
|
||||||
|
PreTrainedTokenizer,
|
||||||
|
PreTrainedTokenizerFast,
|
||||||
|
)
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_ONNX_BERT_MODEL_PATHS, Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
||||||
|
__loaded_models: dict[Languages, onnxruntime.InferenceSession] = {}
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
|
__loaded_tokenizers: dict[
|
||||||
|
Languages,
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast],
|
||||||
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = ["CPUExecutionProvider"],
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> onnxruntime.InferenceSession: # fmt: skip
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
|
||||||
|
一度ロードされていれば、ロード済みの ONNX 版 BERT モデルを即座に返す。
|
||||||
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
|
Style-Bert-VITS2 では、ONNX 版 BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の ONNX 版 BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: tsukumijima/deberta-v2-large-japanese-char-wwm-onnx
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
onnxruntime.InferenceSession: ロード済みの BERT モデル
|
||||||
|
"""
|
||||||
|
|
||||||
|
# すでにロード済みの場合はそのまま返す
|
||||||
|
if language in __loaded_models:
|
||||||
|
return __loaded_models[language]
|
||||||
|
|
||||||
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
|
if pretrained_model_name_or_path is None:
|
||||||
|
assert DEFAULT_ONNX_BERT_MODEL_PATHS[language].exists(), \
|
||||||
|
f"The default {language.name} ONNX BERT model does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip
|
||||||
|
pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
|
# pretrained_model_name_or_path に Hugging Face のリポジトリ名が指定された場合 (aaaa/bbbb のフォーマットを想定):
|
||||||
|
# 指定された revision の ONNX 版 BERT モデルを cache_dir にダウンロードする (既にダウンロード済みの場合は何も行われない)
|
||||||
|
if len(pretrained_model_name_or_path.split("/")) == 2:
|
||||||
|
model_path = Path(
|
||||||
|
hf_hub_download(
|
||||||
|
repo_id=pretrained_model_name_or_path,
|
||||||
|
filename="model.onnx",
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# pretrained_model_name_or_path にファイルパスが指定された場合:
|
||||||
|
# 既にダウンロード済みという前提のもと、モデルへのローカルパスを model_path に格納する
|
||||||
|
else:
|
||||||
|
model_path = Path(pretrained_model_name_or_path).resolve() / "model.onnx"
|
||||||
|
|
||||||
|
start_time = time.time()
|
||||||
|
sess_options = onnxruntime.SessionOptions()
|
||||||
|
# ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
|
||||||
|
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
|
||||||
|
# エラー以外のログを出力しない
|
||||||
|
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
|
||||||
|
sess_options.log_severity_level = 3
|
||||||
|
onnxruntime.set_default_logger_severity(3)
|
||||||
|
|
||||||
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
|
__loaded_models[language] = onnxruntime.InferenceSession(
|
||||||
|
model_path,
|
||||||
|
sess_options=sess_options,
|
||||||
|
providers=onnx_providers,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language.name} ONNX BERT model from {pretrained_model_name_or_path} ({time.time() - start_time:.2f}s)"
|
||||||
|
)
|
||||||
|
|
||||||
|
return __loaded_models[language]
|
||||||
|
|
||||||
|
|
||||||
|
def load_tokenizer(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT トークナイザーをロードし、ロード済みの ONNX 版 BERT トークナイザーを返す。
|
||||||
|
一度ロードされていれば、ロード済みの ONNX 版 BERT トークナイザーを即座に返す。
|
||||||
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
|
Style-Bert-VITS2 では、ONNX 版 BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の ONNX 版 BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: tsukumijima/deberta-v2-large-japanese-char-wwm-onnx
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー
|
||||||
|
"""
|
||||||
|
|
||||||
|
# すでにロード済みの場合はそのまま返す
|
||||||
|
if language in __loaded_tokenizers:
|
||||||
|
return __loaded_tokenizers[language]
|
||||||
|
|
||||||
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
|
if pretrained_model_name_or_path is None:
|
||||||
|
assert DEFAULT_ONNX_BERT_MODEL_PATHS[language].exists(), \
|
||||||
|
f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip
|
||||||
|
pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
|
## 英語のみ DebertaV2TokenizerFast でロードする必要がある
|
||||||
|
if language == Languages.EN:
|
||||||
|
__loaded_tokenizers[language] = DebertaV2TokenizerFast.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
__loaded_tokenizers[language] = AutoTokenizer.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language.name} ONNX BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return __loaded_tokenizers[language]
|
||||||
|
|
||||||
|
|
||||||
|
def is_model_loaded(language: Languages) -> bool:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT モデルがロード済みかどうかを返す。
|
||||||
|
"""
|
||||||
|
|
||||||
|
return language in __loaded_models
|
||||||
|
|
||||||
|
|
||||||
|
def is_tokenizer_loaded(language: Languages) -> bool:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT トークナイザーがロード済みかどうかを返す。
|
||||||
|
"""
|
||||||
|
|
||||||
|
return language in __loaded_tokenizers
|
||||||
|
|
||||||
|
|
||||||
|
def unload_model(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT モデルをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT モデルの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_models:
|
||||||
|
del __loaded_models[language]
|
||||||
|
gc.collect()
|
||||||
|
logger.info(f"Unloaded the {language.name} ONNX BERT model")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT トークナイザーをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT トークナイザーの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_tokenizers:
|
||||||
|
del __loaded_tokenizers[language]
|
||||||
|
gc.collect()
|
||||||
|
logger.info(f"Unloaded the {language.name} ONNX BERT tokenizer")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_models() -> None:
|
||||||
|
"""
|
||||||
|
すべての ONNX 版 BERT モデルをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_models.keys()):
|
||||||
|
unload_model(language)
|
||||||
|
logger.info("Unloaded all ONNX BERT models")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_tokenizers() -> None:
|
||||||
|
"""
|
||||||
|
すべての ONNX 版 BERT トークナイザーをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_tokenizers.keys()):
|
||||||
|
unload_tokenizer(language)
|
||||||
|
logger.info("Unloaded all ONNX BERT tokenizers")
|
||||||
@@ -1,10 +1,14 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional, Union
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import onnxruntime
|
||||||
from numpy.typing import NDArray
|
from numpy.typing import NDArray
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from style_bert_vits2.constants import (
|
from style_bert_vits2.constants import (
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
@@ -20,24 +24,33 @@ from style_bert_vits2.constants import (
|
|||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
from style_bert_vits2.models.infer import get_net_g, infer
|
|
||||||
from style_bert_vits2.models.models import SynthesizerTrn
|
|
||||||
from style_bert_vits2.models.models_jp_extra import (
|
|
||||||
SynthesizerTrn as SynthesizerTrnJPExtra,
|
|
||||||
)
|
|
||||||
from style_bert_vits2.voice import adjust_voice
|
from style_bert_vits2.voice import adjust_voice
|
||||||
|
|
||||||
|
|
||||||
# Gradio の import は重いため、ここでは型チェック時のみ import する
|
if TYPE_CHECKING:
|
||||||
# ライブラリとしての利用を考慮し、TTSModelHolder の _for_gradio() 系メソッド以外では Gradio に依存しないようにする
|
from style_bert_vits2.models.models import SynthesizerTrn
|
||||||
# _for_gradio() 系メソッドの戻り値の型アノテーションを文字列としているのは、Gradio なしで実行できるようにするため
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
# if TYPE_CHECKING:
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
# import gradio as gr
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class NullModelParam(BaseModel):
|
||||||
|
"""
|
||||||
|
ヌルモデルのパラメータを表す Pydantic モデル。
|
||||||
|
各パラメータは 0.0 から 1.0 の範囲で指定する。
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: str # モデル名
|
||||||
|
path: Path # モデルファイルのパス
|
||||||
|
weight: float = Field(ge=0.0, le=1.0) # 声質の重み
|
||||||
|
pitch: float = Field(ge=0.0, le=1.0) # 声の高さの重み
|
||||||
|
style: float = Field(ge=0.0, le=1.0) # 話し方の重み
|
||||||
|
tempo: float = Field(ge=0.0, le=1.0) # テンポの重み
|
||||||
|
|
||||||
|
|
||||||
class TTSModel:
|
class TTSModel:
|
||||||
"""
|
"""
|
||||||
Style-Bert-Vits2 の音声合成モデルを操作するクラス。
|
Style-Bert-VITS2 の音声合成モデルを操作するクラス。
|
||||||
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
|
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -46,22 +59,30 @@ class TTSModel:
|
|||||||
model_path: Path,
|
model_path: Path,
|
||||||
config_path: Union[Path, HyperParameters],
|
config_path: Union[Path, HyperParameters],
|
||||||
style_vec_path: Union[Path, NDArray[Any]],
|
style_vec_path: Union[Path, NDArray[Any]],
|
||||||
device: str,
|
device: str = "cpu",
|
||||||
) -> None:
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = ["CPUExecutionProvider"],
|
||||||
|
) -> None: # fmt: skip
|
||||||
"""
|
"""
|
||||||
Style-Bert-Vits2 の音声合成モデルを初期化する。
|
Style-Bert-VITS2 の音声合成モデルを初期化する。
|
||||||
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_path (Path): モデル (.safetensors) のパス
|
model_path (Path): モデル (.safetensors / .onnx) のパス
|
||||||
config_path (Union[Path, HyperParameters]): ハイパーパラメータ (config.json) のパス (直接 HyperParameters を指定することも可能)
|
config_path (Union[Path, HyperParameters]): ハイパーパラメータ (config.json) のパス (直接 HyperParameters を指定することも可能)
|
||||||
style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能)
|
style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能)
|
||||||
device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など)
|
device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など)
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
self.model_path: Path = model_path
|
self.model_path: Path = model_path
|
||||||
self.device: str = device
|
self.device: str = device
|
||||||
self.null_model_params: dict[int, dict[str, Union[float, str]]] = {}
|
self.onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = onnx_providers # fmt: skip
|
||||||
|
|
||||||
|
# ONNX 形式のモデルかどうか
|
||||||
|
if self.model_path.suffix == ".onnx":
|
||||||
|
self.is_onnx_model = True
|
||||||
|
else:
|
||||||
|
self.is_onnx_model = False
|
||||||
|
|
||||||
# ハイパーパラメータの Pydantic モデルが直接指定された
|
# ハイパーパラメータの Pydantic モデルが直接指定された
|
||||||
if isinstance(config_path, HyperParameters):
|
if isinstance(config_path, HyperParameters):
|
||||||
@@ -77,11 +98,11 @@ class TTSModel:
|
|||||||
# スタイルベクトルの NDArray が直接指定された
|
# スタイルベクトルの NDArray が直接指定された
|
||||||
if isinstance(style_vec_path, np.ndarray):
|
if isinstance(style_vec_path, np.ndarray):
|
||||||
self.style_vec_path: Path = Path("") # 互換性のため空の Path を設定
|
self.style_vec_path: Path = Path("") # 互換性のため空の Path を設定
|
||||||
self.__style_vectors: NDArray[Any] = style_vec_path
|
self.style_vectors: NDArray[Any] = style_vec_path
|
||||||
# スタイルベクトルのパスが指定された
|
# スタイルベクトルのパスが指定された
|
||||||
else:
|
else:
|
||||||
self.style_vec_path: Path = style_vec_path
|
self.style_vec_path: Path = style_vec_path
|
||||||
self.__style_vectors: NDArray[Any] = np.load(self.style_vec_path)
|
self.style_vectors: NDArray[Any] = np.load(self.style_vec_path)
|
||||||
|
|
||||||
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
||||||
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
||||||
@@ -96,59 +117,141 @@ class TTSModel:
|
|||||||
f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.__style_vectors.shape[0] != num_styles:
|
if self.style_vectors.shape[0] != num_styles:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The number of styles ({num_styles}) does not match the number of style vectors ({self.__style_vectors.shape[0]})"
|
f"The number of styles ({num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})"
|
||||||
)
|
)
|
||||||
self.__style_vector_inference: Optional[Any] = None
|
self.style_vector_inference: Optional[Any] = None
|
||||||
|
|
||||||
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
# net_g / null_model_params は PyTorch 推論時のみ遅延初期化される
|
||||||
|
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
||||||
|
self.null_model_params: Optional[dict[int, NullModelParam]] = None
|
||||||
|
|
||||||
|
# onnx_session は ONNX 推論時のみ遅延初期化される
|
||||||
|
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
|
||||||
|
|
||||||
def load(self) -> None:
|
def load(self) -> None:
|
||||||
"""
|
"""
|
||||||
音声合成モデルをデバイスにロードする。
|
音声合成モデルをデバイスにロードする。
|
||||||
"""
|
"""
|
||||||
self.__net_g = get_net_g(
|
|
||||||
model_path=str(self.model_path),
|
|
||||||
version=self.hyper_parameters.version,
|
|
||||||
device=self.device,
|
|
||||||
hps=self.hyper_parameters,
|
|
||||||
)
|
|
||||||
if len(self.null_model_params.keys()) == 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
for null_model_info in self.null_model_params.values():
|
start_time = time.time()
|
||||||
logger.info(f"Adding null model: {null_model_info['path']}...")
|
|
||||||
null_model_add = get_net_g(
|
# PyTorch 推論時
|
||||||
model_path=str(null_model_info["path"]),
|
if not self.is_onnx_model:
|
||||||
|
from style_bert_vits2.models.infer import get_net_g
|
||||||
|
|
||||||
|
self.net_g = get_net_g(
|
||||||
|
model_path=str(self.model_path),
|
||||||
version=self.hyper_parameters.version,
|
version=self.hyper_parameters.version,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
)
|
)
|
||||||
# 愚直。もっと上手い方法ありそう
|
logger.info(
|
||||||
params = zip(self.__net_g.dec.parameters(), null_model_add.dec.parameters())
|
f'Model loaded successfully from {self.model_path} to "{self.device}" device ({time.time() - start_time:.2f}s)'
|
||||||
for v in params:
|
|
||||||
v[0].data.add_(v[1].data, alpha=float(null_model_info["weight"]))
|
|
||||||
params = zip(
|
|
||||||
self.__net_g.flow.parameters(), null_model_add.flow.parameters()
|
|
||||||
)
|
)
|
||||||
for v in params:
|
|
||||||
v[0].data.add_(v[1].data, alpha=float(null_model_info["pitch"]))
|
|
||||||
|
|
||||||
params = zip(
|
# ここからはヌルモデルのロード用パラメータが指定されている場合のみ
|
||||||
self.__net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
|
if self.null_model_params is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 推論対象のモデルの重みとヌルモデルの重みをマージ
|
||||||
|
for null_model_info in self.null_model_params.values():
|
||||||
|
logger.info(f"Adding null model: {null_model_info.path}...")
|
||||||
|
null_model_add = get_net_g(
|
||||||
|
model_path=str(null_model_info.path),
|
||||||
|
version=self.hyper_parameters.version,
|
||||||
|
device=self.device,
|
||||||
|
hps=self.hyper_parameters,
|
||||||
|
)
|
||||||
|
# 愚直。もっと上手い方法ありそう
|
||||||
|
params = zip(
|
||||||
|
self.net_g.dec.parameters(), null_model_add.dec.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info.weight))
|
||||||
|
params = zip(
|
||||||
|
self.net_g.flow.parameters(), null_model_add.flow.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info.pitch))
|
||||||
|
|
||||||
|
params = zip(
|
||||||
|
self.net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info.style))
|
||||||
|
# テンポは sdp と dp 二つあるからとりあえずどっちも足す
|
||||||
|
params = zip(
|
||||||
|
self.net_g.sdp.parameters(), null_model_add.sdp.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info.tempo))
|
||||||
|
params = zip(self.net_g.dp.parameters(), null_model_add.dp.parameters())
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info.tempo))
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Null models merged successfully ({time.time() - start_time:.2f}s)"
|
||||||
)
|
)
|
||||||
for v in params:
|
|
||||||
v[0].data.add_(v[1].data, alpha=float(null_model_info["style"]))
|
|
||||||
# テンポはsdpとdp二つあるからとりあえずどっちも足す
|
|
||||||
params = zip(self.__net_g.sdp.parameters(), null_model_add.sdp.parameters())
|
|
||||||
for v in params:
|
|
||||||
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
|
|
||||||
params = zip(self.__net_g.dp.parameters(), null_model_add.dp.parameters())
|
|
||||||
for v in params:
|
|
||||||
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
|
|
||||||
|
|
||||||
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
# ONNX 推論時
|
||||||
|
else:
|
||||||
|
sess_options = onnxruntime.SessionOptions()
|
||||||
|
# ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
|
||||||
|
## DmlExecutionProvider が先頭に指定されているときのみ、DirectML 推論の高速化のためすべての最適化を有効にする
|
||||||
|
assert len(self.onnx_providers) > 0
|
||||||
|
first_provider_name = (
|
||||||
|
self.onnx_providers[0]
|
||||||
|
if type(self.onnx_providers[0]) is str
|
||||||
|
else self.onnx_providers[0][0]
|
||||||
|
)
|
||||||
|
if first_provider_name == "DmlExecutionProvider":
|
||||||
|
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_ALL # fmt: skip
|
||||||
|
else:
|
||||||
|
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
|
||||||
|
# エラー以外のログを出力しない
|
||||||
|
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
|
||||||
|
sess_options.log_severity_level = 3
|
||||||
|
onnxruntime.set_default_logger_severity(3)
|
||||||
|
|
||||||
|
self.onnx_session = onnxruntime.InferenceSession(
|
||||||
|
str(self.model_path),
|
||||||
|
sess_options=sess_options,
|
||||||
|
providers=self.onnx_providers,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
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:
|
||||||
|
"""
|
||||||
|
音声合成モデルをデバイスからアンロードする。
|
||||||
|
PyTorch モデルの場合は CUDA メモリも解放される。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
# PyTorch 推論時
|
||||||
|
if self.net_g is not None:
|
||||||
|
del self.net_g
|
||||||
|
self.net_g = None
|
||||||
|
|
||||||
|
# CUDA キャッシュをクリア
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
# ONNX 推論時
|
||||||
|
if self.onnx_session is not None:
|
||||||
|
del self.onnx_session
|
||||||
|
self.onnx_session = None
|
||||||
|
|
||||||
|
gc.collect()
|
||||||
|
logger.info(f"Model unloaded successfully ({time.time() - start_time:.2f}s)")
|
||||||
|
|
||||||
|
def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
スタイルベクトルを取得する。
|
スタイルベクトルを取得する。
|
||||||
|
|
||||||
@@ -159,12 +262,12 @@ class TTSModel:
|
|||||||
Returns:
|
Returns:
|
||||||
NDArray[Any]: スタイルベクトル
|
NDArray[Any]: スタイルベクトル
|
||||||
"""
|
"""
|
||||||
mean = self.__style_vectors[0]
|
mean = self.style_vectors[0]
|
||||||
style_vec = self.__style_vectors[style_id]
|
style_vec = self.style_vectors[style_id]
|
||||||
style_vec = mean + (style_vec - mean) * weight
|
style_vec = mean + (style_vec - mean) * weight
|
||||||
return style_vec
|
return style_vec
|
||||||
|
|
||||||
def __get_style_vector_from_audio(
|
def get_style_vector_from_audio(
|
||||||
self, audio_path: str, weight: float = 1.0
|
self, audio_path: str, weight: float = 1.0
|
||||||
) -> NDArray[Any]:
|
) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
@@ -177,7 +280,7 @@ class TTSModel:
|
|||||||
NDArray[Any]: スタイルベクトル
|
NDArray[Any]: スタイルベクトル
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if self.__style_vector_inference is None:
|
if self.style_vector_inference is None:
|
||||||
|
|
||||||
# pyannote.audio は scikit-learn などの大量の重量級ライブラリに依存しているため、
|
# pyannote.audio は scikit-learn などの大量の重量級ライブラリに依存しているため、
|
||||||
# TTSModel.infer() に reference_audio_path を指定し音声からスタイルベクトルを推論する場合のみ遅延 import する
|
# TTSModel.infer() に reference_audio_path を指定し音声からスタイルベクトルを推論する場合のみ遅延 import する
|
||||||
@@ -189,21 +292,24 @@ class TTSModel:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# スタイルベクトルを取得するための推論モデルを初期化
|
# スタイルベクトルを取得するための推論モデルを初期化
|
||||||
self.__style_vector_inference = pyannote.audio.Inference(
|
import torch
|
||||||
|
|
||||||
|
self.style_vector_inference = pyannote.audio.Inference(
|
||||||
model=pyannote.audio.Model.from_pretrained(
|
model=pyannote.audio.Model.from_pretrained(
|
||||||
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
||||||
),
|
),
|
||||||
window="whole",
|
window="whole",
|
||||||
)
|
)
|
||||||
self.__style_vector_inference.to(torch.device(self.device))
|
self.style_vector_inference.to(torch.device(self.device))
|
||||||
|
|
||||||
# 音声からスタイルベクトルを推論
|
# 音声からスタイルベクトルを推論
|
||||||
xvec = self.__style_vector_inference(audio_path)
|
xvec = self.style_vector_inference(audio_path)
|
||||||
mean = self.__style_vectors[0]
|
mean = self.style_vectors[0]
|
||||||
xvec = mean + (xvec - mean) * weight
|
xvec = mean + (xvec - mean) * weight
|
||||||
return xvec
|
return xvec
|
||||||
|
|
||||||
def __convert_to_16_bit_wav(self, data: NDArray[Any]) -> NDArray[Any]:
|
@staticmethod
|
||||||
|
def convert_to_16_bit_wav(data: NDArray[Any]) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
音声データを 16-bit int 形式に変換する。
|
音声データを 16-bit int 形式に変換する。
|
||||||
gradio.processing_utils.convert_to_16_bit_wav() を移植したもの。
|
gradio.processing_utils.convert_to_16_bit_wav() を移植したもの。
|
||||||
@@ -214,6 +320,7 @@ class TTSModel:
|
|||||||
Returns:
|
Returns:
|
||||||
NDArray[Any]: 16-bit int 形式の音声データ
|
NDArray[Any]: 16-bit int 形式の音声データ
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Based on: https://docs.scipy.org/doc/scipy/reference/generated/scipy.io.wavfile.write.html
|
# Based on: https://docs.scipy.org/doc/scipy/reference/generated/scipy.io.wavfile.write.html
|
||||||
if data.dtype in [np.float64, np.float32, np.float16]: # type: ignore
|
if data.dtype in [np.float64, np.float32, np.float16]: # type: ignore
|
||||||
data = data / np.abs(data).max()
|
data = data / np.abs(data).max()
|
||||||
@@ -238,6 +345,7 @@ class TTSModel:
|
|||||||
"Audio data cannot be converted automatically from "
|
"Audio data cannot be converted automatically from "
|
||||||
f"{data.dtype} to 16-bit int format."
|
f"{data.dtype} to 16-bit int format."
|
||||||
)
|
)
|
||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
def infer(
|
def infer(
|
||||||
@@ -261,7 +369,7 @@ class TTSModel:
|
|||||||
given_tone: Optional[list[int]] = None,
|
given_tone: Optional[list[int]] = None,
|
||||||
pitch_scale: float = 1.0,
|
pitch_scale: float = 1.0,
|
||||||
intonation_scale: float = 1.0,
|
intonation_scale: float = 1.0,
|
||||||
null_model_params: dict[int, dict[str, Union[str, float]]] = {},
|
null_model_params: Optional[dict[int, NullModelParam]] = None,
|
||||||
force_reload_model: bool = False,
|
force_reload_model: bool = False,
|
||||||
) -> tuple[int, NDArray[Any]]:
|
) -> tuple[int, NDArray[Any]]:
|
||||||
"""
|
"""
|
||||||
@@ -287,7 +395,7 @@ class TTSModel:
|
|||||||
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
||||||
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
null_model_params (dict[int, dict[str, Union[str, float]]], optional): 推論時に使用するヌルモデルの名前、適用割合のdictが入ったdict。
|
null_model_params (Optional[dict[int, NullModelParam]], optional): 推論時に使用するヌルモデルの情報。ONNX 推論では無視される。
|
||||||
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
|
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
|
||||||
Returns:
|
Returns:
|
||||||
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
||||||
@@ -302,25 +410,101 @@ class TTSModel:
|
|||||||
reference_audio_path = None
|
reference_audio_path = None
|
||||||
if assist_text == "" or not use_assist_text:
|
if assist_text == "" or not use_assist_text:
|
||||||
assist_text = None
|
assist_text = None
|
||||||
if null_model_params is not {}:
|
|
||||||
self.null_model_params = null_model_params
|
# スタイルベクトルを取得
|
||||||
else:
|
|
||||||
self.null_model_params = {}
|
|
||||||
if force_reload_model is True:
|
|
||||||
self.__net_g = None
|
|
||||||
if self.__net_g is None:
|
|
||||||
self.load()
|
|
||||||
assert self.__net_g is not None
|
|
||||||
if reference_audio_path is None:
|
if reference_audio_path is None:
|
||||||
style_id = self.style2id[style]
|
style_id = self.style2id[style]
|
||||||
style_vector = self.__get_style_vector(style_id, style_weight)
|
style_vector = self.get_style_vector(style_id, style_weight)
|
||||||
else:
|
else:
|
||||||
style_vector = self.__get_style_vector_from_audio(
|
style_vector = self.get_style_vector_from_audio(
|
||||||
reference_audio_path, style_weight
|
reference_audio_path, style_weight
|
||||||
)
|
)
|
||||||
if not line_split:
|
|
||||||
with torch.no_grad():
|
# PyTorch 推論時
|
||||||
audio = infer(
|
start_time = time.time()
|
||||||
|
if not self.is_onnx_model:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from style_bert_vits2.models.infer import infer
|
||||||
|
|
||||||
|
if null_model_params is not None:
|
||||||
|
self.null_model_params = null_model_params
|
||||||
|
else:
|
||||||
|
self.null_model_params = None
|
||||||
|
|
||||||
|
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
||||||
|
if force_reload_model is True:
|
||||||
|
self.net_g = None
|
||||||
|
|
||||||
|
# モデルがロードされていない場合はロードする
|
||||||
|
if self.net_g is None:
|
||||||
|
self.load()
|
||||||
|
assert self.net_g is not None
|
||||||
|
|
||||||
|
# 通常のテキストから音声を生成
|
||||||
|
if not line_split:
|
||||||
|
with torch.no_grad():
|
||||||
|
audio = infer(
|
||||||
|
text=text,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise,
|
||||||
|
noise_scale_w=noise_w,
|
||||||
|
length_scale=length,
|
||||||
|
sid=speaker_id,
|
||||||
|
language=language,
|
||||||
|
hps=self.hyper_parameters,
|
||||||
|
net_g=self.net_g,
|
||||||
|
device=self.device,
|
||||||
|
assist_text=assist_text,
|
||||||
|
assist_text_weight=assist_text_weight,
|
||||||
|
style_vec=style_vector,
|
||||||
|
given_phone=given_phone,
|
||||||
|
given_tone=given_tone,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 改行ごとに分割して音声を生成
|
||||||
|
else:
|
||||||
|
texts = [t for t in text.split("\n") if t != ""]
|
||||||
|
audios = []
|
||||||
|
with torch.no_grad():
|
||||||
|
for i, t in enumerate(texts):
|
||||||
|
audios.append(
|
||||||
|
infer(
|
||||||
|
text=t,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise,
|
||||||
|
noise_scale_w=noise_w,
|
||||||
|
length_scale=length,
|
||||||
|
sid=speaker_id,
|
||||||
|
language=language,
|
||||||
|
hps=self.hyper_parameters,
|
||||||
|
net_g=self.net_g,
|
||||||
|
device=self.device,
|
||||||
|
assist_text=assist_text,
|
||||||
|
assist_text_weight=assist_text_weight,
|
||||||
|
style_vec=style_vector,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if i != len(texts) - 1:
|
||||||
|
audios.append(np.zeros(int(44100 * split_interval)))
|
||||||
|
audio = np.concatenate(audios)
|
||||||
|
|
||||||
|
# ONNX 推論時
|
||||||
|
else:
|
||||||
|
from style_bert_vits2.models.infer_onnx import infer_onnx
|
||||||
|
|
||||||
|
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
||||||
|
if force_reload_model is True:
|
||||||
|
self.onnx_session = None
|
||||||
|
|
||||||
|
# モデルがロードされていない場合はロードする
|
||||||
|
if self.onnx_session is None:
|
||||||
|
self.load()
|
||||||
|
assert self.onnx_session is not None
|
||||||
|
|
||||||
|
# 通常のテキストから音声を生成
|
||||||
|
if not line_split:
|
||||||
|
audio = infer_onnx(
|
||||||
text=text,
|
text=text,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise,
|
noise_scale=noise,
|
||||||
@@ -329,22 +513,22 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g=self.__net_g,
|
onnx_session=self.onnx_session,
|
||||||
device=self.device,
|
onnx_providers=self.onnx_providers,
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec=style_vector,
|
style_vec=style_vector,
|
||||||
given_phone=given_phone,
|
given_phone=given_phone,
|
||||||
given_tone=given_tone,
|
given_tone=given_tone,
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
texts = text.split("\n")
|
# 改行ごとに分割して音声を生成
|
||||||
texts = [t for t in texts if t != ""]
|
else:
|
||||||
audios = []
|
texts = [t for t in text.split("\n") if t != ""]
|
||||||
with torch.no_grad():
|
audios = []
|
||||||
for i, t in enumerate(texts):
|
for i, t in enumerate(texts):
|
||||||
audios.append(
|
audios.append(
|
||||||
infer(
|
infer_onnx(
|
||||||
text=t,
|
text=t,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise,
|
noise_scale=noise,
|
||||||
@@ -353,8 +537,8 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g=self.__net_g,
|
onnx_session=self.onnx_session,
|
||||||
device=self.device,
|
onnx_providers=self.onnx_providers,
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec=style_vector,
|
style_vec=style_vector,
|
||||||
@@ -363,7 +547,11 @@ class TTSModel:
|
|||||||
if i != len(texts) - 1:
|
if i != len(texts) - 1:
|
||||||
audios.append(np.zeros(int(44100 * split_interval)))
|
audios.append(np.zeros(int(44100 * split_interval)))
|
||||||
audio = np.concatenate(audios)
|
audio = np.concatenate(audios)
|
||||||
logger.info("Audio data generated successfully")
|
|
||||||
|
logger.info(
|
||||||
|
f"Audio data generated successfully ({time.time() - start_time:.2f}s)"
|
||||||
|
)
|
||||||
|
|
||||||
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
||||||
_, audio = adjust_voice(
|
_, audio = adjust_voice(
|
||||||
fs=self.hyper_parameters.data.sampling_rate,
|
fs=self.hyper_parameters.data.sampling_rate,
|
||||||
@@ -371,7 +559,7 @@ class TTSModel:
|
|||||||
pitch_scale=pitch_scale,
|
pitch_scale=pitch_scale,
|
||||||
intonation_scale=intonation_scale,
|
intonation_scale=intonation_scale,
|
||||||
)
|
)
|
||||||
audio = self.__convert_to_16_bit_wav(audio)
|
audio = self.convert_to_16_bit_wav(audio)
|
||||||
return (self.hyper_parameters.data.sampling_rate, audio)
|
return (self.hyper_parameters.data.sampling_rate, audio)
|
||||||
|
|
||||||
|
|
||||||
@@ -384,14 +572,20 @@ class TTSModelInfo(BaseModel):
|
|||||||
|
|
||||||
class TTSModelHolder:
|
class TTSModelHolder:
|
||||||
"""
|
"""
|
||||||
Style-Bert-Vits2 の音声合成モデルを管理するクラス。
|
Style-Bert-VITS2 の音声合成モデルを管理するクラス。
|
||||||
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
|
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, model_root_dir: Path, device: str) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_root_dir: Path,
|
||||||
|
device: str,
|
||||||
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
|
ignore_onnx: bool = False,
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
|
Style-Bert-VITS2 の音声合成モデルを管理するクラスを初期化する。
|
||||||
音声合成モデルは下記のように配置されていることを前提とする (.safetensors のファイル名は自由) 。
|
音声合成モデルは下記のように配置されていることを前提とする (.safetensors / .onnx のファイル名は自由) 。
|
||||||
```
|
```
|
||||||
model_root_dir
|
model_root_dir
|
||||||
├── model-name-1
|
├── model-name-1
|
||||||
@@ -407,11 +601,15 @@ class TTSModelHolder:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス
|
model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス
|
||||||
device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など)
|
device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など)
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
ignore_onnx (bool, optional): ONNX モデルを除外するかどうか. Defaults to False.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
self.root_dir: Path = model_root_dir
|
self.root_dir: Path = model_root_dir
|
||||||
self.device: str = device
|
self.device: str = device
|
||||||
|
self.onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = onnx_providers # fmt: skip
|
||||||
|
self.ignore_onnx: bool = ignore_onnx
|
||||||
self.model_files_dict: dict[str, list[Path]] = {}
|
self.model_files_dict: dict[str, list[Path]] = {}
|
||||||
self.current_model: Optional[TTSModel] = None
|
self.current_model: Optional[TTSModel] = None
|
||||||
self.model_names: list[str] = []
|
self.model_names: list[str] = []
|
||||||
@@ -428,13 +626,21 @@ class TTSModelHolder:
|
|||||||
self.current_model = None
|
self.current_model = None
|
||||||
self.models_info = []
|
self.models_info = []
|
||||||
|
|
||||||
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
model_dirs = sorted([d for d in self.root_dir.iterdir() if d.is_dir()])
|
||||||
for model_dir in model_dirs:
|
for model_dir in model_dirs:
|
||||||
model_files = [
|
if model_dir.name.startswith("."):
|
||||||
f
|
continue
|
||||||
for f in model_dir.iterdir()
|
suffixes = [".pth", ".pt", ".safetensors"]
|
||||||
if f.suffix in [".pth", ".pt", ".safetensors"]
|
if self.ignore_onnx is False:
|
||||||
]
|
suffixes.append(".onnx")
|
||||||
|
model_files = sorted(
|
||||||
|
[
|
||||||
|
f
|
||||||
|
for f in model_dir.iterdir()
|
||||||
|
# 上記 suffixes にマッチするファイルのみを取得し、. から始まるファイルは除外
|
||||||
|
if f.suffix in suffixes and not f.name.startswith(".")
|
||||||
|
]
|
||||||
|
)
|
||||||
if len(model_files) == 0:
|
if len(model_files) == 0:
|
||||||
logger.warning(f"No model files found in {model_dir}, so skip it")
|
logger.warning(f"No model files found in {model_dir}, so skip it")
|
||||||
continue
|
continue
|
||||||
@@ -484,6 +690,7 @@ class TTSModelHolder:
|
|||||||
config_path=self.root_dir / model_name / "config.json",
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
device=self.device,
|
device=self.device,
|
||||||
|
onnx_providers=self.onnx_providers,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self.current_model
|
return self.current_model
|
||||||
@@ -504,29 +711,30 @@ class TTSModelHolder:
|
|||||||
speakers = list(self.current_model.spk2id.keys())
|
speakers = list(self.current_model.spk2id.keys())
|
||||||
styles = list(self.current_model.style2id.keys())
|
styles = list(self.current_model.style2id.keys())
|
||||||
return (
|
return (
|
||||||
gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
|
gr.Dropdown(choices=styles, value=styles[0]),
|
||||||
gr.Button(interactive=True, value="音声合成"),
|
gr.Button(interactive=True, value="音声合成"),
|
||||||
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
gr.Dropdown(choices=speakers, value=speakers[0]),
|
||||||
)
|
)
|
||||||
self.current_model = TTSModel(
|
self.current_model = TTSModel(
|
||||||
model_path=model_path,
|
model_path=model_path,
|
||||||
config_path=self.root_dir / model_name / "config.json",
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
device=self.device,
|
device=self.device,
|
||||||
|
onnx_providers=self.onnx_providers,
|
||||||
)
|
)
|
||||||
speakers = list(self.current_model.spk2id.keys())
|
speakers = list(self.current_model.spk2id.keys())
|
||||||
styles = list(self.current_model.style2id.keys())
|
styles = list(self.current_model.style2id.keys())
|
||||||
return (
|
return (
|
||||||
gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
|
gr.Dropdown(choices=styles, value=styles[0]),
|
||||||
gr.Button(interactive=True, value="音声合成"),
|
gr.Button(interactive=True, value="音声合成"),
|
||||||
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
gr.Dropdown(choices=speakers, value=speakers[0]),
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_model_files_for_gradio(self, model_name: str):
|
def update_model_files_for_gradio(self, model_name: str):
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
|
|
||||||
model_files = [str(f) for f in self.model_files_dict[model_name]]
|
model_files = [str(f) for f in self.model_files_dict[model_name]]
|
||||||
return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
|
return gr.Dropdown(choices=model_files, value=model_files[0])
|
||||||
|
|
||||||
def update_model_names_for_gradio(
|
def update_model_names_for_gradio(
|
||||||
self,
|
self,
|
||||||
@@ -539,7 +747,7 @@ class TTSModelHolder:
|
|||||||
str(f) for f in self.model_files_dict[initial_model_name]
|
str(f) for f in self.model_files_dict[initial_model_name]
|
||||||
]
|
]
|
||||||
return (
|
return (
|
||||||
gr.Dropdown(choices=self.model_names, value=initial_model_name), # type: ignore
|
gr.Dropdown(choices=self.model_names, value=initial_model_name),
|
||||||
gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), # type: ignore
|
gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]),
|
||||||
gr.Button(interactive=False), # For tts_button
|
gr.Button(interactive=False), # For tts_button
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
from typing import Any, Sequence, Union
|
||||||
|
|
||||||
|
|
||||||
|
def torch_device_to_onnx_providers(
|
||||||
|
device: str,
|
||||||
|
) -> Sequence[Union[str, tuple[str, dict[str, Any]]]]:
|
||||||
|
if device.startswith("cuda"):
|
||||||
|
return [
|
||||||
|
# cudnn_conv_algo_search を DEFAULT にすると推論速度が大幅に向上する
|
||||||
|
# ref: https://medium.com/neuml/debug-onnx-gpu-performance-c9290fe07459
|
||||||
|
("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"}),
|
||||||
|
# CUDA が利用できない場合、可能であれば DirectML を利用する
|
||||||
|
("DmlExecutionProvider", {"device_id": 0}),
|
||||||
|
("CPUExecutionProvider", {}),
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
return ["CPUExecutionProvider"]
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
from typing import Any, Literal, Sequence
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from scipy.io import wavfile
|
from scipy.io import wavfile
|
||||||
|
|
||||||
@@ -5,52 +7,130 @@ from style_bert_vits2.constants import BASE_DIR, Languages
|
|||||||
from style_bert_vits2.tts_model import TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModelHolder
|
||||||
|
|
||||||
|
|
||||||
def synthesize(device: str = "cpu"):
|
def synthesize(
|
||||||
|
inference_type: Literal["torch", "onnx"] = "torch",
|
||||||
|
device: str = "cpu",
|
||||||
|
onnx_providers: Sequence[tuple[str, dict[str, Any]]] = [
|
||||||
|
("CPUExecutionProvider", {}),
|
||||||
|
],
|
||||||
|
):
|
||||||
|
|
||||||
# 音声合成モデルが配置されていれば、音声合成を実行
|
# 音声合成モデルが配置されていれば、音声合成を実行
|
||||||
model_holder = TTSModelHolder(BASE_DIR / "model_assets", device)
|
model_holder = TTSModelHolder(BASE_DIR / "model_assets", device, onnx_providers)
|
||||||
if len(model_holder.models_info) > 0:
|
if len(model_holder.models_info) > 0:
|
||||||
|
|
||||||
# jvnv-F2-jp モデルを探す
|
# "koharune-ami" または "amitaro" モデルを探す
|
||||||
for model_info in model_holder.models_info:
|
for model_info in model_holder.models_info:
|
||||||
if model_info.name == "jvnv-F2-jp":
|
if model_info.name == "koharune-ami" or model_info.name == "amitaro":
|
||||||
|
|
||||||
|
# Safetensors 形式または ONNX 形式のモデルファイルに絞り込む
|
||||||
|
if inference_type == "torch":
|
||||||
|
model_files = [
|
||||||
|
f
|
||||||
|
for f in model_info.files
|
||||||
|
if f.endswith(".safetensors") and not f.startswith(".")
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
model_files = [
|
||||||
|
f
|
||||||
|
for f in model_info.files
|
||||||
|
if f.endswith(".onnx") and not f.startswith(".")
|
||||||
|
]
|
||||||
|
if len(model_files) == 0:
|
||||||
|
pytest.skip(
|
||||||
|
f'音声合成モデル "{model_info.name}" のモデルファイルが見つかりませんでした。'
|
||||||
|
)
|
||||||
|
|
||||||
|
# モデルをロード
|
||||||
|
model = model_holder.get_model(model_info.name, model_files[0])
|
||||||
|
model.load()
|
||||||
|
|
||||||
|
# ロードされた InferenceSession の ExecutionProvider が一致するか確認
|
||||||
|
# 一致しない場合、指定された ExecutionProvider で推論できない状態
|
||||||
|
if inference_type == "onnx":
|
||||||
|
assert model.onnx_session is not None
|
||||||
|
assert model.onnx_session.get_providers()[0] == onnx_providers[0][0]
|
||||||
|
|
||||||
# すべてのスタイルに対して音声合成を実行
|
# すべてのスタイルに対して音声合成を実行
|
||||||
for style in model_info.styles:
|
for style in model_info.styles:
|
||||||
|
|
||||||
# 音声合成を実行
|
# 音声合成を実行
|
||||||
model = model_holder.get_model(model_info.name, model_info.files[0])
|
|
||||||
model.load()
|
|
||||||
sample_rate, audio_data = model.infer(
|
sample_rate, audio_data = model.infer(
|
||||||
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
||||||
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
|
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
|
||||||
language=Languages.JP,
|
language=Languages.JP,
|
||||||
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
|
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
|
||||||
speaker_id=0,
|
speaker_id=0,
|
||||||
# 感情表現の強さ (0.0 〜 1.0)
|
# テンポの緩急 (0.0 〜 1.0)
|
||||||
sdp_ratio=0.4,
|
sdp_ratio=0.4,
|
||||||
# スタイル (Neutral, Happy など)
|
# スタイル (Neutral, Happy など)
|
||||||
style=style,
|
style=style,
|
||||||
# スタイルの強さ (0.0 〜 100.0)
|
# スタイルの強さ (0.0 〜 100.0)
|
||||||
style_weight=6.0,
|
style_weight=2.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 音声データを保存
|
# 音声データを保存
|
||||||
(BASE_DIR / "tests/wavs").mkdir(exist_ok=True, parents=True)
|
(BASE_DIR / f"tests/wavs/{model_info.name}").mkdir(
|
||||||
wav_file_path = BASE_DIR / f"tests/wavs/{style}.wav"
|
exist_ok=True, parents=True
|
||||||
|
)
|
||||||
|
wav_file_path = (
|
||||||
|
BASE_DIR / f"tests/wavs/{model_info.name}/{style}.wav"
|
||||||
|
)
|
||||||
with open(wav_file_path, "wb") as f:
|
with open(wav_file_path, "wb") as f:
|
||||||
wavfile.write(f, sample_rate, audio_data)
|
wavfile.write(f, sample_rate, audio_data)
|
||||||
|
|
||||||
# 音声データが保存されたことを確認
|
# 音声データが保存されたことを確認
|
||||||
assert wav_file_path.exists()
|
assert wav_file_path.exists()
|
||||||
# wav_file_path.unlink()
|
|
||||||
|
# モデルをアンロード
|
||||||
|
model.unload()
|
||||||
else:
|
else:
|
||||||
pytest.skip("音声合成モデルが見つかりませんでした。")
|
pytest.skip("音声合成モデルが見つかりませんでした。")
|
||||||
|
|
||||||
|
|
||||||
def test_synthesize_cpu():
|
def test_synthesize_cpu():
|
||||||
synthesize(device="cpu")
|
synthesize(inference_type="torch", device="cpu")
|
||||||
|
|
||||||
|
|
||||||
# Windows環境ではtorchのcudaが簡単に入らないため、テストをスキップ
|
def test_synthesize_cuda():
|
||||||
# def test_synthesize_cuda():
|
synthesize(inference_type="torch", device="cuda")
|
||||||
# synthesize(device="cuda")
|
|
||||||
|
|
||||||
|
def test_synthesize_onnx_cpu():
|
||||||
|
synthesize(
|
||||||
|
inference_type="onnx",
|
||||||
|
onnx_providers=[
|
||||||
|
("CPUExecutionProvider", {}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_synthesize_onnx_cuda():
|
||||||
|
synthesize(
|
||||||
|
inference_type="onnx",
|
||||||
|
onnx_providers=[
|
||||||
|
("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_synthesize_onnx_directml():
|
||||||
|
synthesize(
|
||||||
|
inference_type="onnx",
|
||||||
|
onnx_providers=[
|
||||||
|
# device_id: 0 は、システムにインストールされているプライマリディスプレイ用 GPU に対応する
|
||||||
|
# プライマリディスプレイ用 GPU (GPU 0) よりも性能の高い GPU が接続されている環境では、
|
||||||
|
# 適宜 device_id を変更する必要がある
|
||||||
|
# ref: https://github.com/w-okada/voice-changer/issues/410#issuecomment-1627994911
|
||||||
|
("DmlExecutionProvider", {"device_id": 0}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_synthesize_onnx_coreml():
|
||||||
|
synthesize(
|
||||||
|
inference_type="onnx",
|
||||||
|
onnx_providers=[
|
||||||
|
("CoreMLExecutionProvider", {}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user