(port: int, rank:int, world_size: int)
| 47 | raise TerminateException |
| 48 | |
| 49 | def model_worker(port: int, rank:int, world_size: int): |
| 50 | |
| 51 | def put_response(response): |
| 52 | dist.gather_object(response, object_gather_list=None, dst=0) |
| 53 | |
| 54 | def get_request() -> List: |
| 55 | _ = [[]] |
| 56 | dist.broadcast_object_list(_, src=0) |
| 57 | assert len(_) == 1 |
| 58 | _ = _[0] |
| 59 | return _ |
| 60 | |
| 61 | # specify random seed to ensure consistent token sampling among model parallel ranks |
| 62 | random.seed(0) |
| 63 | torch.random.manual_seed(0) |
| 64 | np.random.seed(0) |
| 65 | |
| 66 | store = dist.TCPStore("127.0.0.1", port, world_size, False) |
| 67 | |
| 68 | dist.init_process_group( |
| 69 | backend="gloo", rank=rank, world_size=world_size, |
| 70 | # init_method=f"tcp://127.0.0.1:{port}", |
| 71 | store=store |
| 72 | ) |
| 73 | |
| 74 | size = dist.get_world_size() |
| 75 | rank = dist.get_rank() |
| 76 | |
| 77 | gpu_ids, from_pretrained_args, from_pretrained_kwargs = get_request() |
| 78 | torch.cuda.set_device(gpu_ids[rank-1]) |
| 79 | |
| 80 | init_print(rank==1) |
| 81 | from_pretrained_kwargs['mp_group'] = dist.new_group(ranks=list(range(1, size)), backend="nccl") |
| 82 | # mp_group identifies which ranks will work collaboratively through model parallelism |
| 83 | model = MetaModel.from_pretrained(*from_pretrained_args, **from_pretrained_kwargs) |
| 84 | |
| 85 | dist.barrier() |
| 86 | |
| 87 | |
| 88 | while True: |
| 89 | try: |
| 90 | request_type, (request_args, request_kwargs) = get_request() |
| 91 | process_special_request(request_type) |
| 92 | |
| 93 | if request_type not in REQUESTS_WITH_STREAM_RESPONSE: |
| 94 | result = getattr(model, request_type)(*request_args, **request_kwargs) |
| 95 | put_response(("SUCCESS", result)) |
| 96 | else: |
| 97 | for stream_response in getattr(model, request_type)(*request_args, **request_kwargs): |
| 98 | put_response(("YIELDING", stream_response)) |
| 99 | yield_request, _ = get_request() |
| 100 | process_special_request(request_type) |
| 101 | if yield_request == "continue_yield": |
| 102 | continue |
| 103 | elif yield_request == "stop_yield": |
| 104 | raise ResetException |
| 105 | else: |
| 106 | raise ValueError(f"Unexpected request type during yield: {yield_request}") |
no test coverage detected