Auto download update, optimize dataloader num_workers (#195)

* update bert

* auto download emo

* fix typo

* fix typo

* fix bert download

* optimize code format

* remove unsued import

* fix a bug

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Isotr0py
2023-11-26 20:16:25 +08:00
committed by GitHub
parent 15babcd739
commit dec3fc0737
6 changed files with 41 additions and 28 deletions

View File

@@ -1,6 +1,6 @@
{ {
"deberta-v2-large-japanese": { "deberta-v2-large-japanese-char-wwm": {
"repo_id": "ku-nlp/deberta-v2-large-japanese", "repo_id": "ku-nlp/deberta-v2-large-japanese-char-wwm",
"files": ["pytorch_model.bin"] "files": ["pytorch_model.bin"]
}, },
"chinese-roberta-wwm-ext-large": { "chinese-roberta-wwm-ext-large": {

View File

@@ -3,7 +3,7 @@ from multiprocessing import Pool
import commons import commons
import utils import utils
from tqdm import tqdm from tqdm import tqdm
from text import check_bert_models, cleaned_text_to_sequence, get_bert from text import cleaned_text_to_sequence, get_bert
import argparse import argparse
import torch.multiprocessing as mp import torch.multiprocessing as mp
from config import config from config import config
@@ -57,7 +57,6 @@ if __name__ == "__main__":
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 = utils.get_hparams_from_file(config_path)
check_bert_models()
lines = [] lines = []
with open(hps.data.training_files, encoding="utf-8") as f: with open(hps.data.training_files, encoding="utf-8") as f:
lines.extend(f.readlines()) lines.extend(f.readlines())

View File

@@ -1,19 +1,21 @@
import argparse
import os
from pathlib import Path
import librosa
import numpy as np
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch.utils.data import Dataset from torch.utils.data import DataLoader, Dataset
from torch.utils.data import DataLoader from tqdm import tqdm
from transformers import Wav2Vec2Processor from transformers import Wav2Vec2Processor
from transformers.models.wav2vec2.modeling_wav2vec2 import ( from transformers.models.wav2vec2.modeling_wav2vec2 import (
Wav2Vec2Model, Wav2Vec2Model,
Wav2Vec2PreTrainedModel, Wav2Vec2PreTrainedModel,
) )
import librosa
import numpy as np
import argparse
from config import config
import utils import utils
import os from config import config
from tqdm import tqdm
class RegressionHead(nn.Module): class RegressionHead(nn.Module):
@@ -78,11 +80,6 @@ class AudioDataset(Dataset):
return torch.from_numpy(processed_data) return torch.from_numpy(processed_data)
model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim"
processor = Wav2Vec2Processor.from_pretrained(model_name)
model = EmotionModel.from_pretrained(model_name)
def process_func( def process_func(
x: np.ndarray, x: np.ndarray,
sampling_rate: int, sampling_rate: int,
@@ -135,16 +132,12 @@ if __name__ == "__main__":
device = config.bert_gen_config.device device = config.bert_gen_config.device
model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim" model_name = "./emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim"
processor = ( REPO_ID = "audeering/wav2vec2-large-robust-12-ft-emotion-msp-dim"
Wav2Vec2Processor.from_pretrained(model_name) if not Path(model_name).joinpath("pytorch_model.bin").exists():
if processor is None utils.download_emo_models(config.mirror, model_name, REPO_ID)
else processor
) processor = Wav2Vec2Processor.from_pretrained(model_name)
model = ( model = EmotionModel.from_pretrained(model_name).to(device)
EmotionModel.from_pretrained(model_name).to(device)
if model is None
else model.to(device)
)
lines = [] lines = []
with open(hps.data.training_files, encoding="utf-8") as f: with open(hps.data.training_files, encoding="utf-8") as f:

View File

@@ -46,3 +46,6 @@ def check_bert_models():
for k, v in models.items(): for k, v in models.items():
local_path = Path("./bert").joinpath(k) local_path = Path("./bert").joinpath(k)
_check_bert(v["repo_id"], v["files"], local_path) _check_bert(v["repo_id"], v["files"], local_path)
check_bert_models()

View File

@@ -130,7 +130,7 @@ def run():
collate_fn = TextAudioSpeakerCollate() collate_fn = TextAudioSpeakerCollate()
train_loader = DataLoader( train_loader = DataLoader(
train_dataset, train_dataset,
num_workers=config.train_ms_config.num_workers, num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1),
shuffle=False, shuffle=False,
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,

View File

@@ -16,6 +16,24 @@ MATPLOTLIB_FLAG = False
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def download_emo_models(mirror, repo_id, model_name):
if mirror == "openi":
import openi
openi.model.download_model(
"Stardust_minus/Bert-VITS2",
repo_id.split("/")[-1],
"./emotional",
)
else:
hf_hub_download(
repo_id,
"pytorch_model.bin",
local_dir=model_name,
local_dir_use_symlinks=False,
)
def download_checkpoint( def download_checkpoint(
dir_path, repo_config, token=None, regex="G_*.pth", mirror="openi" dir_path, repo_config, token=None, regex="G_*.pth", mirror="openi"
): ):