forked from sgl-project/rbg
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsgl_topo_client.py
More file actions
159 lines (127 loc) · 5.73 KB
/
Copy pathsgl_topo_client.py
File metadata and controls
159 lines (127 loc) · 5.73 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
# -*- coding: utf-8 -*-
# @Author: zibai.gj
import os
import traceback
import requests
from typing import Optional
from patio.logger import init_logger
from patio.topo.client.base_topo_client import GroupTopoClient
from patio.topo import utils
logger = init_logger(__name__)
def get_sgl_router_endpoint(worker_info: dict) -> Optional[str]:
rbg_group_name = os.getenv("GROUP_NAME")
if rbg_group_name is None:
raise Exception("GROUP_NAME is not set")
router_role_name = os.getenv("SGL_ROUTER_ROLE_NAME")
if router_role_name is None:
raise Exception("SGL_ROUTER_ROLE_NAME is not set")
sgl_router_port = os.getenv("SGL_ROUTER_PORT")
if sgl_router_port is None:
raise Exception("SGL_ROUTER_PORT is not set")
return f"{rbg_group_name}-{router_role_name}-0.s-{rbg_group_name}-{router_role_name}:{sgl_router_port}"
def get_worker_endpoint(worker_info: dict) -> Optional[str]:
port = worker_info.get("port", "8000")
worker_endpoint = os.getenv("POD_IP")
if worker_endpoint is not None:
return f"{worker_endpoint}:{port}"
# Use headless service pod domain if POD_IP is not set
rbg_group_name = os.getenv("GROUP_NAME")
if rbg_group_name is None:
raise Exception("GROUP_NAME is not set")
role_name = os.getenv("ROLE_NAME")
if role_name is None:
raise Exception("ROLE_NAME is not set")
role_index = os.getenv("ROLE_INDEX")
if role_index is None:
raise Exception("ROLE_INDEX is not set")
return f"{rbg_group_name}-{role_name}-{role_index}.s-{rbg_group_name}-{role_name}:{port}"
def get_health_check_endpoint(worker_info: dict) -> Optional[str]:
port = worker_info.get("port", "8000")
local_url = os.getenv("POD_IP")
if local_url is None:
local_url = "localhost"
return f"{local_url}:{port}"
class SGLangGroupTopoClient(GroupTopoClient):
_instance = None
# Singleton
def __new__(cls, *args, **kwargs):
if cls._instance is None:
cls._instance = super(SGLangGroupTopoClient, cls).__new__(cls)
cls._instance.__initialized = False
return cls._instance
def __init__(self, worker_info: dict):
if not self.__initialized:
self.health_check_endpoint = get_health_check_endpoint(worker_info)
self.worker_endpoint = get_worker_endpoint(worker_info)
self.sgl_router_endpoint = get_sgl_router_endpoint(worker_info)
self.worker_id = None
self.__initialized = True
def wait_engine_ready(self, worker_info: dict) -> bool:
def f():
health_check_url = f"http://{self.health_check_endpoint}/health"
resp = requests.get(health_check_url)
if resp.status_code == 200:
logger.info("Health check OK, inference engine is now ready.")
else:
raise Exception(
f"health check failed, url: {health_check_url}, status_code: {resp.status_code}, content: {resp.text}")
try:
utils.retry(f, retry_times=60, interval=3)
return True
except Exception as e:
logger.error(f"failed to check if worker engine is ready: {e}")
traceback.print_exc()
return False
def register(self, url: str, worker_info: dict, file_path: Optional[str] = None) -> bool:
"""
worker_info example:
port: 8000
worker_type: "prefill"
bootstrap_port: 34000
"""
worker_info = worker_info.copy()
url = f"http://{self.worker_endpoint}"
worker_info["url"] = url
# port is already included in self.worker_endpoint
del worker_info["port"]
def f():
worker_registration_url = f"http://{self.sgl_router_endpoint}/workers"
resp = requests.post(worker_registration_url, json=worker_info, headers={"Content-Type": "application/json"})
if resp.status_code == 202:
# Status Code 202 Accepted
self.worker_id = resp.json().get("worker_id")
if self.worker_id is None:
raise Exception(
f"register failed: missing worker_id in response body, "
f"url: {worker_registration_url}, status_code: {resp.status_code}, content: {resp.text}"
)
logger.info(f"registered worker successfully. worker_id: {self.worker_id}")
else:
raise Exception(f"register failed, url: {worker_registration_url}, status_code: {resp.status_code}, content: {resp.text}")
try:
utils.retry(f, retry_times=60, interval=3)
return True
except Exception as e:
logger.error(f"failed to register worker: {e}")
traceback.print_exc()
return False
def unregister(self):
if self.worker_id is None:
logger.warning("worker_id is not set, skipping unregister (registration may have failed)")
return False
def f():
worker_registration_url = f"http://{self.sgl_router_endpoint}/workers/{self.worker_id}"
resp = requests.delete(worker_registration_url)
if resp.status_code == 202:
# Status Code 202 Accepted
logger.info(f"unregistered worker successfully. worker_id: {self.worker_id}")
else:
raise Exception(
f"unregister failed, url: {worker_registration_url}, status_code: {resp.status_code}, content: {resp.text}")
try:
utils.retry(f, retry_times=60, interval=3)
return True
except Exception as e:
logger.error(f"failed to unregister worker: {e}")
traceback.print_exc()
return False