Repository navigation
Expand file tree
/
Copy pathapi_server.py
More file actions
237 lines (212 loc) · 10.4 KB
/
Copy pathapi_server.py
File metadata and controls
237 lines (212 loc) · 10.4 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
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
"""OpenAI 兼容的 HTTP API 服务器(基于标准库 http.server)。
单卡: python api_server.py --model /path/to/model --port 8000
多卡: torchrun --nproc_per_node=2 api_server.py --model /path/to/model --tp-size 2
"""
import os, argparse, json, time, uuid, threading
os.environ.setdefault("TILELANG_CACHE_DIR",
os.path.join(os.path.dirname(os.path.abspath(__file__)), ".tilelang_cache"))
from dataclasses import asdict
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import torch, torch.distributed as dist
from engine.llm import LLM
from engine.sampling_params import SamplingParams
from engine.parallel import get_tp_rank, is_tp_active
llm: LLM = None # 全局引擎
MODEL_NAME = "" # 模型名(用于响应)
infer_lock = threading.Lock() # 串行化推理(TP 下必须)
def build_prompt(messages, response_format=None):
"""用 tokenizer 的 chat template 拼接 prompt。"""
prompt = llm.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
if response_format and isinstance(response_format, dict) \
and response_format.get("type") == "json_object":
prompt += "\n请以JSON格式输出。" # 结构化输出指令
return prompt
def make_sampling_params(data):
"""从请求体构造 SamplingParams。"""
stop = data.get("stop")
stop = [stop] if isinstance(stop, str) else (stop or [])
# OpenAI 的 response_format: {"type": "json_object"} 映射为 JSON 约束解码
guided_json = data.get("guided_json")
if guided_json is None and isinstance(data.get("response_format"), dict) \
and data["response_format"].get("type") == "json_object":
guided_json = True
return SamplingParams(
temperature=float(data.get("temperature", 0.6)),
top_p=float(data.get("top_p", 1.0)),
top_k=int(data.get("top_k", -1)),
max_tokens=int(data.get("max_tokens", 256)),
repetition_penalty=float(data.get("repetition_penalty", 1.0)),
stop=stop,
logit_bias=data.get("logit_bias") or {},
allowed_token_ids=data.get("allowed_token_ids") or [],
bad_token_ids=data.get("bad_token_ids") or [],
guided_choice=data.get("guided_choice") or [],
guided_json=guided_json,
guided_regex=data.get("guided_regex"),
)
def _broadcast_bytes(data: bytes):
"""广播一段字节流(先长度后内容)。"""
length = torch.tensor([len(data)], dtype=torch.long, device="cuda")
dist.broadcast(length, src=0)
if data:
buf = torch.frombuffer(bytearray(data), dtype=torch.uint8).to("cuda")
dist.broadcast(buf, src=0)
def _recv_body(n: int) -> bytes:
"""接收已知长度为 n 的广播内容(n<=0 返回空)。"""
if n <= 0:
return b""
buf = torch.empty(n, dtype=torch.uint8, device="cuda")
dist.broadcast(buf, src=0)
return buf.cpu().numpy().tobytes()
def _recv_bytes() -> bytes:
length = torch.tensor([0], dtype=torch.long, device="cuda")
dist.broadcast(length, src=0)
return _recv_body(int(length.item()))
def _params_to_bytes(params) -> bytes:
"""完整序列化 SamplingParams(含 stop / 约束解码参数),保证 TP 各 rank 一致。"""
return json.dumps(asdict(params), ensure_ascii=False).encode("utf-8")
def _params_from_bytes(data: bytes) -> SamplingParams:
d = json.loads(data.decode("utf-8"))
# JSON 会把 logit_bias 的 int 键变成字符串,恢复为 token id
d["logit_bias"] = {int(k): v for k, v in (d.get("logit_bias") or {}).items()}
return SamplingParams(**d)
def broadcast_work(prompt, params):
"""rank 0 广播 prompt 与完整采样参数,唤醒 worker 参与本轮推理。"""
_broadcast_bytes(prompt.encode("utf-8"))
_broadcast_bytes(_params_to_bytes(params))
def worker_loop():
"""非 rank 0 的 worker:循环接收广播并参与推理。"""
while True:
length = torch.tensor([0], dtype=torch.long, device="cuda")
dist.broadcast(length, src=0)
if int(length.item()) == -1:
break # 退出信号
# prompt 的长度已在上面接收(用于区分退出信号),只补收内容
prompt = _recv_body(int(length.item())).decode("utf-8")
params = _params_from_bytes(_recv_bytes())
llm.generate([prompt], params)
def run_inference(prompt, params):
"""非流式推理。"""
with infer_lock:
if is_tp_active():
broadcast_work(prompt, params)
return llm.generate([prompt], params)[0]
def run_inference_stream(prompt, params):
"""流式推理(生成器)。"""
with infer_lock:
if is_tp_active():
broadcast_work(prompt, params)
for _i, text, _f in llm.generate_stream([prompt], params):
yield text
class Handler(BaseHTTPRequestHandler):
def _send_json(self, code, obj):
body = json.dumps(obj, ensure_ascii=False).encode("utf-8")
self.send_response(code)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def _err(self, code, message):
self._send_json(code, {"error": {"message": message, "type": "invalid_request_error"}})
def do_GET(self):
if self.path == "/v1/models":
self._send_json(200, {"object": "list", "data": [
{"id": MODEL_NAME, "object": "model", "created": int(time.time())}]})
else:
self._err(404, f"Not found: {self.path}")
def do_POST(self):
if self.path != "/v1/chat/completions":
self._err(404, f"Not found: {self.path}")
return
try:
data = json.loads(self.rfile.read(int(self.headers.get("Content-Length", 0))).decode("utf-8"))
except Exception as e:
self._err(400, f"Invalid JSON: {e}")
return
if not data.get("messages"):
self._err(400, "messages is required")
return
try:
prompt = build_prompt(data["messages"], data.get("response_format"))
params = make_sampling_params(data)
except Exception as e:
self._err(400, f"Invalid params: {e}")
return
rid, created = f"chatcmpl-{uuid.uuid4().hex[:24]}", int(time.time())
if data.get("stream", False):
self._stream(rid, created, prompt, params)
else:
self._nonstream(rid, created, prompt, params)
def _nonstream(self, rid, created, prompt, params):
try:
out = run_inference(prompt, params)
except Exception as e:
self._err(500, f"Internal error: {e}")
return
pt, ct = len(out["prompt_token_ids"]), len(out["token_ids"])
self._send_json(200, {
"id": rid, "object": "chat.completion", "created": created, "model": MODEL_NAME,
"choices": [{"index": 0, "message": {"role": "assistant", "content": out["text"]},
"finish_reason": "stop"}],
"usage": {"prompt_tokens": pt, "completion_tokens": ct, "total_tokens": pt + ct}})
def _stream(self, rid, created, prompt, params):
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.send_header("Cache-Control", "no-cache")
self.end_headers()
def send(obj):
self.wfile.write(f"data: {json.dumps(obj, ensure_ascii=False)}\n\n".encode("utf-8"))
self.wfile.flush()
try:
for text in run_inference_stream(prompt, params):
send({"id": rid, "object": "chat.completion.chunk", "created": created,
"model": MODEL_NAME,
"choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}]})
send({"id": rid, "object": "chat.completion.chunk", "created": created,
"model": MODEL_NAME,
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]})
self.wfile.write(b"data: [DONE]\n\n")
self.wfile.flush()
except Exception as e:
send({"error": {"message": str(e), "type": "internal_error"}})
# OpenAI 兼容客户端依赖 [DONE] 判断流结束,异常路径也要发送
self.wfile.write(b"data: [DONE]\n\n")
self.wfile.flush()
def log_message(self, *a):
pass # 静默默认日志
def main():
global llm, MODEL_NAME
ap = argparse.ArgumentParser(description="OpenAI 兼容 API 服务器")
ap.add_argument("--host", default="0.0.0.0")
ap.add_argument("--port", type=int, default=8000)
ap.add_argument("--model", required=True, help="模型路径")
ap.add_argument("--tp-size", type=int, default=1, help="张量并行大小")
ap.add_argument("--max-num-seqs", type=int, default=64)
ap.add_argument("--max-model-len", type=int, default=4096)
ap.add_argument("--gpu-memory-utilization", type=float, default=0.9)
ap.add_argument("--enforce-eager", action="store_true")
ap.add_argument("--num-scheduler-steps", type=int, default=1,
help="多步调度:每次调度执行的解码步数(1 = 关闭)")
ap.add_argument("--no-prefix-caching", action="store_true", help="禁用前缀缓存")
args = ap.parse_args()
MODEL_NAME = os.path.basename(os.path.normpath(args.model)) or args.model
print(f"初始化引擎: model={args.model}, tp_size={args.tp_size}")
llm = LLM(model=args.model, max_num_seqs=args.max_num_seqs, max_model_len=args.max_model_len,
gpu_memory_utilization=args.gpu_memory_utilization, enforce_eager=args.enforce_eager,
tp_size=args.tp_size, num_scheduler_steps=args.num_scheduler_steps,
enable_prefix_caching=not args.no_prefix_caching)
# TP>1 时:rank 0 启动 HTTP 服务器,其余 rank 进入 worker 循环参与推理
if args.tp_size > 1 and get_tp_rank() != 0:
print(f"[worker] rank {get_tp_rank()} 进入推理循环")
worker_loop()
else:
server = ThreadingHTTPServer((args.host, args.port), Handler)
print(f"API 服务器: http://{args.host}:{args.port} (rank {get_tp_rank()})")
try:
server.serve_forever()
except KeyboardInterrupt:
pass
if is_tp_active(): # 通知 worker 退出
dist.broadcast(torch.tensor([-1], dtype=torch.long, device="cuda"), src=0)
if __name__ == "__main__":
main()