()
| 270 | |
| 271 | |
| 272 | def main(): |
| 273 | args = parse_args() |
| 274 | validate_args(args) |
| 275 | print_config(args) |
| 276 | |
| 277 | total_lines = count_lines(args.input_file_path) |
| 278 | error_path = args.output_file_path.replace(".jsonl", "_error.jsonl") |
| 279 | skip_lines, existing_success, existing_errors = ( |
| 280 | find_resume_offset(args.output_file_path, error_path) |
| 281 | if args.resume |
| 282 | else (0, 0, 0) |
| 283 | ) |
| 284 | if skip_lines >= total_lines: |
| 285 | print(f"All {total_lines} samples are already processed.") |
| 286 | return |
| 287 | |
| 288 | if args.resume and skip_lines > 0: |
| 289 | print( |
| 290 | "Resume mode: " |
| 291 | f"{existing_success} success, {existing_errors} errors, skip {skip_lines}" |
| 292 | ) |
| 293 | |
| 294 | valid_servers = validate_servers(args) |
| 295 | print(f"Using servers: {valid_servers}") |
| 296 | |
| 297 | file_mode = "a" if args.resume and skip_lines > 0 else "w" |
| 298 | stats = { |
| 299 | "success": 0, |
| 300 | "errors": 0, |
| 301 | "context_sum": 0, |
| 302 | "context_min": None, |
| 303 | "context_max": 0, |
| 304 | } |
| 305 | queues = {server_address: [] for server_address in valid_servers} |
| 306 | next_server_index = 0 |
| 307 | submitted_count = 0 |
| 308 | |
| 309 | with ( |
| 310 | open(args.input_file_path, "r", encoding="utf-8") as input_handle, |
| 311 | open(args.output_file_path, file_mode, encoding="utf-8") as output_handle, |
| 312 | open(error_path, file_mode, encoding="utf-8") as error_handle, |
| 313 | ThreadPoolExecutor(max_workers=args.concurrency * len(valid_servers)) as executor, |
| 314 | ): |
| 315 | for _ in range(skip_lines): |
| 316 | next(input_handle, None) |
| 317 | |
| 318 | progress_total = ( |
| 319 | total_lines |
| 320 | if args.num_samples is None |
| 321 | else min(total_lines, skip_lines + args.num_samples) |
| 322 | ) |
| 323 | progress = tqdm(total=progress_total, initial=skip_lines, desc="Processing") |
| 324 | for line in input_handle: |
| 325 | if args.num_samples is not None and submitted_count >= args.num_samples: |
| 326 | break |
| 327 | |
| 328 | sample = json.loads(line) |
| 329 | server_address = valid_servers[next_server_index] |
no test coverage detected