| 286 | |
| 287 | |
| 288 | def run_inference(current_rank, args): |
| 289 | # random sleep 0-3s: |
| 290 | time.sleep(current_rank * 2 + random.random()) |
| 291 | |
| 292 | args.current_rank = current_rank |
| 293 | |
| 294 | gpu_list = args.gpus.strip().split(',') |
| 295 | current_gpu = gpu_list[current_rank] |
| 296 | |
| 297 | dir_path = args.output_json_file_path = args.inference_model_path + '/intermediate_predictions' |
| 298 | print("mkdir -p " + dir_path) |
| 299 | os.system("mkdir -p " + dir_path) |
| 300 | |
| 301 | args.output_json_file_path = args.inference_model_path + '/intermediate_predictions/round{}_rank{}'.format(args.round, current_rank) + '.json' |
| 302 | command = [ |
| 303 | "deepspeed", |
| 304 | f"--include=localhost:{current_gpu}", |
| 305 | f"--master_port {args.master_port + current_rank}", |
| 306 | "inference/infer_part.py", |
| 307 | "--deepspeed" |
| 308 | ] |
| 309 | # add other args: |
| 310 | for key, value in vars(args).items(): |
| 311 | if key not in ["master_port", "gpus"]: |
| 312 | command.append(f"--{key} {value}") |
| 313 | |
| 314 | # print red command: |
| 315 | print("\033[31m" + " ".join(command) + "\033[0m") |
| 316 | |
| 317 | print('') |
| 318 | # result = subprocess.run(command, shell=True, capture_output=True, text=True) |
| 319 | # os.system(" ".join(command)) |
| 320 | |
| 321 | # os.popen(" ".join(command)).readlines() |
| 322 | |
| 323 | fail_num = 0 |
| 324 | result = None |
| 325 | |
| 326 | if fail_num < 3: |
| 327 | try: |
| 328 | # only need last line: |
| 329 | _ = os.popen(" ".join(command)).readlines()[-2] |
| 330 | # read str in args.output_json_file_path: |
| 331 | with open(args.output_json_file_path, "r") as file: |
| 332 | result = file.read() |
| 333 | except Exception as e: |
| 334 | fail_num += 1 |
| 335 | print(f"Fail for {fail_num} times. Error: {e}") |
| 336 | result = None |
| 337 | |
| 338 | # # print red result: |
| 339 | # print("\033[31m" + str(eval(result.strip())) + "\033[0m") |
| 340 | |
| 341 | # # result = subprocess.run(command, capture_output=True, text=True) |
| 342 | # result = "".join(os.popen(" ".join(command)).readlines()) |
| 343 | return eval(result.strip()) |
| 344 | |
| 345 | |