From edb5d05e340d70e27c5f303e1c9c8317e24883ab Mon Sep 17 00:00:00 2001 From: Isotr0py <41363108+Isotr0py@users.noreply.github.com> Date: Tue, 28 Nov 2023 12:29:08 +0800 Subject: [PATCH] Fix multiprocess (#201) * fix nultiprocess * fix format * fix multiprocess=0 --- bert_gen.py | 14 ++++++++------ emo_gen.py | 7 ++++++- text/__init__.py | 7 +++++++ 3 files changed, 21 insertions(+), 7 deletions(-) diff --git a/bert_gen.py b/bert_gen.py index f067054..588c768 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -1,12 +1,14 @@ +import argparse +from multiprocessing import Pool, cpu_count + import torch -from multiprocessing import Pool +import torch.multiprocessing as mp +from tqdm import tqdm + import commons 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 text import cleaned_text_to_sequence, get_bert def process_line(line): @@ -64,7 +66,7 @@ if __name__ == "__main__": with open(hps.data.validation_files, encoding="utf-8") as f: lines.extend(f.readlines()) if len(lines) != 0: - num_processes = args.num_processes + num_processes = min(args.num_processes, cpu_count()) with Pool(processes=num_processes) as pool: for _ in tqdm(pool.imap_unordered(process_line, lines), total=len(lines)): pass diff --git a/emo_gen.py b/emo_gen.py index fda55c6..78ba59d 100644 --- a/emo_gen.py +++ b/emo_gen.py @@ -136,7 +136,12 @@ if __name__ == "__main__": wavnames = [line.split("|")[0] for line in lines] 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(): for i, data in tqdm(enumerate(data_loader), total=len(data_loader)): diff --git a/text/__init__.py b/text/__init__.py index 45592e0..d297e9a 100644 --- a/text/__init__.py +++ b/text/__init__.py @@ -48,4 +48,11 @@ def check_bert_models(): _check_bert(v["repo_id"], v["files"], local_path) +def init_openjtalk(): + import pyopenjtalk + + pyopenjtalk.g2p("こんにちは,世界。") + + +init_openjtalk() check_bert_models()