96 lines
4.8 KiB
Python
96 lines
4.8 KiB
Python
"""sparkle OpenAI 兼容 server 入口:python -m mlx_streaming.server
|
||
|
||
模型路径参数与 sparkle chat 一致(默认值同样来自 mlx_streaming/config.py 的环境变量)。
|
||
SPARKLE_FAKE=1 时用 FakeBackend,免模型开发/测试。
|
||
"""
|
||
import argparse
|
||
import os
|
||
|
||
from mlx_streaming import config
|
||
|
||
|
||
def _build_parser():
|
||
p = argparse.ArgumentParser(
|
||
prog="sparkle-server",
|
||
description="sparkle:OpenAI 兼容的流式 MoE 推理 server")
|
||
p.add_argument("--host", default="127.0.0.1", help="监听地址(默认 127.0.0.1)")
|
||
p.add_argument("--port", type=int, default=8317, help="监听端口(默认 8317)")
|
||
p.add_argument("--model", default=config.model_path(), help="主模型路径(MLX 量化)")
|
||
p.add_argument("--expert-dir", default=config.expert_dir(),
|
||
help="拆分后的 per-expert 目录")
|
||
p.add_argument("--mtp-out", default=config.mtp_out(), help="MTP 权重文件")
|
||
p.add_argument("--qn-config", default=config.qn_config(),
|
||
help="Qwen3-Next 配置 JSON")
|
||
p.add_argument("-k", "--k", type=int, default=3, help="MTP 投机宽度(默认 3)")
|
||
p.add_argument("-n", "--max-tokens", type=int, default=4096,
|
||
help="每轮最多生成的新 token 数(默认 4096)")
|
||
p.add_argument("--expert-slots", type=int, default=32,
|
||
help="常驻专家池容量(默认 32,同时作为侧区行数默认)")
|
||
p.add_argument("--spec-slots", type=int, default=None,
|
||
help="侧区行数 POOL_SPEC_SLOTS(默认跟随 --expert-slots)")
|
||
return p
|
||
|
||
|
||
def _safe_prefill_chunk(args) -> int:
|
||
"""按专家池实际可安放容量算 prefill 步长上限,取 min(启动方给的值, 上限)。
|
||
|
||
一次前向路由的唯一专家最多 chunk × top_k 个,真实区能安放的只有
|
||
expert_slots − AUTOPIN 钉死数(钉死的永不驱逐)。超出会被 block.py 分流到
|
||
host/fetch 路径:结果仍正确,但每层每 chunk 重读盘,prefill 慢几倍。
|
||
读不到 top_k 时按 Qwen3-Next 的 10 兜底(偏保守,宁可 chunk 小)。
|
||
"""
|
||
import json
|
||
|
||
top_k = 10
|
||
try:
|
||
with open(args.qn_config, encoding="utf-8") as f:
|
||
top_k = max(1, int(json.load(f).get("num_experts_per_tok", top_k)))
|
||
except Exception: # noqa: BLE001 配置读不到/字段缺失都退到保守默认
|
||
pass
|
||
slots = max(1, int(args.expert_slots))
|
||
pinned = int(slots * config.autopin_budget_frac()) if config.autopin() else 0
|
||
limit = max(1, (slots - pinned) // top_k)
|
||
want = int(os.environ.get("PREFILL_CHUNK") or limit)
|
||
if want > limit:
|
||
print(f"[PREFILL_CHUNK] {want} → {limit}:{slots} 槽扣掉 AUTOPIN 钉死 {pinned} 后只能安放 "
|
||
f"{slots - pinned} 个专家,chunk×top_k({top_k}) 不得超过它,否则 prefill 退到读盘慢路径。",
|
||
flush=True)
|
||
return min(want, limit)
|
||
|
||
|
||
def main(argv=None):
|
||
args = _build_parser().parse_args(argv)
|
||
|
||
# KV_QUANT 维持 Sparkle 原始配方(开):此前"不按规则回答"已证实是 tools 未渲染所致,
|
||
# 与 KV 量化无关(且其通过质量门:token 一致率≥95%、logits cosine≥0.99)。
|
||
# KV_QUANT=0 会使 prefill 慢 ~3 倍(fp16 KV 带宽),业务长提示下代价不可接受。
|
||
# 显式 export KV_QUANT=0 可关闭做质量对照。
|
||
|
||
# 前缀快照头部的**上限**(实际长度由引擎按 system+tools 的真实边界自动取,见
|
||
# MLXBackend._fixed_prefix_len)。写死数字曾是主要的慢因:业务 system+4 工具实测
|
||
# 3972 token,固定切 2048 只覆盖一半,每个新会话都要白 prefill 1900+ token
|
||
# (prefill 约 30 tok/s,即多等一分钟)。这里只兜住极端情况下快照过大。
|
||
os.environ.setdefault("PREFIX_SNAPSHOT_HEAD", "8192")
|
||
|
||
# AUTOPIN 用一半槽位钉历史热专家(冷启动不慢热),另一半留给 demand 动态换入。
|
||
os.environ.setdefault("AUTOPIN_BUDGET_FRAC", "0.5")
|
||
|
||
# PREFILL_CHUNK 必须由槽位推导,不能让启动方拍数字:每个 chunk 的一次前向最多路由
|
||
# chunk × top_k 个唯一专家,而真实区能安放的只有 expert_slots − AUTOPIN 钉死数。超了
|
||
# 会被分流到 host/fetch 路径(正确但每层重读盘,慢几倍),槽位调小后写死的 chunk 就失效。
|
||
os.environ["PREFILL_CHUNK"] = str(_safe_prefill_chunk(args))
|
||
|
||
import uvicorn
|
||
|
||
from mlx_streaming.server.app import create_app
|
||
|
||
# timeout_graceful_shutdown:SIGTERM 后最多等 3 秒就强制关闭连接。
|
||
# 外部客户端的轮询持有 keep-alive 长连接,不设超时 uvicorn 会永远
|
||
# 卡在 "Waiting for connections to close",进程不退、内存不释放。
|
||
uvicorn.run(create_app(args), host=args.host, port=args.port,
|
||
timeout_graceful_shutdown=3)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|