sparkle/mlx_streaming/server/state.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

102 lines
4.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""引擎单例管理:加载状态、生成串行锁、运行期调参与后台重建。"""
from __future__ import annotations
import asyncio
import resource
import sys
from mlx_streaming.tui.backend import FakeBackend
def rss_bytes() -> int:
"""当前进程常驻内存峰值(字节)。macOS 的 ru_maxrss 直接是字节;Linux 是 KB。"""
ru = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
return int(ru if sys.platform == "darwin" else ru * 1024)
class EngineManager:
"""持有一个 ChatBackend 单例;引擎是单对话的,同一时刻只允许一轮生成。
asyncio.Lock 做排队:第二个并发请求等锁而非直接失败。k / max_tokens / expert_slots
是运行期可调参数,由 /api/engine/config 修改;k 与 max_tokens 即时生效(作为后续
OpenAI 请求的缺省值),expert_slots 变化触发后台重建。
"""
def __init__(self, args, backend):
self.args = args # argparse Namespace 风格,重建 MLXBackend 时复用
self.backend = backend
self.loaded = False
self.reloading = False
self.load_error: str | None = None
self.lock = asyncio.Lock()
self.last_tok_per_s = 0.0
self.k = getattr(args, "k", 3)
self.max_tokens = getattr(args, "max_tokens", 4096)
self.expert_slots = getattr(args, "expert_slots", 32)
def load_initial(self, on_status=lambda m: None) -> None:
"""启动加载(放线程里跑,真模型以分钟计);就绪前 loaded=False。"""
try:
self.backend.load(on_status)
self.loaded = True
except Exception as e: # noqa: BLE001 加载失败不拖垮 server,状态暴露在 /api/stats
self.load_error = f"{type(e).__name__}: {e}"
def rebuild(self, on_status=lambda m: None) -> None:
"""后台重建引擎(fake 模式只换实例;真后端按更新后的 args 重新装配 + 预热)。"""
self.reloading = True
self.loaded = False
try:
self.args.expert_slots = self.expert_slots
old = self.backend
if isinstance(old, FakeBackend):
self.backend = FakeBackend(reply=old.reply)
else:
from mlx_streaming.tui.backend import MLXBackend
self.backend = MLXBackend(self.args)
self.backend.load(on_status)
self.loaded = True
except Exception as e: # noqa: BLE001
self.load_error = f"{type(e).__name__}: {e}"
finally:
self.reloading = False
def _resident_pool(self):
"""拿常驻专家池(fake 后端 / 非双源路径下为 None)。"""
store = getattr(getattr(self.backend, "_model", None), "_expert_store", None)
return getattr(store, "_resident", None)
def unplaced_experts(self) -> int:
"""专家池累计「分不到槽落 0 号槽」数(见 ResidentExpertPool.unplaced)。
取不到(fake 后端 / 非双源路径)返回 0。这是数值正确性的看门指标:非 0 就说明
某些前向是拿错专家权重算的,回答不可信。
"""
return int(getattr(self._resident_pool(), "unplaced", 0) or 0)
def pool_metrics(self) -> dict:
"""专家池命中/读盘指标:定位「慢」到底是池不够大还是走了慢路径。
decode 每步要取 层数×top_k 个专家miss 一个就得读一份专家权重(8bit ≈3.3MB)
所以 hit_rate 直接决定速度。gpu_fallback 记的是本该零 host 往返的快路径退回读盘的次数。
"""
rp = self._resident_pool()
if rp is None:
return {}
hits, misses = int(getattr(rp, "hits", 0)), int(getattr(rp, "misses", 0))
return {"hits": hits, "misses": misses,
"hit_rate": round(hits / (hits + misses), 4) if hits + misses else 0.0,
"gpu_fastpath": int(getattr(rp, "gpu_fastpath", 0)),
"gpu_fallback": int(getattr(rp, "gpu_fallback", 0)),
"prefetch_hits": int(getattr(rp, "prefetch_hits", 0)),
"prefetch_loads": int(getattr(rp, "prefetch_loads", 0))}
def apply_gen_params(self, req_max_tokens=None) -> int:
"""把当前调参写入后端 args(单飞前提下安全),返回本轮有效 max_tokens。"""
eff = req_max_tokens or self.max_tokens
if hasattr(self.backend, "args"):
self.backend.args.k = self.k
self.backend.args.max_tokens = eff
return eff