Fix multiprocess (#201)

* fix nultiprocess

* fix format

* fix multiprocess=0
This commit is contained in:
Isotr0py
2023-11-28 12:29:08 +08:00
committed by GitHub
parent e92f67fe43
commit edb5d05e34
3 changed files with 21 additions and 7 deletions

View File

@@ -1,12 +1,14 @@
import argparse
from multiprocessing import Pool, cpu_count
import torch import torch
from multiprocessing import Pool import torch.multiprocessing as mp
from tqdm import tqdm
import commons import commons
import utils import utils
from tqdm import tqdm
from text import cleaned_text_to_sequence, get_bert
import argparse
import torch.multiprocessing as mp
from config import config from config import config
from text import cleaned_text_to_sequence, get_bert
def process_line(line): def process_line(line):
@@ -64,7 +66,7 @@ if __name__ == "__main__":
with open(hps.data.validation_files, encoding="utf-8") as f: with open(hps.data.validation_files, encoding="utf-8") as f:
lines.extend(f.readlines()) lines.extend(f.readlines())
if len(lines) != 0: if len(lines) != 0:
num_processes = args.num_processes num_processes = min(args.num_processes, cpu_count())
with Pool(processes=num_processes) as pool: with Pool(processes=num_processes) as pool:
for _ in tqdm(pool.imap_unordered(process_line, lines), total=len(lines)): for _ in tqdm(pool.imap_unordered(process_line, lines), total=len(lines)):
pass pass

View File

@@ -136,7 +136,12 @@ if __name__ == "__main__":
wavnames = [line.split("|")[0] for line in lines] wavnames = [line.split("|")[0] for line in lines]
dataset = AudioDataset(wavnames, 16000, processor) dataset = AudioDataset(wavnames, 16000, processor)
data_loader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=16) data_loader = DataLoader(
dataset,
batch_size=1,
shuffle=False,
num_workers=min(args.num_processes, os.cpu_count() - 1),
)
with torch.no_grad(): with torch.no_grad():
for i, data in tqdm(enumerate(data_loader), total=len(data_loader)): for i, data in tqdm(enumerate(data_loader), total=len(data_loader)):

View File

@@ -48,4 +48,11 @@ def check_bert_models():
_check_bert(v["repo_id"], v["files"], local_path) _check_bert(v["repo_id"], v["files"], local_path)
def init_openjtalk():
import pyopenjtalk
pyopenjtalk.g2p("こんにちは,世界。")
init_openjtalk()
check_bert_models() check_bert_models()