()
| 205 | return results |
| 206 | |
| 207 | def 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 |
no test coverage detected