Change to use thread pool executor for multiprocess
This commit is contained in:
18
bert_gen.py
18
bert_gen.py
@@ -1,5 +1,5 @@
|
|||||||
import argparse
|
import argparse
|
||||||
from multiprocessing import Pool
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
@@ -73,13 +73,13 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
if len(lines) != 0:
|
if len(lines) != 0:
|
||||||
num_processes = args.num_processes
|
num_processes = args.num_processes
|
||||||
with Pool(processes=num_processes) as pool:
|
with ThreadPoolExecutor(max_workers=num_processes) as executor:
|
||||||
for _ in tqdm(
|
_ = list(
|
||||||
pool.imap_unordered(process_line, zip(lines, add_blank)),
|
tqdm(
|
||||||
total=len(lines),
|
executor.map(process_line, zip(lines, add_blank)),
|
||||||
file=SAFE_STDOUT,
|
total=len(lines),
|
||||||
):
|
file=SAFE_STDOUT,
|
||||||
# 这里是缩进的代码块,表示循环体
|
)
|
||||||
pass # 使用pass语句作为占位符
|
)
|
||||||
|
|
||||||
logger.info(f"bert.pt is generated! total: {len(lines)} bert.pt files.")
|
logger.info(f"bert.pt is generated! total: {len(lines)} bert.pt files.")
|
||||||
|
|||||||
@@ -42,7 +42,7 @@
|
|||||||
"%cd Style-Bert-VITS2/\n",
|
"%cd Style-Bert-VITS2/\n",
|
||||||
"!pip install -r requirements.txt\n",
|
"!pip install -r requirements.txt\n",
|
||||||
"!apt install libcublas11\n",
|
"!apt install libcublas11\n",
|
||||||
"!python initialize.py"
|
"!python initialize.py --skip_jvnv"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
25
resample.py
25
resample.py
@@ -1,6 +1,6 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
from multiprocessing import Pool, cpu_count
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
import librosa
|
import librosa
|
||||||
import pyloudnorm as pyln
|
import pyloudnorm as pyln
|
||||||
@@ -110,13 +110,22 @@ if __name__ == "__main__":
|
|||||||
logger.error(f"No wav files found in {args.in_dir}")
|
logger.error(f"No wav files found in {args.in_dir}")
|
||||||
raise ValueError(f"No wav files found in {args.in_dir}")
|
raise ValueError(f"No wav files found in {args.in_dir}")
|
||||||
|
|
||||||
pool = Pool(processes=processes)
|
# pool = Pool(processes=processes)
|
||||||
for _ in tqdm(
|
# for _ in tqdm(
|
||||||
pool.imap_unordered(process, tasks), file=SAFE_STDOUT, total=len(tasks)
|
# pool.imap_unordered(process, tasks), file=SAFE_STDOUT, total=len(tasks)
|
||||||
):
|
# ):
|
||||||
pass
|
# pass
|
||||||
|
|
||||||
pool.close()
|
# pool.close()
|
||||||
pool.join()
|
# pool.join()
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=processes) as executor:
|
||||||
|
_ = list(
|
||||||
|
tqdm(
|
||||||
|
executor.map(process, tasks),
|
||||||
|
total=len(tasks),
|
||||||
|
file=SAFE_STDOUT,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
logger.info("Resampling Done!")
|
logger.info("Resampling Done!")
|
||||||
|
|||||||
Reference in New Issue
Block a user