@@ -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
130130if __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