MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / model_worker

Function model_worker

accessory/model/multi_gpu_wrapper.py:49–116  ·  view source on GitHub ↗
(port: int, rank:int, world_size: int)

Source from the content-addressed store, hash-verified

47 raise TerminateException
48
49def 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}")

Callers 1

Calls 5

get_requestFunction · 0.85
init_printFunction · 0.85
process_special_requestFunction · 0.85
put_responseFunction · 0.85
from_pretrainedMethod · 0.80

Tested by

no test coverage detected