@@ -3,23 +3,15 @@
|
|||||||
|
|
||||||
*
|
*
|
||||||
|
|
||||||
|
!/style_bert_vits2/
|
||||||
|
|
||||||
!/bert/deberta-v2-large-japanese-char-wwm/
|
!/bert/deberta-v2-large-japanese-char-wwm/
|
||||||
!/common/
|
!/common/
|
||||||
!/configs/
|
!/configs/
|
||||||
!/dict_data/default.csv
|
!/dict_data/default.csv
|
||||||
!/model_assets/
|
!/model_assets/
|
||||||
!/monotonic_align/
|
|
||||||
!/text/
|
|
||||||
|
|
||||||
!/attentions.py
|
|
||||||
!/commons.py
|
|
||||||
!/config.py
|
!/config.py
|
||||||
!/default_config.yml
|
!/default_config.yml
|
||||||
!/infer.py
|
|
||||||
!/models.py
|
|
||||||
!/models_jp_extra.py
|
|
||||||
!/modules.py
|
|
||||||
!/requirements.txt
|
!/requirements.txt
|
||||||
!/server_editor.py
|
!/server_editor.py
|
||||||
!/transforms.py
|
|
||||||
!/utils.py
|
|
||||||
|
|||||||
9
.gitignore
vendored
9
.gitignore
vendored
@@ -1,8 +1,13 @@
|
|||||||
.vscode/
|
|
||||||
|
|
||||||
__pycache__/
|
__pycache__/
|
||||||
venv/
|
venv/
|
||||||
|
.venv/
|
||||||
|
dist/
|
||||||
|
.coverage
|
||||||
.ipynb_checkpoints/
|
.ipynb_checkpoints/
|
||||||
|
.ruff_cache/
|
||||||
|
|
||||||
|
/Data/
|
||||||
|
/model_assets/
|
||||||
|
|
||||||
/*.yml
|
/*.yml
|
||||||
!/default_config.yml
|
!/default_config.yml
|
||||||
|
|||||||
6
.vscode/extensions.json
vendored
Normal file
6
.vscode/extensions.json
vendored
Normal file
@@ -0,0 +1,6 @@
|
|||||||
|
{
|
||||||
|
"recommendations": [
|
||||||
|
"ms-python.python",
|
||||||
|
"ms-python.vscode-pylance"
|
||||||
|
]
|
||||||
|
}
|
||||||
26
.vscode/settings.json
vendored
Normal file
26
.vscode/settings.json
vendored
Normal file
@@ -0,0 +1,26 @@
|
|||||||
|
{
|
||||||
|
// Pylance の Type Checking を有効化
|
||||||
|
"python.languageServer": "Pylance",
|
||||||
|
"python.analysis.typeCheckingMode": "strict",
|
||||||
|
// Pylance の Type Checking のうち、いくつかのエラー報告を抑制する
|
||||||
|
"python.analysis.diagnosticSeverityOverrides": {
|
||||||
|
"reportConstantRedefinition": "none",
|
||||||
|
"reportGeneralTypeIssues": "warning",
|
||||||
|
"reportMissingParameterType": "warning",
|
||||||
|
"reportMissingTypeStubs": "none",
|
||||||
|
"reportPrivateImportUsage": "none",
|
||||||
|
"reportPrivateUsage": "warning",
|
||||||
|
"reportShadowedImports": "none",
|
||||||
|
"reportUnnecessaryComparison": "none",
|
||||||
|
"reportUnknownArgumentType": "none",
|
||||||
|
"reportUnknownMemberType": "none",
|
||||||
|
"reportUnknownParameterType": "warning",
|
||||||
|
"reportUnknownVariableType": "none",
|
||||||
|
"reportUnusedFunction": "none",
|
||||||
|
"reportUnusedVariable": "information",
|
||||||
|
},
|
||||||
|
"[python]": {
|
||||||
|
"editor.defaultFormatter": "ms-python.black-formatter",
|
||||||
|
"editor.formatOnType": true,
|
||||||
|
},
|
||||||
|
}
|
||||||
11
Dataset.bat
11
Dataset.bat
@@ -1,11 +0,0 @@
|
|||||||
chcp 65001 > NUL
|
|
||||||
@echo off
|
|
||||||
|
|
||||||
pushd %~dp0
|
|
||||||
echo Running webui_dataset.py...
|
|
||||||
venv\Scripts\python webui_dataset.py
|
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
|
||||||
|
|
||||||
popd
|
|
||||||
pause
|
|
||||||
@@ -90,7 +90,7 @@ ENV LD_LIBRARY_PATH=/usr/local/cuda/lib64:$LD_LIBRARY_PATH
|
|||||||
|
|
||||||
# Based on https://pytorch.org/get-started/locally/
|
# Based on https://pytorch.org/get-started/locally/
|
||||||
|
|
||||||
RUN $PIP_INSTALL torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu118
|
RUN $PIP_INSTALL torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118
|
||||||
|
|
||||||
|
|
||||||
RUN $PIP_INSTALL jupyterlab
|
RUN $PIP_INSTALL jupyterlab
|
||||||
|
|||||||
13
Merge.bat
13
Merge.bat
@@ -1,13 +0,0 @@
|
|||||||
chcp 65001 > NUL
|
|
||||||
|
|
||||||
@echo off
|
|
||||||
|
|
||||||
pushd %~dp0
|
|
||||||
|
|
||||||
echo Running webui_merge.py...
|
|
||||||
venv\Scripts\python webui_merge.py
|
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
|
||||||
|
|
||||||
popd
|
|
||||||
pause
|
|
||||||
25
README.md
25
README.md
@@ -4,6 +4,8 @@ Bert-VITS2 with more controllable voice styles.
|
|||||||
|
|
||||||
https://github.com/litagin02/Style-Bert-VITS2/assets/139731664/e853f9a2-db4a-4202-a1dd-56ded3c562a0
|
https://github.com/litagin02/Style-Bert-VITS2/assets/139731664/e853f9a2-db4a-4202-a1dd-56ded3c562a0
|
||||||
|
|
||||||
|
You can install via `pip install style-bert-vits2` (inference only), see [library.ipynb](/library.ipynb) for example usage.
|
||||||
|
|
||||||
- **解説チュートリアル動画** [YouTube](https://youtu.be/aTUSzgDl1iY) [ニコニコ動画](https://www.nicovideo.jp/watch/sm43391524)
|
- **解説チュートリアル動画** [YouTube](https://youtu.be/aTUSzgDl1iY) [ニコニコ動画](https://www.nicovideo.jp/watch/sm43391524)
|
||||||
- [English README](docs/README_en.md)
|
- [English README](docs/README_en.md)
|
||||||
- [](http://colab.research.google.com/github/litagin02/Style-Bert-VITS2/blob/master/colab.ipynb)
|
- [](http://colab.research.google.com/github/litagin02/Style-Bert-VITS2/blob/master/colab.ipynb)
|
||||||
@@ -12,6 +14,7 @@ https://github.com/litagin02/Style-Bert-VITS2/assets/139731664/e853f9a2-db4a-420
|
|||||||
|
|
||||||
- [**リリースページ**](https://github.com/litagin02/Style-Bert-VITS2/releases/)、[更新履歴](/docs/CHANGELOG.md)
|
- [**リリースページ**](https://github.com/litagin02/Style-Bert-VITS2/releases/)、[更新履歴](/docs/CHANGELOG.md)
|
||||||
|
|
||||||
|
- 2024-03-15: ver 2.4.0 (大規模リファクタリングや種々の改良、ライブラリ化)
|
||||||
- 2024-02-26: ver 2.3 (辞書機能とエディター機能)
|
- 2024-02-26: ver 2.3 (辞書機能とエディター機能)
|
||||||
- 2024-02-09: ver 2.2
|
- 2024-02-09: ver 2.2
|
||||||
- 2024-02-07: ver 2.1
|
- 2024-02-07: ver 2.1
|
||||||
@@ -42,6 +45,8 @@ CLIでの使い方は[こちら](/docs/CLI.md)を参照してください。
|
|||||||
|
|
||||||
### インストール
|
### インストール
|
||||||
|
|
||||||
|
Pythonライブラリとしてのpipでのインストールや使用例は[library.ipynb](/library.ipynb)を参照してください。
|
||||||
|
|
||||||
#### GitやPythonに馴染みが無い方
|
#### GitやPythonに馴染みが無い方
|
||||||
|
|
||||||
Windowsを前提としています。
|
Windowsを前提としています。
|
||||||
@@ -62,8 +67,7 @@ git clone https://github.com/litagin02/Style-Bert-VITS2.git
|
|||||||
cd Style-Bert-VITS2
|
cd Style-Bert-VITS2
|
||||||
python -m venv venv
|
python -m venv venv
|
||||||
venv\Scripts\activate
|
venv\Scripts\activate
|
||||||
# PyTorch 2.2.x系は今のところは学習エラーが出るので前のバージョンを使う
|
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118
|
||||||
pip install torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu118
|
|
||||||
pip install -r requirements.txt
|
pip install -r requirements.txt
|
||||||
python initialize.py # 必要なモデルとデフォルトTTSモデルをダウンロード
|
python initialize.py # 必要なモデルとデフォルトTTSモデルをダウンロード
|
||||||
```
|
```
|
||||||
@@ -107,8 +111,8 @@ model_assets
|
|||||||
|
|
||||||
#### データセット作り
|
#### データセット作り
|
||||||
|
|
||||||
- `Dataset.bat`をダブルクリックか`python webui_dataset.py`すると、音声ファイルからデータセットを作るためのWebUIが起動します(音声ファイルを適切な長さにスライスし、その後に文字の書き起こしを自動で行います)。
|
- `App.bat`をダブルクリックか`python app.py`したところの「データセット作成」タブから、音声ファイルを適切な長さにスライスし、その後に文字の書き起こしを自動で行えます。
|
||||||
- 指示に従った後、閉じて下の「学習WebUI」でそのまま学習を行うことができます。
|
- 指示に従った後、下の「学習」タブでそのまま学習を行うことができます。
|
||||||
|
|
||||||
注意: データセットの手動修正やノイズ除去等、細かい修正を行いたい場合は[Aivis](https://github.com/tsukumijima/Aivis)や、そのデータセット部分のWindows対応版 [Aivis Dataset](https://github.com/litagin02/Aivis-Dataset) を使うといいかもしれません。ですがファイル数が多い場合などは、このツールで簡易的に切り出してデータセットを作るだけでも十分という気もしています。
|
注意: データセットの手動修正やノイズ除去等、細かい修正を行いたい場合は[Aivis](https://github.com/tsukumijima/Aivis)や、そのデータセット部分のWindows対応版 [Aivis Dataset](https://github.com/litagin02/Aivis-Dataset) を使うといいかもしれません。ですがファイル数が多い場合などは、このツールで簡易的に切り出してデータセットを作るだけでも十分という気もしています。
|
||||||
|
|
||||||
@@ -116,12 +120,12 @@ model_assets
|
|||||||
|
|
||||||
#### 学習WebUI
|
#### 学習WebUI
|
||||||
|
|
||||||
- `Train.bat`をダブルクリックか`python webui_train.py`するとWebUIが起動するので指示に従ってください。
|
- `App.bat`をダブルクリックか`python app.py`して開くWebUIの「学習」タブから指示に従ってください。
|
||||||
|
|
||||||
### スタイルの生成
|
### スタイルの生成
|
||||||
|
|
||||||
- デフォルトスタイル「Neutral」以外のスタイルを使いたい人向けです。
|
- デフォルトスタイル「Neutral」以外のスタイルを使いたい人向けです。
|
||||||
- `Style.bat`をダブルクリックか`python webui_style_vectors.py`するとWebUIが起動します。
|
- `App.bat`をダブルクリックか`python app.py`して開くWebUIの「スタイル作成」タブから、音声ファイルを使ってスタイルを生成できます。
|
||||||
- 学習とは独立しているので、学習中でもできるし、学習が終わっても何度もやりなおせます(前処理は終わらせている必要があります)。
|
- 学習とは独立しているので、学習中でもできるし、学習が終わっても何度もやりなおせます(前処理は終わらせている必要があります)。
|
||||||
- スタイルについての仕様の詳細は[clustering.ipynb](clustering.ipynb)を参照してください。
|
- スタイルについての仕様の詳細は[clustering.ipynb](clustering.ipynb)を参照してください。
|
||||||
|
|
||||||
@@ -140,7 +144,7 @@ API仕様は起動後に`/docs`にて確認ください。
|
|||||||
### マージ
|
### マージ
|
||||||
|
|
||||||
2つのモデルを、「声質」「声の高さ」「感情表現」「テンポ」の4点で混ぜ合わせて、新しいモデルを作ることが出来ます。
|
2つのモデルを、「声質」「声の高さ」「感情表現」「テンポ」の4点で混ぜ合わせて、新しいモデルを作ることが出来ます。
|
||||||
`Merge.bat`をダブルクリックか`python webui_merge.py`するとWebUIが起動します。
|
`App.bat`をダブルクリックか`python app.py`して開くWebUIの「マージ」タブから、2つのモデルを選択してマージすることができます。
|
||||||
|
|
||||||
### 自然性評価
|
### 自然性評価
|
||||||
|
|
||||||
@@ -165,13 +169,6 @@ python speech_mos.py -m <model_name>
|
|||||||
- safetensors形式のサポート、デフォルトでsafetensorsを使用するように
|
- safetensors形式のサポート、デフォルトでsafetensorsを使用するように
|
||||||
- その他軽微なbugfixやリファクタリング
|
- その他軽微なbugfixやリファクタリング
|
||||||
|
|
||||||
## TODO
|
|
||||||
- [x] デフォルトのJVNVモデルにJP-Extra版のものを追加
|
|
||||||
- [x] LinuxやWSL等、Windowsの通常環境以外でのサポート ← おそらく問題ないとの報告あり
|
|
||||||
- [x] 複数話者学習での音声合成対応(学習は現在でも可能)
|
|
||||||
- [x] `server_fastapi.py`の対応、とくにAPIで使えるようになると嬉しい人が増えるのかもしれない
|
|
||||||
- [x] モデルのマージで声音と感情表現を混ぜる機能の実装
|
|
||||||
- [ ] 英語等多言語対応?
|
|
||||||
|
|
||||||
## References
|
## References
|
||||||
In addition to the original reference (written below), I used the following repositories:
|
In addition to the original reference (written below), I used the following repositories:
|
||||||
|
|||||||
12
Style.bat
12
Style.bat
@@ -1,12 +0,0 @@
|
|||||||
chcp 65001 > NUL
|
|
||||||
|
|
||||||
@echo off
|
|
||||||
|
|
||||||
pushd %~dp0
|
|
||||||
echo Running webui_style_vectors.py...
|
|
||||||
venv\Scripts\python webui_style_vectors.py
|
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
|
||||||
|
|
||||||
popd
|
|
||||||
pause
|
|
||||||
13
Train.bat
13
Train.bat
@@ -1,13 +0,0 @@
|
|||||||
chcp 65001 > NUL
|
|
||||||
|
|
||||||
@echo off
|
|
||||||
|
|
||||||
pushd %~dp0
|
|
||||||
|
|
||||||
echo Running webui_train.py...
|
|
||||||
venv\Scripts\python webui_train.py
|
|
||||||
|
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
|
||||||
|
|
||||||
popd
|
|
||||||
pause
|
|
||||||
537
app.py
537
app.py
@@ -1,502 +1,65 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import datetime
|
|
||||||
import json
|
|
||||||
import os
|
|
||||||
import sys
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
import torch
|
import torch
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
from common.constants import (
|
from style_bert_vits2.constants import GRADIO_THEME, VERSION
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
||||||
DEFAULT_LENGTH,
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
DEFAULT_LINE_SPLIT,
|
from style_bert_vits2.tts_model import TTSModelHolder
|
||||||
DEFAULT_NOISE,
|
from webui.dataset import create_dataset_app
|
||||||
DEFAULT_NOISEW,
|
from webui.inference import create_inference_app
|
||||||
DEFAULT_SDP_RATIO,
|
from webui.merge import create_merge_app
|
||||||
DEFAULT_SPLIT_INTERVAL,
|
from webui.style_vectors import create_style_vectors_app
|
||||||
DEFAULT_STYLE,
|
from webui.train import create_train_app
|
||||||
DEFAULT_STYLE_WEIGHT,
|
|
||||||
GRADIO_THEME,
|
|
||||||
LATEST_VERSION,
|
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
||||||
Languages,
|
pyopenjtalk_worker.initialize_worker()
|
||||||
)
|
|
||||||
from common.log import logger
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
from common.tts_model import ModelHolder
|
update_dict()
|
||||||
from infer import InvalidToneError
|
|
||||||
from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize
|
|
||||||
|
|
||||||
# Get path settings
|
# Get path settings
|
||||||
with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f:
|
with Path("configs/paths.yml").open("r", encoding="utf-8") as f:
|
||||||
path_config: dict[str, str] = yaml.safe_load(f.read())
|
path_config: dict[str, str] = yaml.safe_load(f.read())
|
||||||
# dataset_root = path_config["dataset_root"]
|
# dataset_root = path_config["dataset_root"]
|
||||||
assets_root = path_config["assets_root"]
|
assets_root = path_config["assets_root"]
|
||||||
|
|
||||||
languages = [l.value for l in Languages]
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--device", type=str, default="cuda")
|
||||||
|
parser.add_argument("--host", type=str, default="127.0.0.1")
|
||||||
|
parser.add_argument("--port", type=int, default=None)
|
||||||
|
parser.add_argument("--no_autolaunch", action="store_true")
|
||||||
|
parser.add_argument("--share", action="store_true")
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
device = args.device
|
||||||
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
|
device = "cpu"
|
||||||
|
|
||||||
|
model_holder = TTSModelHolder(Path(assets_root), device)
|
||||||
|
|
||||||
|
with gr.Blocks(theme=GRADIO_THEME) as app:
|
||||||
|
gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})")
|
||||||
|
with gr.Tabs():
|
||||||
|
with gr.Tab("音声合成"):
|
||||||
|
create_inference_app(model_holder=model_holder)
|
||||||
|
with gr.Tab("データセット作成"):
|
||||||
|
create_dataset_app()
|
||||||
|
with gr.Tab("学習"):
|
||||||
|
create_train_app()
|
||||||
|
with gr.Tab("スタイル作成"):
|
||||||
|
create_style_vectors_app()
|
||||||
|
with gr.Tab("マージ"):
|
||||||
|
create_merge_app(model_holder=model_holder)
|
||||||
|
|
||||||
|
|
||||||
def tts_fn(
|
app.launch(
|
||||||
model_name,
|
server_name=args.host,
|
||||||
model_path,
|
server_port=args.port,
|
||||||
text,
|
inbrowser=not args.no_autolaunch,
|
||||||
language,
|
share=args.share,
|
||||||
reference_audio_path,
|
)
|
||||||
sdp_ratio,
|
|
||||||
noise_scale,
|
|
||||||
noise_scale_w,
|
|
||||||
length_scale,
|
|
||||||
line_split,
|
|
||||||
split_interval,
|
|
||||||
assist_text,
|
|
||||||
assist_text_weight,
|
|
||||||
use_assist_text,
|
|
||||||
style,
|
|
||||||
style_weight,
|
|
||||||
kata_tone_json_str,
|
|
||||||
use_tone,
|
|
||||||
speaker,
|
|
||||||
pitch_scale,
|
|
||||||
intonation_scale,
|
|
||||||
):
|
|
||||||
model_holder.load_model_gr(model_name, model_path)
|
|
||||||
|
|
||||||
wrong_tone_message = ""
|
|
||||||
kata_tone: Optional[list[tuple[str, int]]] = None
|
|
||||||
if use_tone and kata_tone_json_str != "":
|
|
||||||
if language != "JP":
|
|
||||||
logger.warning("Only Japanese is supported for tone generation.")
|
|
||||||
wrong_tone_message = "アクセント指定は現在日本語のみ対応しています。"
|
|
||||||
if line_split:
|
|
||||||
logger.warning("Tone generation is not supported for line split.")
|
|
||||||
wrong_tone_message = (
|
|
||||||
"アクセント指定は改行で分けて生成を使わない場合のみ対応しています。"
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
kata_tone = []
|
|
||||||
json_data = json.loads(kata_tone_json_str)
|
|
||||||
# tupleを使うように変換
|
|
||||||
for kana, tone in json_data:
|
|
||||||
assert isinstance(kana, str) and tone in (0, 1), f"{kana}, {tone}"
|
|
||||||
kata_tone.append((kana, tone))
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Error occurred when parsing kana_tone_json: {e}")
|
|
||||||
wrong_tone_message = f"アクセント指定が不正です: {e}"
|
|
||||||
kata_tone = None
|
|
||||||
|
|
||||||
# toneは実際に音声合成に代入される際のみnot Noneになる
|
|
||||||
tone: Optional[list[int]] = None
|
|
||||||
if kata_tone is not None:
|
|
||||||
phone_tone = kata_tone2phone_tone(kata_tone)
|
|
||||||
tone = [t for _, t in phone_tone]
|
|
||||||
|
|
||||||
speaker_id = model_holder.current_model.spk2id[speaker]
|
|
||||||
|
|
||||||
start_time = datetime.datetime.now()
|
|
||||||
|
|
||||||
assert model_holder.current_model is not None
|
|
||||||
|
|
||||||
try:
|
|
||||||
sr, audio = model_holder.current_model.infer(
|
|
||||||
text=text,
|
|
||||||
language=language,
|
|
||||||
reference_audio_path=reference_audio_path,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
noise=noise_scale,
|
|
||||||
noisew=noise_scale_w,
|
|
||||||
length=length_scale,
|
|
||||||
line_split=line_split,
|
|
||||||
split_interval=split_interval,
|
|
||||||
assist_text=assist_text,
|
|
||||||
assist_text_weight=assist_text_weight,
|
|
||||||
use_assist_text=use_assist_text,
|
|
||||||
style=style,
|
|
||||||
style_weight=style_weight,
|
|
||||||
given_tone=tone,
|
|
||||||
sid=speaker_id,
|
|
||||||
pitch_scale=pitch_scale,
|
|
||||||
intonation_scale=intonation_scale,
|
|
||||||
)
|
|
||||||
except InvalidToneError as e:
|
|
||||||
logger.error(f"Tone error: {e}")
|
|
||||||
return f"Error: アクセント指定が不正です:\n{e}", None, kata_tone_json_str
|
|
||||||
except ValueError as e:
|
|
||||||
logger.error(f"Value error: {e}")
|
|
||||||
return f"Error: {e}", None, kata_tone_json_str
|
|
||||||
|
|
||||||
end_time = datetime.datetime.now()
|
|
||||||
duration = (end_time - start_time).total_seconds()
|
|
||||||
|
|
||||||
if tone is None and language == "JP":
|
|
||||||
# アクセント指定に使えるようにアクセント情報を返す
|
|
||||||
norm_text = text_normalize(text)
|
|
||||||
kata_tone = g2kata_tone(norm_text)
|
|
||||||
kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False)
|
|
||||||
elif tone is None:
|
|
||||||
kata_tone_json_str = ""
|
|
||||||
message = f"Success, time: {duration} seconds."
|
|
||||||
if wrong_tone_message != "":
|
|
||||||
message = wrong_tone_message + "\n" + message
|
|
||||||
return message, (sr, audio), kata_tone_json_str
|
|
||||||
|
|
||||||
|
|
||||||
initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?"
|
|
||||||
|
|
||||||
examples = [
|
|
||||||
[initial_text, "JP"],
|
|
||||||
[
|
|
||||||
"""あなたがそんなこと言うなんて、私はとっても嬉しい。
|
|
||||||
あなたがそんなこと言うなんて、私はとっても怒ってる。
|
|
||||||
あなたがそんなこと言うなんて、私はとっても驚いてる。
|
|
||||||
あなたがそんなこと言うなんて、私はとっても辛い。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # ChatGPTに考えてもらった告白セリフ
|
|
||||||
"""私、ずっと前からあなたのことを見てきました。あなたの笑顔、優しさ、強さに、心惹かれていたんです。
|
|
||||||
友達として過ごす中で、あなたのことがだんだんと特別な存在になっていくのがわかりました。
|
|
||||||
えっと、私、あなたのことが好きです!もしよければ、私と付き合ってくれませんか?""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # 夏目漱石『吾輩は猫である』
|
|
||||||
"""吾輩は猫である。名前はまだ無い。
|
|
||||||
どこで生れたかとんと見当がつかぬ。なんでも薄暗いじめじめした所でニャーニャー泣いていた事だけは記憶している。
|
|
||||||
吾輩はここで初めて人間というものを見た。しかもあとで聞くと、それは書生という、人間中で一番獰悪な種族であったそうだ。
|
|
||||||
この書生というのは時々我々を捕まえて煮て食うという話である。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # 梶井基次郎『桜の樹の下には』
|
|
||||||
"""桜の樹の下には屍体が埋まっている!これは信じていいことなんだよ。
|
|
||||||
何故って、桜の花があんなにも見事に咲くなんて信じられないことじゃないか。俺はあの美しさが信じられないので、このにさんにち不安だった。
|
|
||||||
しかしいま、やっとわかるときが来た。桜の樹の下には屍体が埋まっている。これは信じていいことだ。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # ChatGPTと考えた、感情を表すセリフ
|
|
||||||
"""やったー!テストで満点取れた!私とっても嬉しいな!
|
|
||||||
どうして私の意見を無視するの?許せない!ムカつく!あんたなんか死ねばいいのに。
|
|
||||||
あはははっ!この漫画めっちゃ笑える、見てよこれ、ふふふ、あはは。
|
|
||||||
あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しい。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # 上の丁寧語バージョン
|
|
||||||
"""やりました!テストで満点取れましたよ!私とっても嬉しいです!
|
|
||||||
どうして私の意見を無視するんですか?許せません!ムカつきます!あんたなんか死んでください。
|
|
||||||
あはははっ!この漫画めっちゃ笑えます、見てくださいこれ、ふふふ、あはは。
|
|
||||||
あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しいです。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[ # ChatGPTに考えてもらった音声合成の説明文章
|
|
||||||
"""音声合成は、機械学習を活用して、テキストから人の声を再現する技術です。この技術は、言語の構造を解析し、それに基づいて音声を生成します。
|
|
||||||
この分野の最新の研究成果を使うと、より自然で表現豊かな音声の生成が可能である。深層学習の応用により、感情やアクセントを含む声質の微妙な変化も再現することが出来る。""",
|
|
||||||
"JP",
|
|
||||||
],
|
|
||||||
[
|
|
||||||
"Speech synthesis is the artificial production of human speech. A computer system used for this purpose is called a speech synthesizer, and can be implemented in software or hardware products.",
|
|
||||||
"EN",
|
|
||||||
],
|
|
||||||
[
|
|
||||||
"语音合成是人工制造人类语音。用于此目的的计算机系统称为语音合成器,可以通过软件或硬件产品实现。",
|
|
||||||
"ZH",
|
|
||||||
],
|
|
||||||
]
|
|
||||||
|
|
||||||
initial_md = f"""
|
|
||||||
# Style-Bert-VITS2 ver {LATEST_VERSION} 音声合成
|
|
||||||
|
|
||||||
- Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py`で起動できます。
|
|
||||||
|
|
||||||
- 初期からある[jvnvのモデル](https://huggingface.co/litagin/style_bert_vits2_jvnv)は、[JVNVコーパス(言語音声と非言語音声を持つ日本語感情音声コーパス)](https://sites.google.com/site/shinnosuketakamichi/research-topics/jvnv_corpus)で学習されたモデルです。ライセンスは[CC BY-SA 4.0](https://creativecommons.org/licenses/by-sa/4.0/deed.ja)です。
|
|
||||||
"""
|
|
||||||
|
|
||||||
how_to_md = """
|
|
||||||
下のように`model_assets`ディレクトリの中にモデルファイルたちを置いてください。
|
|
||||||
```
|
|
||||||
model_assets
|
|
||||||
├── your_model
|
|
||||||
│ ├── config.json
|
|
||||||
│ ├── your_model_file1.safetensors
|
|
||||||
│ ├── your_model_file2.safetensors
|
|
||||||
│ ├── ...
|
|
||||||
│ └── style_vectors.npy
|
|
||||||
└── another_model
|
|
||||||
├── ...
|
|
||||||
```
|
|
||||||
各モデルにはファイルたちが必要です:
|
|
||||||
- `config.json`:学習時の設定ファイル
|
|
||||||
- `*.safetensors`:学習済みモデルファイル(1つ以上が必要、複数可)
|
|
||||||
- `style_vectors.npy`:スタイルベクトルファイル
|
|
||||||
|
|
||||||
上2つは`Train.bat`による学習で自動的に正しい位置に保存されます。`style_vectors.npy`は`Style.bat`を実行して指示に従って生成してください。
|
|
||||||
"""
|
|
||||||
|
|
||||||
style_md = f"""
|
|
||||||
- プリセットまたは音声ファイルから読み上げの声音・感情・スタイルのようなものを制御できます。
|
|
||||||
- デフォルトの{DEFAULT_STYLE}でも、十分に読み上げる文に応じた感情で感情豊かに読み上げられます。このスタイル制御は、それを重み付きで上書きするような感じです。
|
|
||||||
- 強さを大きくしすぎると発音が変になったり声にならなかったりと崩壊することがあります。
|
|
||||||
- どのくらいに強さがいいかはモデルやスタイルによって異なるようです。
|
|
||||||
- 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
def make_interactive():
|
|
||||||
return gr.update(interactive=True, value="音声合成")
|
|
||||||
|
|
||||||
|
|
||||||
def make_non_interactive():
|
|
||||||
return gr.update(interactive=False, value="音声合成(モデルをロードしてください)")
|
|
||||||
|
|
||||||
|
|
||||||
def gr_util(item):
|
|
||||||
if item == "プリセットから選ぶ":
|
|
||||||
return (gr.update(visible=True), gr.Audio(visible=False, value=None))
|
|
||||||
else:
|
|
||||||
return (gr.update(visible=False), gr.update(visible=True))
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
parser = argparse.ArgumentParser()
|
|
||||||
parser.add_argument("--cpu", action="store_true", help="Use CPU instead of GPU")
|
|
||||||
parser.add_argument(
|
|
||||||
"--dir", "-d", type=str, help="Model directory", default=assets_root
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--share", action="store_true", help="Share this app publicly", default=False
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--server-name",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="Server name for Gradio app",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--no-autolaunch",
|
|
||||||
action="store_true",
|
|
||||||
default=False,
|
|
||||||
help="Do not launch app automatically",
|
|
||||||
)
|
|
||||||
args = parser.parse_args()
|
|
||||||
model_dir = Path(args.dir)
|
|
||||||
|
|
||||||
if args.cpu:
|
|
||||||
device = "cpu"
|
|
||||||
else:
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
||||||
|
|
||||||
model_holder = ModelHolder(model_dir, device)
|
|
||||||
|
|
||||||
model_names = model_holder.model_names
|
|
||||||
if len(model_names) == 0:
|
|
||||||
logger.error(
|
|
||||||
f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。"
|
|
||||||
)
|
|
||||||
sys.exit(1)
|
|
||||||
initial_id = 0
|
|
||||||
initial_pth_files = model_holder.model_files_dict[model_names[initial_id]]
|
|
||||||
|
|
||||||
with gr.Blocks(theme=GRADIO_THEME) as app:
|
|
||||||
gr.Markdown(initial_md)
|
|
||||||
with gr.Accordion(label="使い方", open=False):
|
|
||||||
gr.Markdown(how_to_md)
|
|
||||||
with gr.Row():
|
|
||||||
with gr.Column():
|
|
||||||
with gr.Row():
|
|
||||||
with gr.Column(scale=3):
|
|
||||||
model_name = gr.Dropdown(
|
|
||||||
label="モデル一覧",
|
|
||||||
choices=model_names,
|
|
||||||
value=model_names[initial_id],
|
|
||||||
)
|
|
||||||
model_path = gr.Dropdown(
|
|
||||||
label="モデルファイル",
|
|
||||||
choices=initial_pth_files,
|
|
||||||
value=initial_pth_files[0],
|
|
||||||
)
|
|
||||||
refresh_button = gr.Button("更新", scale=1, visible=True)
|
|
||||||
load_button = gr.Button("ロード", scale=1, variant="primary")
|
|
||||||
text_input = gr.TextArea(label="テキスト", value=initial_text)
|
|
||||||
pitch_scale = gr.Slider(
|
|
||||||
minimum=0.8,
|
|
||||||
maximum=1.5,
|
|
||||||
value=1,
|
|
||||||
step=0.05,
|
|
||||||
label="音程(1以外では音質劣化)",
|
|
||||||
visible=False, # pyworldが必要
|
|
||||||
)
|
|
||||||
intonation_scale = gr.Slider(
|
|
||||||
minimum=0,
|
|
||||||
maximum=2,
|
|
||||||
value=1,
|
|
||||||
step=0.1,
|
|
||||||
label="抑揚(1以外では音質劣化)",
|
|
||||||
visible=False, # pyworldが必要
|
|
||||||
)
|
|
||||||
|
|
||||||
line_split = gr.Checkbox(
|
|
||||||
label="改行で分けて生成(分けたほうが感情が乗ります)",
|
|
||||||
value=DEFAULT_LINE_SPLIT,
|
|
||||||
)
|
|
||||||
split_interval = gr.Slider(
|
|
||||||
minimum=0.0,
|
|
||||||
maximum=2,
|
|
||||||
value=DEFAULT_SPLIT_INTERVAL,
|
|
||||||
step=0.1,
|
|
||||||
label="改行ごとに挟む無音の長さ(秒)",
|
|
||||||
)
|
|
||||||
line_split.change(
|
|
||||||
lambda x: (gr.Slider(visible=x)),
|
|
||||||
inputs=[line_split],
|
|
||||||
outputs=[split_interval],
|
|
||||||
)
|
|
||||||
tone = gr.Textbox(
|
|
||||||
label="アクセント調整(数値は 0=低 か1=高 のみ)",
|
|
||||||
info="改行で分けない場合のみ使えます。万能ではありません。",
|
|
||||||
)
|
|
||||||
use_tone = gr.Checkbox(label="アクセント調整を使う", value=False)
|
|
||||||
use_tone.change(
|
|
||||||
lambda x: (gr.Checkbox(value=False) if x else gr.Checkbox()),
|
|
||||||
inputs=[use_tone],
|
|
||||||
outputs=[line_split],
|
|
||||||
)
|
|
||||||
language = gr.Dropdown(choices=languages, value="JP", label="Language")
|
|
||||||
speaker = gr.Dropdown(label="話者")
|
|
||||||
with gr.Accordion(label="詳細設定", open=False):
|
|
||||||
sdp_ratio = gr.Slider(
|
|
||||||
minimum=0,
|
|
||||||
maximum=1,
|
|
||||||
value=DEFAULT_SDP_RATIO,
|
|
||||||
step=0.1,
|
|
||||||
label="SDP Ratio",
|
|
||||||
)
|
|
||||||
noise_scale = gr.Slider(
|
|
||||||
minimum=0.1,
|
|
||||||
maximum=2,
|
|
||||||
value=DEFAULT_NOISE,
|
|
||||||
step=0.1,
|
|
||||||
label="Noise",
|
|
||||||
)
|
|
||||||
noise_scale_w = gr.Slider(
|
|
||||||
minimum=0.1,
|
|
||||||
maximum=2,
|
|
||||||
value=DEFAULT_NOISEW,
|
|
||||||
step=0.1,
|
|
||||||
label="Noise_W",
|
|
||||||
)
|
|
||||||
length_scale = gr.Slider(
|
|
||||||
minimum=0.1,
|
|
||||||
maximum=2,
|
|
||||||
value=DEFAULT_LENGTH,
|
|
||||||
step=0.1,
|
|
||||||
label="Length",
|
|
||||||
)
|
|
||||||
use_assist_text = gr.Checkbox(
|
|
||||||
label="Assist textを使う", value=False
|
|
||||||
)
|
|
||||||
assist_text = gr.Textbox(
|
|
||||||
label="Assist text",
|
|
||||||
placeholder="どうして私の意見を無視するの?許せない、ムカつく!死ねばいいのに。",
|
|
||||||
info="このテキストの読み上げと似た声音・感情になりやすくなります。ただ抑揚やテンポ等が犠牲になる傾向があります。",
|
|
||||||
visible=False,
|
|
||||||
)
|
|
||||||
assist_text_weight = gr.Slider(
|
|
||||||
minimum=0,
|
|
||||||
maximum=1,
|
|
||||||
value=DEFAULT_ASSIST_TEXT_WEIGHT,
|
|
||||||
step=0.1,
|
|
||||||
label="Assist textの強さ",
|
|
||||||
visible=False,
|
|
||||||
)
|
|
||||||
use_assist_text.change(
|
|
||||||
lambda x: (gr.Textbox(visible=x), gr.Slider(visible=x)),
|
|
||||||
inputs=[use_assist_text],
|
|
||||||
outputs=[assist_text, assist_text_weight],
|
|
||||||
)
|
|
||||||
with gr.Column():
|
|
||||||
with gr.Accordion("スタイルについて詳細", open=False):
|
|
||||||
gr.Markdown(style_md)
|
|
||||||
style_mode = gr.Radio(
|
|
||||||
["プリセットから選ぶ", "音声ファイルを入力"],
|
|
||||||
label="スタイルの指定方法",
|
|
||||||
value="プリセットから選ぶ",
|
|
||||||
)
|
|
||||||
style = gr.Dropdown(
|
|
||||||
label=f"スタイル({DEFAULT_STYLE}が平均スタイル)",
|
|
||||||
choices=["モデルをロードしてください"],
|
|
||||||
value="モデルをロードしてください",
|
|
||||||
)
|
|
||||||
style_weight = gr.Slider(
|
|
||||||
minimum=0,
|
|
||||||
maximum=50,
|
|
||||||
value=DEFAULT_STYLE_WEIGHT,
|
|
||||||
step=0.1,
|
|
||||||
label="スタイルの強さ",
|
|
||||||
)
|
|
||||||
ref_audio_path = gr.Audio(
|
|
||||||
label="参照音声", type="filepath", visible=False
|
|
||||||
)
|
|
||||||
tts_button = gr.Button(
|
|
||||||
"音声合成(モデルをロードしてください)",
|
|
||||||
variant="primary",
|
|
||||||
interactive=False,
|
|
||||||
)
|
|
||||||
text_output = gr.Textbox(label="情報")
|
|
||||||
audio_output = gr.Audio(label="結果")
|
|
||||||
with gr.Accordion("テキスト例", open=False):
|
|
||||||
gr.Examples(examples, inputs=[text_input, language])
|
|
||||||
|
|
||||||
tts_button.click(
|
|
||||||
tts_fn,
|
|
||||||
inputs=[
|
|
||||||
model_name,
|
|
||||||
model_path,
|
|
||||||
text_input,
|
|
||||||
language,
|
|
||||||
ref_audio_path,
|
|
||||||
sdp_ratio,
|
|
||||||
noise_scale,
|
|
||||||
noise_scale_w,
|
|
||||||
length_scale,
|
|
||||||
line_split,
|
|
||||||
split_interval,
|
|
||||||
assist_text,
|
|
||||||
assist_text_weight,
|
|
||||||
use_assist_text,
|
|
||||||
style,
|
|
||||||
style_weight,
|
|
||||||
tone,
|
|
||||||
use_tone,
|
|
||||||
speaker,
|
|
||||||
pitch_scale,
|
|
||||||
intonation_scale,
|
|
||||||
],
|
|
||||||
outputs=[text_output, audio_output, tone],
|
|
||||||
)
|
|
||||||
|
|
||||||
model_name.change(
|
|
||||||
model_holder.update_model_files_gr,
|
|
||||||
inputs=[model_name],
|
|
||||||
outputs=[model_path],
|
|
||||||
)
|
|
||||||
|
|
||||||
model_path.change(make_non_interactive, outputs=[tts_button])
|
|
||||||
|
|
||||||
refresh_button.click(
|
|
||||||
model_holder.update_model_names_gr,
|
|
||||||
outputs=[model_name, model_path, tts_button],
|
|
||||||
)
|
|
||||||
|
|
||||||
load_button.click(
|
|
||||||
model_holder.load_model_gr,
|
|
||||||
inputs=[model_name, model_path],
|
|
||||||
outputs=[style, tts_button, speaker],
|
|
||||||
)
|
|
||||||
|
|
||||||
style_mode.change(
|
|
||||||
gr_util,
|
|
||||||
inputs=[style_mode],
|
|
||||||
outputs=[style, ref_audio_path],
|
|
||||||
)
|
|
||||||
|
|
||||||
app.launch(
|
|
||||||
inbrowser=not args.no_autolaunch, share=args.share, server_name=args.server_name
|
|
||||||
)
|
|
||||||
|
|||||||
50
bert_gen.py
50
bert_gen.py
@@ -5,15 +5,29 @@ import torch
|
|||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
import commons
|
|
||||||
import utils
|
|
||||||
from common.log import logger
|
|
||||||
from common.stdout_wrapper import SAFE_STDOUT
|
|
||||||
from config import config
|
from config import config
|
||||||
from text import cleaned_text_to_sequence, get_bert
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.models import commons
|
||||||
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
|
from style_bert_vits2.nlp import (
|
||||||
|
bert_models,
|
||||||
|
cleaned_text_to_sequence,
|
||||||
|
extract_bert_feature,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
||||||
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
|
|
||||||
def process_line(x):
|
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
||||||
|
pyopenjtalk_worker.initialize_worker()
|
||||||
|
|
||||||
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
|
update_dict()
|
||||||
|
|
||||||
|
|
||||||
|
def process_line(x: tuple[str, bool]):
|
||||||
line, add_blank = x
|
line, add_blank = x
|
||||||
device = config.bert_gen_config.device
|
device = config.bert_gen_config.device
|
||||||
if config.bert_gen_config.use_multi_device:
|
if config.bert_gen_config.use_multi_device:
|
||||||
@@ -21,15 +35,17 @@ def process_line(x):
|
|||||||
rank = rank[0] if len(rank) > 0 else 0
|
rank = rank[0] if len(rank) > 0 else 0
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
gpu_id = rank % torch.cuda.device_count()
|
gpu_id = rank % torch.cuda.device_count()
|
||||||
device = torch.device(f"cuda:{gpu_id}")
|
device = f"cuda:{gpu_id}"
|
||||||
else:
|
else:
|
||||||
device = torch.device("cpu")
|
device = "cpu"
|
||||||
wav_path, _, language_str, text, phones, tone, word2ph = line.strip().split("|")
|
wav_path, _, language_str, text, phones, tone, word2ph = line.strip().split("|")
|
||||||
phone = phones.split(" ")
|
phone = phones.split(" ")
|
||||||
tone = [int(i) for i in tone.split(" ")]
|
tone = [int(i) for i in tone.split(" ")]
|
||||||
word2ph = [int(i) for i in word2ph.split(" ")]
|
word2ph = [int(i) for i in word2ph.split(" ")]
|
||||||
word2ph = [i for i in word2ph]
|
word2ph = [i for i in word2ph]
|
||||||
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
phone, tone, language = cleaned_text_to_sequence(
|
||||||
|
phone, tone, Languages[language_str]
|
||||||
|
)
|
||||||
|
|
||||||
if add_blank:
|
if add_blank:
|
||||||
phone = commons.intersperse(phone, 0)
|
phone = commons.intersperse(phone, 0)
|
||||||
@@ -45,7 +61,7 @@ def process_line(x):
|
|||||||
bert = torch.load(bert_path)
|
bert = torch.load(bert_path)
|
||||||
assert bert.shape[-1] == len(phone)
|
assert bert.shape[-1] == len(phone)
|
||||||
except Exception:
|
except Exception:
|
||||||
bert = get_bert(text, word2ph, language_str, device)
|
bert = extract_bert_feature(text, word2ph, language_str, device)
|
||||||
assert bert.shape[-1] == len(phone)
|
assert bert.shape[-1] == len(phone)
|
||||||
torch.save(bert, bert_path)
|
torch.save(bert, bert_path)
|
||||||
|
|
||||||
@@ -57,22 +73,20 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"-c", "--config", type=str, default=config.bert_gen_config.config_path
|
"-c", "--config", type=str, default=config.bert_gen_config.config_path
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
|
||||||
"--num_processes", type=int, default=config.bert_gen_config.num_processes
|
|
||||||
)
|
|
||||||
args, _ = parser.parse_known_args()
|
args, _ = parser.parse_known_args()
|
||||||
config_path = args.config
|
config_path = args.config
|
||||||
hps = utils.get_hparams_from_file(config_path)
|
hps = HyperParameters.load_from_json(config_path)
|
||||||
lines = []
|
lines: list[str] = []
|
||||||
with open(hps.data.training_files, encoding="utf-8") as f:
|
with open(hps.data.training_files, "r", encoding="utf-8") as f:
|
||||||
lines.extend(f.readlines())
|
lines.extend(f.readlines())
|
||||||
|
|
||||||
with open(hps.data.validation_files, encoding="utf-8") as f:
|
with open(hps.data.validation_files, "r", encoding="utf-8") as f:
|
||||||
lines.extend(f.readlines())
|
lines.extend(f.readlines())
|
||||||
add_blank = [hps.data.add_blank] * len(lines)
|
add_blank = [hps.data.add_blank] * len(lines)
|
||||||
|
|
||||||
if len(lines) != 0:
|
if len(lines) != 0:
|
||||||
num_processes = args.num_processes
|
# pyopenjtalkの別ワーカー化により、並列処理でエラーがでる模様なので、一旦シングルスレッド強制にする
|
||||||
|
num_processes = 1
|
||||||
with ThreadPoolExecutor(max_workers=num_processes) as executor:
|
with ThreadPoolExecutor(max_workers=num_processes) as executor:
|
||||||
_ = list(
|
_ = list(
|
||||||
tqdm(
|
tqdm(
|
||||||
|
|||||||
53
colab.ipynb
53
colab.ipynb
@@ -4,7 +4,7 @@
|
|||||||
"cell_type": "markdown",
|
"cell_type": "markdown",
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"source": [
|
"source": [
|
||||||
"# Style-Bert-VITS2 (ver 2.3.1) のGoogle Colabでの学習\n",
|
"# Style-Bert-VITS2 (ver 2.4.0) のGoogle Colabでの学習\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
||||||
"\n",
|
"\n",
|
||||||
@@ -26,7 +26,9 @@
|
|||||||
"source": [
|
"source": [
|
||||||
"## 0. 環境構築\n",
|
"## 0. 環境構築\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Style-Bert-VITS2の環境をcolab上に構築します。グラボモードが有効になっていることを確認し、以下のセルを順に実行してください。"
|
"Style-Bert-VITS2の環境をcolab上に構築します。グラボモードが有効になっていることを確認し、以下のセルを順に実行してください。\n",
|
||||||
|
"\n",
|
||||||
|
"**最近のcolabのアップデートにより、エラーダイアログ「WARNING: The following packages were previously imported in this runtime: [pydevd_plugins]」が出るが、「キャンセル」を選択して続行してください。**"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -35,13 +37,12 @@
|
|||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"#@title このセルを実行して環境構築してください。\n",
|
"# このセルを実行して環境構築してください。\n",
|
||||||
"#@markdown 最後に赤文字でエラーや警告が出ても何故かうまくいくみたいです。\n",
|
"# エラーダイアログ「WARNING: The following packages were previously imported in this runtime: [pydevd_plugins]」が出るが「キャンセル」を選択して続行してください。\n",
|
||||||
"\n",
|
"\n",
|
||||||
"!git clone https://github.com/litagin02/Style-Bert-VITS2.git\n",
|
"!git clone https://github.com/litagin02/Style-Bert-VITS2.git\n",
|
||||||
"%cd Style-Bert-VITS2/\n",
|
"%cd Style-Bert-VITS2/\n",
|
||||||
"!pip install -r requirements.txt\n",
|
"!pip install -r requirements.txt\n",
|
||||||
"!apt install libcublas11\n",
|
|
||||||
"!python initialize.py --skip_jvnv"
|
"!python initialize.py --skip_jvnv"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
@@ -119,7 +120,7 @@
|
|||||||
"initial_prompt = \"こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!\"\n",
|
"initial_prompt = \"こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!\"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"!python slice.py -i {input_dir} --model_name {model_name}\n",
|
"!python slice.py -i {input_dir} --model_name {model_name}\n",
|
||||||
"!python transcribe.py --model_name {model_name} --compute_type float16 --initial_prompt {initial_prompt}"
|
"!python transcribe.py --model_name {model_name} --initial_prompt {initial_prompt} --use_hf_whisper"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -219,7 +220,7 @@
|
|||||||
"batch_size = 4\n",
|
"batch_size = 4\n",
|
||||||
"\n",
|
"\n",
|
||||||
"# 学習のエポック数(データセットを合計何周するか)。\n",
|
"# 学習のエポック数(データセットを合計何周するか)。\n",
|
||||||
"# 100ぐらいで十分かもしれませんが、もっと多くやると質が上がるのかもしれません。\n",
|
"# 100で多すぎるほどかもしれませんが、もっと多くやると質が上がるのかもしれません。\n",
|
||||||
"epochs = 100\n",
|
"epochs = 100\n",
|
||||||
"\n",
|
"\n",
|
||||||
"# 保存頻度。何ステップごとにモデルを保存するか。分からなければデフォルトのままで。\n",
|
"# 保存頻度。何ステップごとにモデルを保存するか。分からなければデフォルトのままで。\n",
|
||||||
@@ -255,7 +256,7 @@
|
|||||||
},
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"from webui_train import preprocess_all\n",
|
"from webui.train import preprocess_all\n",
|
||||||
"\n",
|
"\n",
|
||||||
"preprocess_all(\n",
|
"preprocess_all(\n",
|
||||||
" model_name=model_name,\n",
|
" model_name=model_name,\n",
|
||||||
@@ -307,7 +308,7 @@
|
|||||||
"\n",
|
"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"import yaml\n",
|
"import yaml\n",
|
||||||
"from webui_train import get_path\n",
|
"from webui.train import get_path\n",
|
||||||
"\n",
|
"\n",
|
||||||
"dataset_path, _, _, _, config_path = get_path(model_name)\n",
|
"dataset_path, _, _, _, config_path = get_path(model_name)\n",
|
||||||
"\n",
|
"\n",
|
||||||
@@ -350,41 +351,9 @@
|
|||||||
},
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"#@title 学習結果を試すならここから\n",
|
"# 学習結果を試す・マージ・スタイル分けはこちらから\n",
|
||||||
"!python app.py --share --dir {assets_root}"
|
"!python app.py --share --dir {assets_root}"
|
||||||
]
|
]
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {},
|
|
||||||
"source": [
|
|
||||||
"## 5. スタイル分け"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"!python webui_style_vectors.py --share"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {},
|
|
||||||
"source": [
|
|
||||||
"## 6. マージ"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"!python webui_merge.py --share"
|
|
||||||
]
|
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"metadata": {
|
"metadata": {
|
||||||
|
|||||||
@@ -1,28 +0,0 @@
|
|||||||
import enum
|
|
||||||
|
|
||||||
# Built-in theme: "default", "base", "monochrome", "soft", "glass"
|
|
||||||
# See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
|
||||||
GRADIO_THEME: str = "NoCrypt/miku"
|
|
||||||
|
|
||||||
LATEST_VERSION: str = "2.3.1"
|
|
||||||
|
|
||||||
USER_DICT_DIR = "dict_data"
|
|
||||||
|
|
||||||
DEFAULT_STYLE: str = "Neutral"
|
|
||||||
DEFAULT_STYLE_WEIGHT: float = 5.0
|
|
||||||
|
|
||||||
|
|
||||||
class Languages(str, enum.Enum):
|
|
||||||
JP = "JP"
|
|
||||||
EN = "EN"
|
|
||||||
ZH = "ZH"
|
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_SDP_RATIO: float = 0.2
|
|
||||||
DEFAULT_NOISE: float = 0.6
|
|
||||||
DEFAULT_NOISEW: float = 0.8
|
|
||||||
DEFAULT_LENGTH: float = 1.0
|
|
||||||
DEFAULT_LINE_SPLIT: bool = True
|
|
||||||
DEFAULT_SPLIT_INTERVAL: float = 0.5
|
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT: float = 0.7
|
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT: float = 1.0
|
|
||||||
@@ -1,17 +0,0 @@
|
|||||||
"""
|
|
||||||
logger封装
|
|
||||||
"""
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from .stdout_wrapper import SAFE_STDOUT
|
|
||||||
|
|
||||||
# 移除所有默认的处理器
|
|
||||||
logger.remove()
|
|
||||||
|
|
||||||
# 自定义格式并添加到标准输出
|
|
||||||
log_format = (
|
|
||||||
"<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}"
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True)
|
|
||||||
@@ -1,33 +0,0 @@
|
|||||||
import subprocess
|
|
||||||
import sys
|
|
||||||
|
|
||||||
from .log import logger
|
|
||||||
from .stdout_wrapper import SAFE_STDOUT
|
|
||||||
|
|
||||||
python = sys.executable
|
|
||||||
|
|
||||||
|
|
||||||
def run_script_with_log(cmd: list[str], ignore_warning=False) -> tuple[bool, str]:
|
|
||||||
logger.info(f"Running: {' '.join(cmd)}")
|
|
||||||
result = subprocess.run(
|
|
||||||
[python] + cmd,
|
|
||||||
stdout=SAFE_STDOUT, # type: ignore
|
|
||||||
stderr=subprocess.PIPE,
|
|
||||||
text=True,
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
if result.returncode != 0:
|
|
||||||
logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}")
|
|
||||||
return False, result.stderr
|
|
||||||
elif result.stderr and not ignore_warning:
|
|
||||||
logger.warning(f"Warning: {' '.join(cmd)}\n{result.stderr}")
|
|
||||||
return True, result.stderr
|
|
||||||
logger.success(f"Success: {' '.join(cmd)}")
|
|
||||||
return True, ""
|
|
||||||
|
|
||||||
|
|
||||||
def second_elem_of(original_function):
|
|
||||||
def inner_function(*args, **kwargs):
|
|
||||||
return original_function(*args, **kwargs)[1]
|
|
||||||
|
|
||||||
return inner_function
|
|
||||||
@@ -1,332 +0,0 @@
|
|||||||
import os
|
|
||||||
import warnings
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Optional, Union
|
|
||||||
|
|
||||||
import gradio as gr
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from gradio.processing_utils import convert_to_16_bit_wav
|
|
||||||
|
|
||||||
import utils
|
|
||||||
from infer import get_net_g, infer
|
|
||||||
from models import SynthesizerTrn
|
|
||||||
from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
|
|
||||||
|
|
||||||
from .constants import (
|
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
|
||||||
DEFAULT_LENGTH,
|
|
||||||
DEFAULT_LINE_SPLIT,
|
|
||||||
DEFAULT_NOISE,
|
|
||||||
DEFAULT_NOISEW,
|
|
||||||
DEFAULT_SDP_RATIO,
|
|
||||||
DEFAULT_SPLIT_INTERVAL,
|
|
||||||
DEFAULT_STYLE,
|
|
||||||
DEFAULT_STYLE_WEIGHT,
|
|
||||||
)
|
|
||||||
from .log import logger
|
|
||||||
|
|
||||||
|
|
||||||
def adjust_voice(fs, wave, pitch_scale, intonation_scale):
|
|
||||||
if pitch_scale == 1.0 and intonation_scale == 1.0:
|
|
||||||
# 初期値の場合は、音質劣化を避けるためにそのまま返す
|
|
||||||
return fs, wave
|
|
||||||
|
|
||||||
try:
|
|
||||||
import pyworld
|
|
||||||
except ImportError:
|
|
||||||
raise ImportError(
|
|
||||||
"pyworld is not installed. Please install it by `pip install pyworld`"
|
|
||||||
)
|
|
||||||
|
|
||||||
# pyworldでf0を加工して合成
|
|
||||||
# pyworldよりもよいのがあるかもしれないが……
|
|
||||||
|
|
||||||
wave = wave.astype(np.double)
|
|
||||||
f0, t = pyworld.harvest(wave, fs)
|
|
||||||
# 質が高そうだしとりあえずharvestにしておく
|
|
||||||
|
|
||||||
sp = pyworld.cheaptrick(wave, f0, t, fs)
|
|
||||||
ap = pyworld.d4c(wave, f0, t, fs)
|
|
||||||
|
|
||||||
non_zero_f0 = [f for f in f0 if f != 0]
|
|
||||||
f0_mean = sum(non_zero_f0) / len(non_zero_f0)
|
|
||||||
|
|
||||||
for i, f in enumerate(f0):
|
|
||||||
if f == 0:
|
|
||||||
continue
|
|
||||||
f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean)
|
|
||||||
|
|
||||||
wave = pyworld.synthesize(f0, sp, ap, fs)
|
|
||||||
return fs, wave
|
|
||||||
|
|
||||||
|
|
||||||
class Model:
|
|
||||||
def __init__(
|
|
||||||
self, model_path: Path, config_path: Path, style_vec_path: Path, device: str
|
|
||||||
):
|
|
||||||
self.model_path: Path = model_path
|
|
||||||
self.config_path: Path = config_path
|
|
||||||
self.style_vec_path: Path = style_vec_path
|
|
||||||
self.device: str = device
|
|
||||||
self.hps: utils.HParams = utils.get_hparams_from_file(self.config_path)
|
|
||||||
self.spk2id: dict[str, int] = self.hps.data.spk2id
|
|
||||||
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
|
||||||
|
|
||||||
self.num_styles: int = self.hps.data.num_styles
|
|
||||||
if hasattr(self.hps.data, "style2id"):
|
|
||||||
self.style2id: dict[str, int] = self.hps.data.style2id
|
|
||||||
else:
|
|
||||||
self.style2id: dict[str, int] = {str(i): i for i in range(self.num_styles)}
|
|
||||||
if len(self.style2id) != self.num_styles:
|
|
||||||
raise ValueError(
|
|
||||||
f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.style_vectors: np.ndarray = np.load(self.style_vec_path)
|
|
||||||
if self.style_vectors.shape[0] != self.num_styles:
|
|
||||||
raise ValueError(
|
|
||||||
f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
|
||||||
|
|
||||||
def load_net_g(self):
|
|
||||||
self.net_g = get_net_g(
|
|
||||||
model_path=str(self.model_path),
|
|
||||||
version=self.hps.version,
|
|
||||||
device=self.device,
|
|
||||||
hps=self.hps,
|
|
||||||
)
|
|
||||||
|
|
||||||
def get_style_vector(self, style_id: int, weight: float = 1.0) -> np.ndarray:
|
|
||||||
mean = self.style_vectors[0]
|
|
||||||
style_vec = self.style_vectors[style_id]
|
|
||||||
style_vec = mean + (style_vec - mean) * weight
|
|
||||||
return style_vec
|
|
||||||
|
|
||||||
def get_style_vector_from_audio(
|
|
||||||
self, audio_path: str, weight: float = 1.0
|
|
||||||
) -> np.ndarray:
|
|
||||||
from style_gen import get_style_vector
|
|
||||||
|
|
||||||
xvec = get_style_vector(audio_path)
|
|
||||||
mean = self.style_vectors[0]
|
|
||||||
xvec = mean + (xvec - mean) * weight
|
|
||||||
return xvec
|
|
||||||
|
|
||||||
def infer(
|
|
||||||
self,
|
|
||||||
text: str,
|
|
||||||
language: str = "JP",
|
|
||||||
sid: int = 0,
|
|
||||||
reference_audio_path: Optional[str] = None,
|
|
||||||
sdp_ratio: float = DEFAULT_SDP_RATIO,
|
|
||||||
noise: float = DEFAULT_NOISE,
|
|
||||||
noisew: float = DEFAULT_NOISEW,
|
|
||||||
length: float = DEFAULT_LENGTH,
|
|
||||||
line_split: bool = DEFAULT_LINE_SPLIT,
|
|
||||||
split_interval: float = DEFAULT_SPLIT_INTERVAL,
|
|
||||||
assist_text: Optional[str] = None,
|
|
||||||
assist_text_weight: float = DEFAULT_ASSIST_TEXT_WEIGHT,
|
|
||||||
use_assist_text: bool = False,
|
|
||||||
style: str = DEFAULT_STYLE,
|
|
||||||
style_weight: float = DEFAULT_STYLE_WEIGHT,
|
|
||||||
given_tone: Optional[list[int]] = None,
|
|
||||||
pitch_scale: float = 1.0,
|
|
||||||
intonation_scale: float = 1.0,
|
|
||||||
) -> tuple[int, np.ndarray]:
|
|
||||||
logger.info(f"Start generating audio data from text:\n{text}")
|
|
||||||
if language != "JP" and self.hps.version.endswith("JP-Extra"):
|
|
||||||
raise ValueError(
|
|
||||||
"The model is trained with JP-Extra, but the language is not JP"
|
|
||||||
)
|
|
||||||
if reference_audio_path == "":
|
|
||||||
reference_audio_path = None
|
|
||||||
if assist_text == "" or not use_assist_text:
|
|
||||||
assist_text = None
|
|
||||||
|
|
||||||
if self.net_g is None:
|
|
||||||
self.load_net_g()
|
|
||||||
if reference_audio_path is None:
|
|
||||||
style_id = self.style2id[style]
|
|
||||||
style_vector = self.get_style_vector(style_id, style_weight)
|
|
||||||
else:
|
|
||||||
style_vector = self.get_style_vector_from_audio(
|
|
||||||
reference_audio_path, style_weight
|
|
||||||
)
|
|
||||||
if not line_split:
|
|
||||||
with torch.no_grad():
|
|
||||||
audio = infer(
|
|
||||||
text=text,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
noise_scale=noise,
|
|
||||||
noise_scale_w=noisew,
|
|
||||||
length_scale=length,
|
|
||||||
sid=sid,
|
|
||||||
language=language,
|
|
||||||
hps=self.hps,
|
|
||||||
net_g=self.net_g,
|
|
||||||
device=self.device,
|
|
||||||
assist_text=assist_text,
|
|
||||||
assist_text_weight=assist_text_weight,
|
|
||||||
style_vec=style_vector,
|
|
||||||
given_tone=given_tone,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
texts = text.split("\n")
|
|
||||||
texts = [t for t in texts 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=noisew,
|
|
||||||
length_scale=length,
|
|
||||||
sid=sid,
|
|
||||||
language=language,
|
|
||||||
hps=self.hps,
|
|
||||||
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)
|
|
||||||
logger.info("Audio data generated successfully")
|
|
||||||
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
|
||||||
_, audio = adjust_voice(
|
|
||||||
fs=self.hps.data.sampling_rate,
|
|
||||||
wave=audio,
|
|
||||||
pitch_scale=pitch_scale,
|
|
||||||
intonation_scale=intonation_scale,
|
|
||||||
)
|
|
||||||
with warnings.catch_warnings():
|
|
||||||
warnings.simplefilter("ignore")
|
|
||||||
audio = convert_to_16_bit_wav(audio)
|
|
||||||
return (self.hps.data.sampling_rate, audio)
|
|
||||||
|
|
||||||
|
|
||||||
class ModelHolder:
|
|
||||||
def __init__(self, root_dir: Path, device: str):
|
|
||||||
self.root_dir: Path = root_dir
|
|
||||||
self.device: str = device
|
|
||||||
self.model_files_dict: dict[str, list[Path]] = {}
|
|
||||||
self.current_model: Optional[Model] = None
|
|
||||||
self.model_names: list[str] = []
|
|
||||||
self.models: list[Model] = []
|
|
||||||
self.refresh()
|
|
||||||
|
|
||||||
def refresh(self):
|
|
||||||
self.model_files_dict = {}
|
|
||||||
self.model_names = []
|
|
||||||
self.current_model = None
|
|
||||||
|
|
||||||
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
|
||||||
for model_dir in model_dirs:
|
|
||||||
model_files = [
|
|
||||||
f
|
|
||||||
for f in model_dir.iterdir()
|
|
||||||
if f.suffix in [".pth", ".pt", ".safetensors"]
|
|
||||||
]
|
|
||||||
if len(model_files) == 0:
|
|
||||||
logger.warning(f"No model files found in {model_dir}, so skip it")
|
|
||||||
continue
|
|
||||||
config_path = model_dir / "config.json"
|
|
||||||
if not config_path.exists():
|
|
||||||
logger.warning(
|
|
||||||
f"Config file {config_path} not found, so skip {model_dir}"
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
self.model_files_dict[model_dir.name] = model_files
|
|
||||||
self.model_names.append(model_dir.name)
|
|
||||||
|
|
||||||
def models_info(self):
|
|
||||||
if hasattr(self, "_models_info"):
|
|
||||||
return self._models_info
|
|
||||||
result = []
|
|
||||||
for name, files in self.model_files_dict.items():
|
|
||||||
# Get styles
|
|
||||||
config_path = self.root_dir / name / "config.json"
|
|
||||||
hps = utils.get_hparams_from_file(config_path)
|
|
||||||
style2id: dict[str, int] = hps.data.style2id
|
|
||||||
styles = list(style2id.keys())
|
|
||||||
result.append(
|
|
||||||
{
|
|
||||||
"name": name,
|
|
||||||
"files": [str(f) for f in files],
|
|
||||||
"styles": styles,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
self._models_info = result
|
|
||||||
return result
|
|
||||||
|
|
||||||
def load_model(self, model_name: str, model_path_str: str):
|
|
||||||
model_path = Path(model_path_str)
|
|
||||||
if model_name not in self.model_files_dict:
|
|
||||||
raise ValueError(f"Model `{model_name}` is not found")
|
|
||||||
if model_path not in self.model_files_dict[model_name]:
|
|
||||||
raise ValueError(f"Model file `{model_path}` is not found")
|
|
||||||
if self.current_model is None or self.current_model.model_path != model_path:
|
|
||||||
self.current_model = Model(
|
|
||||||
model_path=model_path,
|
|
||||||
config_path=self.root_dir / model_name / "config.json",
|
|
||||||
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
return self.current_model
|
|
||||||
|
|
||||||
def load_model_gr(
|
|
||||||
self, model_name: str, model_path_str: str
|
|
||||||
) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]:
|
|
||||||
model_path = Path(model_path_str)
|
|
||||||
if model_name not in self.model_files_dict:
|
|
||||||
raise ValueError(f"Model `{model_name}` is not found")
|
|
||||||
if model_path not in self.model_files_dict[model_name]:
|
|
||||||
raise ValueError(f"Model file `{model_path}` is not found")
|
|
||||||
if (
|
|
||||||
self.current_model is not None
|
|
||||||
and self.current_model.model_path == model_path
|
|
||||||
):
|
|
||||||
# Already loaded
|
|
||||||
speakers = list(self.current_model.spk2id.keys())
|
|
||||||
styles = list(self.current_model.style2id.keys())
|
|
||||||
return (
|
|
||||||
gr.Dropdown(choices=styles, value=styles[0]),
|
|
||||||
gr.Button(interactive=True, value="音声合成"),
|
|
||||||
gr.Dropdown(choices=speakers, value=speakers[0]),
|
|
||||||
)
|
|
||||||
self.current_model = Model(
|
|
||||||
model_path=model_path,
|
|
||||||
config_path=self.root_dir / model_name / "config.json",
|
|
||||||
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
speakers = list(self.current_model.spk2id.keys())
|
|
||||||
styles = list(self.current_model.style2id.keys())
|
|
||||||
return (
|
|
||||||
gr.Dropdown(choices=styles, value=styles[0]),
|
|
||||||
gr.Button(interactive=True, value="音声合成"),
|
|
||||||
gr.Dropdown(choices=speakers, value=speakers[0]),
|
|
||||||
)
|
|
||||||
|
|
||||||
def update_model_files_gr(self, model_name: str) -> gr.Dropdown:
|
|
||||||
model_files = self.model_files_dict[model_name]
|
|
||||||
return gr.Dropdown(choices=model_files, value=model_files[0])
|
|
||||||
|
|
||||||
def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
|
|
||||||
self.refresh()
|
|
||||||
initial_model_name = self.model_names[0]
|
|
||||||
initial_model_files = self.model_files_dict[initial_model_name]
|
|
||||||
return (
|
|
||||||
gr.Dropdown(choices=self.model_names, value=initial_model_name),
|
|
||||||
gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]),
|
|
||||||
gr.Button(interactive=False), # For tts_button
|
|
||||||
)
|
|
||||||
152
commons.py
152
commons.py
@@ -1,152 +0,0 @@
|
|||||||
import math
|
|
||||||
import torch
|
|
||||||
from torch.nn import functional as F
|
|
||||||
|
|
||||||
|
|
||||||
def init_weights(m, mean=0.0, std=0.01):
|
|
||||||
classname = m.__class__.__name__
|
|
||||||
if classname.find("Conv") != -1:
|
|
||||||
m.weight.data.normal_(mean, std)
|
|
||||||
|
|
||||||
|
|
||||||
def get_padding(kernel_size, dilation=1):
|
|
||||||
return int((kernel_size * dilation - dilation) / 2)
|
|
||||||
|
|
||||||
|
|
||||||
def convert_pad_shape(pad_shape):
|
|
||||||
layer = pad_shape[::-1]
|
|
||||||
pad_shape = [item for sublist in layer for item in sublist]
|
|
||||||
return pad_shape
|
|
||||||
|
|
||||||
|
|
||||||
def intersperse(lst, item):
|
|
||||||
result = [item] * (len(lst) * 2 + 1)
|
|
||||||
result[1::2] = lst
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def kl_divergence(m_p, logs_p, m_q, logs_q):
|
|
||||||
"""KL(P||Q)"""
|
|
||||||
kl = (logs_q - logs_p) - 0.5
|
|
||||||
kl += (
|
|
||||||
0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q)
|
|
||||||
)
|
|
||||||
return kl
|
|
||||||
|
|
||||||
|
|
||||||
def rand_gumbel(shape):
|
|
||||||
"""Sample from the Gumbel distribution, protect from overflows."""
|
|
||||||
uniform_samples = torch.rand(shape) * 0.99998 + 0.00001
|
|
||||||
return -torch.log(-torch.log(uniform_samples))
|
|
||||||
|
|
||||||
|
|
||||||
def rand_gumbel_like(x):
|
|
||||||
g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device)
|
|
||||||
return g
|
|
||||||
|
|
||||||
|
|
||||||
def slice_segments(x, ids_str, segment_size=4):
|
|
||||||
gather_indices = ids_str.view(x.size(0), 1, 1).repeat(
|
|
||||||
1, x.size(1), 1
|
|
||||||
) + torch.arange(segment_size, device=x.device)
|
|
||||||
return torch.gather(x, 2, gather_indices)
|
|
||||||
|
|
||||||
|
|
||||||
def rand_slice_segments(x, x_lengths=None, segment_size=4):
|
|
||||||
b, d, t = x.size()
|
|
||||||
if x_lengths is None:
|
|
||||||
x_lengths = t
|
|
||||||
ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0)
|
|
||||||
ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long)
|
|
||||||
ret = slice_segments(x, ids_str, segment_size)
|
|
||||||
return ret, ids_str
|
|
||||||
|
|
||||||
|
|
||||||
def get_timing_signal_1d(length, channels, min_timescale=1.0, max_timescale=1.0e4):
|
|
||||||
position = torch.arange(length, dtype=torch.float)
|
|
||||||
num_timescales = channels // 2
|
|
||||||
log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / (
|
|
||||||
num_timescales - 1
|
|
||||||
)
|
|
||||||
inv_timescales = min_timescale * torch.exp(
|
|
||||||
torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment
|
|
||||||
)
|
|
||||||
scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1)
|
|
||||||
signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0)
|
|
||||||
signal = F.pad(signal, [0, 0, 0, channels % 2])
|
|
||||||
signal = signal.view(1, channels, length)
|
|
||||||
return signal
|
|
||||||
|
|
||||||
|
|
||||||
def add_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4):
|
|
||||||
b, channels, length = x.size()
|
|
||||||
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
|
|
||||||
return x + signal.to(dtype=x.dtype, device=x.device)
|
|
||||||
|
|
||||||
|
|
||||||
def cat_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4, axis=1):
|
|
||||||
b, channels, length = x.size()
|
|
||||||
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
|
|
||||||
return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis)
|
|
||||||
|
|
||||||
|
|
||||||
def subsequent_mask(length):
|
|
||||||
mask = torch.tril(torch.ones(length, length)).unsqueeze(0).unsqueeze(0)
|
|
||||||
return mask
|
|
||||||
|
|
||||||
|
|
||||||
@torch.jit.script
|
|
||||||
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
|
|
||||||
n_channels_int = n_channels[0]
|
|
||||||
in_act = input_a + input_b
|
|
||||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
|
||||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
|
||||||
acts = t_act * s_act
|
|
||||||
return acts
|
|
||||||
|
|
||||||
|
|
||||||
def shift_1d(x):
|
|
||||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
def sequence_mask(length, max_length=None):
|
|
||||||
if max_length is None:
|
|
||||||
max_length = length.max()
|
|
||||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
|
||||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
|
||||||
|
|
||||||
|
|
||||||
def generate_path(duration, mask):
|
|
||||||
"""
|
|
||||||
duration: [b, 1, t_x]
|
|
||||||
mask: [b, 1, t_y, t_x]
|
|
||||||
"""
|
|
||||||
|
|
||||||
b, _, t_y, t_x = mask.shape
|
|
||||||
cum_duration = torch.cumsum(duration, -1)
|
|
||||||
|
|
||||||
cum_duration_flat = cum_duration.view(b * t_x)
|
|
||||||
path = sequence_mask(cum_duration_flat, t_y).to(mask.dtype)
|
|
||||||
path = path.view(b, t_x, t_y)
|
|
||||||
path = path - F.pad(path, convert_pad_shape([[0, 0], [1, 0], [0, 0]]))[:, :-1]
|
|
||||||
path = path.unsqueeze(1).transpose(2, 3) * mask
|
|
||||||
return path
|
|
||||||
|
|
||||||
|
|
||||||
def clip_grad_value_(parameters, clip_value, norm_type=2):
|
|
||||||
if isinstance(parameters, torch.Tensor):
|
|
||||||
parameters = [parameters]
|
|
||||||
parameters = list(filter(lambda p: p.grad is not None, parameters))
|
|
||||||
norm_type = float(norm_type)
|
|
||||||
if clip_value is not None:
|
|
||||||
clip_value = float(clip_value)
|
|
||||||
|
|
||||||
total_norm = 0
|
|
||||||
for p in parameters:
|
|
||||||
param_norm = p.grad.data.norm(norm_type)
|
|
||||||
total_norm += param_norm.item() ** norm_type
|
|
||||||
if clip_value is not None:
|
|
||||||
p.grad.data.clamp_(min=-clip_value, max=clip_value)
|
|
||||||
total_norm = total_norm ** (1.0 / norm_type)
|
|
||||||
return total_norm
|
|
||||||
@@ -9,7 +9,8 @@ from typing import Dict, List
|
|||||||
import torch
|
import torch
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
from common.log import logger
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
# If not cuda available, set possible devices to cpu
|
# If not cuda available, set possible devices to cpu
|
||||||
cuda_available = torch.cuda.is_available()
|
cuda_available = torch.cuda.is_available()
|
||||||
@@ -91,7 +92,7 @@ class Bert_gen_config:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config_path: str,
|
config_path: str,
|
||||||
num_processes: int = 2,
|
num_processes: int = 1,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
use_multi_device: bool = False,
|
use_multi_device: bool = False,
|
||||||
):
|
):
|
||||||
@@ -237,7 +238,7 @@ class Config:
|
|||||||
"If you have no special needs, please do not modify default_config.yml."
|
"If you have no special needs, please do not modify default_config.yml."
|
||||||
)
|
)
|
||||||
# sys.exit(0)
|
# sys.exit(0)
|
||||||
with open(file=config_path, mode="r", encoding="utf-8") as file:
|
with open(config_path, "r", encoding="utf-8") as file:
|
||||||
yaml_config: Dict[str, any] = yaml.safe_load(file.read())
|
yaml_config: Dict[str, any] = yaml.safe_load(file.read())
|
||||||
model_name: str = yaml_config["model_name"]
|
model_name: str = yaml_config["model_name"]
|
||||||
self.model_name: str = model_name
|
self.model_name: str = model_name
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
{
|
{
|
||||||
"model_name": "your_model_name",
|
"model_name": "Dummy",
|
||||||
"train": {
|
"train": {
|
||||||
"log_interval": 200,
|
"log_interval": 200,
|
||||||
"eval_interval": 1000,
|
"eval_interval": 1000,
|
||||||
@@ -24,8 +24,9 @@
|
|||||||
"freeze_encoder": false
|
"freeze_encoder": false
|
||||||
},
|
},
|
||||||
"data": {
|
"data": {
|
||||||
"training_files": "Data/your_model_name/filelists/train.list",
|
"use_jp_extra": false,
|
||||||
"validation_files": "Data/your_model_name/filelists/val.list",
|
"training_files": "Data/Dummy/train.list",
|
||||||
|
"validation_files": "Data/Dummy/val.list",
|
||||||
"max_wav_value": 32768.0,
|
"max_wav_value": 32768.0,
|
||||||
"sampling_rate": 44100,
|
"sampling_rate": 44100,
|
||||||
"filter_length": 2048,
|
"filter_length": 2048,
|
||||||
@@ -68,5 +69,5 @@
|
|||||||
"use_spectral_norm": false,
|
"use_spectral_norm": false,
|
||||||
"gin_channels": 256
|
"gin_channels": 256
|
||||||
},
|
},
|
||||||
"version": "2.3.1"
|
"version": "2.4.0"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
{
|
{
|
||||||
|
"model_name": "Dummy",
|
||||||
"train": {
|
"train": {
|
||||||
"log_interval": 200,
|
"log_interval": 200,
|
||||||
"eval_interval": 1000,
|
"eval_interval": 1000,
|
||||||
@@ -27,8 +28,8 @@
|
|||||||
},
|
},
|
||||||
"data": {
|
"data": {
|
||||||
"use_jp_extra": true,
|
"use_jp_extra": true,
|
||||||
"training_files": "filelists/train.list",
|
"training_files": "Data/Dummy/train.list",
|
||||||
"validation_files": "filelists/val.list",
|
"validation_files": "Data/Dummy/val.list",
|
||||||
"max_wav_value": 32768.0,
|
"max_wav_value": 32768.0,
|
||||||
"sampling_rate": 44100,
|
"sampling_rate": 44100,
|
||||||
"filter_length": 2048,
|
"filter_length": 2048,
|
||||||
@@ -75,5 +76,5 @@
|
|||||||
"initial_channel": 64
|
"initial_channel": 64
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"version": "2.3.1-JP-Extra"
|
"version": "2.4.0-JP-Extra"
|
||||||
}
|
}
|
||||||
@@ -7,12 +7,14 @@ import torch
|
|||||||
import torch.utils.data
|
import torch.utils.data
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
import commons
|
|
||||||
from config import config
|
from config import config
|
||||||
from mel_processing import mel_spectrogram_torch, spectrogram_torch
|
from mel_processing import mel_spectrogram_torch, spectrogram_torch
|
||||||
from text import cleaned_text_to_sequence
|
from style_bert_vits2.logging import logger
|
||||||
from common.log import logger
|
from style_bert_vits2.models import commons
|
||||||
from utils import load_filepaths_and_text, load_wav_to_torch
|
from style_bert_vits2.models.hyper_parameters import HyperParametersData
|
||||||
|
from style_bert_vits2.models.utils import load_filepaths_and_text, load_wav_to_torch
|
||||||
|
from style_bert_vits2.nlp import cleaned_text_to_sequence
|
||||||
|
|
||||||
|
|
||||||
"""Multi speaker version"""
|
"""Multi speaker version"""
|
||||||
|
|
||||||
@@ -24,7 +26,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
|||||||
3) computes spectrograms from audio files.
|
3) computes spectrograms from audio files.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, audiopaths_sid_text, hparams):
|
def __init__(self, audiopaths_sid_text: str, hparams: HyperParametersData):
|
||||||
self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)
|
self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)
|
||||||
self.max_wav_value = hparams.max_wav_value
|
self.max_wav_value = hparams.max_wav_value
|
||||||
self.sampling_rate = hparams.sampling_rate
|
self.sampling_rate = hparams.sampling_rate
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ preprocess_text:
|
|||||||
|
|
||||||
bert_gen:
|
bert_gen:
|
||||||
config_path: "config.json"
|
config_path: "config.json"
|
||||||
num_processes: 2
|
num_processes: 1
|
||||||
device: "cuda"
|
device: "cuda"
|
||||||
use_multi_device: false
|
use_multi_device: false
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
|
import json
|
||||||
import os
|
import os
|
||||||
from common.log import logger
|
|
||||||
from common.constants import DEFAULT_STYLE
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import json
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_STYLE
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
def set_style_config(json_path, output_path):
|
def set_style_config(json_path, output_path):
|
||||||
|
|||||||
@@ -1,5 +1,59 @@
|
|||||||
# Changelog
|
# Changelog
|
||||||
|
|
||||||
|
## v2.4.0 (2024-03-15)
|
||||||
|
|
||||||
|
大規模リファクタリング・日本語処理のワーカー化と機能追加等。データセット作り・学習・音声合成・マージ・スタイルWebUIは全て`app.py` (`App.bat`) へ統一されましたのでご注意ください。
|
||||||
|
|
||||||
|
### アップデート手順
|
||||||
|
- 2.3未満(辞書・エディター追加前)からのアップデートの場合は、[Update-to-Dict-Editor.bat](https://github.com/litagin02/Style-Bert-VITS2/releases/download/2.4.0/Update-to-Dict-Editor.bat)をダウンロードし、`Style-Bert-VITS2`フォルダがある場所(インストールbatファイルとかがあったところ)においてダブルクリックしてください。
|
||||||
|
- それ以外の場合は、単純に今までの`Update-Style-Bert-VITS2.bat`でアップデートできます。
|
||||||
|
- ただしアップデートにより多くのファイルが移動したり不要になったりしたので、それらを削除したい場合は[Clean.bat](https://github.com/litagin02/Style-Bert-VITS2/releases/download/2.4.0/Clean.bat)を`Update-Style-Bert-VITS2.bat`と同じ場所に保存して実行してください。
|
||||||
|
|
||||||
|
### 内部改善
|
||||||
|
|
||||||
|
- [tsukumijimaさんによる大規模リファクタリングのプルリク](https://github.com/litagin02/Style-Bert-VITS2/pull/92) によって、内部コードが非常に整理され可読性が高まりライブラリ化もされた。[tsukumijimaさん](https://github.com/tsukumijima) 大変な作業を本当にありがとうございます!
|
||||||
|
- ライブラリとして`pip install style-bert-vits2`によりすぐにインストールでき、音声合成部分の機能が使えます(使用例は[/library.ipynb](/library.ipynb)を参照してください)
|
||||||
|
- その他このプルリクに動機づけられ、多くのコードのリファクタリング・型アノテーションの追加等を行った
|
||||||
|
- 日本語処理のpyopenjtalkをソケット通信を用いて別プロセス化し、複数同時に学習や音声合成を立ち上げても辞書の競合エラーが起きないように。[kale4eat](https://github.com/kale4eat) さんによる[PR](https://github.com/litagin02/Style-Bert-VITS2/pull/89) で
|
||||||
|
|
||||||
|
### バグ修正
|
||||||
|
|
||||||
|
- 上記にもある通り、音声合成と学習前処理など、日本語処理を扱うものを2つ以上起動しようとするとエラーが発生する仕様の解決。ユーザー辞書は追加すれば常にどこからでも適応されます。
|
||||||
|
- `raw`フォルダの直下でなくサブフォルダ内に音声ファイルがある場合に、`wavs`フォルダでもその構造が保たれてしまい、書き起こしファイルとの整合性が取れなくなる挙動を修正し、常に`wav`フォルダ直下へ`wav`ファイルを保存するように変更
|
||||||
|
- スライス時に元ファイル名にピリオド `.` が含まれると、スライス後のファイル名がおかしくなるバグの修正
|
||||||
|
|
||||||
|
### 機能改善・追加
|
||||||
|
|
||||||
|
- 各種WebUIを一つ`app.py` `App.bat` に統一
|
||||||
|
- その他以下の変更や、軽微なUI・説明文の改善等
|
||||||
|
|
||||||
|
**データセット作成**
|
||||||
|
|
||||||
|
- スライス処理の高速化(マルチスレッドにした、大量にスライス元ファイルファイルがある場合に高速になります)、またスライス元のファイルを`wav`以外の`mp3`や`ogg`などの形式にも対応
|
||||||
|
- スライス処理時に、ファイル名にスライスされた開始終了区間を含めるオプションを追加([aka7774](https://github.com/aka7774) さんによるPRです、ありがとうございます!)
|
||||||
|
- 書き起こしの高速化、またHugging FaceのWhisperモデルを使うオプションを追加。バッチサイズを上げることでVRAMを食う代わりに速度が大幅に向上します。
|
||||||
|
|
||||||
|
**学習**
|
||||||
|
|
||||||
|
- 学習元の音声ファイル(`Data/モデル名/raw`にいれるやつ)を、`wav`以外の`mp3`や`ogg`などの形式にも対応(前処理段階で自動的に`wav`ファイルに変換されます)(ただし変わらず1ファイル2-12秒程度の範囲の長さが望ましい)
|
||||||
|
|
||||||
|
**音声合成**
|
||||||
|
|
||||||
|
- 音声合成時に、生成音声の音の高さ(音高)と抑揚の幅を調整できるように(ただし音質が少し劣化する)。`App.bat`や`Editor.bat`のどちらからでも使えます。
|
||||||
|
- `Editor.bat`の複数話者モデルでの話者指定を可能に
|
||||||
|
- `Editor.bat`で、改行を含む文字列をペーストすると自動的に欄が増えるように。また「↑↓」キーで欄を追加・行き来できるように(エディター側で以前に既にアプデしていました)
|
||||||
|
- `Editor.bat`でモデル一覧のリロードをメニューに追加
|
||||||
|
|
||||||
|
**API**
|
||||||
|
|
||||||
|
- `server_fastapi.py`の実行時に全てのモデルファイルを読み込もうとする挙動を修正。音声合成がリクエストされて初めてそのモデルを読み込むように変更(APIを使わない音声合成のときと同じ挙動)
|
||||||
|
- `server_fastapi.py`の音声合成エンドポイント`/voice`について、GETメソッドに加えてPOSTメソッドを追加。GETメソッドでは多くの制約があるようなのでPOSTを使うことが推奨されます。
|
||||||
|
|
||||||
|
**CLI**
|
||||||
|
|
||||||
|
- `preprocess_text.py`で、書き起こしファイルでの音声ファイル名を自動的に正しい`Data/モデル名/wavs/`へ書き換える`--correct_path`オプションの追加(WebUIでは今までもこの挙動でした)
|
||||||
|
- その他上述のデータセット作成の機能追加に伴うCLIのオプションの追加(詳しくは[CLI.md](/docs/CLI.md)を参照)
|
||||||
|
|
||||||
## v2.3.1 (2024-02-27)
|
## v2.3.1 (2024-02-27)
|
||||||
|
|
||||||
### バグ修正
|
### バグ修正
|
||||||
|
|||||||
17
docs/CLI.md
17
docs/CLI.md
@@ -7,7 +7,7 @@ git clone https://github.com/litagin02/Style-Bert-VITS2.git
|
|||||||
cd Style-Bert-VITS2
|
cd Style-Bert-VITS2
|
||||||
python -m venv venv
|
python -m venv venv
|
||||||
venv\Scripts\activate
|
venv\Scripts\activate
|
||||||
pip install torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu118
|
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118
|
||||||
pip install -r requirements.txt
|
pip install -r requirements.txt
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -24,9 +24,11 @@ Optional:
|
|||||||
|
|
||||||
## 1. Dataset preparation
|
## 1. Dataset preparation
|
||||||
|
|
||||||
### 1.1. Slice wavs
|
### 1.1. Slice audio files
|
||||||
|
|
||||||
|
The following audio formats are supported: ".wav", ".flac", ".mp3", ".ogg", ".opus".
|
||||||
```bash
|
```bash
|
||||||
python slice.py --model_name <model_name> [-i <input_dir>] [-m <min_sec>] [-M <max_sec>]
|
python slice.py --model_name <model_name> [-i <input_dir>] [-m <min_sec>] [-M <max_sec>] [--time_suffix]
|
||||||
```
|
```
|
||||||
|
|
||||||
Required:
|
Required:
|
||||||
@@ -36,8 +38,9 @@ Optional:
|
|||||||
- `input_dir`: Path to the directory containing the audio files to slice (default: `inputs`)
|
- `input_dir`: Path to the directory containing the audio files to slice (default: `inputs`)
|
||||||
- `min_sec`: Minimum duration of the sliced audio files in seconds (default: 2).
|
- `min_sec`: Minimum duration of the sliced audio files in seconds (default: 2).
|
||||||
- `max_sec`: Maximum duration of the sliced audio files in seconds (default: 12).
|
- `max_sec`: Maximum duration of the sliced audio files in seconds (default: 12).
|
||||||
|
- `--time_suffix`: Make the filename end with -start_ms-end_ms when saving wav.
|
||||||
|
|
||||||
### 1.2. Transcribe wavs
|
### 1.2. Transcribe audio files
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python transcribe.py --model_name <model_name>
|
python transcribe.py --model_name <model_name>
|
||||||
@@ -50,7 +53,11 @@ Optional
|
|||||||
- `--device`: `cuda` or `cpu` (default: `cuda`).
|
- `--device`: `cuda` or `cpu` (default: `cuda`).
|
||||||
- `--language`: `jp`, `en`, or `en` (default: `jp`).
|
- `--language`: `jp`, `en`, or `en` (default: `jp`).
|
||||||
- `--model`: Whisper model, default: `large-v3`
|
- `--model`: Whisper model, default: `large-v3`
|
||||||
- `--compute_type`: default: `bfloat16`
|
- `--compute_type`: default: `bfloat16`. Only used if not `--use_hf_whisper`.
|
||||||
|
- `--use_hf_whisper`: Use Hugging Face's whisper model instead of default faster-whisper (HF whisper is faster but requires more VRAM).
|
||||||
|
- `--batch_size`: Batch size (default: 16). Only used if `--use_hf_whisper`.
|
||||||
|
- `--num_beams`: Beam size (default: 1).
|
||||||
|
- `--no_repeat_ngram_size`: N-gram size for no repeat (default: 10).
|
||||||
|
|
||||||
## 2. Preprocess
|
## 2. Preprocess
|
||||||
|
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ git clone https://github.com/litagin02/Style-Bert-VITS2.git
|
|||||||
環境構築(デフォルトはPyTorch 1.x系、Python 3.9の模様)
|
環境構築(デフォルトはPyTorch 1.x系、Python 3.9の模様)
|
||||||
```bash
|
```bash
|
||||||
cd /storage/sbv2/Style-Bert-VITS2
|
cd /storage/sbv2/Style-Bert-VITS2
|
||||||
pip install torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu118 && pip install -r requirements.txt
|
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 && pip install -r requirements.txt
|
||||||
```
|
```
|
||||||
事前学習済みモデル等のダウンロード、またパスを`/notebooks/`以下のものに設定
|
事前学習済みモデル等のダウンロード、またパスを`/notebooks/`以下のものに設定
|
||||||
```bash
|
```bash
|
||||||
@@ -57,7 +57,7 @@ mkdir inputs
|
|||||||
unzip Foo.zip -d inputs
|
unzip Foo.zip -d inputs
|
||||||
cd /storage/sbv2/Style-Bert-VITS2
|
cd /storage/sbv2/Style-Bert-VITS2
|
||||||
python slice.py --model_name Foo -i /notebooks/inputs
|
python slice.py --model_name Foo -i /notebooks/inputs
|
||||||
python transcribe.py --model_name Foo
|
python transcribe.py --model_name Foo --use_hf_whisper
|
||||||
```
|
```
|
||||||
|
|
||||||
それが終わったら、以下のコマンドで一括前処理を行う(パラメータは各自お好み、バッチサイズ5か6でVRAM 16GBギリくらい)。
|
それが終わったら、以下のコマンドで一括前処理を行う(パラメータは各自お好み、バッチサイズ5か6でVRAM 16GBギリくらい)。
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
|
import argparse
|
||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
import argparse
|
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。"
|
description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。"
|
||||||
|
|||||||
314
infer.py
314
infer.py
@@ -1,314 +0,0 @@
|
|||||||
import torch
|
|
||||||
|
|
||||||
import commons
|
|
||||||
import utils
|
|
||||||
from models import SynthesizerTrn
|
|
||||||
from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
|
|
||||||
from text import cleaned_text_to_sequence, get_bert
|
|
||||||
from text.cleaner import clean_text
|
|
||||||
from text.symbols import symbols
|
|
||||||
from common.log import logger
|
|
||||||
|
|
||||||
|
|
||||||
class InvalidToneError(ValueError):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
def get_net_g(model_path: str, version: str, device: str, hps):
|
|
||||||
if version.endswith("JP-Extra"):
|
|
||||||
logger.info("Using JP-Extra model")
|
|
||||||
net_g = SynthesizerTrnJPExtra(
|
|
||||||
len(symbols),
|
|
||||||
hps.data.filter_length // 2 + 1,
|
|
||||||
hps.train.segment_size // hps.data.hop_length,
|
|
||||||
n_speakers=hps.data.n_speakers,
|
|
||||||
**hps.model,
|
|
||||||
).to(device)
|
|
||||||
else:
|
|
||||||
logger.info("Using normal model")
|
|
||||||
net_g = SynthesizerTrn(
|
|
||||||
len(symbols),
|
|
||||||
hps.data.filter_length // 2 + 1,
|
|
||||||
hps.train.segment_size // hps.data.hop_length,
|
|
||||||
n_speakers=hps.data.n_speakers,
|
|
||||||
**hps.model,
|
|
||||||
).to(device)
|
|
||||||
net_g.state_dict()
|
|
||||||
_ = net_g.eval()
|
|
||||||
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
|
||||||
_ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
|
|
||||||
elif model_path.endswith(".safetensors"):
|
|
||||||
_ = utils.load_safetensors(model_path, net_g, True)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unknown model format: {model_path}")
|
|
||||||
return net_g
|
|
||||||
|
|
||||||
|
|
||||||
def get_text(
|
|
||||||
text,
|
|
||||||
language_str,
|
|
||||||
hps,
|
|
||||||
device,
|
|
||||||
assist_text=None,
|
|
||||||
assist_text_weight=0.7,
|
|
||||||
given_tone=None,
|
|
||||||
):
|
|
||||||
use_jp_extra = hps.version.endswith("JP-Extra")
|
|
||||||
# 推論のときにのみ呼び出されるので、raise_yomi_errorはFalseに設定
|
|
||||||
norm_text, phone, tone, word2ph = clean_text(
|
|
||||||
text, language_str, use_jp_extra, raise_yomi_error=False
|
|
||||||
)
|
|
||||||
if given_tone is not None:
|
|
||||||
if len(given_tone) != len(phone):
|
|
||||||
raise InvalidToneError(
|
|
||||||
f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})"
|
|
||||||
)
|
|
||||||
tone = given_tone
|
|
||||||
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
|
||||||
|
|
||||||
if hps.data.add_blank:
|
|
||||||
phone = commons.intersperse(phone, 0)
|
|
||||||
tone = commons.intersperse(tone, 0)
|
|
||||||
language = commons.intersperse(language, 0)
|
|
||||||
for i in range(len(word2ph)):
|
|
||||||
word2ph[i] = word2ph[i] * 2
|
|
||||||
word2ph[0] += 1
|
|
||||||
bert_ori = get_bert(
|
|
||||||
norm_text,
|
|
||||||
word2ph,
|
|
||||||
language_str,
|
|
||||||
device,
|
|
||||||
assist_text,
|
|
||||||
assist_text_weight,
|
|
||||||
)
|
|
||||||
del word2ph
|
|
||||||
assert bert_ori.shape[-1] == len(phone), phone
|
|
||||||
|
|
||||||
if language_str == "ZH":
|
|
||||||
bert = bert_ori
|
|
||||||
ja_bert = torch.zeros(1024, len(phone))
|
|
||||||
en_bert = torch.zeros(1024, len(phone))
|
|
||||||
elif language_str == "JP":
|
|
||||||
bert = torch.zeros(1024, len(phone))
|
|
||||||
ja_bert = bert_ori
|
|
||||||
en_bert = torch.zeros(1024, len(phone))
|
|
||||||
elif language_str == "EN":
|
|
||||||
bert = torch.zeros(1024, len(phone))
|
|
||||||
ja_bert = torch.zeros(1024, len(phone))
|
|
||||||
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 = torch.LongTensor(phone)
|
|
||||||
tone = torch.LongTensor(tone)
|
|
||||||
language = torch.LongTensor(language)
|
|
||||||
return bert, ja_bert, en_bert, phone, tone, language
|
|
||||||
|
|
||||||
|
|
||||||
def infer(
|
|
||||||
text,
|
|
||||||
style_vec,
|
|
||||||
sdp_ratio,
|
|
||||||
noise_scale,
|
|
||||||
noise_scale_w,
|
|
||||||
length_scale,
|
|
||||||
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
|
|
||||||
language,
|
|
||||||
hps,
|
|
||||||
net_g,
|
|
||||||
device,
|
|
||||||
skip_start=False,
|
|
||||||
skip_end=False,
|
|
||||||
assist_text=None,
|
|
||||||
assist_text_weight=0.7,
|
|
||||||
given_tone=None,
|
|
||||||
):
|
|
||||||
is_jp_extra = hps.version.endswith("JP-Extra")
|
|
||||||
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
|
||||||
text,
|
|
||||||
language,
|
|
||||||
hps,
|
|
||||||
device,
|
|
||||||
assist_text=assist_text,
|
|
||||||
assist_text_weight=assist_text_weight,
|
|
||||||
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]
|
|
||||||
with torch.no_grad():
|
|
||||||
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 = torch.from_numpy(style_vec).to(device).unsqueeze(0)
|
|
||||||
del phones
|
|
||||||
sid_tensor = torch.LongTensor([sid]).to(device)
|
|
||||||
if is_jp_extra:
|
|
||||||
output = net_g.infer(
|
|
||||||
x_tst,
|
|
||||||
x_tst_lengths,
|
|
||||||
sid_tensor,
|
|
||||||
tones,
|
|
||||||
lang_ids,
|
|
||||||
ja_bert,
|
|
||||||
style_vec=style_vec,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
noise_scale=noise_scale,
|
|
||||||
noise_scale_w=noise_scale_w,
|
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
output = net_g.infer(
|
|
||||||
x_tst,
|
|
||||||
x_tst_lengths,
|
|
||||||
sid_tensor,
|
|
||||||
tones,
|
|
||||||
lang_ids,
|
|
||||||
bert,
|
|
||||||
ja_bert,
|
|
||||||
en_bert,
|
|
||||||
style_vec=style_vec,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
noise_scale=noise_scale,
|
|
||||||
noise_scale_w=noise_scale_w,
|
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
|
||||||
audio = output[0][0, 0].data.cpu().float().numpy()
|
|
||||||
del (
|
|
||||||
x_tst,
|
|
||||||
tones,
|
|
||||||
lang_ids,
|
|
||||||
bert,
|
|
||||||
x_tst_lengths,
|
|
||||||
sid_tensor,
|
|
||||||
ja_bert,
|
|
||||||
en_bert,
|
|
||||||
style_vec,
|
|
||||||
) # , emo
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
return audio
|
|
||||||
|
|
||||||
|
|
||||||
def infer_multilang(
|
|
||||||
text,
|
|
||||||
style_vec,
|
|
||||||
sdp_ratio,
|
|
||||||
noise_scale,
|
|
||||||
noise_scale_w,
|
|
||||||
length_scale,
|
|
||||||
sid,
|
|
||||||
language,
|
|
||||||
hps,
|
|
||||||
net_g,
|
|
||||||
device,
|
|
||||||
skip_start=False,
|
|
||||||
skip_end=False,
|
|
||||||
):
|
|
||||||
bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], []
|
|
||||||
# emo = get_emo_(reference_audio, emotion, sid)
|
|
||||||
# if isinstance(reference_audio, np.ndarray):
|
|
||||||
# emo = get_clap_audio_feature(reference_audio, device)
|
|
||||||
# else:
|
|
||||||
# emo = get_clap_text_feature(emotion, device)
|
|
||||||
# emo = torch.squeeze(emo, dim=1)
|
|
||||||
for idx, (txt, lang) in enumerate(zip(text, language)):
|
|
||||||
_skip_start = (idx != 0) or (skip_start and idx == 0)
|
|
||||||
_skip_end = (idx != len(language) - 1) or skip_end
|
|
||||||
(
|
|
||||||
temp_bert,
|
|
||||||
temp_ja_bert,
|
|
||||||
temp_en_bert,
|
|
||||||
temp_phones,
|
|
||||||
temp_tones,
|
|
||||||
temp_lang_ids,
|
|
||||||
) = get_text(txt, lang, hps, device)
|
|
||||||
if _skip_start:
|
|
||||||
temp_bert = temp_bert[:, 3:]
|
|
||||||
temp_ja_bert = temp_ja_bert[:, 3:]
|
|
||||||
temp_en_bert = temp_en_bert[:, 3:]
|
|
||||||
temp_phones = temp_phones[3:]
|
|
||||||
temp_tones = temp_tones[3:]
|
|
||||||
temp_lang_ids = temp_lang_ids[3:]
|
|
||||||
if _skip_end:
|
|
||||||
temp_bert = temp_bert[:, :-2]
|
|
||||||
temp_ja_bert = temp_ja_bert[:, :-2]
|
|
||||||
temp_en_bert = temp_en_bert[:, :-2]
|
|
||||||
temp_phones = temp_phones[:-2]
|
|
||||||
temp_tones = temp_tones[:-2]
|
|
||||||
temp_lang_ids = temp_lang_ids[:-2]
|
|
||||||
bert.append(temp_bert)
|
|
||||||
ja_bert.append(temp_ja_bert)
|
|
||||||
en_bert.append(temp_en_bert)
|
|
||||||
phones.append(temp_phones)
|
|
||||||
tones.append(temp_tones)
|
|
||||||
lang_ids.append(temp_lang_ids)
|
|
||||||
bert = torch.concatenate(bert, dim=1)
|
|
||||||
ja_bert = torch.concatenate(ja_bert, dim=1)
|
|
||||||
en_bert = torch.concatenate(en_bert, dim=1)
|
|
||||||
phones = torch.concatenate(phones, dim=0)
|
|
||||||
tones = torch.concatenate(tones, dim=0)
|
|
||||||
lang_ids = torch.concatenate(lang_ids, dim=0)
|
|
||||||
with torch.no_grad():
|
|
||||||
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)
|
|
||||||
# emo = emo.to(device).unsqueeze(0)
|
|
||||||
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
|
|
||||||
del phones
|
|
||||||
speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device)
|
|
||||||
audio = (
|
|
||||||
net_g.infer(
|
|
||||||
x_tst,
|
|
||||||
x_tst_lengths,
|
|
||||||
speakers,
|
|
||||||
tones,
|
|
||||||
lang_ids,
|
|
||||||
bert,
|
|
||||||
ja_bert,
|
|
||||||
en_bert,
|
|
||||||
style_vec=style_vec,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
noise_scale=noise_scale,
|
|
||||||
noise_scale_w=noise_scale_w,
|
|
||||||
length_scale=length_scale,
|
|
||||||
)[0][0, 0]
|
|
||||||
.data.cpu()
|
|
||||||
.float()
|
|
||||||
.numpy()
|
|
||||||
)
|
|
||||||
del (
|
|
||||||
x_tst,
|
|
||||||
tones,
|
|
||||||
lang_ids,
|
|
||||||
bert,
|
|
||||||
x_tst_lengths,
|
|
||||||
speakers,
|
|
||||||
ja_bert,
|
|
||||||
en_bert,
|
|
||||||
) # , emo
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
return audio
|
|
||||||
@@ -5,23 +5,18 @@ from pathlib import Path
|
|||||||
import yaml
|
import yaml
|
||||||
from huggingface_hub import hf_hub_download
|
from huggingface_hub import hf_hub_download
|
||||||
|
|
||||||
from common.log import logger
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
def download_bert_models():
|
def download_bert_models():
|
||||||
with open("bert/bert_models.json", "r") as fp:
|
with open("bert/bert_models.json", "r", encoding="utf-8") as fp:
|
||||||
models = json.load(fp)
|
models = json.load(fp)
|
||||||
for k, v in models.items():
|
for k, v in models.items():
|
||||||
local_path = Path("bert").joinpath(k)
|
local_path = Path("bert").joinpath(k)
|
||||||
for file in v["files"]:
|
for file in v["files"]:
|
||||||
if not Path(local_path).joinpath(file).exists():
|
if not Path(local_path).joinpath(file).exists():
|
||||||
logger.info(f"Downloading {k} {file}")
|
logger.info(f"Downloading {k} {file}")
|
||||||
hf_hub_download(
|
hf_hub_download(v["repo_id"], file, local_dir=local_path)
|
||||||
v["repo_id"],
|
|
||||||
file,
|
|
||||||
local_dir=local_path,
|
|
||||||
local_dir_use_symlinks=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def download_slm_model():
|
def download_slm_model():
|
||||||
@@ -29,12 +24,7 @@ def download_slm_model():
|
|||||||
file = "pytorch_model.bin"
|
file = "pytorch_model.bin"
|
||||||
if not Path(local_path).joinpath(file).exists():
|
if not Path(local_path).joinpath(file).exists():
|
||||||
logger.info(f"Downloading wavlm-base-plus {file}")
|
logger.info(f"Downloading wavlm-base-plus {file}")
|
||||||
hf_hub_download(
|
hf_hub_download("microsoft/wavlm-base-plus", file, local_dir=local_path)
|
||||||
"microsoft/wavlm-base-plus",
|
|
||||||
file,
|
|
||||||
local_dir=local_path,
|
|
||||||
local_dir_use_symlinks=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def download_pretrained_models():
|
def download_pretrained_models():
|
||||||
@@ -44,10 +34,7 @@ def download_pretrained_models():
|
|||||||
if not Path(local_path).joinpath(file).exists():
|
if not Path(local_path).joinpath(file).exists():
|
||||||
logger.info(f"Downloading pretrained {file}")
|
logger.info(f"Downloading pretrained {file}")
|
||||||
hf_hub_download(
|
hf_hub_download(
|
||||||
"litagin/Style-Bert-VITS2-1.0-base",
|
"litagin/Style-Bert-VITS2-1.0-base", file, local_dir=local_path
|
||||||
file,
|
|
||||||
local_dir=local_path,
|
|
||||||
local_dir_use_symlinks=False,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -58,10 +45,7 @@ def download_jp_extra_pretrained_models():
|
|||||||
if not Path(local_path).joinpath(file).exists():
|
if not Path(local_path).joinpath(file).exists():
|
||||||
logger.info(f"Downloading JP-Extra pretrained {file}")
|
logger.info(f"Downloading JP-Extra pretrained {file}")
|
||||||
hf_hub_download(
|
hf_hub_download(
|
||||||
"litagin/Style-Bert-VITS2-2.0-base-JP-Extra",
|
"litagin/Style-Bert-VITS2-2.0-base-JP-Extra", file, local_dir=local_path
|
||||||
file,
|
|
||||||
local_dir=local_path,
|
|
||||||
local_dir_use_symlinks=False,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
138
library.ipynb
Normal file
138
library.ipynb
Normal file
@@ -0,0 +1,138 @@
|
|||||||
|
{
|
||||||
|
"cells": [
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"metadata": {},
|
||||||
|
"source": [
|
||||||
|
"# Style-Bert-VITS2ライブラリの使用例\n",
|
||||||
|
"\n",
|
||||||
|
"`pip install style-bert-vits2`を使った、jupyter notebookでの使用例です。Google colab等でも動きます。"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# PyTorch環境の構築(ない場合)\n",
|
||||||
|
"# 参照: https://pytorch.org/get-started/locally/\n",
|
||||||
|
"\n",
|
||||||
|
"!pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "LLrngKcQEAyP"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# style-bert-vits2のインストール\n",
|
||||||
|
"\n",
|
||||||
|
"!pip install style-bert-vits2"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "9xRtfUg5EZkx"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# BERTモデルをロード(ローカルに手動でダウンロードする必要はありません)\n",
|
||||||
|
"\n",
|
||||||
|
"from style_bert_vits2.nlp import bert_models\n",
|
||||||
|
"from style_bert_vits2.constants import Languages\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"bert_models.load_model(Languages.JP, \"ku-nlp/deberta-v2-large-japanese-char-wwm\")\n",
|
||||||
|
"bert_models.load_tokenizer(Languages.JP, \"ku-nlp/deberta-v2-large-japanese-char-wwm\")\n",
|
||||||
|
"# bert_models.load_model(Languages.EN, \"microsoft/deberta-v3-large\")\n",
|
||||||
|
"# bert_models.load_tokenizer(Languages.EN, \"microsoft/deberta-v3-large\")\n",
|
||||||
|
"# bert_models.load_model(Languages.ZH, \"hfl/chinese-roberta-wwm-ext-large\")\n",
|
||||||
|
"# bert_models.load_tokenizer(Languages.ZH, \"hfl/chinese-roberta-wwm-ext-large\")"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "q2V9d3HyFAr_"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# Hugging Faceから試しにデフォルトモデルをダウンロードしてみて、それを音声合成に使ってみる\n",
|
||||||
|
"# model_assetsディレクトリにダウンロードされます\n",
|
||||||
|
"\n",
|
||||||
|
"from pathlib import Path\n",
|
||||||
|
"from huggingface_hub import hf_hub_download\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"model_file = \"jvnv-F1-jp/jvnv-F1-jp_e160_s14000.safetensors\"\n",
|
||||||
|
"config_file = \"jvnv-F1-jp/config.json\"\n",
|
||||||
|
"style_file = \"jvnv-F1-jp/style_vectors.npy\"\n",
|
||||||
|
"\n",
|
||||||
|
"for file in [model_file, config_file, style_file]:\n",
|
||||||
|
" print(file)\n",
|
||||||
|
" hf_hub_download(\n",
|
||||||
|
" \"litagin/style_bert_vits2_jvnv\",\n",
|
||||||
|
" file,\n",
|
||||||
|
" local_dir=\"model_assets\"\n",
|
||||||
|
" )"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "hJa31MEUFhe4"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# 上でダウンロードしたモデルファイルを指定して音声合成のテスト\n",
|
||||||
|
"\n",
|
||||||
|
"from style_bert_vits2.tts_model import TTSModel\n",
|
||||||
|
"\n",
|
||||||
|
"assets_root = Path(\"model_assets\")\n",
|
||||||
|
"\n",
|
||||||
|
"model = TTSModel(\n",
|
||||||
|
" model_path=assets_root / model_file,\n",
|
||||||
|
" config_path=assets_root / config_file,\n",
|
||||||
|
" style_vec_path=assets_root / style_file,\n",
|
||||||
|
" device=\"cpu\"\n",
|
||||||
|
")"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "Gal0tqrtGXZx"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"from IPython.display import Audio, display\n",
|
||||||
|
"\n",
|
||||||
|
"sr, audio = model.infer(text=\"こんにちは\")\n",
|
||||||
|
"display(Audio(audio, rate=sr))"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"colab": {
|
||||||
|
"provenance": []
|
||||||
|
},
|
||||||
|
"kernelspec": {
|
||||||
|
"display_name": "Python 3",
|
||||||
|
"name": "python3"
|
||||||
|
},
|
||||||
|
"language_info": {
|
||||||
|
"name": "python"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"nbformat": 4,
|
||||||
|
"nbformat_minor": 0
|
||||||
|
}
|
||||||
@@ -2,8 +2,6 @@ import torch
|
|||||||
import torchaudio
|
import torchaudio
|
||||||
from transformers import AutoModel
|
from transformers import AutoModel
|
||||||
|
|
||||||
from common.log import logger
|
|
||||||
|
|
||||||
|
|
||||||
def feature_loss(fmap_r, fmap_g):
|
def feature_loss(fmap_r, fmap_g):
|
||||||
loss = 0
|
loss = 0
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
|
import warnings
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.utils.data
|
import torch.utils.data
|
||||||
from librosa.filters import mel as librosa_mel_fn
|
from librosa.filters import mel as librosa_mel_fn
|
||||||
import warnings
|
|
||||||
|
|
||||||
# warnings.simplefilter(action='ignore', category=FutureWarning)
|
# warnings.simplefilter(action='ignore', category=FutureWarning)
|
||||||
warnings.filterwarnings(action="ignore")
|
warnings.filterwarnings(action="ignore")
|
||||||
|
|||||||
@@ -1,16 +0,0 @@
|
|||||||
from numpy import zeros, int32, float32
|
|
||||||
from torch import from_numpy
|
|
||||||
|
|
||||||
from .core import maximum_path_jit
|
|
||||||
|
|
||||||
|
|
||||||
def maximum_path(neg_cent, mask):
|
|
||||||
device = neg_cent.device
|
|
||||||
dtype = neg_cent.dtype
|
|
||||||
neg_cent = neg_cent.data.cpu().numpy().astype(float32)
|
|
||||||
path = zeros(neg_cent.shape, dtype=int32)
|
|
||||||
|
|
||||||
t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32)
|
|
||||||
t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32)
|
|
||||||
maximum_path_jit(path, neg_cent, t_t_max, t_s_max)
|
|
||||||
return from_numpy(path).to(device=device, dtype=dtype)
|
|
||||||
@@ -1,46 +0,0 @@
|
|||||||
import numba
|
|
||||||
|
|
||||||
|
|
||||||
@numba.jit(
|
|
||||||
numba.void(
|
|
||||||
numba.int32[:, :, ::1],
|
|
||||||
numba.float32[:, :, ::1],
|
|
||||||
numba.int32[::1],
|
|
||||||
numba.int32[::1],
|
|
||||||
),
|
|
||||||
nopython=True,
|
|
||||||
nogil=True,
|
|
||||||
)
|
|
||||||
def maximum_path_jit(paths, values, t_ys, t_xs):
|
|
||||||
b = paths.shape[0]
|
|
||||||
max_neg_val = -1e9
|
|
||||||
for i in range(int(b)):
|
|
||||||
path = paths[i]
|
|
||||||
value = values[i]
|
|
||||||
t_y = t_ys[i]
|
|
||||||
t_x = t_xs[i]
|
|
||||||
|
|
||||||
v_prev = v_cur = 0.0
|
|
||||||
index = t_x - 1
|
|
||||||
|
|
||||||
for y in range(t_y):
|
|
||||||
for x in range(max(0, t_x + y - t_y), min(t_x, y + 1)):
|
|
||||||
if x == y:
|
|
||||||
v_cur = max_neg_val
|
|
||||||
else:
|
|
||||||
v_cur = value[y - 1, x]
|
|
||||||
if x == 0:
|
|
||||||
if y == 0:
|
|
||||||
v_prev = 0.0
|
|
||||||
else:
|
|
||||||
v_prev = max_neg_val
|
|
||||||
else:
|
|
||||||
v_prev = value[y - 1, x - 1]
|
|
||||||
value[y, x] += max(v_prev, v_cur)
|
|
||||||
|
|
||||||
for y in range(t_y - 1, -1, -1):
|
|
||||||
path[y, index] = 1
|
|
||||||
if index != 0 and (
|
|
||||||
index == y or value[y - 1, index] < value[y - 1, index - 1]
|
|
||||||
):
|
|
||||||
index = index - 1
|
|
||||||
@@ -1,7 +1,17 @@
|
|||||||
import argparse
|
import argparse
|
||||||
from webui_train import preprocess_all
|
|
||||||
from multiprocessing import cpu_count
|
from multiprocessing import cpu_count
|
||||||
|
|
||||||
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
||||||
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
|
from webui.train import preprocess_all
|
||||||
|
|
||||||
|
|
||||||
|
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
||||||
|
pyopenjtalk_worker.initialize_worker()
|
||||||
|
|
||||||
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
|
update_dict()
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser()
|
parser = argparse.ArgumentParser()
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -75,7 +85,10 @@ if __name__ == "__main__":
|
|||||||
default=200,
|
default=200,
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--yomi_error", type=str, help="Yomi error. raise, skip, use", default="raise"
|
"--yomi_error",
|
||||||
|
type=str,
|
||||||
|
help="Yomi error. Options: raise, skip, use",
|
||||||
|
default="raise",
|
||||||
)
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|||||||
@@ -1,136 +1,162 @@
|
|||||||
|
import argparse
|
||||||
import json
|
import json
|
||||||
import os
|
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
|
from pathlib import Path
|
||||||
from random import shuffle
|
from random import shuffle
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import click
|
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from common.log import logger
|
from config import Preprocess_text_config, config
|
||||||
from common.stdout_wrapper import SAFE_STDOUT
|
from style_bert_vits2.logging import logger
|
||||||
from config import config
|
from style_bert_vits2.nlp import clean_text
|
||||||
from text.cleaner import clean_text
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker
|
||||||
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
preprocess_text_config = config.preprocess_text_config
|
|
||||||
|
# このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化
|
||||||
|
pyopenjtalk_worker.initialize_worker()
|
||||||
|
|
||||||
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
|
update_dict()
|
||||||
|
|
||||||
|
|
||||||
|
preprocess_text_config: Preprocess_text_config = config.preprocess_text_config
|
||||||
|
|
||||||
|
|
||||||
# Count lines for tqdm
|
# Count lines for tqdm
|
||||||
def count_lines(file_path: str):
|
def count_lines(file_path: Path):
|
||||||
with open(file_path, "r", encoding="utf-8") as file:
|
with file_path.open("r", encoding="utf-8") as file:
|
||||||
return sum(1 for _ in file)
|
return sum(1 for _ in file)
|
||||||
|
|
||||||
|
|
||||||
@click.command()
|
def write_error_log(error_log_path: Path, line: str, error: Exception):
|
||||||
@click.option(
|
with error_log_path.open("a", encoding="utf-8") as error_log:
|
||||||
"--transcription-path",
|
error_log.write(f"{line.strip()}\n{error}\n\n")
|
||||||
default=preprocess_text_config.transcription_path,
|
|
||||||
type=click.Path(exists=True, file_okay=True, dir_okay=False),
|
|
||||||
)
|
def process_line(
|
||||||
@click.option("--cleaned-path", default=preprocess_text_config.cleaned_path)
|
line: str,
|
||||||
@click.option("--train-path", default=preprocess_text_config.train_path)
|
transcription_path: Path,
|
||||||
@click.option("--val-path", default=preprocess_text_config.val_path)
|
correct_path: bool,
|
||||||
@click.option(
|
|
||||||
"--config-path",
|
|
||||||
default=preprocess_text_config.config_path,
|
|
||||||
type=click.Path(exists=True, file_okay=True, dir_okay=False),
|
|
||||||
)
|
|
||||||
@click.option("--val-per-lang", default=preprocess_text_config.val_per_lang)
|
|
||||||
@click.option("--max-val-total", default=preprocess_text_config.max_val_total)
|
|
||||||
@click.option("--clean/--no-clean", default=preprocess_text_config.clean)
|
|
||||||
@click.option("-y", "--yml_config")
|
|
||||||
@click.option("--use_jp_extra", is_flag=True)
|
|
||||||
@click.option("--yomi_error", default="raise")
|
|
||||||
def preprocess(
|
|
||||||
transcription_path: str,
|
|
||||||
cleaned_path: Optional[str],
|
|
||||||
train_path: str,
|
|
||||||
val_path: str,
|
|
||||||
config_path: str,
|
|
||||||
val_per_lang: int,
|
|
||||||
max_val_total: int,
|
|
||||||
clean: bool,
|
|
||||||
yml_config: str, # 这个不要删
|
|
||||||
use_jp_extra: bool,
|
use_jp_extra: bool,
|
||||||
yomi_error: str,
|
yomi_error: str,
|
||||||
|
):
|
||||||
|
splitted_line = line.strip().split("|")
|
||||||
|
if len(splitted_line) != 4:
|
||||||
|
raise ValueError(f"Invalid line format: {line.strip()}")
|
||||||
|
utt, spk, language, text = splitted_line
|
||||||
|
norm_text, phones, tones, word2ph = clean_text(
|
||||||
|
text=text,
|
||||||
|
language=language, # type: ignore
|
||||||
|
use_jp_extra=use_jp_extra,
|
||||||
|
raise_yomi_error=(yomi_error != "use"),
|
||||||
|
)
|
||||||
|
if correct_path:
|
||||||
|
utt = str(transcription_path.parent / "wavs" / utt)
|
||||||
|
|
||||||
|
return "{}|{}|{}|{}|{}|{}|{}\n".format(
|
||||||
|
utt,
|
||||||
|
spk,
|
||||||
|
language,
|
||||||
|
norm_text,
|
||||||
|
" ".join(phones),
|
||||||
|
" ".join([str(i) for i in tones]),
|
||||||
|
" ".join([str(i) for i in word2ph]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess(
|
||||||
|
transcription_path: Path,
|
||||||
|
cleaned_path: Optional[Path],
|
||||||
|
train_path: Path,
|
||||||
|
val_path: Path,
|
||||||
|
config_path: Path,
|
||||||
|
val_per_lang: int,
|
||||||
|
max_val_total: int,
|
||||||
|
# clean: bool,
|
||||||
|
use_jp_extra: bool,
|
||||||
|
yomi_error: str,
|
||||||
|
correct_path: bool,
|
||||||
):
|
):
|
||||||
assert yomi_error in ["raise", "skip", "use"]
|
assert yomi_error in ["raise", "skip", "use"]
|
||||||
if cleaned_path == "" or cleaned_path is None:
|
if cleaned_path == "" or cleaned_path is None:
|
||||||
cleaned_path = transcription_path + ".cleaned"
|
cleaned_path = transcription_path.with_name(
|
||||||
|
transcription_path.name + ".cleaned"
|
||||||
|
)
|
||||||
|
|
||||||
error_log_path = os.path.join(os.path.dirname(cleaned_path), "text_error.log")
|
error_log_path = transcription_path.parent / "text_error.log"
|
||||||
if os.path.exists(error_log_path):
|
if error_log_path.exists():
|
||||||
os.remove(error_log_path)
|
error_log_path.unlink()
|
||||||
error_count = 0
|
error_count = 0
|
||||||
|
|
||||||
if clean:
|
total_lines = count_lines(transcription_path)
|
||||||
total_lines = count_lines(transcription_path)
|
|
||||||
with open(cleaned_path, "w", encoding="utf-8") as out_file:
|
|
||||||
with open(transcription_path, "r", encoding="utf-8") as trans_file:
|
|
||||||
for line in tqdm(trans_file, file=SAFE_STDOUT, total=total_lines):
|
|
||||||
try:
|
|
||||||
utt, spk, language, text = line.strip().split("|")
|
|
||||||
norm_text, phones, tones, word2ph = clean_text(
|
|
||||||
text=text,
|
|
||||||
language=language,
|
|
||||||
use_jp_extra=use_jp_extra,
|
|
||||||
raise_yomi_error=(yomi_error != "use"),
|
|
||||||
)
|
|
||||||
|
|
||||||
out_file.write(
|
# transcription_path から 1行ずつ読み込んで文章処理して cleaned_path に書き込む
|
||||||
"{}|{}|{}|{}|{}|{}|{}\n".format(
|
with (
|
||||||
utt,
|
transcription_path.open("r", encoding="utf-8") as trans_file,
|
||||||
spk,
|
cleaned_path.open("w", encoding="utf-8") as out_file,
|
||||||
language,
|
):
|
||||||
norm_text,
|
for line in tqdm(trans_file, file=SAFE_STDOUT, total=total_lines):
|
||||||
" ".join(phones),
|
try:
|
||||||
" ".join([str(i) for i in tones]),
|
processed_line = process_line(
|
||||||
" ".join([str(i) for i in word2ph]),
|
line,
|
||||||
)
|
transcription_path,
|
||||||
)
|
correct_path,
|
||||||
except Exception as e:
|
use_jp_extra,
|
||||||
logger.error(
|
yomi_error,
|
||||||
f"An error occurred at line:\n{line.strip()}\n{e}",
|
)
|
||||||
encoding="utf-8",
|
out_file.write(processed_line)
|
||||||
)
|
except Exception as e:
|
||||||
with open(error_log_path, "a", encoding="utf-8") as error_log:
|
logger.error(
|
||||||
error_log.write(f"{line.strip()}\n{e}\n\n")
|
f"An error occurred at line:\n{line.strip()}\n{e}", encoding="utf-8"
|
||||||
error_count += 1
|
)
|
||||||
|
write_error_log(error_log_path, line, e)
|
||||||
|
error_count += 1
|
||||||
|
|
||||||
transcription_path = cleaned_path
|
transcription_path = cleaned_path
|
||||||
spk_utt_map = defaultdict(list)
|
|
||||||
spk_id_map = {}
|
|
||||||
current_sid = 0
|
|
||||||
|
|
||||||
with open(transcription_path, "r", encoding="utf-8") as f:
|
# 各話者ごとのlineの辞書
|
||||||
audioPaths = set()
|
spk_utt_map: dict[str, list[str]] = defaultdict(list)
|
||||||
countSame = 0
|
|
||||||
countNotFound = 0
|
# 話者からIDへの写像
|
||||||
|
spk_id_map: dict[str, int] = {}
|
||||||
|
|
||||||
|
# 話者ID
|
||||||
|
current_sid: int = 0
|
||||||
|
|
||||||
|
# 音源ファイルのチェックや、spk_id_mapの作成
|
||||||
|
with transcription_path.open("r", encoding="utf-8") as f:
|
||||||
|
audio_paths: set[str] = set()
|
||||||
|
count_same = 0
|
||||||
|
count_not_found = 0
|
||||||
for line in f.readlines():
|
for line in f.readlines():
|
||||||
utt, spk, language, text, phones, tones, word2ph = line.strip().split("|")
|
utt, spk = line.strip().split("|")[:2]
|
||||||
if utt in audioPaths:
|
if utt in audio_paths:
|
||||||
# 过滤数据集错误:相同的音频匹配多个文本,导致后续bert出问题
|
logger.warning(f"Same audio file appears multiple times: {utt}")
|
||||||
logger.warning(f"Same audio matches multiple texts: {line}")
|
count_same += 1
|
||||||
countSame += 1
|
|
||||||
continue
|
continue
|
||||||
if not os.path.isfile(utt):
|
if not Path(utt).is_file():
|
||||||
# 过滤数据集错误:不存在对应音频
|
|
||||||
logger.warning(f"Audio not found: {utt}")
|
logger.warning(f"Audio not found: {utt}")
|
||||||
countNotFound += 1
|
count_not_found += 1
|
||||||
continue
|
continue
|
||||||
audioPaths.add(utt)
|
audio_paths.add(utt)
|
||||||
spk_utt_map[language].append(line)
|
spk_utt_map[spk].append(line)
|
||||||
|
|
||||||
|
# 新しい話者が出てきたら話者IDを割り当て、current_sidを1増やす
|
||||||
if spk not in spk_id_map.keys():
|
if spk not in spk_id_map.keys():
|
||||||
spk_id_map[spk] = current_sid
|
spk_id_map[spk] = current_sid
|
||||||
current_sid += 1
|
current_sid += 1
|
||||||
if countSame > 0 or countNotFound > 0:
|
if count_same > 0 or count_not_found > 0:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Total repeated audios: {countSame}, Total number of audio not found: {countNotFound}"
|
f"Total repeated audios: {count_same}, Total number of audio not found: {count_not_found}"
|
||||||
)
|
)
|
||||||
|
|
||||||
train_list = []
|
train_list: list[str] = []
|
||||||
val_list = []
|
val_list: list[str] = []
|
||||||
|
|
||||||
|
# 各話者ごとにシャッフルして、val_per_lang個をval_listに、残りをtrain_listに追加
|
||||||
for spk, utts in spk_utt_map.items():
|
for spk, utts in spk_utt_map.items():
|
||||||
shuffle(utts)
|
shuffle(utts)
|
||||||
val_list += utts[:val_per_lang]
|
val_list += utts[:val_per_lang]
|
||||||
@@ -141,26 +167,21 @@ def preprocess(
|
|||||||
train_list += val_list[max_val_total:]
|
train_list += val_list[max_val_total:]
|
||||||
val_list = val_list[:max_val_total]
|
val_list = val_list[:max_val_total]
|
||||||
|
|
||||||
with open(train_path, "w", encoding="utf-8") as f:
|
with train_path.open("w", encoding="utf-8") as f:
|
||||||
for line in train_list:
|
for line in train_list:
|
||||||
f.write(line)
|
f.write(line)
|
||||||
|
|
||||||
with open(val_path, "w", encoding="utf-8") as f:
|
with val_path.open("w", encoding="utf-8") as f:
|
||||||
for line in val_list:
|
for line in val_list:
|
||||||
f.write(line)
|
f.write(line)
|
||||||
|
|
||||||
json_config = json.load(open(config_path, encoding="utf-8"))
|
with config_path.open("r", encoding="utf-8") as f:
|
||||||
|
json_config = json.load(f)
|
||||||
|
|
||||||
json_config["data"]["spk2id"] = spk_id_map
|
json_config["data"]["spk2id"] = spk_id_map
|
||||||
json_config["data"]["n_speakers"] = len(spk_id_map)
|
json_config["data"]["n_speakers"] = len(spk_id_map)
|
||||||
# 新增写入:写入训练版本、数据集路径
|
|
||||||
# json_config["version"] = latest_version
|
with config_path.open("w", encoding="utf-8") as f:
|
||||||
json_config["data"]["training_files"] = os.path.normpath(train_path).replace(
|
|
||||||
"\\", "/"
|
|
||||||
)
|
|
||||||
json_config["data"]["validation_files"] = os.path.normpath(val_path).replace(
|
|
||||||
"\\", "/"
|
|
||||||
)
|
|
||||||
with open(config_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(json_config, f, indent=2, ensure_ascii=False)
|
json.dump(json_config, f, indent=2, ensure_ascii=False)
|
||||||
if error_count > 0:
|
if error_count > 0:
|
||||||
if yomi_error == "skip":
|
if yomi_error == "skip":
|
||||||
@@ -185,4 +206,49 @@ def preprocess(
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
preprocess()
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument(
|
||||||
|
"--transcription-path", default=preprocess_text_config.transcription_path
|
||||||
|
)
|
||||||
|
parser.add_argument("--cleaned-path", default=preprocess_text_config.cleaned_path)
|
||||||
|
parser.add_argument("--train-path", default=preprocess_text_config.train_path)
|
||||||
|
parser.add_argument("--val-path", default=preprocess_text_config.val_path)
|
||||||
|
parser.add_argument("--config-path", default=preprocess_text_config.config_path)
|
||||||
|
|
||||||
|
# 「話者ごと」のバリデーションデータ数、言語ごとではない!
|
||||||
|
# 元のコードや設定ファイルでval_per_langとなっていたので名前をそのままにしている
|
||||||
|
parser.add_argument(
|
||||||
|
"--val-per-lang",
|
||||||
|
default=preprocess_text_config.val_per_lang,
|
||||||
|
help="Number of validation data per SPEAKER, not per language (due to compatibility with the original code).",
|
||||||
|
)
|
||||||
|
parser.add_argument("--max-val-total", default=preprocess_text_config.max_val_total)
|
||||||
|
parser.add_argument("--use_jp_extra", action="store_true")
|
||||||
|
parser.add_argument("--yomi_error", default="raise")
|
||||||
|
parser.add_argument("--correct_path", action="store_true")
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
transcription_path = Path(args.transcription_path)
|
||||||
|
cleaned_path = Path(args.cleaned_path) if args.cleaned_path else None
|
||||||
|
train_path = Path(args.train_path)
|
||||||
|
val_path = Path(args.val_path)
|
||||||
|
config_path = Path(args.config_path)
|
||||||
|
val_per_lang = int(args.val_per_lang)
|
||||||
|
max_val_total = int(args.max_val_total)
|
||||||
|
use_jp_extra: bool = args.use_jp_extra
|
||||||
|
yomi_error: str = args.yomi_error
|
||||||
|
correct_path: bool = args.correct_path
|
||||||
|
|
||||||
|
preprocess(
|
||||||
|
transcription_path=transcription_path,
|
||||||
|
cleaned_path=cleaned_path,
|
||||||
|
train_path=train_path,
|
||||||
|
val_path=val_path,
|
||||||
|
config_path=config_path,
|
||||||
|
val_per_lang=val_per_lang,
|
||||||
|
max_val_total=max_val_total,
|
||||||
|
use_jp_extra=use_jp_extra,
|
||||||
|
yomi_error=yomi_error,
|
||||||
|
correct_path=correct_path,
|
||||||
|
)
|
||||||
|
|||||||
133
pyproject.toml
Normal file
133
pyproject.toml
Normal file
@@ -0,0 +1,133 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["hatchling"]
|
||||||
|
build-backend = "hatchling.build"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
name = "style-bert-vits2"
|
||||||
|
dynamic = ["version"]
|
||||||
|
description = 'Style-Bert-VITS2: Bert-VITS2 with more controllable voice styles.'
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.9"
|
||||||
|
license = "AGPL-3.0"
|
||||||
|
keywords = []
|
||||||
|
authors = [
|
||||||
|
{ name = "litagin02", email = "139731664+litagin02@users.noreply.github.com" },
|
||||||
|
]
|
||||||
|
classifiers = [
|
||||||
|
"Development Status :: 4 - Beta",
|
||||||
|
"Programming Language :: Python",
|
||||||
|
"Programming Language :: Python :: 3.9",
|
||||||
|
"Programming Language :: Python :: 3.10",
|
||||||
|
"Programming Language :: Python :: 3.11",
|
||||||
|
"Programming Language :: Python :: Implementation :: CPython",
|
||||||
|
]
|
||||||
|
dependencies = [
|
||||||
|
'cmudict',
|
||||||
|
'cn2an',
|
||||||
|
'g2p_en',
|
||||||
|
'gradio',
|
||||||
|
'jieba',
|
||||||
|
'librosa==0.9.2',
|
||||||
|
'loguru',
|
||||||
|
'num2words',
|
||||||
|
'numba',
|
||||||
|
'numpy',
|
||||||
|
'pyannote.audio>=3.1.0',
|
||||||
|
'pydantic>=2.0',
|
||||||
|
'pyopenjtalk-dict',
|
||||||
|
'pypinyin',
|
||||||
|
'pyworld-prebuilt',
|
||||||
|
'safetensors',
|
||||||
|
'scipy',
|
||||||
|
'torch>=2.1',
|
||||||
|
'transformers',
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
Documentation = "https://github.com/litagin02/Style-Bert-VITS2#readme"
|
||||||
|
Issues = "https://github.com/litagin02/Style-Bert-VITS2/issues"
|
||||||
|
Source = "https://github.com/litagin02/Style-Bert-VITS2"
|
||||||
|
|
||||||
|
[tool.hatch.version]
|
||||||
|
path = "style_bert_vits2/constants.py"
|
||||||
|
|
||||||
|
[tool.hatch.build.targets.sdist]
|
||||||
|
only-include = [
|
||||||
|
".vscode",
|
||||||
|
"dict_data/default.csv",
|
||||||
|
"docs",
|
||||||
|
"style_bert_vits2",
|
||||||
|
"tests",
|
||||||
|
"LGPL_LICENSE",
|
||||||
|
"LICENSE",
|
||||||
|
"pyproject.toml",
|
||||||
|
"README.md",
|
||||||
|
]
|
||||||
|
exclude = [
|
||||||
|
".git",
|
||||||
|
".gitignore",
|
||||||
|
".gitattributes",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.hatch.build.targets.wheel]
|
||||||
|
packages = ["style_bert_vits2"]
|
||||||
|
|
||||||
|
[tool.hatch.envs.test]
|
||||||
|
dependencies = [
|
||||||
|
"coverage[toml]>=6.5",
|
||||||
|
"pytest",
|
||||||
|
]
|
||||||
|
[tool.hatch.envs.test.scripts]
|
||||||
|
# Usage: `hatch run test:test`
|
||||||
|
test = "pytest {args:tests}"
|
||||||
|
# Usage: `hatch run test:coverage`
|
||||||
|
test-cov = "coverage run -m pytest {args:tests}"
|
||||||
|
# Usage: `hatch run test:cov-report`
|
||||||
|
cov-report = [
|
||||||
|
"- coverage combine",
|
||||||
|
"coverage report",
|
||||||
|
]
|
||||||
|
# Usage: `hatch run test:cov`
|
||||||
|
cov = [
|
||||||
|
"test-cov",
|
||||||
|
"cov-report",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.hatch.envs.style]
|
||||||
|
detached = true
|
||||||
|
dependencies = [
|
||||||
|
"black",
|
||||||
|
"isort",
|
||||||
|
]
|
||||||
|
[tool.hatch.envs.style.scripts]
|
||||||
|
check = [
|
||||||
|
"black --check --diff .",
|
||||||
|
"isort --check-only --diff --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
|
]
|
||||||
|
fmt = [
|
||||||
|
"black .",
|
||||||
|
"isort --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
|
"check",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[tool.hatch.envs.test.matrix]]
|
||||||
|
python = ["3.9", "3.10", "3.11"]
|
||||||
|
|
||||||
|
[tool.coverage.run]
|
||||||
|
source_pkgs = ["style_bert_vits2", "tests"]
|
||||||
|
branch = true
|
||||||
|
parallel = true
|
||||||
|
omit = [
|
||||||
|
"style_bert_vits2/constants.py",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.coverage.paths]
|
||||||
|
style_bert_vits2 = ["style_bert_vits2", "*/style-bert-vits2/style_bert_vits2"]
|
||||||
|
tests = ["tests", "*/style-bert-vits2/tests"]
|
||||||
|
|
||||||
|
[tool.coverage.report]
|
||||||
|
exclude_lines = [
|
||||||
|
"no cov",
|
||||||
|
"if __name__ == .__main__.:",
|
||||||
|
"if TYPE_CHECKING:",
|
||||||
|
]
|
||||||
@@ -1,81 +0,0 @@
|
|||||||
import re
|
|
||||||
|
|
||||||
|
|
||||||
def extract_language_and_text_updated(speaker, dialogue):
|
|
||||||
# 使用正则表达式匹配<语言>标签和其后的文本
|
|
||||||
pattern_language_text = r"<(\S+?)>([^<]+)"
|
|
||||||
matches = re.findall(pattern_language_text, dialogue, re.DOTALL)
|
|
||||||
speaker = speaker[1:-1]
|
|
||||||
# 清理文本:去除两边的空白字符
|
|
||||||
matches_cleaned = [(lang.upper(), text.strip()) for lang, text in matches]
|
|
||||||
matches_cleaned.append(speaker)
|
|
||||||
return matches_cleaned
|
|
||||||
|
|
||||||
|
|
||||||
def validate_text(input_text):
|
|
||||||
# 验证说话人的正则表达式
|
|
||||||
pattern_speaker = r"(\[\S+?\])((?:\s*<\S+?>[^<\[\]]+?)+)"
|
|
||||||
|
|
||||||
# 使用re.DOTALL标志使.匹配包括换行符在内的所有字符
|
|
||||||
matches = re.findall(pattern_speaker, input_text, re.DOTALL)
|
|
||||||
|
|
||||||
# 对每个匹配到的说话人内容进行进一步验证
|
|
||||||
for _, dialogue in matches:
|
|
||||||
language_text_matches = extract_language_and_text_updated(_, dialogue)
|
|
||||||
if not language_text_matches:
|
|
||||||
return (
|
|
||||||
False,
|
|
||||||
"Error: Invalid format detected in dialogue content. Please check your input.",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果输入的文本中没有找到任何匹配项
|
|
||||||
if not matches:
|
|
||||||
return (
|
|
||||||
False,
|
|
||||||
"Error: No valid speaker format detected. Please check your input.",
|
|
||||||
)
|
|
||||||
|
|
||||||
return True, "Input is valid."
|
|
||||||
|
|
||||||
|
|
||||||
def text_matching(text: str) -> list:
|
|
||||||
speaker_pattern = r"(\[\S+?\])(.+?)(?=\[\S+?\]|$)"
|
|
||||||
matches = re.findall(speaker_pattern, text, re.DOTALL)
|
|
||||||
result = []
|
|
||||||
for speaker, dialogue in matches:
|
|
||||||
result.append(extract_language_and_text_updated(speaker, dialogue))
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def cut_para(text):
|
|
||||||
splitted_para = re.split("[\n]", text) # 按段分
|
|
||||||
splitted_para = [
|
|
||||||
sentence.strip() for sentence in splitted_para if sentence.strip()
|
|
||||||
] # 删除空字符串
|
|
||||||
return splitted_para
|
|
||||||
|
|
||||||
|
|
||||||
def cut_sent(para):
|
|
||||||
para = re.sub("([。!;?\?])([^”’])", r"\1\n\2", para) # 单字符断句符
|
|
||||||
para = re.sub("(\.{6})([^”’])", r"\1\n\2", para) # 英文省略号
|
|
||||||
para = re.sub("(\…{2})([^”’])", r"\1\n\2", para) # 中文省略号
|
|
||||||
para = re.sub("([。!?\?][”’])([^,。!?\?])", r"\1\n\2", para)
|
|
||||||
para = para.rstrip() # 段尾如果有多余的\n就去掉它
|
|
||||||
return para.split("\n")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
text = """
|
|
||||||
[说话人1]
|
|
||||||
[说话人2]<zh>你好吗?<jp>元気ですか?<jp>こんにちは,世界。<zh>你好吗?
|
|
||||||
[说话人3]<zh>谢谢。<jp>どういたしまして。
|
|
||||||
"""
|
|
||||||
text_matching(text)
|
|
||||||
# 测试函数
|
|
||||||
test_text = """
|
|
||||||
[说话人1]<zh>你好,こんにちは!<jp>こんにちは,世界。
|
|
||||||
[说话人2]<zh>你好吗?
|
|
||||||
"""
|
|
||||||
text_matching(test_text)
|
|
||||||
res = validate_text(test_text)
|
|
||||||
print(res)
|
|
||||||
@@ -14,11 +14,12 @@ numba
|
|||||||
numpy
|
numpy
|
||||||
psutil
|
psutil
|
||||||
pyannote.audio>=3.1.0
|
pyannote.audio>=3.1.0
|
||||||
|
pydantic>=2.0
|
||||||
pyloudnorm
|
pyloudnorm
|
||||||
# pyopenjtalk-prebuilt # Should be manually uninstalled
|
# pyopenjtalk-prebuilt # Should be manually uninstalled
|
||||||
pyopenjtalk-dict
|
pyopenjtalk-dict
|
||||||
pypinyin
|
pypinyin
|
||||||
# pyworld # Not supported on Windows without Cython...
|
pyworld-prebuilt
|
||||||
PyYAML
|
PyYAML
|
||||||
requests
|
requests
|
||||||
safetensors
|
safetensors
|
||||||
|
|||||||
100
resample.py
100
resample.py
@@ -1,15 +1,19 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import os
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from multiprocessing import cpu_count
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import librosa
|
import librosa
|
||||||
import pyloudnorm as pyln
|
import pyloudnorm as pyln
|
||||||
import soundfile
|
import soundfile
|
||||||
|
from numpy.typing import NDArray
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from common.log import logger
|
|
||||||
from common.stdout_wrapper import SAFE_STDOUT
|
|
||||||
from config import config
|
from config import config
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_BLOCK_SIZE: float = 0.400 # seconds
|
DEFAULT_BLOCK_SIZE: float = 0.400 # seconds
|
||||||
|
|
||||||
@@ -18,32 +22,40 @@ class BlockSizeException(Exception):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
def normalize_audio(data, sr):
|
def normalize_audio(data: NDArray[Any], sr: int):
|
||||||
meter = pyln.Meter(sr, block_size=DEFAULT_BLOCK_SIZE) # create BS.1770 meter
|
meter = pyln.Meter(sr, block_size=DEFAULT_BLOCK_SIZE) # create BS.1770 meter
|
||||||
try:
|
try:
|
||||||
loudness = meter.integrated_loudness(data)
|
loudness = meter.integrated_loudness(data)
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
raise BlockSizeException(e)
|
raise BlockSizeException(e)
|
||||||
# logger.info(f"loudness: {loudness}")
|
|
||||||
data = pyln.normalize.loudness(data, loudness, -23.0)
|
data = pyln.normalize.loudness(data, loudness, -23.0)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
def process(item):
|
def resample(file: Path, output_dir: Path, target_sr: int, normalize: bool, trim: bool):
|
||||||
spkdir, wav_name, args = item
|
"""
|
||||||
wav_path = os.path.join(args.in_dir, spkdir, wav_name)
|
fileを読み込んで、target_srなwavファイルに変換してoutput_dir直下に保存する
|
||||||
if os.path.exists(wav_path) and wav_path.lower().endswith(".wav"):
|
"""
|
||||||
wav, sr = librosa.load(wav_path, sr=args.sr)
|
try:
|
||||||
if args.normalize:
|
# librosaが読めるファイルかチェック
|
||||||
|
# wav以外にもmp3やoggやflacなども読める
|
||||||
|
wav: NDArray[Any]
|
||||||
|
sr: int
|
||||||
|
wav, sr = librosa.load(file, sr=target_sr)
|
||||||
|
if normalize:
|
||||||
try:
|
try:
|
||||||
wav = normalize_audio(wav, sr)
|
wav = normalize_audio(wav, sr)
|
||||||
except BlockSizeException:
|
except BlockSizeException:
|
||||||
|
print("")
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Skip normalize due to less than {DEFAULT_BLOCK_SIZE} second audio: {wav_path}"
|
f"Skip normalize due to less than {DEFAULT_BLOCK_SIZE} second audio: {file}"
|
||||||
)
|
)
|
||||||
if args.trim:
|
if trim:
|
||||||
wav, _ = librosa.effects.trim(wav, top_db=30)
|
wav, _ = librosa.effects.trim(wav, top_db=30)
|
||||||
soundfile.write(os.path.join(args.out_dir, spkdir, wav_name), wav, sr)
|
soundfile.write(output_dir / file.with_suffix(".wav").name, wav, sr)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Cannot load file, so skipping: {file}, {e}")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
@@ -55,14 +67,14 @@ if __name__ == "__main__":
|
|||||||
help="sampling rate",
|
help="sampling rate",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--in_dir",
|
"--input_dir",
|
||||||
"-i",
|
"-i",
|
||||||
type=str,
|
type=str,
|
||||||
default=config.resample_config.in_dir,
|
default=config.resample_config.in_dir,
|
||||||
help="path to source dir",
|
help="path to source dir",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--out_dir",
|
"--output_dir",
|
||||||
"-o",
|
"-o",
|
||||||
type=str,
|
type=str,
|
||||||
default=config.resample_config.out_dir,
|
default=config.resample_config.out_dir,
|
||||||
@@ -86,46 +98,36 @@ if __name__ == "__main__":
|
|||||||
default=False,
|
default=False,
|
||||||
help="trim silence (start and end only)",
|
help="trim silence (start and end only)",
|
||||||
)
|
)
|
||||||
args, _ = parser.parse_known_args()
|
args = parser.parse_args()
|
||||||
# autodl 无卡模式会识别出46个cpu
|
|
||||||
if args.num_processes == 0:
|
if args.num_processes == 0:
|
||||||
processes = cpu_count() - 2 if cpu_count() > 4 else 1
|
processes = cpu_count() - 2 if cpu_count() > 4 else 1
|
||||||
else:
|
else:
|
||||||
processes = args.num_processes
|
processes: int = args.num_processes
|
||||||
|
|
||||||
tasks = []
|
input_dir = Path(args.input_dir)
|
||||||
|
output_dir = Path(args.output_dir)
|
||||||
|
sr = int(args.sr)
|
||||||
|
normalize: bool = args.normalize
|
||||||
|
trim: bool = args.trim
|
||||||
|
|
||||||
for dirpath, _, filenames in os.walk(args.in_dir):
|
# 後でlibrosaに読ませて有効な音声ファイルかチェックするので、全てのファイルを取得
|
||||||
# 子级目录
|
original_files = [f for f in input_dir.rglob("*") if f.is_file()]
|
||||||
spk_dir = os.path.relpath(dirpath, args.in_dir)
|
|
||||||
spk_dir_out = os.path.join(args.out_dir, spk_dir)
|
|
||||||
if not os.path.isdir(spk_dir_out):
|
|
||||||
os.makedirs(spk_dir_out, exist_ok=True)
|
|
||||||
for filename in filenames:
|
|
||||||
if filename.lower().endswith(".wav"):
|
|
||||||
twople = (spk_dir, filename, args)
|
|
||||||
tasks.append(twople)
|
|
||||||
|
|
||||||
if len(tasks) == 0:
|
if len(original_files) == 0:
|
||||||
logger.error(f"No wav files found in {args.in_dir}")
|
logger.error(f"No files found in {input_dir}")
|
||||||
raise ValueError(f"No wav files found in {args.in_dir}")
|
raise ValueError(f"No files found in {input_dir}")
|
||||||
|
|
||||||
# pool = Pool(processes=processes)
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
# for _ in tqdm(
|
|
||||||
# pool.imap_unordered(process, tasks), file=SAFE_STDOUT, total=len(tasks)
|
|
||||||
# ):
|
|
||||||
# pass
|
|
||||||
|
|
||||||
# pool.close()
|
|
||||||
# pool.join()
|
|
||||||
|
|
||||||
with ThreadPoolExecutor(max_workers=processes) as executor:
|
with ThreadPoolExecutor(max_workers=processes) as executor:
|
||||||
_ = list(
|
futures = [
|
||||||
tqdm(
|
executor.submit(resample, file, output_dir, sr, normalize, trim)
|
||||||
executor.map(process, tasks),
|
for file in original_files
|
||||||
total=len(tasks),
|
]
|
||||||
file=SAFE_STDOUT,
|
for future in tqdm(
|
||||||
)
|
as_completed(futures), total=len(original_files), file=SAFE_STDOUT
|
||||||
)
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
logger.info("Resampling Done!")
|
logger.info("Resampling Done!")
|
||||||
|
|||||||
66
scripts/Clean.bat
Normal file
66
scripts/Clean.bat
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
chcp 65001 > NUL
|
||||||
|
@echo off
|
||||||
|
setlocal
|
||||||
|
echo 不要になった以下のフォルダ・ファイルを削除します:
|
||||||
|
echo 注: 学習やマージ等はApp.batへ統合されました。
|
||||||
|
echo Style-Bert-VITS2\common\
|
||||||
|
echo Style-Bert-VITS2\monotonic_align\
|
||||||
|
echo Style-Bert-VITS2\text\
|
||||||
|
echo Style-Bert-VITS2\tools\
|
||||||
|
echo Style-Bert-VITS2\attentions.py
|
||||||
|
echo Style-Bert-VITS2\commons.py
|
||||||
|
echo Style-Bert-VITS2\Dataset.bat
|
||||||
|
echo Style-Bert-VITS2\infer.py
|
||||||
|
echo Style-Bert-VITS2\Merge.bat
|
||||||
|
echo Style-Bert-VITS2\models_jp_extra.py
|
||||||
|
echo Style-Bert-VITS2\models.py
|
||||||
|
echo Style-Bert-VITS2\modules.py
|
||||||
|
echo Style-Bert-VITS2\re_matching.py
|
||||||
|
echo Style-Bert-VITS2\spec_gen.py
|
||||||
|
echo Style-Bert-VITS2\Style.bat
|
||||||
|
echo Style-Bert-VITS2\Train.bat
|
||||||
|
echo Style-Bert-VITS2\transforms.py
|
||||||
|
echo Style-Bert-VITS2\update_status.py
|
||||||
|
echo Style-Bert-VITS2\utils.py
|
||||||
|
echo Style-Bert-VITS2\webui_dataset.py
|
||||||
|
echo Style-Bert-VITS2\webui_merge.py
|
||||||
|
echo Style-Bert-VITS2\webui_style_vectors.py
|
||||||
|
echo Style-Bert-VITS2\webui_train.py
|
||||||
|
echo Style-Bert-VITS2\webui.py
|
||||||
|
echo.
|
||||||
|
set /p delConfirm=以上のフォルダファイルを削除しますか? (y/n):
|
||||||
|
if /I "%delConfirm%"=="Y" goto proceed
|
||||||
|
if /I "%delConfirm%"=="y" goto proceed
|
||||||
|
if "%delConfirm%"=="" goto proceed
|
||||||
|
goto end
|
||||||
|
|
||||||
|
:proceed
|
||||||
|
rd /s /q "Style-Bert-VITS2\common"
|
||||||
|
rd /s /q "Style-Bert-VITS2\monotonic_align"
|
||||||
|
rd /s /q "Style-Bert-VITS2\text"
|
||||||
|
rd /s /q "Style-Bert-VITS2\tools"
|
||||||
|
del /q "Style-Bert-VITS2\attentions.py"
|
||||||
|
del /q "Style-Bert-VITS2\commons.py"
|
||||||
|
del /q "Style-Bert-VITS2\Dataset.bat"
|
||||||
|
del /q "Style-Bert-VITS2\infer.py"
|
||||||
|
del /q "Style-Bert-VITS2\Merge.bat"
|
||||||
|
del /q "Style-Bert-VITS2\models_jp_extra.py"
|
||||||
|
del /q "Style-Bert-VITS2\models.py"
|
||||||
|
del /q "Style-Bert-VITS2\modules.py"
|
||||||
|
del /q "Style-Bert-VITS2\re_matching.py"
|
||||||
|
del /q "Style-Bert-VITS2\spec_gen.py"
|
||||||
|
del /q "Style-Bert-VITS2\Style.bat"
|
||||||
|
del /q "Style-Bert-VITS2\Train.bat"
|
||||||
|
del /q "Style-Bert-VITS2\transforms.py"
|
||||||
|
del /q "Style-Bert-VITS2\update_status.py"
|
||||||
|
del /q "Style-Bert-VITS2\utils.py"
|
||||||
|
del /q "Style-Bert-VITS2\webui_dataset.py"
|
||||||
|
del /q "Style-Bert-VITS2\webui_merge.py"
|
||||||
|
del /q "Style-Bert-VITS2\webui_style_vectors.py"
|
||||||
|
del /q "Style-Bert-VITS2\webui_train.py"
|
||||||
|
del /q "Style-Bert-VITS2\webui.py"
|
||||||
|
echo 完了しました。
|
||||||
|
pause
|
||||||
|
|
||||||
|
:end
|
||||||
|
endlocal
|
||||||
@@ -41,7 +41,8 @@ call Style-Bert-VITS2\scripts\Setup-Python.bat ..\..\lib\python ..\venv
|
|||||||
if %errorlevel% neq 0 ( popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
@REM 依存関係インストール
|
@REM 依存関係インストール
|
||||||
pip install torch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 --index-url https://download.pytorch.org/whl/cu118
|
pip3 install torch torchvision torchaudio --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% )
|
||||||
|
|
||||||
pip install -r Style-Bert-VITS2\requirements.txt
|
pip install -r Style-Bert-VITS2\requirements.txt
|
||||||
|
|||||||
@@ -37,11 +37,18 @@ if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
|||||||
pip install -U -r Style-Bert-VITS2\requirements.txt
|
pip install -U -r Style-Bert-VITS2\requirements.txt
|
||||||
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% )
|
||||||
|
|
||||||
|
echo ----------------------------------------
|
||||||
echo Update completed. Running Style-Bert-VITS2 Editor...
|
echo Update completed. Running Style-Bert-VITS2 Editor...
|
||||||
|
echo ----------------------------------------
|
||||||
|
|
||||||
|
@REM Style-Bert-VITS2フォルダに移動
|
||||||
|
pushd Style-Bert-VITS2
|
||||||
|
|
||||||
@REM Style-Bert-VITS2 Editorを起動
|
@REM Style-Bert-VITS2 Editorを起動
|
||||||
python server_editor.py --inbrowser
|
python server_editor.py --inbrowser
|
||||||
|
|
||||||
pause
|
pause
|
||||||
|
|
||||||
popd
|
popd
|
||||||
|
|
||||||
|
popd
|
||||||
|
|||||||
@@ -16,13 +16,13 @@ import zipfile
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import yaml
|
from typing import Optional
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pyopenjtalk
|
|
||||||
import requests
|
import requests
|
||||||
import torch
|
import torch
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
import yaml
|
||||||
from fastapi import APIRouter, FastAPI, HTTPException, status
|
from fastapi import APIRouter, FastAPI, HTTPException, status
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import JSONResponse, Response
|
from fastapi.responses import JSONResponse, Response
|
||||||
@@ -30,20 +30,29 @@ from fastapi.staticfiles import StaticFiles
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from scipy.io import wavfile
|
from scipy.io import wavfile
|
||||||
|
|
||||||
from common.constants import (
|
from style_bert_vits2.constants import (
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
DEFAULT_NOISE,
|
DEFAULT_NOISE,
|
||||||
DEFAULT_NOISEW,
|
DEFAULT_NOISEW,
|
||||||
DEFAULT_SDP_RATIO,
|
DEFAULT_SDP_RATIO,
|
||||||
DEFAULT_STYLE,
|
DEFAULT_STYLE,
|
||||||
DEFAULT_STYLE_WEIGHT,
|
DEFAULT_STYLE_WEIGHT,
|
||||||
LATEST_VERSION,
|
VERSION,
|
||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from common.log import logger
|
from style_bert_vits2.logging import logger
|
||||||
from common.tts_model import ModelHolder
|
from style_bert_vits2.nlp import bert_models
|
||||||
from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from text.user_dict import apply_word, update_dict, read_dict, rewrite_word, delete_word
|
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.user_dict import (
|
||||||
|
apply_word,
|
||||||
|
delete_word,
|
||||||
|
read_dict,
|
||||||
|
rewrite_word,
|
||||||
|
update_dict,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo
|
||||||
|
|
||||||
|
|
||||||
# ---フロントエンド部分に関する処理---
|
# ---フロントエンド部分に関する処理---
|
||||||
@@ -140,6 +149,19 @@ def save_last_download(latest_release):
|
|||||||
# ---フロントエンド部分に関する処理ここまで---
|
# ---フロントエンド部分に関する処理ここまで---
|
||||||
# 以降はAPIの設定
|
# 以降はAPIの設定
|
||||||
|
|
||||||
|
# pyopenjtalk_worker を起動
|
||||||
|
## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する
|
||||||
|
pyopenjtalk.initialize_worker()
|
||||||
|
|
||||||
|
# pyopenjtalk の辞書を更新
|
||||||
|
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"
|
||||||
@@ -176,7 +198,7 @@ if device == "cuda" and not torch.cuda.is_available():
|
|||||||
model_dir = Path(args.model_dir)
|
model_dir = Path(args.model_dir)
|
||||||
port = int(args.port)
|
port = int(args.port)
|
||||||
|
|
||||||
model_holder = ModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, 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)
|
||||||
@@ -197,7 +219,7 @@ router = APIRouter()
|
|||||||
|
|
||||||
@router.get("/version")
|
@router.get("/version")
|
||||||
def version() -> str:
|
def version() -> str:
|
||||||
return LATEST_VERSION
|
return VERSION
|
||||||
|
|
||||||
|
|
||||||
class MoraTone(BaseModel):
|
class MoraTone(BaseModel):
|
||||||
@@ -213,7 +235,7 @@ class TextRequest(BaseModel):
|
|||||||
async def read_item(item: TextRequest):
|
async def read_item(item: TextRequest):
|
||||||
try:
|
try:
|
||||||
# 最初に正規化しないと整合性がとれない
|
# 最初に正規化しないと整合性がとれない
|
||||||
text = text_normalize(item.text)
|
text = normalize_text(item.text)
|
||||||
kata_tone_list = g2kata_tone(text)
|
kata_tone_list = g2kata_tone(text)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -224,13 +246,13 @@ async def read_item(item: TextRequest):
|
|||||||
|
|
||||||
|
|
||||||
@router.post("/normalize")
|
@router.post("/normalize")
|
||||||
async def normalize_text(item: TextRequest):
|
async def normalize(item: TextRequest):
|
||||||
return text_normalize(item.text)
|
return normalize_text(item.text)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/models_info")
|
@router.get("/models_info", response_model=list[TTSModelInfo])
|
||||||
def models_info():
|
def models_info():
|
||||||
return model_holder.models_info()
|
return model_holder.models_info
|
||||||
|
|
||||||
|
|
||||||
class SynthesisRequest(BaseModel):
|
class SynthesisRequest(BaseModel):
|
||||||
@@ -250,6 +272,7 @@ class SynthesisRequest(BaseModel):
|
|||||||
silenceAfter: float = 0.5
|
silenceAfter: float = 0.5
|
||||||
pitchScale: float = 1.0
|
pitchScale: float = 1.0
|
||||||
intonationScale: float = 1.0
|
intonationScale: float = 1.0
|
||||||
|
speaker: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@router.post("/synthesis", response_class=AudioResponse)
|
@router.post("/synthesis", response_class=AudioResponse)
|
||||||
@@ -260,7 +283,7 @@ def synthesis(request: SynthesisRequest):
|
|||||||
detail=f"1行の文字数は{args.line_length}文字以下にしてください。",
|
detail=f"1行の文字数は{args.line_length}文字以下にしてください。",
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
model = model_holder.load_model(
|
model = model_holder.get_model(
|
||||||
model_name=request.model, model_path_str=request.modelFile
|
model_name=request.model, model_path_str=request.modelFile
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -275,12 +298,19 @@ def synthesis(request: SynthesisRequest):
|
|||||||
]
|
]
|
||||||
phone_tone = kata_tone2phone_tone(kata_tone_list)
|
phone_tone = kata_tone2phone_tone(kata_tone_list)
|
||||||
tone = [t for _, t in phone_tone]
|
tone = [t for _, t in phone_tone]
|
||||||
|
try:
|
||||||
|
sid = 0 if request.speaker is None else model.spk2id[request.speaker]
|
||||||
|
except KeyError:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail=f"Speaker {request.speaker} not found in {model.spk2id}",
|
||||||
|
)
|
||||||
sr, audio = model.infer(
|
sr, audio = model.infer(
|
||||||
text=text,
|
text=text,
|
||||||
language=request.language.value,
|
language=request.language,
|
||||||
sdp_ratio=request.sdpRatio,
|
sdp_ratio=request.sdpRatio,
|
||||||
noise=request.noise,
|
noise=request.noise,
|
||||||
noisew=request.noisew,
|
noise_w=request.noisew,
|
||||||
length=1 / request.speed,
|
length=1 / request.speed,
|
||||||
given_tone=tone,
|
given_tone=tone,
|
||||||
style=request.style,
|
style=request.style,
|
||||||
@@ -291,6 +321,7 @@ def synthesis(request: SynthesisRequest):
|
|||||||
line_split=False,
|
line_split=False,
|
||||||
pitch_scale=request.pitchScale,
|
pitch_scale=request.pitchScale,
|
||||||
intonation_scale=request.intonationScale,
|
intonation_scale=request.intonationScale,
|
||||||
|
speaker_id=sid,
|
||||||
)
|
)
|
||||||
|
|
||||||
with BytesIO() as wavContent:
|
with BytesIO() as wavContent:
|
||||||
@@ -311,6 +342,7 @@ def multi_synthesis(request: MultiSynthesisRequest):
|
|||||||
detail=f"行数は{args.line_count}行以下にしてください。",
|
detail=f"行数は{args.line_count}行以下にしてください。",
|
||||||
)
|
)
|
||||||
audios = []
|
audios = []
|
||||||
|
sr = None
|
||||||
for i, req in enumerate(lines):
|
for i, req in enumerate(lines):
|
||||||
if args.line_length is not None and len(req.text) > args.line_length:
|
if args.line_length is not None and len(req.text) > args.line_length:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -318,7 +350,7 @@ def multi_synthesis(request: MultiSynthesisRequest):
|
|||||||
detail=f"1行の文字数は{args.line_length}文字以下にしてください。",
|
detail=f"1行の文字数は{args.line_length}文字以下にしてください。",
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
model = model_holder.load_model(
|
model = model_holder.get_model(
|
||||||
model_name=req.model, model_path_str=req.modelFile
|
model_name=req.model, model_path_str=req.modelFile
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -335,10 +367,10 @@ def multi_synthesis(request: MultiSynthesisRequest):
|
|||||||
tone = [t for _, t in phone_tone]
|
tone = [t for _, t in phone_tone]
|
||||||
sr, audio = model.infer(
|
sr, audio = model.infer(
|
||||||
text=text,
|
text=text,
|
||||||
language=req.language.value,
|
language=req.language,
|
||||||
sdp_ratio=req.sdpRatio,
|
sdp_ratio=req.sdpRatio,
|
||||||
noise=req.noise,
|
noise=req.noise,
|
||||||
noisew=req.noisew,
|
noise_w=req.noisew,
|
||||||
length=1 / req.speed,
|
length=1 / req.speed,
|
||||||
given_tone=tone,
|
given_tone=tone,
|
||||||
style=req.style,
|
style=req.style,
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import os
|
|||||||
import sys
|
import sys
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, Optional, Union
|
from typing import Any, Optional
|
||||||
from urllib.parse import unquote
|
from urllib.parse import unquote
|
||||||
|
|
||||||
import GPUtil
|
import GPUtil
|
||||||
@@ -20,7 +20,8 @@ from fastapi.middleware.cors import CORSMiddleware
|
|||||||
from fastapi.responses import FileResponse, Response
|
from fastapi.responses import FileResponse, Response
|
||||||
from scipy.io import wavfile
|
from scipy.io import wavfile
|
||||||
|
|
||||||
from common.constants import (
|
from config import config
|
||||||
|
from style_bert_vits2.constants import (
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
DEFAULT_LENGTH,
|
DEFAULT_LENGTH,
|
||||||
DEFAULT_LINE_SPLIT,
|
DEFAULT_LINE_SPLIT,
|
||||||
@@ -32,13 +33,33 @@ from common.constants import (
|
|||||||
DEFAULT_STYLE_WEIGHT,
|
DEFAULT_STYLE_WEIGHT,
|
||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from common.log import logger
|
from style_bert_vits2.logging import logger
|
||||||
from common.tts_model import Model, ModelHolder
|
from style_bert_vits2.nlp import bert_models
|
||||||
from config import config
|
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.tts_model import TTSModel, TTSModelHolder
|
||||||
|
|
||||||
|
|
||||||
ln = config.server_config.language
|
ln = config.server_config.language
|
||||||
|
|
||||||
|
|
||||||
|
# pyopenjtalk_worker を起動
|
||||||
|
## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する
|
||||||
|
pyopenjtalk.initialize_worker()
|
||||||
|
|
||||||
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
|
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}")
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -51,17 +72,22 @@ class AudioResponse(Response):
|
|||||||
media_type = "audio/wav"
|
media_type = "audio/wav"
|
||||||
|
|
||||||
|
|
||||||
def load_models(model_holder: ModelHolder):
|
loaded_models: list[TTSModel] = []
|
||||||
model_holder.models = []
|
|
||||||
|
|
||||||
|
def load_models(model_holder: TTSModelHolder):
|
||||||
|
global loaded_models
|
||||||
|
loaded_models = []
|
||||||
for model_name, model_paths in model_holder.model_files_dict.items():
|
for model_name, model_paths in model_holder.model_files_dict.items():
|
||||||
model = Model(
|
model = TTSModel(
|
||||||
model_path=model_paths[0],
|
model_path=model_paths[0],
|
||||||
config_path=model_holder.root_dir / model_name / "config.json",
|
config_path=model_holder.root_dir / model_name / "config.json",
|
||||||
style_vec_path=model_holder.root_dir / model_name / "style_vectors.npy",
|
style_vec_path=model_holder.root_dir / model_name / "style_vectors.npy",
|
||||||
device=model_holder.device,
|
device=model_holder.device,
|
||||||
)
|
)
|
||||||
model.load_net_g()
|
# 起動時に全てのモデルを読み込むのは時間がかかりメモリを食うのでやめる
|
||||||
model_holder.models.append(model)
|
# model.load()
|
||||||
|
loaded_models.append(model)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
@@ -78,13 +104,14 @@ if __name__ == "__main__":
|
|||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
model_dir = Path(args.dir)
|
model_dir = Path(args.dir)
|
||||||
model_holder = ModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, 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)
|
||||||
|
|
||||||
logger.info("Loading models...")
|
logger.info("Loading models...")
|
||||||
load_models(model_holder)
|
load_models(model_holder)
|
||||||
|
|
||||||
limit = config.server_config.limit
|
limit = config.server_config.limit
|
||||||
app = FastAPI()
|
app = FastAPI()
|
||||||
allow_origins = config.server_config.origins
|
allow_origins = config.server_config.origins
|
||||||
@@ -99,12 +126,13 @@ if __name__ == "__main__":
|
|||||||
allow_methods=["*"],
|
allow_methods=["*"],
|
||||||
allow_headers=["*"],
|
allow_headers=["*"],
|
||||||
)
|
)
|
||||||
app.logger = logger
|
# app.logger = logger
|
||||||
|
# ↑効いていなさそう。loggerをどうやって上書きするかはよく分からなかった。
|
||||||
|
|
||||||
@app.api_route("/voice", methods=["GET", "POST"], response_class=AudioResponse)
|
@app.api_route("/voice", methods=["GET", "POST"], response_class=AudioResponse)
|
||||||
async def voice(
|
async def voice(
|
||||||
request: Request,
|
request: Request,
|
||||||
text: str = Query(..., min_length=1, max_length=limit, description=f"セリフ"),
|
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_id: int = Query(
|
model_id: int = Query(
|
||||||
0, description="モデルID。`GET /models/info`のkeyの値を指定ください"
|
0, description="モデルID。`GET /models/info`のkeyの値を指定ください"
|
||||||
@@ -132,7 +160,7 @@ if __name__ == "__main__":
|
|||||||
DEFAULT_LENGTH,
|
DEFAULT_LENGTH,
|
||||||
description="話速。基準は1で大きくするほど音声は長くなり読み上げが遅まる",
|
description="話速。基準は1で大きくするほど音声は長くなり読み上げが遅まる",
|
||||||
),
|
),
|
||||||
language: Languages = Query(ln, description=f"textの言語"),
|
language: Languages = Query(ln, description="textの言語"),
|
||||||
auto_split: bool = Query(DEFAULT_LINE_SPLIT, description="改行で分けて生成"),
|
auto_split: bool = Query(DEFAULT_LINE_SPLIT, description="改行で分けて生成"),
|
||||||
split_interval: float = Query(
|
split_interval: float = Query(
|
||||||
DEFAULT_SPLIT_INTERVAL, description="分けた場合に挟む無音の長さ(秒)"
|
DEFAULT_SPLIT_INTERVAL, description="分けた場合に挟む無音の長さ(秒)"
|
||||||
@@ -144,7 +172,7 @@ if __name__ == "__main__":
|
|||||||
assist_text_weight: float = Query(
|
assist_text_weight: float = Query(
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT, description="assist_textの強さ"
|
DEFAULT_ASSIST_TEXT_WEIGHT, description="assist_textの強さ"
|
||||||
),
|
),
|
||||||
style: Optional[Union[int, str]] = Query(DEFAULT_STYLE, description="スタイル"),
|
style: Optional[str] = Query(DEFAULT_STYLE, description="スタイル"),
|
||||||
style_weight: float = Query(DEFAULT_STYLE_WEIGHT, description="スタイルの強さ"),
|
style_weight: float = Query(DEFAULT_STYLE_WEIGHT, description="スタイルの強さ"),
|
||||||
reference_audio_path: Optional[str] = Query(
|
reference_audio_path: Optional[str] = Query(
|
||||||
None, description="スタイルを音声ファイルで行う"
|
None, description="スタイルを音声ファイルで行う"
|
||||||
@@ -159,11 +187,11 @@ if __name__ == "__main__":
|
|||||||
"The GET method is not recommended for this endpoint due to various restrictions. Please use the POST method."
|
"The GET method is not recommended for this endpoint due to various restrictions. Please use the POST method."
|
||||||
)
|
)
|
||||||
if model_id >= len(
|
if model_id >= len(
|
||||||
model_holder.models
|
model_holder.model_names
|
||||||
): # /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")
|
||||||
|
|
||||||
model = model_holder.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():
|
||||||
raise_validation_error(
|
raise_validation_error(
|
||||||
@@ -177,16 +205,17 @@ if __name__ == "__main__":
|
|||||||
speaker_id = model.spk2id[speaker_name]
|
speaker_id = model.spk2id[speaker_name]
|
||||||
if style not in model.style2id.keys():
|
if style not in model.style2id.keys():
|
||||||
raise_validation_error(f"style={style} not found", "style")
|
raise_validation_error(f"style={style} not found", "style")
|
||||||
|
assert style is not None
|
||||||
if encoding is not None:
|
if encoding is not None:
|
||||||
text = unquote(text, encoding=encoding)
|
text = unquote(text, encoding=encoding)
|
||||||
sr, audio = model.infer(
|
sr, audio = model.infer(
|
||||||
text=text,
|
text=text,
|
||||||
language=language,
|
language=language,
|
||||||
sid=speaker_id,
|
speaker_id=speaker_id,
|
||||||
reference_audio_path=reference_audio_path,
|
reference_audio_path=reference_audio_path,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise=noise,
|
noise=noise,
|
||||||
noisew=noisew,
|
noise_w=noisew,
|
||||||
length=length,
|
length=length,
|
||||||
line_split=auto_split,
|
line_split=auto_split,
|
||||||
split_interval=split_interval,
|
split_interval=split_interval,
|
||||||
@@ -205,8 +234,8 @@ if __name__ == "__main__":
|
|||||||
def get_loaded_models_info():
|
def get_loaded_models_info():
|
||||||
"""ロードされたモデル情報の取得"""
|
"""ロードされたモデル情報の取得"""
|
||||||
|
|
||||||
result: Dict[str, Dict] = dict()
|
result: dict[str, dict[str, Any]] = dict()
|
||||||
for model_id, model in enumerate(model_holder.models):
|
for model_id, model in enumerate(loaded_models):
|
||||||
result[str(model_id)] = {
|
result[str(model_id)] = {
|
||||||
"config_path": model.config_path,
|
"config_path": model.config_path,
|
||||||
"model_path": model.model_path,
|
"model_path": model.model_path,
|
||||||
|
|||||||
188
slice.py
188
slice.py
@@ -1,28 +1,31 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import os
|
|
||||||
import shutil
|
import shutil
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from queue import Queue
|
||||||
|
from threading import Thread
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
import soundfile as sf
|
import soundfile as sf
|
||||||
import torch
|
import torch
|
||||||
import yaml
|
import yaml
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from common.log import logger
|
from style_bert_vits2.logging import logger
|
||||||
from common.stdout_wrapper import SAFE_STDOUT
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
vad_model, utils = torch.hub.load(
|
|
||||||
repo_or_dir="snakers4/silero-vad",
|
|
||||||
model="silero_vad",
|
|
||||||
onnx=True,
|
|
||||||
trust_repo=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
(get_speech_timestamps, _, read_audio, *_) = utils
|
def is_audio_file(file: Path) -> bool:
|
||||||
|
supported_extensions = [".wav", ".flac", ".mp3", ".ogg", ".opus"]
|
||||||
|
return file.suffix.lower() in supported_extensions
|
||||||
|
|
||||||
|
|
||||||
def get_stamps(
|
def get_stamps(
|
||||||
audio_file, min_silence_dur_ms: int = 700, min_sec: float = 2, max_sec: float = 12
|
vad_model: Any,
|
||||||
|
utils: Any,
|
||||||
|
audio_file: Path,
|
||||||
|
min_silence_dur_ms: int = 700,
|
||||||
|
min_sec: float = 2,
|
||||||
|
max_sec: float = 12,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
min_silence_dur_ms: int (ミリ秒):
|
min_silence_dur_ms: int (ミリ秒):
|
||||||
@@ -37,11 +40,12 @@ def get_stamps(
|
|||||||
この秒数より大きい発話は無視する。
|
この秒数より大きい発話は無視する。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
(get_speech_timestamps, _, read_audio, *_) = utils
|
||||||
sampling_rate = 16000 # 16kHzか8kHzのみ対応
|
sampling_rate = 16000 # 16kHzか8kHzのみ対応
|
||||||
|
|
||||||
min_ms = int(min_sec * 1000)
|
min_ms = int(min_sec * 1000)
|
||||||
|
|
||||||
wav = read_audio(audio_file, sampling_rate=sampling_rate)
|
wav = read_audio(str(audio_file), sampling_rate=sampling_rate)
|
||||||
speech_timestamps = get_speech_timestamps(
|
speech_timestamps = get_speech_timestamps(
|
||||||
wav,
|
wav,
|
||||||
vad_model,
|
vad_model,
|
||||||
@@ -55,16 +59,20 @@ def get_stamps(
|
|||||||
|
|
||||||
|
|
||||||
def split_wav(
|
def split_wav(
|
||||||
audio_file,
|
vad_model: Any,
|
||||||
target_dir="raw",
|
utils: Any,
|
||||||
min_sec=2,
|
audio_file: Path,
|
||||||
max_sec=12,
|
target_dir: Path,
|
||||||
min_silence_dur_ms=700,
|
min_sec: float = 2,
|
||||||
time_suffix=False,
|
max_sec: float = 12,
|
||||||
):
|
min_silence_dur_ms: int = 700,
|
||||||
margin = 200 # ミリ秒単位で、音声の前後に余裕を持たせる
|
time_suffix: bool = False,
|
||||||
|
) -> tuple[float, int]:
|
||||||
|
margin: int = 200 # ミリ秒単位で、音声の前後に余裕を持たせる
|
||||||
speech_timestamps = get_stamps(
|
speech_timestamps = get_stamps(
|
||||||
audio_file,
|
vad_model=vad_model,
|
||||||
|
utils=utils,
|
||||||
|
audio_file=audio_file,
|
||||||
min_silence_dur_ms=min_silence_dur_ms,
|
min_silence_dur_ms=min_silence_dur_ms,
|
||||||
min_sec=min_sec,
|
min_sec=min_sec,
|
||||||
max_sec=max_sec,
|
max_sec=max_sec,
|
||||||
@@ -74,10 +82,10 @@ def split_wav(
|
|||||||
|
|
||||||
total_ms = len(data) / sr * 1000
|
total_ms = len(data) / sr * 1000
|
||||||
|
|
||||||
file_name = os.path.basename(audio_file).split(".")[0]
|
file_name = audio_file.stem
|
||||||
os.makedirs(target_dir, exist_ok=True)
|
target_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
total_time_ms = 0
|
total_time_ms: float = 0
|
||||||
count = 0
|
count = 0
|
||||||
|
|
||||||
# タイムスタンプに従って分割し、ファイルに保存
|
# タイムスタンプに従って分割し、ファイルに保存
|
||||||
@@ -93,7 +101,7 @@ def split_wav(
|
|||||||
file = f"{file_name}-{int(start_ms)}-{int(end_ms)}.wav"
|
file = f"{file_name}-{int(start_ms)}-{int(end_ms)}.wav"
|
||||||
else:
|
else:
|
||||||
file = f"{file_name}-{i}.wav"
|
file = f"{file_name}-{i}.wav"
|
||||||
sf.write(os.path.join(target_dir, file), segment, sr)
|
sf.write(str(target_dir / file), segment, sr)
|
||||||
total_time_ms += end_ms - start_ms
|
total_time_ms += end_ms - start_ms
|
||||||
count += 1
|
count += 1
|
||||||
|
|
||||||
@@ -131,42 +139,128 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--time_suffix",
|
"--time_suffix",
|
||||||
"-t",
|
"-t",
|
||||||
action='store_true',
|
action="store_true",
|
||||||
help="Make the filename end with -start_ms-end_ms when saving wav.",
|
help="Make the filename end with -start_ms-end_ms when saving wav.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num_processes",
|
||||||
|
type=int,
|
||||||
|
default=3,
|
||||||
|
help="Number of processes to use. Default 3 seems to be the best.",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f:
|
with open(Path("configs/paths.yml"), "r", encoding="utf-8") as f:
|
||||||
path_config: dict[str, str] = yaml.safe_load(f.read())
|
path_config: dict[str, str] = yaml.safe_load(f.read())
|
||||||
dataset_root = path_config["dataset_root"]
|
dataset_root = path_config["dataset_root"]
|
||||||
|
|
||||||
input_dir = args.input_dir
|
model_name = str(args.model_name)
|
||||||
output_dir = os.path.join(dataset_root, args.model_name, "raw")
|
input_dir = Path(args.input_dir)
|
||||||
min_sec = args.min_sec
|
output_dir = Path(dataset_root) / model_name / "raw"
|
||||||
max_sec = args.max_sec
|
min_sec: float = args.min_sec
|
||||||
min_silence_dur_ms = args.min_silence_dur_ms
|
max_sec: float = args.max_sec
|
||||||
time_suffix = args.time_suffix
|
min_silence_dur_ms: int = args.min_silence_dur_ms
|
||||||
|
time_suffix: bool = args.time_suffix
|
||||||
|
num_processes: int = args.num_processes
|
||||||
|
|
||||||
wav_files = Path(input_dir).glob("**/*.wav")
|
audio_files = [file for file in input_dir.rglob("*") if is_audio_file(file)]
|
||||||
wav_files = list(wav_files)
|
|
||||||
logger.info(f"Found {len(wav_files)} wav files.")
|
logger.info(f"Found {len(audio_files)} audio files.")
|
||||||
if os.path.exists(output_dir):
|
if output_dir.exists():
|
||||||
logger.warning(f"Output directory {output_dir} already exists, deleting...")
|
logger.warning(f"Output directory {output_dir} already exists, deleting...")
|
||||||
shutil.rmtree(output_dir)
|
shutil.rmtree(output_dir)
|
||||||
|
|
||||||
|
# モデルをダウンロードしておく
|
||||||
|
_ = torch.hub.load(
|
||||||
|
repo_or_dir="litagin02/silero-vad",
|
||||||
|
model="silero_vad",
|
||||||
|
onnx=True,
|
||||||
|
trust_repo=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Silero VADのモデルは、同じインスタンスで並列処理するとおかしくなるらしい
|
||||||
|
# ワーカーごとにモデルをロードするようにするため、Queueを使って処理する
|
||||||
|
def process_queue(
|
||||||
|
q: Queue[Optional[Path]],
|
||||||
|
result_queue: Queue[tuple[float, int]],
|
||||||
|
error_queue: Queue[tuple[Path, Exception]],
|
||||||
|
):
|
||||||
|
# logger.debug("Worker started.")
|
||||||
|
vad_model, utils = torch.hub.load(
|
||||||
|
repo_or_dir="litagin02/silero-vad",
|
||||||
|
model="silero_vad",
|
||||||
|
onnx=True,
|
||||||
|
trust_repo=True,
|
||||||
|
)
|
||||||
|
while True:
|
||||||
|
file = q.get()
|
||||||
|
if file is None: # 終了シグナルを確認
|
||||||
|
q.task_done()
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
time_sec, count = split_wav(
|
||||||
|
vad_model=vad_model,
|
||||||
|
utils=utils,
|
||||||
|
audio_file=file,
|
||||||
|
target_dir=output_dir,
|
||||||
|
min_sec=min_sec,
|
||||||
|
max_sec=max_sec,
|
||||||
|
min_silence_dur_ms=min_silence_dur_ms,
|
||||||
|
time_suffix=time_suffix,
|
||||||
|
)
|
||||||
|
result_queue.put((time_sec, count))
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error processing {file}: {e}")
|
||||||
|
error_queue.put((file, e))
|
||||||
|
result_queue.put((0, 0))
|
||||||
|
finally:
|
||||||
|
q.task_done()
|
||||||
|
|
||||||
|
q: Queue[Optional[Path]] = Queue()
|
||||||
|
result_queue: Queue[tuple[float, int]] = Queue()
|
||||||
|
error_queue: Queue[tuple[Path, Exception]] = Queue()
|
||||||
|
|
||||||
|
# ファイル数が少ない場合は、ワーカー数をファイル数に合わせる
|
||||||
|
num_processes = min(num_processes, len(audio_files))
|
||||||
|
|
||||||
|
threads = [
|
||||||
|
Thread(target=process_queue, args=(q, result_queue, error_queue))
|
||||||
|
for _ in range(num_processes)
|
||||||
|
]
|
||||||
|
for t in threads:
|
||||||
|
t.start()
|
||||||
|
|
||||||
|
pbar = tqdm(total=len(audio_files), file=SAFE_STDOUT)
|
||||||
|
for file in audio_files:
|
||||||
|
q.put(file)
|
||||||
|
|
||||||
|
# result_queueを監視し、要素が追加されるごとに結果を加算しプログレスバーを更新
|
||||||
total_sec = 0
|
total_sec = 0
|
||||||
total_count = 0
|
total_count = 0
|
||||||
for wav_file in tqdm(wav_files, file=SAFE_STDOUT):
|
for _ in range(len(audio_files)):
|
||||||
time_sec, count = split_wav(
|
time, count = result_queue.get()
|
||||||
audio_file=str(wav_file),
|
total_sec += time
|
||||||
target_dir=output_dir,
|
|
||||||
min_sec=min_sec,
|
|
||||||
max_sec=max_sec,
|
|
||||||
min_silence_dur_ms=min_silence_dur_ms,
|
|
||||||
time_suffix=time_suffix,
|
|
||||||
)
|
|
||||||
total_sec += time_sec
|
|
||||||
total_count += count
|
total_count += count
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
# 全ての処理が終わるまで待つ
|
||||||
|
q.join()
|
||||||
|
|
||||||
|
# 終了シグナル None を送る
|
||||||
|
for _ in range(num_processes):
|
||||||
|
q.put(None)
|
||||||
|
|
||||||
|
for t in threads:
|
||||||
|
t.join()
|
||||||
|
|
||||||
|
pbar.close()
|
||||||
|
|
||||||
|
if not error_queue.empty():
|
||||||
|
error_str = "Error slicing some files:"
|
||||||
|
while not error_queue.empty():
|
||||||
|
file, e = error_queue.get()
|
||||||
|
error_str += f"\n{file}: {e}"
|
||||||
|
raise RuntimeError(error_str)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Slice done! Total time: {total_sec / 60:.2f} min, {total_count} files."
|
f"Slice done! Total time: {total_sec / 60:.2f} min, {total_count} files."
|
||||||
|
|||||||
87
spec_gen.py
87
spec_gen.py
@@ -1,87 +0,0 @@
|
|||||||
import torch
|
|
||||||
from tqdm import tqdm
|
|
||||||
from multiprocessing import Pool
|
|
||||||
from mel_processing import spectrogram_torch, mel_spectrogram_torch
|
|
||||||
from utils import load_wav_to_torch
|
|
||||||
|
|
||||||
|
|
||||||
class AudioProcessor:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
max_wav_value,
|
|
||||||
use_mel_spec_posterior,
|
|
||||||
filter_length,
|
|
||||||
n_mel_channels,
|
|
||||||
sampling_rate,
|
|
||||||
hop_length,
|
|
||||||
win_length,
|
|
||||||
mel_fmin,
|
|
||||||
mel_fmax,
|
|
||||||
):
|
|
||||||
self.max_wav_value = max_wav_value
|
|
||||||
self.use_mel_spec_posterior = use_mel_spec_posterior
|
|
||||||
self.filter_length = filter_length
|
|
||||||
self.n_mel_channels = n_mel_channels
|
|
||||||
self.sampling_rate = sampling_rate
|
|
||||||
self.hop_length = hop_length
|
|
||||||
self.win_length = win_length
|
|
||||||
self.mel_fmin = mel_fmin
|
|
||||||
self.mel_fmax = mel_fmax
|
|
||||||
|
|
||||||
def process_audio(self, filename):
|
|
||||||
audio, sampling_rate = load_wav_to_torch(filename)
|
|
||||||
audio_norm = audio / self.max_wav_value
|
|
||||||
audio_norm = audio_norm.unsqueeze(0)
|
|
||||||
spec_filename = filename.replace(".wav", ".spec.pt")
|
|
||||||
if self.use_mel_spec_posterior:
|
|
||||||
spec_filename = spec_filename.replace(".spec.pt", ".mel.pt")
|
|
||||||
try:
|
|
||||||
spec = torch.load(spec_filename)
|
|
||||||
except:
|
|
||||||
if self.use_mel_spec_posterior:
|
|
||||||
spec = mel_spectrogram_torch(
|
|
||||||
audio_norm,
|
|
||||||
self.filter_length,
|
|
||||||
self.n_mel_channels,
|
|
||||||
self.sampling_rate,
|
|
||||||
self.hop_length,
|
|
||||||
self.win_length,
|
|
||||||
self.mel_fmin,
|
|
||||||
self.mel_fmax,
|
|
||||||
center=False,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
spec = spectrogram_torch(
|
|
||||||
audio_norm,
|
|
||||||
self.filter_length,
|
|
||||||
self.sampling_rate,
|
|
||||||
self.hop_length,
|
|
||||||
self.win_length,
|
|
||||||
center=False,
|
|
||||||
)
|
|
||||||
spec = torch.squeeze(spec, 0)
|
|
||||||
torch.save(spec, spec_filename)
|
|
||||||
return spec, audio_norm
|
|
||||||
|
|
||||||
|
|
||||||
# 使用示例
|
|
||||||
processor = AudioProcessor(
|
|
||||||
max_wav_value=32768.0,
|
|
||||||
use_mel_spec_posterior=False,
|
|
||||||
filter_length=2048,
|
|
||||||
n_mel_channels=128,
|
|
||||||
sampling_rate=44100,
|
|
||||||
hop_length=512,
|
|
||||||
win_length=2048,
|
|
||||||
mel_fmin=0.0,
|
|
||||||
mel_fmax="null",
|
|
||||||
)
|
|
||||||
|
|
||||||
with open("filelists/train.list", "r") as f:
|
|
||||||
filepaths = [line.split("|")[0] for line in f] # 取每一行的第一部分作为audiopath
|
|
||||||
|
|
||||||
# 使用多进程处理
|
|
||||||
with Pool(processes=32) as pool: # 使用4个进程
|
|
||||||
with tqdm(total=len(filepaths)) as pbar:
|
|
||||||
for i, _ in enumerate(pool.imap_unordered(processor.process_audio, filepaths)):
|
|
||||||
pbar.update()
|
|
||||||
@@ -10,9 +10,10 @@ import pandas as pd
|
|||||||
import torch
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from common.log import logger
|
|
||||||
from common.tts_model import Model
|
|
||||||
from config import config
|
from config import config
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.tts_model import TTSModel
|
||||||
|
|
||||||
|
|
||||||
warnings.filterwarnings("ignore")
|
warnings.filterwarnings("ignore")
|
||||||
|
|
||||||
@@ -54,7 +55,7 @@ safetensors_files = model_path.glob("*.safetensors")
|
|||||||
|
|
||||||
|
|
||||||
def get_model(model_file: Path):
|
def get_model(model_file: Path):
|
||||||
return Model(
|
return TTSModel(
|
||||||
model_path=str(model_file),
|
model_path=str(model_file),
|
||||||
config_path=str(model_file.parent / "config.json"),
|
config_path=str(model_file.parent / "config.json"),
|
||||||
style_vec_path=str(model_file.parent / "style_vectors.npy"),
|
style_vec_path=str(model_file.parent / "style_vectors.npy"),
|
||||||
|
|||||||
15
style_bert_vits2/.editorconfig
Normal file
15
style_bert_vits2/.editorconfig
Normal file
@@ -0,0 +1,15 @@
|
|||||||
|
root = true
|
||||||
|
|
||||||
|
[*]
|
||||||
|
charset = utf-8
|
||||||
|
end_of_line = lf
|
||||||
|
insert_final_newline = true
|
||||||
|
indent_size = 4
|
||||||
|
indent_style = space
|
||||||
|
trim_trailing_whitespace = true
|
||||||
|
|
||||||
|
[*.md]
|
||||||
|
trim_trailing_whitespace = false
|
||||||
|
|
||||||
|
[*.yml]
|
||||||
|
indent_size = 2
|
||||||
0
style_bert_vits2/__init__.py
Normal file
0
style_bert_vits2/__init__.py
Normal file
48
style_bert_vits2/constants.py
Normal file
48
style_bert_vits2/constants.py
Normal file
@@ -0,0 +1,48 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from style_bert_vits2.utils.strenum import StrEnum
|
||||||
|
|
||||||
|
|
||||||
|
# Style-Bert-VITS2 のバージョン
|
||||||
|
VERSION = "2.4.0"
|
||||||
|
|
||||||
|
# Style-Bert-VITS2 のベースディレクトリ
|
||||||
|
BASE_DIR = Path(__file__).parent.parent
|
||||||
|
|
||||||
|
|
||||||
|
# 利用可能な言語
|
||||||
|
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
|
||||||
|
class Languages(StrEnum):
|
||||||
|
JP = "JP"
|
||||||
|
EN = "EN"
|
||||||
|
ZH = "ZH"
|
||||||
|
|
||||||
|
|
||||||
|
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
||||||
|
DEFAULT_BERT_TOKENIZER_PATHS = {
|
||||||
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
||||||
|
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
||||||
|
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
||||||
|
}
|
||||||
|
|
||||||
|
# デフォルトのユーザー辞書ディレクトリ
|
||||||
|
## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される
|
||||||
|
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
||||||
|
DEFAULT_USER_DICT_DIR = BASE_DIR / "dict_data"
|
||||||
|
|
||||||
|
# デフォルトの推論パラメータ
|
||||||
|
DEFAULT_STYLE = "Neutral"
|
||||||
|
DEFAULT_STYLE_WEIGHT = 5.0
|
||||||
|
DEFAULT_SDP_RATIO = 0.2
|
||||||
|
DEFAULT_NOISE = 0.6
|
||||||
|
DEFAULT_NOISEW = 0.8
|
||||||
|
DEFAULT_LENGTH = 1.0
|
||||||
|
DEFAULT_LINE_SPLIT = True
|
||||||
|
DEFAULT_SPLIT_INTERVAL = 0.5
|
||||||
|
DEFAULT_ASSIST_TEXT_WEIGHT = 0.7
|
||||||
|
DEFAULT_ASSIST_TEXT_WEIGHT = 1.0
|
||||||
|
|
||||||
|
# Gradio のテーマ
|
||||||
|
## Built-in theme: "default", "base", "monochrome", "soft", "glass"
|
||||||
|
## See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
||||||
|
GRADIO_THEME = "NoCrypt/miku"
|
||||||
15
style_bert_vits2/logging.py
Normal file
15
style_bert_vits2/logging.py
Normal file
@@ -0,0 +1,15 @@
|
|||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
|
|
||||||
|
# Remove all default handlers
|
||||||
|
logger.remove()
|
||||||
|
|
||||||
|
# Add a new handler
|
||||||
|
logger.add(
|
||||||
|
SAFE_STDOUT,
|
||||||
|
format="<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}",
|
||||||
|
backtrace=True,
|
||||||
|
diagnose=True,
|
||||||
|
)
|
||||||
0
style_bert_vits2/models/__init__.py
Normal file
0
style_bert_vits2/models/__init__.py
Normal file
@@ -1,14 +1,15 @@
|
|||||||
import math
|
import math
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
|
|
||||||
import commons
|
from style_bert_vits2.models import commons
|
||||||
from common.log import logger as logging
|
|
||||||
|
|
||||||
|
|
||||||
class LayerNorm(nn.Module):
|
class LayerNorm(nn.Module):
|
||||||
def __init__(self, channels, eps=1e-5):
|
def __init__(self, channels: int, eps: float = 1e-5) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.eps = eps
|
self.eps = eps
|
||||||
@@ -16,14 +17,16 @@ class LayerNorm(nn.Module):
|
|||||||
self.gamma = nn.Parameter(torch.ones(channels))
|
self.gamma = nn.Parameter(torch.ones(channels))
|
||||||
self.beta = nn.Parameter(torch.zeros(channels))
|
self.beta = nn.Parameter(torch.zeros(channels))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
x = x.transpose(1, -1)
|
x = x.transpose(1, -1)
|
||||||
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
|
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
|
||||||
return x.transpose(1, -1)
|
return x.transpose(1, -1)
|
||||||
|
|
||||||
|
|
||||||
@torch.jit.script
|
@torch.jit.script # type: ignore
|
||||||
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
|
def fused_add_tanh_sigmoid_multiply(
|
||||||
|
input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int]
|
||||||
|
) -> torch.Tensor:
|
||||||
n_channels_int = n_channels[0]
|
n_channels_int = n_channels[0]
|
||||||
in_act = input_a + input_b
|
in_act = input_a + input_b
|
||||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||||
@@ -35,16 +38,16 @@ def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
|
|||||||
class Encoder(nn.Module):
|
class Encoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size=1,
|
kernel_size: int = 1,
|
||||||
p_dropout=0.0,
|
p_dropout: float = 0.0,
|
||||||
window_size=4,
|
window_size: int = 4,
|
||||||
isflow=True,
|
isflow: bool = True,
|
||||||
**kwargs
|
**kwargs: Any,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
self.filter_channels = filter_channels
|
self.filter_channels = filter_channels
|
||||||
@@ -67,7 +70,7 @@ class Encoder(nn.Module):
|
|||||||
self.cond_layer_idx = (
|
self.cond_layer_idx = (
|
||||||
kwargs["cond_layer_idx"] if "cond_layer_idx" in kwargs else 2
|
kwargs["cond_layer_idx"] if "cond_layer_idx" in kwargs else 2
|
||||||
)
|
)
|
||||||
# logging.debug(self.gin_channels, self.cond_layer_idx)
|
# logger.debug(self.gin_channels, self.cond_layer_idx)
|
||||||
assert (
|
assert (
|
||||||
self.cond_layer_idx < self.n_layers
|
self.cond_layer_idx < self.n_layers
|
||||||
), "cond_layer_idx should be less than n_layers"
|
), "cond_layer_idx should be less than n_layers"
|
||||||
@@ -98,12 +101,15 @@ class Encoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
||||||
x = x * x_mask
|
x = x * x_mask
|
||||||
for i in range(self.n_layers):
|
for i in range(self.n_layers):
|
||||||
if i == self.cond_layer_idx and g is not None:
|
if i == self.cond_layer_idx and g is not None:
|
||||||
g = self.spk_emb_linear(g.transpose(1, 2))
|
g = self.spk_emb_linear(g.transpose(1, 2))
|
||||||
|
assert g is not None
|
||||||
g = g.transpose(1, 2)
|
g = g.transpose(1, 2)
|
||||||
x = x + g
|
x = x + g
|
||||||
x = x * x_mask
|
x = x * x_mask
|
||||||
@@ -121,16 +127,16 @@ class Encoder(nn.Module):
|
|||||||
class Decoder(nn.Module):
|
class Decoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size=1,
|
kernel_size: int = 1,
|
||||||
p_dropout=0.0,
|
p_dropout: float = 0.0,
|
||||||
proximal_bias=False,
|
proximal_bias: bool = False,
|
||||||
proximal_init=True,
|
proximal_init: bool = True,
|
||||||
**kwargs
|
**kwargs: Any,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
self.filter_channels = filter_channels
|
self.filter_channels = filter_channels
|
||||||
@@ -178,7 +184,13 @@ class Decoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||||
|
|
||||||
def forward(self, x, x_mask, h, h_mask):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
h: torch.Tensor,
|
||||||
|
h_mask: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: decoder input
|
x: decoder input
|
||||||
h: encoder output
|
h: encoder output
|
||||||
@@ -207,16 +219,16 @@ class Decoder(nn.Module):
|
|||||||
class MultiHeadAttention(nn.Module):
|
class MultiHeadAttention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
p_dropout=0.0,
|
p_dropout: float = 0.0,
|
||||||
window_size=None,
|
window_size: Optional[int] = None,
|
||||||
heads_share=True,
|
heads_share: bool = True,
|
||||||
block_length=None,
|
block_length: Optional[int] = None,
|
||||||
proximal_bias=False,
|
proximal_bias: bool = False,
|
||||||
proximal_init=False,
|
proximal_init: bool = False,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert channels % n_heads == 0
|
assert channels % n_heads == 0
|
||||||
|
|
||||||
@@ -256,9 +268,13 @@ class MultiHeadAttention(nn.Module):
|
|||||||
if proximal_init:
|
if proximal_init:
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
self.conv_k.weight.copy_(self.conv_q.weight)
|
self.conv_k.weight.copy_(self.conv_q.weight)
|
||||||
|
assert self.conv_k.bias is not None
|
||||||
|
assert self.conv_q.bias is not None
|
||||||
self.conv_k.bias.copy_(self.conv_q.bias)
|
self.conv_k.bias.copy_(self.conv_q.bias)
|
||||||
|
|
||||||
def forward(self, x, c, attn_mask=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
q = self.conv_q(x)
|
q = self.conv_q(x)
|
||||||
k = self.conv_k(c)
|
k = self.conv_k(c)
|
||||||
v = self.conv_v(c)
|
v = self.conv_v(c)
|
||||||
@@ -268,7 +284,13 @@ class MultiHeadAttention(nn.Module):
|
|||||||
x = self.conv_o(x)
|
x = self.conv_o(x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def attention(self, query, key, value, mask=None):
|
def attention(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
key: torch.Tensor,
|
||||||
|
value: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
# reshape [b, d, t] -> [b, n_h, t, d_k]
|
# reshape [b, d, t] -> [b, n_h, t, d_k]
|
||||||
b, d, t_s, t_t = (*key.size(), query.size(2))
|
b, d, t_s, t_t = (*key.size(), query.size(2))
|
||||||
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
|
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
|
||||||
@@ -319,7 +341,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
||||||
return output, p_attn
|
return output, p_attn
|
||||||
|
|
||||||
def _matmul_with_relative_values(self, x, y):
|
def _matmul_with_relative_values(
|
||||||
|
self, x: torch.Tensor, y: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, m]
|
x: [b, h, l, m]
|
||||||
y: [h or 1, m, d]
|
y: [h or 1, m, d]
|
||||||
@@ -328,7 +352,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
ret = torch.matmul(x, y.unsqueeze(0))
|
ret = torch.matmul(x, y.unsqueeze(0))
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def _matmul_with_relative_keys(self, x, y):
|
def _matmul_with_relative_keys(
|
||||||
|
self, x: torch.Tensor, y: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, d]
|
x: [b, h, l, d]
|
||||||
y: [h or 1, m, d]
|
y: [h or 1, m, d]
|
||||||
@@ -337,8 +363,11 @@ class MultiHeadAttention(nn.Module):
|
|||||||
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def _get_relative_embeddings(self, relative_embeddings, length):
|
def _get_relative_embeddings(
|
||||||
2 * self.window_size + 1
|
self, relative_embeddings: torch.Tensor, length: int
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert self.window_size is not None
|
||||||
|
2 * self.window_size + 1 # type: ignore
|
||||||
# Pad first before slice to avoid using cond ops.
|
# Pad first before slice to avoid using cond ops.
|
||||||
pad_length = max(length - (self.window_size + 1), 0)
|
pad_length = max(length - (self.window_size + 1), 0)
|
||||||
slice_start_position = max((self.window_size + 1) - length, 0)
|
slice_start_position = max((self.window_size + 1) - length, 0)
|
||||||
@@ -355,7 +384,7 @@ class MultiHeadAttention(nn.Module):
|
|||||||
]
|
]
|
||||||
return used_relative_embeddings
|
return used_relative_embeddings
|
||||||
|
|
||||||
def _relative_position_to_absolute_position(self, x):
|
def _relative_position_to_absolute_position(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, 2*l-1]
|
x: [b, h, l, 2*l-1]
|
||||||
ret: [b, h, l, l]
|
ret: [b, h, l, l]
|
||||||
@@ -376,7 +405,7 @@ class MultiHeadAttention(nn.Module):
|
|||||||
]
|
]
|
||||||
return x_final
|
return x_final
|
||||||
|
|
||||||
def _absolute_position_to_relative_position(self, x):
|
def _absolute_position_to_relative_position(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, l]
|
x: [b, h, l, l]
|
||||||
ret: [b, h, l, 2*l-1]
|
ret: [b, h, l, 2*l-1]
|
||||||
@@ -392,7 +421,7 @@ class MultiHeadAttention(nn.Module):
|
|||||||
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
|
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
|
||||||
return x_final
|
return x_final
|
||||||
|
|
||||||
def _attention_bias_proximal(self, length):
|
def _attention_bias_proximal(self, length: int) -> torch.Tensor:
|
||||||
"""Bias for self-attention to encourage attention to close positions.
|
"""Bias for self-attention to encourage attention to close positions.
|
||||||
Args:
|
Args:
|
||||||
length: an integer scalar.
|
length: an integer scalar.
|
||||||
@@ -407,14 +436,14 @@ class MultiHeadAttention(nn.Module):
|
|||||||
class FFN(nn.Module):
|
class FFN(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout=0.0,
|
p_dropout: float = 0.0,
|
||||||
activation=None,
|
activation: Optional[str] = None,
|
||||||
causal=False,
|
causal: bool = False,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
@@ -433,7 +462,7 @@ class FFN(nn.Module):
|
|||||||
self.conv_2 = nn.Conv1d(filter_channels, out_channels, kernel_size)
|
self.conv_2 = nn.Conv1d(filter_channels, out_channels, kernel_size)
|
||||||
self.drop = nn.Dropout(p_dropout)
|
self.drop = nn.Dropout(p_dropout)
|
||||||
|
|
||||||
def forward(self, x, x_mask):
|
def forward(self, x: torch.Tensor, x_mask: torch.Tensor) -> torch.Tensor:
|
||||||
x = self.conv_1(self.padding(x * x_mask))
|
x = self.conv_1(self.padding(x * x_mask))
|
||||||
if self.activation == "gelu":
|
if self.activation == "gelu":
|
||||||
x = x * torch.sigmoid(1.702 * x)
|
x = x * torch.sigmoid(1.702 * x)
|
||||||
@@ -443,7 +472,7 @@ class FFN(nn.Module):
|
|||||||
x = self.conv_2(self.padding(x * x_mask))
|
x = self.conv_2(self.padding(x * x_mask))
|
||||||
return x * x_mask
|
return x * x_mask
|
||||||
|
|
||||||
def _causal_padding(self, x):
|
def _causal_padding(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
if self.kernel_size == 1:
|
if self.kernel_size == 1:
|
||||||
return x
|
return x
|
||||||
pad_l = self.kernel_size - 1
|
pad_l = self.kernel_size - 1
|
||||||
@@ -452,7 +481,7 @@ class FFN(nn.Module):
|
|||||||
x = F.pad(x, commons.convert_pad_shape(padding))
|
x = F.pad(x, commons.convert_pad_shape(padding))
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def _same_padding(self, x):
|
def _same_padding(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
if self.kernel_size == 1:
|
if self.kernel_size == 1:
|
||||||
return x
|
return x
|
||||||
pad_l = (self.kernel_size - 1) // 2
|
pad_l = (self.kernel_size - 1) // 2
|
||||||
223
style_bert_vits2/models/commons.py
Normal file
223
style_bert_vits2/models/commons.py
Normal file
@@ -0,0 +1,223 @@
|
|||||||
|
"""
|
||||||
|
以下に記述されている関数のコメントはリファクタリング時に GPT-4 に生成させたもので、
|
||||||
|
コードと完全に一致している保証はない。あくまで参考程度とすること。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import functional as F
|
||||||
|
|
||||||
|
|
||||||
|
def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None:
|
||||||
|
"""
|
||||||
|
モジュールの重みを初期化する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
m (torch.nn.Module): 重みを初期化する対象のモジュール
|
||||||
|
mean (float): 正規分布の平均
|
||||||
|
std (float): 正規分布の標準偏差
|
||||||
|
"""
|
||||||
|
classname = m.__class__.__name__
|
||||||
|
if classname.find("Conv") != -1:
|
||||||
|
m.weight.data.normal_(mean, std)
|
||||||
|
|
||||||
|
|
||||||
|
def get_padding(kernel_size: int, dilation: int = 1) -> int:
|
||||||
|
"""
|
||||||
|
カーネルサイズと膨張率からパディングの大きさを計算する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
kernel_size (int): カーネルのサイズ
|
||||||
|
dilation (int): 膨張率
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
int: 計算されたパディングの大きさ
|
||||||
|
"""
|
||||||
|
return int((kernel_size * dilation - dilation) / 2)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_pad_shape(pad_shape: list[list[Any]]) -> list[Any]:
|
||||||
|
"""
|
||||||
|
パディングの形状を変換する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pad_shape (list[list[Any]]): 変換前のパディングの形状
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[Any]: 変換後のパディングの形状
|
||||||
|
"""
|
||||||
|
layer = pad_shape[::-1]
|
||||||
|
new_pad_shape = [item for sublist in layer for item in sublist]
|
||||||
|
return new_pad_shape
|
||||||
|
|
||||||
|
|
||||||
|
def intersperse(lst: list[Any], item: Any) -> list[Any]:
|
||||||
|
"""
|
||||||
|
リストの要素の間に特定のアイテムを挿入する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
lst (list[Any]): 元のリスト
|
||||||
|
item (Any): 挿入するアイテム
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[Any]: 新しいリスト
|
||||||
|
"""
|
||||||
|
result = [item] * (len(lst) * 2 + 1)
|
||||||
|
result[1::2] = lst
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def slice_segments(
|
||||||
|
x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
テンソルからセグメントをスライスする
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x (torch.Tensor): 入力テンソル
|
||||||
|
ids_str (torch.Tensor): スライスを開始するインデックス
|
||||||
|
segment_size (int, optional): スライスのサイズ (デフォルト: 4)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: スライスされたセグメント
|
||||||
|
"""
|
||||||
|
gather_indices = ids_str.view(x.size(0), 1, 1).repeat(
|
||||||
|
1, x.size(1), 1
|
||||||
|
) + torch.arange(segment_size, device=x.device)
|
||||||
|
return torch.gather(x, 2, gather_indices)
|
||||||
|
|
||||||
|
|
||||||
|
def rand_slice_segments(
|
||||||
|
x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""
|
||||||
|
ランダムなセグメントをスライスする
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x (torch.Tensor): 入力テンソル
|
||||||
|
x_lengths (Optional[torch.Tensor], optional): 各バッチの長さ (デフォルト: None)
|
||||||
|
segment_size (int, optional): スライスのサイズ (デフォルト: 4)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[torch.Tensor, torch.Tensor]: スライスされたセグメントと開始インデックス
|
||||||
|
"""
|
||||||
|
b, d, t = x.size()
|
||||||
|
if x_lengths is None:
|
||||||
|
x_lengths = t # type: ignore
|
||||||
|
ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0) # type: ignore
|
||||||
|
ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long)
|
||||||
|
ret = slice_segments(x, ids_str, segment_size)
|
||||||
|
return ret, ids_str
|
||||||
|
|
||||||
|
|
||||||
|
def subsequent_mask(length: int) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
後続のマスクを生成する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
length (int): マスクのサイズ
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: 生成されたマスク
|
||||||
|
"""
|
||||||
|
mask = torch.tril(torch.ones(length, length)).unsqueeze(0).unsqueeze(0)
|
||||||
|
return mask
|
||||||
|
|
||||||
|
|
||||||
|
@torch.jit.script # type: ignore
|
||||||
|
def fused_add_tanh_sigmoid_multiply(
|
||||||
|
input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_a (torch.Tensor): 入力テンソル A
|
||||||
|
input_b (torch.Tensor): 入力テンソル B
|
||||||
|
n_channels (torch.Tensor): チャネル数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: 演算結果
|
||||||
|
"""
|
||||||
|
n_channels_int = n_channels[0]
|
||||||
|
in_act = input_a + input_b
|
||||||
|
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||||
|
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||||
|
acts = t_act * s_act
|
||||||
|
return acts
|
||||||
|
|
||||||
|
|
||||||
|
def sequence_mask(
|
||||||
|
length: torch.Tensor, max_length: Optional[int] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
シーケンスマスクを生成する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
length (torch.Tensor): 各シーケンスの長さ
|
||||||
|
max_length (Optional[int]): 最大のシーケンス長さ。指定されていない場合は length の最大値を使用
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: 生成されたシーケンスマスク
|
||||||
|
"""
|
||||||
|
if max_length is None:
|
||||||
|
max_length = length.max() # type: ignore
|
||||||
|
x = torch.arange(max_length, dtype=length.dtype, device=length.device) # type: ignore
|
||||||
|
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
パスを生成する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
duration (torch.Tensor): 各時間ステップの持続時間
|
||||||
|
mask (torch.Tensor): マスクテンソル
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: 生成されたパス
|
||||||
|
"""
|
||||||
|
b, _, t_y, t_x = mask.shape
|
||||||
|
cum_duration = torch.cumsum(duration, -1)
|
||||||
|
|
||||||
|
cum_duration_flat = cum_duration.view(b * t_x)
|
||||||
|
path = sequence_mask(cum_duration_flat, t_y).to(mask.dtype)
|
||||||
|
path = path.view(b, t_x, t_y)
|
||||||
|
path = path - F.pad(path, convert_pad_shape([[0, 0], [1, 0], [0, 0]]))[:, :-1]
|
||||||
|
path = path.unsqueeze(1).transpose(2, 3) * mask
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def clip_grad_value_(
|
||||||
|
parameters: Union[torch.Tensor, list[torch.Tensor]],
|
||||||
|
clip_value: Optional[float],
|
||||||
|
norm_type: float = 2.0,
|
||||||
|
) -> float:
|
||||||
|
"""
|
||||||
|
勾配の値をクリップする
|
||||||
|
|
||||||
|
Args:
|
||||||
|
parameters (Union[torch.Tensor, list[torch.Tensor]]): クリップするパラメータ
|
||||||
|
clip_value (Optional[float]): クリップする値。None の場合はクリップしない
|
||||||
|
norm_type (float): ノルムの種類
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
float: 総ノルム
|
||||||
|
"""
|
||||||
|
if isinstance(parameters, torch.Tensor):
|
||||||
|
parameters = [parameters]
|
||||||
|
parameters = list(filter(lambda p: p.grad is not None, parameters))
|
||||||
|
norm_type = float(norm_type)
|
||||||
|
if clip_value is not None:
|
||||||
|
clip_value = float(clip_value)
|
||||||
|
|
||||||
|
total_norm = 0.0
|
||||||
|
for p in parameters:
|
||||||
|
assert p.grad is not None
|
||||||
|
param_norm = p.grad.data.norm(norm_type)
|
||||||
|
total_norm += param_norm.item() ** norm_type
|
||||||
|
if clip_value is not None:
|
||||||
|
p.grad.data.clamp_(min=-clip_value, max=clip_value)
|
||||||
|
total_norm = total_norm ** (1.0 / norm_type)
|
||||||
|
return total_norm
|
||||||
129
style_bert_vits2/models/hyper_parameters.py
Normal file
129
style_bert_vits2/models/hyper_parameters.py
Normal file
@@ -0,0 +1,129 @@
|
|||||||
|
"""
|
||||||
|
Style-Bert-VITS2 モデルのハイパーパラメータを表す Pydantic モデル。
|
||||||
|
デフォルト値は configs/config_jp_extra.json 内の定義と概ね同一で、
|
||||||
|
万が一ロードした config.json に存在しないキーがあった際のフェイルセーフとして適用される。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional, Union
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
|
||||||
|
|
||||||
|
class HyperParametersTrain(BaseModel):
|
||||||
|
log_interval: int = 200
|
||||||
|
eval_interval: int = 1000
|
||||||
|
seed: int = 42
|
||||||
|
epochs: int = 1000
|
||||||
|
learning_rate: float = 0.0001
|
||||||
|
betas: tuple[float, float] = (0.8, 0.99)
|
||||||
|
eps: float = 1e-9
|
||||||
|
batch_size: int = 2
|
||||||
|
bf16_run: bool = False
|
||||||
|
fp16_run: bool = False
|
||||||
|
lr_decay: float = 0.99996
|
||||||
|
segment_size: int = 16384
|
||||||
|
init_lr_ratio: int = 1
|
||||||
|
warmup_epochs: int = 0
|
||||||
|
c_mel: int = 45
|
||||||
|
c_kl: float = 1.0
|
||||||
|
c_commit: int = 100
|
||||||
|
skip_optimizer: bool = False
|
||||||
|
freeze_ZH_bert: bool = False
|
||||||
|
freeze_JP_bert: bool = False
|
||||||
|
freeze_EN_bert: bool = False
|
||||||
|
freeze_emo: bool = False
|
||||||
|
freeze_style: bool = False
|
||||||
|
freeze_decoder: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
class HyperParametersData(BaseModel):
|
||||||
|
use_jp_extra: bool = True
|
||||||
|
training_files: str = "Data/Dummy/train.list"
|
||||||
|
validation_files: str = "Data/Dummy/val.list"
|
||||||
|
max_wav_value: float = 32768.0
|
||||||
|
sampling_rate: int = 44100
|
||||||
|
filter_length: int = 2048
|
||||||
|
hop_length: int = 512
|
||||||
|
win_length: int = 2048
|
||||||
|
n_mel_channels: int = 128
|
||||||
|
mel_fmin: float = 0.0
|
||||||
|
mel_fmax: Optional[float] = None
|
||||||
|
add_blank: bool = True
|
||||||
|
n_speakers: int = 1
|
||||||
|
cleaned_text: bool = True
|
||||||
|
spk2id: dict[str, int] = {
|
||||||
|
"Dummy": 0,
|
||||||
|
}
|
||||||
|
num_styles: int = 1
|
||||||
|
style2id: dict[str, int] = {
|
||||||
|
"Neutral": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class HyperParametersModelSLM(BaseModel):
|
||||||
|
model: str = "./slm/wavlm-base-plus"
|
||||||
|
sr: int = 16000
|
||||||
|
hidden: int = 768
|
||||||
|
nlayers: int = 13
|
||||||
|
initial_channel: int = 64
|
||||||
|
|
||||||
|
|
||||||
|
class HyperParametersModel(BaseModel):
|
||||||
|
use_spk_conditioned_encoder: bool = True
|
||||||
|
use_noise_scaled_mas: bool = True
|
||||||
|
use_mel_posterior_encoder: bool = False
|
||||||
|
use_duration_discriminator: bool = False
|
||||||
|
use_wavlm_discriminator: bool = True
|
||||||
|
inter_channels: int = 192
|
||||||
|
hidden_channels: int = 192
|
||||||
|
filter_channels: int = 768
|
||||||
|
n_heads: int = 2
|
||||||
|
n_layers: int = 6
|
||||||
|
kernel_size: int = 3
|
||||||
|
p_dropout: float = 0.1
|
||||||
|
resblock: str = "1"
|
||||||
|
resblock_kernel_sizes: list[int] = [3, 7, 11]
|
||||||
|
resblock_dilation_sizes: list[list[int]] = [
|
||||||
|
[1, 3, 5],
|
||||||
|
[1, 3, 5],
|
||||||
|
[1, 3, 5],
|
||||||
|
]
|
||||||
|
upsample_rates: list[int] = [8, 8, 2, 2, 2]
|
||||||
|
upsample_initial_channel: int = 512
|
||||||
|
upsample_kernel_sizes: list[int] = [16, 16, 8, 2, 2]
|
||||||
|
n_layers_q: int = 3
|
||||||
|
use_spectral_norm: bool = False
|
||||||
|
gin_channels: int = 512
|
||||||
|
slm: HyperParametersModelSLM = HyperParametersModelSLM()
|
||||||
|
|
||||||
|
|
||||||
|
class HyperParameters(BaseModel):
|
||||||
|
model_name: str = "Dummy"
|
||||||
|
version: str = "2.0-JP-Extra"
|
||||||
|
train: HyperParametersTrain = HyperParametersTrain()
|
||||||
|
data: HyperParametersData = HyperParametersData()
|
||||||
|
model: HyperParametersModel = HyperParametersModel()
|
||||||
|
|
||||||
|
# 以下は学習時にのみ動的に設定されるパラメータ (通常 config.json には存在しない)
|
||||||
|
model_dir: Optional[str] = None
|
||||||
|
speedup: bool = False
|
||||||
|
repo_id: Optional[str] = None
|
||||||
|
|
||||||
|
# model_ 以下を Pydantic の保護対象から除外する
|
||||||
|
model_config = ConfigDict(protected_namespaces=())
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def load_from_json(json_path: Union[str, Path]) -> "HyperParameters":
|
||||||
|
"""
|
||||||
|
与えられた JSON ファイルからハイパーパラメータを読み込む。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
json_path (Union[str, Path]): JSON ファイルのパス
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
HyperParameters: ハイパーパラメータ
|
||||||
|
"""
|
||||||
|
|
||||||
|
with open(json_path, "r", encoding="utf-8") as f:
|
||||||
|
return HyperParameters.model_validate_json(f.read())
|
||||||
267
style_bert_vits2/models/infer.py
Normal file
267
style_bert_vits2/models/infer.py
Normal file
@@ -0,0 +1,267 @@
|
|||||||
|
from typing import Any, Optional, Union, cast
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.models import commons, utils
|
||||||
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
|
from style_bert_vits2.models.models import SynthesizerTrn
|
||||||
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp import (
|
||||||
|
clean_text,
|
||||||
|
cleaned_text_to_sequence,
|
||||||
|
extract_bert_feature,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp.symbols import SYMBOLS
|
||||||
|
|
||||||
|
|
||||||
|
def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
|
||||||
|
if version.endswith("JP-Extra"):
|
||||||
|
logger.info("Using JP-Extra model")
|
||||||
|
net_g = SynthesizerTrnJPExtra(
|
||||||
|
n_vocab=len(SYMBOLS),
|
||||||
|
spec_channels=hps.data.filter_length // 2 + 1,
|
||||||
|
segment_size=hps.train.segment_size // hps.data.hop_length,
|
||||||
|
n_speakers=hps.data.n_speakers,
|
||||||
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
|
inter_channels=hps.model.inter_channels,
|
||||||
|
hidden_channels=hps.model.hidden_channels,
|
||||||
|
filter_channels=hps.model.filter_channels,
|
||||||
|
n_heads=hps.model.n_heads,
|
||||||
|
n_layers=hps.model.n_layers,
|
||||||
|
kernel_size=hps.model.kernel_size,
|
||||||
|
p_dropout=hps.model.p_dropout,
|
||||||
|
resblock=hps.model.resblock,
|
||||||
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
|
upsample_rates=hps.model.upsample_rates,
|
||||||
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
|
n_layers_q=hps.model.n_layers_q,
|
||||||
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
|
gin_channels=hps.model.gin_channels,
|
||||||
|
slm=hps.model.slm,
|
||||||
|
).to(device)
|
||||||
|
else:
|
||||||
|
logger.info("Using normal model")
|
||||||
|
net_g = SynthesizerTrn(
|
||||||
|
n_vocab=len(SYMBOLS),
|
||||||
|
spec_channels=hps.data.filter_length // 2 + 1,
|
||||||
|
segment_size=hps.train.segment_size // hps.data.hop_length,
|
||||||
|
n_speakers=hps.data.n_speakers,
|
||||||
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
|
inter_channels=hps.model.inter_channels,
|
||||||
|
hidden_channels=hps.model.hidden_channels,
|
||||||
|
filter_channels=hps.model.filter_channels,
|
||||||
|
n_heads=hps.model.n_heads,
|
||||||
|
n_layers=hps.model.n_layers,
|
||||||
|
kernel_size=hps.model.kernel_size,
|
||||||
|
p_dropout=hps.model.p_dropout,
|
||||||
|
resblock=hps.model.resblock,
|
||||||
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
|
upsample_rates=hps.model.upsample_rates,
|
||||||
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
|
n_layers_q=hps.model.n_layers_q,
|
||||||
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
|
gin_channels=hps.model.gin_channels,
|
||||||
|
slm=hps.model.slm,
|
||||||
|
).to(device)
|
||||||
|
net_g.state_dict()
|
||||||
|
_ = net_g.eval()
|
||||||
|
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
||||||
|
_ = utils.checkpoints.load_checkpoint(
|
||||||
|
model_path, net_g, None, skip_optimizer=True
|
||||||
|
)
|
||||||
|
elif model_path.endswith(".safetensors"):
|
||||||
|
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown model format: {model_path}")
|
||||||
|
return net_g
|
||||||
|
|
||||||
|
|
||||||
|
def get_text(
|
||||||
|
text: str,
|
||||||
|
language_str: Languages,
|
||||||
|
hps: HyperParameters,
|
||||||
|
device: str,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
):
|
||||||
|
use_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
|
# 推論時のみ呼び出されるので、raise_yomi_error は False に設定
|
||||||
|
norm_text, phone, tone, word2ph = clean_text(
|
||||||
|
text,
|
||||||
|
language_str,
|
||||||
|
use_jp_extra=use_jp_extra,
|
||||||
|
raise_yomi_error=False,
|
||||||
|
)
|
||||||
|
if given_tone is not None:
|
||||||
|
if len(given_tone) != len(phone):
|
||||||
|
raise InvalidToneError(
|
||||||
|
f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})"
|
||||||
|
)
|
||||||
|
tone = given_tone
|
||||||
|
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
||||||
|
|
||||||
|
if hps.data.add_blank:
|
||||||
|
phone = commons.intersperse(phone, 0)
|
||||||
|
tone = commons.intersperse(tone, 0)
|
||||||
|
language = commons.intersperse(language, 0)
|
||||||
|
for i in range(len(word2ph)):
|
||||||
|
word2ph[i] = word2ph[i] * 2
|
||||||
|
word2ph[0] += 1
|
||||||
|
bert_ori = extract_bert_feature(
|
||||||
|
norm_text,
|
||||||
|
word2ph,
|
||||||
|
language_str,
|
||||||
|
device,
|
||||||
|
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 = torch.zeros(1024, len(phone))
|
||||||
|
en_bert = torch.zeros(1024, len(phone))
|
||||||
|
elif language_str == Languages.JP:
|
||||||
|
bert = torch.zeros(1024, len(phone))
|
||||||
|
ja_bert = bert_ori
|
||||||
|
en_bert = torch.zeros(1024, len(phone))
|
||||||
|
elif language_str == Languages.EN:
|
||||||
|
bert = torch.zeros(1024, len(phone))
|
||||||
|
ja_bert = torch.zeros(1024, len(phone))
|
||||||
|
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 = torch.LongTensor(phone)
|
||||||
|
tone = torch.LongTensor(tone)
|
||||||
|
language = torch.LongTensor(language)
|
||||||
|
return bert, ja_bert, en_bert, phone, tone, language
|
||||||
|
|
||||||
|
|
||||||
|
def infer(
|
||||||
|
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,
|
||||||
|
net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra],
|
||||||
|
device: str,
|
||||||
|
skip_start: bool = False,
|
||||||
|
skip_end: bool = False,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
):
|
||||||
|
is_jp_extra = hps.version.endswith("JP-Extra")
|
||||||
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
|
text,
|
||||||
|
language,
|
||||||
|
hps,
|
||||||
|
device,
|
||||||
|
assist_text=assist_text,
|
||||||
|
assist_text_weight=assist_text_weight,
|
||||||
|
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]
|
||||||
|
with torch.no_grad():
|
||||||
|
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_vec).to(device).unsqueeze(0)
|
||||||
|
del phones
|
||||||
|
sid_tensor = torch.LongTensor([sid]).to(device)
|
||||||
|
if is_jp_extra:
|
||||||
|
output = cast(SynthesizerTrnJPExtra, net_g).infer(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
ja_bert,
|
||||||
|
style_vec=style_vec_tensor,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise_scale,
|
||||||
|
noise_scale_w=noise_scale_w,
|
||||||
|
length_scale=length_scale,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
output = cast(SynthesizerTrn, net_g).infer(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec=style_vec_tensor,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
noise_scale=noise_scale,
|
||||||
|
noise_scale_w=noise_scale_w,
|
||||||
|
length_scale=length_scale,
|
||||||
|
)
|
||||||
|
audio = output[0][0, 0].data.cpu().float().numpy()
|
||||||
|
del (
|
||||||
|
x_tst,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec,
|
||||||
|
) # , emo
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
return audio
|
||||||
|
|
||||||
|
|
||||||
|
class InvalidToneError(ValueError):
|
||||||
|
pass
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
import math
|
import math
|
||||||
import warnings
|
from typing import Any, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
@@ -7,18 +7,19 @@ from torch.nn import Conv1d, Conv2d, ConvTranspose1d
|
|||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
|
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
|
||||||
|
|
||||||
import attentions
|
from style_bert_vits2.models import attentions, commons, modules, monotonic_alignment
|
||||||
import commons
|
from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS
|
||||||
import modules
|
|
||||||
import monotonic_align
|
|
||||||
from commons import get_padding, init_weights
|
|
||||||
from text import num_languages, num_tones, symbols
|
|
||||||
|
|
||||||
|
|
||||||
class DurationDiscriminator(nn.Module): # vits2
|
class DurationDiscriminator(nn.Module): # vits2
|
||||||
def __init__(
|
def __init__(
|
||||||
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
|
self,
|
||||||
):
|
in_channels: int,
|
||||||
|
filter_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
p_dropout: float,
|
||||||
|
gin_channels: int = 0,
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -52,7 +53,13 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
|
|
||||||
self.output_layer = nn.Sequential(nn.Linear(filter_channels, 1), nn.Sigmoid())
|
self.output_layer = nn.Sequential(nn.Linear(filter_channels, 1), nn.Sigmoid())
|
||||||
|
|
||||||
def forward_probability(self, x, x_mask, dur, g=None):
|
def forward_probability(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
dur: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
dur = self.dur_proj(dur)
|
dur = self.dur_proj(dur)
|
||||||
x = torch.cat([x, dur], dim=1)
|
x = torch.cat([x, dur], dim=1)
|
||||||
x = self.pre_out_conv_1(x * x_mask)
|
x = self.pre_out_conv_1(x * x_mask)
|
||||||
@@ -68,7 +75,14 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
output_prob = self.output_layer(x)
|
output_prob = self.output_layer(x)
|
||||||
return output_prob
|
return output_prob
|
||||||
|
|
||||||
def forward(self, x, x_mask, dur_r, dur_hat, g=None):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
dur_r: torch.Tensor,
|
||||||
|
dur_hat: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> list[torch.Tensor]:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -93,17 +107,17 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
class TransformerCouplingBlock(nn.Module):
|
class TransformerCouplingBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
share_parameter=False,
|
share_parameter: bool = False,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -115,16 +129,17 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
self.flows = nn.ModuleList()
|
self.flows = nn.ModuleList()
|
||||||
|
|
||||||
self.wn = (
|
self.wn = (
|
||||||
attentions.FFT(
|
# attentions.FFT(
|
||||||
hidden_channels,
|
# hidden_channels,
|
||||||
filter_channels,
|
# filter_channels,
|
||||||
n_heads,
|
# n_heads,
|
||||||
n_layers,
|
# n_layers,
|
||||||
kernel_size,
|
# kernel_size,
|
||||||
p_dropout,
|
# p_dropout,
|
||||||
isflow=True,
|
# isflow=True,
|
||||||
gin_channels=self.gin_channels,
|
# gin_channels=self.gin_channels,
|
||||||
)
|
# )
|
||||||
|
None
|
||||||
if share_parameter
|
if share_parameter
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -146,7 +161,13 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
self.flows.append(modules.Flip())
|
self.flows.append(modules.Flip())
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
for flow in self.flows:
|
for flow in self.flows:
|
||||||
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
||||||
@@ -159,13 +180,13 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
class StochasticDurationPredictor(nn.Module):
|
class StochasticDurationPredictor(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
filter_channels = in_channels # it needs to be removed from future version.
|
filter_channels = in_channels # it needs to be removed from future version.
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -205,7 +226,15 @@ class StochasticDurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, filter_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, filter_channels, 1)
|
||||||
|
|
||||||
def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
w: Optional[torch.Tensor] = None,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
noise_scale: float = 1.0,
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
x = self.pre(x)
|
x = self.pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
@@ -269,8 +298,13 @@ class StochasticDurationPredictor(nn.Module):
|
|||||||
|
|
||||||
class DurationPredictor(nn.Module):
|
class DurationPredictor(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
|
self,
|
||||||
):
|
in_channels: int,
|
||||||
|
filter_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
p_dropout: float,
|
||||||
|
gin_channels: int = 0,
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -293,7 +327,9 @@ class DurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -313,17 +349,17 @@ class DurationPredictor(nn.Module):
|
|||||||
class TextEncoder(nn.Module):
|
class TextEncoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_vocab,
|
n_vocab: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
n_speakers,
|
n_speakers: int,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.n_vocab = n_vocab
|
self.n_vocab = n_vocab
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
@@ -334,11 +370,11 @@ class TextEncoder(nn.Module):
|
|||||||
self.kernel_size = kernel_size
|
self.kernel_size = kernel_size
|
||||||
self.p_dropout = p_dropout
|
self.p_dropout = p_dropout
|
||||||
self.gin_channels = gin_channels
|
self.gin_channels = gin_channels
|
||||||
self.emb = nn.Embedding(len(symbols), hidden_channels)
|
self.emb = nn.Embedding(len(SYMBOLS), hidden_channels)
|
||||||
nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.tone_emb = nn.Embedding(num_tones, hidden_channels)
|
self.tone_emb = nn.Embedding(NUM_TONES, hidden_channels)
|
||||||
nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.language_emb = nn.Embedding(num_languages, hidden_channels)
|
self.language_emb = nn.Embedding(NUM_LANGUAGES, hidden_channels)
|
||||||
nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
||||||
self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
||||||
@@ -358,17 +394,17 @@ class TextEncoder(nn.Module):
|
|||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
x,
|
x: torch.Tensor,
|
||||||
x_lengths,
|
x_lengths: torch.Tensor,
|
||||||
tone,
|
tone: torch.Tensor,
|
||||||
language,
|
language: torch.Tensor,
|
||||||
bert,
|
bert: torch.Tensor,
|
||||||
ja_bert,
|
ja_bert: torch.Tensor,
|
||||||
en_bert,
|
en_bert: torch.Tensor,
|
||||||
style_vec,
|
style_vec: torch.Tensor,
|
||||||
sid,
|
sid: torch.Tensor,
|
||||||
g=None,
|
g: Optional[torch.Tensor] = None,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
bert_emb = self.bert_proj(bert).transpose(1, 2)
|
bert_emb = self.bert_proj(bert).transpose(1, 2)
|
||||||
ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
|
ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
|
||||||
en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
|
en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
|
||||||
@@ -400,14 +436,14 @@ class TextEncoder(nn.Module):
|
|||||||
class ResidualCouplingBlock(nn.Module):
|
class ResidualCouplingBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -432,7 +468,13 @@ class ResidualCouplingBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
self.flows.append(modules.Flip())
|
self.flows.append(modules.Flip())
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
for flow in self.flows:
|
for flow in self.flows:
|
||||||
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
||||||
@@ -445,14 +487,14 @@ class ResidualCouplingBlock(nn.Module):
|
|||||||
class PosteriorEncoder(nn.Module):
|
class PosteriorEncoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
@@ -472,7 +514,12 @@ class PosteriorEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
||||||
|
|
||||||
def forward(self, x, x_lengths, g=None):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
|
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
|
||||||
x.dtype
|
x.dtype
|
||||||
)
|
)
|
||||||
@@ -487,22 +534,22 @@ class PosteriorEncoder(nn.Module):
|
|||||||
class Generator(torch.nn.Module):
|
class Generator(torch.nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
initial_channel,
|
initial_channel: int,
|
||||||
resblock,
|
resblock_str: str,
|
||||||
resblock_kernel_sizes,
|
resblock_kernel_sizes: list[int],
|
||||||
resblock_dilation_sizes,
|
resblock_dilation_sizes: list[list[int]],
|
||||||
upsample_rates,
|
upsample_rates: list[int],
|
||||||
upsample_initial_channel,
|
upsample_initial_channel: int,
|
||||||
upsample_kernel_sizes,
|
upsample_kernel_sizes: list[int],
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super(Generator, self).__init__()
|
super(Generator, self).__init__()
|
||||||
self.num_kernels = len(resblock_kernel_sizes)
|
self.num_kernels = len(resblock_kernel_sizes)
|
||||||
self.num_upsamples = len(upsample_rates)
|
self.num_upsamples = len(upsample_rates)
|
||||||
self.conv_pre = Conv1d(
|
self.conv_pre = Conv1d(
|
||||||
initial_channel, upsample_initial_channel, 7, 1, padding=3
|
initial_channel, upsample_initial_channel, 7, 1, padding=3
|
||||||
)
|
)
|
||||||
resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
|
resblock = modules.ResBlock1 if resblock_str == "1" else modules.ResBlock2
|
||||||
|
|
||||||
self.ups = nn.ModuleList()
|
self.ups = nn.ModuleList()
|
||||||
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
||||||
@@ -519,20 +566,24 @@ class Generator(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.resblocks = nn.ModuleList()
|
self.resblocks = nn.ModuleList()
|
||||||
|
ch = None
|
||||||
for i in range(len(self.ups)):
|
for i in range(len(self.ups)):
|
||||||
ch = upsample_initial_channel // (2 ** (i + 1))
|
ch = upsample_initial_channel // (2 ** (i + 1))
|
||||||
for j, (k, d) in enumerate(
|
for j, (k, d) in enumerate(
|
||||||
zip(resblock_kernel_sizes, resblock_dilation_sizes)
|
zip(resblock_kernel_sizes, resblock_dilation_sizes)
|
||||||
):
|
):
|
||||||
self.resblocks.append(resblock(ch, k, d))
|
self.resblocks.append(resblock(ch, k, d)) # type: ignore
|
||||||
|
|
||||||
|
assert ch is not None
|
||||||
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
|
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
|
||||||
self.ups.apply(init_weights)
|
self.ups.apply(commons.init_weights)
|
||||||
|
|
||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
||||||
|
|
||||||
def forward(self, x, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = self.conv_pre(x)
|
x = self.conv_pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + self.cond(g)
|
x = x + self.cond(g)
|
||||||
@@ -546,6 +597,7 @@ class Generator(torch.nn.Module):
|
|||||||
xs = self.resblocks[i * self.num_kernels + j](x)
|
xs = self.resblocks[i * self.num_kernels + j](x)
|
||||||
else:
|
else:
|
||||||
xs += self.resblocks[i * self.num_kernels + j](x)
|
xs += self.resblocks[i * self.num_kernels + j](x)
|
||||||
|
assert xs is not None
|
||||||
x = xs / self.num_kernels
|
x = xs / self.num_kernels
|
||||||
x = F.leaky_relu(x)
|
x = F.leaky_relu(x)
|
||||||
x = self.conv_post(x)
|
x = self.conv_post(x)
|
||||||
@@ -553,7 +605,7 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def remove_weight_norm(self):
|
def remove_weight_norm(self) -> None:
|
||||||
print("Removing weight norm...")
|
print("Removing weight norm...")
|
||||||
for layer in self.ups:
|
for layer in self.ups:
|
||||||
remove_weight_norm(layer)
|
remove_weight_norm(layer)
|
||||||
@@ -562,7 +614,13 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorP(torch.nn.Module):
|
class DiscriminatorP(torch.nn.Module):
|
||||||
def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
|
def __init__(
|
||||||
|
self,
|
||||||
|
period: int,
|
||||||
|
kernel_size: int = 5,
|
||||||
|
stride: int = 3,
|
||||||
|
use_spectral_norm: bool = False,
|
||||||
|
) -> None:
|
||||||
super(DiscriminatorP, self).__init__()
|
super(DiscriminatorP, self).__init__()
|
||||||
self.period = period
|
self.period = period
|
||||||
self.use_spectral_norm = use_spectral_norm
|
self.use_spectral_norm = use_spectral_norm
|
||||||
@@ -575,7 +633,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
32,
|
32,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -584,7 +642,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
128,
|
128,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -593,7 +651,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
512,
|
512,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -602,7 +660,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
1024,
|
1024,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -611,14 +669,14 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
1024,
|
1024,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
1,
|
1,
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
|
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
|
||||||
fmap = []
|
fmap = []
|
||||||
|
|
||||||
# 1d to 2d
|
# 1d to 2d
|
||||||
@@ -641,7 +699,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorS(torch.nn.Module):
|
class DiscriminatorS(torch.nn.Module):
|
||||||
def __init__(self, use_spectral_norm=False):
|
def __init__(self, use_spectral_norm: bool = False) -> None:
|
||||||
super(DiscriminatorS, self).__init__()
|
super(DiscriminatorS, self).__init__()
|
||||||
norm_f = weight_norm if use_spectral_norm is False else spectral_norm
|
norm_f = weight_norm if use_spectral_norm is False else spectral_norm
|
||||||
self.convs = nn.ModuleList(
|
self.convs = nn.ModuleList(
|
||||||
@@ -656,7 +714,7 @@ class DiscriminatorS(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
|
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
|
||||||
fmap = []
|
fmap = []
|
||||||
|
|
||||||
for layer in self.convs:
|
for layer in self.convs:
|
||||||
@@ -671,7 +729,7 @@ class DiscriminatorS(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class MultiPeriodDiscriminator(torch.nn.Module):
|
class MultiPeriodDiscriminator(torch.nn.Module):
|
||||||
def __init__(self, use_spectral_norm=False):
|
def __init__(self, use_spectral_norm: bool = False) -> None:
|
||||||
super(MultiPeriodDiscriminator, self).__init__()
|
super(MultiPeriodDiscriminator, self).__init__()
|
||||||
periods = [2, 3, 5, 7, 11]
|
periods = [2, 3, 5, 7, 11]
|
||||||
|
|
||||||
@@ -681,7 +739,13 @@ class MultiPeriodDiscriminator(torch.nn.Module):
|
|||||||
]
|
]
|
||||||
self.discriminators = nn.ModuleList(discs)
|
self.discriminators = nn.ModuleList(discs)
|
||||||
|
|
||||||
def forward(self, y, y_hat):
|
def forward(
|
||||||
|
self,
|
||||||
|
y: torch.Tensor,
|
||||||
|
y_hat: torch.Tensor,
|
||||||
|
) -> tuple[
|
||||||
|
list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]
|
||||||
|
]:
|
||||||
y_d_rs = []
|
y_d_rs = []
|
||||||
y_d_gs = []
|
y_d_gs = []
|
||||||
fmap_rs = []
|
fmap_rs = []
|
||||||
@@ -703,7 +767,7 @@ class ReferenceEncoder(nn.Module):
|
|||||||
outputs --- [N, ref_enc_gru_size]
|
outputs --- [N, ref_enc_gru_size]
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, spec_channels, gin_channels=0):
|
def __init__(self, spec_channels: int, gin_channels: int = 0) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.spec_channels = spec_channels
|
self.spec_channels = spec_channels
|
||||||
ref_enc_filters = [32, 32, 64, 64, 128, 128]
|
ref_enc_filters = [32, 32, 64, 64, 128, 128]
|
||||||
@@ -732,7 +796,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Linear(128, gin_channels)
|
self.proj = nn.Linear(128, gin_channels)
|
||||||
|
|
||||||
def forward(self, inputs, mask=None):
|
def forward(
|
||||||
|
self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
N = inputs.size(0)
|
N = inputs.size(0)
|
||||||
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
||||||
for conv in self.convs:
|
for conv in self.convs:
|
||||||
@@ -750,7 +816,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
|
|
||||||
return self.proj(out.squeeze(0))
|
return self.proj(out.squeeze(0))
|
||||||
|
|
||||||
def calculate_channels(self, L, kernel_size, stride, pad, n_convs):
|
def calculate_channels(
|
||||||
|
self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int
|
||||||
|
) -> int:
|
||||||
for i in range(n_convs):
|
for i in range(n_convs):
|
||||||
L = (L - kernel_size + 2 * pad) // stride + 1
|
L = (L - kernel_size + 2 * pad) // stride + 1
|
||||||
return L
|
return L
|
||||||
@@ -763,31 +831,31 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_vocab,
|
n_vocab: int,
|
||||||
spec_channels,
|
spec_channels: int,
|
||||||
segment_size,
|
segment_size: int,
|
||||||
inter_channels,
|
inter_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
resblock,
|
resblock: str,
|
||||||
resblock_kernel_sizes,
|
resblock_kernel_sizes: list[int],
|
||||||
resblock_dilation_sizes,
|
resblock_dilation_sizes: list[list[int]],
|
||||||
upsample_rates,
|
upsample_rates: list[int],
|
||||||
upsample_initial_channel,
|
upsample_initial_channel: int,
|
||||||
upsample_kernel_sizes,
|
upsample_kernel_sizes: list[int],
|
||||||
n_speakers=256,
|
n_speakers: int = 256,
|
||||||
gin_channels=256,
|
gin_channels: int = 256,
|
||||||
use_sdp=True,
|
use_sdp: bool = True,
|
||||||
n_flow_layer=4,
|
n_flow_layer: int = 4,
|
||||||
n_layers_trans_flow=4,
|
n_layers_trans_flow: int = 4,
|
||||||
flow_share_parameter=False,
|
flow_share_parameter: bool = False,
|
||||||
use_transformer_flow=True,
|
use_transformer_flow: bool = True,
|
||||||
**kwargs,
|
**kwargs: Any,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.n_vocab = n_vocab
|
self.n_vocab = n_vocab
|
||||||
self.spec_channels = spec_channels
|
self.spec_channels = spec_channels
|
||||||
@@ -885,18 +953,27 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
x,
|
x: torch.Tensor,
|
||||||
x_lengths,
|
x_lengths: torch.Tensor,
|
||||||
y,
|
y: torch.Tensor,
|
||||||
y_lengths,
|
y_lengths: torch.Tensor,
|
||||||
sid,
|
sid: torch.Tensor,
|
||||||
tone,
|
tone: torch.Tensor,
|
||||||
language,
|
language: torch.Tensor,
|
||||||
bert,
|
bert: torch.Tensor,
|
||||||
ja_bert,
|
ja_bert: torch.Tensor,
|
||||||
en_bert,
|
en_bert: torch.Tensor,
|
||||||
style_vec,
|
style_vec: torch.Tensor,
|
||||||
):
|
) -> tuple[
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
tuple[torch.Tensor, ...],
|
||||||
|
tuple[torch.Tensor, ...],
|
||||||
|
]:
|
||||||
if self.n_speakers > 0:
|
if self.n_speakers > 0:
|
||||||
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
||||||
else:
|
else:
|
||||||
@@ -933,7 +1010,7 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
|
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
|
||||||
attn = (
|
attn = (
|
||||||
monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1))
|
monotonic_alignment.maximum_path(neg_cent, attn_mask.squeeze(1))
|
||||||
.unsqueeze(1)
|
.unsqueeze(1)
|
||||||
.detach()
|
.detach()
|
||||||
)
|
)
|
||||||
@@ -974,27 +1051,28 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
def infer(
|
def infer(
|
||||||
self,
|
self,
|
||||||
x,
|
x: torch.Tensor,
|
||||||
x_lengths,
|
x_lengths: torch.Tensor,
|
||||||
sid,
|
sid: torch.Tensor,
|
||||||
tone,
|
tone: torch.Tensor,
|
||||||
language,
|
language: torch.Tensor,
|
||||||
bert,
|
bert: torch.Tensor,
|
||||||
ja_bert,
|
ja_bert: torch.Tensor,
|
||||||
en_bert,
|
en_bert: torch.Tensor,
|
||||||
style_vec,
|
style_vec: torch.Tensor,
|
||||||
noise_scale=0.667,
|
noise_scale: float = 0.667,
|
||||||
length_scale=1,
|
length_scale: float = 1.0,
|
||||||
noise_scale_w=0.8,
|
noise_scale_w: float = 0.8,
|
||||||
max_len=None,
|
max_len: Optional[int] = None,
|
||||||
sdp_ratio=0,
|
sdp_ratio: float = 0.0,
|
||||||
y=None,
|
y: Optional[torch.Tensor] = None,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]:
|
||||||
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)
|
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)
|
||||||
# g = self.gst(y)
|
# g = self.gst(y)
|
||||||
if self.n_speakers > 0:
|
if self.n_speakers > 0:
|
||||||
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
||||||
else:
|
else:
|
||||||
|
assert y is not None
|
||||||
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
|
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
|
||||||
x, m_p, logs_p, x_mask = self.enc_p(
|
x, m_p, logs_p, x_mask = self.enc_p(
|
||||||
x, x_lengths, tone, language, bert, ja_bert, en_bert, style_vec, sid, g=g
|
x, x_lengths, tone, language, bert, ja_bert, en_bert, style_vec, sid, g=g
|
||||||
@@ -1,24 +1,25 @@
|
|||||||
import math
|
import math
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
from torch.nn import Conv1d, Conv2d, ConvTranspose1d
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
|
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
|
||||||
|
|
||||||
import commons
|
from style_bert_vits2.models import attentions, commons, modules, monotonic_alignment
|
||||||
import modules
|
from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS
|
||||||
import attentions
|
|
||||||
import monotonic_align
|
|
||||||
|
|
||||||
from torch.nn import Conv1d, ConvTranspose1d, Conv2d
|
|
||||||
from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm
|
|
||||||
|
|
||||||
from commons import init_weights, get_padding
|
|
||||||
from text import symbols, num_tones, num_languages
|
|
||||||
|
|
||||||
|
|
||||||
class DurationDiscriminator(nn.Module): # vits2
|
class DurationDiscriminator(nn.Module): # vits2
|
||||||
def __init__(
|
def __init__(
|
||||||
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
|
self,
|
||||||
):
|
in_channels: int,
|
||||||
|
filter_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
p_dropout: float,
|
||||||
|
gin_channels: int = 0,
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -49,7 +50,7 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
nn.Linear(2 * filter_channels, 1), nn.Sigmoid()
|
nn.Linear(2 * filter_channels, 1), nn.Sigmoid()
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_probability(self, x, dur):
|
def forward_probability(self, x: torch.Tensor, dur: torch.Tensor) -> torch.Tensor:
|
||||||
dur = self.dur_proj(dur)
|
dur = self.dur_proj(dur)
|
||||||
x = torch.cat([x, dur], dim=1)
|
x = torch.cat([x, dur], dim=1)
|
||||||
x = x.transpose(1, 2)
|
x = x.transpose(1, 2)
|
||||||
@@ -57,7 +58,14 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
output_prob = self.output_layer(x)
|
output_prob = self.output_layer(x)
|
||||||
return output_prob
|
return output_prob
|
||||||
|
|
||||||
def forward(self, x, x_mask, dur_r, dur_hat, g=None):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
dur_r: torch.Tensor,
|
||||||
|
dur_hat: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> list[torch.Tensor]:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -82,17 +90,17 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
class TransformerCouplingBlock(nn.Module):
|
class TransformerCouplingBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
share_parameter=False,
|
share_parameter: bool = False,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -104,16 +112,17 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
self.flows = nn.ModuleList()
|
self.flows = nn.ModuleList()
|
||||||
|
|
||||||
self.wn = (
|
self.wn = (
|
||||||
attentions.FFT(
|
# attentions.FFT(
|
||||||
hidden_channels,
|
# hidden_channels,
|
||||||
filter_channels,
|
# filter_channels,
|
||||||
n_heads,
|
# n_heads,
|
||||||
n_layers,
|
# n_layers,
|
||||||
kernel_size,
|
# kernel_size,
|
||||||
p_dropout,
|
# p_dropout,
|
||||||
isflow=True,
|
# isflow=True,
|
||||||
gin_channels=self.gin_channels,
|
# gin_channels=self.gin_channels,
|
||||||
)
|
# )
|
||||||
|
None
|
||||||
if share_parameter
|
if share_parameter
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -135,7 +144,13 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
self.flows.append(modules.Flip())
|
self.flows.append(modules.Flip())
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
for flow in self.flows:
|
for flow in self.flows:
|
||||||
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
||||||
@@ -148,13 +163,13 @@ class TransformerCouplingBlock(nn.Module):
|
|||||||
class StochasticDurationPredictor(nn.Module):
|
class StochasticDurationPredictor(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
filter_channels = in_channels # it needs to be removed from future version.
|
filter_channels = in_channels # it needs to be removed from future version.
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -194,7 +209,15 @@ class StochasticDurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, filter_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, filter_channels, 1)
|
||||||
|
|
||||||
def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
w: Optional[torch.Tensor] = None,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
noise_scale: float = 1.0,
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
x = self.pre(x)
|
x = self.pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
@@ -258,8 +281,13 @@ class StochasticDurationPredictor(nn.Module):
|
|||||||
|
|
||||||
class DurationPredictor(nn.Module):
|
class DurationPredictor(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
|
self,
|
||||||
):
|
in_channels: int,
|
||||||
|
filter_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
p_dropout: float,
|
||||||
|
gin_channels: int = 0,
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
@@ -282,7 +310,9 @@ class DurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -300,14 +330,14 @@ class DurationPredictor(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class Bottleneck(nn.Sequential):
|
class Bottleneck(nn.Sequential):
|
||||||
def __init__(self, in_dim, hidden_dim):
|
def __init__(self, in_dim: int, hidden_dim: int) -> None:
|
||||||
c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
|
c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
|
||||||
c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
|
c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
|
||||||
super().__init__(*[c_fc1, c_fc2])
|
super().__init__(c_fc1, c_fc2)
|
||||||
|
|
||||||
|
|
||||||
class Block(nn.Module):
|
class Block(nn.Module):
|
||||||
def __init__(self, in_dim, hidden_dim) -> None:
|
def __init__(self, in_dim: int, hidden_dim: int) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.norm = nn.LayerNorm(in_dim)
|
self.norm = nn.LayerNorm(in_dim)
|
||||||
self.mlp = MLP(in_dim, hidden_dim)
|
self.mlp = MLP(in_dim, hidden_dim)
|
||||||
@@ -318,13 +348,13 @@ class Block(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
class MLP(nn.Module):
|
||||||
def __init__(self, in_dim, hidden_dim):
|
def __init__(self, in_dim: int, hidden_dim: int) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
|
self.c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False)
|
||||||
self.c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
|
self.c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False)
|
||||||
self.c_proj = nn.Linear(hidden_dim, in_dim, bias=False)
|
self.c_proj = nn.Linear(hidden_dim, in_dim, bias=False)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
x = F.silu(self.c_fc1(x)) * self.c_fc2(x)
|
x = F.silu(self.c_fc1(x)) * self.c_fc2(x)
|
||||||
x = self.c_proj(x)
|
x = self.c_proj(x)
|
||||||
return x
|
return x
|
||||||
@@ -333,16 +363,16 @@ class MLP(nn.Module):
|
|||||||
class TextEncoder(nn.Module):
|
class TextEncoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_vocab,
|
n_vocab: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.n_vocab = n_vocab
|
self.n_vocab = n_vocab
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
@@ -353,11 +383,11 @@ class TextEncoder(nn.Module):
|
|||||||
self.kernel_size = kernel_size
|
self.kernel_size = kernel_size
|
||||||
self.p_dropout = p_dropout
|
self.p_dropout = p_dropout
|
||||||
self.gin_channels = gin_channels
|
self.gin_channels = gin_channels
|
||||||
self.emb = nn.Embedding(len(symbols), hidden_channels)
|
self.emb = nn.Embedding(len(SYMBOLS), hidden_channels)
|
||||||
nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.tone_emb = nn.Embedding(num_tones, hidden_channels)
|
self.tone_emb = nn.Embedding(NUM_TONES, hidden_channels)
|
||||||
nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.language_emb = nn.Embedding(num_languages, hidden_channels)
|
self.language_emb = nn.Embedding(NUM_LANGUAGES, hidden_channels)
|
||||||
nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5)
|
nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5)
|
||||||
self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
|
||||||
|
|
||||||
@@ -375,7 +405,16 @@ class TextEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
||||||
|
|
||||||
def forward(self, x, x_lengths, tone, language, bert, style_vec, g=None):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
tone: torch.Tensor,
|
||||||
|
language: torch.Tensor,
|
||||||
|
bert: torch.Tensor,
|
||||||
|
style_vec: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
bert_emb = self.bert_proj(bert).transpose(1, 2)
|
bert_emb = self.bert_proj(bert).transpose(1, 2)
|
||||||
style_emb = self.style_proj(style_vec.unsqueeze(1))
|
style_emb = self.style_proj(style_vec.unsqueeze(1))
|
||||||
x = (
|
x = (
|
||||||
@@ -402,14 +441,14 @@ class TextEncoder(nn.Module):
|
|||||||
class ResidualCouplingBlock(nn.Module):
|
class ResidualCouplingBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
n_flows=4,
|
n_flows: int = 4,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -434,7 +473,13 @@ class ResidualCouplingBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
self.flows.append(modules.Flip())
|
self.flows.append(modules.Flip())
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
for flow in self.flows:
|
for flow in self.flows:
|
||||||
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
x, _ = flow(x, x_mask, g=g, reverse=reverse)
|
||||||
@@ -447,14 +492,14 @@ class ResidualCouplingBlock(nn.Module):
|
|||||||
class PosteriorEncoder(nn.Module):
|
class PosteriorEncoder(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
@@ -474,7 +519,12 @@ class PosteriorEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
|
||||||
|
|
||||||
def forward(self, x, x_lengths, g=None):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
|
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
|
||||||
x.dtype
|
x.dtype
|
||||||
)
|
)
|
||||||
@@ -489,22 +539,22 @@ class PosteriorEncoder(nn.Module):
|
|||||||
class Generator(torch.nn.Module):
|
class Generator(torch.nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
initial_channel,
|
initial_channel: int,
|
||||||
resblock,
|
resblock_str: str,
|
||||||
resblock_kernel_sizes,
|
resblock_kernel_sizes: list[int],
|
||||||
resblock_dilation_sizes,
|
resblock_dilation_sizes: list[list[int]],
|
||||||
upsample_rates,
|
upsample_rates: list[int],
|
||||||
upsample_initial_channel,
|
upsample_initial_channel: int,
|
||||||
upsample_kernel_sizes,
|
upsample_kernel_sizes: list[int],
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
super(Generator, self).__init__()
|
super(Generator, self).__init__()
|
||||||
self.num_kernels = len(resblock_kernel_sizes)
|
self.num_kernels = len(resblock_kernel_sizes)
|
||||||
self.num_upsamples = len(upsample_rates)
|
self.num_upsamples = len(upsample_rates)
|
||||||
self.conv_pre = Conv1d(
|
self.conv_pre = Conv1d(
|
||||||
initial_channel, upsample_initial_channel, 7, 1, padding=3
|
initial_channel, upsample_initial_channel, 7, 1, padding=3
|
||||||
)
|
)
|
||||||
resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
|
resblock = modules.ResBlock1 if resblock_str == "1" else modules.ResBlock2
|
||||||
|
|
||||||
self.ups = nn.ModuleList()
|
self.ups = nn.ModuleList()
|
||||||
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
||||||
@@ -521,20 +571,24 @@ class Generator(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.resblocks = nn.ModuleList()
|
self.resblocks = nn.ModuleList()
|
||||||
|
ch = None
|
||||||
for i in range(len(self.ups)):
|
for i in range(len(self.ups)):
|
||||||
ch = upsample_initial_channel // (2 ** (i + 1))
|
ch = upsample_initial_channel // (2 ** (i + 1))
|
||||||
for j, (k, d) in enumerate(
|
for j, (k, d) in enumerate(
|
||||||
zip(resblock_kernel_sizes, resblock_dilation_sizes)
|
zip(resblock_kernel_sizes, resblock_dilation_sizes)
|
||||||
):
|
):
|
||||||
self.resblocks.append(resblock(ch, k, d))
|
self.resblocks.append(resblock(ch, k, d)) # type: ignore
|
||||||
|
|
||||||
|
assert ch is not None
|
||||||
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
|
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
|
||||||
self.ups.apply(init_weights)
|
self.ups.apply(commons.init_weights)
|
||||||
|
|
||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
||||||
|
|
||||||
def forward(self, x, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = self.conv_pre(x)
|
x = self.conv_pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + self.cond(g)
|
x = x + self.cond(g)
|
||||||
@@ -548,6 +602,7 @@ class Generator(torch.nn.Module):
|
|||||||
xs = self.resblocks[i * self.num_kernels + j](x)
|
xs = self.resblocks[i * self.num_kernels + j](x)
|
||||||
else:
|
else:
|
||||||
xs += self.resblocks[i * self.num_kernels + j](x)
|
xs += self.resblocks[i * self.num_kernels + j](x)
|
||||||
|
assert xs is not None
|
||||||
x = xs / self.num_kernels
|
x = xs / self.num_kernels
|
||||||
x = F.leaky_relu(x)
|
x = F.leaky_relu(x)
|
||||||
x = self.conv_post(x)
|
x = self.conv_post(x)
|
||||||
@@ -555,7 +610,7 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def remove_weight_norm(self):
|
def remove_weight_norm(self) -> None:
|
||||||
print("Removing weight norm...")
|
print("Removing weight norm...")
|
||||||
for layer in self.ups:
|
for layer in self.ups:
|
||||||
remove_weight_norm(layer)
|
remove_weight_norm(layer)
|
||||||
@@ -564,7 +619,13 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorP(torch.nn.Module):
|
class DiscriminatorP(torch.nn.Module):
|
||||||
def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
|
def __init__(
|
||||||
|
self,
|
||||||
|
period: int,
|
||||||
|
kernel_size: int = 5,
|
||||||
|
stride: int = 3,
|
||||||
|
use_spectral_norm: bool = False,
|
||||||
|
) -> None:
|
||||||
super(DiscriminatorP, self).__init__()
|
super(DiscriminatorP, self).__init__()
|
||||||
self.period = period
|
self.period = period
|
||||||
self.use_spectral_norm = use_spectral_norm
|
self.use_spectral_norm = use_spectral_norm
|
||||||
@@ -577,7 +638,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
32,
|
32,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -586,7 +647,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
128,
|
128,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -595,7 +656,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
512,
|
512,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -604,7 +665,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
1024,
|
1024,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
(stride, 1),
|
(stride, 1),
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
norm_f(
|
norm_f(
|
||||||
@@ -613,14 +674,14 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
1024,
|
1024,
|
||||||
(kernel_size, 1),
|
(kernel_size, 1),
|
||||||
1,
|
1,
|
||||||
padding=(get_padding(kernel_size, 1), 0),
|
padding=(commons.get_padding(kernel_size, 1), 0),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
|
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
|
||||||
fmap = []
|
fmap = []
|
||||||
|
|
||||||
# 1d to 2d
|
# 1d to 2d
|
||||||
@@ -643,7 +704,7 @@ class DiscriminatorP(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorS(torch.nn.Module):
|
class DiscriminatorS(torch.nn.Module):
|
||||||
def __init__(self, use_spectral_norm=False):
|
def __init__(self, use_spectral_norm: bool = False) -> None:
|
||||||
super(DiscriminatorS, self).__init__()
|
super(DiscriminatorS, self).__init__()
|
||||||
norm_f = weight_norm if use_spectral_norm is False else spectral_norm
|
norm_f = weight_norm if use_spectral_norm is False else spectral_norm
|
||||||
self.convs = nn.ModuleList(
|
self.convs = nn.ModuleList(
|
||||||
@@ -658,7 +719,7 @@ class DiscriminatorS(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
|
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
|
||||||
fmap = []
|
fmap = []
|
||||||
|
|
||||||
for layer in self.convs:
|
for layer in self.convs:
|
||||||
@@ -673,7 +734,7 @@ class DiscriminatorS(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class MultiPeriodDiscriminator(torch.nn.Module):
|
class MultiPeriodDiscriminator(torch.nn.Module):
|
||||||
def __init__(self, use_spectral_norm=False):
|
def __init__(self, use_spectral_norm: bool = False) -> None:
|
||||||
super(MultiPeriodDiscriminator, self).__init__()
|
super(MultiPeriodDiscriminator, self).__init__()
|
||||||
periods = [2, 3, 5, 7, 11]
|
periods = [2, 3, 5, 7, 11]
|
||||||
|
|
||||||
@@ -683,7 +744,13 @@ class MultiPeriodDiscriminator(torch.nn.Module):
|
|||||||
]
|
]
|
||||||
self.discriminators = nn.ModuleList(discs)
|
self.discriminators = nn.ModuleList(discs)
|
||||||
|
|
||||||
def forward(self, y, y_hat):
|
def forward(
|
||||||
|
self,
|
||||||
|
y: torch.Tensor,
|
||||||
|
y_hat: torch.Tensor,
|
||||||
|
) -> tuple[
|
||||||
|
list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]
|
||||||
|
]:
|
||||||
y_d_rs = []
|
y_d_rs = []
|
||||||
y_d_gs = []
|
y_d_gs = []
|
||||||
fmap_rs = []
|
fmap_rs = []
|
||||||
@@ -703,8 +770,12 @@ class WavLMDiscriminator(nn.Module):
|
|||||||
"""docstring for Discriminator."""
|
"""docstring for Discriminator."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, slm_hidden=768, slm_layers=13, initial_channel=64, use_spectral_norm=False
|
self,
|
||||||
):
|
slm_hidden: int = 768,
|
||||||
|
slm_layers: int = 13,
|
||||||
|
initial_channel: int = 64,
|
||||||
|
use_spectral_norm: bool = False,
|
||||||
|
) -> None:
|
||||||
super(WavLMDiscriminator, self).__init__()
|
super(WavLMDiscriminator, self).__init__()
|
||||||
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
||||||
self.pre = norm_f(
|
self.pre = norm_f(
|
||||||
@@ -734,7 +805,7 @@ class WavLMDiscriminator(nn.Module):
|
|||||||
|
|
||||||
self.conv_post = norm_f(Conv1d(initial_channel * 4, 1, 3, 1, padding=1))
|
self.conv_post = norm_f(Conv1d(initial_channel * 4, 1, 3, 1, padding=1))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
x = self.pre(x)
|
x = self.pre(x)
|
||||||
|
|
||||||
fmap = []
|
fmap = []
|
||||||
@@ -754,7 +825,7 @@ class ReferenceEncoder(nn.Module):
|
|||||||
outputs --- [N, ref_enc_gru_size]
|
outputs --- [N, ref_enc_gru_size]
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, spec_channels, gin_channels=0):
|
def __init__(self, spec_channels: int, gin_channels: int = 0) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.spec_channels = spec_channels
|
self.spec_channels = spec_channels
|
||||||
ref_enc_filters = [32, 32, 64, 64, 128, 128]
|
ref_enc_filters = [32, 32, 64, 64, 128, 128]
|
||||||
@@ -783,7 +854,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Linear(128, gin_channels)
|
self.proj = nn.Linear(128, gin_channels)
|
||||||
|
|
||||||
def forward(self, inputs, mask=None):
|
def forward(
|
||||||
|
self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
N = inputs.size(0)
|
N = inputs.size(0)
|
||||||
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
||||||
for conv in self.convs:
|
for conv in self.convs:
|
||||||
@@ -801,7 +874,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
|
|
||||||
return self.proj(out.squeeze(0))
|
return self.proj(out.squeeze(0))
|
||||||
|
|
||||||
def calculate_channels(self, L, kernel_size, stride, pad, n_convs):
|
def calculate_channels(
|
||||||
|
self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int
|
||||||
|
) -> int:
|
||||||
for i in range(n_convs):
|
for i in range(n_convs):
|
||||||
L = (L - kernel_size + 2 * pad) // stride + 1
|
L = (L - kernel_size + 2 * pad) // stride + 1
|
||||||
return L
|
return L
|
||||||
@@ -814,31 +889,31 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_vocab,
|
n_vocab: int,
|
||||||
spec_channels,
|
spec_channels: int,
|
||||||
segment_size,
|
segment_size: int,
|
||||||
inter_channels,
|
inter_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
resblock,
|
resblock: str,
|
||||||
resblock_kernel_sizes,
|
resblock_kernel_sizes: list[int],
|
||||||
resblock_dilation_sizes,
|
resblock_dilation_sizes: list[list[int]],
|
||||||
upsample_rates,
|
upsample_rates: list[int],
|
||||||
upsample_initial_channel,
|
upsample_initial_channel: int,
|
||||||
upsample_kernel_sizes,
|
upsample_kernel_sizes: list[int],
|
||||||
n_speakers=256,
|
n_speakers: int = 256,
|
||||||
gin_channels=256,
|
gin_channels: int = 256,
|
||||||
use_sdp=True,
|
use_sdp: bool = True,
|
||||||
n_flow_layer=4,
|
n_flow_layer: int = 4,
|
||||||
n_layers_trans_flow=6,
|
n_layers_trans_flow: int = 6,
|
||||||
flow_share_parameter=False,
|
flow_share_parameter: bool = False,
|
||||||
use_transformer_flow=True,
|
use_transformer_flow: bool = True,
|
||||||
**kwargs
|
**kwargs: Any,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.n_vocab = n_vocab
|
self.n_vocab = n_vocab
|
||||||
self.spec_channels = spec_channels
|
self.spec_channels = spec_channels
|
||||||
@@ -935,16 +1010,26 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
x,
|
x: torch.Tensor,
|
||||||
x_lengths,
|
x_lengths: torch.Tensor,
|
||||||
y,
|
y: torch.Tensor,
|
||||||
y_lengths,
|
y_lengths: torch.Tensor,
|
||||||
sid,
|
sid: torch.Tensor,
|
||||||
tone,
|
tone: torch.Tensor,
|
||||||
language,
|
language: torch.Tensor,
|
||||||
bert,
|
bert: torch.Tensor,
|
||||||
style_vec,
|
style_vec: torch.Tensor,
|
||||||
):
|
) -> tuple[
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
torch.Tensor,
|
||||||
|
tuple[torch.Tensor, ...],
|
||||||
|
tuple[torch.Tensor, ...],
|
||||||
|
]:
|
||||||
if self.n_speakers > 0:
|
if self.n_speakers > 0:
|
||||||
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
||||||
else:
|
else:
|
||||||
@@ -981,7 +1066,7 @@ class SynthesizerTrn(nn.Module):
|
|||||||
|
|
||||||
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
|
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
|
||||||
attn = (
|
attn = (
|
||||||
monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1))
|
monotonic_alignment.maximum_path(neg_cent, attn_mask.squeeze(1))
|
||||||
.unsqueeze(1)
|
.unsqueeze(1)
|
||||||
.detach()
|
.detach()
|
||||||
)
|
)
|
||||||
@@ -1016,32 +1101,33 @@ class SynthesizerTrn(nn.Module):
|
|||||||
ids_slice,
|
ids_slice,
|
||||||
x_mask,
|
x_mask,
|
||||||
y_mask,
|
y_mask,
|
||||||
(z, z_p, m_p, logs_p, m_q, logs_q),
|
(z, z_p, m_p, logs_p, m_q, logs_q), # type: ignore
|
||||||
(x, logw, logw_), # , logw_sdp),
|
(x, logw, logw_), # , logw_sdp),
|
||||||
g,
|
g,
|
||||||
)
|
)
|
||||||
|
|
||||||
def infer(
|
def infer(
|
||||||
self,
|
self,
|
||||||
x,
|
x: torch.Tensor,
|
||||||
x_lengths,
|
x_lengths: torch.Tensor,
|
||||||
sid,
|
sid: torch.Tensor,
|
||||||
tone,
|
tone: torch.Tensor,
|
||||||
language,
|
language: torch.Tensor,
|
||||||
bert,
|
bert: torch.Tensor,
|
||||||
style_vec,
|
style_vec: torch.Tensor,
|
||||||
noise_scale=0.667,
|
noise_scale: float = 0.667,
|
||||||
length_scale=1,
|
length_scale: float = 1.0,
|
||||||
noise_scale_w=0.8,
|
noise_scale_w: float = 0.8,
|
||||||
max_len=None,
|
max_len: Optional[int] = None,
|
||||||
sdp_ratio=0,
|
sdp_ratio: float = 0.0,
|
||||||
y=None,
|
y: Optional[torch.Tensor] = None,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]:
|
||||||
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)
|
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)
|
||||||
# g = self.gst(y)
|
# g = self.gst(y)
|
||||||
if self.n_speakers > 0:
|
if self.n_speakers > 0:
|
||||||
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
||||||
else:
|
else:
|
||||||
|
assert y is not None
|
||||||
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
|
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
|
||||||
x, m_p, logs_p, x_mask = self.enc_p(
|
x, m_p, logs_p, x_mask = self.enc_p(
|
||||||
x, x_lengths, tone, language, bert, style_vec, g=g
|
x, x_lengths, tone, language, bert, style_vec, g=g
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
import math
|
import math
|
||||||
import warnings
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
@@ -7,16 +7,16 @@ from torch.nn import Conv1d
|
|||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from torch.nn.utils import remove_weight_norm, weight_norm
|
from torch.nn.utils import remove_weight_norm, weight_norm
|
||||||
|
|
||||||
import commons
|
from style_bert_vits2.models import commons
|
||||||
from attentions import Encoder
|
from style_bert_vits2.models.attentions import Encoder
|
||||||
from commons import get_padding, init_weights
|
from style_bert_vits2.models.transforms import piecewise_rational_quadratic_transform
|
||||||
from transforms import piecewise_rational_quadratic_transform
|
|
||||||
|
|
||||||
LRELU_SLOPE = 0.1
|
LRELU_SLOPE = 0.1
|
||||||
|
|
||||||
|
|
||||||
class LayerNorm(nn.Module):
|
class LayerNorm(nn.Module):
|
||||||
def __init__(self, channels, eps=1e-5):
|
def __init__(self, channels: int, eps: float = 1e-5) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.eps = eps
|
self.eps = eps
|
||||||
@@ -24,7 +24,7 @@ class LayerNorm(nn.Module):
|
|||||||
self.gamma = nn.Parameter(torch.ones(channels))
|
self.gamma = nn.Parameter(torch.ones(channels))
|
||||||
self.beta = nn.Parameter(torch.zeros(channels))
|
self.beta = nn.Parameter(torch.zeros(channels))
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
x = x.transpose(1, -1)
|
x = x.transpose(1, -1)
|
||||||
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
|
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
|
||||||
return x.transpose(1, -1)
|
return x.transpose(1, -1)
|
||||||
@@ -33,13 +33,13 @@ class LayerNorm(nn.Module):
|
|||||||
class ConvReluNorm(nn.Module):
|
class ConvReluNorm(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
out_channels,
|
out_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
p_dropout,
|
p_dropout: float,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -70,9 +70,10 @@ class ConvReluNorm(nn.Module):
|
|||||||
self.norm_layers.append(LayerNorm(hidden_channels))
|
self.norm_layers.append(LayerNorm(hidden_channels))
|
||||||
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
|
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
|
||||||
self.proj.weight.data.zero_()
|
self.proj.weight.data.zero_()
|
||||||
|
assert self.proj.bias is not None
|
||||||
self.proj.bias.data.zero_()
|
self.proj.bias.data.zero_()
|
||||||
|
|
||||||
def forward(self, x, x_mask):
|
def forward(self, x: torch.Tensor, x_mask: torch.Tensor) -> torch.Tensor:
|
||||||
x_org = x
|
x_org = x
|
||||||
for i in range(self.n_layers):
|
for i in range(self.n_layers):
|
||||||
x = self.conv_layers[i](x * x_mask)
|
x = self.conv_layers[i](x * x_mask)
|
||||||
@@ -87,7 +88,9 @@ class DDSConv(nn.Module):
|
|||||||
Dialted and Depth-Separable Convolution
|
Dialted and Depth-Separable Convolution
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, channels, kernel_size, n_layers, p_dropout=0.0):
|
def __init__(
|
||||||
|
self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.kernel_size = kernel_size
|
self.kernel_size = kernel_size
|
||||||
@@ -116,7 +119,9 @@ class DDSConv(nn.Module):
|
|||||||
self.norms_1.append(LayerNorm(channels))
|
self.norms_1.append(LayerNorm(channels))
|
||||||
self.norms_2.append(LayerNorm(channels))
|
self.norms_2.append(LayerNorm(channels))
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + g
|
x = x + g
|
||||||
for i in range(self.n_layers):
|
for i in range(self.n_layers):
|
||||||
@@ -134,13 +139,13 @@ class DDSConv(nn.Module):
|
|||||||
class WN(torch.nn.Module):
|
class WN(torch.nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
p_dropout=0,
|
p_dropout: float = 0,
|
||||||
):
|
) -> None:
|
||||||
super(WN, self).__init__()
|
super(WN, self).__init__()
|
||||||
assert kernel_size % 2 == 1
|
assert kernel_size % 2 == 1
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -183,7 +188,13 @@ class WN(torch.nn.Module):
|
|||||||
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
|
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
|
||||||
self.res_skip_layers.append(res_skip_layer)
|
self.res_skip_layers.append(res_skip_layer)
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, **kwargs):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> torch.Tensor:
|
||||||
output = torch.zeros_like(x)
|
output = torch.zeros_like(x)
|
||||||
n_channels_tensor = torch.IntTensor([self.hidden_channels])
|
n_channels_tensor = torch.IntTensor([self.hidden_channels])
|
||||||
|
|
||||||
@@ -210,7 +221,7 @@ class WN(torch.nn.Module):
|
|||||||
output = output + res_skip_acts
|
output = output + res_skip_acts
|
||||||
return output * x_mask
|
return output * x_mask
|
||||||
|
|
||||||
def remove_weight_norm(self):
|
def remove_weight_norm(self) -> None:
|
||||||
if self.gin_channels != 0:
|
if self.gin_channels != 0:
|
||||||
torch.nn.utils.remove_weight_norm(self.cond_layer)
|
torch.nn.utils.remove_weight_norm(self.cond_layer)
|
||||||
for l in self.in_layers:
|
for l in self.in_layers:
|
||||||
@@ -220,7 +231,12 @@ class WN(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ResBlock1(torch.nn.Module):
|
class ResBlock1(torch.nn.Module):
|
||||||
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
|
def __init__(
|
||||||
|
self,
|
||||||
|
channels: int,
|
||||||
|
kernel_size: int = 3,
|
||||||
|
dilation: tuple[int, int, int] = (1, 3, 5),
|
||||||
|
) -> None:
|
||||||
super(ResBlock1, self).__init__()
|
super(ResBlock1, self).__init__()
|
||||||
self.convs1 = nn.ModuleList(
|
self.convs1 = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -231,7 +247,7 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=dilation[0],
|
dilation=dilation[0],
|
||||||
padding=get_padding(kernel_size, dilation[0]),
|
padding=commons.get_padding(kernel_size, dilation[0]),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
weight_norm(
|
weight_norm(
|
||||||
@@ -241,7 +257,7 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=dilation[1],
|
dilation=dilation[1],
|
||||||
padding=get_padding(kernel_size, dilation[1]),
|
padding=commons.get_padding(kernel_size, dilation[1]),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
weight_norm(
|
weight_norm(
|
||||||
@@ -251,12 +267,12 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=dilation[2],
|
dilation=dilation[2],
|
||||||
padding=get_padding(kernel_size, dilation[2]),
|
padding=commons.get_padding(kernel_size, dilation[2]),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.convs1.apply(init_weights)
|
self.convs1.apply(commons.init_weights)
|
||||||
|
|
||||||
self.convs2 = nn.ModuleList(
|
self.convs2 = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -267,7 +283,7 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=1,
|
dilation=1,
|
||||||
padding=get_padding(kernel_size, 1),
|
padding=commons.get_padding(kernel_size, 1),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
weight_norm(
|
weight_norm(
|
||||||
@@ -277,7 +293,7 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=1,
|
dilation=1,
|
||||||
padding=get_padding(kernel_size, 1),
|
padding=commons.get_padding(kernel_size, 1),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
weight_norm(
|
weight_norm(
|
||||||
@@ -287,14 +303,16 @@ class ResBlock1(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=1,
|
dilation=1,
|
||||||
padding=get_padding(kernel_size, 1),
|
padding=commons.get_padding(kernel_size, 1),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.convs2.apply(init_weights)
|
self.convs2.apply(commons.init_weights)
|
||||||
|
|
||||||
def forward(self, x, x_mask=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
for c1, c2 in zip(self.convs1, self.convs2):
|
for c1, c2 in zip(self.convs1, self.convs2):
|
||||||
xt = F.leaky_relu(x, LRELU_SLOPE)
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
||||||
if x_mask is not None:
|
if x_mask is not None:
|
||||||
@@ -309,7 +327,7 @@ class ResBlock1(torch.nn.Module):
|
|||||||
x = x * x_mask
|
x = x * x_mask
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def remove_weight_norm(self):
|
def remove_weight_norm(self) -> None:
|
||||||
for l in self.convs1:
|
for l in self.convs1:
|
||||||
remove_weight_norm(l)
|
remove_weight_norm(l)
|
||||||
for l in self.convs2:
|
for l in self.convs2:
|
||||||
@@ -317,7 +335,9 @@ class ResBlock1(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ResBlock2(torch.nn.Module):
|
class ResBlock2(torch.nn.Module):
|
||||||
def __init__(self, channels, kernel_size=3, dilation=(1, 3)):
|
def __init__(
|
||||||
|
self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)
|
||||||
|
) -> None:
|
||||||
super(ResBlock2, self).__init__()
|
super(ResBlock2, self).__init__()
|
||||||
self.convs = nn.ModuleList(
|
self.convs = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -328,7 +348,7 @@ class ResBlock2(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=dilation[0],
|
dilation=dilation[0],
|
||||||
padding=get_padding(kernel_size, dilation[0]),
|
padding=commons.get_padding(kernel_size, dilation[0]),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
weight_norm(
|
weight_norm(
|
||||||
@@ -338,14 +358,16 @@ class ResBlock2(torch.nn.Module):
|
|||||||
kernel_size,
|
kernel_size,
|
||||||
1,
|
1,
|
||||||
dilation=dilation[1],
|
dilation=dilation[1],
|
||||||
padding=get_padding(kernel_size, dilation[1]),
|
padding=commons.get_padding(kernel_size, dilation[1]),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.convs.apply(init_weights)
|
self.convs.apply(commons.init_weights)
|
||||||
|
|
||||||
def forward(self, x, x_mask=None):
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
for c in self.convs:
|
for c in self.convs:
|
||||||
xt = F.leaky_relu(x, LRELU_SLOPE)
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
||||||
if x_mask is not None:
|
if x_mask is not None:
|
||||||
@@ -356,13 +378,19 @@ class ResBlock2(torch.nn.Module):
|
|||||||
x = x * x_mask
|
x = x * x_mask
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def remove_weight_norm(self):
|
def remove_weight_norm(self) -> None:
|
||||||
for l in self.convs:
|
for l in self.convs:
|
||||||
remove_weight_norm(l)
|
remove_weight_norm(l)
|
||||||
|
|
||||||
|
|
||||||
class Log(nn.Module):
|
class Log(nn.Module):
|
||||||
def forward(self, x, x_mask, reverse=False, **kwargs):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
reverse: bool = False,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask
|
y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask
|
||||||
logdet = torch.sum(-y, [1, 2])
|
logdet = torch.sum(-y, [1, 2])
|
||||||
@@ -373,7 +401,13 @@ class Log(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class Flip(nn.Module):
|
class Flip(nn.Module):
|
||||||
def forward(self, x, *args, reverse=False, **kwargs):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
*args: Any,
|
||||||
|
reverse: bool = False,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
x = torch.flip(x, [1])
|
x = torch.flip(x, [1])
|
||||||
if not reverse:
|
if not reverse:
|
||||||
logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device)
|
logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device)
|
||||||
@@ -383,13 +417,19 @@ class Flip(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ElementwiseAffine(nn.Module):
|
class ElementwiseAffine(nn.Module):
|
||||||
def __init__(self, channels):
|
def __init__(self, channels: int) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.m = nn.Parameter(torch.zeros(channels, 1))
|
self.m = nn.Parameter(torch.zeros(channels, 1))
|
||||||
self.logs = nn.Parameter(torch.zeros(channels, 1))
|
self.logs = nn.Parameter(torch.zeros(channels, 1))
|
||||||
|
|
||||||
def forward(self, x, x_mask, reverse=False, **kwargs):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
reverse: bool = False,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
if not reverse:
|
if not reverse:
|
||||||
y = self.m + torch.exp(self.logs) * x
|
y = self.m + torch.exp(self.logs) * x
|
||||||
y = y * x_mask
|
y = y * x_mask
|
||||||
@@ -403,15 +443,15 @@ class ElementwiseAffine(nn.Module):
|
|||||||
class ResidualCouplingLayer(nn.Module):
|
class ResidualCouplingLayer(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
dilation_rate,
|
dilation_rate: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
p_dropout=0,
|
p_dropout: float = 0,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
mean_only=False,
|
mean_only: bool = False,
|
||||||
):
|
) -> None:
|
||||||
assert channels % 2 == 0, "channels should be divisible by 2"
|
assert channels % 2 == 0, "channels should be divisible by 2"
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
@@ -433,9 +473,16 @@ class ResidualCouplingLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
|
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
|
||||||
self.post.weight.data.zero_()
|
self.post.weight.data.zero_()
|
||||||
|
assert self.post.bias is not None
|
||||||
self.post.bias.data.zero_()
|
self.post.bias.data.zero_()
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
||||||
h = self.pre(x0) * x_mask
|
h = self.pre(x0) * x_mask
|
||||||
h = self.enc(h, x_mask, g=g)
|
h = self.enc(h, x_mask, g=g)
|
||||||
@@ -460,13 +507,13 @@ class ResidualCouplingLayer(nn.Module):
|
|||||||
class ConvFlow(nn.Module):
|
class ConvFlow(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
in_channels,
|
in_channels: int,
|
||||||
filter_channels,
|
filter_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
num_bins=10,
|
num_bins: int = 10,
|
||||||
tail_bound=5.0,
|
tail_bound: float = 5.0,
|
||||||
):
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.in_channels = in_channels
|
self.in_channels = in_channels
|
||||||
self.filter_channels = filter_channels
|
self.filter_channels = filter_channels
|
||||||
@@ -482,9 +529,16 @@ class ConvFlow(nn.Module):
|
|||||||
filter_channels, self.half_channels * (num_bins * 3 - 1), 1
|
filter_channels, self.half_channels * (num_bins * 3 - 1), 1
|
||||||
)
|
)
|
||||||
self.proj.weight.data.zero_()
|
self.proj.weight.data.zero_()
|
||||||
|
assert self.proj.bias is not None
|
||||||
self.proj.bias.data.zero_()
|
self.proj.bias.data.zero_()
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
||||||
h = self.pre(x0)
|
h = self.pre(x0)
|
||||||
h = self.convs(h, x_mask, g=g)
|
h = self.convs(h, x_mask, g=g)
|
||||||
@@ -520,17 +574,17 @@ class ConvFlow(nn.Module):
|
|||||||
class TransformerCouplingLayer(nn.Module):
|
class TransformerCouplingLayer(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
channels,
|
channels: int,
|
||||||
hidden_channels,
|
hidden_channels: int,
|
||||||
kernel_size,
|
kernel_size: int,
|
||||||
n_layers,
|
n_layers: int,
|
||||||
n_heads,
|
n_heads: int,
|
||||||
p_dropout=0,
|
p_dropout: float = 0,
|
||||||
filter_channels=0,
|
filter_channels: int = 0,
|
||||||
mean_only=False,
|
mean_only: bool = False,
|
||||||
wn_sharing_parameter=None,
|
wn_sharing_parameter: Optional[nn.Module] = None,
|
||||||
gin_channels=0,
|
gin_channels: int = 0,
|
||||||
):
|
) -> None:
|
||||||
assert channels % 2 == 0, "channels should be divisible by 2"
|
assert channels % 2 == 0, "channels should be divisible by 2"
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
@@ -557,9 +611,16 @@ class TransformerCouplingLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
|
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
|
||||||
self.post.weight.data.zero_()
|
self.post.weight.data.zero_()
|
||||||
|
assert self.post.bias is not None
|
||||||
self.post.bias.data.zero_()
|
self.post.bias.data.zero_()
|
||||||
|
|
||||||
def forward(self, x, x_mask, g=None, reverse=False):
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
reverse: bool = False,
|
||||||
|
) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
|
||||||
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
|
||||||
h = self.pre(x0) * x_mask
|
h = self.pre(x0) * x_mask
|
||||||
h = self.enc(h, x_mask, g=g)
|
h = self.enc(h, x_mask, g=g)
|
||||||
89
style_bert_vits2/models/monotonic_alignment.py
Normal file
89
style_bert_vits2/models/monotonic_alignment.py
Normal file
@@ -0,0 +1,89 @@
|
|||||||
|
"""
|
||||||
|
以下に記述されている関数のコメントはリファクタリング時に GPT-4 に生成させたもので、
|
||||||
|
コードと完全に一致している保証はない。あくまで参考程度とすること。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import numba
|
||||||
|
import torch
|
||||||
|
from numpy import float32, int32, zeros
|
||||||
|
|
||||||
|
|
||||||
|
def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
与えられた負の中心とマスクを使用して最大パスを計算する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
neg_cent (torch.Tensor): 負の中心を表すテンソル
|
||||||
|
mask (torch.Tensor): マスクを表すテンソル
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor: 計算された最大パスを表すテンソル
|
||||||
|
"""
|
||||||
|
|
||||||
|
device = neg_cent.device
|
||||||
|
dtype = neg_cent.dtype
|
||||||
|
neg_cent = neg_cent.data.cpu().numpy().astype(float32)
|
||||||
|
path = zeros(neg_cent.shape, dtype=int32)
|
||||||
|
|
||||||
|
t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32)
|
||||||
|
t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32)
|
||||||
|
__maximum_path_jit(path, neg_cent, t_t_max, t_s_max)
|
||||||
|
|
||||||
|
return torch.from_numpy(path).to(device=device, dtype=dtype)
|
||||||
|
|
||||||
|
|
||||||
|
@numba.jit(
|
||||||
|
numba.void(
|
||||||
|
numba.int32[:, :, ::1],
|
||||||
|
numba.float32[:, :, ::1],
|
||||||
|
numba.int32[::1],
|
||||||
|
numba.int32[::1],
|
||||||
|
),
|
||||||
|
nopython=True,
|
||||||
|
nogil=True,
|
||||||
|
) # type: ignore
|
||||||
|
def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None:
|
||||||
|
"""
|
||||||
|
与えられたパス、値、およびターゲットの y と x 座標を使用して JIT で最大パスを計算する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
paths: 計算されたパスを格納するための整数型の 3 次元配列
|
||||||
|
values: 値を格納するための浮動小数点型の 3 次元配列
|
||||||
|
t_ys: ターゲットの y 座標を格納するための整数型の 1 次元配列
|
||||||
|
t_xs: ターゲットの x 座標を格納するための整数型の 1 次元配列
|
||||||
|
"""
|
||||||
|
|
||||||
|
b = paths.shape[0]
|
||||||
|
max_neg_val = -1e9
|
||||||
|
for i in range(int(b)):
|
||||||
|
path = paths[i]
|
||||||
|
value = values[i]
|
||||||
|
t_y = t_ys[i]
|
||||||
|
t_x = t_xs[i]
|
||||||
|
|
||||||
|
v_prev = v_cur = 0.0
|
||||||
|
index = t_x - 1
|
||||||
|
|
||||||
|
for y in range(t_y):
|
||||||
|
for x in range(max(0, t_x + y - t_y), min(t_x, y + 1)):
|
||||||
|
if x == y:
|
||||||
|
v_cur = max_neg_val
|
||||||
|
else:
|
||||||
|
v_cur = value[y - 1, x]
|
||||||
|
if x == 0:
|
||||||
|
if y == 0:
|
||||||
|
v_prev = 0.0
|
||||||
|
else:
|
||||||
|
v_prev = max_neg_val
|
||||||
|
else:
|
||||||
|
v_prev = value[y - 1, x - 1]
|
||||||
|
value[y, x] += max(v_prev, v_cur)
|
||||||
|
|
||||||
|
for y in range(t_y - 1, -1, -1):
|
||||||
|
path[y, index] = 1
|
||||||
|
if index != 0 and (
|
||||||
|
index == y or value[y - 1, index] < value[y - 1, index - 1]
|
||||||
|
):
|
||||||
|
index = index - 1
|
||||||
@@ -1,7 +1,8 @@
|
|||||||
import torch
|
from typing import Optional
|
||||||
from torch.nn import functional as F
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from torch.nn import functional as F
|
||||||
|
|
||||||
|
|
||||||
DEFAULT_MIN_BIN_WIDTH = 1e-3
|
DEFAULT_MIN_BIN_WIDTH = 1e-3
|
||||||
@@ -10,17 +11,18 @@ DEFAULT_MIN_DERIVATIVE = 1e-3
|
|||||||
|
|
||||||
|
|
||||||
def piecewise_rational_quadratic_transform(
|
def piecewise_rational_quadratic_transform(
|
||||||
inputs,
|
inputs: torch.Tensor,
|
||||||
unnormalized_widths,
|
unnormalized_widths: torch.Tensor,
|
||||||
unnormalized_heights,
|
unnormalized_heights: torch.Tensor,
|
||||||
unnormalized_derivatives,
|
unnormalized_derivatives: torch.Tensor,
|
||||||
inverse=False,
|
inverse: bool = False,
|
||||||
tails=None,
|
tails: Optional[str] = None,
|
||||||
tail_bound=1.0,
|
tail_bound: float = 1.0,
|
||||||
min_bin_width=DEFAULT_MIN_BIN_WIDTH,
|
min_bin_width: float = DEFAULT_MIN_BIN_WIDTH,
|
||||||
min_bin_height=DEFAULT_MIN_BIN_HEIGHT,
|
min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT,
|
||||||
min_derivative=DEFAULT_MIN_DERIVATIVE,
|
min_derivative: float = DEFAULT_MIN_DERIVATIVE,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
|
||||||
if tails is None:
|
if tails is None:
|
||||||
spline_fn = rational_quadratic_spline
|
spline_fn = rational_quadratic_spline
|
||||||
spline_kwargs = {}
|
spline_kwargs = {}
|
||||||
@@ -37,28 +39,31 @@ def piecewise_rational_quadratic_transform(
|
|||||||
min_bin_width=min_bin_width,
|
min_bin_width=min_bin_width,
|
||||||
min_bin_height=min_bin_height,
|
min_bin_height=min_bin_height,
|
||||||
min_derivative=min_derivative,
|
min_derivative=min_derivative,
|
||||||
**spline_kwargs
|
**spline_kwargs, # type: ignore
|
||||||
)
|
)
|
||||||
return outputs, logabsdet
|
return outputs, logabsdet
|
||||||
|
|
||||||
|
|
||||||
def searchsorted(bin_locations, inputs, eps=1e-6):
|
def searchsorted(
|
||||||
|
bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6
|
||||||
|
) -> torch.Tensor:
|
||||||
bin_locations[..., -1] += eps
|
bin_locations[..., -1] += eps
|
||||||
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1
|
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1
|
||||||
|
|
||||||
|
|
||||||
def unconstrained_rational_quadratic_spline(
|
def unconstrained_rational_quadratic_spline(
|
||||||
inputs,
|
inputs: torch.Tensor,
|
||||||
unnormalized_widths,
|
unnormalized_widths: torch.Tensor,
|
||||||
unnormalized_heights,
|
unnormalized_heights: torch.Tensor,
|
||||||
unnormalized_derivatives,
|
unnormalized_derivatives: torch.Tensor,
|
||||||
inverse=False,
|
inverse: bool = False,
|
||||||
tails="linear",
|
tails: str = "linear",
|
||||||
tail_bound=1.0,
|
tail_bound: float = 1.0,
|
||||||
min_bin_width=DEFAULT_MIN_BIN_WIDTH,
|
min_bin_width: float = DEFAULT_MIN_BIN_WIDTH,
|
||||||
min_bin_height=DEFAULT_MIN_BIN_HEIGHT,
|
min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT,
|
||||||
min_derivative=DEFAULT_MIN_DERIVATIVE,
|
min_derivative: float = DEFAULT_MIN_DERIVATIVE,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
|
||||||
inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound)
|
inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound)
|
||||||
outside_interval_mask = ~inside_interval_mask
|
outside_interval_mask = ~inside_interval_mask
|
||||||
|
|
||||||
@@ -74,7 +79,7 @@ def unconstrained_rational_quadratic_spline(
|
|||||||
outputs[outside_interval_mask] = inputs[outside_interval_mask]
|
outputs[outside_interval_mask] = inputs[outside_interval_mask]
|
||||||
logabsdet[outside_interval_mask] = 0
|
logabsdet[outside_interval_mask] = 0
|
||||||
else:
|
else:
|
||||||
raise RuntimeError("{} tails are not implemented.".format(tails))
|
raise RuntimeError(f"{tails} tails are not implemented.")
|
||||||
|
|
||||||
(
|
(
|
||||||
outputs[inside_interval_mask],
|
outputs[inside_interval_mask],
|
||||||
@@ -98,19 +103,20 @@ def unconstrained_rational_quadratic_spline(
|
|||||||
|
|
||||||
|
|
||||||
def rational_quadratic_spline(
|
def rational_quadratic_spline(
|
||||||
inputs,
|
inputs: torch.Tensor,
|
||||||
unnormalized_widths,
|
unnormalized_widths: torch.Tensor,
|
||||||
unnormalized_heights,
|
unnormalized_heights: torch.Tensor,
|
||||||
unnormalized_derivatives,
|
unnormalized_derivatives: torch.Tensor,
|
||||||
inverse=False,
|
inverse: bool = False,
|
||||||
left=0.0,
|
left: float = 0.0,
|
||||||
right=1.0,
|
right: float = 1.0,
|
||||||
bottom=0.0,
|
bottom: float = 0.0,
|
||||||
top=1.0,
|
top: float = 1.0,
|
||||||
min_bin_width=DEFAULT_MIN_BIN_WIDTH,
|
min_bin_width: float = DEFAULT_MIN_BIN_WIDTH,
|
||||||
min_bin_height=DEFAULT_MIN_BIN_HEIGHT,
|
min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT,
|
||||||
min_derivative=DEFAULT_MIN_DERIVATIVE,
|
min_derivative: float = DEFAULT_MIN_DERIVATIVE,
|
||||||
):
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
|
||||||
if torch.min(inputs) < left or torch.max(inputs) > right:
|
if torch.min(inputs) < left or torch.max(inputs) > right:
|
||||||
raise ValueError("Input to a transform is not within its domain")
|
raise ValueError("Input to a transform is not within its domain")
|
||||||
|
|
||||||
262
style_bert_vits2/models/utils/__init__.py
Normal file
262
style_bert_vits2/models/utils/__init__.py
Normal file
@@ -0,0 +1,262 @@
|
|||||||
|
import glob
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import subprocess
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
from scipy.io.wavfile import read
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.models.utils import checkpoints # type: ignore
|
||||||
|
from style_bert_vits2.models.utils import safetensors # type: ignore
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
# tensorboard はライブラリとしてインストールされている場合は依存関係に含まれないため、型チェック時のみインポートする
|
||||||
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
|
|
||||||
|
|
||||||
|
__is_matplotlib_imported = False
|
||||||
|
|
||||||
|
|
||||||
|
def summarize(
|
||||||
|
writer: "SummaryWriter",
|
||||||
|
global_step: int,
|
||||||
|
scalars: dict[str, float] = {},
|
||||||
|
histograms: dict[str, Any] = {},
|
||||||
|
images: dict[str, Any] = {},
|
||||||
|
audios: dict[str, Any] = {},
|
||||||
|
audio_sampling_rate: int = 22050,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
指定されたデータを TensorBoard にまとめて追加する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
writer (SummaryWriter): TensorBoard への書き込みを行うオブジェクト
|
||||||
|
global_step (int): グローバルステップ数
|
||||||
|
scalars (dict[str, float]): スカラー値の辞書
|
||||||
|
histograms (dict[str, Any]): ヒストグラムの辞書
|
||||||
|
images (dict[str, Any]): 画像データの辞書
|
||||||
|
audios (dict[str, Any]): 音声データの辞書
|
||||||
|
audio_sampling_rate (int): 音声データのサンプリングレート
|
||||||
|
"""
|
||||||
|
for k, v in scalars.items():
|
||||||
|
writer.add_scalar(k, v, global_step)
|
||||||
|
for k, v in histograms.items():
|
||||||
|
writer.add_histogram(k, v, global_step)
|
||||||
|
for k, v in images.items():
|
||||||
|
writer.add_image(k, v, global_step, dataformats="HWC")
|
||||||
|
for k, v in audios.items():
|
||||||
|
writer.add_audio(k, v, global_step, audio_sampling_rate)
|
||||||
|
|
||||||
|
|
||||||
|
def is_resuming(dir_path: Union[str, Path]) -> bool:
|
||||||
|
"""
|
||||||
|
指定されたディレクトリパスに再開可能なモデルが存在するかどうかを返す
|
||||||
|
|
||||||
|
Args:
|
||||||
|
dir_path: チェックするディレクトリのパス
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 再開可能なモデルが存在するかどうか
|
||||||
|
"""
|
||||||
|
# JP-ExtraバージョンではDURがなくWDがあったり変わるため、Gのみで判断する
|
||||||
|
g_list = glob.glob(os.path.join(dir_path, "G_*.pth"))
|
||||||
|
# d_list = glob.glob(os.path.join(dir_path, "D_*.pth"))
|
||||||
|
# dur_list = glob.glob(os.path.join(dir_path, "DUR_*.pth"))
|
||||||
|
return len(g_list) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
指定されたスペクトログラムを画像データに変換する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
spectrogram (NDArray[Any]): スペクトログラム
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: 画像データ
|
||||||
|
"""
|
||||||
|
|
||||||
|
global __is_matplotlib_imported
|
||||||
|
if not __is_matplotlib_imported:
|
||||||
|
import matplotlib
|
||||||
|
|
||||||
|
matplotlib.use("Agg")
|
||||||
|
__is_matplotlib_imported = True
|
||||||
|
mpl_logger = logging.getLogger("matplotlib")
|
||||||
|
mpl_logger.setLevel(logging.WARNING)
|
||||||
|
import matplotlib.pylab as plt
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
fig, ax = plt.subplots(figsize=(10, 2))
|
||||||
|
im = ax.imshow(spectrogram, aspect="auto", origin="lower", interpolation="none")
|
||||||
|
plt.colorbar(im, ax=ax)
|
||||||
|
plt.xlabel("Frames")
|
||||||
|
plt.ylabel("Channels")
|
||||||
|
plt.tight_layout()
|
||||||
|
|
||||||
|
fig.canvas.draw()
|
||||||
|
data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") # type: ignore
|
||||||
|
data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))
|
||||||
|
plt.close()
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def plot_alignment_to_numpy(
|
||||||
|
alignment: NDArray[Any], info: Optional[str] = None
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
指定されたアライメントを画像データに変換する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
alignment (NDArray[Any]): アライメント
|
||||||
|
info (Optional[str]): 画像に追加する情報
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: 画像データ
|
||||||
|
"""
|
||||||
|
|
||||||
|
global __is_matplotlib_imported
|
||||||
|
if not __is_matplotlib_imported:
|
||||||
|
import matplotlib
|
||||||
|
|
||||||
|
matplotlib.use("Agg")
|
||||||
|
__is_matplotlib_imported = True
|
||||||
|
mpl_logger = logging.getLogger("matplotlib")
|
||||||
|
mpl_logger.setLevel(logging.WARNING)
|
||||||
|
import matplotlib.pylab as plt
|
||||||
|
|
||||||
|
fig, ax = plt.subplots(figsize=(6, 4))
|
||||||
|
im = ax.imshow(
|
||||||
|
alignment.transpose(), aspect="auto", origin="lower", interpolation="none"
|
||||||
|
)
|
||||||
|
fig.colorbar(im, ax=ax)
|
||||||
|
xlabel = "Decoder timestep"
|
||||||
|
if info is not None:
|
||||||
|
xlabel += "\n\n" + info
|
||||||
|
plt.xlabel(xlabel)
|
||||||
|
plt.ylabel("Encoder timestep")
|
||||||
|
plt.tight_layout()
|
||||||
|
|
||||||
|
fig.canvas.draw()
|
||||||
|
data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") # type: ignore
|
||||||
|
data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))
|
||||||
|
plt.close()
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, int]:
|
||||||
|
"""
|
||||||
|
指定された音声ファイルを読み込み、PyTorch のテンソルに変換して返す
|
||||||
|
|
||||||
|
Args:
|
||||||
|
full_path (Union[str, Path]): 音声ファイルのパス
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[torch.FloatTensor, int]: 音声データのテンソルとサンプリングレート
|
||||||
|
"""
|
||||||
|
|
||||||
|
sampling_rate, data = read(full_path)
|
||||||
|
return torch.FloatTensor(data.astype(np.float32)), sampling_rate
|
||||||
|
|
||||||
|
|
||||||
|
def load_filepaths_and_text(
|
||||||
|
filename: Union[str, Path], split: str = "|"
|
||||||
|
) -> list[list[str]]:
|
||||||
|
"""
|
||||||
|
指定されたファイルからファイルパスとテキストを読み込む
|
||||||
|
|
||||||
|
Args:
|
||||||
|
filename (Union[str, Path]): ファイルのパス
|
||||||
|
split (str): ファイルの区切り文字 (デフォルト: "|")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[list[str]]: ファイルパスとテキストのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
with open(filename, "r", encoding="utf-8") as f:
|
||||||
|
filepaths_and_text = [line.strip().split(split) for line in f]
|
||||||
|
return filepaths_and_text
|
||||||
|
|
||||||
|
|
||||||
|
def get_logger(
|
||||||
|
model_dir_path: Union[str, Path], filename: str = "train.log"
|
||||||
|
) -> logging.Logger:
|
||||||
|
"""
|
||||||
|
ロガーを取得する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_dir_path (Union[str, Path]): ログを保存するディレクトリのパス
|
||||||
|
filename (str): ログファイルの名前 (デフォルト: "train.log")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
logging.Logger: ロガー
|
||||||
|
"""
|
||||||
|
|
||||||
|
global logger
|
||||||
|
logger = logging.getLogger(os.path.basename(model_dir_path))
|
||||||
|
logger.setLevel(logging.DEBUG)
|
||||||
|
|
||||||
|
formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s")
|
||||||
|
if not os.path.exists(model_dir_path):
|
||||||
|
os.makedirs(model_dir_path)
|
||||||
|
h = logging.FileHandler(os.path.join(model_dir_path, filename))
|
||||||
|
h.setLevel(logging.DEBUG)
|
||||||
|
h.setFormatter(formatter)
|
||||||
|
logger.addHandler(h)
|
||||||
|
return logger
|
||||||
|
|
||||||
|
|
||||||
|
def get_steps(model_path: Union[str, Path]) -> Optional[int]:
|
||||||
|
"""
|
||||||
|
モデルのパスからイテレーション回数を取得する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_path (Union[str, Path]): モデルのパス
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Optional[int]: イテレーション回数
|
||||||
|
"""
|
||||||
|
|
||||||
|
matches = re.findall(r"\d+", model_path) # type: ignore
|
||||||
|
return matches[-1] if matches else None
|
||||||
|
|
||||||
|
|
||||||
|
def check_git_hash(model_dir_path: Union[str, Path]) -> None:
|
||||||
|
"""
|
||||||
|
モデルのディレクトリに .git ディレクトリが存在する場合、ハッシュ値を比較する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_dir_path (Union[str, Path]): モデルのディレクトリのパス
|
||||||
|
"""
|
||||||
|
|
||||||
|
source_dir = os.path.dirname(os.path.realpath(__file__))
|
||||||
|
if not os.path.exists(os.path.join(source_dir, ".git")):
|
||||||
|
logger.warning(
|
||||||
|
"{} is not a git repository, therefore hash value comparison will be ignored.".format(
|
||||||
|
source_dir
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
cur_hash = subprocess.getoutput("git rev-parse HEAD")
|
||||||
|
|
||||||
|
path = os.path.join(model_dir_path, "githash")
|
||||||
|
if os.path.exists(path):
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
saved_hash = f.read()
|
||||||
|
if saved_hash != cur_hash:
|
||||||
|
logger.warning(
|
||||||
|
"git hash values are different. {}(saved) != {}(current)".format(
|
||||||
|
saved_hash[:8], cur_hash[:8]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(cur_hash)
|
||||||
202
style_bert_vits2/models/utils/checkpoints.py
Normal file
202
style_bert_vits2/models/utils/checkpoints.py
Normal file
@@ -0,0 +1,202 @@
|
|||||||
|
import glob
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
|
def load_checkpoint(
|
||||||
|
checkpoint_path: Union[str, Path],
|
||||||
|
model: torch.nn.Module,
|
||||||
|
optimizer: Optional[torch.optim.Optimizer] = None,
|
||||||
|
skip_optimizer: bool = False,
|
||||||
|
for_infer: bool = False,
|
||||||
|
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
|
||||||
|
"""
|
||||||
|
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
checkpoint_path (Union[str, Path]): チェックポイントファイルのパス
|
||||||
|
model (torch.nn.Module): 更新するモデル
|
||||||
|
optimizer (Optional[torch.optim.Optimizer]): 更新するオプティマイザー。None の場合は更新しない
|
||||||
|
skip_optimizer (bool): オプティマイザーの更新をスキップするかどうかのフラグ
|
||||||
|
for_infer (bool): 推論用に読み込むかどうかのフラグ
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: 更新されたモデルとオプティマイザー、学習率、イテレーション回数
|
||||||
|
"""
|
||||||
|
|
||||||
|
assert os.path.isfile(checkpoint_path)
|
||||||
|
checkpoint_dict = torch.load(checkpoint_path, map_location="cpu")
|
||||||
|
iteration = checkpoint_dict["iteration"]
|
||||||
|
learning_rate = checkpoint_dict["learning_rate"]
|
||||||
|
logger.info(
|
||||||
|
f"Loading model and optimizer at iteration {iteration} from {checkpoint_path}"
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
optimizer is not None
|
||||||
|
and not skip_optimizer
|
||||||
|
and checkpoint_dict["optimizer"] is not None
|
||||||
|
):
|
||||||
|
optimizer.load_state_dict(checkpoint_dict["optimizer"])
|
||||||
|
elif optimizer is None and not skip_optimizer:
|
||||||
|
# else: Disable this line if Infer and resume checkpoint,then enable the line upper
|
||||||
|
new_opt_dict = optimizer.state_dict() # type: ignore
|
||||||
|
new_opt_dict_params = new_opt_dict["param_groups"][0]["params"]
|
||||||
|
new_opt_dict["param_groups"] = checkpoint_dict["optimizer"]["param_groups"]
|
||||||
|
new_opt_dict["param_groups"][0]["params"] = new_opt_dict_params
|
||||||
|
optimizer.load_state_dict(new_opt_dict) # type: ignore
|
||||||
|
|
||||||
|
saved_state_dict = checkpoint_dict["model"]
|
||||||
|
if hasattr(model, "module"):
|
||||||
|
state_dict = model.module.state_dict()
|
||||||
|
else:
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
|
||||||
|
new_state_dict = {}
|
||||||
|
for k, v in state_dict.items():
|
||||||
|
try:
|
||||||
|
# assert "emb_g" not in k
|
||||||
|
new_state_dict[k] = saved_state_dict[k]
|
||||||
|
assert saved_state_dict[k].shape == v.shape, (
|
||||||
|
saved_state_dict[k].shape,
|
||||||
|
v.shape,
|
||||||
|
)
|
||||||
|
except:
|
||||||
|
# For upgrading from the old version
|
||||||
|
if "ja_bert_proj" in k:
|
||||||
|
v = torch.zeros_like(v)
|
||||||
|
logger.warning(
|
||||||
|
f"Seems you are using the old version of the model, the {k} is automatically set to zero for backward compatibility"
|
||||||
|
)
|
||||||
|
elif "enc_q" in k and for_infer:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
logger.error(f"{k} is not in the checkpoint {checkpoint_path}")
|
||||||
|
|
||||||
|
new_state_dict[k] = v
|
||||||
|
|
||||||
|
if hasattr(model, "module"):
|
||||||
|
model.module.load_state_dict(new_state_dict, strict=False)
|
||||||
|
else:
|
||||||
|
model.load_state_dict(new_state_dict, strict=False)
|
||||||
|
|
||||||
|
logger.info(f"Loaded '{checkpoint_path}' (iteration {iteration})")
|
||||||
|
|
||||||
|
return model, optimizer, learning_rate, iteration
|
||||||
|
|
||||||
|
|
||||||
|
def save_checkpoint(
|
||||||
|
model: torch.nn.Module,
|
||||||
|
optimizer: Union[torch.optim.Optimizer, torch.optim.AdamW],
|
||||||
|
learning_rate: float,
|
||||||
|
iteration: int,
|
||||||
|
checkpoint_path: Union[str, Path],
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
モデルとオプティマイザーの状態を指定されたパスに保存する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model (torch.nn.Module): 保存するモデル
|
||||||
|
optimizer (Union[torch.optim.Optimizer, torch.optim.AdamW]): 保存するオプティマイザー
|
||||||
|
learning_rate (float): 学習率
|
||||||
|
iteration (int): イテレーション回数
|
||||||
|
checkpoint_path (Union[str, Path]): 保存先のパス
|
||||||
|
"""
|
||||||
|
logger.info(
|
||||||
|
f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}"
|
||||||
|
)
|
||||||
|
if hasattr(model, "module"):
|
||||||
|
state_dict = model.module.state_dict()
|
||||||
|
else:
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
torch.save(
|
||||||
|
{
|
||||||
|
"model": state_dict,
|
||||||
|
"iteration": iteration,
|
||||||
|
"optimizer": optimizer.state_dict(),
|
||||||
|
"learning_rate": learning_rate,
|
||||||
|
},
|
||||||
|
checkpoint_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def clean_checkpoints(
|
||||||
|
model_dir_path: Union[str, Path] = "logs/44k/",
|
||||||
|
n_ckpts_to_keep: int = 2,
|
||||||
|
sort_by_time: bool = True,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_dir_path (Union[str, Path]): モデルが保存されているディレクトリのパス
|
||||||
|
n_ckpts_to_keep (int): 保持するチェックポイントの数(G_0.pth と D_0.pth を除く)
|
||||||
|
sort_by_time (bool): True の場合、時間順に削除。False の場合、名前順に削除
|
||||||
|
"""
|
||||||
|
|
||||||
|
ckpts_files = [
|
||||||
|
f
|
||||||
|
for f in os.listdir(model_dir_path)
|
||||||
|
if os.path.isfile(os.path.join(model_dir_path, f))
|
||||||
|
]
|
||||||
|
|
||||||
|
def name_key(_f: str) -> int:
|
||||||
|
return int(re.compile("._(\\d+)\\.pth").match(_f).group(1)) # type: ignore
|
||||||
|
|
||||||
|
def time_key(_f: str) -> float:
|
||||||
|
return os.path.getmtime(os.path.join(model_dir_path, _f))
|
||||||
|
|
||||||
|
sort_key = time_key if sort_by_time else name_key
|
||||||
|
|
||||||
|
def x_sorted(_x: str) -> list[str]:
|
||||||
|
return sorted(
|
||||||
|
[f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")],
|
||||||
|
key=sort_key,
|
||||||
|
)
|
||||||
|
|
||||||
|
to_del = [
|
||||||
|
os.path.join(model_dir_path, fn)
|
||||||
|
for fn in (
|
||||||
|
x_sorted("G_")[:-n_ckpts_to_keep]
|
||||||
|
+ x_sorted("D_")[:-n_ckpts_to_keep]
|
||||||
|
+ x_sorted("WD_")[:-n_ckpts_to_keep]
|
||||||
|
+ x_sorted("DUR_")[:-n_ckpts_to_keep]
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
def del_info(fn: str) -> None:
|
||||||
|
return logger.info(f"Free up space by deleting ckpt {fn}")
|
||||||
|
|
||||||
|
def del_routine(x: str) -> list[Any]:
|
||||||
|
return [os.remove(x), del_info(x)]
|
||||||
|
|
||||||
|
[del_routine(fn) for fn in to_del]
|
||||||
|
|
||||||
|
|
||||||
|
def get_latest_checkpoint_path(
|
||||||
|
model_dir_path: Union[str, Path], regex: str = "G_*.pth"
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
指定されたディレクトリから最新のチェックポイントのパスを取得する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_dir_path (Union[str, Path]): モデルが保存されているディレクトリのパス
|
||||||
|
regex (str): チェックポイントのファイル名の正規表現
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: 最新のチェックポイントのパス
|
||||||
|
"""
|
||||||
|
|
||||||
|
f_list = glob.glob(os.path.join(str(model_dir_path), regex))
|
||||||
|
f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f))))
|
||||||
|
try:
|
||||||
|
x = f_list[-1]
|
||||||
|
except IndexError:
|
||||||
|
raise ValueError(f"No checkpoint found in {model_dir_path} with regex {regex}")
|
||||||
|
|
||||||
|
return x
|
||||||
91
style_bert_vits2/models/utils/safetensors.py
Normal file
91
style_bert_vits2/models/utils/safetensors.py
Normal file
@@ -0,0 +1,91 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors import safe_open
|
||||||
|
from safetensors.torch import save_file
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
|
def load_safetensors(
|
||||||
|
checkpoint_path: Union[str, Path],
|
||||||
|
model: torch.nn.Module,
|
||||||
|
for_infer: bool = False,
|
||||||
|
) -> tuple[torch.nn.Module, Optional[int]]:
|
||||||
|
"""
|
||||||
|
指定されたパスから safetensors モデルを読み込み、モデルとイテレーションを返す。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
checkpoint_path (Union[str, Path]): モデルのチェックポイントファイルのパス
|
||||||
|
model (torch.nn.Module): 読み込む対象のモデル
|
||||||
|
for_infer (bool): 推論用に読み込むかどうかのフラグ
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[torch.nn.Module, Optional[int]]: 読み込まれたモデルとイテレーション回数(存在する場合)
|
||||||
|
"""
|
||||||
|
|
||||||
|
tensors: dict[str, Any] = {}
|
||||||
|
iteration: Optional[int] = None
|
||||||
|
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as f: # type: ignore
|
||||||
|
for key in f.keys():
|
||||||
|
if key == "iteration":
|
||||||
|
iteration = f.get_tensor(key).item()
|
||||||
|
tensors[key] = f.get_tensor(key)
|
||||||
|
if hasattr(model, "module"):
|
||||||
|
result = model.module.load_state_dict(tensors, strict=False)
|
||||||
|
else:
|
||||||
|
result = model.load_state_dict(tensors, strict=False)
|
||||||
|
for key in result.missing_keys:
|
||||||
|
if key.startswith("enc_q") and for_infer:
|
||||||
|
continue
|
||||||
|
logger.warning(f"Missing key: {key}")
|
||||||
|
for key in result.unexpected_keys:
|
||||||
|
if key == "iteration":
|
||||||
|
continue
|
||||||
|
logger.warning(f"Unexpected key: {key}")
|
||||||
|
if iteration is None:
|
||||||
|
logger.info(f"Loaded '{checkpoint_path}'")
|
||||||
|
else:
|
||||||
|
logger.info(f"Loaded '{checkpoint_path}' (iteration {iteration})")
|
||||||
|
|
||||||
|
return model, iteration
|
||||||
|
|
||||||
|
|
||||||
|
def save_safetensors(
|
||||||
|
model: torch.nn.Module,
|
||||||
|
iteration: int,
|
||||||
|
checkpoint_path: Union[str, Path],
|
||||||
|
is_half: bool = False,
|
||||||
|
for_infer: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
モデルを safetensors 形式で保存する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model (torch.nn.Module): 保存するモデル
|
||||||
|
iteration (int): イテレーション回数
|
||||||
|
checkpoint_path (Union[str, Path]): 保存先のパス
|
||||||
|
is_half (bool): モデルを半精度で保存するかどうかのフラグ
|
||||||
|
for_infer (bool): 推論用に保存するかどうかのフラグ
|
||||||
|
"""
|
||||||
|
|
||||||
|
if hasattr(model, "module"):
|
||||||
|
state_dict = model.module.state_dict()
|
||||||
|
else:
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
keys = []
|
||||||
|
for k in state_dict:
|
||||||
|
if "enc_q" in k and for_infer:
|
||||||
|
continue # noqa: E701
|
||||||
|
keys.append(k)
|
||||||
|
|
||||||
|
new_dict = (
|
||||||
|
{k: state_dict[k].half() for k in keys}
|
||||||
|
if is_half
|
||||||
|
else {k: state_dict[k] for k in keys}
|
||||||
|
)
|
||||||
|
new_dict["iteration"] = torch.LongTensor([iteration])
|
||||||
|
logger.info(f"Saved safetensors to {checkpoint_path}")
|
||||||
|
|
||||||
|
save_file(new_dict, checkpoint_path)
|
||||||
120
style_bert_vits2/nlp/__init__.py
Normal file
120
style_bert_vits2/nlp/__init__.py
Normal file
@@ -0,0 +1,120 @@
|
|||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.nlp.symbols import (
|
||||||
|
LANGUAGE_ID_MAP,
|
||||||
|
LANGUAGE_TONE_START_MAP,
|
||||||
|
SYMBOLS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# __init__.py は配下のモジュールをインポートした時点で実行される
|
||||||
|
# PyTorch のインポートは重いので、型チェック時以外はインポートしない
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
__symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)}
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
language: Languages,
|
||||||
|
device: str,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> "torch.Tensor":
|
||||||
|
"""
|
||||||
|
テキストから BERT の特徴量を抽出する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): テキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
language (Languages): テキストの言語
|
||||||
|
device (str): 推論に利用するデバイス
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language == Languages.JP:
|
||||||
|
from style_bert_vits2.nlp.japanese.bert_feature import extract_bert_feature
|
||||||
|
elif language == Languages.EN:
|
||||||
|
from style_bert_vits2.nlp.english.bert_feature import extract_bert_feature
|
||||||
|
elif language == Languages.ZH:
|
||||||
|
from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Language {language} not supported")
|
||||||
|
|
||||||
|
return extract_bert_feature(text, word2ph, device, assist_text, assist_text_weight)
|
||||||
|
|
||||||
|
|
||||||
|
def clean_text(
|
||||||
|
text: str,
|
||||||
|
language: Languages,
|
||||||
|
use_jp_extra: bool = True,
|
||||||
|
raise_yomi_error: bool = False,
|
||||||
|
) -> tuple[str, list[str], list[int], list[int]]:
|
||||||
|
"""
|
||||||
|
テキストをクリーニングし、音素に変換する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): クリーニングするテキスト
|
||||||
|
language (Languages): テキストの言語
|
||||||
|
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]]: クリーニングされたテキストと、音素・アクセント・元のテキストの各文字に音素が何個割り当てられるかのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Changed to import inside if condition to avoid unnecessary import
|
||||||
|
if language == Languages.JP:
|
||||||
|
from style_bert_vits2.nlp.japanese.g2p import g2p
|
||||||
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
|
|
||||||
|
norm_text = normalize_text(text)
|
||||||
|
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
|
||||||
|
elif language == Languages.EN:
|
||||||
|
from style_bert_vits2.nlp.english.g2p import g2p
|
||||||
|
from style_bert_vits2.nlp.english.normalizer import normalize_text
|
||||||
|
|
||||||
|
norm_text = normalize_text(text)
|
||||||
|
phones, tones, word2ph = g2p(norm_text)
|
||||||
|
elif language == Languages.ZH:
|
||||||
|
from style_bert_vits2.nlp.chinese.g2p import g2p
|
||||||
|
from style_bert_vits2.nlp.chinese.normalizer import normalize_text
|
||||||
|
|
||||||
|
norm_text = normalize_text(text)
|
||||||
|
phones, tones, word2ph = g2p(norm_text)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Language {language} not supported")
|
||||||
|
|
||||||
|
return norm_text, phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
|
def cleaned_text_to_sequence(
|
||||||
|
cleaned_phones: list[str], tones: list[int], language: Languages
|
||||||
|
) -> tuple[list[int], list[int], list[int]]:
|
||||||
|
"""
|
||||||
|
音素リスト・アクセントリスト・言語を、テキスト内の対応する ID に変換する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
cleaned_phones (list[str]): clean_text() でクリーニングされた音素のリスト
|
||||||
|
tones (list[int]): 各音素のアクセント
|
||||||
|
language (Languages): テキストの言語
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text
|
||||||
|
"""
|
||||||
|
|
||||||
|
phones = [__symbol_to_id[symbol] for symbol in cleaned_phones]
|
||||||
|
tone_start = LANGUAGE_TONE_START_MAP[language]
|
||||||
|
tones = [i + tone_start for i in tones]
|
||||||
|
lang_id = LANGUAGE_ID_MAP[language]
|
||||||
|
lang_ids = [lang_id for i in phones]
|
||||||
|
|
||||||
|
return phones, tones, lang_ids
|
||||||
210
style_bert_vits2/nlp/bert_models.py
Normal file
210
style_bert_vits2/nlp/bert_models.py
Normal file
@@ -0,0 +1,210 @@
|
|||||||
|
"""
|
||||||
|
Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデルをロード/取得するためのモジュール。
|
||||||
|
|
||||||
|
オリジナルの Bert-VITS2 では各言語ごとの BERT モデルが初回インポート時にハードコードされたパスから「暗黙的に」ロードされているが、
|
||||||
|
場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。
|
||||||
|
|
||||||
|
そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。
|
||||||
|
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import gc
|
||||||
|
from typing import Optional, Union, cast
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from transformers import (
|
||||||
|
AutoModelForMaskedLM,
|
||||||
|
AutoTokenizer,
|
||||||
|
DebertaV2Model,
|
||||||
|
DebertaV2Tokenizer,
|
||||||
|
PreTrainedModel,
|
||||||
|
PreTrainedTokenizer,
|
||||||
|
PreTrainedTokenizerFast,
|
||||||
|
)
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_BERT_TOKENIZER_PATHS, Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
||||||
|
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
|
__loaded_tokenizers: dict[
|
||||||
|
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
|
||||||
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> Union[PreTrainedModel, DebertaV2Model]:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す。
|
||||||
|
一度ロードされていれば、ロード済みの 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 では、BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: ku-nlp/deberta-v2-large-japanese-char-wwm
|
||||||
|
- 英語: microsoft/deberta-v3-large
|
||||||
|
- 中国語: hfl/chinese-roberta-wwm-ext-large
|
||||||
|
|
||||||
|
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[PreTrainedModel, DebertaV2Model]: ロード済みの 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_BERT_TOKENIZER_PATHS[
|
||||||
|
language
|
||||||
|
].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_TOKENIZER_PATHS[language])
|
||||||
|
|
||||||
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
|
## 英語のみ DebertaV2Model でロードする必要がある
|
||||||
|
if language == Languages.EN:
|
||||||
|
model = cast(
|
||||||
|
DebertaV2Model,
|
||||||
|
DebertaV2Model.from_pretrained(
|
||||||
|
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
||||||
|
),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
model = AutoModelForMaskedLM.from_pretrained(
|
||||||
|
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
||||||
|
)
|
||||||
|
__loaded_models[language] = model
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def load_tokenizer(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
||||||
|
一度ロードされていれば、ロード済みの 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 では、BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: ku-nlp/deberta-v2-large-japanese-char-wwm
|
||||||
|
- 英語: microsoft/deberta-v3-large
|
||||||
|
- 中国語: hfl/chinese-roberta-wwm-ext-large
|
||||||
|
|
||||||
|
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_BERT_TOKENIZER_PATHS[
|
||||||
|
language
|
||||||
|
].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])
|
||||||
|
|
||||||
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
|
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
||||||
|
if language == Languages.EN:
|
||||||
|
tokenizer = DebertaV2Tokenizer.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
__loaded_tokenizers[language] = tokenizer
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return tokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def unload_model(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT モデルの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_models:
|
||||||
|
del __loaded_models[language]
|
||||||
|
gc.collect()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
logger.info(f"Unloaded the {language} BERT model")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT トークナイザーをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT トークナイザーの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_tokenizers:
|
||||||
|
del __loaded_tokenizers[language]
|
||||||
|
gc.collect()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
logger.info(f"Unloaded the {language} BERT tokenizer")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_models() -> None:
|
||||||
|
"""
|
||||||
|
すべての BERT モデルをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_models.keys()):
|
||||||
|
unload_model(language)
|
||||||
|
logger.info("Unloaded all BERT models")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_tokenizers() -> None:
|
||||||
|
"""
|
||||||
|
すべての BERT トークナイザーをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_tokenizers.keys()):
|
||||||
|
unload_tokenizer(language)
|
||||||
|
logger.info("Unloaded all BERT tokenizers")
|
||||||
0
style_bert_vits2/nlp/chinese/__init__.py
Normal file
0
style_bert_vits2/nlp/chinese/__init__.py
Normal file
@@ -1,54 +1,58 @@
|
|||||||
import sys
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from transformers import AutoModelForMaskedLM, AutoTokenizer
|
|
||||||
|
|
||||||
from config import config
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
LOCAL_PATH = "./bert/chinese-roberta-wwm-ext-large"
|
|
||||||
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(LOCAL_PATH)
|
|
||||||
|
|
||||||
models = dict()
|
|
||||||
|
|
||||||
|
|
||||||
def get_bert_feature(
|
def extract_bert_feature(
|
||||||
text,
|
text: str,
|
||||||
word2ph,
|
word2ph: list[int],
|
||||||
device=config.bert_gen_config.device,
|
device: str,
|
||||||
assist_text=None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight=0.7,
|
assist_text_weight: float = 0.7,
|
||||||
):
|
) -> torch.Tensor:
|
||||||
if (
|
"""
|
||||||
sys.platform == "darwin"
|
中国語のテキストから BERT の特徴量を抽出する
|
||||||
and torch.backends.mps.is_available()
|
|
||||||
and device == "cpu"
|
Args:
|
||||||
):
|
text (str): 中国語のテキスト
|
||||||
device = "mps"
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
if not device:
|
device (str): 推論に利用するデバイス
|
||||||
device = "cuda"
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
if device not in models.keys():
|
model = bert_models.load_model(Languages.ZH).to(device) # type: ignore
|
||||||
models[device] = AutoModelForMaskedLM.from_pretrained(LOCAL_PATH).to(device)
|
|
||||||
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
|
tokenizer = bert_models.load_tokenizer(Languages.ZH)
|
||||||
inputs = tokenizer(text, return_tensors="pt")
|
inputs = tokenizer(text, return_tensors="pt")
|
||||||
for i in inputs:
|
for i in inputs:
|
||||||
inputs[i] = inputs[i].to(device)
|
inputs[i] = inputs[i].to(device) # type: ignore
|
||||||
res = models[device](**inputs, output_hidden_states=True)
|
res = model(**inputs, output_hidden_states=True)
|
||||||
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
if assist_text:
|
if assist_text:
|
||||||
style_inputs = tokenizer(assist_text, return_tensors="pt")
|
style_inputs = tokenizer(assist_text, return_tensors="pt")
|
||||||
for i in style_inputs:
|
for i in style_inputs:
|
||||||
style_inputs[i] = style_inputs[i].to(device)
|
style_inputs[i] = style_inputs[i].to(device) # type: ignore
|
||||||
style_res = models[device](**style_inputs, output_hidden_states=True)
|
style_res = model(**style_inputs, output_hidden_states=True)
|
||||||
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
style_res_mean = style_res.mean(0)
|
style_res_mean = style_res.mean(0)
|
||||||
|
|
||||||
assert len(word2ph) == len(text) + 2
|
assert len(word2ph) == len(text) + 2
|
||||||
word2phone = word2ph
|
word2phone = word2ph
|
||||||
phone_level_feature = []
|
phone_level_feature = []
|
||||||
for i in range(len(word2phone)):
|
for i in range(len(word2phone)):
|
||||||
if assist_text:
|
if assist_text:
|
||||||
|
assert style_res_mean is not None
|
||||||
repeat_feature = (
|
repeat_feature = (
|
||||||
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||||
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||||
@@ -1,75 +1,23 @@
|
|||||||
import os
|
|
||||||
import re
|
import re
|
||||||
|
from pathlib import Path
|
||||||
import cn2an
|
|
||||||
from pypinyin import lazy_pinyin, Style
|
|
||||||
|
|
||||||
from text.symbols import punctuation
|
|
||||||
from text.tone_sandhi import ToneSandhi
|
|
||||||
|
|
||||||
current_file_path = os.path.dirname(__file__)
|
|
||||||
pinyin_to_symbol_map = {
|
|
||||||
line.split("\t")[0]: line.strip().split("\t")[1]
|
|
||||||
for line in open(os.path.join(current_file_path, "opencpop-strict.txt")).readlines()
|
|
||||||
}
|
|
||||||
|
|
||||||
import jieba.posseg as psg
|
import jieba.posseg as psg
|
||||||
|
from pypinyin import Style, lazy_pinyin
|
||||||
|
|
||||||
|
from style_bert_vits2.nlp.chinese.tone_sandhi import ToneSandhi
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||||
|
|
||||||
|
|
||||||
rep_map = {
|
with open(Path(__file__).parent / "opencpop-strict.txt", "r", encoding="utf-8") as f:
|
||||||
":": ",",
|
__PINYIN_TO_SYMBOL_MAP = {
|
||||||
";": ",",
|
line.split("\t")[0]: line.strip().split("\t")[1] for line in f.readlines()
|
||||||
",": ",",
|
}
|
||||||
"。": ".",
|
|
||||||
"!": "!",
|
|
||||||
"?": "?",
|
|
||||||
"\n": ".",
|
|
||||||
"·": ",",
|
|
||||||
"、": ",",
|
|
||||||
"...": "…",
|
|
||||||
"$": ".",
|
|
||||||
"“": "'",
|
|
||||||
"”": "'",
|
|
||||||
'"': "'",
|
|
||||||
"‘": "'",
|
|
||||||
"’": "'",
|
|
||||||
"(": "'",
|
|
||||||
")": "'",
|
|
||||||
"(": "'",
|
|
||||||
")": "'",
|
|
||||||
"《": "'",
|
|
||||||
"》": "'",
|
|
||||||
"【": "'",
|
|
||||||
"】": "'",
|
|
||||||
"[": "'",
|
|
||||||
"]": "'",
|
|
||||||
"—": "-",
|
|
||||||
"~": "-",
|
|
||||||
"~": "-",
|
|
||||||
"「": "'",
|
|
||||||
"」": "'",
|
|
||||||
}
|
|
||||||
|
|
||||||
tone_modifier = ToneSandhi()
|
|
||||||
|
|
||||||
|
|
||||||
def replace_punctuation(text):
|
def g2p(text: str) -> tuple[list[str], list[int], list[int]]:
|
||||||
text = text.replace("嗯", "恩").replace("呣", "母")
|
pattern = r"(?<=[{0}])\s*".format("".join(PUNCTUATIONS))
|
||||||
pattern = re.compile("|".join(re.escape(p) for p in rep_map.keys()))
|
|
||||||
|
|
||||||
replaced_text = pattern.sub(lambda x: rep_map[x.group()], text)
|
|
||||||
|
|
||||||
replaced_text = re.sub(
|
|
||||||
r"[^\u4e00-\u9fa5" + "".join(punctuation) + r"]+", "", replaced_text
|
|
||||||
)
|
|
||||||
|
|
||||||
return replaced_text
|
|
||||||
|
|
||||||
|
|
||||||
def g2p(text):
|
|
||||||
pattern = r"(?<=[{0}])\s*".format("".join(punctuation))
|
|
||||||
sentences = [i for i in re.split(pattern, text) if i.strip() != ""]
|
sentences = [i for i in re.split(pattern, text) if i.strip() != ""]
|
||||||
phones, tones, word2ph = _g2p(sentences)
|
phones, tones, word2ph = __g2p(sentences)
|
||||||
assert sum(word2ph) == len(phones)
|
assert sum(word2ph) == len(phones)
|
||||||
assert len(word2ph) == len(text) # Sometimes it will crash,you can add a try-catch.
|
assert len(word2ph) == len(text) # Sometimes it will crash,you can add a try-catch.
|
||||||
phones = ["_"] + phones + ["_"]
|
phones = ["_"] + phones + ["_"]
|
||||||
@@ -78,34 +26,22 @@ def g2p(text):
|
|||||||
return phones, tones, word2ph
|
return phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
def _get_initials_finals(word):
|
def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]:
|
||||||
initials = []
|
|
||||||
finals = []
|
|
||||||
orig_initials = lazy_pinyin(word, neutral_tone_with_five=True, style=Style.INITIALS)
|
|
||||||
orig_finals = lazy_pinyin(
|
|
||||||
word, neutral_tone_with_five=True, style=Style.FINALS_TONE3
|
|
||||||
)
|
|
||||||
for c, v in zip(orig_initials, orig_finals):
|
|
||||||
initials.append(c)
|
|
||||||
finals.append(v)
|
|
||||||
return initials, finals
|
|
||||||
|
|
||||||
|
|
||||||
def _g2p(segments):
|
|
||||||
phones_list = []
|
phones_list = []
|
||||||
tones_list = []
|
tones_list = []
|
||||||
word2ph = []
|
word2ph = []
|
||||||
|
tone_modifier = ToneSandhi()
|
||||||
for seg in segments:
|
for seg in segments:
|
||||||
# Replace all English words in the sentence
|
# Replace all English words in the sentence
|
||||||
seg = re.sub("[a-zA-Z]+", "", seg)
|
seg = re.sub("[a-zA-Z]+", "", seg)
|
||||||
seg_cut = psg.lcut(seg)
|
seg_cut = psg.lcut(seg)
|
||||||
initials = []
|
initials = []
|
||||||
finals = []
|
finals = []
|
||||||
seg_cut = tone_modifier.pre_merge_for_modify(seg_cut)
|
seg_cut = tone_modifier.pre_merge_for_modify(seg_cut) # type: ignore
|
||||||
for word, pos in seg_cut:
|
for word, pos in seg_cut:
|
||||||
if pos == "eng":
|
if pos == "eng":
|
||||||
continue
|
continue
|
||||||
sub_initials, sub_finals = _get_initials_finals(word)
|
sub_initials, sub_finals = __get_initials_finals(word)
|
||||||
sub_finals = tone_modifier.modified_tone(word, pos, sub_finals)
|
sub_finals = tone_modifier.modified_tone(word, pos, sub_finals)
|
||||||
initials.append(sub_initials)
|
initials.append(sub_initials)
|
||||||
finals.append(sub_finals)
|
finals.append(sub_finals)
|
||||||
@@ -119,7 +55,7 @@ def _g2p(segments):
|
|||||||
# NOTE: post process for pypinyin outputs
|
# NOTE: post process for pypinyin outputs
|
||||||
# we discriminate i, ii and iii
|
# we discriminate i, ii and iii
|
||||||
if c == v:
|
if c == v:
|
||||||
assert c in punctuation
|
assert c in PUNCTUATIONS
|
||||||
phone = [c]
|
phone = [c]
|
||||||
tone = "0"
|
tone = "0"
|
||||||
word2ph.append(1)
|
word2ph.append(1)
|
||||||
@@ -159,8 +95,12 @@ def _g2p(segments):
|
|||||||
if pinyin[0] in single_rep_map.keys():
|
if pinyin[0] in single_rep_map.keys():
|
||||||
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
|
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
|
||||||
|
|
||||||
assert pinyin in pinyin_to_symbol_map.keys(), (pinyin, seg, raw_pinyin)
|
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (
|
||||||
phone = pinyin_to_symbol_map[pinyin].split(" ")
|
pinyin,
|
||||||
|
seg,
|
||||||
|
raw_pinyin,
|
||||||
|
)
|
||||||
|
phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
|
||||||
word2ph.append(len(phone))
|
word2ph.append(len(phone))
|
||||||
|
|
||||||
phones_list += phone
|
phones_list += phone
|
||||||
@@ -168,32 +108,32 @@ def _g2p(segments):
|
|||||||
return phones_list, tones_list, word2ph
|
return phones_list, tones_list, word2ph
|
||||||
|
|
||||||
|
|
||||||
def text_normalize(text):
|
def __get_initials_finals(word: str) -> tuple[list[str], list[str]]:
|
||||||
numbers = re.findall(r"\d+(?:\.?\d+)?", text)
|
initials = []
|
||||||
for number in numbers:
|
finals = []
|
||||||
text = text.replace(number, cn2an.an2cn(number), 1)
|
orig_initials = lazy_pinyin(word, neutral_tone_with_five=True, style=Style.INITIALS)
|
||||||
text = replace_punctuation(text)
|
orig_finals = lazy_pinyin(
|
||||||
return text
|
word, neutral_tone_with_five=True, style=Style.FINALS_TONE3
|
||||||
|
)
|
||||||
|
for c, v in zip(orig_initials, orig_finals):
|
||||||
def get_bert_feature(text, word2ph):
|
initials.append(c)
|
||||||
from text import chinese_bert
|
finals.append(v)
|
||||||
|
return initials, finals
|
||||||
return chinese_bert.get_bert_feature(text, word2ph)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
from text.chinese_bert import get_bert_feature
|
from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature
|
||||||
|
from style_bert_vits2.nlp.chinese.normalizer import normalize_text
|
||||||
|
|
||||||
text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏"
|
text = "啊!但是《原神》是由,米哈游自主, [研发]的一款全.新开放世界.冒险游戏"
|
||||||
text = text_normalize(text)
|
text = normalize_text(text)
|
||||||
print(text)
|
print(text)
|
||||||
phones, tones, word2ph = g2p(text)
|
phones, tones, word2ph = g2p(text)
|
||||||
bert = get_bert_feature(text, word2ph)
|
bert = extract_bert_feature(text, word2ph, "cuda")
|
||||||
|
|
||||||
print(phones, tones, word2ph, bert.shape)
|
print(phones, tones, word2ph, bert.shape)
|
||||||
|
|
||||||
|
|
||||||
# # 示例用法
|
# 示例用法
|
||||||
# text = "这是一个示例文本:,你好!这是一个测试...."
|
# text = "这是一个示例文本:,你好!这是一个测试...."
|
||||||
# print(g2p_paddle(text)) # 输出: 这是一个示例文本你好这是一个测试
|
# print(g2p_paddle(text)) # 输出: 这是一个示例文本你好这是一个测试
|
||||||
61
style_bert_vits2/nlp/chinese/normalizer.py
Normal file
61
style_bert_vits2/nlp/chinese/normalizer.py
Normal file
@@ -0,0 +1,61 @@
|
|||||||
|
import re
|
||||||
|
|
||||||
|
import cn2an
|
||||||
|
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
numbers = re.findall(r"\d+(?:\.?\d+)?", text)
|
||||||
|
for number in numbers:
|
||||||
|
text = text.replace(number, cn2an.an2cn(number), 1)
|
||||||
|
text = replace_punctuation(text)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def replace_punctuation(text: str) -> str:
|
||||||
|
|
||||||
|
REPLACE_MAP = {
|
||||||
|
":": ",",
|
||||||
|
";": ",",
|
||||||
|
",": ",",
|
||||||
|
"。": ".",
|
||||||
|
"!": "!",
|
||||||
|
"?": "?",
|
||||||
|
"\n": ".",
|
||||||
|
"·": ",",
|
||||||
|
"、": ",",
|
||||||
|
"...": "…",
|
||||||
|
"$": ".",
|
||||||
|
"“": "'",
|
||||||
|
"”": "'",
|
||||||
|
'"': "'",
|
||||||
|
"‘": "'",
|
||||||
|
"’": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"《": "'",
|
||||||
|
"》": "'",
|
||||||
|
"【": "'",
|
||||||
|
"】": "'",
|
||||||
|
"[": "'",
|
||||||
|
"]": "'",
|
||||||
|
"—": "-",
|
||||||
|
"~": "-",
|
||||||
|
"~": "-",
|
||||||
|
"「": "'",
|
||||||
|
"」": "'",
|
||||||
|
}
|
||||||
|
|
||||||
|
text = text.replace("嗯", "恩").replace("呣", "母")
|
||||||
|
pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys()))
|
||||||
|
|
||||||
|
replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text)
|
||||||
|
|
||||||
|
replaced_text = re.sub(
|
||||||
|
r"[^\u4e00-\u9fa5" + "".join(PUNCTUATIONS) + r"]+", "", replaced_text
|
||||||
|
)
|
||||||
|
|
||||||
|
return replaced_text
|
||||||
@@ -11,12 +11,9 @@
|
|||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
from typing import List
|
|
||||||
from typing import Tuple
|
|
||||||
|
|
||||||
import jieba
|
import jieba
|
||||||
from pypinyin import lazy_pinyin
|
from pypinyin import Style, lazy_pinyin
|
||||||
from pypinyin import Style
|
|
||||||
|
|
||||||
|
|
||||||
class ToneSandhi:
|
class ToneSandhi:
|
||||||
@@ -463,7 +460,7 @@ class ToneSandhi:
|
|||||||
# word: "家里"
|
# word: "家里"
|
||||||
# pos: "s"
|
# pos: "s"
|
||||||
# finals: ['ia1', 'i3']
|
# finals: ['ia1', 'i3']
|
||||||
def _neural_sandhi(self, word: str, pos: str, finals: List[str]) -> List[str]:
|
def _neural_sandhi(self, word: str, pos: str, finals: list[str]) -> list[str]:
|
||||||
# reduplication words for n. and v. e.g. 奶奶, 试试, 旺旺
|
# reduplication words for n. and v. e.g. 奶奶, 试试, 旺旺
|
||||||
for j, item in enumerate(word):
|
for j, item in enumerate(word):
|
||||||
if (
|
if (
|
||||||
@@ -522,7 +519,7 @@ class ToneSandhi:
|
|||||||
finals = sum(finals_list, [])
|
finals = sum(finals_list, [])
|
||||||
return finals
|
return finals
|
||||||
|
|
||||||
def _bu_sandhi(self, word: str, finals: List[str]) -> List[str]:
|
def _bu_sandhi(self, word: str, finals: list[str]) -> list[str]:
|
||||||
# e.g. 看不懂
|
# e.g. 看不懂
|
||||||
if len(word) == 3 and word[1] == "不":
|
if len(word) == 3 and word[1] == "不":
|
||||||
finals[1] = finals[1][:-1] + "5"
|
finals[1] = finals[1][:-1] + "5"
|
||||||
@@ -533,7 +530,7 @@ class ToneSandhi:
|
|||||||
finals[i] = finals[i][:-1] + "2"
|
finals[i] = finals[i][:-1] + "2"
|
||||||
return finals
|
return finals
|
||||||
|
|
||||||
def _yi_sandhi(self, word: str, finals: List[str]) -> List[str]:
|
def _yi_sandhi(self, word: str, finals: list[str]) -> list[str]:
|
||||||
# "一" in number sequences, e.g. 一零零, 二一零
|
# "一" in number sequences, e.g. 一零零, 二一零
|
||||||
if word.find("一") != -1 and all(
|
if word.find("一") != -1 and all(
|
||||||
[item.isnumeric() for item in word if item != "一"]
|
[item.isnumeric() for item in word if item != "一"]
|
||||||
@@ -558,9 +555,9 @@ class ToneSandhi:
|
|||||||
finals[i] = finals[i][:-1] + "4"
|
finals[i] = finals[i][:-1] + "4"
|
||||||
return finals
|
return finals
|
||||||
|
|
||||||
def _split_word(self, word: str) -> List[str]:
|
def _split_word(self, word: str) -> list[str]:
|
||||||
word_list = jieba.cut_for_search(word)
|
word_list = jieba.cut_for_search(word)
|
||||||
word_list = sorted(word_list, key=lambda i: len(i), reverse=False)
|
word_list = sorted(word_list, key=lambda i: len(i), reverse=False) # type: ignore
|
||||||
first_subword = word_list[0]
|
first_subword = word_list[0]
|
||||||
first_begin_idx = word.find(first_subword)
|
first_begin_idx = word.find(first_subword)
|
||||||
if first_begin_idx == 0:
|
if first_begin_idx == 0:
|
||||||
@@ -571,7 +568,7 @@ class ToneSandhi:
|
|||||||
new_word_list = [second_subword, first_subword]
|
new_word_list = [second_subword, first_subword]
|
||||||
return new_word_list
|
return new_word_list
|
||||||
|
|
||||||
def _three_sandhi(self, word: str, finals: List[str]) -> List[str]:
|
def _three_sandhi(self, word: str, finals: list[str]) -> list[str]:
|
||||||
if len(word) == 2 and self._all_tone_three(finals):
|
if len(word) == 2 and self._all_tone_three(finals):
|
||||||
finals[0] = finals[0][:-1] + "2"
|
finals[0] = finals[0][:-1] + "2"
|
||||||
elif len(word) == 3:
|
elif len(word) == 3:
|
||||||
@@ -611,12 +608,12 @@ class ToneSandhi:
|
|||||||
|
|
||||||
return finals
|
return finals
|
||||||
|
|
||||||
def _all_tone_three(self, finals: List[str]) -> bool:
|
def _all_tone_three(self, finals: list[str]) -> bool:
|
||||||
return all(x[-1] == "3" for x in finals)
|
return all(x[-1] == "3" for x in finals)
|
||||||
|
|
||||||
# merge "不" and the word behind it
|
# merge "不" and the word behind it
|
||||||
# if don't merge, "不" sometimes appears alone according to jieba, which may occur sandhi error
|
# if don't merge, "不" sometimes appears alone according to jieba, which may occur sandhi error
|
||||||
def _merge_bu(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
def _merge_bu(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||||
new_seg = []
|
new_seg = []
|
||||||
last_word = ""
|
last_word = ""
|
||||||
for word, pos in seg:
|
for word, pos in seg:
|
||||||
@@ -636,7 +633,7 @@ class ToneSandhi:
|
|||||||
# e.g.
|
# e.g.
|
||||||
# input seg: [('听', 'v'), ('一', 'm'), ('听', 'v')]
|
# input seg: [('听', 'v'), ('一', 'm'), ('听', 'v')]
|
||||||
# output seg: [['听一听', 'v']]
|
# output seg: [['听一听', 'v']]
|
||||||
def _merge_yi(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
def _merge_yi(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||||
new_seg = [] * len(seg)
|
new_seg = [] * len(seg)
|
||||||
# function 1
|
# function 1
|
||||||
i = 0
|
i = 0
|
||||||
@@ -674,8 +671,8 @@ class ToneSandhi:
|
|||||||
|
|
||||||
# the first and the second words are all_tone_three
|
# the first and the second words are all_tone_three
|
||||||
def _merge_continuous_three_tones(
|
def _merge_continuous_three_tones(
|
||||||
self, seg: List[Tuple[str, str]]
|
self, seg: list[tuple[str, str]]
|
||||||
) -> List[Tuple[str, str]]:
|
) -> list[tuple[str, str]]:
|
||||||
new_seg = []
|
new_seg = []
|
||||||
sub_finals_list = [
|
sub_finals_list = [
|
||||||
lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3)
|
lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3)
|
||||||
@@ -709,8 +706,8 @@ class ToneSandhi:
|
|||||||
|
|
||||||
# the last char of first word and the first char of second word is tone_three
|
# the last char of first word and the first char of second word is tone_three
|
||||||
def _merge_continuous_three_tones_2(
|
def _merge_continuous_three_tones_2(
|
||||||
self, seg: List[Tuple[str, str]]
|
self, seg: list[tuple[str, str]]
|
||||||
) -> List[Tuple[str, str]]:
|
) -> list[tuple[str, str]]:
|
||||||
new_seg = []
|
new_seg = []
|
||||||
sub_finals_list = [
|
sub_finals_list = [
|
||||||
lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3)
|
lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3)
|
||||||
@@ -738,7 +735,7 @@ class ToneSandhi:
|
|||||||
new_seg.append([word, pos])
|
new_seg.append([word, pos])
|
||||||
return new_seg
|
return new_seg
|
||||||
|
|
||||||
def _merge_er(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
def _merge_er(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||||
new_seg = []
|
new_seg = []
|
||||||
for i, (word, pos) in enumerate(seg):
|
for i, (word, pos) in enumerate(seg):
|
||||||
if i - 1 >= 0 and word == "儿" and seg[i - 1][0] != "#":
|
if i - 1 >= 0 and word == "儿" and seg[i - 1][0] != "#":
|
||||||
@@ -747,7 +744,7 @@ class ToneSandhi:
|
|||||||
new_seg.append([word, pos])
|
new_seg.append([word, pos])
|
||||||
return new_seg
|
return new_seg
|
||||||
|
|
||||||
def _merge_reduplication(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
def _merge_reduplication(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||||
new_seg = []
|
new_seg = []
|
||||||
for i, (word, pos) in enumerate(seg):
|
for i, (word, pos) in enumerate(seg):
|
||||||
if new_seg and word == new_seg[-1][0]:
|
if new_seg and word == new_seg[-1][0]:
|
||||||
@@ -756,7 +753,7 @@ class ToneSandhi:
|
|||||||
new_seg.append([word, pos])
|
new_seg.append([word, pos])
|
||||||
return new_seg
|
return new_seg
|
||||||
|
|
||||||
def pre_merge_for_modify(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
def pre_merge_for_modify(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]:
|
||||||
seg = self._merge_bu(seg)
|
seg = self._merge_bu(seg)
|
||||||
try:
|
try:
|
||||||
seg = self._merge_yi(seg)
|
seg = self._merge_yi(seg)
|
||||||
@@ -768,7 +765,7 @@ class ToneSandhi:
|
|||||||
seg = self._merge_er(seg)
|
seg = self._merge_er(seg)
|
||||||
return seg
|
return seg
|
||||||
|
|
||||||
def modified_tone(self, word: str, pos: str, finals: List[str]) -> List[str]:
|
def modified_tone(self, word: str, pos: str, finals: list[str]) -> list[str]:
|
||||||
finals = self._bu_sandhi(word, finals)
|
finals = self._bu_sandhi(word, finals)
|
||||||
finals = self._yi_sandhi(word, finals)
|
finals = self._yi_sandhi(word, finals)
|
||||||
finals = self._neural_sandhi(word, pos, finals)
|
finals = self._neural_sandhi(word, pos, finals)
|
||||||
0
style_bert_vits2/nlp/english/__init__.py
Normal file
0
style_bert_vits2/nlp/english/__init__.py
Normal file
66
style_bert_vits2/nlp/english/bert_feature.py
Normal file
66
style_bert_vits2/nlp/english/bert_feature.py
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
device: str,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
英語のテキストから BERT の特徴量を抽出する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 英語のテキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
device (str): 推論に利用するデバイス
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: BERT の特徴量
|
||||||
|
"""
|
||||||
|
|
||||||
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
|
device = "cpu"
|
||||||
|
model = bert_models.load_model(Languages.EN).to(device) # type: ignore
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
with torch.no_grad():
|
||||||
|
tokenizer = bert_models.load_tokenizer(Languages.EN)
|
||||||
|
inputs = tokenizer(text, return_tensors="pt")
|
||||||
|
for i in inputs:
|
||||||
|
inputs[i] = inputs[i].to(device) # type: ignore
|
||||||
|
res = model(**inputs, output_hidden_states=True)
|
||||||
|
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="pt")
|
||||||
|
for i in style_inputs:
|
||||||
|
style_inputs[i] = style_inputs[i].to(device) # type: ignore
|
||||||
|
style_res = model(**style_inputs, output_hidden_states=True)
|
||||||
|
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
style_res_mean = style_res.mean(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 = (
|
||||||
|
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||||
|
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||||
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
|
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||||
|
|
||||||
|
return phone_level_feature.T
|
||||||
46
style_bert_vits2/nlp/english/cmudict.py
Normal file
46
style_bert_vits2/nlp/english/cmudict.py
Normal file
@@ -0,0 +1,46 @@
|
|||||||
|
import pickle
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
CMU_DICT_PATH = Path(__file__).parent / "cmudict.rep"
|
||||||
|
CACHE_PATH = Path(__file__).parent / "cmudict_cache.pickle"
|
||||||
|
|
||||||
|
|
||||||
|
def get_dict() -> dict[str, list[list[str]]]:
|
||||||
|
if CACHE_PATH.exists():
|
||||||
|
with open(CACHE_PATH, "rb") as pickle_file:
|
||||||
|
g2p_dict = pickle.load(pickle_file)
|
||||||
|
else:
|
||||||
|
g2p_dict = read_dict()
|
||||||
|
cache_dict(g2p_dict, CACHE_PATH)
|
||||||
|
|
||||||
|
return g2p_dict
|
||||||
|
|
||||||
|
|
||||||
|
def read_dict() -> dict[str, list[list[str]]]:
|
||||||
|
g2p_dict = {}
|
||||||
|
start_line = 49
|
||||||
|
with open(CMU_DICT_PATH, "r", encoding="utf-8") as f:
|
||||||
|
line = f.readline()
|
||||||
|
line_index = 1
|
||||||
|
while line:
|
||||||
|
if line_index >= start_line:
|
||||||
|
line = line.strip()
|
||||||
|
word_split = line.split(" ")
|
||||||
|
word = word_split[0]
|
||||||
|
|
||||||
|
syllable_split = word_split[1].split(" - ")
|
||||||
|
g2p_dict[word] = []
|
||||||
|
for syllable in syllable_split:
|
||||||
|
phone_split = syllable.split(" ")
|
||||||
|
g2p_dict[word].append(phone_split)
|
||||||
|
|
||||||
|
line_index = line_index + 1
|
||||||
|
line = f.readline()
|
||||||
|
|
||||||
|
return g2p_dict
|
||||||
|
|
||||||
|
|
||||||
|
def cache_dict(g2p_dict: dict[str, list[list[str]]], file_path: Path) -> None:
|
||||||
|
with open(file_path, "wb") as pickle_file:
|
||||||
|
pickle.dump(g2p_dict, pickle_file)
|
||||||
240
style_bert_vits2/nlp/english/g2p.py
Normal file
240
style_bert_vits2/nlp/english/g2p.py
Normal file
@@ -0,0 +1,240 @@
|
|||||||
|
import re
|
||||||
|
|
||||||
|
from g2p_en import G2p
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
from style_bert_vits2.nlp.english.cmudict import get_dict
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS, SYMBOLS
|
||||||
|
|
||||||
|
|
||||||
|
def g2p(text: str) -> tuple[list[str], list[int], list[int]]:
|
||||||
|
|
||||||
|
ARPA = {
|
||||||
|
"AH0",
|
||||||
|
"S",
|
||||||
|
"AH1",
|
||||||
|
"EY2",
|
||||||
|
"AE2",
|
||||||
|
"EH0",
|
||||||
|
"OW2",
|
||||||
|
"UH0",
|
||||||
|
"NG",
|
||||||
|
"B",
|
||||||
|
"G",
|
||||||
|
"AY0",
|
||||||
|
"M",
|
||||||
|
"AA0",
|
||||||
|
"F",
|
||||||
|
"AO0",
|
||||||
|
"ER2",
|
||||||
|
"UH1",
|
||||||
|
"IY1",
|
||||||
|
"AH2",
|
||||||
|
"DH",
|
||||||
|
"IY0",
|
||||||
|
"EY1",
|
||||||
|
"IH0",
|
||||||
|
"K",
|
||||||
|
"N",
|
||||||
|
"W",
|
||||||
|
"IY2",
|
||||||
|
"T",
|
||||||
|
"AA1",
|
||||||
|
"ER1",
|
||||||
|
"EH2",
|
||||||
|
"OY0",
|
||||||
|
"UH2",
|
||||||
|
"UW1",
|
||||||
|
"Z",
|
||||||
|
"AW2",
|
||||||
|
"AW1",
|
||||||
|
"V",
|
||||||
|
"UW2",
|
||||||
|
"AA2",
|
||||||
|
"ER",
|
||||||
|
"AW0",
|
||||||
|
"UW0",
|
||||||
|
"R",
|
||||||
|
"OW1",
|
||||||
|
"EH1",
|
||||||
|
"ZH",
|
||||||
|
"AE0",
|
||||||
|
"IH2",
|
||||||
|
"IH",
|
||||||
|
"Y",
|
||||||
|
"JH",
|
||||||
|
"P",
|
||||||
|
"AY1",
|
||||||
|
"EY0",
|
||||||
|
"OY2",
|
||||||
|
"TH",
|
||||||
|
"HH",
|
||||||
|
"D",
|
||||||
|
"ER0",
|
||||||
|
"CH",
|
||||||
|
"AO1",
|
||||||
|
"AE1",
|
||||||
|
"AO2",
|
||||||
|
"OY1",
|
||||||
|
"AY2",
|
||||||
|
"IH1",
|
||||||
|
"OW0",
|
||||||
|
"L",
|
||||||
|
"SH",
|
||||||
|
}
|
||||||
|
|
||||||
|
_g2p = G2p()
|
||||||
|
|
||||||
|
phones = []
|
||||||
|
tones = []
|
||||||
|
phone_len = []
|
||||||
|
# tokens = [tokenizer.tokenize(i) for i in words]
|
||||||
|
words = __text_to_words(text)
|
||||||
|
eng_dict = get_dict()
|
||||||
|
|
||||||
|
for word in words:
|
||||||
|
temp_phones, temp_tones = [], []
|
||||||
|
if len(word) > 1:
|
||||||
|
if "'" in word:
|
||||||
|
word = ["".join(word)]
|
||||||
|
for w in word:
|
||||||
|
if w in PUNCTUATIONS:
|
||||||
|
temp_phones.append(w)
|
||||||
|
temp_tones.append(0)
|
||||||
|
continue
|
||||||
|
if w.upper() in eng_dict:
|
||||||
|
phns, tns = __refine_syllables(eng_dict[w.upper()])
|
||||||
|
temp_phones += [__post_replace_ph(i) for i in phns]
|
||||||
|
temp_tones += tns
|
||||||
|
# w2ph.append(len(phns))
|
||||||
|
else:
|
||||||
|
phone_list = list(filter(lambda p: p != " ", _g2p(w))) # type: ignore
|
||||||
|
phns = []
|
||||||
|
tns = []
|
||||||
|
for ph in phone_list:
|
||||||
|
if ph in ARPA:
|
||||||
|
ph, tn = __refine_ph(ph)
|
||||||
|
phns.append(ph)
|
||||||
|
tns.append(tn)
|
||||||
|
else:
|
||||||
|
phns.append(ph)
|
||||||
|
tns.append(0)
|
||||||
|
temp_phones += [__post_replace_ph(i) for i in phns]
|
||||||
|
temp_tones += tns
|
||||||
|
phones += temp_phones
|
||||||
|
tones += temp_tones
|
||||||
|
phone_len.append(len(temp_phones))
|
||||||
|
# phones = [post_replace_ph(i) for i in phones]
|
||||||
|
|
||||||
|
word2ph = []
|
||||||
|
for token, pl in zip(words, phone_len):
|
||||||
|
word_len = len(token)
|
||||||
|
|
||||||
|
aaa = __distribute_phone(pl, word_len)
|
||||||
|
word2ph += aaa
|
||||||
|
|
||||||
|
phones = ["_"] + phones + ["_"]
|
||||||
|
tones = [0] + tones + [0]
|
||||||
|
word2ph = [1] + word2ph + [1]
|
||||||
|
assert len(phones) == len(tones), text
|
||||||
|
assert len(phones) == sum(word2ph), text
|
||||||
|
|
||||||
|
return phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
|
def __post_replace_ph(ph: str) -> str:
|
||||||
|
REPLACE_MAP = {
|
||||||
|
":": ",",
|
||||||
|
";": ",",
|
||||||
|
",": ",",
|
||||||
|
"。": ".",
|
||||||
|
"!": "!",
|
||||||
|
"?": "?",
|
||||||
|
"\n": ".",
|
||||||
|
"·": ",",
|
||||||
|
"、": ",",
|
||||||
|
"…": "...",
|
||||||
|
"···": "...",
|
||||||
|
"・・・": "...",
|
||||||
|
"v": "V",
|
||||||
|
}
|
||||||
|
if ph in REPLACE_MAP.keys():
|
||||||
|
ph = REPLACE_MAP[ph]
|
||||||
|
if ph in SYMBOLS:
|
||||||
|
return ph
|
||||||
|
if ph not in SYMBOLS:
|
||||||
|
ph = "UNK"
|
||||||
|
return ph
|
||||||
|
|
||||||
|
|
||||||
|
def __refine_ph(phn: str) -> tuple[str, int]:
|
||||||
|
tone = 0
|
||||||
|
if re.search(r"\d$", phn):
|
||||||
|
tone = int(phn[-1]) + 1
|
||||||
|
phn = phn[:-1]
|
||||||
|
else:
|
||||||
|
tone = 3
|
||||||
|
return phn.lower(), tone
|
||||||
|
|
||||||
|
|
||||||
|
def __refine_syllables(syllables: list[list[str]]) -> tuple[list[str], list[int]]:
|
||||||
|
tones = []
|
||||||
|
phonemes = []
|
||||||
|
for phn_list in syllables:
|
||||||
|
for i in range(len(phn_list)):
|
||||||
|
phn = phn_list[i]
|
||||||
|
phn, tone = __refine_ph(phn)
|
||||||
|
phonemes.append(phn)
|
||||||
|
tones.append(tone)
|
||||||
|
return phonemes, tones
|
||||||
|
|
||||||
|
|
||||||
|
def __distribute_phone(n_phone: int, n_word: int) -> list[int]:
|
||||||
|
phones_per_word = [0] * n_word
|
||||||
|
for task in range(n_phone):
|
||||||
|
min_tasks = min(phones_per_word)
|
||||||
|
min_index = phones_per_word.index(min_tasks)
|
||||||
|
phones_per_word[min_index] += 1
|
||||||
|
return phones_per_word
|
||||||
|
|
||||||
|
|
||||||
|
def __text_to_words(text: str) -> list[list[str]]:
|
||||||
|
tokenizer = bert_models.load_tokenizer(Languages.EN)
|
||||||
|
tokens = tokenizer.tokenize(text)
|
||||||
|
words = []
|
||||||
|
for idx, t in enumerate(tokens):
|
||||||
|
if t.startswith("▁"):
|
||||||
|
words.append([t[1:]])
|
||||||
|
else:
|
||||||
|
if t in PUNCTUATIONS:
|
||||||
|
if idx == len(tokens) - 1:
|
||||||
|
words.append([f"{t}"])
|
||||||
|
else:
|
||||||
|
if (
|
||||||
|
not tokens[idx + 1].startswith("▁")
|
||||||
|
and tokens[idx + 1] not in PUNCTUATIONS
|
||||||
|
):
|
||||||
|
if idx == 0:
|
||||||
|
words.append([])
|
||||||
|
words[-1].append(f"{t}")
|
||||||
|
else:
|
||||||
|
words.append([f"{t}"])
|
||||||
|
else:
|
||||||
|
if idx == 0:
|
||||||
|
words.append([])
|
||||||
|
words[-1].append(f"{t}")
|
||||||
|
return words
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# print(get_dict())
|
||||||
|
# print(eng_word_to_phoneme("hello"))
|
||||||
|
print(g2p("In this paper, we propose 1 DSPGAN, a GAN-based universal vocoder."))
|
||||||
|
# all_phones = set()
|
||||||
|
# eng_dict = get_dict()
|
||||||
|
# for k, syllables in eng_dict.items():
|
||||||
|
# for group in syllables:
|
||||||
|
# for ph in group:
|
||||||
|
# all_phones.add(ph)
|
||||||
|
# print(all_phones)
|
||||||
132
style_bert_vits2/nlp/english/normalizer.py
Normal file
132
style_bert_vits2/nlp/english/normalizer.py
Normal file
@@ -0,0 +1,132 @@
|
|||||||
|
import re
|
||||||
|
|
||||||
|
import inflect
|
||||||
|
|
||||||
|
|
||||||
|
__INFLECT = inflect.engine()
|
||||||
|
__COMMA_NUMBER_PATTERN = re.compile(r"([0-9][0-9\,]+[0-9])")
|
||||||
|
__DECIMAL_NUMBER_PATTERN = re.compile(r"([0-9]+\.[0-9]+)")
|
||||||
|
__POUNDS_PATTERN = re.compile(r"£([0-9\,]*[0-9]+)")
|
||||||
|
__DOLLARS_PATTERN = re.compile(r"\$([0-9\.\,]*[0-9]+)")
|
||||||
|
__ORDINAL_PATTERN = re.compile(r"[0-9]+(st|nd|rd|th)")
|
||||||
|
__NUMBER_PATTERN = re.compile(r"[0-9]+")
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
text = __normalize_numbers(text)
|
||||||
|
text = replace_punctuation(text)
|
||||||
|
text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def replace_punctuation(text: str) -> str:
|
||||||
|
REPLACE_MAP = {
|
||||||
|
":": ",",
|
||||||
|
";": ",",
|
||||||
|
",": ",",
|
||||||
|
"。": ".",
|
||||||
|
"!": "!",
|
||||||
|
"?": "?",
|
||||||
|
"\n": ".",
|
||||||
|
".": ".",
|
||||||
|
"…": "...",
|
||||||
|
"···": "...",
|
||||||
|
"・・・": "...",
|
||||||
|
"·": ",",
|
||||||
|
"・": ",",
|
||||||
|
"、": ",",
|
||||||
|
"$": ".",
|
||||||
|
"“": "'",
|
||||||
|
"”": "'",
|
||||||
|
'"': "'",
|
||||||
|
"‘": "'",
|
||||||
|
"’": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"《": "'",
|
||||||
|
"》": "'",
|
||||||
|
"【": "'",
|
||||||
|
"】": "'",
|
||||||
|
"[": "'",
|
||||||
|
"]": "'",
|
||||||
|
"—": "-",
|
||||||
|
"−": "-",
|
||||||
|
"~": "-",
|
||||||
|
"~": "-",
|
||||||
|
"「": "'",
|
||||||
|
"」": "'",
|
||||||
|
}
|
||||||
|
pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys()))
|
||||||
|
replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text)
|
||||||
|
# replaced_text = re.sub(
|
||||||
|
# r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005"
|
||||||
|
# + "".join(punctuation)
|
||||||
|
# + r"]+",
|
||||||
|
# "",
|
||||||
|
# replaced_text,
|
||||||
|
# )
|
||||||
|
return replaced_text
|
||||||
|
|
||||||
|
|
||||||
|
def __normalize_numbers(text: str) -> str:
|
||||||
|
text = re.sub(__COMMA_NUMBER_PATTERN, __remove_commas, text)
|
||||||
|
text = re.sub(__POUNDS_PATTERN, r"\1 pounds", text)
|
||||||
|
text = re.sub(__DOLLARS_PATTERN, __expand_dollars, text)
|
||||||
|
text = re.sub(__DECIMAL_NUMBER_PATTERN, __expand_decimal_point, text)
|
||||||
|
text = re.sub(__ORDINAL_PATTERN, __expand_ordinal, text)
|
||||||
|
text = re.sub(__NUMBER_PATTERN, __expand_number, text)
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def __expand_dollars(m: re.Match[str]) -> str:
|
||||||
|
match = m.group(1)
|
||||||
|
parts = match.split(".")
|
||||||
|
if len(parts) > 2:
|
||||||
|
return match + " dollars" # Unexpected format
|
||||||
|
dollars = int(parts[0]) if parts[0] else 0
|
||||||
|
cents = int(parts[1]) if len(parts) > 1 and parts[1] else 0
|
||||||
|
if dollars and cents:
|
||||||
|
dollar_unit = "dollar" if dollars == 1 else "dollars"
|
||||||
|
cent_unit = "cent" if cents == 1 else "cents"
|
||||||
|
return "%s %s, %s %s" % (dollars, dollar_unit, cents, cent_unit)
|
||||||
|
elif dollars:
|
||||||
|
dollar_unit = "dollar" if dollars == 1 else "dollars"
|
||||||
|
return "%s %s" % (dollars, dollar_unit)
|
||||||
|
elif cents:
|
||||||
|
cent_unit = "cent" if cents == 1 else "cents"
|
||||||
|
return "%s %s" % (cents, cent_unit)
|
||||||
|
else:
|
||||||
|
return "zero dollars"
|
||||||
|
|
||||||
|
|
||||||
|
def __remove_commas(m: re.Match[str]) -> str:
|
||||||
|
return m.group(1).replace(",", "")
|
||||||
|
|
||||||
|
|
||||||
|
def __expand_ordinal(m: re.Match[str]) -> str:
|
||||||
|
return __INFLECT.number_to_words(m.group(0)) # type: ignore
|
||||||
|
|
||||||
|
|
||||||
|
def __expand_number(m: re.Match[str]) -> str:
|
||||||
|
num = int(m.group(0))
|
||||||
|
if num > 1000 and num < 3000:
|
||||||
|
if num == 2000:
|
||||||
|
return "two thousand"
|
||||||
|
elif num > 2000 and num < 2010:
|
||||||
|
return "two thousand " + __INFLECT.number_to_words(num % 100) # type: ignore
|
||||||
|
elif num % 100 == 0:
|
||||||
|
return __INFLECT.number_to_words(num // 100) + " hundred" # type: ignore
|
||||||
|
else:
|
||||||
|
return __INFLECT.number_to_words(
|
||||||
|
num, andword="", zero="oh", group=2 # type: ignore
|
||||||
|
).replace(
|
||||||
|
", ", " "
|
||||||
|
) # type: ignore
|
||||||
|
else:
|
||||||
|
return __INFLECT.number_to_words(num, andword="") # type: ignore
|
||||||
|
|
||||||
|
|
||||||
|
def __expand_decimal_point(m: re.Match[str]) -> str:
|
||||||
|
return m.group(1).replace(".", " point ")
|
||||||
0
style_bert_vits2/nlp/japanese/__init__.py
Normal file
0
style_bert_vits2/nlp/japanese/__init__.py
Normal file
73
style_bert_vits2/nlp/japanese/bert_feature.py
Normal file
73
style_bert_vits2/nlp/japanese/bert_feature.py
Normal file
@@ -0,0 +1,73 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
||||||
|
|
||||||
|
|
||||||
|
def extract_bert_feature(
|
||||||
|
text: str,
|
||||||
|
word2ph: list[int],
|
||||||
|
device: str,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = 0.7,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
日本語のテキストから BERT の特徴量を抽出する
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 日本語のテキスト
|
||||||
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
device (str): 推論に利用するデバイス
|
||||||
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
torch.Tensor: 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])
|
||||||
|
|
||||||
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
|
device = "cpu"
|
||||||
|
model = bert_models.load_model(Languages.JP).to(device) # type: ignore
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
with torch.no_grad():
|
||||||
|
tokenizer = bert_models.load_tokenizer(Languages.JP)
|
||||||
|
inputs = tokenizer(text, return_tensors="pt")
|
||||||
|
for i in inputs:
|
||||||
|
inputs[i] = inputs[i].to(device) # type: ignore
|
||||||
|
res = model(**inputs, output_hidden_states=True)
|
||||||
|
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="pt")
|
||||||
|
for i in style_inputs:
|
||||||
|
style_inputs[i] = style_inputs[i].to(device) # type: ignore
|
||||||
|
style_res = model(**style_inputs, output_hidden_states=True)
|
||||||
|
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
style_res_mean = style_res.mean(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 = (
|
||||||
|
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||||
|
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||||
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
|
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||||
|
|
||||||
|
return phone_level_feature.T
|
||||||
489
style_bert_vits2/nlp/japanese/g2p.py
Normal file
489
style_bert_vits2/nlp/japanese/g2p.py
Normal file
@@ -0,0 +1,489 @@
|
|||||||
|
import re
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
|
from style_bert_vits2.nlp.japanese.mora_list import MORA_KATA_TO_MORA_PHONEMES
|
||||||
|
from style_bert_vits2.nlp.japanese.normalizer import replace_punctuation
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||||
|
|
||||||
|
|
||||||
|
def g2p(
|
||||||
|
norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False
|
||||||
|
) -> tuple[list[str], list[int], list[int]]:
|
||||||
|
"""
|
||||||
|
他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、
|
||||||
|
- phones: 音素のリスト(ただし `!` や `,` や `.` など punctuation が含まれうる)
|
||||||
|
- tones: アクセントのリスト、0(低)と1(高)からなり、phones と同じ長さ
|
||||||
|
- word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
|
のタプルを返す。
|
||||||
|
ただし `phones` と `tones` の最初と終わりに `_` が入り、応じて `word2ph` の最初と最後に 1 が追加される。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
norm_text (str): 正規化されたテキスト
|
||||||
|
use_jp_extra (bool, optional): False の場合、「ん」の音素を「N」ではなく「n」とする。Defaults to True.
|
||||||
|
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[list[str], list[int], list[int]]: 音素のリスト、アクセントのリスト、word2ph のリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# pyopenjtalk のフルコンテキストラベルを使ってアクセントを取り出すと、punctuation の位置が消えてしまい情報が失われてしまう:
|
||||||
|
# 「こんにちは、世界。」と「こんにちは!世界。」と「こんにちは!!!???世界……。」は全て同じになる。
|
||||||
|
# よって、まず punctuation 無しの音素とアクセントのリストを作り、
|
||||||
|
# それとは別に pyopenjtalk.run_frontend() で得られる音素リスト(こちらは punctuation が保持される)を使い、
|
||||||
|
# アクセント割当をしなおすことによって punctuation を含めた音素とアクセントのリストを作る。
|
||||||
|
|
||||||
|
# punctuation がすべて消えた、音素とアクセントのタプルのリスト(「ん」は「N」)
|
||||||
|
phone_tone_list_wo_punct = __g2phone_tone_wo_punct(norm_text)
|
||||||
|
|
||||||
|
# sep_text: 単語単位の単語のリスト、読めない文字があったら raise_yomi_error なら例外、そうでないなら読めない文字が消えて返ってくる
|
||||||
|
# sep_kata: 単語単位の単語のカタカナ読みのリスト
|
||||||
|
sep_text, sep_kata = text_to_sep_kata(norm_text, raise_yomi_error=raise_yomi_error)
|
||||||
|
|
||||||
|
# sep_phonemes: 各単語ごとの音素のリストのリスト
|
||||||
|
sep_phonemes = __handle_long([__kata_to_phoneme_list(i) for i in sep_kata])
|
||||||
|
|
||||||
|
# phone_w_punct: sep_phonemes を結合した、punctuation を元のまま保持した音素列
|
||||||
|
phone_w_punct: list[str] = []
|
||||||
|
for i in sep_phonemes:
|
||||||
|
phone_w_punct += i
|
||||||
|
|
||||||
|
# punctuation 無しのアクセント情報を使って、punctuation を含めたアクセント情報を作る
|
||||||
|
phone_tone_list = __align_tones(phone_w_punct, phone_tone_list_wo_punct)
|
||||||
|
# logger.debug(f"phone_tone_list:\n{phone_tone_list}")
|
||||||
|
|
||||||
|
# word2ph は厳密な解答は不可能なので(「今日」「眼鏡」等の熟字訓が存在)、
|
||||||
|
# Bert-VITS2 では、単語単位の分割を使って、単語の文字ごとにだいたい均等に音素を分配する
|
||||||
|
|
||||||
|
# sep_text から、各単語を1文字1文字分割して、文字のリスト(のリスト)を作る
|
||||||
|
sep_tokenized: list[list[str]] = []
|
||||||
|
for i in sep_text:
|
||||||
|
if i not in PUNCTUATIONS:
|
||||||
|
sep_tokenized.append(
|
||||||
|
bert_models.load_tokenizer(Languages.JP).tokenize(i)
|
||||||
|
) # ここでおそらく`i`が文字単位に分割される
|
||||||
|
else:
|
||||||
|
sep_tokenized.append([i])
|
||||||
|
|
||||||
|
# 各単語について、音素の数と文字の数を比較して、均等っぽく分配する
|
||||||
|
word2ph = []
|
||||||
|
for token, phoneme in zip(sep_tokenized, sep_phonemes):
|
||||||
|
phone_len = len(phoneme)
|
||||||
|
word_len = len(token)
|
||||||
|
word2ph += __distribute_phone(phone_len, word_len)
|
||||||
|
|
||||||
|
# 最初と最後に `_` 記号を追加、アクセントは 0(低)、word2ph もそれに合わせて追加
|
||||||
|
phone_tone_list = [("_", 0)] + phone_tone_list + [("_", 0)]
|
||||||
|
word2ph = [1] + word2ph + [1]
|
||||||
|
|
||||||
|
phones = [phone for phone, _ in phone_tone_list]
|
||||||
|
tones = [tone for _, tone in phone_tone_list]
|
||||||
|
|
||||||
|
assert len(phones) == sum(word2ph), f"{len(phones)} != {sum(word2ph)}"
|
||||||
|
|
||||||
|
# use_jp_extra でない場合は「N」を「n」に変換
|
||||||
|
if not use_jp_extra:
|
||||||
|
phones = [phone if phone != "N" else "n" for phone in phones]
|
||||||
|
|
||||||
|
return phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
|
def text_to_sep_kata(
|
||||||
|
norm_text: str, raise_yomi_error: bool = False
|
||||||
|
) -> tuple[list[str], list[str]]:
|
||||||
|
"""
|
||||||
|
`normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、
|
||||||
|
分割された単語リストとその読み(カタカナ or 記号1文字)のリストのタプルを返す。
|
||||||
|
単語分割結果は、`g2p()` の `word2ph` で1文字あたりに割り振る音素記号の数を決めるために使う。
|
||||||
|
例:
|
||||||
|
`私はそう思う!って感じ?` →
|
||||||
|
["私", "は", "そう", "思う", "!", "って", "感じ", "?"], ["ワタシ", "ワ", "ソー", "オモウ", "!", "ッテ", "カンジ", "?"]
|
||||||
|
|
||||||
|
Args:
|
||||||
|
norm_text (str): 正規化されたテキスト
|
||||||
|
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[list[str], list[str]]: 分割された単語リストと、その読み(カタカナ or 記号1文字)のリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# parsed: OpenJTalkの解析結果
|
||||||
|
parsed = pyopenjtalk.run_frontend(norm_text)
|
||||||
|
sep_text: list[str] = []
|
||||||
|
sep_kata: list[str] = []
|
||||||
|
|
||||||
|
for parts in parsed:
|
||||||
|
# word: 実際の単語の文字列
|
||||||
|
# yomi: その読み、但し無声化サインの`’`は除去
|
||||||
|
word, yomi = replace_punctuation(parts["string"]), parts["pron"].replace(
|
||||||
|
"’", ""
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
ここで `yomi` の取りうる値は以下の通りのはず。
|
||||||
|
- `word` が通常単語 → 通常の読み(カタカナ)
|
||||||
|
(カタカナからなり、長音記号も含みうる、`アー` 等)
|
||||||
|
- `word` が `ー` から始まる → `ーラー` や `ーーー` など
|
||||||
|
- `word` が句読点や空白等 → `、`
|
||||||
|
- `word` が punctuation の繰り返し → 全角にしたもの
|
||||||
|
基本的に punctuation は1文字ずつ分かれるが、何故かある程度連続すると1つにまとまる。
|
||||||
|
他にも `word` が読めないキリル文字アラビア文字等が来ると `、` になるが、正規化でこの場合は起きないはず。
|
||||||
|
また元のコードでは `yomi` が空白の場合の処理があったが、これは起きないはず。
|
||||||
|
処理すべきは `yomi` が `、` の場合のみのはず。
|
||||||
|
"""
|
||||||
|
assert yomi != "", f"Empty yomi: {word}"
|
||||||
|
if yomi == "、":
|
||||||
|
# word は正規化されているので、`.`, `,`, `!`, `'`, `-`, `--` のいずれか
|
||||||
|
if not set(word).issubset(set(PUNCTUATIONS)): # 記号繰り返しか判定
|
||||||
|
# ここは pyopenjtalk が読めない文字等のときに起こる
|
||||||
|
if raise_yomi_error:
|
||||||
|
raise YomiError(f"Cannot read: {word} in:\n{norm_text}")
|
||||||
|
logger.warning(f"Ignoring unknown: {word} in:\n{norm_text}")
|
||||||
|
continue
|
||||||
|
# yomi は元の記号のままに変更
|
||||||
|
yomi = word
|
||||||
|
elif yomi == "?":
|
||||||
|
assert word == "?", f"yomi `?` comes from: {word}"
|
||||||
|
yomi = "?"
|
||||||
|
sep_text.append(word)
|
||||||
|
sep_kata.append(yomi)
|
||||||
|
|
||||||
|
return sep_text, sep_kata
|
||||||
|
|
||||||
|
|
||||||
|
def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
テキストに対して、音素とアクセント(0か1)のペアのリストを返す。
|
||||||
|
ただし「!」「.」「?」等の非音素記号 (punctuation) は全て消える(ポーズ記号も残さない)。
|
||||||
|
非音素記号を含める処理は `align_tones()` で行われる。
|
||||||
|
また「っ」は「q」に、「ん」は「N」に変換される。
|
||||||
|
例: "こんにちは、世界ー。。元気?!" →
|
||||||
|
[('k', 0), ('o', 0), ('N', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('N', 0), ('k', 0), ('i', 0)]
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): テキスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[tuple[str, int]]: 音素とアクセントのペアのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
prosodies = __pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True)
|
||||||
|
# logger.debug(f"prosodies: {prosodies}")
|
||||||
|
result: list[tuple[str, int]] = []
|
||||||
|
current_phrase: list[tuple[str, int]] = []
|
||||||
|
current_tone = 0
|
||||||
|
|
||||||
|
for i, letter in enumerate(prosodies):
|
||||||
|
# 特殊記号の処理
|
||||||
|
|
||||||
|
# 文頭記号、無視する
|
||||||
|
if letter == "^":
|
||||||
|
assert i == 0, "Unexpected ^"
|
||||||
|
# アクセント句の終わりに来る記号
|
||||||
|
elif letter in ("$", "?", "_", "#"):
|
||||||
|
# 保持しているフレーズを、アクセント数値を 0-1 に修正し結果に追加
|
||||||
|
result.extend(__fix_phone_tone(current_phrase))
|
||||||
|
# 末尾に来る終了記号、無視(文中の疑問文は `_` になる)
|
||||||
|
if letter in ("$", "?"):
|
||||||
|
assert i == len(prosodies) - 1, f"Unexpected {letter}"
|
||||||
|
# あとは "_"(ポーズ)と "#"(アクセント句の境界)のみ
|
||||||
|
# これらは残さず、次のアクセント句に備える。
|
||||||
|
current_phrase = []
|
||||||
|
# 0 を基準点にしてそこから上昇・下降する(負の場合は上の `fix_phone_tone` で直る)
|
||||||
|
current_tone = 0
|
||||||
|
# アクセント上昇記号
|
||||||
|
elif letter == "[":
|
||||||
|
current_tone = current_tone + 1
|
||||||
|
# アクセント下降記号
|
||||||
|
elif letter == "]":
|
||||||
|
current_tone = current_tone - 1
|
||||||
|
# それ以外は通常の音素
|
||||||
|
else:
|
||||||
|
if letter == "cl": # 「っ」の処理
|
||||||
|
letter = "q"
|
||||||
|
# elif letter == "N": # 「ん」の処理
|
||||||
|
# letter = "n"
|
||||||
|
current_phrase.append((letter, current_tone))
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def __pyopenjtalk_g2p_prosody(
|
||||||
|
text: str, drop_unvoiced_vowels: bool = True
|
||||||
|
) -> list[str]:
|
||||||
|
"""
|
||||||
|
ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。
|
||||||
|
ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py
|
||||||
|
------------------------------------------------------------------------------------------
|
||||||
|
|
||||||
|
Extract phoneme + prosoody symbol sequence from input full-context labels.
|
||||||
|
|
||||||
|
The algorithm is based on `Prosodic features control by symbols as input of
|
||||||
|
sequence-to-sequence acoustic modeling for neural TTS`_ with some r9y9's tweaks.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): Input text.
|
||||||
|
drop_unvoiced_vowels (bool): whether to drop unvoiced vowels.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[str]: List of phoneme + prosody symbols.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
>>> from espnet2.text.phoneme_tokenizer import pyopenjtalk_g2p_prosody
|
||||||
|
>>> pyopenjtalk_g2p_prosody("こんにちは。")
|
||||||
|
['^', 'k', 'o', '[', 'N', 'n', 'i', 'ch', 'i', 'w', 'a', '$']
|
||||||
|
|
||||||
|
.. _`Prosodic features control by symbols as input of sequence-to-sequence acoustic
|
||||||
|
modeling for neural TTS`: https://doi.org/10.1587/transinf.2020EDP7104
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _numeric_feature_by_regex(regex: str, s: str) -> int:
|
||||||
|
match = re.search(regex, s)
|
||||||
|
if match is None:
|
||||||
|
return -50
|
||||||
|
return int(match.group(1))
|
||||||
|
|
||||||
|
labels = pyopenjtalk.make_label(pyopenjtalk.run_frontend(text))
|
||||||
|
N = len(labels)
|
||||||
|
|
||||||
|
phones = []
|
||||||
|
for n in range(N):
|
||||||
|
lab_curr = labels[n]
|
||||||
|
|
||||||
|
# current phoneme
|
||||||
|
p3 = re.search(r"\-(.*?)\+", lab_curr).group(1) # type: ignore
|
||||||
|
# deal unvoiced vowels as normal vowels
|
||||||
|
if drop_unvoiced_vowels and p3 in "AEIOU":
|
||||||
|
p3 = p3.lower()
|
||||||
|
|
||||||
|
# deal with sil at the beginning and the end of text
|
||||||
|
if p3 == "sil":
|
||||||
|
assert n == 0 or n == N - 1
|
||||||
|
if n == 0:
|
||||||
|
phones.append("^")
|
||||||
|
elif n == N - 1:
|
||||||
|
# check question form or not
|
||||||
|
e3 = _numeric_feature_by_regex(r"!(\d+)_", lab_curr)
|
||||||
|
if e3 == 0:
|
||||||
|
phones.append("$")
|
||||||
|
elif e3 == 1:
|
||||||
|
phones.append("?")
|
||||||
|
continue
|
||||||
|
elif p3 == "pau":
|
||||||
|
phones.append("_")
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
phones.append(p3)
|
||||||
|
|
||||||
|
# accent type and position info (forward or backward)
|
||||||
|
a1 = _numeric_feature_by_regex(r"/A:([0-9\-]+)\+", lab_curr)
|
||||||
|
a2 = _numeric_feature_by_regex(r"\+(\d+)\+", lab_curr)
|
||||||
|
a3 = _numeric_feature_by_regex(r"\+(\d+)/", lab_curr)
|
||||||
|
|
||||||
|
# number of mora in accent phrase
|
||||||
|
f1 = _numeric_feature_by_regex(r"/F:(\d+)_", lab_curr)
|
||||||
|
|
||||||
|
a2_next = _numeric_feature_by_regex(r"\+(\d+)\+", labels[n + 1])
|
||||||
|
# accent phrase border
|
||||||
|
if a3 == 1 and a2_next == 1 and p3 in "aeiouAEIOUNcl":
|
||||||
|
phones.append("#")
|
||||||
|
# pitch falling
|
||||||
|
elif a1 == 0 and a2_next == a2 + 1 and a2 != f1:
|
||||||
|
phones.append("]")
|
||||||
|
# pitch rising
|
||||||
|
elif a2 == 1 and a2_next == 2:
|
||||||
|
phones.append("[")
|
||||||
|
|
||||||
|
return phones
|
||||||
|
|
||||||
|
|
||||||
|
def __fix_phone_tone(phone_tone_list: list[tuple[str, int]]) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
`phone_tone_list` の tone(アクセントの値)を 0 か 1 の範囲に修正する。
|
||||||
|
例: [(a, 0), (i, -1), (u, -1)] → [(a, 1), (i, 0), (u, 0)]
|
||||||
|
|
||||||
|
Args:
|
||||||
|
phone_tone_list (list[tuple[str, int]]): 音素とアクセントのペアのリスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[tuple[str, int]]: 修正された音素とアクセントのペアのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
tone_values = set(tone for _, tone in phone_tone_list)
|
||||||
|
if len(tone_values) == 1:
|
||||||
|
assert tone_values == {0}, tone_values
|
||||||
|
return phone_tone_list
|
||||||
|
elif len(tone_values) == 2:
|
||||||
|
if tone_values == {0, 1}:
|
||||||
|
return phone_tone_list
|
||||||
|
elif tone_values == {-1, 0}:
|
||||||
|
return [
|
||||||
|
(letter, 0 if tone == -1 else 1) for letter, tone in phone_tone_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unexpected tone values: {tone_values}")
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unexpected tone values: {tone_values}")
|
||||||
|
|
||||||
|
|
||||||
|
def __handle_long(sep_phonemes: list[list[str]]) -> list[list[str]]:
|
||||||
|
"""
|
||||||
|
フレーズごとに分かれた音素(長音記号がそのまま)のリストのリスト `sep_phonemes` を受け取り、
|
||||||
|
その長音記号を処理して、音素のリストのリストを返す。
|
||||||
|
基本的には直前の音素を伸ばすが、直前の音素が母音でない場合もしくは冒頭の場合は、
|
||||||
|
おそらく長音記号とダッシュを勘違いしていると思われるので、ダッシュに対応する音素 `-` に変換する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sep_phonemes (list[list[str]]): フレーズごとに分かれた音素のリストのリスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[list[str]]: 長音記号を処理した音素のリストのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 母音の集合 (便宜上「ん」を含める)
|
||||||
|
VOWELS = {"a", "i", "u", "e", "o", "N"}
|
||||||
|
|
||||||
|
for i in range(len(sep_phonemes)):
|
||||||
|
if len(sep_phonemes[i]) == 0:
|
||||||
|
# 空白文字等でリストが空の場合
|
||||||
|
continue
|
||||||
|
if sep_phonemes[i][0] == "ー":
|
||||||
|
if i != 0:
|
||||||
|
prev_phoneme = sep_phonemes[i - 1][-1]
|
||||||
|
if prev_phoneme in VOWELS:
|
||||||
|
# 母音と「ん」のあとの伸ばし棒なので、その母音に変換
|
||||||
|
sep_phonemes[i][0] = sep_phonemes[i - 1][-1]
|
||||||
|
else:
|
||||||
|
# 「。ーー」等おそらく予期しない長音記号
|
||||||
|
# ダッシュの勘違いだと思われる
|
||||||
|
sep_phonemes[i][0] = "-"
|
||||||
|
else:
|
||||||
|
# 冒頭に長音記号が来ていおり、これはダッシュの勘違いと思われる
|
||||||
|
sep_phonemes[i][0] = "-"
|
||||||
|
if "ー" in sep_phonemes[i]:
|
||||||
|
for j in range(len(sep_phonemes[i])):
|
||||||
|
if sep_phonemes[i][j] == "ー":
|
||||||
|
sep_phonemes[i][j] = sep_phonemes[i][j - 1][-1]
|
||||||
|
|
||||||
|
return sep_phonemes
|
||||||
|
|
||||||
|
|
||||||
|
def __kata_to_phoneme_list(text: str) -> list[str]:
|
||||||
|
"""
|
||||||
|
原則カタカナの `text` を受け取り、それをそのままいじらずに音素記号のリストに変換。
|
||||||
|
注意点:
|
||||||
|
- punctuation かその繰り返しが来た場合、punctuation たちをそのままリストにして返す。
|
||||||
|
- 冒頭に続く「ー」はそのまま「ー」のままにする(`handle_long()` で処理される)
|
||||||
|
- 文中の「ー」は前の音素記号の最後の音素記号に変換される。
|
||||||
|
例:
|
||||||
|
`ーーソーナノカーー` → ["ー", "ー", "s", "o", "o", "n", "a", "n", "o", "k", "a", "a", "a"]
|
||||||
|
`?` → ["?"]
|
||||||
|
`!?!?!?!?!` → ["!", "?", "!", "?", "!", "?", "!", "?", "!"]
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): カタカナのテキスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[str]: 音素記号のリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
if set(text).issubset(set(PUNCTUATIONS)):
|
||||||
|
return list(text)
|
||||||
|
# `text` がカタカナ(`ー`含む)のみからなるかどうかをチェック
|
||||||
|
if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None:
|
||||||
|
raise ValueError(f"Input must be katakana only: {text}")
|
||||||
|
sorted_keys = sorted(MORA_KATA_TO_MORA_PHONEMES.keys(), key=len, reverse=True)
|
||||||
|
pattern = "|".join(map(re.escape, sorted_keys))
|
||||||
|
|
||||||
|
def mora2phonemes(mora: str) -> str:
|
||||||
|
cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora]
|
||||||
|
if cosonant is None:
|
||||||
|
return f" {vowel}"
|
||||||
|
return f" {cosonant} {vowel}"
|
||||||
|
|
||||||
|
spaced_phonemes = re.sub(pattern, lambda m: mora2phonemes(m.group()), text)
|
||||||
|
|
||||||
|
# 長音記号「ー」の処理
|
||||||
|
long_pattern = r"(\w)(ー*)"
|
||||||
|
long_replacement = lambda m: m.group(1) + (" " + m.group(1)) * len(m.group(2)) # type: ignore
|
||||||
|
spaced_phonemes = re.sub(long_pattern, long_replacement, spaced_phonemes)
|
||||||
|
|
||||||
|
return spaced_phonemes.strip().split(" ")
|
||||||
|
|
||||||
|
|
||||||
|
def __align_tones(
|
||||||
|
phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]]
|
||||||
|
) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
例: …私は、、そう思う。
|
||||||
|
phones_with_punct:
|
||||||
|
[".", ".", ".", "w", "a", "t", "a", "sh", "i", "w", "a", ",", ",", "s", "o", "o", "o", "m", "o", "u", "."]
|
||||||
|
phone_tone_list:
|
||||||
|
[("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), ("_", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0))]
|
||||||
|
Return:
|
||||||
|
[(".", 0), (".", 0), (".", 0), ("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), (",", 0), (",", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0), (".", 0)]
|
||||||
|
|
||||||
|
Args:
|
||||||
|
phones_with_punct (list[str]): punctuation を含む音素のリスト
|
||||||
|
phone_tone_list (list[tuple[str, int]]): punctuation を含まない音素とアクセントのペアのリスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[tuple[str, int]]: punctuation を含む音素とアクセントのペアのリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
result: list[tuple[str, int]] = []
|
||||||
|
tone_index = 0
|
||||||
|
for phone in phones_with_punct:
|
||||||
|
if tone_index >= len(phone_tone_list):
|
||||||
|
# 余った punctuation がある場合 → (punctuation, 0) を追加
|
||||||
|
result.append((phone, 0))
|
||||||
|
elif phone == phone_tone_list[tone_index][0]:
|
||||||
|
# phone_tone_list の現在の音素と一致する場合 → tone をそこから取得、(phone, tone) を追加
|
||||||
|
result.append((phone, phone_tone_list[tone_index][1]))
|
||||||
|
# 探す index を1つ進める
|
||||||
|
tone_index += 1
|
||||||
|
elif phone in PUNCTUATIONS:
|
||||||
|
# phone が punctuation の場合 → (phone, 0) を追加
|
||||||
|
result.append((phone, 0))
|
||||||
|
else:
|
||||||
|
logger.debug(f"phones: {phones_with_punct}")
|
||||||
|
logger.debug(f"phone_tone_list: {phone_tone_list}")
|
||||||
|
logger.debug(f"result: {result}")
|
||||||
|
logger.debug(f"tone_index: {tone_index}")
|
||||||
|
logger.debug(f"phone: {phone}")
|
||||||
|
raise ValueError(f"Unexpected phone: {phone}")
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def __distribute_phone(n_phone: int, n_word: int) -> list[int]:
|
||||||
|
"""
|
||||||
|
左から右に 1 ずつ振り分け、次にまた左から右に1ずつ増やし、というふうに、
|
||||||
|
音素の数 `n_phone` を単語の数 `n_word` に分配する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_phone (int): 音素の数
|
||||||
|
n_word (int): 単語の数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list[int]: 単語ごとの音素の数のリスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
phones_per_word = [0] * n_word
|
||||||
|
for _ in range(n_phone):
|
||||||
|
min_tasks = min(phones_per_word)
|
||||||
|
min_index = phones_per_word.index(min_tasks)
|
||||||
|
phones_per_word[min_index] += 1
|
||||||
|
|
||||||
|
return phones_per_word
|
||||||
|
|
||||||
|
|
||||||
|
class YomiError(Exception):
|
||||||
|
"""
|
||||||
|
OpenJTalk で、読みが正しく取得できない箇所があるときに発生する例外。
|
||||||
|
基本的に「学習の前処理のテキスト処理時」には発生させ、そうでない場合は、
|
||||||
|
ignore_yomi_error=True にしておいて、この例外を発生させないようにする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
pass
|
||||||
92
style_bert_vits2/nlp/japanese/g2p_utils.py
Normal file
92
style_bert_vits2/nlp/japanese/g2p_utils.py
Normal file
@@ -0,0 +1,92 @@
|
|||||||
|
from style_bert_vits2.nlp.japanese.g2p import g2p
|
||||||
|
from style_bert_vits2.nlp.japanese.mora_list import (
|
||||||
|
MORA_KATA_TO_MORA_PHONEMES,
|
||||||
|
MORA_PHONEMES_TO_MORA_KATA,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||||
|
|
||||||
|
|
||||||
|
def g2kata_tone(norm_text: str) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
テキストからカタカナとアクセントのペアのリストを返す。
|
||||||
|
推論時のみに使われる関数のため、常に `raise_yomi_error=False` を指定して g2p() を呼ぶ仕様になっている。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
norm_text: 正規化されたテキスト。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
カタカナと音高のリスト。
|
||||||
|
"""
|
||||||
|
|
||||||
|
phones, tones, _ = g2p(norm_text, use_jp_extra=True, raise_yomi_error=False)
|
||||||
|
return phone_tone2kata_tone(list(zip(phones, tones)))
|
||||||
|
|
||||||
|
|
||||||
|
def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
phone_tone の phone 部分をカタカナに変換する。ただし最初と最後の ("_", 0) は無視する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
phone_tone: 音素と音高のリスト。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
カタカナと音高のリスト。
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 子音の集合
|
||||||
|
CONSONANTS = set(
|
||||||
|
[
|
||||||
|
consonant
|
||||||
|
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
|
||||||
|
if consonant is not None
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
phone_tone = phone_tone[1:] # 最初の("_", 0)を無視
|
||||||
|
phones = [phone for phone, _ in phone_tone]
|
||||||
|
tones = [tone for _, tone in phone_tone]
|
||||||
|
result: list[tuple[str, int]] = []
|
||||||
|
current_mora = ""
|
||||||
|
for phone, next_phone, tone, next_tone in zip(phones, phones[1:], tones, tones[1:]):
|
||||||
|
# zip の関係で最後の ("_", 0) は無視されている
|
||||||
|
if phone in PUNCTUATIONS:
|
||||||
|
result.append((phone, tone))
|
||||||
|
continue
|
||||||
|
if phone in CONSONANTS: # n以外の子音の場合
|
||||||
|
assert current_mora == "", f"Unexpected {phone} after {current_mora}"
|
||||||
|
assert tone == next_tone, f"Unexpected {phone} tone {tone} != {next_tone}"
|
||||||
|
current_mora = phone
|
||||||
|
else:
|
||||||
|
# phoneが母音もしくは「N」
|
||||||
|
current_mora += phone
|
||||||
|
result.append((MORA_PHONEMES_TO_MORA_KATA[current_mora], tone))
|
||||||
|
current_mora = ""
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def kata_tone2phone_tone(kata_tone: list[tuple[str, int]]) -> list[tuple[str, int]]:
|
||||||
|
"""
|
||||||
|
`phone_tone2kata_tone()` の逆の変換を行う。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
kata_tone: カタカナと音高のリスト。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
音素と音高のリスト。
|
||||||
|
"""
|
||||||
|
|
||||||
|
result: list[tuple[str, int]] = [("_", 0)]
|
||||||
|
for mora, tone in kata_tone:
|
||||||
|
if mora in PUNCTUATIONS:
|
||||||
|
result.append((mora, tone))
|
||||||
|
else:
|
||||||
|
consonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora]
|
||||||
|
if consonant is None:
|
||||||
|
result.append((vowel, tone))
|
||||||
|
else:
|
||||||
|
result.append((consonant, tone))
|
||||||
|
result.append((vowel, tone))
|
||||||
|
result.append(("_", 0))
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -1,10 +1,10 @@
|
|||||||
"""
|
"""
|
||||||
VOICEVOXのソースコードからお借りして最低限に改造したコード。
|
以下のコードは VOICEVOX のソースコードからお借りし最低限の改造を行ったもの。
|
||||||
https://github.com/VOICEVOX/voicevox_engine/blob/master/voicevox_engine/tts_pipeline/mora_list.py
|
https://github.com/VOICEVOX/voicevox_engine/blob/master/voicevox_engine/tts_pipeline/mora_list.py
|
||||||
"""
|
"""
|
||||||
|
|
||||||
"""
|
"""
|
||||||
以下のモーラ対応表はOpenJTalkのソースコードから取得し、
|
以下のモーラ対応表は OpenJTalk のソースコードから取得し、
|
||||||
カタカナ表記とモーラが一対一対応するように改造した。
|
カタカナ表記とモーラが一対一対応するように改造した。
|
||||||
ライセンス表記:
|
ライセンス表記:
|
||||||
-----------------------------------------------------------------
|
-----------------------------------------------------------------
|
||||||
@@ -46,13 +46,15 @@ OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY
|
|||||||
OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
|
OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
|
||||||
POSSIBILITY OF SUCH DAMAGE.
|
POSSIBILITY OF SUCH DAMAGE.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
# (カタカナ, 子音, 母音)の順。子音がない場合はNoneを入れる。
|
|
||||||
|
# (カタカナ, 子音, 母音)の順。子音がない場合は None を入れる。
|
||||||
# 但し「ン」と「ッ」は母音のみという扱いで、「ン」は「N」、「ッ」は「q」とする。
|
# 但し「ン」と「ッ」は母音のみという扱いで、「ン」は「N」、「ッ」は「q」とする。
|
||||||
# (元々「ッ」は「cl」)
|
# (元々「ッ」は「cl」)
|
||||||
# また「デェ = dy e」はpyopenjtalkの出力(de e)と合わないため削除
|
# また「デェ = dy e」は pyopenjtalk の出力(de e)と合わないため削除
|
||||||
_mora_list_minimum: list[tuple[str, Optional[str], str]] = [
|
__MORA_LIST_MINIMUM: list[tuple[str, Optional[str], str]] = [
|
||||||
("ヴォ", "v", "o"),
|
("ヴォ", "v", "o"),
|
||||||
("ヴェ", "v", "e"),
|
("ヴェ", "v", "e"),
|
||||||
("ヴィ", "v", "i"),
|
("ヴィ", "v", "i"),
|
||||||
@@ -199,7 +201,7 @@ _mora_list_minimum: list[tuple[str, Optional[str], str]] = [
|
|||||||
("イ", None, "i"),
|
("イ", None, "i"),
|
||||||
("ア", None, "a"),
|
("ア", None, "a"),
|
||||||
]
|
]
|
||||||
_mora_list_additional: list[tuple[str, Optional[str], str]] = [
|
__MORA_LIST_ADDITIONAL: list[tuple[str, Optional[str], str]] = [
|
||||||
("ヴョ", "by", "o"),
|
("ヴョ", "by", "o"),
|
||||||
("ヴュ", "by", "u"),
|
("ヴュ", "by", "u"),
|
||||||
("ヴャ", "by", "a"),
|
("ヴャ", "by", "a"),
|
||||||
@@ -220,13 +222,15 @@ _mora_list_additional: list[tuple[str, Optional[str], str]] = [
|
|||||||
("ァ", None, "a"),
|
("ァ", None, "a"),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# モーラの音素表記とカタカナの対応表
|
||||||
# 例: "vo" -> "ヴォ", "a" -> "ア"
|
# 例: "vo" -> "ヴォ", "a" -> "ア"
|
||||||
mora_phonemes_to_mora_kata: dict[str, str] = {
|
MORA_PHONEMES_TO_MORA_KATA: dict[str, str] = {
|
||||||
(consonant or "") + vowel: kana for [kana, consonant, vowel] in _mora_list_minimum
|
(consonant or "") + vowel: kana for [kana, consonant, vowel] in __MORA_LIST_MINIMUM
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# モーラのカタカナ表記と音素の対応表
|
||||||
# 例: "ヴォ" -> ("v", "o"), "ア" -> (None, "a")
|
# 例: "ヴォ" -> ("v", "o"), "ア" -> (None, "a")
|
||||||
mora_kata_to_mora_phonemes: dict[str, tuple[Optional[str], str]] = {
|
MORA_KATA_TO_MORA_PHONEMES: dict[str, tuple[Optional[str], str]] = {
|
||||||
kana: (consonant, vowel)
|
kana: (consonant, vowel)
|
||||||
for [kana, consonant, vowel] in _mora_list_minimum + _mora_list_additional
|
for [kana, consonant, vowel] in __MORA_LIST_MINIMUM + __MORA_LIST_ADDITIONAL
|
||||||
}
|
}
|
||||||
163
style_bert_vits2/nlp/japanese/normalizer.py
Normal file
163
style_bert_vits2/nlp/japanese/normalizer.py
Normal file
@@ -0,0 +1,163 @@
|
|||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
|
||||||
|
from num2words import num2words
|
||||||
|
|
||||||
|
from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
"""
|
||||||
|
日本語のテキストを正規化する。
|
||||||
|
結果は、ちょうど次の文字のみからなる:
|
||||||
|
- ひらがな
|
||||||
|
- カタカナ(全角長音記号「ー」が入る!)
|
||||||
|
- 漢字
|
||||||
|
- 半角アルファベット(大文字と小文字)
|
||||||
|
- ギリシャ文字
|
||||||
|
- `.` (句点`。`や`…`の一部や改行等)
|
||||||
|
- `,` (読点`、`や`:`等)
|
||||||
|
- `?` (疑問符`?`)
|
||||||
|
- `!` (感嘆符`!`)
|
||||||
|
- `'` (`「`や`」`等)
|
||||||
|
- `-` (`―`(ダッシュ、長音記号ではない)や`-`等)
|
||||||
|
|
||||||
|
注意点:
|
||||||
|
- 三点リーダー`…`は`...`に変換される(`なるほど…。` → `なるほど....`)
|
||||||
|
- 数字は漢字に変換される(`1,100円` → `千百円`、`52.34` → `五十二点三四`)
|
||||||
|
- 読点や疑問符等の位置・個数等は保持される(`??あ、、!!!` → `??あ,,!!!`)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 正規化するテキスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: 正規化されたテキスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
res = unicodedata.normalize("NFKC", text) # ここでアルファベットは半角になる
|
||||||
|
res = __convert_numbers_to_words(res) # 「100円」→「百円」等
|
||||||
|
# 「~」と「〜」と「~」も長音記号として扱う
|
||||||
|
res = res.replace("~", "ー")
|
||||||
|
res = res.replace("~", "ー")
|
||||||
|
res = res.replace("〜", "ー")
|
||||||
|
|
||||||
|
res = replace_punctuation(res) # 句読点等正規化、読めない文字を削除
|
||||||
|
|
||||||
|
# 結合文字の濁点・半濁点を削除
|
||||||
|
# 通常の「ば」等はそのままのこされる、「あ゛」は上で「あ゙」になりここで「あ」になる
|
||||||
|
res = res.replace("\u3099", "") # 結合文字の濁点を削除、る゙ → る
|
||||||
|
res = res.replace("\u309A", "") # 結合文字の半濁点を削除、な゚ → な
|
||||||
|
return res
|
||||||
|
|
||||||
|
|
||||||
|
def replace_punctuation(text: str) -> str:
|
||||||
|
"""
|
||||||
|
句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す:
|
||||||
|
漢字・平仮名・カタカナ、アルファベット、ギリシャ文字
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 正規化するテキスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: 正規化されたテキスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 記号類の正規化変換マップ
|
||||||
|
REPLACE_MAP = {
|
||||||
|
":": ",",
|
||||||
|
";": ",",
|
||||||
|
",": ",",
|
||||||
|
"。": ".",
|
||||||
|
"!": "!",
|
||||||
|
"?": "?",
|
||||||
|
"\n": ".",
|
||||||
|
".": ".",
|
||||||
|
"…": "...",
|
||||||
|
"···": "...",
|
||||||
|
"・・・": "...",
|
||||||
|
"·": ",",
|
||||||
|
"・": ",",
|
||||||
|
"、": ",",
|
||||||
|
"$": ".",
|
||||||
|
"“": "'",
|
||||||
|
"”": "'",
|
||||||
|
'"': "'",
|
||||||
|
"‘": "'",
|
||||||
|
"’": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"(": "'",
|
||||||
|
")": "'",
|
||||||
|
"《": "'",
|
||||||
|
"》": "'",
|
||||||
|
"【": "'",
|
||||||
|
"】": "'",
|
||||||
|
"[": "'",
|
||||||
|
"]": "'",
|
||||||
|
# NFKC 正規化後のハイフン・ダッシュの変種を全て通常半角ハイフン - \u002d に変換
|
||||||
|
"\u02d7": "\u002d", # ˗, Modifier Letter Minus Sign
|
||||||
|
"\u2010": "\u002d", # ‐, Hyphen,
|
||||||
|
# "\u2011": "\u002d", # ‑, Non-Breaking Hyphen, NFKC により \u2010 に変換される
|
||||||
|
"\u2012": "\u002d", # ‒, Figure Dash
|
||||||
|
"\u2013": "\u002d", # –, En Dash
|
||||||
|
"\u2014": "\u002d", # —, Em Dash
|
||||||
|
"\u2015": "\u002d", # ―, Horizontal Bar
|
||||||
|
"\u2043": "\u002d", # ⁃, Hyphen Bullet
|
||||||
|
"\u2212": "\u002d", # −, Minus Sign
|
||||||
|
"\u23af": "\u002d", # ⎯, Horizontal Line Extension
|
||||||
|
"\u23e4": "\u002d", # ⏤, Straightness
|
||||||
|
"\u2500": "\u002d", # ─, Box Drawings Light Horizontal
|
||||||
|
"\u2501": "\u002d", # ━, Box Drawings Heavy Horizontal
|
||||||
|
"\u2e3a": "\u002d", # ⸺, Two-Em Dash
|
||||||
|
"\u2e3b": "\u002d", # ⸻, Three-Em Dash
|
||||||
|
# "~": "-", # これは長音記号「ー」として扱うよう変更
|
||||||
|
# "~": "-", # これも長音記号「ー」として扱うよう変更
|
||||||
|
"「": "'",
|
||||||
|
"」": "'",
|
||||||
|
}
|
||||||
|
|
||||||
|
pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys()))
|
||||||
|
|
||||||
|
# 句読点を辞書で置換
|
||||||
|
replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text)
|
||||||
|
|
||||||
|
replaced_text = re.sub(
|
||||||
|
# ↓ ひらがな、カタカナ、漢字
|
||||||
|
r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005"
|
||||||
|
# ↓ 半角アルファベット(大文字と小文字)
|
||||||
|
+ r"\u0041-\u005A\u0061-\u007A"
|
||||||
|
# ↓ 全角アルファベット(大文字と小文字)
|
||||||
|
+ r"\uFF21-\uFF3A\uFF41-\uFF5A"
|
||||||
|
# ↓ ギリシャ文字
|
||||||
|
+ r"\u0370-\u03FF\u1F00-\u1FFF"
|
||||||
|
# ↓ "!", "?", "…", ",", ".", "'", "-", 但し`…`はすでに`...`に変換されている
|
||||||
|
+ "".join(PUNCTUATIONS) + r"]+",
|
||||||
|
# 上述以外の文字を削除
|
||||||
|
"",
|
||||||
|
replaced_text,
|
||||||
|
)
|
||||||
|
|
||||||
|
return replaced_text
|
||||||
|
|
||||||
|
|
||||||
|
def __convert_numbers_to_words(text: str) -> str:
|
||||||
|
"""
|
||||||
|
記号や数字を日本語の文字表現に変換する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 変換するテキスト
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: 変換されたテキスト
|
||||||
|
"""
|
||||||
|
|
||||||
|
NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+")
|
||||||
|
CURRENCY_MAP = {"$": "ドル", "¥": "円", "£": "ポンド", "€": "ユーロ"}
|
||||||
|
CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])")
|
||||||
|
NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?")
|
||||||
|
|
||||||
|
res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text)
|
||||||
|
res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res)
|
||||||
|
res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res)
|
||||||
|
|
||||||
|
return res
|
||||||
160
style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py
Normal file
160
style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py
Normal file
@@ -0,0 +1,160 @@
|
|||||||
|
"""
|
||||||
|
Run the pyopenjtalk worker in a separate process
|
||||||
|
to avoid user dictionary access error
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_client import WorkerClient
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT
|
||||||
|
|
||||||
|
|
||||||
|
WORKER_CLIENT: Optional[WorkerClient] = None
|
||||||
|
|
||||||
|
|
||||||
|
# pyopenjtalk interface
|
||||||
|
# g2p(): not used
|
||||||
|
|
||||||
|
|
||||||
|
def run_frontend(text: str) -> list[dict[str, Any]]:
|
||||||
|
if WORKER_CLIENT is not None:
|
||||||
|
ret = WORKER_CLIENT.dispatch_pyopenjtalk("run_frontend", text)
|
||||||
|
assert isinstance(ret, list)
|
||||||
|
return ret
|
||||||
|
else:
|
||||||
|
# without worker
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
return pyopenjtalk.run_frontend(text)
|
||||||
|
|
||||||
|
|
||||||
|
def make_label(njd_features: Any) -> list[str]:
|
||||||
|
if WORKER_CLIENT is not None:
|
||||||
|
ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features)
|
||||||
|
assert isinstance(ret, list)
|
||||||
|
return ret
|
||||||
|
else:
|
||||||
|
# without worker
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
return pyopenjtalk.make_label(njd_features)
|
||||||
|
|
||||||
|
|
||||||
|
def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None) -> None:
|
||||||
|
if WORKER_CLIENT is not None:
|
||||||
|
WORKER_CLIENT.dispatch_pyopenjtalk("mecab_dict_index", path, out_path, dn_mecab)
|
||||||
|
else:
|
||||||
|
# without worker
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab)
|
||||||
|
|
||||||
|
|
||||||
|
def update_global_jtalk_with_user_dict(path: str) -> None:
|
||||||
|
if WORKER_CLIENT is not None:
|
||||||
|
WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path)
|
||||||
|
else:
|
||||||
|
# without worker
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
pyopenjtalk.update_global_jtalk_with_user_dict(path)
|
||||||
|
|
||||||
|
|
||||||
|
def unset_user_dict() -> None:
|
||||||
|
if WORKER_CLIENT is not None:
|
||||||
|
WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict")
|
||||||
|
else:
|
||||||
|
# without worker
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
pyopenjtalk.unset_user_dict()
|
||||||
|
|
||||||
|
|
||||||
|
# initialize module when imported
|
||||||
|
|
||||||
|
|
||||||
|
def initialize_worker(port: int = WORKER_PORT) -> None:
|
||||||
|
import atexit
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
|
||||||
|
global WORKER_CLIENT
|
||||||
|
if WORKER_CLIENT:
|
||||||
|
return
|
||||||
|
|
||||||
|
client = None
|
||||||
|
try:
|
||||||
|
client = WorkerClient(port)
|
||||||
|
except (socket.timeout, socket.error):
|
||||||
|
logger.debug("try starting pyopenjtalk worker server")
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
worker_pkg_path = os.path.relpath(
|
||||||
|
os.path.dirname(__file__), os.getcwd()
|
||||||
|
).replace(os.sep, ".")
|
||||||
|
args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)]
|
||||||
|
# new session, new process group
|
||||||
|
if sys.platform.startswith("win"):
|
||||||
|
cf = subprocess.CREATE_NEW_CONSOLE | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore
|
||||||
|
si = subprocess.STARTUPINFO() # type: ignore
|
||||||
|
si.dwFlags |= subprocess.STARTF_USESHOWWINDOW # type: ignore
|
||||||
|
si.wShowWindow = subprocess.SW_HIDE # type: ignore
|
||||||
|
subprocess.Popen(args, creationflags=cf, startupinfo=si)
|
||||||
|
else:
|
||||||
|
# align with Windows behavior
|
||||||
|
# start_new_session is same as specifying setsid in preexec_fn
|
||||||
|
subprocess.Popen(
|
||||||
|
args,
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# wait until server listening
|
||||||
|
count = 0
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
client = WorkerClient(port)
|
||||||
|
break
|
||||||
|
except socket.error:
|
||||||
|
time.sleep(0.5)
|
||||||
|
count += 1
|
||||||
|
# 20: max number of retries
|
||||||
|
if count == 20:
|
||||||
|
raise TimeoutError("サーバーに接続できませんでした")
|
||||||
|
|
||||||
|
logger.debug("pyopenjtalk worker server started")
|
||||||
|
WORKER_CLIENT = client
|
||||||
|
atexit.register(terminate_worker)
|
||||||
|
|
||||||
|
# when the process is killed
|
||||||
|
def signal_handler(signum: int, frame: Any):
|
||||||
|
terminate_worker()
|
||||||
|
|
||||||
|
try:
|
||||||
|
signal.signal(signal.SIGTERM, signal_handler)
|
||||||
|
except ValueError:
|
||||||
|
# signal only works in main thread
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# top-level declaration
|
||||||
|
def terminate_worker() -> None:
|
||||||
|
logger.debug("pyopenjtalk worker server terminated")
|
||||||
|
global WORKER_CLIENT
|
||||||
|
if not WORKER_CLIENT:
|
||||||
|
return
|
||||||
|
|
||||||
|
# prepare for unexpected errors
|
||||||
|
try:
|
||||||
|
if WORKER_CLIENT.status() == 1:
|
||||||
|
WORKER_CLIENT.quit_server()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(e)
|
||||||
|
|
||||||
|
WORKER_CLIENT.close()
|
||||||
|
WORKER_CLIENT = None
|
||||||
16
style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__main__.py
Normal file
16
style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__main__.py
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
import argparse
|
||||||
|
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_server import WorkerServer
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--port", type=int, default=WORKER_PORT)
|
||||||
|
args = parser.parse_args()
|
||||||
|
server = WorkerServer()
|
||||||
|
server.start_server(port=args.port)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
import socket
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import (
|
||||||
|
RequestType,
|
||||||
|
receive_data,
|
||||||
|
send_data,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerClient:
|
||||||
|
"""pyopenjtalk worker client"""
|
||||||
|
|
||||||
|
def __init__(self, port: int) -> None:
|
||||||
|
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||||
|
# timeout: 60 seconds
|
||||||
|
sock.settimeout(60)
|
||||||
|
sock.connect((socket.gethostname(), port))
|
||||||
|
self.sock = sock
|
||||||
|
|
||||||
|
def __enter__(self) -> "WorkerClient":
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
|
||||||
|
self.close()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
self.sock.close()
|
||||||
|
|
||||||
|
def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any:
|
||||||
|
data = {
|
||||||
|
"request-type": RequestType.PYOPENJTALK,
|
||||||
|
"func": func,
|
||||||
|
"args": args,
|
||||||
|
"kwargs": kwargs,
|
||||||
|
}
|
||||||
|
logger.trace(f"client sends request: {data}")
|
||||||
|
send_data(self.sock, data)
|
||||||
|
logger.trace("client sent request successfully")
|
||||||
|
response = receive_data(self.sock)
|
||||||
|
logger.trace(f"client received response: {response}")
|
||||||
|
return response.get("return")
|
||||||
|
|
||||||
|
def status(self) -> int:
|
||||||
|
data = {"request-type": RequestType.STATUS}
|
||||||
|
logger.trace(f"client sends request: {data}")
|
||||||
|
send_data(self.sock, data)
|
||||||
|
logger.trace("client sent request successfully")
|
||||||
|
response = receive_data(self.sock)
|
||||||
|
logger.trace(f"client received response: {response}")
|
||||||
|
return cast(int, response.get("client-count"))
|
||||||
|
|
||||||
|
def quit_server(self) -> None:
|
||||||
|
data = {"request-type": RequestType.QUIT_SERVER}
|
||||||
|
logger.trace(f"client sends request: {data}")
|
||||||
|
send_data(self.sock, data)
|
||||||
|
logger.trace("client sent request successfully")
|
||||||
|
response = receive_data(self.sock)
|
||||||
|
logger.trace(f"client received response: {response}")
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
import json
|
||||||
|
import socket
|
||||||
|
from enum import IntEnum, auto
|
||||||
|
from typing import Any, Final
|
||||||
|
|
||||||
|
|
||||||
|
WORKER_PORT: Final[int] = 7861
|
||||||
|
HEADER_SIZE: Final[int] = 4
|
||||||
|
|
||||||
|
|
||||||
|
class RequestType(IntEnum):
|
||||||
|
STATUS = auto()
|
||||||
|
QUIT_SERVER = auto()
|
||||||
|
PYOPENJTALK = auto()
|
||||||
|
|
||||||
|
|
||||||
|
class ConnectionClosedException(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# socket communication
|
||||||
|
|
||||||
|
|
||||||
|
def send_data(sock: socket.socket, data: dict[str, Any]):
|
||||||
|
json_data = json.dumps(data).encode()
|
||||||
|
header = len(json_data).to_bytes(HEADER_SIZE, byteorder="big")
|
||||||
|
sock.sendall(header + json_data)
|
||||||
|
|
||||||
|
|
||||||
|
def __receive_until(sock: socket.socket, size: int):
|
||||||
|
data = b""
|
||||||
|
while len(data) < size:
|
||||||
|
part = sock.recv(size - len(data))
|
||||||
|
if part == b"":
|
||||||
|
raise ConnectionClosedException("接続が閉じられました")
|
||||||
|
data += part
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def receive_data(sock: socket.socket) -> dict[str, Any]:
|
||||||
|
header = __receive_until(sock, HEADER_SIZE)
|
||||||
|
data_length = int.from_bytes(header, byteorder="big")
|
||||||
|
body = __receive_until(sock, data_length)
|
||||||
|
return json.loads(body.decode())
|
||||||
@@ -0,0 +1,123 @@
|
|||||||
|
import select
|
||||||
|
import socket
|
||||||
|
import time
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
import pyopenjtalk
|
||||||
|
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import (
|
||||||
|
ConnectionClosedException,
|
||||||
|
RequestType,
|
||||||
|
receive_data,
|
||||||
|
send_data,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# To make it as fast as possible
|
||||||
|
# Probably faster than calling getattr every time
|
||||||
|
PYOPENJTALK_FUNC_DICT = {
|
||||||
|
"run_frontend": pyopenjtalk.run_frontend,
|
||||||
|
"make_label": pyopenjtalk.make_label,
|
||||||
|
"mecab_dict_index": pyopenjtalk.mecab_dict_index,
|
||||||
|
"update_global_jtalk_with_user_dict": pyopenjtalk.update_global_jtalk_with_user_dict,
|
||||||
|
"unset_user_dict": pyopenjtalk.unset_user_dict,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerServer:
|
||||||
|
"""pyopenjtalk worker server"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.client_count: int = 0
|
||||||
|
self.quit: bool = False
|
||||||
|
|
||||||
|
def handle_request(self, request: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
request_type = None
|
||||||
|
try:
|
||||||
|
request_type = RequestType(cast(int, request.get("request-type")))
|
||||||
|
except Exception:
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"reason": "request-type is invalid",
|
||||||
|
}
|
||||||
|
|
||||||
|
response: dict[str, Any] = {}
|
||||||
|
if request_type:
|
||||||
|
if request_type == RequestType.STATUS:
|
||||||
|
response = {
|
||||||
|
"success": True,
|
||||||
|
"client-count": self.client_count,
|
||||||
|
}
|
||||||
|
elif request_type == RequestType.QUIT_SERVER:
|
||||||
|
self.quit = True
|
||||||
|
response = {"success": True}
|
||||||
|
elif request_type == RequestType.PYOPENJTALK:
|
||||||
|
func_name = request.get("func")
|
||||||
|
assert isinstance(func_name, str)
|
||||||
|
func = PYOPENJTALK_FUNC_DICT[func_name]
|
||||||
|
args = request.get("args")
|
||||||
|
kwargs = request.get("kwargs")
|
||||||
|
assert isinstance(args, list)
|
||||||
|
assert isinstance(kwargs, dict)
|
||||||
|
ret = func(*args, **kwargs)
|
||||||
|
response = {"success": True, "return": ret}
|
||||||
|
else:
|
||||||
|
# NOT REACHED
|
||||||
|
response = request
|
||||||
|
|
||||||
|
return response
|
||||||
|
|
||||||
|
def start_server(self, port: int, no_client_timeout: int = 30) -> None:
|
||||||
|
logger.info("start pyopenjtalk worker server")
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
|
||||||
|
server_socket.bind((socket.gethostname(), port))
|
||||||
|
server_socket.listen()
|
||||||
|
sockets = [server_socket]
|
||||||
|
no_client_since = time.time()
|
||||||
|
while True:
|
||||||
|
if self.client_count == 0:
|
||||||
|
if no_client_since is None:
|
||||||
|
no_client_since = time.time()
|
||||||
|
elif (time.time() - no_client_since) > no_client_timeout:
|
||||||
|
logger.info("quit because there is no client")
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
no_client_since = None
|
||||||
|
|
||||||
|
ready_sockets, _, _ = select.select(sockets, [], [], 0.1)
|
||||||
|
for sock in ready_sockets:
|
||||||
|
if sock is server_socket:
|
||||||
|
logger.info("new client connected")
|
||||||
|
client_socket, _ = server_socket.accept()
|
||||||
|
sockets.append(client_socket)
|
||||||
|
self.client_count += 1
|
||||||
|
else:
|
||||||
|
# client
|
||||||
|
try:
|
||||||
|
request = receive_data(sock)
|
||||||
|
except Exception as e:
|
||||||
|
sock.close()
|
||||||
|
sockets.remove(sock)
|
||||||
|
self.client_count -= 1
|
||||||
|
# unexpected disconnections
|
||||||
|
if not isinstance(e, ConnectionClosedException):
|
||||||
|
logger.error(e)
|
||||||
|
|
||||||
|
logger.info("close connection")
|
||||||
|
continue
|
||||||
|
|
||||||
|
logger.trace(f"server received request: {request}")
|
||||||
|
|
||||||
|
response = self.handle_request(request)
|
||||||
|
logger.trace(f"server sends response: {response}")
|
||||||
|
try:
|
||||||
|
send_data(sock, response)
|
||||||
|
logger.trace("server sent response successfully")
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"an exception occurred during sending responce"
|
||||||
|
)
|
||||||
|
if self.quit:
|
||||||
|
logger.info("quit pyopenjtalk worker server")
|
||||||
|
return
|
||||||
27
style_bert_vits2/nlp/japanese/user_dict/README.md
Normal file
27
style_bert_vits2/nlp/japanese/user_dict/README.md
Normal file
@@ -0,0 +1,27 @@
|
|||||||
|
|
||||||
|
## ユーザー辞書関連のコードについて
|
||||||
|
|
||||||
|
このフォルダに含まれるユーザー辞書関連のコードは、[VOICEVOX ENGINE](https://github.com/VOICEVOX/voicevox_engine) プロジェクトのコードを改変したものを使用しています。
|
||||||
|
VOICEVOX プロジェクトのチームに深く感謝し、その貢献を尊重します。
|
||||||
|
|
||||||
|
### 元のコード
|
||||||
|
|
||||||
|
- [voicevox_engine/user_dict/](https://github.com/VOICEVOX/voicevox_engine/tree/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict)
|
||||||
|
- [voicevox_engine/model.py](https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207)
|
||||||
|
|
||||||
|
### 改変の詳細
|
||||||
|
|
||||||
|
- ファイル名の書き換えおよびそれに伴う import 文の書き換え。
|
||||||
|
- VOICEVOX 固有の部分をコメントアウト。
|
||||||
|
- mutex を使用している部分をコメントアウト。
|
||||||
|
- 参照している pyopenjtalk の違いによるメソッド名の書き換え。
|
||||||
|
- UserDictWord の mora_count のデフォルト値を None に指定。
|
||||||
|
- `model.py` のうち、必要な Pydantic モデルのみを抽出。
|
||||||
|
|
||||||
|
### ライセンス
|
||||||
|
|
||||||
|
元の VOICEVOX ENGINE のリポジトリのコードは、LGPL v3 と、ソースコードの公開が不要な別ライセンスのデュアルライセンスの下で使用されています。
|
||||||
|
当プロジェクトにおけるこのモジュールも LGPL ライセンスの下にあります。
|
||||||
|
|
||||||
|
詳細については、プロジェクトのルートディレクトリにある [LGPL_LICENSE](/LGPL_LICENSE) ファイルをご参照ください。
|
||||||
|
また、元の VOICEVOX ENGINE プロジェクトのライセンスについては、[こちら](https://github.com/VOICEVOX/voicevox_engine/blob/master/LICENSE) をご覧ください。
|
||||||
@@ -1,39 +1,43 @@
|
|||||||
# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。
|
"""
|
||||||
# 引用元:
|
このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。
|
||||||
# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/user_dict.py
|
引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/user_dict.py
|
||||||
# ライセンス: LGPL-3.0
|
ライセンス: LGPL-3.0
|
||||||
# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。
|
詳しくは、このファイルと同じフォルダにある README.md を参照してください。
|
||||||
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
import threading
|
|
||||||
import traceback
|
import traceback
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional
|
from typing import Dict, List, Optional
|
||||||
from uuid import UUID, uuid4
|
from uuid import UUID, uuid4
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pyopenjtalk
|
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException
|
||||||
|
|
||||||
from .word_model import UserDictWord, WordTypes
|
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR
|
||||||
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
|
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import (
|
||||||
|
MAX_PRIORITY,
|
||||||
|
MIN_PRIORITY,
|
||||||
|
part_of_speech_data,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes
|
||||||
|
|
||||||
# from ..utility.mutex_utility import mutex_wrapper
|
|
||||||
# from ..utility.path_utility import engine_root, get_save_dir
|
|
||||||
from .part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data
|
|
||||||
from common.constants import USER_DICT_DIR
|
|
||||||
|
|
||||||
# root_dir = engine_root()
|
# root_dir = engine_root()
|
||||||
# save_dir = get_save_dir()
|
# save_dir = get_save_dir()
|
||||||
root_dir = Path(USER_DICT_DIR)
|
|
||||||
save_dir = Path(USER_DICT_DIR)
|
|
||||||
|
|
||||||
|
# if not save_dir.is_dir():
|
||||||
|
# save_dir.mkdir(parents=True)
|
||||||
|
|
||||||
if not save_dir.is_dir():
|
default_dict_path = (
|
||||||
save_dir.mkdir(parents=True)
|
DEFAULT_USER_DICT_DIR / "default.csv"
|
||||||
|
) # VOICEVOXデフォルト辞書ファイルのパス
|
||||||
default_dict_path = root_dir / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス
|
user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス
|
||||||
user_dict_path = save_dir / "user_dict.json" # ユーザー辞書ファイルのパス
|
compiled_dict_path = (
|
||||||
compiled_dict_path = save_dir / "user.dic" # コンパイル済み辞書ファイルのパス
|
DEFAULT_USER_DICT_DIR / "user.dic"
|
||||||
|
) # コンパイル済み辞書ファイルのパス
|
||||||
|
|
||||||
|
|
||||||
# # 同時書き込みの制御
|
# # 同時書き込みの制御
|
||||||
@@ -54,7 +58,7 @@ def _write_to_json(user_dict: Dict[str, UserDictWord], user_dict_path: Path) ->
|
|||||||
"""
|
"""
|
||||||
converted_user_dict = {}
|
converted_user_dict = {}
|
||||||
for word_uuid, word in user_dict.items():
|
for word_uuid, word in user_dict.items():
|
||||||
word_dict = word.dict()
|
word_dict = word.model_dump()
|
||||||
word_dict["cost"] = _priority2cost(
|
word_dict["cost"] = _priority2cost(
|
||||||
word_dict["context_id"], word_dict["priority"]
|
word_dict["context_id"], word_dict["priority"]
|
||||||
)
|
)
|
||||||
@@ -84,6 +88,7 @@ def update_dict(
|
|||||||
compiled_dict_path : Path
|
compiled_dict_path : Path
|
||||||
コンパイル済み辞書ファイルのパス
|
コンパイル済み辞書ファイルのパス
|
||||||
"""
|
"""
|
||||||
|
|
||||||
random_string = uuid4()
|
random_string = uuid4()
|
||||||
tmp_csv_path = compiled_dict_path.with_suffix(
|
tmp_csv_path = compiled_dict_path.with_suffix(
|
||||||
f".dict_csv-{random_string}.tmp"
|
f".dict_csv-{random_string}.tmp"
|
||||||
@@ -1,18 +1,20 @@
|
|||||||
# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。
|
"""
|
||||||
# 引用元:
|
このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。
|
||||||
# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/part_of_speech_data.py
|
引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/part_of_speech_data.py
|
||||||
# ライセンス: LGPL-3.0
|
ライセンス: LGPL-3.0
|
||||||
# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。
|
詳しくは、このファイルと同じフォルダにある README.md を参照してください。
|
||||||
|
"""
|
||||||
|
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
|
||||||
from .word_model import (
|
from style_bert_vits2.nlp.japanese.user_dict.word_model import (
|
||||||
USER_DICT_MAX_PRIORITY,
|
USER_DICT_MAX_PRIORITY,
|
||||||
USER_DICT_MIN_PRIORITY,
|
USER_DICT_MIN_PRIORITY,
|
||||||
PartOfSpeechDetail,
|
PartOfSpeechDetail,
|
||||||
WordTypes,
|
WordTypes,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
MIN_PRIORITY = USER_DICT_MIN_PRIORITY
|
MIN_PRIORITY = USER_DICT_MIN_PRIORITY
|
||||||
MAX_PRIORITY = USER_DICT_MAX_PRIORITY
|
MAX_PRIORITY = USER_DICT_MAX_PRIORITY
|
||||||
|
|
||||||
@@ -1,14 +1,17 @@
|
|||||||
# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。
|
"""
|
||||||
# 引用元:
|
このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。
|
||||||
# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207
|
引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207
|
||||||
# ライセンス: LGPL-3.0
|
ライセンス: LGPL-3.0
|
||||||
# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。
|
詳しくは、このファイルと同じフォルダにある README.md を参照してください。
|
||||||
|
"""
|
||||||
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from re import findall, fullmatch
|
from re import findall, fullmatch
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
from pydantic import BaseModel, Field, validator
|
from pydantic import BaseModel, Field, validator
|
||||||
|
|
||||||
|
|
||||||
USER_DICT_MIN_PRIORITY = 0
|
USER_DICT_MIN_PRIORITY = 0
|
||||||
USER_DICT_MAX_PRIORITY = 10
|
USER_DICT_MAX_PRIORITY = 10
|
||||||
|
|
||||||
@@ -1,9 +1,14 @@
|
|||||||
punctuation = ["!", "?", "…", ",", ".", "'", "-"]
|
# Punctuations
|
||||||
pu_symbols = punctuation + ["SP", "UNK"]
|
PUNCTUATIONS = ["!", "?", "…", ",", ".", "'", "-"]
|
||||||
pad = "_"
|
|
||||||
|
|
||||||
# chinese
|
# Punctuations and special tokens
|
||||||
zh_symbols = [
|
PUNCTUATION_SYMBOLS = PUNCTUATIONS + ["SP", "UNK"]
|
||||||
|
|
||||||
|
# Padding
|
||||||
|
PAD = "_"
|
||||||
|
|
||||||
|
# Chinese symbols
|
||||||
|
ZH_SYMBOLS = [
|
||||||
"E",
|
"E",
|
||||||
"En",
|
"En",
|
||||||
"a",
|
"a",
|
||||||
@@ -70,10 +75,10 @@ zh_symbols = [
|
|||||||
"EE",
|
"EE",
|
||||||
"OO",
|
"OO",
|
||||||
]
|
]
|
||||||
num_zh_tones = 6
|
NUM_ZH_TONES = 6
|
||||||
|
|
||||||
# japanese
|
# Japanese
|
||||||
ja_symbols = [
|
JP_SYMBOLS = [
|
||||||
"N",
|
"N",
|
||||||
"a",
|
"a",
|
||||||
"a:",
|
"a:",
|
||||||
@@ -117,10 +122,10 @@ ja_symbols = [
|
|||||||
"z",
|
"z",
|
||||||
"zy",
|
"zy",
|
||||||
]
|
]
|
||||||
num_ja_tones = 2
|
NUM_JP_TONES = 2
|
||||||
|
|
||||||
# English
|
# English
|
||||||
en_symbols = [
|
EN_SYMBOLS = [
|
||||||
"aa",
|
"aa",
|
||||||
"ae",
|
"ae",
|
||||||
"ah",
|
"ah",
|
||||||
@@ -161,27 +166,29 @@ en_symbols = [
|
|||||||
"z",
|
"z",
|
||||||
"zh",
|
"zh",
|
||||||
]
|
]
|
||||||
num_en_tones = 4
|
NUM_EN_TONES = 4
|
||||||
|
|
||||||
# combine all symbols
|
# Combine all symbols
|
||||||
normal_symbols = sorted(set(zh_symbols + ja_symbols + en_symbols))
|
NORMAL_SYMBOLS = sorted(set(ZH_SYMBOLS + JP_SYMBOLS + EN_SYMBOLS))
|
||||||
symbols = [pad] + normal_symbols + pu_symbols
|
SYMBOLS = [PAD] + NORMAL_SYMBOLS + PUNCTUATION_SYMBOLS
|
||||||
sil_phonemes_ids = [symbols.index(i) for i in pu_symbols]
|
SIL_PHONEMES_IDS = [SYMBOLS.index(i) for i in PUNCTUATION_SYMBOLS]
|
||||||
|
|
||||||
# combine all tones
|
# Combine all tones
|
||||||
num_tones = num_zh_tones + num_ja_tones + num_en_tones
|
NUM_TONES = NUM_ZH_TONES + NUM_JP_TONES + NUM_EN_TONES
|
||||||
|
|
||||||
# language maps
|
# Language maps
|
||||||
language_id_map = {"ZH": 0, "JP": 1, "EN": 2}
|
LANGUAGE_ID_MAP = {"ZH": 0, "JP": 1, "EN": 2}
|
||||||
num_languages = len(language_id_map.keys())
|
NUM_LANGUAGES = len(LANGUAGE_ID_MAP.keys())
|
||||||
|
|
||||||
language_tone_start_map = {
|
# Language tone start map
|
||||||
|
LANGUAGE_TONE_START_MAP = {
|
||||||
"ZH": 0,
|
"ZH": 0,
|
||||||
"JP": num_zh_tones,
|
"JP": NUM_ZH_TONES,
|
||||||
"EN": num_zh_tones + num_ja_tones,
|
"EN": NUM_ZH_TONES + NUM_JP_TONES,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
a = set(zh_symbols)
|
a = set(ZH_SYMBOLS)
|
||||||
b = set(en_symbols)
|
b = set(EN_SYMBOLS)
|
||||||
print(sorted(a & b))
|
print(sorted(a & b))
|
||||||
428
style_bert_vits2/tts_model.py
Normal file
428
style_bert_vits2/tts_model.py
Normal file
@@ -0,0 +1,428 @@
|
|||||||
|
import warnings
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
|
import gradio as gr
|
||||||
|
import numpy as np
|
||||||
|
import pyannote.audio
|
||||||
|
import torch
|
||||||
|
from gradio.processing_utils import convert_to_16_bit_wav
|
||||||
|
from numpy.typing import NDArray
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import (
|
||||||
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
DEFAULT_LENGTH,
|
||||||
|
DEFAULT_LINE_SPLIT,
|
||||||
|
DEFAULT_NOISE,
|
||||||
|
DEFAULT_NOISEW,
|
||||||
|
DEFAULT_SDP_RATIO,
|
||||||
|
DEFAULT_SPLIT_INTERVAL,
|
||||||
|
DEFAULT_STYLE,
|
||||||
|
DEFAULT_STYLE_WEIGHT,
|
||||||
|
Languages,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
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.nlp import bert_models
|
||||||
|
from style_bert_vits2.voice import adjust_voice
|
||||||
|
|
||||||
|
|
||||||
|
class TTSModel:
|
||||||
|
"""
|
||||||
|
Style-Bert-Vits2 の音声合成モデルを操作するクラス。
|
||||||
|
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, model_path: Path, config_path: Path, style_vec_path: Path, device: str
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Style-Bert-Vits2 の音声合成モデルを初期化する。
|
||||||
|
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_path (Path): モデル (.safetensors) のパス
|
||||||
|
config_path (Path): ハイパーパラメータ (config.json) のパス
|
||||||
|
style_vec_path (Path): スタイルベクトル (style_vectors.npy) のパス
|
||||||
|
device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など)
|
||||||
|
"""
|
||||||
|
|
||||||
|
self.model_path: Path = model_path
|
||||||
|
self.config_path: Path = config_path
|
||||||
|
self.style_vec_path: Path = style_vec_path
|
||||||
|
self.device: str = device
|
||||||
|
self.hyper_parameters: HyperParameters = HyperParameters.load_from_json(
|
||||||
|
self.config_path
|
||||||
|
)
|
||||||
|
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
||||||
|
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
||||||
|
|
||||||
|
num_styles: int = self.hyper_parameters.data.num_styles
|
||||||
|
if hasattr(self.hyper_parameters.data, "style2id"):
|
||||||
|
self.style2id: dict[str, int] = self.hyper_parameters.data.style2id
|
||||||
|
else:
|
||||||
|
self.style2id: dict[str, int] = {str(i): i for i in range(num_styles)}
|
||||||
|
if len(self.style2id) != num_styles:
|
||||||
|
raise ValueError(
|
||||||
|
f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.__style_vector_inference: Optional[pyannote.audio.Inference] = None
|
||||||
|
self.__style_vectors: NDArray[Any] = np.load(self.style_vec_path)
|
||||||
|
if self.__style_vectors.shape[0] != num_styles:
|
||||||
|
raise ValueError(
|
||||||
|
f"The number of styles ({num_styles}) does not match the number of style vectors ({self.__style_vectors.shape[0]})"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
スタイルベクトルを取得する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
style_id (int): スタイル ID (0 から始まるインデックス)
|
||||||
|
weight (float, optional): スタイルベクトルの重み. Defaults to 1.0.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: スタイルベクトル
|
||||||
|
"""
|
||||||
|
mean = self.__style_vectors[0]
|
||||||
|
style_vec = self.__style_vectors[style_id]
|
||||||
|
style_vec = mean + (style_vec - mean) * weight
|
||||||
|
return style_vec
|
||||||
|
|
||||||
|
def __get_style_vector_from_audio(
|
||||||
|
self, audio_path: str, weight: float = 1.0
|
||||||
|
) -> NDArray[Any]:
|
||||||
|
"""
|
||||||
|
音声からスタイルベクトルを推論する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
audio_path (str): 音声ファイルのパス
|
||||||
|
weight (float, optional): スタイルベクトルの重み. Defaults to 1.0.
|
||||||
|
Returns:
|
||||||
|
NDArray[Any]: スタイルベクトル
|
||||||
|
"""
|
||||||
|
|
||||||
|
# スタイルベクトルを取得するための推論モデルを初期化
|
||||||
|
if self.__style_vector_inference is None:
|
||||||
|
self.__style_vector_inference = pyannote.audio.Inference(
|
||||||
|
model=pyannote.audio.Model.from_pretrained(
|
||||||
|
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
||||||
|
),
|
||||||
|
window="whole",
|
||||||
|
)
|
||||||
|
self.__style_vector_inference.to(torch.device(self.device))
|
||||||
|
|
||||||
|
# 音声からスタイルベクトルを推論
|
||||||
|
xvec = self.__style_vector_inference(audio_path)
|
||||||
|
mean = self.__style_vectors[0]
|
||||||
|
xvec = mean + (xvec - mean) * weight
|
||||||
|
return xvec
|
||||||
|
|
||||||
|
def infer(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
language: Languages = Languages.JP,
|
||||||
|
speaker_id: int = 0,
|
||||||
|
reference_audio_path: Optional[str] = None,
|
||||||
|
sdp_ratio: float = DEFAULT_SDP_RATIO,
|
||||||
|
noise: float = DEFAULT_NOISE,
|
||||||
|
noise_w: float = DEFAULT_NOISEW,
|
||||||
|
length: float = DEFAULT_LENGTH,
|
||||||
|
line_split: bool = DEFAULT_LINE_SPLIT,
|
||||||
|
split_interval: float = DEFAULT_SPLIT_INTERVAL,
|
||||||
|
assist_text: Optional[str] = None,
|
||||||
|
assist_text_weight: float = DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
use_assist_text: bool = False,
|
||||||
|
style: str = DEFAULT_STYLE,
|
||||||
|
style_weight: float = DEFAULT_STYLE_WEIGHT,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
|
pitch_scale: float = 1.0,
|
||||||
|
intonation_scale: float = 1.0,
|
||||||
|
) -> tuple[int, NDArray[Any]]:
|
||||||
|
"""
|
||||||
|
テキストから音声を合成する。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text (str): 読み上げるテキスト
|
||||||
|
language (Languages, optional): 言語. Defaults to Languages.JP.
|
||||||
|
speaker_id (int, optional): 話者 ID. Defaults to 0.
|
||||||
|
reference_audio_path (Optional[str], optional): 音声スタイルの参照元の音声ファイルのパス. Defaults to None.
|
||||||
|
sdp_ratio (float, optional): DP と SDP の混合比。0 で DP のみ、1で SDP のみを使用 (値を大きくするとテンポに緩急がつく). Defaults to DEFAULT_SDP_RATIO.
|
||||||
|
noise (float, optional): DP に与えられるノイズ. Defaults to DEFAULT_NOISE.
|
||||||
|
noise_w (float, optional): SDP に与えられるノイズ. Defaults to DEFAULT_NOISEW.
|
||||||
|
length (float, optional): 生成音声の長さ(話速)のパラメータ。大きいほど生成音声が長くゆっくり、小さいほど短く早くなる。 Defaults to DEFAULT_LENGTH.
|
||||||
|
line_split (bool, optional): テキストを改行ごとに分割して生成するかどうか. Defaults to DEFAULT_LINE_SPLIT.
|
||||||
|
split_interval (float, optional): 改行ごとに分割する場合の無音 (秒). Defaults to DEFAULT_SPLIT_INTERVAL.
|
||||||
|
assist_text (Optional[str], optional): 感情表現の参照元の補助テキスト. Defaults to None.
|
||||||
|
assist_text_weight (float, optional): 感情表現の補助テキストを適用する強さ. Defaults to DEFAULT_ASSIST_TEXT_WEIGHT.
|
||||||
|
use_assist_text (bool, optional): 音声合成時に感情表現の補助テキストを使用するかどうか. Defaults to False.
|
||||||
|
style (str, optional): 音声スタイル (Neutral, Happy など). Defaults to DEFAULT_STYLE.
|
||||||
|
style_weight (float, optional): 音声スタイルを適用する強さ. Defaults to DEFAULT_STYLE_WEIGHT.
|
||||||
|
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
||||||
|
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
|
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
||||||
|
"""
|
||||||
|
|
||||||
|
logger.info(f"Start generating audio data from text:\n{text}")
|
||||||
|
if language != "JP" and self.hyper_parameters.version.endswith("JP-Extra"):
|
||||||
|
raise ValueError(
|
||||||
|
"The model is trained with JP-Extra, but the language is not JP"
|
||||||
|
)
|
||||||
|
if reference_audio_path == "":
|
||||||
|
reference_audio_path = None
|
||||||
|
if assist_text == "" or not use_assist_text:
|
||||||
|
assist_text = None
|
||||||
|
|
||||||
|
if self.__net_g is None:
|
||||||
|
self.load()
|
||||||
|
assert self.__net_g is not None
|
||||||
|
if reference_audio_path is None:
|
||||||
|
style_id = self.style2id[style]
|
||||||
|
style_vector = self.__get_style_vector(style_id, style_weight)
|
||||||
|
else:
|
||||||
|
style_vector = self.__get_style_vector_from_audio(
|
||||||
|
reference_audio_path, style_weight
|
||||||
|
)
|
||||||
|
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_tone=given_tone,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
texts = text.split("\n")
|
||||||
|
texts = [t for t in texts 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)
|
||||||
|
logger.info("Audio data generated successfully")
|
||||||
|
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
||||||
|
_, audio = adjust_voice(
|
||||||
|
fs=self.hyper_parameters.data.sampling_rate,
|
||||||
|
wave=audio,
|
||||||
|
pitch_scale=pitch_scale,
|
||||||
|
intonation_scale=intonation_scale,
|
||||||
|
)
|
||||||
|
with warnings.catch_warnings():
|
||||||
|
warnings.simplefilter("ignore")
|
||||||
|
audio = convert_to_16_bit_wav(audio)
|
||||||
|
return (self.hyper_parameters.data.sampling_rate, audio)
|
||||||
|
|
||||||
|
|
||||||
|
class TTSModelInfo(BaseModel):
|
||||||
|
name: str
|
||||||
|
files: list[str]
|
||||||
|
styles: list[str]
|
||||||
|
speakers: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class TTSModelHolder:
|
||||||
|
"""
|
||||||
|
Style-Bert-Vits2 の音声合成モデルを管理するクラス。
|
||||||
|
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, model_root_dir: Path, device: str) -> None:
|
||||||
|
"""
|
||||||
|
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
|
||||||
|
音声合成モデルは下記のように配置されていることを前提とする (.safetensors のファイル名は自由) 。
|
||||||
|
```
|
||||||
|
model_root_dir
|
||||||
|
├── model-name-1
|
||||||
|
│ ├── config.json
|
||||||
|
│ ├── model-name-1_e160_s14000.safetensors
|
||||||
|
│ └── style_vectors.npy
|
||||||
|
├── model-name-2
|
||||||
|
│ ├── config.json
|
||||||
|
│ ├── model-name-2_e160_s14000.safetensors
|
||||||
|
│ └── style_vectors.npy
|
||||||
|
└── ...
|
||||||
|
```
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス
|
||||||
|
device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など)
|
||||||
|
"""
|
||||||
|
|
||||||
|
self.root_dir: Path = model_root_dir
|
||||||
|
self.device: str = device
|
||||||
|
self.model_files_dict: dict[str, list[Path]] = {}
|
||||||
|
self.current_model: Optional[TTSModel] = None
|
||||||
|
self.model_names: list[str] = []
|
||||||
|
self.models_info: list[TTSModelInfo] = []
|
||||||
|
self.refresh()
|
||||||
|
|
||||||
|
def refresh(self) -> None:
|
||||||
|
"""
|
||||||
|
音声合成モデルの一覧を更新する。
|
||||||
|
"""
|
||||||
|
|
||||||
|
self.model_files_dict = {}
|
||||||
|
self.model_names = []
|
||||||
|
self.current_model = None
|
||||||
|
self.models_info = []
|
||||||
|
|
||||||
|
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
||||||
|
for model_dir in model_dirs:
|
||||||
|
model_files = [
|
||||||
|
f
|
||||||
|
for f in model_dir.iterdir()
|
||||||
|
if f.suffix in [".pth", ".pt", ".safetensors"]
|
||||||
|
]
|
||||||
|
if len(model_files) == 0:
|
||||||
|
logger.warning(f"No model files found in {model_dir}, so skip it")
|
||||||
|
continue
|
||||||
|
config_path = model_dir / "config.json"
|
||||||
|
if not config_path.exists():
|
||||||
|
logger.warning(
|
||||||
|
f"Config file {config_path} not found, so skip {model_dir}"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
self.model_files_dict[model_dir.name] = model_files
|
||||||
|
self.model_names.append(model_dir.name)
|
||||||
|
hyper_parameters = HyperParameters.load_from_json(config_path)
|
||||||
|
style2id: dict[str, int] = hyper_parameters.data.style2id
|
||||||
|
styles = list(style2id.keys())
|
||||||
|
spk2id: dict[str, int] = hyper_parameters.data.spk2id
|
||||||
|
speakers = list(spk2id.keys())
|
||||||
|
self.models_info.append(
|
||||||
|
TTSModelInfo(
|
||||||
|
name=model_dir.name,
|
||||||
|
files=[str(f) for f in model_files],
|
||||||
|
styles=styles,
|
||||||
|
speakers=speakers,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_model(self, model_name: str, model_path_str: str) -> TTSModel:
|
||||||
|
"""
|
||||||
|
指定された音声合成モデルのインスタンスを取得する。
|
||||||
|
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_name (str): 音声合成モデルの名前
|
||||||
|
model_path_str (str): 音声合成モデルのファイルパス (.safetensors)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
TTSModel: 音声合成モデルのインスタンス
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_path = Path(model_path_str)
|
||||||
|
if model_name not in self.model_files_dict:
|
||||||
|
raise ValueError(f"Model `{model_name}` is not found")
|
||||||
|
if model_path not in self.model_files_dict[model_name]:
|
||||||
|
raise ValueError(f"Model file `{model_path}` is not found")
|
||||||
|
if self.current_model is None or self.current_model.model_path != model_path:
|
||||||
|
self.current_model = TTSModel(
|
||||||
|
model_path=model_path,
|
||||||
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
return self.current_model
|
||||||
|
|
||||||
|
def get_model_for_gradio(
|
||||||
|
self, model_name: str, model_path_str: str
|
||||||
|
) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]:
|
||||||
|
model_path = Path(model_path_str)
|
||||||
|
if model_name not in self.model_files_dict:
|
||||||
|
raise ValueError(f"Model `{model_name}` is not found")
|
||||||
|
if model_path not in self.model_files_dict[model_name]:
|
||||||
|
raise ValueError(f"Model file `{model_path}` is not found")
|
||||||
|
if (
|
||||||
|
self.current_model is not None
|
||||||
|
and self.current_model.model_path == model_path
|
||||||
|
):
|
||||||
|
# Already loaded
|
||||||
|
speakers = list(self.current_model.spk2id.keys())
|
||||||
|
styles = list(self.current_model.style2id.keys())
|
||||||
|
return (
|
||||||
|
gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
|
||||||
|
gr.Button(interactive=True, value="音声合成"),
|
||||||
|
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
||||||
|
)
|
||||||
|
self.current_model = TTSModel(
|
||||||
|
model_path=model_path,
|
||||||
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
speakers = list(self.current_model.spk2id.keys())
|
||||||
|
styles = list(self.current_model.style2id.keys())
|
||||||
|
return (
|
||||||
|
gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
|
||||||
|
gr.Button(interactive=True, value="音声合成"),
|
||||||
|
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
||||||
|
)
|
||||||
|
|
||||||
|
def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown:
|
||||||
|
model_files = self.model_files_dict[model_name]
|
||||||
|
return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
|
||||||
|
|
||||||
|
def update_model_names_for_gradio(
|
||||||
|
self,
|
||||||
|
) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
|
||||||
|
self.refresh()
|
||||||
|
initial_model_name = self.model_names[0]
|
||||||
|
initial_model_files = self.model_files_dict[initial_model_name]
|
||||||
|
return (
|
||||||
|
gr.Dropdown(choices=self.model_names, value=initial_model_name), # type: ignore
|
||||||
|
gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), # type: ignore
|
||||||
|
gr.Button(interactive=False), # For tts_button
|
||||||
|
)
|
||||||
0
style_bert_vits2/utils/__init__.py
Normal file
0
style_bert_vits2/utils/__init__.py
Normal file
@@ -1,39 +1,41 @@
|
|||||||
"""
|
|
||||||
`sys.stdout` wrapper for both Google Colab and local environment.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from typing import TextIO
|
||||||
|
|
||||||
|
|
||||||
class StdoutWrapper:
|
class StdoutWrapper(TextIO):
|
||||||
def __init__(self):
|
"""
|
||||||
|
`sys.stdout` wrapper for both Google Colab and local environment.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
self.temp_file = tempfile.NamedTemporaryFile(
|
self.temp_file = tempfile.NamedTemporaryFile(
|
||||||
mode="w+", delete=False, encoding="utf-8"
|
mode="w+", delete=False, encoding="utf-8"
|
||||||
)
|
)
|
||||||
self.original_stdout = sys.stdout
|
self.original_stdout = sys.stdout
|
||||||
|
|
||||||
def write(self, message: str):
|
def write(self, message: str) -> int:
|
||||||
self.temp_file.write(message)
|
result = self.temp_file.write(message)
|
||||||
self.temp_file.flush()
|
self.temp_file.flush()
|
||||||
print(message, end="", file=self.original_stdout)
|
print(message, end="", file=self.original_stdout)
|
||||||
|
return result
|
||||||
|
|
||||||
def flush(self):
|
def flush(self) -> None:
|
||||||
self.temp_file.flush()
|
self.temp_file.flush()
|
||||||
|
|
||||||
def read(self):
|
def read(self, n: int = -1) -> str:
|
||||||
self.temp_file.seek(0)
|
self.temp_file.seek(0)
|
||||||
return self.temp_file.read()
|
return self.temp_file.read(n)
|
||||||
|
|
||||||
def close(self):
|
def close(self) -> None:
|
||||||
self.temp_file.close()
|
self.temp_file.close()
|
||||||
|
|
||||||
def fileno(self):
|
def fileno(self) -> int:
|
||||||
return self.temp_file.fileno()
|
return self.temp_file.fileno()
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import google.colab
|
import google.colab # type: ignore
|
||||||
|
|
||||||
SAFE_STDOUT = StdoutWrapper()
|
SAFE_STDOUT = StdoutWrapper()
|
||||||
except ImportError:
|
except ImportError:
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user