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

330 lines
14 KiB
Python

"""sparkle 命令行:面向用户的交互式多轮对话(MTP 自投机快路径)。
用法示例:
sparkle # 直接进入交互式对话(默认子命令 chat)
sparkle chat # 同上
sparkle -k 4 -n 800 --stats # 调宽投机、加长生成、每轮打印吞吐
sparkle --system "你是一个简洁的助手"
sparkle --model models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit --expert-slots 32
只做「生成」一件事:走 MTP 自投机 + 零拷贝双源侧区快路径。关键参数做成命令行 flag,
其余调优项仍从环境变量读取(见 mlx_streaming/config.py)。
交互期间可用命令:
/exit 或 /quit 退出
/reset 清空对话历史(保留 system)
/help 打印帮助
"""
import argparse
import sys
import time
from mlx_streaming import config
# MTP 快路径环境变量兜底配方(benchmark 验证过的最优组合)。
# 用 setdefault 兜底:用户显式导出的环境变量优先级更高,不会被覆盖。
# MTP_ADAPTIVE_DEPTH:置信度门控动态深度。逐位累计置信度跌破 tau 即停,低置信步抽浅省专家加载
# (本系统 IO 瓶颈)。消融(reports/adaptive-depth-2026-07-05)证 τ=0.3、depth_max=3 纯向下收缩
# +5~6% tok/s 且 bit-lossless、零额外显存,稳定优于最小树 pos0 救回 → 设为用户主路径默认。
# depth_max=3 与基础 K 一致,在生产 EXPERT_SLOTS=32 下 seq·top_k 不溢出 cap(扩到 4 须 slots>=40)。
# 注:动态深度与 TREE_TOP2 互斥(adaptive 仅在非 tree 的 plain 路径生效),故此处不开 TREE_TOP2。
# KV_QUANT:IsoQuant K4/V3 + SO(4) 块旋转,仅作用于 12 个全注意力层(线性层递归态不动)。
# 极致压缩长上下文 KV:128k 3.0→~0.68 GiB;短会话收益小但无害。K4/V3/旋转均取默认值,
# 开 KV_QUANT=1 即整套生效。质量验收:token 一致率≥95% + logits cosine≥0.99。
_FASTPATH_ENV = {
"STREAM_BLOB_LOADER": "1",
"NATIVE_FUSED_PREFETCH": "1",
"ZEROCOPY_DUAL_SOURCE": "1",
"SIDEREGION_LFU": "1",
"KV_QUANT": "1",
"MTP_ADAPTIVE_DEPTH": "1",
"MTP_CONF_TAU": "0.3",
"MTP_DEPTH_MAX": "3",
}
def _build_engine(args, on_status=None):
"""按 MTP 快路径装配 model / tokenizer / drafter。
注意:model_builder 在 import 时就读 MODEL/EXPERT_DIR 等环境变量,所以必须
先把命令行参数写进 os.environ,再 import build_streaming_model。
on_status:可选进度回调;为 None 时进度打到 stderr(保持旧行为)。
"""
def _emit(msg):
if on_status is not None:
on_status(msg)
else:
print(msg, file=sys.stderr, flush=True)
import os
os.environ["MODEL"] = args.model
os.environ["QN_CONFIG"] = args.qn_config
os.environ["MTP_OUT"] = args.mtp_out
os.environ["EXPERT_DIR"] = args.expert_dir
os.environ["EXPERT_SLOTS"] = str(args.expert_slots)
spec = args.spec_slots if args.spec_slots is not None else args.expert_slots
os.environ["POOL_SPEC_SLOTS"] = str(spec)
for k, v in _FASTPATH_ENV.items():
os.environ.setdefault(k, v)
import json
import mlx.core as mx # noqa: F401 确保 MLX 已就绪
from mlx_lm.models.qwen3_next import ModelArgs
from mlx_streaming.core.mem import setup_memory_hygiene
from mlx_streaming.model_builder import build_streaming_model
from mlx_streaming.mtp.drafter import MTPDrafter
from mlx_streaming.mtp.qwen3_next_mtp import load_mtp
# 长会话内存防御:封顶 MLX 可回收缓冲(默认 1GB),防长对话里缓冲缓存膨胀把常驻推过墙 /
# 触发 macOS 压缩器抖动。TUI 是典型长会话场景,故在用户主路径启动时就设上。
_applied = setup_memory_hygiene(cache_gb=config.mlx_cache_limit_gb(),
wired_gb=config.mlx_wired_limit_gb())
if _applied:
_emit(f"内存防御: {_applied}")
_emit("正在加载主模型 + 专家(流式)...")
model, tok, _store = build_streaming_model()
_emit("正在加载 MTP drafter...")
with open(args.qn_config) as f:
margs = ModelArgs.from_dict(json.load(f))
mtp = load_mtp(margs, args.mtp_out, quantize=True, bits=config.mtp_bits())
mtp.embed_tokens = model.model.embed_tokens # 共享主模型 embedding
drafter = MTPDrafter(mtp, model.lm_head)
return model, tok, drafter
def _warmup(model, tok, drafter, args):
"""跑一次生成做预热:首轮的明显卡顿主要来自现编译 Metal kernel + 填 MoE 专家 resident 池,
提前把这部分一次性开销移到加载阶段。
覆盖增强:用一段**较长、token id 跨大跨度词表分散**的合成 prompt——MoE 路由依赖 token 内容,
分散的 id 会命中更多专家、更充分地预填专家池;较长 prompt 又能走通多块分块 prefill。
合成 id 直接走 ids_mode(不依赖 tokenizer),分块 prefill(chunk=2)保证长 prompt 也不抬高显存峰值。
预热失败不致命(直接吞掉异常),不影响后续真实生成。
"""
import mlx.core as mx
from mlx_streaming.mtp.generate import mtp_generate
try:
vocab = int(model.model.embed_tokens.weight.shape[0])
n = min(64, vocab) # 预热 prompt 长度(兼顾覆盖与耗时)
step = max(1, vocab // n) # 在词表内均匀取样,最大化专家覆盖
ids = [(1 + i * step) % vocab for i in range(n)]
mtp_generate(model, drafter, tok, mx.array([ids]), 8,
K=args.k, ids_mode=True)
except Exception: # noqa: BLE001 预热仅为压首轮延迟,失败不应中断启动
pass
def _encode_chat(tok, messages, tools=None):
"""把多轮对话按聊天模板编码成 token id 列表。tools(OpenAI 格式)交给模板渲染工具定义。"""
tmpl = getattr(tok, "chat_template", None)
if tmpl:
kwargs = {"add_generation_prompt": True}
if tools:
kwargs["tools"] = tools
out = tok.apply_chat_template(messages, **kwargs)
# 返回形态随 transformers 版本而变,三种都要兜住:
# transformers 4.x → list[int]
# transformers 5.x(默认) → BatchEncoding,即 {"input_ids": [[...]], ...}
# 部分实现 → str
# 只写 list(out) 会在 5.x 下取到字典的**键名** ["input_ids", "attention_mask"],
# 于是 prompt 退化成两个字符串:主生成路径(backend.py 的 ids)喂进去全是乱码,
# 且 _fixed_prefix_len 恒为 2 → head < 256 → 前缀快照永久失效。
# 单元测试全部 monkeypatch 掉了本函数,捕不到这个回归,故在此显式判型。
if isinstance(out, str):
return tok.encode(out)
ids = out if isinstance(out, (list, tuple)) else out["input_ids"]
if len(ids) > 0 and isinstance(ids[0], (list, tuple)):
ids = ids[0] # 批次维度
return [int(t) for t in ids]
# 无聊天模板:朴素拼接兜底。
text = ""
for m in messages:
text += f"{m['role']}: {m['content']}\n"
text += "assistant: "
return tok.encode(text)
def _eos_set(tok):
"""收集所有可能的结束符 token id。"""
eos = set()
ids = getattr(tok, "eos_token_ids", None)
if ids:
eos |= set(ids)
one = getattr(tok, "eos_token_id", None)
if one is not None:
eos.add(one)
return eos
def _truncate_eos(produced, eos):
"""遇到第一个结束符即截断(不含结束符本身)。"""
for i, t in enumerate(produced):
if t in eos:
return produced[:i]
return produced
_HELP = """可用命令:
/exit, /quit 退出
/reset 清空对话历史(保留 system)
/help 显示本帮助
直接输入文本即可对话。"""
def cmd_chat(args):
"""默认启动全屏 TUI;--plain 走纯文本 REPL;--demo 用假后端免模型预览界面。"""
if getattr(args, "plain", False):
return _chat_repl(args)
from mlx_streaming.tui import run_tui
if getattr(args, "demo", False):
# 免模型预览:秒开 TUI,假流式回答,用于验证界面/占位符/状态栏
from mlx_streaming.tui.backend import FakeBackend
demo = FakeBackend(
reply="这是 --demo 测试回答:界面、逐字流式、状态栏(token 数 / tok·s)"
"均为模拟,不加载模型。按 Esc 可中断,/help 看命令。",
delay=0.03)
return run_tui(demo, args)
from mlx_streaming.tui.backend import MLXBackend
return run_tui(MLXBackend(args), args)
def _chat_repl(args):
model, tok, drafter = _build_engine(args)
import mlx.core as mx
from mlx_streaming.mtp.generate import mtp_generate
from mlx_streaming.tui.backend import _reuse_prefix_len
eos = _eos_set(tok)
base_messages = []
if args.system:
base_messages.append({"role": "system", "content": args.system})
messages = list(base_messages)
# 跨轮复用:持久化上一轮 main_cache 及其对应的 token 序列(与 MLXBackend 同机制)。
main_cache = None
cached_ids: list[int] = []
print("\n正在预热(编译 kernel + 填专家池)…", file=sys.stderr, flush=True)
_warmup(model, tok, drafter, args)
print("模型已就绪。输入 /help 查看命令,/exit 退出。", file=sys.stderr, flush=True)
while True:
try:
user = input("\n你 > ").strip()
except (EOFError, KeyboardInterrupt):
print()
break
if not user:
continue
if user in ("/exit", "/quit"):
break
if user == "/reset":
messages = list(base_messages)
main_cache, cached_ids = None, [] # 清历史同时弃用旧 cache
print("对话历史已清空。", file=sys.stderr)
continue
if user == "/help":
print(_HELP, file=sys.stderr)
continue
messages.append({"role": "user", "content": user})
ids = _encode_chat(tok, messages)
# 旧 cache 是本轮 prompt 的严格前缀时只 prefill 新增后缀,否则全量重建。
cached_len = (_reuse_prefix_len(cached_ids, ids)
if main_cache is not None else 0)
cur_cache = main_cache if cached_len else model.make_cache()
# EOS 提前停止:否则会空跑到 max_tokens,且 produced 会含 EOS 后垃圾 token,
# 使 cached_ids 与下轮编码前缀断裂、复用永不触发。
produced_all: list[int] = []
def _on_tokens(new_ids):
produced_all.extend(new_ids)
truncated = _truncate_eos(produced_all, eos)
return len(truncated) < len(produced_all)
t0 = time.perf_counter()
produced, stats = mtp_generate(
model, drafter, tok, mx.array([ids]),
args.max_tokens, K=args.k, ids_mode=True, profile=args.stats,
on_tokens=_on_tokens, main_cache=cur_cache, cached_len=cached_len)
dt = time.perf_counter() - t0
# offset 对账:仅无 over-commit 时记录 cache 供下轮复用(见 MLXBackend.generate)。
if stats.get("resident_tokens") == len(ids) + len(produced) - 1:
main_cache = cur_cache
cached_ids = list(ids) + list(produced[:-1])
else:
main_cache, cached_ids = None, []
out_ids = _truncate_eos(produced, eos)
text = tok.decode(out_ids)
print(f"\n助手 > {text}")
messages.append({"role": "assistant", "content": text})
if args.stats:
tps = len(out_ids) / dt if dt > 0 else 0.0
print(f"[{len(out_ids)} tok, {tps:.1f} tok/s, "
f"accept_len={stats.get('avg_accept_len')}]", file=sys.stderr)
print("再见。", file=sys.stderr)
return 0
def _add_chat_args(p):
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)")
p.add_argument("--system", default=None, help="可选 system 提示词")
p.add_argument("--stats", action="store_true",
help="每轮结束在 stderr 打印 token 数 / tok·s / 接受长度")
p.add_argument("--plain", action="store_true",
help="用纯文本 REPL,不启动全屏 TUI(终端不兼容/调试时用)")
p.add_argument("--demo", action="store_true",
help="免模型预览:用假后端秒开 TUI,验证界面/流式/状态栏")
p.set_defaults(func=cmd_chat)
def _build_parser():
parser = argparse.ArgumentParser(
prog="sparkle",
description="sparkle:Apple Silicon 上的流式 MoE + Qwen3-Next MTP 自投机推理")
sub = parser.add_subparsers(dest="cmd")
chat = sub.add_parser("chat", help="进入交互式多轮对话(MTP 自投机快路径)")
_add_chat_args(chat)
return parser
def main(argv=None):
argv = list(sys.argv[1:] if argv is None else argv)
subcmds = {"chat"}
# 让 chat 成为默认子命令:不带子命令(或首参是 flag)时自动补上 chat;
# 但保留顶层 -h/--help 直接显示总帮助。
if not argv or (argv[0] not in subcmds and argv[0] not in ("-h", "--help")):
argv = ["chat"] + argv
args = _build_parser().parse_args(argv)
return args.func(args)
if __name__ == "__main__":
raise SystemExit(main())