MCPcopy Create free account
hub / github.com/AlayaLab/Hive / main

Function main

pipeline/code/06_superres_apollo.py:207–276  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

205 return results
206
207def main():
208 parser = argparse.ArgumentParser()
209 parser.add_argument("--input_path", type=str, required=True, help="Input JSON file path")
210 parser.add_argument("--output_path", type=str, required=True, help="Output JSON file path")
211 parser.add_argument("--output_audio_dir", type=str, required=True, help="Output audio directory")
212 parser.add_argument("--batch_size", type=int, default=1, help="Number of audio files to process in parallel")
213 args = parser.parse_args()
214
215 input_path = args.input_path
216 output_path = args.output_path
217 output_audio_dir = args.output_audio_dir
218 batch_size = args.batch_size
219 original_batch_size = batch_size
220
221 os.makedirs(output_audio_dir, exist_ok=True)
222
223 apollo_uni_config = get_config('configs/config_apollo_uni.yaml')
224 model = look2hear.models.BaseModel.from_pretrain('my_weights/apollo_model_uni.ckpt', **apollo_uni_config['model']).to(device)
225
226 if torch.cuda.device_count() > 1:
227 model = torch.nn.DataParallel(model)
228 else:
229 logger.info(f"Using single GPU: {device}")
230
231 with open(input_path, "r") as f:
232 data = json.load(f)
233
234 results = []
235 processed_count = 0
236
237 i = 0
238 pbar = tqdm(total=len(data), desc="Batch processing audio files")
239
240 while i < len(data):
241 batch_data = data[i:i + batch_size]
242 batch_results = process_data_batch(batch_data, model, output_audio_dir)
243
244 if batch_results == "CUDA_OOM":
245 if batch_size == 1:
246 logger.error(f"\nError: CUDA out of memory even with batch_size=1, cannot continue")
247 logger.error(f"Please try smaller audio files or GPU with more memory")
248 return
249
250 batch_size = max(1, batch_size // 2)
251 logger.warning(f"\nCUDA out of memory detected, reducing batch_size from {original_batch_size} to {batch_size}, retrying...")
252 original_batch_size = batch_size
253 continue
254
255 results.extend(batch_results)
256 i += len(batch_data)
257 processed_count += len(batch_data)
258 pbar.update(len(batch_data))
259
260 if processed_count > 0 and processed_count % 1024 == 0:
261 if batch_size < args.batch_size:
262 new_batch_size = min(batch_size * 2, args.batch_size)
263 logger.info(f"\nProcessed {processed_count} files, attempting to increase batch_size from {batch_size} to {new_batch_size}...")
264 batch_size = new_batch_size

Callers 1

Calls 5

get_configFunction · 0.85
process_data_batchFunction · 0.85
dumpMethod · 0.80
toMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected