Improve: Support ONNX inference, add ONNX conversion script
This commit is contained in:
1
.gitignore
vendored
1
.gitignore
vendored
@@ -13,6 +13,7 @@ dist/
|
|||||||
/bert/*/*.model
|
/bert/*/*.model
|
||||||
/bert/*/*.safetensors
|
/bert/*/*.safetensors
|
||||||
/bert/*/*.msgpack
|
/bert/*/*.msgpack
|
||||||
|
/bert/*/*.onnx
|
||||||
|
|
||||||
/configs/paths.yml
|
/configs/paths.yml
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +0,0 @@
|
|||||||
*.bin.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,57 +0,0 @@
|
|||||||
---
|
|
||||||
language:
|
|
||||||
- zh
|
|
||||||
tags:
|
|
||||||
- bert
|
|
||||||
license: "apache-2.0"
|
|
||||||
---
|
|
||||||
|
|
||||||
# Please use 'Bert' related functions to load this model!
|
|
||||||
|
|
||||||
## Chinese BERT with Whole Word Masking
|
|
||||||
For further accelerating Chinese natural language processing, we provide **Chinese pre-trained BERT with Whole Word Masking**.
|
|
||||||
|
|
||||||
**[Pre-Training with Whole Word Masking for Chinese BERT](https://arxiv.org/abs/1906.08101)**
|
|
||||||
Yiming Cui, Wanxiang Che, Ting Liu, Bing Qin, Ziqing Yang, Shijin Wang, Guoping Hu
|
|
||||||
|
|
||||||
This repository is developed based on:https://github.com/google-research/bert
|
|
||||||
|
|
||||||
You may also interested in,
|
|
||||||
- Chinese BERT series: https://github.com/ymcui/Chinese-BERT-wwm
|
|
||||||
- Chinese MacBERT: https://github.com/ymcui/MacBERT
|
|
||||||
- Chinese ELECTRA: https://github.com/ymcui/Chinese-ELECTRA
|
|
||||||
- Chinese XLNet: https://github.com/ymcui/Chinese-XLNet
|
|
||||||
- Knowledge Distillation Toolkit - TextBrewer: https://github.com/airaria/TextBrewer
|
|
||||||
|
|
||||||
More resources by HFL: https://github.com/ymcui/HFL-Anthology
|
|
||||||
|
|
||||||
## Citation
|
|
||||||
If you find the technical report or resource is useful, please cite the following technical report in your paper.
|
|
||||||
- Primary: https://arxiv.org/abs/2004.13922
|
|
||||||
```
|
|
||||||
@inproceedings{cui-etal-2020-revisiting,
|
|
||||||
title = "Revisiting Pre-Trained Models for {C}hinese Natural Language Processing",
|
|
||||||
author = "Cui, Yiming and
|
|
||||||
Che, Wanxiang and
|
|
||||||
Liu, Ting and
|
|
||||||
Qin, Bing and
|
|
||||||
Wang, Shijin and
|
|
||||||
Hu, Guoping",
|
|
||||||
booktitle = "Proceedings of the 2020 Conference on Empirical Methods in Natural Language Processing: Findings",
|
|
||||||
month = nov,
|
|
||||||
year = "2020",
|
|
||||||
address = "Online",
|
|
||||||
publisher = "Association for Computational Linguistics",
|
|
||||||
url = "https://www.aclweb.org/anthology/2020.findings-emnlp.58",
|
|
||||||
pages = "657--668",
|
|
||||||
}
|
|
||||||
```
|
|
||||||
- Secondary: https://arxiv.org/abs/1906.08101
|
|
||||||
```
|
|
||||||
@article{chinese-bert-wwm,
|
|
||||||
title={Pre-Training with Whole Word Masking for Chinese BERT},
|
|
||||||
author={Cui, Yiming and Che, Wanxiang and Liu, Ting and Qin, Bing and Yang, Ziqing and Wang, Shijin and Hu, Guoping},
|
|
||||||
journal={arXiv preprint arXiv:1906.08101},
|
|
||||||
year={2019}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
37
bert/deberta-v2-large-japanese-char-wwm-onnx/config.json
Normal file
37
bert/deberta-v2-large-japanese-char-wwm-onnx/config.json
Normal file
@@ -0,0 +1,37 @@
|
|||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"DebertaV2ForMaskedLM"
|
||||||
|
],
|
||||||
|
"attention_head_size": 64,
|
||||||
|
"attention_probs_dropout_prob": 0.1,
|
||||||
|
"conv_act": "gelu",
|
||||||
|
"conv_kernel_size": 3,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_dropout_prob": 0.1,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"layer_norm_eps": 1e-07,
|
||||||
|
"max_position_embeddings": 512,
|
||||||
|
"max_relative_positions": -1,
|
||||||
|
"model_type": "deberta-v2",
|
||||||
|
"norm_rel_ebd": "layer_norm",
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"pad_token_id": 0,
|
||||||
|
"pooler_dropout": 0,
|
||||||
|
"pooler_hidden_act": "gelu",
|
||||||
|
"pooler_hidden_size": 1024,
|
||||||
|
"pos_att_type": [
|
||||||
|
"p2c",
|
||||||
|
"c2p"
|
||||||
|
],
|
||||||
|
"position_biased_input": false,
|
||||||
|
"position_buckets": 256,
|
||||||
|
"relative_attention": true,
|
||||||
|
"share_att_key": true,
|
||||||
|
"torch_dtype": "float16",
|
||||||
|
"transformers_version": "4.25.1",
|
||||||
|
"type_vocab_size": 0,
|
||||||
|
"vocab_size": 22012
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
{
|
||||||
|
"cls_token": "[CLS]",
|
||||||
|
"mask_token": "[MASK]",
|
||||||
|
"pad_token": "[PAD]",
|
||||||
|
"sep_token": "[SEP]",
|
||||||
|
"unk_token": "[UNK]"
|
||||||
|
}
|
||||||
22116
bert/deberta-v2-large-japanese-char-wwm-onnx/tokenizer.json
Normal file
22116
bert/deberta-v2-large-japanese-char-wwm-onnx/tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,19 @@
|
|||||||
|
{
|
||||||
|
"cls_token": "[CLS]",
|
||||||
|
"do_lower_case": false,
|
||||||
|
"do_subword_tokenize": true,
|
||||||
|
"do_word_tokenize": true,
|
||||||
|
"jumanpp_kwargs": null,
|
||||||
|
"mask_token": "[MASK]",
|
||||||
|
"mecab_kwargs": null,
|
||||||
|
"model_max_length": 1000000000000000019884624838656,
|
||||||
|
"never_split": null,
|
||||||
|
"pad_token": "[PAD]",
|
||||||
|
"sep_token": "[SEP]",
|
||||||
|
"special_tokens_map_file": null,
|
||||||
|
"subword_tokenizer_type": "character",
|
||||||
|
"sudachi_kwargs": null,
|
||||||
|
"tokenizer_class": "BertJapaneseTokenizer",
|
||||||
|
"unk_token": "[UNK]",
|
||||||
|
"word_tokenizer_type": "basic"
|
||||||
|
}
|
||||||
22012
bert/deberta-v2-large-japanese-char-wwm-onnx/vocab.txt
Normal file
22012
bert/deberta-v2-large-japanese-char-wwm-onnx/vocab.txt
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,34 +0,0 @@
|
|||||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.model filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.npy filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.npz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pickle filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pkl filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
|
||||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.wasm filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zst filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,89 +0,0 @@
|
|||||||
---
|
|
||||||
language: ja
|
|
||||||
license: cc-by-sa-4.0
|
|
||||||
library_name: transformers
|
|
||||||
tags:
|
|
||||||
- deberta
|
|
||||||
- deberta-v2
|
|
||||||
- fill-mask
|
|
||||||
- character
|
|
||||||
- wwm
|
|
||||||
datasets:
|
|
||||||
- wikipedia
|
|
||||||
- cc100
|
|
||||||
- oscar
|
|
||||||
metrics:
|
|
||||||
- accuracy
|
|
||||||
mask_token: "[MASK]"
|
|
||||||
widget:
|
|
||||||
- text: "京都大学で自然言語処理を[MASK][MASK]する。"
|
|
||||||
---
|
|
||||||
|
|
||||||
# Model Card for Japanese character-level DeBERTa V2 large
|
|
||||||
|
|
||||||
## Model description
|
|
||||||
|
|
||||||
This is a Japanese DeBERTa V2 large model pre-trained on Japanese Wikipedia, the Japanese portion of CC-100, and the Japanese portion of OSCAR.
|
|
||||||
This model is trained with character-level tokenization and whole word masking.
|
|
||||||
|
|
||||||
## How to use
|
|
||||||
|
|
||||||
You can use this model for masked language modeling as follows:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from transformers import AutoTokenizer, AutoModelForMaskedLM
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained('ku-nlp/deberta-v2-large-japanese-char-wwm')
|
|
||||||
model = AutoModelForMaskedLM.from_pretrained('ku-nlp/deberta-v2-large-japanese-char-wwm')
|
|
||||||
|
|
||||||
sentence = '京都大学で自然言語処理を[MASK][MASK]する。'
|
|
||||||
encoding = tokenizer(sentence, return_tensors='pt')
|
|
||||||
...
|
|
||||||
```
|
|
||||||
|
|
||||||
You can also fine-tune this model on downstream tasks.
|
|
||||||
|
|
||||||
## Tokenization
|
|
||||||
|
|
||||||
There is no need to tokenize texts in advance, and you can give raw texts to the tokenizer.
|
|
||||||
The texts are tokenized into character-level tokens by [sentencepiece](https://github.com/google/sentencepiece).
|
|
||||||
|
|
||||||
## Training data
|
|
||||||
|
|
||||||
We used the following corpora for pre-training:
|
|
||||||
|
|
||||||
- Japanese Wikipedia (as of 20221020, 3.2GB, 27M sentences, 1.3M documents)
|
|
||||||
- Japanese portion of CC-100 (85GB, 619M sentences, 66M documents)
|
|
||||||
- Japanese portion of OSCAR (54GB, 326M sentences, 25M documents)
|
|
||||||
|
|
||||||
Note that we filtered out documents annotated with "header", "footer", or "noisy" tags in OSCAR.
|
|
||||||
Also note that Japanese Wikipedia was duplicated 10 times to make the total size of the corpus comparable to that of CC-100 and OSCAR. As a result, the total size of the training data is 171GB.
|
|
||||||
|
|
||||||
## Training procedure
|
|
||||||
|
|
||||||
We first segmented texts in the corpora into words using [Juman++ 2.0.0-rc3](https://github.com/ku-nlp/jumanpp/releases/tag/v2.0.0-rc3) for whole word masking.
|
|
||||||
Then, we built a sentencepiece model with 22,012 tokens including all characters that appear in the training corpus.
|
|
||||||
|
|
||||||
We tokenized raw corpora into character-level subwords using the sentencepiece model and trained the Japanese DeBERTa model using [transformers](https://github.com/huggingface/transformers) library.
|
|
||||||
The training took 26 days using 16 NVIDIA A100-SXM4-40GB GPUs.
|
|
||||||
|
|
||||||
The following hyperparameters were used during pre-training:
|
|
||||||
|
|
||||||
- learning_rate: 1e-4
|
|
||||||
- per_device_train_batch_size: 26
|
|
||||||
- distributed_type: multi-GPU
|
|
||||||
- num_devices: 16
|
|
||||||
- gradient_accumulation_steps: 8
|
|
||||||
- total_train_batch_size: 3,328
|
|
||||||
- max_seq_length: 512
|
|
||||||
- optimizer: Adam with betas=(0.9,0.999) and epsilon=1e-06
|
|
||||||
- lr_scheduler_type: linear schedule with warmup (lr = 0 at 300k steps)
|
|
||||||
- training_steps: 260,000
|
|
||||||
- warmup_steps: 10,000
|
|
||||||
|
|
||||||
The accuracy of the trained model on the masked language modeling task was 0.795.
|
|
||||||
The evaluation set consists of 5,000 randomly sampled documents from each of the training corpora.
|
|
||||||
|
|
||||||
## Acknowledgments
|
|
||||||
|
|
||||||
This work was supported by Joint Usage/Research Center for Interdisciplinary Large-scale Information Infrastructures (JHPCN) through General Collaboration Project no. jh221004, "Developing a Platform for Constructing and Sharing of Large-Scale Japanese Language Models".
|
|
||||||
For training models, we used the mdx: a platform for the data-driven future.
|
|
||||||
27
bert/deberta-v3-large/.gitattributes
vendored
27
bert/deberta-v3-large/.gitattributes
vendored
@@ -1,27 +0,0 @@
|
|||||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bin.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.model filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
|
||||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*.zstandard filter=lfs diff=lfs merge=lfs -text
|
|
||||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
||||||
@@ -1,93 +0,0 @@
|
|||||||
---
|
|
||||||
language: en
|
|
||||||
tags:
|
|
||||||
- deberta
|
|
||||||
- deberta-v3
|
|
||||||
- fill-mask
|
|
||||||
thumbnail: https://huggingface.co/front/thumbnails/microsoft.png
|
|
||||||
license: mit
|
|
||||||
---
|
|
||||||
|
|
||||||
## DeBERTaV3: Improving DeBERTa using ELECTRA-Style Pre-Training with Gradient-Disentangled Embedding Sharing
|
|
||||||
|
|
||||||
[DeBERTa](https://arxiv.org/abs/2006.03654) improves the BERT and RoBERTa models using disentangled attention and enhanced mask decoder. With those two improvements, DeBERTa out perform RoBERTa on a majority of NLU tasks with 80GB training data.
|
|
||||||
|
|
||||||
In [DeBERTa V3](https://arxiv.org/abs/2111.09543), we further improved the efficiency of DeBERTa using ELECTRA-Style pre-training with Gradient Disentangled Embedding Sharing. Compared to DeBERTa, our V3 version significantly improves the model performance on downstream tasks. You can find more technique details about the new model from our [paper](https://arxiv.org/abs/2111.09543).
|
|
||||||
|
|
||||||
Please check the [official repository](https://github.com/microsoft/DeBERTa) for more implementation details and updates.
|
|
||||||
|
|
||||||
The DeBERTa V3 large model comes with 24 layers and a hidden size of 1024. It has 304M backbone parameters with a vocabulary containing 128K tokens which introduces 131M parameters in the Embedding layer. This model was trained using the 160GB data as DeBERTa V2.
|
|
||||||
|
|
||||||
|
|
||||||
#### Fine-tuning on NLU tasks
|
|
||||||
|
|
||||||
We present the dev results on SQuAD 2.0 and MNLI tasks.
|
|
||||||
|
|
||||||
| Model |Vocabulary(K)|Backbone #Params(M)| SQuAD 2.0(F1/EM) | MNLI-m/mm(ACC)|
|
|
||||||
|-------------------|----------|-------------------|-----------|----------|
|
|
||||||
| RoBERTa-large |50 |304 | 89.4/86.5 | 90.2 |
|
|
||||||
| XLNet-large |32 |- | 90.6/87.9 | 90.8 |
|
|
||||||
| DeBERTa-large |50 |- | 90.7/88.0 | 91.3 |
|
|
||||||
| **DeBERTa-v3-large**|128|304 | **91.5/89.0**| **91.8/91.9**|
|
|
||||||
|
|
||||||
|
|
||||||
#### Fine-tuning with HF transformers
|
|
||||||
|
|
||||||
```bash
|
|
||||||
#!/bin/bash
|
|
||||||
|
|
||||||
cd transformers/examples/pytorch/text-classification/
|
|
||||||
|
|
||||||
pip install datasets
|
|
||||||
export TASK_NAME=mnli
|
|
||||||
|
|
||||||
output_dir="ds_results"
|
|
||||||
|
|
||||||
num_gpus=8
|
|
||||||
|
|
||||||
batch_size=8
|
|
||||||
|
|
||||||
python -m torch.distributed.launch --nproc_per_node=${num_gpus} \
|
|
||||||
run_glue.py \
|
|
||||||
--model_name_or_path microsoft/deberta-v3-large \
|
|
||||||
--task_name $TASK_NAME \
|
|
||||||
--do_train \
|
|
||||||
--do_eval \
|
|
||||||
--evaluation_strategy steps \
|
|
||||||
--max_seq_length 256 \
|
|
||||||
--warmup_steps 50 \
|
|
||||||
--per_device_train_batch_size ${batch_size} \
|
|
||||||
--learning_rate 6e-6 \
|
|
||||||
--num_train_epochs 2 \
|
|
||||||
--output_dir $output_dir \
|
|
||||||
--overwrite_output_dir \
|
|
||||||
--logging_steps 1000 \
|
|
||||||
--logging_dir $output_dir
|
|
||||||
|
|
||||||
```
|
|
||||||
|
|
||||||
### Citation
|
|
||||||
|
|
||||||
If you find DeBERTa useful for your work, please cite the following papers:
|
|
||||||
|
|
||||||
``` latex
|
|
||||||
@misc{he2021debertav3,
|
|
||||||
title={DeBERTaV3: Improving DeBERTa using ELECTRA-Style Pre-Training with Gradient-Disentangled Embedding Sharing},
|
|
||||||
author={Pengcheng He and Jianfeng Gao and Weizhu Chen},
|
|
||||||
year={2021},
|
|
||||||
eprint={2111.09543},
|
|
||||||
archivePrefix={arXiv},
|
|
||||||
primaryClass={cs.CL}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
``` latex
|
|
||||||
@inproceedings{
|
|
||||||
he2021deberta,
|
|
||||||
title={DEBERTA: DECODING-ENHANCED BERT WITH DISENTANGLED ATTENTION},
|
|
||||||
author={Pengcheng He and Xiaodong Liu and Jianfeng Gao and Weizhu Chen},
|
|
||||||
booktitle={International Conference on Learning Representations},
|
|
||||||
year={2021},
|
|
||||||
url={https://openreview.net/forum?id=XPZIaotutsD}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
75
convert_bert_onnx.py
Normal file
75
convert_bert_onnx.py
Normal file
@@ -0,0 +1,75 @@
|
|||||||
|
# usage: .venv/bin/python convert_bert_onnx.py --language JP
|
||||||
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_deberta.py
|
||||||
|
|
||||||
|
from argparse import ArgumentParser
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import onnx
|
||||||
|
import torch
|
||||||
|
from onnxsim import simplify
|
||||||
|
from torch import nn
|
||||||
|
from transformers.convert_slow_tokenizer import BertConverter
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
|
||||||
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = ArgumentParser()
|
||||||
|
parser.add_argument("--language", default=Languages.JP)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# モデルの入出力先ファイルパスを取得
|
||||||
|
language = Languages(args.language)
|
||||||
|
pretrained_model_name_or_path = DEFAULT_BERT_MODEL_PATHS[language]
|
||||||
|
onnx_temp_model_path = Path(pretrained_model_name_or_path) / f"model_temp.onnx"
|
||||||
|
onnx_optimized_model_path = Path(pretrained_model_name_or_path) / f"model.onnx"
|
||||||
|
tokenizer_json_path = Path(pretrained_model_name_or_path) / "tokenizer.json"
|
||||||
|
|
||||||
|
# トークナイザーを Fast Tokenizer 用形式に変換して保存
|
||||||
|
tokenizer = bert_models.load_tokenizer(language)
|
||||||
|
converter = BertConverter(tokenizer)
|
||||||
|
converter.converted().save(str(tokenizer_json_path))
|
||||||
|
|
||||||
|
# TODO: JP, ZH は変換できるが、EN は途中で強制終了されてしまい変換できない
|
||||||
|
class ONNXBert(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super(ONNXBert, self).__init__()
|
||||||
|
self.model = bert_models.load_model(language)
|
||||||
|
|
||||||
|
def forward(self, input_ids, token_type_ids, attention_mask):
|
||||||
|
inputs = {
|
||||||
|
"input_ids": input_ids,
|
||||||
|
"token_type_ids": token_type_ids,
|
||||||
|
"attention_mask": attention_mask,
|
||||||
|
}
|
||||||
|
res = self.model(**inputs, output_hidden_states=True)
|
||||||
|
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||||
|
return res
|
||||||
|
|
||||||
|
# ONNX 変換用の BERT モデルをロード
|
||||||
|
model = ONNXBert()
|
||||||
|
inputs = tokenizer("今日はいい天気ですね", return_tensors="pt")
|
||||||
|
|
||||||
|
# モデルを ONNX に変換
|
||||||
|
torch.onnx.export(
|
||||||
|
model=model,
|
||||||
|
args=(inputs["input_ids"], inputs["token_type_ids"], inputs["attention_mask"]),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
input_names=["input_ids", "token_type_ids", "attention_mask"],
|
||||||
|
output_names=["output"],
|
||||||
|
verbose=True,
|
||||||
|
dynamic_axes={
|
||||||
|
"input_ids": {1: "batch_size"},
|
||||||
|
"attention_mask": {1: "batch_size"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# ONNX モデルを最適化
|
||||||
|
onnx_model = onnx.load(onnx_temp_model_path)
|
||||||
|
simplified_onnx_model, check = simplify(onnx_model)
|
||||||
|
onnx.save(simplified_onnx_model, onnx_optimized_model_path)
|
||||||
|
|
||||||
|
# 最適化前の ONNX モデルを削除
|
||||||
|
onnx_temp_model_path.unlink()
|
||||||
|
print(f"ONNX model optimized and saved to {onnx_optimized_model_path}")
|
||||||
150
convert_onnx.py
Normal file
150
convert_onnx.py
Normal file
@@ -0,0 +1,150 @@
|
|||||||
|
# usage: .venv/bin/python convert_onnx.py --model model_assets/amitaro/amitaro.safetensors
|
||||||
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py
|
||||||
|
|
||||||
|
from argparse import ArgumentParser
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import cast
|
||||||
|
|
||||||
|
import onnx
|
||||||
|
import torch
|
||||||
|
from onnxsim import simplify
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import (
|
||||||
|
DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
DEFAULT_STYLE,
|
||||||
|
DEFAULT_STYLE_WEIGHT,
|
||||||
|
Languages,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.models.infer import get_text
|
||||||
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.tts_model import TTSModel
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = ArgumentParser()
|
||||||
|
parser.add_argument("--model", required=True)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# モデルの入出力先ファイルパスを取得
|
||||||
|
model_path = Path(args.model)
|
||||||
|
onnx_temp_model_path = Path(args.model).parent / f"{model_path.stem}_temp.onnx"
|
||||||
|
onnx_optimized_model_path = Path(args.model).parent / f"{model_path.stem}.onnx"
|
||||||
|
config_path = Path(args.model).parent / "config.json"
|
||||||
|
style_vec_path = Path(args.model).parent / "style_vectors.npy"
|
||||||
|
assert model_path.exists(), "Model file does not exist"
|
||||||
|
assert config_path.exists(), "Config file does not exist"
|
||||||
|
assert style_vec_path.exists(), "Style vector file does not exist"
|
||||||
|
assert model_path.suffix != ".onnx", "Model file is already ONNX"
|
||||||
|
|
||||||
|
# PyTorch モデルを読み込む
|
||||||
|
device = "cpu"
|
||||||
|
tts_model = TTSModel(
|
||||||
|
model_path=model_path,
|
||||||
|
config_path=config_path,
|
||||||
|
style_vec_path=style_vec_path,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
tts_model.load()
|
||||||
|
style_id = tts_model.style2id[DEFAULT_STYLE]
|
||||||
|
assert tts_model.net_g is not None, "Model is not loaded"
|
||||||
|
assert (
|
||||||
|
tts_model.hyper_parameters.data.use_jp_extra is True
|
||||||
|
), "Normal model is not supported yet"
|
||||||
|
|
||||||
|
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
||||||
|
def forward(
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_lengths: torch.Tensor,
|
||||||
|
sid: torch.Tensor,
|
||||||
|
tone: torch.Tensor,
|
||||||
|
language: torch.Tensor,
|
||||||
|
bert: torch.Tensor,
|
||||||
|
style_vec: torch.Tensor,
|
||||||
|
length_scale: float = 1.0,
|
||||||
|
sdp_ratio: float = 0.0,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]:
|
||||||
|
return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer(
|
||||||
|
x,
|
||||||
|
x_lengths,
|
||||||
|
sid,
|
||||||
|
tone,
|
||||||
|
language,
|
||||||
|
bert,
|
||||||
|
style_vec,
|
||||||
|
sdp_ratio=sdp_ratio,
|
||||||
|
length_scale=length_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
tts_model.net_g.forward = forward # type: ignore
|
||||||
|
|
||||||
|
# 音声合成に必要な BERT 特徴量・音素列・アクセント列・言語 ID を取得
|
||||||
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
|
"今日はいい天気ですね。",
|
||||||
|
Languages.JP,
|
||||||
|
tts_model.hyper_parameters,
|
||||||
|
device,
|
||||||
|
assist_text=None,
|
||||||
|
assist_text_weight=DEFAULT_ASSIST_TEXT_WEIGHT,
|
||||||
|
given_phone=None,
|
||||||
|
given_tone=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# スタイルベクトルを取得
|
||||||
|
style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT)
|
||||||
|
|
||||||
|
x_tst = phones.to(device).unsqueeze(0)
|
||||||
|
tones = tones.to(device).unsqueeze(0)
|
||||||
|
lang_ids = lang_ids.to(device).unsqueeze(0)
|
||||||
|
bert = bert.to(device).unsqueeze(0)
|
||||||
|
ja_bert = ja_bert.to(device).unsqueeze(0)
|
||||||
|
en_bert = en_bert.to(device).unsqueeze(0)
|
||||||
|
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
|
||||||
|
style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0)
|
||||||
|
|
||||||
|
# モデルを ONNX に変換
|
||||||
|
torch.onnx.export(
|
||||||
|
model=tts_model.net_g,
|
||||||
|
args=(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
torch.LongTensor([0]).to(device),
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
torch.tensor(1.0),
|
||||||
|
torch.tensor(0.0),
|
||||||
|
),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
verbose=True,
|
||||||
|
dynamic_axes={
|
||||||
|
"x_tst": {1: "batch_size"},
|
||||||
|
"x_tst_lengths": {0: "batch_size"},
|
||||||
|
"tones": {1: "batch_size"},
|
||||||
|
"language": {1: "batch_size"},
|
||||||
|
"bert": {2: "batch_size"},
|
||||||
|
},
|
||||||
|
input_names=[
|
||||||
|
"x_tst",
|
||||||
|
"x_tst_lengths",
|
||||||
|
"sid",
|
||||||
|
"tones",
|
||||||
|
"language",
|
||||||
|
"bert",
|
||||||
|
"style_vec",
|
||||||
|
"length_scale",
|
||||||
|
"sdp_ratio",
|
||||||
|
],
|
||||||
|
output_names=["output"],
|
||||||
|
)
|
||||||
|
|
||||||
|
# ONNX モデルを最適化
|
||||||
|
onnx_model = onnx.load(onnx_temp_model_path)
|
||||||
|
simplified_onnx_model, check = simplify(onnx_model)
|
||||||
|
onnx.save(simplified_onnx_model, onnx_optimized_model_path)
|
||||||
|
|
||||||
|
# 最適化前の ONNX モデルを削除
|
||||||
|
onnx_temp_model_path.unlink()
|
||||||
|
print(f"ONNX model optimized and saved to {onnx_optimized_model_path}")
|
||||||
@@ -41,7 +41,7 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
@@ -193,6 +193,12 @@ skip_static_files = bool(args.skip_static_files)
|
|||||||
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
||||||
bert_models.load_model(Languages.JP, device_map=device)
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
|
if device == "cpu":
|
||||||
|
onnx_provider = "CPUExecutionProvider"
|
||||||
|
else:
|
||||||
|
onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
|
||||||
|
onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider])
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, device)
|
||||||
if len(model_holder.model_names) == 0:
|
if len(model_holder.model_names) == 0:
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
from style_bert_vits2.nlp.japanese.user_dict import update_dict
|
||||||
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
|
||||||
@@ -103,6 +103,12 @@ if __name__ == "__main__":
|
|||||||
bert_models.load_tokenizer(Languages.EN)
|
bert_models.load_tokenizer(Languages.EN)
|
||||||
bert_models.load_model(Languages.ZH, device_map=device)
|
bert_models.load_model(Languages.ZH, device_map=device)
|
||||||
bert_models.load_tokenizer(Languages.ZH)
|
bert_models.load_tokenizer(Languages.ZH)
|
||||||
|
if device == "cpu":
|
||||||
|
onnx_provider = "CPUExecutionProvider"
|
||||||
|
else:
|
||||||
|
onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
|
||||||
|
onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider])
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_dir = Path(args.dir)
|
model_dir = Path(args.dir)
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, device)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import Any, Optional, Sequence
|
from typing import Any, Optional, Sequence, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import onnxruntime
|
import onnxruntime
|
||||||
@@ -35,8 +35,7 @@ def get_text_onnx(
|
|||||||
text: str,
|
text: str,
|
||||||
language_str: Languages,
|
language_str: Languages,
|
||||||
hps: HyperParameters,
|
hps: HyperParameters,
|
||||||
onnx_providers: list[str],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
|
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
given_phone: Optional[list[str]] = None,
|
given_phone: Optional[list[str]] = None,
|
||||||
@@ -68,7 +67,6 @@ def get_text_onnx(
|
|||||||
word2ph,
|
word2ph,
|
||||||
language_str,
|
language_str,
|
||||||
onnx_providers,
|
onnx_providers,
|
||||||
onnx_provider_options,
|
|
||||||
assist_text,
|
assist_text,
|
||||||
assist_text_weight,
|
assist_text_weight,
|
||||||
)
|
)
|
||||||
@@ -111,8 +109,7 @@ def infer_onnx(
|
|||||||
language: Languages,
|
language: Languages,
|
||||||
hps: HyperParameters,
|
hps: HyperParameters,
|
||||||
onnx_session: onnxruntime.InferenceSession,
|
onnx_session: onnxruntime.InferenceSession,
|
||||||
onnx_providers: list[str],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
|
|
||||||
skip_start: bool = False,
|
skip_start: bool = False,
|
||||||
skip_end: bool = False,
|
skip_end: bool = False,
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
@@ -126,7 +123,6 @@ def infer_onnx(
|
|||||||
language,
|
language,
|
||||||
hps,
|
hps,
|
||||||
onnx_providers=onnx_providers,
|
onnx_providers=onnx_providers,
|
||||||
onnx_provider_options=onnx_provider_options,
|
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
given_phone=given_phone,
|
given_phone=given_phone,
|
||||||
@@ -171,8 +167,8 @@ def infer_onnx(
|
|||||||
input_names[4]: lang_ids,
|
input_names[4]: lang_ids,
|
||||||
input_names[5]: ja_bert,
|
input_names[5]: ja_bert,
|
||||||
input_names[6]: style_vec_tensor,
|
input_names[6]: style_vec_tensor,
|
||||||
input_names[7]: length_scale,
|
input_names[7]: np.array([length_scale], dtype=np.float32),
|
||||||
input_names[8]: sdp_ratio,
|
input_names[8]: np.array([sdp_ratio], dtype=np.float32),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Optional, Sequence
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
from numpy.typing import NDArray
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
@@ -60,8 +60,7 @@ def extract_bert_feature_onnx(
|
|||||||
text: str,
|
text: str,
|
||||||
word2ph: list[int],
|
word2ph: list[int],
|
||||||
language: Languages,
|
language: Languages,
|
||||||
onnx_providers: list[str],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
|
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> NDArray[Any]:
|
) -> NDArray[Any]:
|
||||||
@@ -73,7 +72,6 @@ def extract_bert_feature_onnx(
|
|||||||
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
language (Languages): テキストの言語
|
language (Languages): テキストの言語
|
||||||
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
|
|
||||||
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
@@ -90,7 +88,6 @@ def extract_bert_feature_onnx(
|
|||||||
text,
|
text,
|
||||||
word2ph,
|
word2ph,
|
||||||
onnx_providers,
|
onnx_providers,
|
||||||
onnx_provider_options,
|
|
||||||
assist_text,
|
assist_text,
|
||||||
assist_text_weight,
|
assist_text_weight,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Any, Optional, Sequence
|
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from numpy.typing import NDArray
|
from numpy.typing import NDArray
|
||||||
@@ -86,8 +86,7 @@ def extract_bert_feature(
|
|||||||
def extract_bert_feature_onnx(
|
def extract_bert_feature_onnx(
|
||||||
text: str,
|
text: str,
|
||||||
word2ph: list[int],
|
word2ph: list[int],
|
||||||
onnx_providers: list[str],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
|
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> NDArray[Any]:
|
) -> NDArray[Any]:
|
||||||
@@ -98,7 +97,6 @@ def extract_bert_feature_onnx(
|
|||||||
text (str): 日本語のテキスト
|
text (str): 日本語のテキスト
|
||||||
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
|
|
||||||
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
@@ -118,7 +116,6 @@ def extract_bert_feature_onnx(
|
|||||||
session = onnx_bert_models.load_model(
|
session = onnx_bert_models.load_model(
|
||||||
language=Languages.JP,
|
language=Languages.JP,
|
||||||
onnx_providers=onnx_providers,
|
onnx_providers=onnx_providers,
|
||||||
onnx_provider_options=onnx_provider_options,
|
|
||||||
)
|
)
|
||||||
output_name = session.get_outputs()[0].name
|
output_name = session.get_outputs()[0].name
|
||||||
res = session.run(
|
res = session.run(
|
||||||
|
|||||||
@@ -39,11 +39,10 @@ __loaded_tokenizers: dict[
|
|||||||
def load_model(
|
def load_model(
|
||||||
language: Languages,
|
language: Languages,
|
||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
onnx_providers: list[str] = ["CPUExecutionProvider"],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = ["CPUExecutionProvider"],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]] = None,
|
|
||||||
cache_dir: Optional[str] = None,
|
cache_dir: Optional[str] = None,
|
||||||
revision: str = "main",
|
revision: str = "main",
|
||||||
) -> onnxruntime.InferenceSession:
|
) -> onnxruntime.InferenceSession: # fmt: skip
|
||||||
"""
|
"""
|
||||||
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
|
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
|
||||||
一度ロードされていれば、ロード済みの ONNX 版 BERT モデルを即座に返す。
|
一度ロードされていれば、ロード済みの ONNX 版 BERT モデルを即座に返す。
|
||||||
@@ -59,7 +58,6 @@ def load_model(
|
|||||||
language (Languages): ロードする学習済みモデルの対象言語
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
|
|
||||||
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
@@ -98,7 +96,6 @@ def load_model(
|
|||||||
__loaded_models[language] = onnxruntime.InferenceSession(
|
__loaded_models[language] = onnxruntime.InferenceSession(
|
||||||
model_path,
|
model_path,
|
||||||
providers=onnx_providers,
|
providers=onnx_providers,
|
||||||
provider_options=onnx_provider_options,
|
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} ONNX BERT model from {pretrained_model_name_or_path}"
|
f"Loaded the {language} ONNX BERT model from {pretrained_model_name_or_path}"
|
||||||
|
|||||||
@@ -44,9 +44,8 @@ class TTSModel:
|
|||||||
config_path: Union[Path, HyperParameters],
|
config_path: Union[Path, HyperParameters],
|
||||||
style_vec_path: Union[Path, NDArray[Any]],
|
style_vec_path: Union[Path, NDArray[Any]],
|
||||||
device: str = "cpu",
|
device: str = "cpu",
|
||||||
onnx_providers: list[str] = ["CPUExecutionProvider"],
|
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = ["CPUExecutionProvider"],
|
||||||
onnx_provider_options: Optional[Sequence[dict[str, Any]]] = None,
|
) -> None: # fmt: skip
|
||||||
) -> None:
|
|
||||||
"""
|
"""
|
||||||
Style-Bert-VITS2 の音声合成モデルを初期化する。
|
Style-Bert-VITS2 の音声合成モデルを初期化する。
|
||||||
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
|
||||||
@@ -57,13 +56,11 @@ class TTSModel:
|
|||||||
style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能)
|
style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能)
|
||||||
device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など)
|
device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など)
|
||||||
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
self.model_path: Path = model_path
|
self.model_path: Path = model_path
|
||||||
self.device: str = device
|
self.device: str = device
|
||||||
self.onnx_providers: list[str] = onnx_providers
|
self.onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = onnx_providers # fmt: skip
|
||||||
self.onnx_provider_options: Optional[Sequence[dict[str, Any]]] = onnx_provider_options # fmt: skip
|
|
||||||
|
|
||||||
# ONNX 形式のモデルかどうか
|
# ONNX 形式のモデルかどうか
|
||||||
if self.model_path.suffix == ".onnx":
|
if self.model_path.suffix == ".onnx":
|
||||||
@@ -85,11 +82,11 @@ class TTSModel:
|
|||||||
# スタイルベクトルの NDArray が直接指定された
|
# スタイルベクトルの NDArray が直接指定された
|
||||||
if isinstance(style_vec_path, np.ndarray):
|
if isinstance(style_vec_path, np.ndarray):
|
||||||
self.style_vec_path: Path = Path("") # 互換性のため空の Path を設定
|
self.style_vec_path: Path = Path("") # 互換性のため空の Path を設定
|
||||||
self.__style_vectors: NDArray[Any] = style_vec_path
|
self.style_vectors: NDArray[Any] = style_vec_path
|
||||||
# スタイルベクトルのパスが指定された
|
# スタイルベクトルのパスが指定された
|
||||||
else:
|
else:
|
||||||
self.style_vec_path: Path = style_vec_path
|
self.style_vec_path: Path = style_vec_path
|
||||||
self.__style_vectors: NDArray[Any] = np.load(self.style_vec_path)
|
self.style_vectors: NDArray[Any] = np.load(self.style_vec_path)
|
||||||
|
|
||||||
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
||||||
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
||||||
@@ -104,17 +101,17 @@ class TTSModel:
|
|||||||
f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})"
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.__style_vectors.shape[0] != num_styles:
|
if self.style_vectors.shape[0] != num_styles:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The number of styles ({num_styles}) does not match the number of style vectors ({self.__style_vectors.shape[0]})"
|
f"The number of styles ({num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})"
|
||||||
)
|
)
|
||||||
self.__style_vector_inference: Optional[Any] = None
|
self.style_vector_inference: Optional[Any] = None
|
||||||
|
|
||||||
# __net_g は PyTorch 推論時のみ遅延初期化される
|
# net_g は PyTorch 推論時のみ遅延初期化される
|
||||||
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
||||||
|
|
||||||
# __onnx_session は ONNX 推論時のみ遅延初期化される
|
# onnx_session は ONNX 推論時のみ遅延初期化される
|
||||||
self.__onnx_session: Optional[onnxruntime.InferenceSession] = None
|
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
|
||||||
|
|
||||||
def load(self) -> None:
|
def load(self) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -125,7 +122,7 @@ class TTSModel:
|
|||||||
if not self.is_onnx_model:
|
if not self.is_onnx_model:
|
||||||
from style_bert_vits2.models.infer import get_net_g
|
from style_bert_vits2.models.infer import get_net_g
|
||||||
|
|
||||||
self.__net_g = get_net_g(
|
self.net_g = get_net_g(
|
||||||
model_path=str(self.model_path),
|
model_path=str(self.model_path),
|
||||||
version=self.hyper_parameters.version,
|
version=self.hyper_parameters.version,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -134,12 +131,13 @@ class TTSModel:
|
|||||||
|
|
||||||
# ONNX 推論時
|
# ONNX 推論時
|
||||||
else:
|
else:
|
||||||
self.__onnx_session = onnxruntime.InferenceSession(
|
self.onnx_session = onnxruntime.InferenceSession(
|
||||||
path_or_bytes=str(self.model_path),
|
path_or_bytes=str(self.model_path),
|
||||||
providers=self.onnx_providers,
|
providers=self.onnx_providers,
|
||||||
provider_options=self.onnx_provider_options,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
logger.info(f"Model loaded successfully from {self.model_path}")
|
||||||
|
|
||||||
def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
スタイルベクトルを取得する。
|
スタイルベクトルを取得する。
|
||||||
@@ -151,8 +149,8 @@ class TTSModel:
|
|||||||
Returns:
|
Returns:
|
||||||
NDArray[Any]: スタイルベクトル
|
NDArray[Any]: スタイルベクトル
|
||||||
"""
|
"""
|
||||||
mean = self.__style_vectors[0]
|
mean = self.style_vectors[0]
|
||||||
style_vec = self.__style_vectors[style_id]
|
style_vec = self.style_vectors[style_id]
|
||||||
style_vec = mean + (style_vec - mean) * weight
|
style_vec = mean + (style_vec - mean) * weight
|
||||||
return style_vec
|
return style_vec
|
||||||
|
|
||||||
@@ -169,7 +167,7 @@ class TTSModel:
|
|||||||
NDArray[Any]: スタイルベクトル
|
NDArray[Any]: スタイルベクトル
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if self.__style_vector_inference is None:
|
if self.style_vector_inference is None:
|
||||||
|
|
||||||
# pyannote.audio は scikit-learn などの大量の重量級ライブラリに依存しているため、
|
# pyannote.audio は scikit-learn などの大量の重量級ライブラリに依存しているため、
|
||||||
# TTSModel.infer() に reference_audio_path を指定し音声からスタイルベクトルを推論する場合のみ遅延 import する
|
# TTSModel.infer() に reference_audio_path を指定し音声からスタイルベクトルを推論する場合のみ遅延 import する
|
||||||
@@ -183,17 +181,17 @@ class TTSModel:
|
|||||||
# スタイルベクトルを取得するための推論モデルを初期化
|
# スタイルベクトルを取得するための推論モデルを初期化
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
self.__style_vector_inference = pyannote.audio.Inference(
|
self.style_vector_inference = pyannote.audio.Inference(
|
||||||
model=pyannote.audio.Model.from_pretrained(
|
model=pyannote.audio.Model.from_pretrained(
|
||||||
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
||||||
),
|
),
|
||||||
window="whole",
|
window="whole",
|
||||||
)
|
)
|
||||||
self.__style_vector_inference.to(torch.device(self.device))
|
self.style_vector_inference.to(torch.device(self.device))
|
||||||
|
|
||||||
# 音声からスタイルベクトルを推論
|
# 音声からスタイルベクトルを推論
|
||||||
xvec = self.__style_vector_inference(audio_path)
|
xvec = self.style_vector_inference(audio_path)
|
||||||
mean = self.__style_vectors[0]
|
mean = self.style_vectors[0]
|
||||||
xvec = mean + (xvec - mean) * weight
|
xvec = mean + (xvec - mean) * weight
|
||||||
return xvec
|
return xvec
|
||||||
|
|
||||||
@@ -310,9 +308,9 @@ class TTSModel:
|
|||||||
from style_bert_vits2.models.infer import infer
|
from style_bert_vits2.models.infer import infer
|
||||||
|
|
||||||
# モデルがロードされていない場合はロードする
|
# モデルがロードされていない場合はロードする
|
||||||
if self.__net_g is None:
|
if self.net_g is None:
|
||||||
self.load()
|
self.load()
|
||||||
assert self.__net_g is not None
|
assert self.net_g is not None
|
||||||
|
|
||||||
# 通常のテキストから音声を生成
|
# 通常のテキストから音声を生成
|
||||||
if not line_split:
|
if not line_split:
|
||||||
@@ -326,7 +324,7 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g=self.__net_g,
|
net_g=self.net_g,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
@@ -352,7 +350,7 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g=self.__net_g,
|
net_g=self.net_g,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
@@ -368,9 +366,9 @@ class TTSModel:
|
|||||||
from style_bert_vits2.models.infer_onnx import infer_onnx
|
from style_bert_vits2.models.infer_onnx import infer_onnx
|
||||||
|
|
||||||
# モデルがロードされていない場合はロードする
|
# モデルがロードされていない場合はロードする
|
||||||
if self.__onnx_session is None:
|
if self.onnx_session is None:
|
||||||
self.load()
|
self.load()
|
||||||
assert self.__onnx_session is not None
|
assert self.onnx_session is not None
|
||||||
|
|
||||||
# 通常のテキストから音声を生成
|
# 通常のテキストから音声を生成
|
||||||
if not line_split:
|
if not line_split:
|
||||||
@@ -383,9 +381,8 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
onnx_session=self.__onnx_session,
|
onnx_session=self.onnx_session,
|
||||||
onnx_providers=self.onnx_providers,
|
onnx_providers=self.onnx_providers,
|
||||||
onnx_provider_options=self.onnx_provider_options,
|
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec=style_vector,
|
style_vec=style_vector,
|
||||||
@@ -409,9 +406,8 @@ class TTSModel:
|
|||||||
sid=speaker_id,
|
sid=speaker_id,
|
||||||
language=language,
|
language=language,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
onnx_session=self.__onnx_session,
|
onnx_session=self.onnx_session,
|
||||||
onnx_providers=self.onnx_providers,
|
onnx_providers=self.onnx_providers,
|
||||||
onnx_provider_options=self.onnx_provider_options,
|
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec=style_vector,
|
style_vec=style_vector,
|
||||||
@@ -487,13 +483,17 @@ class TTSModelHolder:
|
|||||||
self.current_model = None
|
self.current_model = None
|
||||||
self.models_info = []
|
self.models_info = []
|
||||||
|
|
||||||
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
model_dirs = sorted([d for d in self.root_dir.iterdir() if d.is_dir()])
|
||||||
for model_dir in model_dirs:
|
for model_dir in model_dirs:
|
||||||
model_files = [
|
if model_dir.name.startswith("."):
|
||||||
|
continue
|
||||||
|
model_files = sorted(
|
||||||
|
[
|
||||||
f
|
f
|
||||||
for f in model_dir.iterdir()
|
for f in model_dir.iterdir()
|
||||||
if f.suffix in [".pth", ".pt", ".safetensors", ".onnx"]
|
if f.suffix in [".pth", ".pt", ".safetensors", ".onnx"]
|
||||||
]
|
]
|
||||||
|
)
|
||||||
if len(model_files) == 0:
|
if len(model_files) == 0:
|
||||||
logger.warning(f"No model files found in {model_dir}, so skip it")
|
logger.warning(f"No model files found in {model_dir}, so skip it")
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user