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:
@@ -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": {
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|||||||
39
emo_gen.py
39
emo_gen.py
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
18
utils.py
18
utils.py
@@ -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"
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user