Skip to content
This repository was archived by the owner on Nov 11, 2023. It is now read-only.

Commit 7c7bb11

Browse files
committed
fix(preprocess): pass device
1 parent 28497ff commit 7c7bb11

1 file changed

Lines changed: 16 additions & 11 deletions

File tree

preprocess_hubert_f0.py

Lines changed: 16 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -103,7 +103,7 @@ def process_one(filename, hmodel,f0p,rank,diff=False,mel_extractor=None):
103103
if not os.path.exists(aug_vol_path):
104104
np.save(aug_vol_path,aug_vol.to('cpu').numpy())
105105

106-
def process_batch(file_chunk, f0p, diff=False, mel_extractor=None):
106+
def process_batch(file_chunk, f0p, diff=False, mel_extractor=None, device="cpu"):
107107
logger.info("Loading speech encoder for content...")
108108
rank = mp.current_process()._identity
109109
rank = rank[0] if len(rank) > 0 else 0
@@ -116,19 +116,20 @@ def process_batch(file_chunk, f0p, diff=False, mel_extractor=None):
116116
for filename in tqdm(file_chunk):
117117
process_one(filename, hmodel, f0p, gpu_id, diff, mel_extractor)
118118

119-
def parallel_process(filenames, num_processes, f0p, diff, mel_extractor):
119+
def parallel_process(filenames, num_processes, f0p, diff, mel_extractor, device):
120120
with ProcessPoolExecutor(max_workers=num_processes) as executor:
121121
tasks = []
122122
for i in range(num_processes):
123123
start = int(i * len(filenames) / num_processes)
124124
end = int((i + 1) * len(filenames) / num_processes)
125125
file_chunk = filenames[start:end]
126-
tasks.append(executor.submit(process_batch, file_chunk, f0p, diff, mel_extractor))
126+
tasks.append(executor.submit(process_batch, file_chunk, f0p, diff, mel_extractor, device=device))
127127
for task in tqdm(tasks):
128128
task.result()
129129

130130
if __name__ == "__main__":
131131
parser = argparse.ArgumentParser()
132+
parser.add_argument('-d', '--device', type=str, default=None)
132133
parser.add_argument(
133134
"--in_dir", type=str, default="dataset/44k", help="path to input dir"
134135
)
@@ -143,15 +144,19 @@ def parallel_process(filenames, num_processes, f0p, diff, mel_extractor):
143144
)
144145
args = parser.parse_args()
145146
f0p = args.f0_predictor
146-
print(speech_encoder)
147-
logger.info("Using " + speech_encoder + " SpeechEncoder")
148-
logger.info("Using " + f0p + "f0 extractor")
149-
logger.info("Using diff Mode:")
150-
print(args.use_diff)
147+
device = args.device
148+
if device is None:
149+
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
150+
151+
logger.info("Using device: " + device)
152+
logger.info("Using SpeechEncoder: " + speech_encoder)
153+
logger.info("Using extractor:" + f0p)
154+
logger.info("Using diff Mode: " + args.use_diff)
155+
151156
if args.use_diff:
152157
print("use_diff")
153158
print("Loading Mel Extractor...")
154-
mel_extractor = Vocoder(dconfig.vocoder.type, dconfig.vocoder.ckpt, device = "cuda:0")
159+
mel_extractor = Vocoder(dconfig.vocoder.type, dconfig.vocoder.ckpt, device=device)
155160
print("Loaded Mel Extractor.")
156161
else:
157162
mel_extractor = None
@@ -162,5 +167,5 @@ def parallel_process(filenames, num_processes, f0p, diff, mel_extractor):
162167
num_processes = args.num_processes
163168
if num_processes == 0:
164169
num_processes = os.cpu_count()
165-
166-
parallel_process(filenames, num_processes, f0p, args.use_diff, mel_extractor)
170+
171+
parallel_process(filenames, num_processes, f0p, args.use_diff, mel_extractor, device)

0 commit comments

Comments
 (0)