This commit is contained in:
fiser_jun 2026-08-04 14:34:00 +08:00
commit 4745f264b2
90 changed files with 13695 additions and 0 deletions

40
.gitignore vendored Normal file
View File

@ -0,0 +1,40 @@
.venv/
__pycache__/
*.pyc
.pytest_cache/
.worktrees/
*.gputrace/
# 模型权重本地放置(仅提交 models/README.md
models/**
!models/README.md
*.safetensors
/experts/
# CMake 缓存
native/ext/build/
# 日志与一次性产物
logs/
*.log
*.bak
benchmarks/reports/*.json
/dl_*.sh
/dl_*.log
/build_*.sh
/build_*.log
/extract_mtp.log
/dl_model.py
# 根目录校验脚本(非运行时)
/verify_*.py
# 本地可留、不入库
mlx_streaming/tests/
mlx_streaming/tools/
benchmarks/*.py
.DS_Store
# 本地可留、不入库
/性能优化与选型*.md
/8bit*.md
/演示*.md

187
README.md Normal file
View File

@ -0,0 +1,187 @@
# Sparkle
本地 **Qwen3-Next-80B-A3B Instruct 8bit** 流式推理引擎,对接 ccSparkle Agent模型源**Sparkle**)。
在约 64GB 统一内存的 Apple Silicon 上,把约 77GB 专家权重放在 SSD运行时按需装入 + 预取,从而跑通 80B 8bitoMLX / Ollama 全量加载做不到的档位)。动机见 [`models/README.md`](./models/README.md):用 80B@64G 验证流式路径,对应下游 **30B 8bit A3B** 进 32G/16G、统一内存省 50%+。
| 项 | 值 |
|----|----|
| Python | **3.14.***(与 `native_moe_ext.cpython-314-darwin.so` 对齐;见 `pyproject.toml` |
| 已验证机器 | MacBook Pro M2 Max / 64GB |
| 安装目录 | 本仓库根目录(下文记为 `$SPARKLE_HOME`,即 `clone` 后的路径) |
| 引擎端口 | `8317`(经中台 `8088` 启停与代理) |
| 模型 / 专家 / MTP | 见 [`models/README.md`](./models/README.md)(不入库,约 160GB |
| 性能基线 | 稳态约 910 tok/s贪心+MTP采样约 46 tok/s常驻约 2333GB 量级 |
依赖安装走清华 PyPI 镜像(`pyproject.toml` 的 `tool.uv.index`)。
---
## 1. 系统组成
```
浏览器前端(8090) ← 日常对话界面
│ 顶部选 Sparkle 源
中台代理(8088) ← 转发请求,规避跨域
Sparkle 引擎(8317) ← 80B 8bit 流式推理
├─ /v1/chat/completions OpenAI 兼容接口
└─ /admin 调参页(槽位 / 投机宽度 / 状态)
```
专家权重留在 SSD运行时按需读取 + 预测式预取;首次启动和切换话题时偶有停顿(读盘)。
**AUTOPIN**(默认开启):记录专家路由热度(`models/experts_8bit_g64/pool_usage.json`),启动时把最热一批 pin 进常驻池(默认占槽位 50%,定死不驱逐)。冷启动命中率可从约 72% 升到 **98%+**。`AUTOPIN=0` 关闭;`AUTOPIN_BUDGET_FRAC` 调比例01
---
## 2. 启动
**推荐**:打开对话界面 → 顶部选 **Sparkle** → 中台自动加载引擎(约 90 秒,状态变绿后再对话)。
1. 中台(桌面应用常会自动拉起;路径按你本机 Agent 工程根目录,下文记 `$AGENT_HOME`
```bash
cd "$AGENT_HOME" && node scripts/zhongtai-server.js
```
2. 前端静态服务:
```bash
cd "$AGENT_HOME/<前端静态目录>" && python3 -m http.server 8090
```
3. 浏览器打开对话页,顶部点 **Sparkle**——未运行时会自动启动(也可在 `配` → Sparkle tab「引擎管理」点**加载**)。
就绪:`配` → Sparkle 状态点变绿,或 `curl -s http://127.0.0.1:8088/sparkle-engine/status``"loaded":true`
**手动启动(调试)**
```bash
cd "$SPARKLE_HOME" # 本仓库根目录
.venv/bin/python -m mlx_streaming.server --port 8317 \
--model models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit \
--expert-dir models/experts_8bit_g64 \
--qn-config models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit/config.json \
--mtp-out models/qn_mtp_weights.safetensors \
--expert-slots 96
```
- Admin`http://127.0.0.1:8317/admin`
- 健康:`curl -s http://127.0.0.1:8317/api/stats``unplaced_experts` 应为 `0`
中台默认 `SPARKLE_ENGINE_DIR` → 本目录。改目录名后若「卸载」失灵,先杀 `8317` 再重启中台。
> **工具白名单**:默认只放画像 4 工具(`profile_list_employees` / `profile_get_portrait_overview` / `profile_get_person_report` / `profile_get_team_report`)。中台动态工具集会改 system 头部哈希,导致前缀快照 miss钉死工具集后快照才稳定。排班 `metro_*` 等会被裁掉。临时放开:`TOOLS_ALLOW=`;换名单:`TOOLS_ALLOW=a,b,c`。裁剪结果见 `logs/server.log``[TOOLS_ALLOW]`
> **不要自己设 `PREFILL_CHUNK`**。安全上限由槽位决定引擎启动时按槽位自算96 槽 → 464 槽 → 3超限会被钳制并打日志。
---
## 3. 日常对话
1. 打开对话页(`http://localhost:8090`),顶部选 **Sparkle**
2. 首次:`配` → Sparkle tabBase URL `http://127.0.0.1:8317`API Key 任意非空(如 `sparkle`),刷新选模型 `Qwen3-Next-80B-A3B-Instruct-MLX-8bit`,保存
3. 正常对话;工具调用链路与 oMLX 源相同
---
## 4. 调参与引擎管理
`配`**本地 Sparkle** tab「引擎管理」`http://127.0.0.1:8317/admin` 等价):
状态点:灰(未启动)→ 黄(加载中,约 90 秒)→ 绿(就绪)。`加载` / `卸载` 启停;就绪时显示 tok/s 与 RSS。
| 参数 | 范围 | 说明 | 生效 |
|------|------|------|------|
| expert-slots | 32160 | 常驻专家数,越大命中率越高、越占内存(每槽 ~160MB。64GB + AUTOPIN64 槽约 9 tok/s**96 槽约 10.3 tok/s推荐**≥128 近红线;**160 会卡死,勿用** | 应用后进程级重启(约 90 秒) |
| K投机宽度 | 15 | MTP 草稿 token 数,默认 3只影响速度 | 即时 |
| max_tokens | — | 单次回答最大长度 | 即时 |
日常 64 够用;长文档可试 96128系统开始 swap 就降回 64。
---
## 5. 故障排查
| 现象 | 原因 | 处理 |
|------|------|------|
| 503 / engine reloading | 加载或重建中 | 等 6090 秒后再试 |
| `EADDRINUSE ... 8317` | 旧 server 未关 | `lsof -tiTCP:8317 -sTCP:LISTEN \| xargs kill` |
| 中台 `EADDRINUSE ... 8088` | 已有健康中台 | 直接用;要重启先 kill |
| 改名后「卸载」无效 | 旧进程挂在旧路径 | 杀 8317重启中台再加载 |
| 明显变慢 | swap 或槽位太低 | 看内存压力;调槽;关其他大应用 |
| 首轮特别慢 | 冷启动 | 正常,第二轮起恢复 |
| 长文中途停顿 | 读盘装冷门专家 | 正常;调高槽位可缓解 |
| 刷新拉不到模型 | server 未起或端口错 | `curl http://127.0.0.1:8317/v1/models -H "Authorization: Bearer x"` |
---
## 6. 常见问题
**Q: 为什么 oMLX / Ollama 里看不到这个 8bit 模型?**
A: 权重远超 64GB 整模上限,只有 Sparkle 流式能跑。oMLX 继续跑中小模型即可。
**Q: 投机宽度 K 是什么?**
A: MTP 每步草稿 token 数;主模型批量验证,猜中白赚。只影响速度,不影响内容。
**Q: 回答乱编 / 不调工具 / 幻觉?**
A: Qwen3-Next 官方要求采样temp 0.7 / top_p 0.8 / top_k 20。贪心+MTP 在该模型上易幻觉。默认走采样(约 46 tok/s显式 `temperature: 0` 才走贪心+MTP约 10 tok/s。Sparkle 客户端对真实对话会把 `0` 改回采样。
**Q: 模块分析比 oMLX 7B 还空?**
A: 多为分析轮传了 `temperature: 0` 进贪心导致指令遵循崩坏,不是花名册数据错。另查 `unplaced_experts`(非 0 = 错专家)与是否需清空 `prefix_snapshots`
**Q: 回答里夹「<E5A4B9>」乱码**
A: 已修byte-level BPE 流式反分词用 `StreamingDetokenizer`,避免半字符 U+FFFD 丢字。
**Q: 8bit 比 7B 还笨?**
A: 曾因 prefill 超槽静默映射到 0 号专家。现自动算 `PREFILL_CHUNK`、超容量走正确慢路径,并用 `unplaced_experts` 告警(健康值恒为 0。非 0 时回答不可信;若错算过,删 `models/prefix_snapshots/` 再预热。
**Q: 和云端略有不同?**
A: 量化 + KV 压缩的正常差异;语义质量不受影响。
**Q: 首轮为什么要等?**
A: 新对话要 prefill 系统提示+工具定义。有前缀快照后:重启后首会话约百秒级,之后新会话约数十秒,同会话追问约数秒。`PREFIX_SNAPSHOT_HEAD` 调长度(默认 20480 关闭)。
**Q: 速度还能再快吗?**
A: 96 槽 + AUTOPIN + prefill 分块约 910 tok/s接近本机 8bit 上限。160 槽会卡死,勿试。
**Q: 能跑多长上下文?**
A: KV 量化默认开;对话侧另有约 24k token 压缩阈值(中台配置)。
**Q: 想换模型或其它量化档?**
A: 本引擎面向 Qwen3-Next-80B 流式 8bit 部署;换档需按当前模型重新准备专家 blob 与 MTP。
---
## 7. 目录要点
```
sparkle/
├── mlx_streaming/ # 引擎 Python 包(含已编译 native_moe_ext*.so
├── native/ext/ # C++/Metal 扩展源码(重编 .so 用)
├── models/ # 权重与专家 blob约 160GB勿删
│ ├── Qwen3-Next-80B-A3B-Instruct-MLX-8bit/
│ ├── experts_8bit_g64/
│ └── qn_mtp_weights.safetensors
├── .venv/
├── pyproject.toml / uv.lock
└── benchmarks/reports/ # 性能基线报告
```
`models/` 三者缺一不可。不要把模型目录移进 oMLX 等其它运行时的模型扫描目录(会列出但整模加载不了)。
换 Python 小版本时需 `make -C native/ext native_moe_ext` 重编 `.so`。`native/ext/build/` 已忽略。
---
## 8. 质量与速度(摘要)
- **默认应对话走采样解码**(约 temperature 0.7 / top_p 0.8 / top_k 20。`temperature=0` 会进贪心+MTP本机实测指令遵循差于同表 oMLX 7B。
- 稳态 decode采样约 **46 tok/s**;贪心+MTP 约 **910 tok/s**(速度更高,但质量风险大,生产默认勿用)。
- 专家池命中率健康时约 **94%+**`unplaced_experts ≠ 0` 表示曾用错专家权重,回答不可信。
---
## 9. 相关文档
| 文档 | 内容 |
|------|------|
| [models/README.md](./models/README.md) | 为何做流式、NAS 下载与专家/MTP 准备 |
| [benchmarks/reports/8bit-g64-baseline-2026-07-25.md](./benchmarks/reports/8bit-g64-baseline-2026-07-25.md) | 8bit 性能基线 |

View File

@ -0,0 +1,61 @@
# Qwen3-Next-80B-A3B 8bit(g64)基线 — M2 Max 64GB
日期:2026-07-25 · 机器:MacBook Pro M2 Max / 64GB / 内置 SSD
模型:`models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit`(**bits=8, group_size=64**,79GB)
数据:`models/experts_8bit_g64/blobs`(48 层,77GB,v1 affine,stride 3,342,336B,page-aligned)+ `models/qn_mtp_weights.safetensors`(3.3GB)
## 性能
| 场景 | tok/s | accept_len | 备注 |
|---|---|---|---|
| 冷启动首轮(32 tok) | 5.0 | 1.889 | 含 kernel 编译 + 池冷填充 |
| 稳态(128 tok,64 槽) | **9.4** | 2.246 | 符合 8-12 tok/s 预期 |
| server 短请求(33 tok,含 prefill) | 6.6 | — | OpenAI 链路实测 |
- 常驻内存:RSS ≈ **23.5 GB**(`--expert-slots 64`,池 ~10GB + dense + 缓冲;64GB 机器余量充足,可上探 96-128 槽)
- 加载时间:约 60-90s(含预热)
## 正确性(三层证据链)
1. **kernel 单元等价**(`verify_8bit_equiv.py`):L0/L15/L47 blob fused kernel vs Python 参考 **diff = 0**
2. **端到端对拍**(`verify_e2e_spec.py`):贪心 vs 投机 token 17 确定性分叉 → 触发对照
3. **KV 量化对照**(`verify_teacher_forced.py`,seq=1 vs seq=2 生产原语):
- `KV_QUANT=0`:**24/24 token 全一致**,cosine 0.99935 → 8bit 权重路径数值干净
- `KV_QUANT=1`(生产默认):23/24(95.8%),cosine 0.99870,翻转处 margin=0 近 tie → 达质量门(≥95%/≥0.99)
- 结论:分叉全部归因于 KV 量化(KV 量化既有有损特性),与权重量化位宽无关
## 已知事项
- `VirtualPool.acquire_host` latent bug:单次前向单层唯一专家数超池容量时 fetch 回退会 AttributeError;生产 prefill 分块 chunk=2 永不触发。勿用整段 seq=N 一次性前向做诊断。
- server `/admin` 调参页池内存估算已按 g64 修正为 ~160MB/槽。
- server 模型 id 取目录名:`Qwen3-Next-80B-A3B-Instruct-MLX-8bit`。
## 复现命令
```bash
# CLI
.venv/bin/sparkle --model models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit \
--expert-dir models/experts_8bit_g64 \
--qn-config models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit/config.json \
--mtp-out models/qn_mtp_weights.safetensors --expert-slots 64 --stats
# OpenAI server(前端 Sparkle 源对接,端口 8317)
.venv/bin/python -m mlx_streaming.server --port 8317 \
--model models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit \
--expert-dir models/experts_8bit_g64 \
--qn-config models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit/config.json \
--mtp-out models/qn_mtp_weights.safetensors --expert-slots 64
```
## 追记(2026-07-25 晚):AUTOPIN + 调优后
| 配置 | decode tok/s | 备注 |
|---|---|---|
| 64 槽 + AUTOPIN | ~8.0-9.4 | 稳定,RSS ~24GB |
| **96 槽 + AUTOPIN(生产)** | **9.3-10.3** | 甜点位,RSS ~32GB;CLI 128tok: 9.3(MTP4)/9.7(MTP8) |
| 160 槽 + AUTOPIN | 不可用 | RSS ~34GB 系统耗尽卡死(AUTOPIN eager 预热多占 ~5GB) |
- AUTOPIN(热度持久化+启动 pin 定死):命中率 72.2%→98.6%,pin 3439 专家/6.5s
- PREFILL_CHUNK=8:676 tok prompt 2m49s→33s
- MTP_BITS 4→8 A/B:accept_len 2.246→2.207,无效,保持 4
- 槽位调整改进程级重启(进程内重建在新 C++ pin 状态下会把引擎拖到 0.24 tok/s)

3
mlx_streaming/.gitignore vendored Normal file
View File

@ -0,0 +1,3 @@
.venv/
__pycache__/
*.pyc

View File

329
mlx_streaming/cli.py Normal file
View File

@ -0,0 +1,329 @@
"""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 跨大跨度词表分散**的合成 promptMoE 路由依赖 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())

298
mlx_streaming/config.py Normal file
View File

@ -0,0 +1,298 @@
"""集中配置所有运行时开关从环境变量读取env 名与默认值与历史一致,便于现有跑法)。
设计目的把原先散落在 core/ 各文件的 `os.environ.get` 全部收口到这里单点可查
避免默认值漂移读取均为运行时每次调用读 env与原行为一致测试/probe
monkeypatch env 仍生效热路径每 token env 的成本与原先相同
同一 env 名在不同调用方有不同默认的项 CROSS_LAYER_PREFETCH_AHEADhook=0
native 预取=1accessor 暴露 `default` 参数由调用方传入绝不擅自统一
"""
import os
def _b(name: str, default: str = "0") -> bool:
return os.environ.get(name, default) == "1"
def _i(name: str, default) -> int:
return int(os.environ.get(name, str(default)))
def _f(name: str, default) -> float:
return float(os.environ.get(name, str(default)))
def _s(name: str, default: str = "") -> str:
return os.environ.get(name, default)
# ============================ 模型 / 目录 ============================
# 默认放持久目录 models/(不要放 /tmp会被系统清空。运行从仓库根目录起相对路径即可。
def model_path() -> str: return _s("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
def qn_config() -> str: return _s("QN_CONFIG", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit/config.json")
def mtp_out() -> str: return _s("MTP_OUT", "models/qn_mtp_weights.safetensors")
def expert_dir(default: str = "models/experts_8bit_g64") -> str: return _s("EXPERT_DIR", default)
def expert_slots() -> int: return _i("EXPERT_SLOTS", 64)
# 长期运行内存防御:封顶 MLX 可回收缓冲(默认 1GB),防长会话缓存膨胀;
# 双源侧区池的专家 buffer 走 C++ owned pool、不经 MLX 缓冲缓存,故 1GB 缓冲复用额度已够,
# 再大只是白占常驻。wired limit 默认 0=关(opt-in),设 >0 则 wire 该 GB 数的 GPU 缓冲防 macOS
# 压缩器,须 < 系统建议工作集(本机 26.8GB)。
def mlx_cache_limit_gb() -> float: return _f("MLX_CACHE_LIMIT_GB", 1.0)
def mlx_wired_limit_gb() -> float: return _f("MLX_WIRED_LIMIT_GB", 0.0)
def expert_pool_profile() -> str: return _s("EXPERT_POOL_PROFILE", "")
def hidden_variant() -> str: return _s("HIDDEN_VARIANT", "pre_final_norm")
def blob_dir() -> str: return _s("BLOB_DIR", "")
def compute_buffer_dir() -> str: return _s("COMPUTE_BUFFER_DIR", "")
# ============================ MoE 热路径 ============================
def resident_pool_enabled() -> bool: return _b("RESIDENT_POOL", "1")
def gpu_remap_enabled() -> bool: return _b("GPU_REMAP", "1")
# 实验:让 verify(seq>1)也走 GPU 侧 slot 重映射(acquire_gpu),消掉 host 路径每层 .tolist() 栅栏。
def verify_gpu_remap() -> bool: return _b("VERIFY_GPU_REMAP", "0")
# A2GPU 重映射路径下 promote 是否用 GPU membership 现算 used 过滤假阳性(默认开;回退设 0
def gpu_remap_promote_filter() -> bool: return _b("GPU_REMAP_PROMOTE_FILTER", "1")
def moe_topk_override() -> "str | None": return os.environ.get("MOE_TOPK_OVERRIDE")
def eager_expert_load() -> bool: return _b("EAGER_EXPERT_LOAD", "0")
# ============================ 缓存 / 驱逐 ============================
def evict_policy() -> str: return _s("EVICT_POLICY", "lfu")
# 专家读盘是否绕过 OS page cache(F_NOCACHE):默认开,保证基准每次都是真实 NVMe 读、
# 结果可复现(不被页缓存冷热污染);EXPERT_NOCACHE=0 恢复 mmap/page-cache(重复跑更快但飘)。
def expert_nocache() -> bool: return _b("EXPERT_NOCACHE", "1")
def lfu_decay_interval() -> int: return _i("LFU_DECAY_INTERVAL", 0)
def expert_bundle() -> bool: return _b("EXPERT_BUNDLE", "0")
def expert_bundle_dir(default: str) -> str: return _s("EXPERT_BUNDLE_DIR", default)
def expert_bundle_cache() -> int: return _i("EXPERT_BUNDLE_CACHE", 4)
def expert_stack_cache() -> int: return _i("EXPERT_STACK_CACHE", 0)
def async_prefetch() -> bool: return _b("ASYNC_PREFETCH", "0")
def prefetch_buffer_experts() -> int: return _i("PREFETCH_BUFFER_EXPERTS", 2048)
# ============================ AUTOPIN(专家热度持久化 + 启动预热钉死) ============================
# 总开关。1(默认):会话中持续累计每层路由热度并节流落盘({EXPERT_DIR}/pool_usage.json)
# 引擎启动时按历史热度把每层最热专家预填进常驻池并 pin 住(不参与任何驱逐),消灭冷启动慢热。
# 0整个特性零行为(不计数、不读写盘、不预热)。
def autopin() -> bool: return _b("AUTOPIN", "1")
# 热度文件路径覆盖:缺省 {EXPERT_DIR}/pool_usage.json结构 {layer: {expert_id: count}}。
def autopin_usage_file() -> str: return _s("AUTOPIN_USAGE_FILE", "")
# 落盘节流:约每这么多次前向写一次盘(前向边界按「MoE 层号回绕」近似)atexit 总会兜底写一次。
# 写同目录临时文件再 rename 防半写损坏;读不到/损坏静默从零开始。
def autopin_save_every() -> int: return _i("AUTOPIN_SAVE_EVERY", 200)
# 预热 pin 预算占比:每层 pin 数量 = floor(该层池容量 cap_for × FRAC),取值 0-1默认 0.5。
# pinned 永占真实区槽FRAC 过大会挤掉 demand 自由槽(驱逐候选变少),不建议 >0.5。
def autopin_budget_frac() -> float: return min(1.0, max(0.0, _f("AUTOPIN_BUDGET_FRAC", 0.5)))
# ============================ 自定义 Metal 算子 ============================
def custom_qproj() -> bool: return _b("CUSTOM_QPROJ", "0")
def custom_qproj_bits() -> int: return _i("CUSTOM_QPROJ_BITS", 6)
def custom_qproj_targets() -> str: return _s("CUSTOM_QPROJ_TARGETS", "gate,up")
def custom_qproj_max_seq() -> int: return _i("CUSTOM_QPROJ_MAX_SEQ", 4)
def custom_qproj_tile() -> int: return _i("CUSTOM_QPROJ_TILE", 4)
def custom_fused_moe() -> bool: return _b("CUSTOM_FUSED_MOE", "0")
def custom_fused_moe_bits() -> int: return _i("CUSTOM_FUSED_MOE_BITS", 6)
def custom_fused_moe_lanes() -> int: return _i("CUSTOM_FUSED_MOE_LANES", 8)
def custom_fused_moe_block() -> int: return _i("CUSTOM_FUSED_MOE_BLOCK", 256)
def custom_fused_moe_max_seq() -> int: return _i("CUSTOM_FUSED_MOE_MAX_SEQ", 4)
# ============================ native MoE 后端 ============================
def native_moe() -> bool: return _b("NATIVE_MOE", "0")
def native_moe_mlx_op() -> bool: return _b("NATIVE_MOE_MLX_OP", "0")
def native_moe_synthetic() -> bool: return _b("NATIVE_MOE_SYNTHETIC", "0")
def native_moe_slot_pool() -> bool: return _b("NATIVE_MOE_SLOT_POOL", "0")
def native_moe_raise() -> bool: return _b("NATIVE_MOE_RAISE", "0")
def native_moe_slot_cap() -> int: return _i("NATIVE_MOE_SLOT_CAP", 96)
def native_moe_stage_cache() -> bool: return _b("NATIVE_MOE_STAGE_CACHE", "1")
def native_moe_stage_cache_experts() -> int: return _i("NATIVE_MOE_STAGE_CACHE_EXPERTS", 96)
def native_moe_stage_bundle_cache() -> bool: return _b("NATIVE_MOE_STAGE_BUNDLE_CACHE", "1")
def native_moe_stage_cache_bundles() -> int: return _i("NATIVE_MOE_STAGE_CACHE_BUNDLES", 16)
def native_moe_stage_prefetch() -> bool: return _b("NATIVE_MOE_STAGE_PREFETCH", "0")
# ============================ 流式 blob ============================
def stream_blob() -> bool: return _b("STREAM_BLOB", "0")
def stream_blob_loader() -> bool: return _b("STREAM_BLOB_LOADER", "0")
def stream_blob_bg() -> bool: return _b("STREAM_BLOB_BG", "0")
def stream_blob_workers() -> int: return _i("STREAM_BLOB_WORKERS", 8)
def stream_blob_window() -> int: return _i("STREAM_BLOB_WINDOW", 3)
def stream_blob_nocache(default: str = "1") -> bool: return _b("STREAM_BLOB_NOCACHE", default)
# 默认 12:双源侧区(cap32+侧区32)K=3 MTP 甜点扫描结果(REPEAT=2)。瓶颈是 SSD 读盘时序而非容量,
# 加大 budget 只会灌满共享读队列→预取到得更晚(timing miss↑)。W=24/B=12 在 hit(0.901)/disk(3874)/
# tok·s(11.61)/确定性(nmis=2)全面优于旧 W=32/B=16(10.0 tok·s、nmis=43),tok·s +16%。
def stream_blob_bg_budget(default: int = 12) -> int: return _i("STREAM_BLOB_BG_BUDGET", default)
# 按需 miss 批量并行读:把本层所有 miss 专家收成一批,一次 load_experts(8-worker 并行 pread)
# 取代逐专家串行 pread实测串行 6.4GB/s vs 并行 8 路 22GB/s。仅影响读取调度结果等价。
# 默认开:剖析实测 decode 时间 ~92% 花在每层 host 回退路径,其中字节物化是大头;批量堆叠
# (每段一次 mx.array 取代逐专家 6N 次构造)在 cap=12 把 decode 从 ~466ms 降到 ~408ms(~12%)。
def batch_miss_read() -> bool: return _b("BATCH_MISS_READ", "1")
# demand miss 走 C++ 原生物化(load_experts_native)blob_load 把字节 pread 直进 MLX 数组、
# eval 时在 C++ 跑、绕 GIL、惰性。比 stacked 再快 ~6%(把物化从 acquire 挪到 eval)。
# 默认关(opt-in):实测它偶发把末尾 token 推进双稳态慢挡(2-3s/token 的"悬崖")
# 其 lazy 物化推高内存压力,稳定性不如 stacked(纯 numpy 同步、4 次跑 0 慢挡)。仅在确认
# 环境内存充裕、需要那 6% 时手动开。不消费 async prefetch buffer故仅在 not async_prefetch 接入。
def native_demand_loader() -> bool: return _b("NATIVE_DEMAND_LOADER", "0")
def staging_ring() -> int: return _i("STAGING_RING", 2) # 安全下限=2MTP 每步 verify+replay 各对同层 submit 一次
def staging_pread_parallel() -> bool: return _b("STAGING_PREAD_PARALLEL", "0") # staging fill 派后台池并行;默认关:实测不降 timing miss(IO 受限)且弱化 buffer 新鲜度不变量,详见 benchmarks/reports/staging-pread-parallel-2026-06-25.md
def stg_verify() -> bool: return _b("STG_VERIFY", "0") # 诊断:acquire_gpu 命中后池槽字节真值校验,默认关、对主路径零影响
def stream_blob_prefetch_budget(default: int) -> int: return _i("STREAM_BLOB_PREFETCH_BUDGET", default)
# ============================ 跨层 / 预取 ============================
# 注意 AHEAD 默认按调用方传:跨层 hook=0、native 预取=1。
def cross_layer_prefetch() -> bool: return _b("CROSS_LAYER_PREFETCH", "0")
def cross_layer_ahead(default: int = 0) -> int: return _i("CROSS_LAYER_PREFETCH_AHEAD", default)
def cross_layer_mult() -> int: return max(1, _i("CROSS_LAYER_PREFETCH_MULT", 2))
# 预测宽度方案B预测 top-N 候选用于"减常驻"N 大→recall 高且不占内存(只是 gate argpartition
# 真正占 staging 的是过滤常驻后、按分截断到 staging budget 的缺口子集。
# 默认 24:双源侧区 K=3 甜点扫描(REPEAT=2)。realized hit 在 W≥24 就到顶(~0.90),再加宽只抬"名义
# recall"——多出的候选在 SSD 时序上落不进侧区且抢带宽,W=48→hit 0.85、W=64→0.75、tok·s 一路下滑。
# 收窄到刚好覆盖可及工作集(W=24)最快最省。详见 benchmarks/reports/prefetch-width-budget-sweep-2026-07-01.md。
def cross_layer_predict_width() -> int: return _i("CROSS_LAYER_PREDICT_WIDTH", 24)
# per-layer 自适应 ahead早层用小 ahead 保召回、晚层用大 ahead 保时序(默认来自实测)。
def cross_layer_cutoff() -> int: return _i("CROSS_LAYER_CUTOFF", 6) # 切点:层号 <cutoff 用 lo否则 hi
def cross_layer_ahead_lo() -> int: return _i("CROSS_LAYER_AHEAD_LO", 1) # 早层 ahead保召回
def cross_layer_ahead_hi() -> int: return _i("CROSS_LAYER_AHEAD_HI", 3) # 晚层 ahead保时序
def predict_use_x() -> bool: return _b("PREDICT_USE_X", "1") # 默认用本层 MoE 输入 x更新鲜+3.6pp recall=0 回退旧 norm 路径
def predict_agg() -> str: return _s("PREDICT_AGG", "max") # K+1 token 聚合max|mean|union
def predict_union_k() -> int: return _i("PREDICT_UNION_K", 8) # union 时每 token 取的 top-k控候选数
def native_fused_prefetch() -> bool: return _b("NATIVE_FUSED_PREFETCH", "0")
# 池侧区零拷贝双源(单/双缓冲)opt-in、默认 off。VirtualPool 收口,消掉 promote 拷贝。
# 侧区有两种淘汰策略(SIDEREGION_LFU 门控):
# - 旧"∉P 全清"(默认):侧区=一次性预取批,不积累→hit 仅 0.709(反低于基线 0.763)。
# - 新"持久 LFU"(SIDEREGION_LFU=1,单代 spec_gens=1):跨步累积热专家,只读新增。
# 实测(80B,cap=32,warmup64,见 report sideregion-lfu-2026-07-01):
# LFU spec=8 → hit 0.73 / active 4.76GB(省内存)LFU spec=32 → hit 0.81 / 6.9GB(提命中,+8% tok/s)。
# 命中在 ~0.81 饱和(加 warmup 无效),0.85+ 仍需真加常驻槽(cap=64→0.869)。
# 注:dual on 各 spec 均有 run-to-run token 漂移(良性时序噪声,字节校验 0 BAD),故默认 off。
def zerocopy_dual_source() -> bool: return _b("ZEROCOPY_DUAL_SOURCE")
def pool_spec_slots() -> int: return _i("POOL_SPEC_SLOTS", 3) # 每层侧区投机槽数(LFU 推荐 8 省内存 / 32 提命中)
def sideregion_lfu() -> bool: return _b("SIDEREGION_LFU", "1") # 侧区持久 LFU 单缓冲二级缓存(默认 on=生产路径);SIDEREGION_LFU=0 回退 legacy 双缓冲
def native_no_submit() -> bool: return _b("NATIVE_NO_SUBMIT", "0")
def native_no_promote() -> bool: return _b("NATIVE_NO_PROMOTE", "0")
def native_materialize() -> bool: return _b("NATIVE_MATERIALIZE", "0")
def file_prefetch_global_budget() -> int: return _i("FILE_PREFETCH_GLOBAL_BUDGET", 64)
def stage_prefetch_global_budget() -> int: return _i("STAGE_PREFETCH_GLOBAL_BUDGET", 64)
def stage_prefetch_min_score() -> float: return _f("STAGE_PREFETCH_MIN_SCORE", 0)
def stage_prefetch_per_layer_budget(default: int = 4) -> int: return _i("STAGE_PREFETCH_PER_LAYER_BUDGET", default)
# ============================ 前缀快照(跨会话 system+tools 复用) ============================
# 每个会话开头都是同一段长 system+tools 提示词(agentic 场景 ~2600 tok),新会话首轮
# 要全量 prefill(实测 158s)。快照前 N 个 token 的 KV/递归态,后续同前缀会话直接跳过。
# 0=关闭。注意:仅当 prompt 明显长于该值才启用(短提示词无收益)。
def prefix_snapshot_head() -> int: return _i("PREFIX_SNAPSHOT_HEAD", 2048)
# 快照条目上限(LRU),每条约占几十 MB(线性层递归态为主)
def prefix_snapshot_entries() -> int: return _i("PREFIX_SNAPSHOT_ENTRIES", 4)
# ============================ 工具白名单 ============================
# 渲染进 chat template 的工具白名单(按工具名),用于把 prompt 头部钉成固定的一段。
#
# 快照命中的前提是头部**字节级**相同(key = ids[:head] 的哈希)。而调用方的工具发现是动态的:
# 同一段 system 提示词会配不同的工具集发来,头部随之变化,快照必然 miss —— 每个新会话都要
# 全量 prefill 4000 token 上下。实测本部署已因此存下两份不同头部的快照(3972 / 4261 token),
# 互不命中。白名单把工具集收敛成固定集合后,头部才稳定、快照才能真正消灭"首轮等待"。
# 省 token 只是附带好处(工具定义约占头部三分之一)。
#
# 默认只放画像(profile)子系统实际在用的 4 个工具;其余(如排班 metro_*)本部署不处理。
# 设为空(TOOLS_ALLOW=)即关闭过滤、原样透传全部工具。
_TOOLS_ALLOW_DEFAULT = ("profile_list_employees,profile_get_portrait_overview,"
"profile_get_person_report,profile_get_team_report")
def tools_allow() -> "list[str]":
spec = _s("TOOLS_ALLOW", _TOOLS_ALLOW_DEFAULT)
return [s for s in (part.strip() for part in spec.split(",")) if s]
# ============================ MTP ============================
def mtp_verify_mode() -> str: return _s("MTP_VERIFY_MODE", "batch")
# MTP drafter 量化位宽(load_mtp bits 参数):4(默认,省内存)| 8(草稿更准,接受率更高,
# 实测杠杆:accept_len 提升即等比例提速;代价 ~0.7GB 内存)
def mtp_bits() -> int: return _i("MTP_BITS", 4)
# 喂给 MTP drafter 的主模型 hidden:pre_norm(默认,与训练/验证一致,接受率高)| post_norm(旧行为)
def mtp_hidden() -> str: return _s("MTP_HIDDEN", "pre_norm")
# 分块 prefill 的每块 token 数:整段 prefill 一次前向的激活峰值 ∝ prompt 长度(每 MoE 层瞬时
# 物化大量唯一专家 + 长序列激活),把 prompt 切成小块逐块喂入(KV/SSM 按 offset 因果累积、末块
# 末位 logits 与整段等价),峰值压回 ∝chunk,使 prefill 与 decode 同稳态。默认 2(沿用 DeepSeek
# 实证:每块唯一专家少、稳走 resident 池便宜路径);PREFILL_CHUNK=0 关闭分块(整段 prefill)。
def prefill_chunk() -> int: return _i("PREFILL_CHUNK", 2)
# ============================ KV 量化(IsoQuant K4/V3)============================
# 只作用于 12 个全注意力层。SO(4) 块旋转去相关 + 非对称 K4/V3 仿射量化:128k KV 3.0→~0.68 GiB。
# 线性层(Gated DeltaNet 递归态)不动。质量验收:token 一致率≥95% + logits cosine≥0.99。
def kv_quant() -> bool: return _b("KV_QUANT", "0")
def kv_k_bits() -> int: return _i("KV_K_BITS", 4) # Key 位宽(默认 4)
def kv_v_bits() -> int: return _i("KV_V_BITS", 3) # Value 位宽(默认 3,非对称)
def kv_group_size() -> int: return _i("KV_GROUP_SIZE", 64) # 仿射量化分组(整除 head_dim=256)
def kv_rotate() -> bool: return _b("KV_ROTATE", "1") # SO(4) 块旋转去相关(默认开)
def kv_rot_seed() -> int: return _i("KV_ROT_SEED", 0) # 旋转随机种子(data-oblivious、固定)
# ============================ profiling / 诊断 ============================
def window_prof() -> bool: return _b("WINDOW_PROF", "0")
def predict_recall_prof() -> bool: return _b("PREDICT_RECALL_PROF", "0")
def miss_attrib() -> bool: return _b("MISS_ATTRIB", "0") # miss 归因A(预测到没到位)/B(没预测到)
def route_trace_enabled() -> bool: return _b("ROUTE_TRACE", "0")
def stream_prof() -> bool: return _b("STREAM_PROF", "0")
def probe_predict_only() -> bool: return _b("PROBE_PREDICT_ONLY", "0")
def probe_perlayer_sync() -> bool: return _b("PROBE_PERLAYER_SYNC", "0")
# 预取 host 墙钟探针:量 predict/submit/promote 各段主线程不可重叠的 CPU 时间。默认关、零开销。
def prefetch_tprof() -> bool: return _b("PREFETCH_TPROF", "0")
# 并集专家数探针:按前向 seq 分桶记每层路由专家并集大小(seq=K 即 MTP verify 的专家并集)。默认关。
def union_prof() -> bool: return _b("UNION_PROF", "0")
# 接受率 top-k 覆盖探针(>0 且 profile 时):记每个草稿位置 MTP 的 top-k 候选,量模型真实 token
# 是否落在 top-2/top-3 里 = 树形展开的救回上界。默认 0=关。
def accept_topk() -> int: return _i("ACCEPT_TOPK", 0)
# 最小树:位置1 展开 top-2。仅当 A 链首草稿被拒且 B 候选=真实 token 时,额外跑一次 B 链前向救回。
# 两分支各为独立 batch=1 seq=K 前向(线性层不能批处理树,故不拍平),per-forward union 不变。默认关。
def tree_top2() -> bool: return _b("TREE_TOP2", "0")
# 第2 草稿位置(pos1)top-2 救回:在 tree_top2 基础上,额外抽 chainC(第2 位次选分支),当第1 位
# 命中但第2 位被拒、且 chainC 第2 token=模型真值时改验 chainC。探针实测 pos1「首选错次选对」比例
# (~11%)高于 pos0(~7%),消融证实接受长度确实 +2.74%(bit-lossless);但每次救回要多一次主模型
# 前向,该成本在当前硬件/批量下恰好抵消 token 收益,净 tok/s 无提升(见 benchmarks/_bench_p1_ablation)。
# 故默认关,保护已验证的 pos0 纯路径 tok/s 收益;留作前向变廉价/批量变大时收益翻正的储备,一行 env 开。
def tree_top2_p1() -> bool: return _b("TREE_TOP2_P1", "0")
# 完整树形验证(batch-of-paths):把 P 条候选路径拍到 batch 维,一次 batched 前向并行验证所有路径,
# 选接受最长的路径提交(提取赢家 row)。每条路径是普通线性序列,故线性层/全注意力层都走成熟的
# batch 前向(无需改 kernel);单次前向的 batch=P 计算加宽了预取窗口,同时多路径提升接受长度
# ——一举拿下 accept_len 与 hit_rate 两只鸟。默认关(与 tree_top2/普通链验证互斥,优先级最高)。
def tree_verify() -> bool: return _b("TREE_VERIFY", "0")
# 树分支数(候选路径条数 P)。当前 drafter 在位置1 展开 top-P,P=2 即 top-2。
def tree_branches() -> int: return _i("TREE_BRANCHES", 2)
# 置信度门控动态深度(P-MTP 风格):逐位抽草稿时累计置信度 C_i=p0·…·p_i,C_i≥tau 且未到 depth_max
# 就继续加深,否则本步只 verify 到当前深度。低置信步收缩(省专家加载/cap 压力),高置信步至多到
# depth_max。探针实测置信度对接受率区分度极强(高低置信接受率差 +42~59pp)。消融结论(见
# benchmarks/reports,_bench_adaptive):收益全部来自"向下收缩",τ=0.3、depth_max=3 即 +5~6% tok/s
# 且 bit-lossless、零额外显存;向上扩到 K=4 反而更慢(第4 位专家加载成本 > 多接受的 token),且在
# 生产 EXPERT_SLOTS=32 下 seq=4·top_k=40>cap 会溢出致有损。故 depth_max 默认 3(=基础 K,slots=32 安全),
# 要扩到 4 必须同时把 EXPERT_SLOTS 提到 ≥40。默认关,作可选加速路径(与 tree_top2 互斥)。
def adaptive_depth() -> bool: return _b("MTP_ADAPTIVE_DEPTH", "0")
def conf_tau() -> float: return _f("MTP_CONF_TAU", 0.3)
def depth_max() -> int: return _i("MTP_DEPTH_MAX", 3)
# 合并路径:在动态深度基础上,对"被保留成深链(n>=2)"的步叠加 pos0 top-2 救回(首 token 被拒且
# top-2=模型真值时改验 B 链)。两机理正交(动态深度压每步成本、救回抬接受长度),但都作用于低置信步、
# 方向相反(depth=1 的浅步无位置可救),故叠加非简单相加,须实测。默认关。
def adaptive_rescue() -> bool: return _b("MTP_ADAPTIVE_RESCUE", "0")
def parse_layers_env(name: str) -> "set[int] | None":
"""解析 "0-3,5,8" 形式的层集合环境变量;为空返回 None表示全部层"""
spec = os.environ.get(name, "").strip()
if not spec:
return None
out: "set[int]" = set()
for part in spec.split(","):
part = part.strip()
if not part:
continue
if "-" in part:
a, b = part.split("-", 1)
out.update(range(int(a), int(b) + 1))
else:
out.add(int(part))
return out

View File

1
mlx_streaming/core/cache/__init__.py vendored Normal file
View File

@ -0,0 +1 @@
"""专家缓存模块:常驻池(resident_pool) / LRU+文件后端存储(expert_store) / blob 流式源(blob_loader)。"""

174
mlx_streaming/core/cache/autopin.py vendored Normal file
View File

@ -0,0 +1,174 @@
"""AUTOPIN专家热度持久化 + 启动预热钉死。
两部分
- **热度计数与持久化**搭已有的物化点零同步计频block host 路径的 flat
acquire_gpu LFU piggyback(全命中) miss 回退 flatdual decode/verify Python
物化点 C++ g_real LFU freq 在落盘时按增量差分合并节流写
{EXPERT_DIR}/pool_usage.json(同目录临时文件 + rename 防半写读不到/损坏静默从零)
- **启动预热钉死**build_streaming_model 末尾按历史热度把每层 top-N 热专家
(N = floor(cap_for × AUTOPIN_BUDGET_FRAC)) blob 并行读预填进常驻池并 pin
(不参与任何驱逐)消灭冷启动慢热
覆盖率口径prefill/host 路径与 GPU remap miss 回退全计 dual decode 全命中
快路径搭 LFU piggyback (inds 已物化零额外同步)dual decode/verify C++ freq
兜底( lfu 策略计频)EVICT_POLICYlfu 时两条 GPU 快路径均不计绝不为计数给
decode 热路径新增 GPUhost 同步
"""
import atexit
import json
import os
import tempfile
from collections import Counter
from typing import Dict, List
from mlx_streaming import config
_USAGE_NAME = "pool_usage.json"
def usage_path() -> str:
"""热度文件路径AUTOPIN_USAGE_FILE 覆盖,缺省 {EXPERT_DIR}/pool_usage.json。"""
p = config.autopin_usage_file()
return p or os.path.join(config.expert_dir(), _USAGE_NAME)
def load_usage(path: str) -> "Dict[int, Counter]":
"""读热度文件;不存在/损坏/结构非法一律静默返回空(从零开始)。"""
try:
with open(path) as f:
raw = json.load(f)
return {int(l): Counter({int(e): int(c) for e, c in d.items()})
for l, d in raw.items()}
except Exception:
return {}
def save_usage(path: str, counts: "Dict[int, Counter]") -> None:
"""原子写:先写同目录临时文件再 rename防半写损坏。失败静默(不阻断推理)。"""
try:
d = os.path.dirname(path) or "."
fd, tmp = tempfile.mkstemp(dir=d, prefix=".pool_usage.", suffix=".tmp")
with os.fdopen(fd, "w") as f:
json.dump({str(l): {str(e): c for e, c in cnt.items()}
for l, cnt in counts.items()}, f)
os.replace(tmp, path)
except Exception:
pass
class HeatCounter:
"""进程内热度累计器note() 增量计频,节流 + atexit 落盘。
前向边界用MoE 层号回绕近似( VirtualPool.begin_forward 同判据)MoE 块按层号
递增被调出现 layer <= 上次即新前向路径在创建时快照避免进程退出时 env 已变
导致 atexit 写错位置
"""
def __init__(self, path: str):
self.path = path
self.counts: "Dict[int, Counter]" = {}
self._last_layer = -1
self._forwards_since_save = 0
# C++ 累计频次的已合并基线(real_freq_dump 是进程累计值,按差分合并防重复计)
self._cpp_baseline: "Dict[int, Dict[int, int]]" = {}
def note(self, layer: int, ids: List[int]) -> None:
layer = int(layer)
if layer <= self._last_layer:
self._forwards_since_save += 1
self._last_layer = layer
c = self.counts.get(layer)
if c is None:
c = self.counts[layer] = Counter()
for e in ids:
c[int(e)] += 1
if self._forwards_since_save >= config.autopin_save_every():
self.flush()
def tick(self) -> None:
"""显式前向边界(dual decode 无 Python 计频点,由 begin_forward 驱动周期落盘)。"""
self._forwards_since_save += 1
if self._forwards_since_save >= config.autopin_save_every():
self.flush()
def merge_cpp_freq(self) -> None:
"""把 dual decode 的 C++ LFU 累计频次按增量差分合并进来。"""
try:
import mlx_streaming.native_moe_ext as _N
flat = _N.real_freq_dump()
except Exception:
return
cur: "Dict[int, Dict[int, int]]" = {}
for i in range(0, len(flat), 3):
l, e, c = int(flat[i]), int(flat[i + 1]), int(flat[i + 2])
cur.setdefault(l, {})[e] = c
for l, em in cur.items():
base = self._cpp_baseline.get(l, {})
cnt = self.counts.setdefault(l, Counter())
for e, c in em.items():
b = base.get(e, 0)
cnt[e] += c - b if c >= b else c # real_reset 后 c<b以 c 为新基线
self._cpp_baseline = cur
def flush(self) -> None:
self.merge_cpp_freq()
save_usage(self.path, self.counts)
self._forwards_since_save = 0
_counter: "HeatCounter | None" = None
def _get_counter() -> HeatCounter:
global _counter
if _counter is None:
_counter = HeatCounter(usage_path())
atexit.register(_counter.flush)
return _counter
def note(layer: int, ids: List[int]) -> None:
"""路由热度计数入口(热路径调用)AUTOPIN=0 立即返回,零行为。"""
if not config.autopin():
return
_get_counter().note(layer, ids)
def tick() -> None:
"""前向边界入口(virtual_pool begin_forward 调用)AUTOPIN=0 立即返回。"""
if not config.autopin():
return
_get_counter().tick()
def counter() -> "HeatCounter | None":
"""测试/诊断用:当前计数器(未触发过计数则为 None)。"""
return _counter
def warm_start_pin(store) -> "dict":
"""启动预热钉死:按 usage 文件每层取 top-N 热专家,经 store.pin_batch 预填进池并 pin 住。
N = floor(该层 cap_for × AUTOPIN_BUDGET_FRAC)pin 加载由 pin_batch 收口 blob_loader
走并行 pread(生产唯一数据源)无则退回逐专家 per-expert 文件(兼容旧部署)
usage 缺失/为空 返回零摘要(后台计数仍在累计供下次启动用)单层失败不阻断整体
"""
summary = {"layers": 0, "pinned": 0, "skipped_layers": 0, "seconds": 0.0}
usage = load_usage(usage_path())
if not usage:
return summary
import time
t0 = time.perf_counter()
frac = config.autopin_budget_frac()
for layer in sorted(usage):
n = int(store.cap_for(layer) * frac)
if n <= 0:
continue
top = [e for e, _ in usage[layer].most_common(n)]
try:
summary["pinned"] += store.pin_batch(layer, top)
summary["layers"] += 1
except Exception:
summary["skipped_layers"] += 1
summary["seconds"] = round(time.perf_counter() - t0, 2)
return summary

238
mlx_streaming/core/cache/blob_loader.py vendored Normal file
View File

@ -0,0 +1,238 @@
"""全流式 blob 专家源:并行 pread + F_NOCACHE + 即时物化为 MLX 量化数组。
字节布局由 prep/blob_layout.py 统一描述单一真相源支持两种格式
- v1 affineexpert_blob_v1 proj 段序 [weight, scales, biases]weight uint32
scales/biases uint16 原始位用时 .view(bfloat16)
- v2 mxfp4expert_blob_v2_mxfp4 proj 段序 [weight, scales] biases
weight uint32scales uint8 原始位绝不 view(bf16)
计算复用 MLX quantized_matmul / gather_qmm不碰已 NO-GO fused kernel
"""
import os
import threading
from collections import OrderedDict
from concurrent.futures import ThreadPoolExecutor
import mlx.core as mx
import numpy as np
from mlx_streaming import config
# macOS fcntlF_NOCACHE 提示内核不要把读过的页留在 page cache支撑低内存目标
_F_NOCACHE = 48
class BlobExpertSource:
def __init__(self, blob_dir: str, hidden: int, inter: int, group: int, bits: int,
num_experts: int, workers: int = 8, nocache: bool = True,
quant_mode: str = "affine", blob_format: "str | None" = None):
self.dir = blob_dir
self.h, self.i, self.g, self.b, self.ne = hidden, inter, group, bits, num_experts
self.workers = workers
self.nocache = nocache
# blob 格式mxfp4 默认 v2无 biases、scales uint8其余默认 v1 affine。
self.quant_mode = quant_mode
self.blob_format = blob_format or (
"expert_blob_v2_mxfp4" if quant_mode == "mxfp4" else "expert_blob_v1")
self._segs, self.stride = self._layout()
self._fds: "dict[int, int]" = {}
# 后台预取:只预读原始字节进 _pf_cache不建 mx.array可后台线程跑
# 滚动窗口只保留最近 window 层的预取字节,支撑低内存。
self.window = config.stream_blob_window()
self._pf_pool = ThreadPoolExecutor(max_workers=workers)
self._pf_cache: "OrderedDict[tuple, bytes]" = OrderedDict()
self._pf_futures: "dict[tuple, object]" = {}
self._pf_layers: "list[int]" = []
self._lock = threading.Lock()
self.preads = 0
self.prefetch_hits = 0
def _layout(self):
# 段表统一由 blob_layout.layout_for 给出dtype 为字符串 "uint32"/"uint16"/"uint8")。
from mlx_streaming.prep.blob_layout import layout_for
return layout_for(self.blob_format, self.h, self.i, self.b, self.g)
def _fd(self, layer: int) -> int:
fd = self._fds.get(layer)
if fd is None:
fd = os.open(os.path.join(self.dir, f"layer{layer:02d}.blob"), os.O_RDONLY)
if self.nocache:
try:
import fcntl
fcntl.fcntl(fd, _F_NOCACHE, 1)
except (OSError, ValueError):
pass
self._fds[layer] = fd
return fd
def _materialize(self, raw: bytes, view_bf16: bool = True) -> dict:
# view_bf16=Falsescales/biases 保留 uint16.view(bfloat16) 在后台线程会报
# no Stream(gpu,0),故后台物化时关掉,由主线程消费时再 view
# 段表 dtype 为字符串:仅 v1 affine 的 uint16 scales/biases 是 bf16 位重解释;
# mxfp4 的 uint8 scales 保持原样,绝不 view(bf16)。
out = {}
off = 0
for proj, tensor, dt, shape, nb in self._segs:
np_dt = np.dtype(dt)
view = np.frombuffer(raw, dtype=np_dt, count=nb // np_dt.itemsize, offset=off).reshape(shape)
arr = mx.array(view)
if view_bf16 and dt == "uint16" and tensor in ("scales", "biases"):
arr = arr.view(mx.bfloat16)
out[f"{proj}.{tensor}"] = arr
off += nb
return out
def _pread(self, layer: int, e: int) -> bytes:
return os.pread(self._fd(layer), self.stride, e * self.stride)
def prefetch_async(self, layer: int, expert_ids) -> None:
"""后台并行预读预测专家的字节(仅 IO不建 mx.array。供跨层预测调用。"""
for e in (int(x) for x in expert_ids):
key = (layer, e)
with self._lock:
if key in self._pf_cache or key in self._pf_futures:
continue
self._pf_futures[key] = self._pf_pool.submit(self._pf_job, layer, e)
self._evict_window(layer)
def _pf_job(self, layer: int, e: int) -> bytes:
raw = self._pread(layer, e)
with self._lock:
self._pf_cache[(layer, e)] = raw
self._pf_futures.pop((layer, e), None)
return raw
def _evict_window(self, layer: int) -> None:
with self._lock:
if layer not in self._pf_layers:
self._pf_layers.append(layer)
while len(self._pf_layers) > self.window:
old = self._pf_layers.pop(0)
for key in [k for k in self._pf_cache if k[0] == old]:
del self._pf_cache[key]
def wait_prefetch(self) -> None:
with self._lock:
futs = list(self._pf_futures.values())
for f in futs:
f.result()
def read_raw(self, layer: int, expert_ids) -> "list[bytes]":
"""取原始字节:优先用预取缓存/在途结果,未命中的并行 pread仅 IO"""
ids = [int(e) for e in expert_ids]
results: "dict[int, bytes]" = {}
misses = []
for e in ids:
key = (layer, e)
with self._lock:
if key in self._pf_cache:
results[e] = self._pf_cache[key]
self.prefetch_hits += 1
continue
fut = self._pf_futures.get(key)
if fut is not None:
results[e] = fut.result()
self.prefetch_hits += 1
else:
misses.append(e)
def rd(e):
return e, self._pread(layer, e)
if misses:
self.preads += len(misses)
if self.workers <= 1 or len(misses) <= 1:
for e in misses:
results[e] = self._pread(layer, e)
else:
# 复用持久线程池(之前每次调用新建/销毁 ThreadPoolExecutor → 大量线程创建 + GIL 争用)。
for e, raw in self._pf_pool.map(rd, misses):
results[e] = raw
return [results[e] for e in ids]
def load_experts(self, layer: int, expert_ids, view_bf16: bool = True) -> dict:
"""返回 {expert_id: {proj.tensor: mx.array}}。并行读字节lazy 物化 mx.array。
view_bf16=Falsescales/biases uint16供后台线程用避免 .view no Stream
零拷贝(preadv MLX buffer)实测在真实 lazy-eval 流里更慢为拿可写 buffer
per-expert mx.eval 同步代价高于省掉的拷贝(见报告)故保留 lazy frombuffer 路径
"""
ids = [int(e) for e in expert_ids]
raws = self.read_raw(layer, ids)
return {e: self._materialize(raw, view_bf16=view_bf16) for e, raw in zip(ids, raws)}
def load_experts_stacked(self, layer: int, expert_ids, view_bf16: bool = True) -> dict:
"""批量物化:每段只建一个 (N,*shape) mx.arraynp.stack 后单次构造),
取代逐专家 6N frombuffer+mx.array返回 {proj.tensor: (N,*shape)} ids 顺序
load_experts 同源同段表 view_bf16 语义只是把"逐专家逐段"折成
"逐段一次堆叠"demand miss 消费的碎 mx.array 构造从 6N 降到 6
"""
ids = [int(e) for e in expert_ids]
raws = self.read_raw(layer, ids)
out = {}
off = 0
for proj, tensor, dt, shape, nb in self._segs:
np_dt = np.dtype(dt)
cnt = nb // np_dt.itemsize
views = [np.frombuffer(raw, dtype=np_dt, count=cnt, offset=off).reshape(shape)
for raw in raws]
arr = mx.array(np.stack(views, axis=0)) # (N, *shape) 单次构造
if view_bf16 and dt == "uint16" and tensor in ("scales", "biases"):
arr = arr.view(mx.bfloat16)
out[f"{proj}.{tensor}"] = arr
off += nb
return out
def load_experts_native(self, layer: int, expert_ids, view_bf16: bool = False) -> dict:
"""native 物化:用 C++ blob_load 把字节 pread 进 MLX 数组eval 时在 C++ 跑、绕 GIL
Python 侧只建惰性切片/view 无数据拷贝目的是把"拷贝"挪出 GIL减少后台线程
对主线程 per-layer 派发的争用
"""
import mlx_streaming.native_moe_ext as _N
ids = [int(e) for e in expert_ids]
path = os.path.join(self.dir, f"layer{layer:02d}.blob")
raw = _N.blob_load(path, mx.array(ids, dtype=mx.uint32), self.stride) # [n, stride] uint8lazy
out = {}
for i, e in enumerate(ids):
row = raw[i]
d = {}
off = 0
for proj, tensor, dt, shape, nb in self._segs:
seg = row[off:off + nb]
# 按段 dtype 重解释weight→uint32v1 affine 的 uint16 scales/biases
# 可再 view(bf16)mxfp4 的 uint8 scales 保持 uint8。
if dt == "uint32":
arr = seg.view(mx.uint32).reshape(shape)
elif dt == "uint16":
arr = seg.view(mx.uint16).reshape(shape)
if view_bf16:
arr = arr.view(mx.bfloat16)
else: # uint8mxfp4 scales
arr = seg.reshape(shape)
d[f"{proj}.{tensor}"] = arr
off += nb
out[e] = d
return out
def keys(self) -> "list[str]":
return [f"{proj}.{tensor}" for proj, tensor, _, _, _ in self._segs]
def acquire(self, layer: int, expert_ids):
"""对齐 ResidentExpertPool.acquire返回 (pool_arrays, slots)。
pool_arrays: {proj.tensor: (n_uniq,...) 堆叠数组}喂给 _sub.forward
slots: 每个 expert_id unique 列表中的下标= localreshape
"""
ids = [int(e) for e in expert_ids]
uniq = list(dict.fromkeys(ids))
experts = self.load_experts(layer, uniq)
pool = {k: mx.stack([experts[e][k] for e in uniq], axis=0) for k in self.keys()}
pos = {e: i for i, e in enumerate(uniq)}
slots = [pos[e] for e in ids]
return pool, slots
def close(self):
self._pf_pool.shutdown(wait=True)
for fd in self._fds.values():
os.close(fd)
self._fds.clear()

604
mlx_streaming/core/cache/expert_store.py vendored Normal file
View File

@ -0,0 +1,604 @@
"""路线 B 专家后端:每层独立 LRU 缓存 + 按需切片,统计命中率。
与早期版本的两处提速改动
- **每层独立 LRU**每个 MoE 层各自一个 LRU容量 = 每层槽数解码自回归时相邻
token 在同一层常路由到重叠专家per-layer 隔离能显著抬高命中率不再被别层挤掉
- **去掉每次驱逐的 clear_cache 抖动**驱逐只丢引用 MLX 自带缓冲复用 + 外部
set_cache_limit 控制上限避免低命中率下疯狂 malloc/free
stacked[layer] = {"weight": (E,O,I) [, "scales", "biases"]}惰性 eval
fetch(layer, expert_ids) 只取这几个专家堆叠成 (k, ...) 返回
"""
import json
import os
import struct
from collections import OrderedDict, Counter
from typing import Dict, List
import mlx.core as mx
import numpy as np
from mlx_streaming import config
from mlx_streaming.core.mem import clear_cache
from mlx_streaming.core.cache.resident_pool import ResidentExpertPool, _POOL_INIT_SLOTS # noqa: F401
# macOS fcntlF_NOCACHE 提示内核不要把读过的页留在 page cache。
_F_NOCACHE = 48
# safetensors dtype -> (numpy 读入类型, 需 view 的 mlx 类型)。BF16 numpy 无原生类型,
# 先读 uint16 原始位再 .view(bfloat16)。其余类型按需补。
_ST_DTYPE = {
"U32": (np.uint32, None),
"U16": (np.uint16, None),
"I32": (np.int32, None),
"F32": (np.float32, None),
"F16": (np.float16, None),
"BF16": (np.uint16, mx.bfloat16),
}
def _pread_all(fd: int, size: int) -> bytes:
"""把整文件读出(os.pread 可能短读,循环补齐)。"""
chunks = []
off = 0
while off < size:
b = os.pread(fd, size - off, off)
if not b:
break
chunks.append(b)
off += len(b)
return b"".join(chunks)
def load_safetensors_nocache(path: str) -> "Dict[str, mx.array]":
"""用 F_NOCACHE 读 safetensors,返回与 mx.load 等价的 {name: mx.array}。
绕过 OS page cache:每次都真实读 NVMe,基准结果可复现(不被页缓存冷热污染)
mx.load 的数值/dtype 完全一致(U32 权重BF16 scales/biases)
"""
fd = os.open(path, os.O_RDONLY)
try:
try:
import fcntl
fcntl.fcntl(fd, _F_NOCACHE, 1)
except (OSError, ValueError, ImportError):
pass
size = os.fstat(fd).st_size
data = _pread_all(fd, size)
finally:
os.close(fd)
n = struct.unpack("<Q", data[:8])[0]
header = json.loads(data[8:8 + n])
base = 8 + n
out: "Dict[str, mx.array]" = {}
for name, info in header.items():
if name == "__metadata__":
continue
np_dt, view_dt = _ST_DTYPE[info["dtype"]]
s, e = info["data_offsets"]
arr = mx.array(np.frombuffer(data[base + s:base + e], dtype=np_dt).reshape(info["shape"]))
if view_dt is not None:
arr = arr.view(view_dt)
out[name] = arr
return out
class _PerLayerLru:
"""每层一个 OrderedDict 的 LRU 集合,容量为「每层槽数」。"""
def __init__(self, capacity: int, clear_on_evict: bool = False):
self.capacity = capacity
self.clear_on_evict = clear_on_evict
self._caches: "Dict[int, OrderedDict[int, Dict[str, mx.array]]]" = {}
self.hits = 0
self.misses = 0
self.prefetch_hits = 0
self.prefetch_loads = 0
def _layer_cache(self, layer: int) -> "OrderedDict[int, Dict[str, mx.array]]":
c = self._caches.get(layer)
if c is None:
c = OrderedDict()
self._caches[layer] = c
return c
def resident_count(self) -> int:
return sum(len(c) for c in self._caches.values())
def get_or_load(self, layer: int, e: int, loader) -> Dict[str, mx.array]:
cache = self._layer_cache(layer)
if e in cache:
self.hits += 1
cache.move_to_end(e)
else:
self.misses += 1
cache[e] = loader(layer, e)
cache.move_to_end(e)
while len(cache) > self.capacity:
cache.popitem(last=False) # 仅丢引用;不在热路径 clear_cache
if self.clear_on_evict:
clear_cache()
return cache[e]
def hit_rate(self) -> float:
tot = self.hits + self.misses
return self.hits / tot if tot else 0.0
def _stack_picked(picked: List[Dict[str, mx.array]]) -> Dict[str, mx.array]:
keys = picked[0].keys()
return {k: mx.stack([p[k] for p in picked]) for k in keys}
def _stack_cache_size(size: int | None) -> int:
if size is not None:
return max(0, int(size))
return max(0, config.expert_stack_cache())
def _stack_cache_get(cache, size: int, layer: int, key: tuple[int, ...]):
if size <= 0:
return None
layer_cache = cache.get(layer)
if layer_cache is None or key not in layer_cache:
return None
layer_cache.move_to_end(key)
return layer_cache[key]
def _stack_cache_put(cache, size: int, layer: int, key: tuple[int, ...],
value: Dict[str, mx.array]) -> None:
if size <= 0:
return
layer_cache = cache.get(layer)
if layer_cache is None:
layer_cache = OrderedDict()
cache[layer] = layer_cache
layer_cache[key] = value
layer_cache.move_to_end(key)
while len(layer_cache) > size:
layer_cache.popitem(last=False)
class LruExpertStore:
"""内存堆叠后端:从常驻(惰性)堆叠张量按专家切片。容量为「每层槽数」。"""
def __init__(self, stacked: Dict[int, Dict[str, mx.array]], capacity: int,
clear_on_evict: bool = False, stack_cache_size: int | None = None):
self._stacked = stacked
self.capacity = capacity
self._lru = _PerLayerLru(capacity, clear_on_evict)
self.stack_cache_size = _stack_cache_size(stack_cache_size)
self._stack_cache: "Dict[int, OrderedDict[tuple[int, ...], Dict[str, mx.array]]]" = {}
@property
def hits(self) -> int:
return self._lru.hits
@property
def misses(self) -> int:
return self._lru.misses
def resident_count(self) -> int:
return self._lru.resident_count()
def cap_for(self, layer: int) -> int:
return self.capacity
def _load_one(self, layer: int, e: int) -> Dict[str, mx.array]:
src = self._stacked[layer]
out = {}
for k, arr in src.items():
sub = arr[e] # 切第 e 个专家
mx.eval(sub) # 显式物化这一个专家
out[k] = sub
return out
def fetch(self, layer: int, expert_ids: List[int]) -> Dict[str, mx.array]:
key = tuple(int(e) for e in expert_ids)
cached = _stack_cache_get(self._stack_cache, self.stack_cache_size, layer, key)
if cached is not None:
for e in key:
self._lru.get_or_load(layer, e, self._load_one)
return cached
picked = [self._lru.get_or_load(layer, int(e), self._load_one) for e in expert_ids]
stacked = _stack_picked(picked)
_stack_cache_put(self._stack_cache, self.stack_cache_size, layer, key, stacked)
return stacked
def hit_rate(self) -> float:
return self._lru.hit_rate()
class FileExpertStore:
"""文件后端专家缓存:从离线拆分的 per-expert safetensors 按需加载,每层独立 LRU。
文件命名{root}/layer{layer:02d}_expert{e:03d}.safetensors
内容为扁平 dict键形如 "gate_proj.weight"/"gate_proj.scales"/"up_proj.weight"...
capacity 语义为每层槽数worst-case 常驻 capacity × MoE 层数
"""
def __init__(self, root: str, capacity: int, clear_on_evict: bool = False,
record: bool = False, stack_cache_size: int | None = None,
layer_caps: "Dict[int, int] | None" = None):
self.root = root
self.capacity = capacity
self._lru = _PerLayerLru(capacity, clear_on_evict)
self.stack_cache_size = _stack_cache_size(stack_cache_size)
self._stack_cache: "Dict[int, OrderedDict[tuple[int, ...], Dict[str, mx.array]]]" = {}
# ③ 热专家常驻:钉住的专家永不进 LRU、永不驱逐
self._pinned: "Dict[int, Dict[int, Dict[str, mx.array]]]" = {}
self.pinned_hits = 0
# 可选 per-layer bundle把同层所有 per-expert safetensors 合并成一个文件,
# 运行时每层只 mx.load 一次,减少大量小文件 open/header parse 开销。
self.use_bundles = config.expert_bundle()
self.bundle_dir = config.expert_bundle_dir(os.path.join(root, "layer_bundles"))
self.bundle_cache_size = config.expert_bundle_cache()
self._bundle_cache: "OrderedDict[int, Dict[str, mx.array]]" = OrderedDict()
self._bundle_key_cache: "Dict[int, Dict[int, list[tuple[str, str]]]]" = {}
self.bundle_loads = 0
# 连续常驻池后端acquire 路径用),与 LRU stack 路径共享 _load_one
# layer_caps每层独立池容量(profile 驱动的无损省内存),缺省层用全局 capacity
# demand miss 批量加载接线(三选一,acquire 优先 stacked > batch > 逐专家):
# - native(默认,且非 async_prefetch):C++ blob_load 直进 MLX 数组、绕 GIL、惰性,返回
# 逐专家 {e:{k}} → 走 _place_experts。native 不消费 async prefetch buffer,故 gate 在
# not async_prefetch(async 路径改用 _batch_load_resident,它会先吃 buffer)。
# - batch_miss_read(默认):numpy 批量堆叠,每段一次 mx.array(取代逐专家 6N 次构造)。
# - 都关:逐专家 loader。
if config.native_demand_loader() and not config.async_prefetch():
_batch_loader = self._batch_load_native
_stacked_loader = None
elif config.batch_miss_read():
_batch_loader = self._batch_load_resident
_stacked_loader = (self._batch_load_stacked
if not config.async_prefetch() else None)
else:
_batch_loader = None
_stacked_loader = None
self._resident = ResidentExpertPool(
capacity, loader=self._load_one_resident, layer_caps=layer_caps,
batch_loader=_batch_loader, stacked_batch_loader=_stacked_loader)
# 激活频率统计(校准阶段用)
self.record = record
self._counts: "Dict[int, Counter]" = {}
self.async_prefetch = config.async_prefetch()
self.prefetch_buffer_size = config.prefetch_buffer_experts()
self._prefetch_buffer: "OrderedDict[tuple[int, int], Dict[str, mx.array]]" = OrderedDict()
self.prefetch_buffer_hits = 0
self.prefetch_submitted = 0
self.prefetch_dropped = 0
# dual-sourcedecode 走 staging 双源命中计数(不写池的零拷贝命中)
self.staging_hits = 0
# 可选 blob miss-loaderSTREAM_BLOB_LOADER=1由 model_builder 注入):
# miss 时用每专家连续 blob 的并行 pread 取代 per-expert safetensors 的 mx.load
# 复用常驻池 GPU-remap 快路径(命中零编排),小 EXPERT_SLOTS 即低内存。
self._blob_loader = None
# 可选后台预取器STREAM_BLOB_BG=1由 model_builder 注入):在独立 stream 上
# 提前物化预测专家promote_prefetched 把它们写进常驻池槽(主线程、便宜 scatter
self._bg = None
def cap_for(self, layer: int) -> int:
"""该层常驻池容量(profile 指定则用之,否则全局 capacity)。"""
return self._resident.cap_for(layer)
def resident_experts(self, layer: int) -> set[int]:
return self._resident.resident_experts(layer)
def resident_lru_scores(self, layer: int) -> dict[int, float]:
return self._resident.resident_lru_scores(layer)
@property
def hits(self) -> int:
return self._lru.hits + self.pinned_hits + self._resident.hits
@property
def misses(self) -> int:
return self._lru.misses + self._resident.misses
def resident_count(self) -> int:
pinned = sum(len(d) for d in self._pinned.values())
return self._lru.resident_count() + pinned
def pinned_count(self) -> int:
return sum(len(d) for d in self._pinned.values())
def path(self, layer: int, e: int) -> str:
import os
return os.path.join(self.root, f"layer{layer:02d}_expert{e:03d}.safetensors")
def bundle_path(self, layer: int) -> str:
return os.path.join(self.bundle_dir, f"layer{layer:02d}.safetensors")
def _load_layer_bundle(self, layer: int) -> Dict[str, mx.array]:
cached = self._bundle_cache.get(layer)
if cached is not None:
self._bundle_cache.move_to_end(layer)
return cached
bundle = mx.load(self.bundle_path(layer))
self.bundle_loads += 1
self._bundle_cache[layer] = bundle
self._bundle_cache.move_to_end(layer)
key_map: "Dict[int, list[tuple[str, str]]]" = {}
for full_key in bundle:
if not full_key.startswith("expert"):
continue
head, rest = full_key.split(".", 1)
e = int(head[len("expert"):])
key_map.setdefault(e, []).append((rest, full_key))
self._bundle_key_cache[layer] = key_map
while len(self._bundle_cache) > self.bundle_cache_size:
old_layer, _ = self._bundle_cache.popitem(last=False)
self._bundle_key_cache.pop(old_layer, None)
return bundle
def _raw_load_one(self, layer: int, e: int) -> Dict[str, mx.array]:
if self._blob_loader is not None:
# blob 单专家1 次连续 pread + 物化,结构与 mx.load 的 per-expert dict 一致。
return self._blob_loader.load_experts(int(layer), [int(e)])[int(e)]
if self.use_bundles and os.path.exists(self.bundle_path(layer)):
bundle = self._load_layer_bundle(layer)
keys = self._bundle_key_cache.get(layer, {}).get(int(e), [])
return {short: bundle[full] for short, full in keys}
# 默认走 F_NOCACHE 读取(绕过 OS page cache):基准每次真实读盘、可复现。
# mx.load 用 mmap 会经页缓存,重复跑被缓存命中污染→速度飘。EXPERT_NOCACHE=0 切回 mmap。
if config.expert_nocache():
w = load_safetensors_nocache(self.path(layer, e))
else:
w = mx.load(self.path(layer, e))
# 默认惰性:不在此 eval,让读盘并入 MLX 异步图、与计算重叠,token 末统一 eval。
# 旧版每专家强制 mx.eval(w) 会每 token 触发 ~42 次同步、打碎流水线(实测 ~2.4× 慢)。
# 数值完全等价(惰性求值仍精确)。EAGER_EXPERT_LOAD=1 可恢复旧的逐专家强制物化。
if config.eager_expert_load():
mx.eval(w)
return w
def _load_one(self, layer: int, e: int) -> Dict[str, mx.array]:
return self._raw_load_one(layer, e)
def _load_one_resident(self, layer: int, e: int) -> Dict[str, mx.array]:
"""resident demand loader优先消费 async prefetch buffer/inflight。"""
key = (int(layer), int(e))
if self.async_prefetch:
cached = self._prefetch_buffer.pop(key, None)
if cached is not None:
self.prefetch_buffer_hits += 1
self._resident.misses = max(0, self._resident.misses - 1)
return cached
return self._raw_load_one(layer, e)
def _batch_load_resident(self, layer: int, ids: List[int]) -> Dict[int, Dict[str, mx.array]]:
"""批量 resident demand loader先消费 async prefetch buffer剩余 miss 一次并行读。
关键blob 路径用一次 load_experts(layer, rest)8-worker 并行 pread取代逐专家
串行 pread实测串行 6.4GB/s 并行 8 22GB/s结果与逐专家加载等价
"""
ids = [int(e) for e in ids]
out: "Dict[int, Dict[str, mx.array]]" = {}
rest: List[int] = []
for e in ids:
if self.async_prefetch:
cached = self._prefetch_buffer.pop((int(layer), e), None)
if cached is not None:
self.prefetch_buffer_hits += 1
self._resident.misses = max(0, self._resident.misses - 1)
out[e] = cached
continue
rest.append(e)
if rest:
if self._blob_loader is not None:
out.update(self._blob_loader.load_experts(int(layer), rest))
else:
for e in rest:
out[e] = self._raw_load_one(layer, e)
return out
def _batch_load_stacked(self, layer: int, ids: List[int]) -> Dict[str, mx.array]:
"""批量+预堆叠 demand loaderblob 路径一次 load_experts_stacked并行 pread + 每段一次
物化返回 {k:(N,*shape)} ids 顺序 blob loader 时退化为逐专家物化后堆叠
仅在 async_prefetch 关闭时接入不消费 prefetch buffer保证不漏 buffer 命中
"""
ids = [int(e) for e in ids]
if self._blob_loader is not None:
return self._blob_loader.load_experts_stacked(int(layer), ids)
per = {e: self._raw_load_one(layer, e) for e in ids}
keys = list(per[ids[0]].keys())
return {k: mx.stack([per[e][k] for e in ids], axis=0) for k in keys}
def _batch_load_native(self, layer: int, ids: List[int]) -> "Dict[int, Dict[str, mx.array]]":
"""批量 native demand loader:blob 路径用 C++ blob_load 把整批字节 pread 直进 MLX 数组
(eval 时在 C++ GIL惰性),返回逐专家 {e:{proj.tensor:arr}}( _place_experts)
相比 numpy 路径(frombuffer+stack+mx.array 同步拷贝),把字节物化从 acquire 挪到 token
eval并绕开 GIL blob loader 时退回逐专家 raw 物化(行为等价)view_bf16=True
兼容 affine 量化(uint16 scales/biasesbf16);mxfp4 uint8 scales 忽略此参数
"""
ids = [int(e) for e in ids]
if self._blob_loader is not None:
return self._blob_loader.load_experts_native(int(layer), ids, view_bf16=True)
return {e: self._raw_load_one(layer, e) for e in ids}
def note(self, layer: int, expert_ids: List[int]) -> None:
"""记录一次激活(校准阶段调用),用于挑热专家。"""
c = self._counts.get(layer)
if c is None:
c = Counter()
self._counts[layer] = c
for e in expert_ids:
c[int(e)] += 1
def recorded_layers(self) -> List[int]:
return list(self._counts.keys())
def hot(self, layer: int, h: int) -> List[int]:
"""返回该层激活最频繁的 h 个专家 id。"""
c = self._counts.get(layer)
if not c:
return []
return [e for e, _ in c.most_common(h)]
def pin(self, layer: int, expert_ids: List[int]) -> None:
"""把这些专家加载并钉为常驻(不计入 LRU、不可驱逐"""
d = self._pinned.setdefault(layer, {})
for e in expert_ids:
e = int(e)
if e not in d:
d[e] = self._load_one(layer, e)
# 连续 resident pool 是 decode/MTP 热路径;同步预取进去并标记不可驱逐,
# 否则旧 _pinned 只会影响 fetch/stack 路径,无法降低 acquire_gpu 的 miss。
self._resident.pin_loaded(layer, {int(e): d[int(e)] for e in expert_ids})
def pin_batch(self, layer: int, expert_ids: List[int]) -> int:
"""批量 pinAUTOPIN 预热用):语义同 pin但加载收口为一次批量读返回 pin 数。
加载路径 blob_loader生产唯一专家源 load_experts 一次 8-worker 并行 pread
取代 pin 的逐专家串行 per-expert safetensors部署里 per-expert 文件已删串行路径
会直接失败 blob_loader 退回逐专家 _load_one兼容旧 per-expert 部署
dual(native demand)模式真实区由 C++ g_real 权威 pin_loaded_dual 注册 real_pin +
C++ 直写池行不留存数组副本避免 pinned 专家双倍内存fetch 路径不受影响
"""
layer = int(layer)
ids = list(dict.fromkeys(int(e) for e in expert_ids))
if not ids:
return 0
if getattr(self._resident, "_native_demand", False):
if self._blob_loader is not None:
loaded = self._blob_loader.load_experts(layer, ids)
else:
loaded = {e: self._load_one(layer, e) for e in ids}
return self._resident.pin_loaded_dual(layer, loaded)
d = self._pinned.setdefault(layer, {})
todo = [e for e in ids if e not in d]
if todo:
if self._blob_loader is not None:
d.update(self._blob_loader.load_experts(layer, todo))
else:
for e in todo:
d[e] = self._load_one(layer, e)
self._resident.pin_loaded(layer, {e: d[e] for e in ids})
return len(ids)
def prefetch(self, layer: int, expert_ids: List[int]) -> None:
"""预取到连续 resident pool不计入 demand miss/hit。"""
if not self.async_prefetch:
self._resident.prefetch(layer, expert_ids)
return
resident = self._resident.resident_experts(layer)
for e in dict.fromkeys(int(x) for x in expert_ids):
key = (int(layer), int(e))
if e in resident:
self._resident.prefetch_hits += 1
continue
if key in self._prefetch_buffer:
self._resident.prefetch_hits += 1
continue
self._prefetch_buffer[key] = self._raw_load_one(int(layer), int(e))
self._prefetch_buffer.move_to_end(key)
self.prefetch_submitted += 1
self._resident.prefetch_loads += 1
while len(self._prefetch_buffer) > self.prefetch_buffer_size:
self._prefetch_buffer.popitem(last=False)
self.prefetch_dropped += 1
def promote_prefetched(self, layer: int) -> int:
"""把后台预取器中该层所有「已物化」专家写进常驻池槽(主线程、默认 stream
同层(AHEAD=0)预取的就是本层将用到的专家应真正放进池槽以转 miss hit
互保护置入时把本批就绪的全部专家当作 current使它们彼此不被驱逐
只挤掉池里陈旧的非本批专家池满且无可驱逐者时跳过剩余留给 demand 回退
返回写入的专家数
"""
if self._bg is None:
return 0
ready = self._bg.take_ready_layer(int(layer))
self._bg.note_promote(int(layer), len(ready))
if not ready:
return 0
self._resident._ensure_layer(int(layer))
slot_of = self._resident._slot_of[int(layer)]
protect = {int(e) for e in ready}
placed = 0
for e, d in ready.items():
e = int(e)
if e in slot_of:
continue
try:
self._resident._place_expert(int(layer), e, d, current=protect)
except ValueError:
# 池满且无非本批可驱逐 → 剩余专家留给 demand 路径
break
placed += 1
return placed
def wait_prefetch(self) -> None:
"""测试/诊断用:等待当前所有 async prefetch 完成。"""
return
def release_fetch_cache(self) -> int:
"""释放 fetch/stack 路径prefill 用)的常驻缓存,返回释放的专家数。
decode token acquire_gpu/_resident 永不再碰 _lru/_stack_cache
prefill 结束后这套缓存每层 cap 个专家 + 堆叠张量是纯死重约占
cap×MoE层数 个专家的内存实测 cap=16 ~9.2GB prefilldecode
边界调用一次即可把这块预算还给系统/MTP token 会按需重新填充
"""
freed = self._lru.resident_count()
self._lru._caches.clear()
self._stack_cache.clear()
clear_cache()
return freed
def reset_stats(self) -> None:
self._lru.hits = 0
self._lru.misses = 0
self.pinned_hits = 0
self._resident.hits = 0
self._resident.misses = 0
self._resident.gpu_fastpath = 0
self._resident.gpu_fallback = 0
self._resident.prefetch_hits = 0
self._resident.prefetch_loads = 0
self.prefetch_buffer_hits = 0
self.prefetch_submitted = 0
self.prefetch_dropped = 0
self.staging_hits = 0
def acquire(self, layer: int, expert_ids: List[int]):
"""连续常驻池取专家:返回 (pool_arrays, slots)。命中零拷贝miss 单槽写。"""
if self.record:
self.note(layer, expert_ids)
return self._resident.acquire(layer, expert_ids)
def acquire_gpu(self, layer: int, inds: mx.array, num_experts: int):
"""GPU 侧 slot 重映射(decode 热路径):命中零 host 往返。record 模式才回 CPU 记账。"""
if self.record:
self.note(layer, [int(i) for i in inds.reshape(-1).tolist()])
return self._resident.acquire_gpu(layer, inds, num_experts)
def fetch(self, layer: int, expert_ids: List[int]) -> Dict[str, mx.array]:
if self.record:
self.note(layer, expert_ids)
key = tuple(int(e) for e in expert_ids)
cached = _stack_cache_get(self._stack_cache, self.stack_cache_size, layer, key)
if cached is not None:
pinned = self._pinned.get(layer)
for e in key:
if pinned is not None and e in pinned:
self.pinned_hits += 1
else:
self._lru.get_or_load(layer, e, self._load_one)
return cached
pinned = self._pinned.get(layer)
picked = []
for e in expert_ids:
e = int(e)
if pinned is not None and e in pinned:
self.pinned_hits += 1
picked.append(pinned[e])
else:
picked.append(self._lru.get_or_load(layer, e, self._load_one))
stacked = _stack_picked(picked)
_stack_cache_put(self._stack_cache, self.stack_cache_size, layer, key, stacked)
return stacked
def hit_rate(self) -> float:
tot = self.hits + self.misses
return self.hits / tot if tot else 0.0

View File

@ -0,0 +1,116 @@
"""把 Qwen3-Next 的 12 个全注意力层就地切换到 K4/V3 旋转量化 KV。
做法(隔离低侵入):
- 子类 `_RotatedQuantAttn` 仅重写 `__call__`:RoPE 后对 q/k 旋转(R_k)v 旋转(R_v),
存非对称量化 cache, asym_quantized_sdpa,输出做 V 逆旋转复原保参数树不变
(只换 `__class__` object.__setattr__ 挂常量,不进 nn.Module 的参数/子模块树)
- 覆盖 `model.make_cache`:全注意力层AsymmetricQuantizedKVCache,线性层ArraysCache(size=2)
"""
from __future__ import annotations
import mlx.core as mx
from mlx_lm.models.cache import ArraysCache
from mlx_lm.models.qwen3_next import Qwen3NextAttention
from mlx_streaming.core.cache.quant_kv import (
AsymmetricQuantizedKVCache,
asym_quantized_sdpa,
build_block_so4,
rotate_last,
)
class _RotatedQuantAttn(Qwen3NextAttention):
"""与原 Qwen3NextAttention.__call__ 数值等价(高 bit 时),叠加旋转 + 非对称量化。"""
def __call__(self, x, mask=None, cache=None):
B, L, D = x.shape
q_proj_output = self.q_proj(x)
queries, gate = mx.split(
q_proj_output.reshape(B, L, self.num_attention_heads, -1), 2, axis=-1)
gate = gate.reshape(B, L, -1)
keys, values = self.k_proj(x), self.v_proj(x)
queries = self.q_norm(queries).transpose(0, 2, 1, 3)
keys = self.k_norm(
keys.reshape(B, L, self.num_key_value_heads, -1)).transpose(0, 2, 1, 3)
values = values.reshape(
B, L, self.num_key_value_heads, -1).transpose(0, 2, 1, 3)
if cache is not None:
queries = self.rope(queries, offset=cache.offset)
keys = self.rope(keys, offset=cache.offset)
else:
queries = self.rope(queries)
keys = self.rope(keys)
# 旋转去相关(RoPE 之后):q、k 用 R_k(分数自动抵消),v 用 R_v(输出后逆旋转)。
Rk = self._kvq_Rk
Rv = self._kvq_Rv
if Rk is not None:
queries = rotate_last(queries, Rk)
keys = rotate_last(keys, Rk)
values = rotate_last(values, Rv)
if cache is not None:
keys, values = cache.update_and_fetch(keys, values)
output = asym_quantized_sdpa(
queries, keys, values, scale=self.scale, mask=mask,
group_size=self._kvq_gs, k_bits=self._kvq_kb, v_bits=self._kvq_vb)
else:
# 无 cache(罕见,如纯前向探针):走稠密 SDPA,V 仍在旋转空间,下方逆旋转复原。
output = mx.fast.scaled_dot_product_attention(
queries, keys, values, scale=self.scale, mask=mask)
if Rv is not None:
output = rotate_last(output, Rv.T)
output = output.transpose(0, 2, 1, 3).reshape(B, L, -1)
return self.o_proj(output * mx.sigmoid(gate))
def _inner(model):
"""取到含 .layers 的内层模块(mlx_lm Model 包了一层 .model)。"""
return model.model if hasattr(model, "model") else model
def patch_kv_quant(model, *, group_size=64, k_bits=4, v_bits=3, rotate=True, seed=0):
"""就地把所有全注意力层切到旋转 + 非对称量化 KV,并覆盖 make_cache。返回 model。"""
inner = _inner(model)
layers = inner.layers
head_dim = None
for l in layers:
if not l.is_linear:
head_dim = l.self_attn.head_dim
break
if head_dim is None:
return model # 没有全注意力层,无需 patch
Rk = build_block_so4(head_dim, seed=seed) if rotate else None
Rv = build_block_so4(head_dim, seed=seed + 1) if rotate else None
for l in layers:
if l.is_linear:
continue
attn = l.self_attn
# 常量挂载用 object.__setattr__,避免进 nn.Module 参数/子模块树(否则破坏权重 save/树遍历)。
object.__setattr__(attn, "_kvq_Rk", Rk)
object.__setattr__(attn, "_kvq_Rv", Rv)
object.__setattr__(attn, "_kvq_gs", group_size)
object.__setattr__(attn, "_kvq_kb", k_bits)
object.__setattr__(attn, "_kvq_vb", v_bits)
attn.__class__ = _RotatedQuantAttn
def make_cache():
return [
AsymmetricQuantizedKVCache(group_size, k_bits, v_bits)
if not l.is_linear else ArraysCache(size=2)
for l in layers
]
# 覆盖实例方法(plain function,nn.Module.__setattr__ 不会把函数纳入参数树)。
model.make_cache = make_cache
return model

223
mlx_streaming/core/cache/quant_kv.py vendored Normal file
View File

@ -0,0 +1,223 @@
"""IsoQuant 风格 K4/V3 非对称量化 KV cache(仅用于全注意力层)。
设计要点(详见 docs/superpowers/specs/2026-06-30-isoquant-kv-quant-design.md):
- SO(4) 块对角旋转去相关:head_dim=256 切成 64 4D ,每块由两个单位四元数构造
一个 SO(4) 旋转(left·right_conj sandwich)旋转正交,在注意力分数里自动抵消
(qk 同旋转 q·kᵀ 不变); V 旋转后需在 SDPA 输出上做逆旋转复原
- 非对称仿射量化:K k_bits(默认 4)V v_bits(默认 3),分别用 mx.quantize 打包,
attention mx.quantized_matmul 快路径(每次调用可传不同 bits),无需自写 Metal
打包宽度用 dim*bits//32( 3-bit 必须如此:mlx_lm 自带 8*4//bits 公式对 3-bit 会算错)
"""
from __future__ import annotations
import mlx.core as mx
import numpy as np
from mlx.utils import tree_map, tree_reduce
# 注意:make_mask 用的是 cache 模块里的 create_attention_mask(签名含 offset/return_array),
# 与 base 模块同名但不同签名的那个不可混用。
from mlx_lm.models.cache import create_attention_mask
# ----------------------------- SO(4) 块旋转 -----------------------------
def _quat_left(q: np.ndarray) -> np.ndarray:
"""单位四元数 q 的左乘矩阵 L(q):把 v 映射到 q*v(Hamilton 积)。"""
w, x, y, z = q
return np.array([
[w, -x, -y, -z],
[x, w, -z, y],
[y, z, w, -x],
[z, -y, x, w],
], dtype=np.float32)
def _quat_right_conj(q: np.ndarray) -> np.ndarray:
"""单位四元数 q 的"右乘其共轭"矩阵 R(conj q):把 v 映射到 v*conj(q)。
L(qL)·R(conj qR) SO(4) 的双边旋转 v -> qL * v * conj(qR),覆盖全部 4D 旋转
"""
w, x, y, z = q
return np.array([
[ w, x, y, z],
[-x, w, z, -y],
[-y, -z, w, x],
[-z, y, -x, w],
], dtype=np.float32)
def build_block_so4(head_dim: int, seed: int = 0, blocks_of: int = 4) -> mx.array:
"""构造 head_dim×head_dim 的块对角正交矩阵,每 blocks_of(=4)维一个 SO(4) 旋转。
data-oblivious:仅由 seed 决定,固定可复现,不依赖数据分布
"""
assert head_dim % blocks_of == 0, "head_dim 必须被 blocks_of 整除"
rng = np.random.default_rng(seed)
M = np.zeros((head_dim, head_dim), dtype=np.float32)
for b in range(head_dim // blocks_of):
qL = rng.standard_normal(4).astype(np.float32)
qL /= np.linalg.norm(qL)
qR = rng.standard_normal(4).astype(np.float32)
qR /= np.linalg.norm(qR)
blk = _quat_left(qL) @ _quat_right_conj(qR) # 4×4 ∈ SO(4)
i = b * blocks_of
M[i:i + blocks_of, i:i + blocks_of] = blk
return mx.array(M)
def rotate_last(x: mx.array, R: mx.array) -> mx.array:
"""对末维做正交旋转:x[..., D] @ R[D, D](按 x 的 dtype 计算)。"""
return x @ R.astype(x.dtype)
# --------------------- 非对称量化 KV cache(K4/V3)---------------------
def _packed_words(dim: int, bits: int) -> int:
"""量化后每行占的 uint32 个数:dim*bits//32(对 2/3/4/5/6/8 bit 均正确)。"""
return dim * bits // 32
class AsymmetricQuantizedKVCache:
"""K、V 各用独立位宽的量化 KV cache(仿 mlx_lm.QuantizedKVCache,但 K/V 不同 bits)。
keys/values 各存为 (wq, scales, biases) 三元组, mx.quantized_matmul 接口对齐
"""
step = 256
def __init__(self, group_size: int = 64, k_bits: int = 4, v_bits: int = 3):
self.keys = None
self.values = None
self.offset = 0
self.group_size = group_size
self.k_bits = k_bits
self.v_bits = v_bits
def _init_quant(self, dim, bits, B, H, steps, dt):
packed = _packed_words(dim, bits)
ngroups = dim // self.group_size
return (
mx.zeros((B, H, steps, packed), dtype=mx.uint32),
mx.zeros((B, H, steps, ngroups), dtype=dt),
mx.zeros((B, H, steps, ngroups), dtype=dt),
)
def update_and_fetch(self, keys, values):
B, H, n, kD = keys.shape
vD = values.shape[-1]
prev = self.offset
if self.keys is None or (prev + n) > self.keys[0].shape[-2]:
new_steps = (self.step + n - 1) // self.step * self.step
if self.keys is not None:
if prev % self.step != 0:
self.keys = tree_map(lambda x: x[..., :prev, :], self.keys)
self.values = tree_map(lambda x: x[..., :prev, :], self.values)
def _expand(x):
z = mx.zeros((B, H, new_steps, x.shape[-1]), dtype=x.dtype)
return mx.concatenate([x, z], axis=-2)
self.keys = tree_map(_expand, self.keys)
self.values = tree_map(_expand, self.values)
else:
self.keys = self._init_quant(kD, self.k_bits, B, H, new_steps, keys.dtype)
self.values = self._init_quant(vD, self.v_bits, B, H, new_steps, values.dtype)
self.offset += n
kq = mx.quantize(keys, group_size=self.group_size, bits=self.k_bits)
vq = mx.quantize(values, group_size=self.group_size, bits=self.v_bits)
for i in range(3):
self.keys[i][..., prev:self.offset, :] = kq[i]
self.values[i][..., prev:self.offset, :] = vq[i]
return (
tree_map(lambda x: x[..., :self.offset, :], self.keys),
tree_map(lambda x: x[..., :self.offset, :], self.values),
)
def make_mask(self, *args, **kwargs):
return create_attention_mask(*args, offset=self.offset, **kwargs)
def is_trimmable(self):
return True
def trim(self, n):
n = min(self.offset, n)
self.offset -= n
return n
def empty(self):
return self.keys is None
@property
def nbytes(self):
if self.keys is None:
return 0
return tree_reduce(lambda a, x: a + x.nbytes, (self.keys, self.values), 0)
# ---- 快照接口(与 KVCache 同口径,供 mtp.kv_cache._snapshot/_restore 与前缀快照使用)----
@property
def state(self):
# buffer 未满时切到 offset(MLX 切片产生新数组,快照不会被后续原地写污染)
if self.offset == self.keys[0].shape[-2]:
return self.keys, self.values
return (
tree_map(lambda x: x[..., : self.offset, :], self.keys),
tree_map(lambda x: x[..., : self.offset, :], self.values),
)
@state.setter
def state(self, v):
self.keys, self.values = v
self.offset = self.keys[0].shape[-2]
@property
def meta_state(self):
return str(self.group_size), str(self.k_bits), str(self.v_bits)
@meta_state.setter
def meta_state(self, v):
self.group_size, self.k_bits, self.v_bits = map(int, v)
# ----------------- 非对称量化 SDPA(K 用 k_bits、V 用 v_bits)-----------------
def asym_quantized_sdpa(queries, q_keys, q_values, scale, mask,
group_size, k_bits, v_bits):
"""基于 mlx_lm.quantized_scaled_dot_product_attention,K/V 各自传 bits。
queries: [B, n_q_heads, L, D];q_keys/q_values: 量化三元组(n_kv_heads)
返回 [B, n_q_heads, L, D](仍在旋转后的 V 空间,调用方负责逆旋转)
"""
B, n_q_heads, L, D = queries.shape
n_kv_heads = q_keys[0].shape[-3]
n_repeats = n_q_heads // n_kv_heads
queries = queries * scale
if n_repeats > 1:
queries = mx.reshape(queries, (B, n_kv_heads, n_repeats, L, D))
q_keys = tree_map(lambda x: mx.expand_dims(x, axis=-3), q_keys)
q_values = tree_map(lambda x: mx.expand_dims(x, axis=-3), q_values)
scores = mx.quantized_matmul(
queries, *q_keys, transpose=True, group_size=group_size, bits=k_bits)
if mask is not None:
if isinstance(mask, str):
qL, kL = scores.shape[-2:]
q_indices = mx.arange(kL - qL, kL)
k_indices = mx.arange(kL)
mask = q_indices[:, None] >= k_indices[None]
if mask.dtype == mx.bool_:
scores = mx.where(mask, scores, mx.finfo(scores.dtype).min)
else:
scores = scores + mask
scores = mx.softmax(scores, axis=-1, precise=True)
out = mx.quantized_matmul(
scores, *q_values, transpose=False, group_size=group_size, bits=v_bits)
if n_repeats > 1:
out = mx.reshape(out, (B, n_q_heads, L, D))
return out

View File

@ -0,0 +1,635 @@
"""常驻专家池:每层一块连续 GPU 张量 + slot LRU支持按需增长、pin、GPU 侧重映射。
设计要点
- 命中只返回槽位不写池miss 只把单个专家原地写进它的槽位`_write_slot`避免拷贝整池
- 池按需增长grow-on-demand起步 `_POOL_INIT_SLOTS` 工作集扩大时 ~1.5× 增至天花板
`cap_for(layer)` 才开始 LRU 淘汰默认无 profile 即自动右尺寸容量内永不超预算
- `acquire_gpu`decode 热路径用 GPU 查找表做纯 GPU slot 重映射命中层零 host 往返
仅一次 miss 标志同步 miss 才回退 host 读盘路径
"""
import os
from collections import OrderedDict, Counter
from typing import Dict, List
import mlx.core as mx
# 诊断门控staging/acquire_gpu 路径消费侧字节级真值校验(混合 ahead 损坏取证用,默认关)。
# 在 acquire_gpu 全命中快路径返回前,对本层真实路由命中的每个专家,把池槽字节与磁盘真值
# 逐 key 比对;不一致即「池槽装错字节」铁证。受 STG_VERIFY=1 控制,对主路径零影响(默认 off
_STG_VERIFY = os.environ.get("STG_VERIFY") == "1"
_stg_verify_state = {"ok": 0, "bad": 0, "printed": 0, "calls": 0, "first_bad_call": None}
from mlx_streaming import config
from mlx_streaming.core.cache import autopin
# Route 3 Phase 1 底座spec/dual 模式下池 buffer 改由 C++ 拥有(mx::allocator + no-op deleter)
# 地址进程内恒定、永不被 MLX donation/迁移,供侧区/demand 后台 pread 安全直写(替代消费侧 MLX scatter
# POOL_OWNED=0 可强制回退 mx.zeros(仅供 A/B 对照)。
_POOL_OWNED = os.environ.get("POOL_OWNED", "1") == "1"
# mx.Dtype -> C++ pool_owned_zeros 接受的 dtype 名。
_DTYPE_NAME = {
mx.uint32: "uint32", mx.uint16: "uint16", mx.uint8: "uint8",
mx.int32: "int32", mx.int16: "int16",
mx.bfloat16: "bfloat16", mx.float16: "float16", mx.float32: "float32",
}
def _owned_pool(sample: "Dict[str, mx.array]", n: int) -> "Dict[str, mx.array]":
"""用 C++-owned buffer 建 (n,*shape) 的 per-key 池数组(地址恒定,供 C++ 直写)。"""
import mlx_streaming.native_moe_ext as _N
out = {}
for k, v in sample.items():
name = _DTYPE_NAME.get(v.dtype)
if name is None:
raise RuntimeError(f"pool_owned_zeros 不支持 dtype {v.dtype} (key={k})")
out[k] = _N.pool_owned_zeros([int(n)] + [int(d) for d in v.shape], name)
return out
# 池按需增长(grow-on-demand)的初始物理行数。
# 起步小,工作集扩大时按 ~1.5× 增长,封顶 cap_for(默认=全局 capacity)。
# 好处:无需 profile 也能自动右尺寸(内存≈实际工作集),且对任何 prompt 自适应、容量内永不超预算。
_POOL_INIT_SLOTS = 16
class ResidentExpertPool:
"""每层一个连续常驻池:(capacity,*shape) 张量 + slot LRU。
命中只返回槽位不写池miss 只把单个专家原地写进它的槽位_write_slot
loader(layer, e) -> Dict[str, mx.array]单个专家的参数未堆叠
"""
def __init__(self, capacity: int, loader, layer_caps: "Dict[int, int] | None" = None,
spec_slots: int = 0, batch_loader=None, stacked_batch_loader=None,
spec_gens: int = 1):
self.capacity = capacity
self.loader = loader
# 可选批量加载器 batch_loader(layer, [ids]) -> {e: expert_dict}acquire 用它把本层所有
# miss 一次并行读8-worker pread取代逐专家串行 loader。None 时退回串行 loader。
self.batch_loader = batch_loader
# 可选「批量+预堆叠」加载器 stacked_batch_loader(layer, [ids]) -> {k:(N,*shape)}
# 在 batch_loader 基础上再把 6N 次碎 mx.array 物化折成每段一次,写池时直接整批 scatter
# 既省 frombuffer/mx.array 构造又省 _write_slots_batch 里的 mx.stack。优先级最高。
self.stacked_batch_loader = stacked_batch_loader
self.spec_slots = int(spec_slots) # >0侧区模式预分配 cap+spec、禁 grow
self.spec_gens = max(1, int(spec_gens)) # 侧区代数:双缓冲=2物理行=cap+spec_gens*spec_slots
# cap_for(layer) 是该层物理行数的「天花板」profile 指定则用之,否则=全局 capacity。
# 池按需增长(grow-on-demand):起步小、随工作集扩大增至天花板才开始 LRU 淘汰。
# 因此默认(无 profile)即自动右尺寸——内存≈实际工作集,对任意 prompt 自适应,
# 且天花板内永不超预算(超预算只是更慢,不崩)。增长偶发(预热期),稳态零拷贝。
self.layer_caps: "Dict[int, int]" = dict(layer_caps or {})
self._pools: "Dict[int, Dict[str, mx.array]]" = {}
self._slot_of: "Dict[int, OrderedDict[int, int]]" = {}
self._free: "Dict[int, list]" = {}
# 每层当前物理已分配行数(按需增长,≤ cap_for)。grow-on-demand 的核心状态。
self._alloc: "Dict[int, int]" = {}
# 每层 GPU 查找表(全局专家 id → slot,-1=不在池),供 acquire_gpu 在命中层
# 做纯 GPU 重映射(消除每层 .tolist 同步)。与 _slot_of 同步维护(仅在表已建时)。
self._slot_table: "Dict[int, mx.array]" = {}
# 每层 pinned 专家集合:预取进 resident pool 后不参与 LRU 驱逐。
# 用于 K>2 / 小槽位时保住热专家工作集,降低真实 miss。
self._pinned: "Dict[int, set[int]]" = {}
# 可选替换策略:默认 lfu短窗口频率 + LRU tie-break实测比 lru 命中更高EVICT_POLICY=lru 回退纯 LRU。
self.eviction_policy = config.evict_policy().lower()
self.lfu_decay_interval = config.lfu_decay_interval()
self._freq: "Dict[int, Counter[int]]" = {}
self._access_count: "Dict[int, int]" = {}
self.hits = 0
self.misses = 0
self.prefetch_hits = 0
self.prefetch_loads = 0
# GPU remap 路径取证:命中层走纯 GPU 快路径次数 vs 有 miss 回退 host 次数
self.gpu_fastpath = 0
self.gpu_fallback = 0
# 累计「分不到槽被迫落 0 号槽」的专家数(来自 demand_dual)。任何 >0 都意味着有前向
# 拿错专家权重算过 → 输出不可信,必须靠调小 PREFILL_CHUNK / 调大槽位消除。
self.unplaced = 0
# dual 路径真实区槽状态由 C++ demand_dual 唯一权威(无 opt-out)。spec 模式 + native 已编译即启用;
# 非 spec(spec_slots==0)或 native 缺失时保持 Python 权威路径(仅 prefill/host/非双源用)。
# 启用后 _slot_of/_free/_freq 在 dual 路径不再维护,resident_experts/_count 改查 C++ g_real(预取过滤要用)。
self._native_demand = False
if int(spec_slots) > 0:
try:
import mlx_streaming.native_moe_ext as _N
self._native_demand = hasattr(_N, "demand_dual")
except Exception:
self._native_demand = False
if self._native_demand:
# 复刻基线 8-worker 并行读demand miss 的 pread 派给 BgReader 并行执行(高优队列)。
import mlx_streaming.native_moe_ext as _N
try:
_N.bg_reader_start(int(os.environ.get("DEMAND_WORKERS", "8")), 0)
except Exception:
pass
if self._native_demand and os.environ.get("DEMAND_TIMING") == "1":
import atexit
import mlx_streaming.native_moe_ext as _N
_N.demand_timing_enable(True)
atexit.register(lambda: print(
"[DEMAND_TIMING ms] inds/pool/side_snap/real_lock/core/build =",
[round(x / 1e3, 1) for x in _N.demand_timings()], flush=True))
def cap_for(self, layer: int) -> int:
"""该层池容量profile 指定则用之(上限 capacity),否则用全局 capacity。"""
return min(self.layer_caps.get(layer, self.capacity), self.capacity)
def _ensure_layer(self, layer: int):
if layer not in self._slot_of:
self._slot_of[layer] = OrderedDict()
# 空闲槽列表懒填:首次 miss 分配池时再灌入,之后按需增长时追加。
# 注意:始终原地 mutate 这个 list,绝不重新绑定,确保 acquire 里持有的本地引用同步可见。
self._free[layer] = []
self._pinned[layer] = set()
self._freq[layer] = Counter()
self._access_count[layer] = 0
def _alloc_pool(self, layer: int, sample: Dict[str, mx.array]):
"""首次为某层分配池张量,起步 _POOL_INIT_SLOTS 行(不超过该层天花板)。"""
if self.spec_slots > 0:
n = self.cap_for(layer) + self.spec_gens * self.spec_slots # 预分配满(含 spec_gens 代侧区)
# spec/dual 模式:池由 C++ 拥有,侧区异步直写落此 bufferRoute 3 底座)。
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, n)
else:
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype) for k, v in sample.items()}
else:
n = min(_POOL_INIT_SLOTS, self.cap_for(layer))
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype)
for k, v in sample.items()
}
self._alloc[layer] = n
real = self.cap_for(layer) if self.spec_slots > 0 else n
self._free[layer].extend(range(real)) # 原地 mutate,不重新绑定;侧区行不进 free永不被 LRU 分配/驱逐
def _grow_pool(self, layer: int, new_n: int):
"""把某层池物理行数扩到 new_n(封顶 cap_for),拼接保留已驻行的数据与 slot 索引。"""
if self.spec_slots > 0:
return # 侧区模式:预分配满,永不 grow保 C++ 写指针稳定)
old_n = self._alloc[layer]
new_n = min(new_n, self.cap_for(layer))
if new_n <= old_n:
return
old = self._pools[layer]
self._pools[layer] = {
k: mx.concatenate(
[v, mx.zeros((new_n - old_n,) + v.shape[1:], dtype=v.dtype)],
axis=0)
for k, v in old.items()
}
self._free[layer].extend(range(old_n, new_n)) # 新行追加为空闲,原地 mutate
self._alloc[layer] = new_n
def preallocate(self, layer: int, sample: "Dict[str, mx.array]", cap: int):
"""满 cap 预分配 typed 池并 eval 固定 data 指针;幂等。
"C++ 直写槽"机制使用调用后该层不再 grow/不再走 _place_expert(MLX scatter)
槽位由调用方直接 mutate _slot_of + _set_table 管理字节由 C++ pread 写入
"""
self._ensure_layer(layer)
if layer in self._pools:
return
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, cap) # C++ 拥有,地址恒定供直写
else:
self._pools[layer] = {
k: mx.zeros((cap,) + v.shape, dtype=v.dtype) for k, v in sample.items()
}
mx.eval(list(self._pools[layer].values())) # 固定 data 指针
self._alloc[layer] = cap
# 登记所有物理行为空闲,供共用分配器(_alloc_slot)按 free 优先分配;
# 已被 slot_of 占用的行不在此列(暖池复用时 slot_of 已满 → free 为空 → 走驱逐)。
occupied = set(self._slot_of[layer].values())
self._free[layer].extend(r for r in range(cap) if r not in occupied)
def allocated_slots(self, layer: int) -> int:
"""该层池张量当前物理行数(按需增长,≤ cap_for),用于内存核算/测试。"""
return self._alloc.get(layer, 0)
def _pool_owned(self, layer: int) -> bool:
"""该层池是否为 C++-owned bufferspec/dual 模式 + POOL_OWNED
owned 池禁止任何 MLX scatter会重分配 buffer孤立侧区直写落池一律走 C++ 直写"""
return _POOL_OWNED and self.spec_slots > 0
def _write_slot(self, layer: int, slot: int, expert: Dict[str, mx.array]):
if self._pool_owned(layer):
self._write_slots_batch(layer, [slot], [expert]) # 单槽也走 C++ 直写,避免 scatter
return
pool = self._pools[layer]
for k, v in expert.items():
pool[k][slot] = v # de-risk 选定:原地写,单槽、不拷贝整池
def _write_slots_batch(self, layer: int, slots: List[int],
experts: "List[Dict[str, mx.array]]") -> None:
"""把多个专家一次性写进各自槽位。
owned C++ memcpy 直写各行 MLX scatter buffer 地址恒定 侧区直写不被孤立
owned每个 key 一次 stack + fancy-index scatter原碎 kernel 最少化路径
"""
if not slots:
return
pool = self._pools[layer]
if self._pool_owned(layer):
import mlx_streaming.native_moe_ext as _N
keys = list(pool.keys())
pool_list = [pool[k] for k in keys]
srcs_flat = [e[k] for k in keys for e in experts] # key-major与 C++ pool_write_rows 约定一致
mx.eval(srcs_flat) # 物化源,固定 data 指针
_N.pool_write_rows(pool_list, srcs_flat, [int(s) for s in slots])
return
idx = mx.array(slots, dtype=mx.int32)
for k in pool:
pool[k][idx] = mx.stack([e[k] for e in experts], axis=0)
def resident_count(self, layer: int) -> int:
if self._native_demand: # 方案BC++ g_real 为真实区权威
import mlx_streaming.native_moe_ext as N
return int(N.real_region_count(int(layer)))
return len(self._slot_of.get(layer, ()))
def resident_experts(self, layer: int) -> set[int]:
"""返回该层 resident pool 里当前已有的专家 id 集合trace/probe + 侧区预取过滤用)。"""
if self._native_demand: # 方案B查 C++ g_real保预取过滤一致
import mlx_streaming.native_moe_ext as N
flat = N.real_region_contents(int(layer))
return {int(flat[i]) for i in range(0, len(flat), 2)}
return set(self._slot_of.get(layer, {}).keys())
def _bootstrap_dual_pool(self, layer: int) -> None:
"""方案B首次为该层预分配满(cap+spec_gens*spec)池、eval 固定 data 指针 + C++ real_init。幂等。
池结构(per-key typed 数组) Python 从一个样本专家的 shape 字节后续由 C++ demand pread 写入
Python _slot_of/_free 在方案B dual 路径不再使用真实区状态归 C++ g_real
"""
self._ensure_layer(layer)
if layer not in self._pools:
sample = self.loader(layer, 0) # 仅取 shape/dtype不写入池
cap = self.cap_for(layer)
n = cap + self.spec_gens * self.spec_slots
if _POOL_OWNED:
self._pools[layer] = _owned_pool(sample, n) # C++ 拥有,地址恒定供直写
else:
self._pools[layer] = {
k: mx.zeros((n,) + v.shape, dtype=v.dtype) for k, v in sample.items()
}
mx.eval(list(self._pools[layer].values())) # 固定 data 指针,供 C++ memcpy
self._alloc[layer] = n
import mlx_streaming.native_moe_ext as N
N.real_init(int(layer), int(cap))
def resident_lru_scores(self, layer: int) -> dict[int, float]:
"""返回 resident 专家的 LRU 新近度分数0=最久未用1=最近使用。"""
keys = list(self._slot_of.get(layer, {}).keys())
if not keys:
return {}
denom = max(1, len(keys) - 1)
return {int(e): i / denom for i, e in enumerate(keys)}
def _note_access(self, layer: int, expert_ids: List[int]) -> None:
if self.eviction_policy != "lfu":
return
self._ensure_layer(layer)
freq = self._freq[layer]
for e in expert_ids:
freq[int(e)] += 1
self._access_count[layer] += len(expert_ids)
if self.lfu_decay_interval > 0 and self._access_count[layer] >= self.lfu_decay_interval:
for e in list(freq):
freq[e] //= 2
if freq[e] <= 0:
del freq[e]
self._access_count[layer] = 0
def _choose_victim(self, layer: int, current: set[int]) -> int:
slot_of = self._slot_of[layer]
pinned = self._pinned.get(layer, set())
# 绝不驱逐当前请求(current)专家:否则 acquire 末尾按 slot_of 取槽会 KeyError。
# 只在「非 current 且非 pinned」中选受害者选不出来说明本次唯一专家+不可驱逐者已超容量,
# 由调用方(acquire 的 len(uniq)>cap 守卫)负责拦截,这里直接报清晰错误。
candidates = [e for e in slot_of if e not in pinned and e not in current]
if not candidates:
raise ValueError(
f"layer {layer} resident pool has no evictable non-current slot "
f"(capacity={self.cap_for(layer)}, pinned={len(pinned)}, current={len(current)})")
if self.eviction_policy != "lfu":
return candidates[0]
freq = self._freq.get(layer, Counter())
return min(enumerate(candidates), key=lambda x: (freq.get(x[1], 0), x[0]))[1]
def _alloc_slot(self, layer: int, e: int, current: "set[int]") -> "tuple[int, bool]":
"""为 e 分配/复用槽free 优先,否则驱逐非 pinned/非 current 的 LRU 最旧),
更新 slot_of + table但不写数据返回 (slot, is_new)
C++ 直写预取(prefetch_cpp) demand 写入(_place_expert)**共用**保证两条路径
从同一套 _free/_slot_of/_slot_table 分配绝不抢占同一物理槽
"""
slot_of, free = self._slot_of[layer], self._free[layer]
if e in slot_of:
slot_of.move_to_end(e)
return slot_of[e], False
cap = self.cap_for(layer)
if not free and self._alloc.get(layer, 0) < cap:
cur = self._alloc[layer]
self._grow_pool(layer, cur + max(1, cur // 2))
if free:
slot = free.pop(0)
else:
evicted_e = self._choose_victim(layer, current)
slot = slot_of.pop(evicted_e)
self._clear_table(layer, evicted_e)
slot_of[e] = slot
slot_of.move_to_end(e)
self._set_table(layer, e, slot)
return slot, True
def _place_expert(self, layer: int, e: int, expert: Dict[str, mx.array],
current: "set[int] | None" = None) -> int:
"""把专家写进 resident pool必要时增长或驱逐非 pinned LRU返回 slot。"""
current = current or {e}
if layer not in self._pools:
self._alloc_pool(layer, expert)
slot, _ = self._alloc_slot(layer, e, current)
self._write_slot(layer, slot, expert)
return slot
def _place_experts(self, layer: int, ids: List[int],
experts: "Dict[int, Dict[str, mx.array]]",
current: "set[int]") -> List[int]:
"""批量把一组 miss 专家写进 resident pool共用槽分配器 + 单次批量 scatter。
与逐个 _place_expert 语义等价同样的 grow/驱逐/槽分配顺序
只把数据写入合并成 _write_slots_batch去掉每 token 上百次碎 scatter
"""
if not ids:
return []
if layer not in self._pools:
self._alloc_pool(layer, experts[ids[0]])
slots = [self._alloc_slot(layer, e, current)[0] for e in ids]
self._write_slots_batch(layer, slots, [experts[e] for e in ids])
return slots
def _place_experts_stacked(self, layer: int, ids: List[int],
stacked: "Dict[str, mx.array]",
current: "set[int]") -> List[int]:
"""批量写一组 missstacked 为按 ids 顺序预堆叠的 {k:(N,*shape)}
key 直接一次 fancy-index scatter _write_slots_batch mx.stack 都省了
"""
if not ids:
return []
if layer not in self._pools:
self._alloc_pool(layer, {k: v[0] for k, v in stacked.items()})
slots = [self._alloc_slot(layer, e, current)[0] for e in ids]
pool = self._pools[layer]
if self._pool_owned(layer):
import mlx_streaming.native_moe_ext as _N
keys = list(pool.keys())
pool_list = [pool[k] for k in keys]
stacked_list = [stacked[k] for k in keys]
mx.eval(stacked_list) # 物化预堆叠源
_N.pool_write_stacked(pool_list, stacked_list, [int(s) for s in slots])
return slots
idx = mx.array(slots, dtype=mx.int32)
for k in pool:
pool[k][idx] = stacked[k]
return slots
def prefetch_cpp(self, layer: int, expert_ids, submit_fn) -> list:
"""用 C++ 直写把预测专家投机预取进池(与 demand fallback 共用槽分配器)。
submit_fn(slot, expert)提交一次异步 C++ 后台读到该物理槽调用方负责后续 wait
已驻 触摸 LRU 跳过未驻 分配槽并 submit返回新提交的 expert 列表
投机性质不计 hit/miss分得的槽可被后续 LRU 驱逐无可驱逐槽时跳过该专家
交给 demand fallback绝不抛错打断热路径
"""
self._ensure_layer(layer)
if layer not in self._pools:
return [] # 池未建(首 token 预热)→ 全交 fallback
current = {int(e) for e in expert_ids}
slot_of = self._slot_of[layer]
submitted = []
for e in expert_ids:
e = int(e)
if e in slot_of:
slot_of.move_to_end(e)
continue
try:
slot, _ = self._alloc_slot(layer, e, current)
except ValueError:
continue # 无可驱逐槽 → 跳过,留给 demand fallback
submit_fn(slot, e)
submitted.append(e)
return submitted
def pin(self, layer: int, expert_ids: List[int]) -> None:
"""预取并钉住专家到 resident poolpinned 专家不被 LRU 驱逐。
pin 是显式预取不计入 hit/miss 统计后续 acquire/acquire_gpu 命中才计 hit
"""
# dual 模式(demand_dual 唯一权威)已弃用 PIN_HOT:真实区归 C++ g_real,Python pinned 不再生效。
# 明确报错而非静默无效——请设 PIN_HOT=0。
if self._native_demand and list(expert_ids):
raise RuntimeError(
"PIN_HOT 在 demand_dual(dual 模式)下已弃用:真实区由 C++ g_real 权威、不支持 pin。"
"请设 PIN_HOT=0。")
loaded = {int(e): self.loader(layer, int(e)) for e in expert_ids}
self.pin_loaded(layer, loaded)
def pin_loaded(self, layer: int, experts: "Dict[int, Dict[str, mx.array]]") -> None:
"""把已加载专家钉进 resident pool避免 FileExpertStore.pin 重复读/持有两份。"""
uniq = list(dict.fromkeys(int(e) for e in experts))
cap = self.cap_for(layer)
self._ensure_layer(layer)
pinned = self._pinned[layer]
if len(pinned | set(uniq)) > cap:
raise ValueError(
f"pin 后 pinned 数量 {len(pinned | set(uniq))} > 该层池容量 {cap}")
slot_of = self._slot_of[layer]
for e in uniq:
if e in slot_of:
pinned.add(e)
slot_of.move_to_end(e)
continue
self._place_expert(layer, e, experts[e], current={e})
pinned.add(e)
def pin_loaded_dual(self, layer: int, experts: "Dict[int, Dict[str, mx.array]]") -> int:
"""dual(native demand)模式的 pinC++ g_real 是真实区唯一权威Python _slot_of 不生效),
必须 real_pin 注册pinned 免驱逐 + free 头取真实区槽再用 pool_write_rows 把字节
C++ 直写池行owned 池禁止 MLX scatter返回成功 pin 的专家数
仅在 _native_demand 下由 FileExpertStore.pin_batch 调用experts 为已加载的
{expert: {key: arr}}key /顺序与池一致同源自 blob 段表
"""
import mlx_streaming.native_moe_ext as _N
ids = list(dict.fromkeys(int(e) for e in experts))
if not ids:
return 0
self._bootstrap_dual_pool(layer) # 建池 + real_init幂等
slots = _N.real_pin(int(layer), ids) # 登记 pinned + 分配真实区槽(-1=分不到)
ok = [(e, s) for e, s in zip(ids, slots) if s >= 0]
if not ok:
return 0
keys = list(self._pools[layer].keys())
pool_list = [self._pools[layer][k] for k in keys]
srcs_flat = [experts[e][k] for k in keys for e, _ in ok] # key-major同 _write_slots_batch
mx.eval(srcs_flat)
_N.pool_write_rows(pool_list, srcs_flat, [s for _, s in ok])
return len(ok)
def prefetch(self, layer: int, expert_ids: List[int]) -> None:
"""把专家预取进 resident pool但不计入 hit/miss且可被后续 LRU 驱逐。"""
uniq = list(dict.fromkeys(int(e) for e in expert_ids))
cap = self.cap_for(layer)
if len(uniq) > cap:
uniq = uniq[:cap]
self._ensure_layer(layer)
slot_of = self._slot_of[layer]
current = set(uniq)
for e in uniq:
if e in slot_of:
slot_of.move_to_end(e)
self.prefetch_hits += 1
continue
expert = self.loader(layer, e)
self._place_expert(layer, e, expert, current=current)
self.prefetch_loads += 1
def acquire(self, layer: int, expert_ids: List[int], protect: "set[int] | None" = None):
"""protect额外的「不可驱逐」专家集除本批 miss 外。dual 路径传入本前向全部 inds
确保落 miss 腾槽时不驱逐本前向仍需的命中专家否则其 slot 被复用 消费侧 gather 错字节
仅影响驱逐候选不改 hit/miss 记账口径"""
# 唯一专家集合(保序去重):池只需同时容纳本次请求的唯一专家
uniq = list(dict.fromkeys(int(e) for e in expert_ids))
cap = self.cap_for(layer)
if len(uniq) > cap:
raise ValueError(
f"本次请求 {len(uniq)} 个唯一专家 > 该层池容量 {cap}")
self._ensure_layer(layer)
self._note_access(layer, uniq)
uniq_set = set(uniq)
# 驱逐保护集 = 本批 miss 调用方额外指定(如本前向全部 inds 命中专家)。
current = uniq_set if protect is None else (uniq_set | protect)
slot_of, free = self._slot_of[layer], self._free[layer]
# 先处理命中(触摸 LRU再把所有 miss 收成一批
misses = []
for e in uniq:
if e in slot_of:
self.hits += 1
slot_of.move_to_end(e) # 触摸为最近使用,避免本次内被驱逐
else:
misses.append(e)
if misses:
self.misses += len(misses)
if self.stacked_batch_loader is not None:
# 批量读 + 批量物化:每段一次 mx.array写池每 key 一次 scatter碎 kernel 最少)
stacked = self.stacked_batch_loader(layer, misses)
self._place_experts_stacked(layer, misses, stacked, current=current)
else:
if self.batch_loader is not None:
# 一次并行读本层所有 miss8-worker pread保序写池LRU/驱逐语义不变
loaded = self.batch_loader(layer, misses)
else:
loaded = {e: self.loader(layer, e) for e in misses}
# 批量写:每 key 仅一次 stacked scatter取代逐专家逐段碎 scatter
self._place_experts(layer, misses, loaded, current=current)
# slots 与原始 expert_ids 一一对应(含重复),便于直接 reshape 成 routing 索引
slots = [slot_of[int(e)] for e in expert_ids]
return self._pools[layer], slots
def _set_table(self, layer: int, e: int, slot: int):
t = self._slot_table.get(layer)
if t is not None:
t[int(e)] = slot # (num_experts,) 小张量原地改,代价可忽略
def _clear_table(self, layer: int, e: int):
t = self._slot_table.get(layer)
if t is not None:
t[int(e)] = -1
def _ensure_table(self, layer: int, num_experts: int) -> mx.array:
t = self._slot_table.get(layer)
if t is None or int(t.shape[0]) != num_experts:
self._ensure_layer(layer)
tab = [-1] * num_experts
for e, slot in self._slot_of[layer].items():
tab[int(e)] = int(slot)
t = mx.array(tab, dtype=mx.int32)
self._slot_table[layer] = t
return t
def acquire_gpu(self, layer: int, inds: mx.array, num_experts: int):
"""GPU 侧 slot 重映射:命中层零 host 往返(仅一次 miss 标志同步)。
返回 (pool_arrays, local)inds 为路由结果(decode 时形如 (1,1,k))
全命中 local = [inds] GPU; miss 回退既有 host acquire 读盘并维护表后重算 local
"""
table = self._ensure_table(layer, num_experts)
local = mx.take(table, inds)
n_miss = int(mx.sum((local < 0).astype(mx.int32))) # 唯一一次 GPU→CPU 同步
if n_miss == 0:
self.gpu_fastpath += 1
self.hits += int(inds.size) # decode top-k 各专家互异 → size 即唯一命中数
if self.eviction_policy == "lfu": # 全命中层也计频:piggyback 上面 n_miss 的 drain,inds 已物化
flat = [int(i) for i in inds.reshape(-1).tolist()]
self._note_access(layer, flat)
autopin.note(layer, flat) # AUTOPIN 热度:搭同一物化点,零额外同步(AUTOPIN=0 立即返回)
return self._pools[layer], local
# 有 miss:回退 host 路径(读盘、写槽、维护表),再用更新后的表重算 local
self.gpu_fallback += 1
flat = [int(i) for i in inds.reshape(-1).tolist()]
autopin.note(layer, flat) # AUTOPIN 热度:miss 回退 flat 已物化,顺带计频
pool_arrays, _ = self.acquire(layer, flat)
local = mx.take(self._slot_table[layer], inds)
return pool_arrays, local
def verify_acquire_bytes(self, layer, inds, stg=None):
"""诊断(STG_VERIFY)acquire 后把本层真实路由命中专家的池槽字节与磁盘真值逐 key 比对。
发现不一致即打印 (call, layer, expert, slot, gen, key)池槽装错字节的铁证
timing/gen 竞态对应call 是全局校验调用序号(token 前向序便于和首分歧 token 关联)
stg可选 NativeStagingManager用于查该专家落池所用 gen
"""
st = _stg_verify_state
st["calls"] += 1
call = st["calls"]
pool = self._pools.get(layer)
slot_of = self._slot_of.get(layer)
if pool is None or slot_of is None:
return
flat = {int(i) for i in inds.reshape(-1).tolist()}
for e in flat:
slot = slot_of.get(e)
if slot is None:
continue
try:
truth = self.loader(layer, e)
except Exception:
continue
bad_key = None
for k in pool:
if k not in truth:
continue
a = pool[k][slot]
b = truth[k]
if a.shape != b.shape or not bool(mx.all(a == b)):
bad_key = k
break
if bad_key is None:
st["ok"] += 1
else:
st["bad"] += 1
if st["first_bad_call"] is None:
st["first_bad_call"] = call
gen = None
if stg is not None:
gen = stg.placed_gen.get((layer, e))
if st["printed"] < 24:
st["printed"] += 1
print(f"[STG_VERIFY] BAD call={call} layer={layer} expert={e} "
f"slot={slot} gen={gen} key={bad_key} "
f"(ok={st['ok']} bad={st['bad']})", flush=True)
def hit_rate(self) -> float:
tot = self.hits + self.misses
return self.hits / tot if tot else 0.0

237
mlx_streaming/core/cache/virtual_pool.py vendored Normal file
View File

@ -0,0 +1,237 @@
"""VirtualPool预取调度器 + 双源双缓冲协调器(统一收口)。
一个对象承担两职因为 block.py 用同一个 `self._vpool` 属性
1. ahead 调度两种模式都用`ahead_for` / `target_for` 决定 L 层应预读哪层
早层用小 ahead 保召回cutoff 起用大 ahead 抢时序`_native_fused_prefetch` 靠它选目标层
2. 双源双缓冲协调 ZEROCOPY_DUAL_SOURCE 模式`begin_forward` / `read_gen` /
`fill_gen` / `acquire` / `prefetch`侧区分两(gen)物理行每前向读上一代fill
另一代根除本前向消费 gather 读的物理行在本前向 eval 期被 fill 覆盖的竞态
spec 2026-06-25-qwen-virtual-pool-double-buffer对外只暴露专家物理行单次
gather 接口消费者每层零 host 同步
构造两种签名并存互斥使用
- 调度器VirtualPool(num_layers=.., cutoff=.., ahead_lo=.., ahead_hi=..)
- 协调器VirtualPool(resident, staging, spec_slots)dual-source 下再补调度参数即可两职合一
"""
import mlx.core as mx
from mlx_streaming.core.cache import autopin
# 方案B STG_VERIFY 校验累计态(诊断用,默认路径不触及)。
_stg_verify_state = {"ok": 0, "bad": 0, "printed": 0, "calls": 0}
class VirtualPool:
def __init__(self, resident=None, staging=None, spec_slots=None, *,
num_layers=None, cutoff=None, ahead_lo=None, ahead_hi=None, store=None):
# --- 双源协调resident/staging 存在时启用)---
self._rp = resident
self._stg = staging
self._store = store # acquire_host 超容量 fetch 回退用(blob 感知);调度器形态可为 None
self._spec = int(spec_slots) if spec_slots is not None else 0
self._gen = 0
self._last_layer = -1 # 前向边界检测:层号回绕(<= 上次) 即新前向
# --- ahead 调度 ---
self._num_layers = int(num_layers) if num_layers is not None else 0
self._cutoff = int(cutoff) if cutoff is not None else 0
self._a_lo = max(1, int(ahead_lo)) if ahead_lo is not None else 1
self._a_hi = max(1, int(ahead_hi)) if ahead_hi is not None else 1
# ---- ahead 调度 ----
def ahead_for(self, src_layer: int) -> int:
# cutoff早层用 lo保召回cutoff 起用 hi保时序
return self._a_lo if int(src_layer) < self._cutoff else self._a_hi
def target_for(self, src_layer: int) -> int:
# 目标层 = src+ahead。末层无可预读 → 返回 0跳过
# 不再 clamp 到末层clamp 会让多个源层同时预读同一末层,对该层每前向 submit 次数 >
# staging ring其环形 buffer 在惰性切片被本前向 eval 消费前就被后续 submit 完成回调覆盖成
# 别的专家字节 → 池槽装错字节。越界即跳过(交给 demand 读真值),保证每层每前向至多 1 次 submit。
L = int(src_layer)
if L >= self._num_layers - 1:
return 0
tgt = L + self.ahead_for(L)
if tgt > self._num_layers - 1: # 越界(本会触发 clamp 堆叠)→ 跳过
return 0
return tgt
# ---- 双源双缓冲协调 ----
def begin_forward(self, layer_idx: int):
"""每个 MoE 块 __call__ 开头调;层号回绕(<= 上次) 判为新前向 → 代 +1。
稳健不依赖首个 MoE 层是 layer 0也不要求 MoE 层连续"""
if layer_idx <= self._last_layer:
self._gen += 1
# AUTOPIN:dual decode 无 Python 计频点,用前向边界驱动周期落盘(AUTOPIN=0 立即返回)。
autopin.tick()
# 新前向开头排空上一前向提交的侧区 fillC++-owned 池 buffer 由后台异步直写,
# 若消费前 fill 未写完GPU gather 会读到半写侧区行DUAL_VERIFY BAD
# drain 阻塞到在途 fill 全部写完 → 本前向要消费的侧区行字节必已就绪。
if self._spec > 0:
import mlx_streaming.native_moe_ext as _N
_N.sideregion_drain()
self._last_layer = layer_idx
def _gens(self) -> int:
# 代数取自常驻池;单代(=1)时读=填=0(持久 LFU 单区),双代(=2)时交替(%2==&1)。
g = getattr(self._rp, "spec_gens", 2) if self._rp is not None else 2
return max(1, int(g))
def read_gen(self) -> int:
return (self._gen - 1) % self._gens() # 读上一前向填好的代;单代恒 0
def fill_gen(self) -> int:
return self._gen % self._gens() # fill 写本代;单代恒 0
def acquire(self, layer, inds, num_experts, *, seq_len=None, layer_cap=None):
"""统一取用入口GPU-remap 路径):对外呈现「所有专家都在」的视角。
返回 (pool_arrays, local, n_experts)计算侧零分支
- dual有侧区 staging spec>0C++ demand_dual 唯一权威真实区 侧区读代单次 gather
n_experts = layer_cap + spec_gens*spec_slots
- dual GPU-remapacquire_gpun_experts = layer_cap
host/fetch 路径见 acquire_host其输入是 host 侧已 .tolist flat语义不同故分开
"""
cap = int(layer_cap) if layer_cap is not None else self._rp.cap_for(layer)
if self._stg is not None and self._spec > 0:
side_gen = self.read_gen()
# C++ demand_dual 是双源 decode 真实区的唯一权威(每层 1 次 inds 同步、零主线程落池/记账)。
# 无 Python 退路native 未编译 demand_dual → 明确报错decode 依赖 native 全套能力)。
if not getattr(self._rp, "_native_demand", False):
raise RuntimeError(
"双源 decode 需要 native demand_dual已成唯一权威但未检出。请编译 native_moe_ext。")
pool, local = self._acquire_native(layer, inds, side_gen, cap)
n_exp = cap + self._rp.spec_gens * self._rp.spec_slots
return pool, local, n_exp
pool, local = self._rp.acquire_gpu(layer, inds, num_experts)
return pool, local, cap
def _native_meta(self, layer):
"""缓存每层 demand_dual 的不变入参pool_list/seg_nbytes/path避免逐层重建 host 胶水。"""
cache = getattr(self, "_nd_meta", None)
if cache is None:
cache = self._nd_meta = {}
m = cache.get(layer)
if m is None:
stg = self._stg
segs = stg.src._segs # (proj, tensor, dt, shape, nb),与池 key 同序
pool_list = [self._rp._pools[layer][f"{p}.{t}"] for p, t, *_ in segs]
m = (pool_list, [int(nb) for *_, nb in segs],
f"{stg.src.dir}/layer{int(layer):02d}.blob", int(stg.stride))
cache[layer] = m
return m
def _acquire_native(self, layer, inds, side_gen, cap):
"""方案B 取用:委派 C++ demand_dual每层 1 次 inds 同步 + 并行 worker pread 落池),
更新 rp 统计计数供报告口径一致"""
import mlx_streaming.native_moe_ext as N
rp = self._rp
rp._bootstrap_dual_pool(layer) # 首次建池 + real_init幂等
# 方案B 容量前提miss 只能落真实区的 cap 槽,其中 pinned 永不可驱逐 → 可安放上限 cappinned。
# 超了 C++ 会把专家落 0 号槽(拿别的专家权重算)→ 逐位不正确。调用方block.py 的分流判据)
# 负责在超容量前把该前向送去 host/fetch这里只做一次性告警兜底。
_placeable = int(cap) - int(N.real_pinned_count(int(layer)))
if int(inds.size) > _placeable and not getattr(self, "_overcap_warned", False):
self._overcap_warned = True
import sys
print(f"[DEMAND_DUAL] 警告inds.size={int(inds.size)} > 可安放槽 {_placeable}"
f"(cap={int(cap)} pinned={int(cap) - _placeable}),真实区超容量,逐位将不正确。"
f"请调小 PREFILL_CHUNK或调大 EXPERT_SLOTS / 调小 AUTOPIN_BUDGET_FRAC。",
file=sys.stderr, flush=True)
pool_list, seg_nbytes, path, stride = self._native_meta(layer)
local = N.demand_dual(inds, pool_list, seg_nbytes, int(layer), int(side_gen), path,
stride, int(cap),
rp.eviction_policy == "lfu", int(rp.lfu_decay_interval))
st = N.demand_last_stats() # [hitpos, misspos, loads, fallback01, unplaced]
rp.hits += st[0]
rp.misses += st[2]
if len(st) > 4:
rp.unplaced += st[4] # >0 即有前向落了 0 号槽,输出不可信(/api/stats 暴露)
if st[3] == 0:
rp.gpu_fastpath += 1
else:
rp.gpu_fallback += 1
from mlx_streaming import config
if config.stg_verify(): # 诊断:方案B 池字节逐 key 真值校验(默认关)
self._verify_native_bytes(layer, inds, local)
return rp._pools[layer], local
def _verify_native_bytes(self, layer, inds, local):
"""诊断(STG_VERIFY方案B):校验「字节落池不变量」——真实区每个占用槽的池字节 == 该槽当前
C++ 属主专家(g_real) blob 真值这是 C++ 接管落池的字节等价铁证发现不一致即池装错字节
不以 localexpert 为判据local 可能因跨调用/多模型共享 g_real 而滞后于 g_real属路由级
问题非落池字节问题逐位权威信号是 e2e n_mismatch
"""
st = _stg_verify_state
st["calls"] += 1
pool = self._rp._pools.get(layer)
if pool is None:
return
import mlx_streaming.native_moe_ext as N
stg = self._stg
path = f"{stg.src.dir}/layer{int(layer):02d}.blob"
segs = stg.src._segs
flat = N.real_region_contents(int(layer)) # [expert0,slot0,expert1,slot1,...]
for j in range(0, len(flat), 2):
e, slot = flat[j], flat[j + 1]
raw = N.blob_load(path, mx.array([e], dtype=mx.uint32), int(stg.stride))[0]
bad, off = None, 0
for p, t, dt, shape, nb in segs:
k = f"{p}.{t}"
pv = pool[k][slot].reshape(-1).view(mx.uint8)
if not bool(mx.all(pv == raw[off:off + nb])):
bad = k
break
off += nb
if bad is None:
st["ok"] += 1
else:
st["bad"] += 1
if st["printed"] < 12:
st["printed"] += 1
print(f"[STG_VERIFY-DUAL] BAD 落池字节错 call={st['calls']} layer={layer} "
f"expert={e} slot={slot} key={bad} (ok={st['ok']} bad={st['bad']})", flush=True)
def acquire_host(self, layer, flat, inds_shape, inds_dtype, layer_cap):
"""host/fetch 路径收口prefill/大 seq 或关 GPU-remapflat 为 host 侧路由 id 列表。
返回 (pool_arrays, local, n_experts) block.py host/fetch 分支逐元素等价
- uniq <= capacquire(flat)local 为槽位n_experts = layer_cap
- uniq > capfetch(uniq_sorted)local remap [0,uniq) 连续索引n_experts = uniq
native demand(dual)下真实区槽状态归 C++ g_realPython _slot_of/_free 从不填充
self._rp.acquire 必然在 _choose_victim "no evictable non-current slot"
dual 下一律走 fetch 旁路不进池正确但慢这条路径本就是超容量前向的兜底
"""
import mlx.core as mx
cap = int(layer_cap)
uniq_set = set(flat)
# 入池的充分条件是 uniq ≤ cap pinned:acquire 把本次全部唯一专家列为不可驱逐(current),
# 腾槽时只能挑「非 pinned 且非 current」的受害者,少一个都会在 _choose_victim 抛
# "no evictable non-current slot"。pinned 是那层实际钉死数,不是拍出来的固定余量——
# 写死余量会在 cap 小于余量时把所有请求推进 fetch 慢路径(每次重读盘)。
_pinned_n = len(getattr(self._rp, "_pinned", {}).get(layer, ()) or ())
if (not getattr(self._rp, "_native_demand", False)
and len(uniq_set) <= max(1, cap - _pinned_n)):
pool, slots = self._rp.acquire(layer, flat)
local = mx.array(slots, dtype=inds_dtype).reshape(inds_shape)
return pool, local, cap
uniq_sorted = sorted(uniq_set)
remap = {g: i for i, g in enumerate(uniq_sorted)}
local = mx.array([remap[i] for i in flat], dtype=inds_dtype).reshape(inds_shape)
# 超容量(或 dual:真实区槽状态归 C++ g_real,Python _slot_of 从不填充 → acquire 必抛)
# 走 fetch 旁路:不进池、按 uniq 现堆一份参与计算,正确但慢。这条路径本就是兜底。
# fetch 由 blob 感知的 FileExpertStore 提供;ResidentExpertPool 本身没有,故优先用注入的 store。
src = self._store if self._store is not None else self._rp
if not hasattr(src, "fetch"):
raise RuntimeError("VirtualPool.acquire_host 超容量回退需要带 fetch 的 store(构造时传入)")
fetched = src.fetch(layer, uniq_sorted)
return fetched, local, len(uniq_sorted)
def prefetch(self, layer, pred, resident, pool_list):
"""向 fill 代 submit 预读base_row = cap_for(layer) + fill_gen*spec。"""
g = self.fill_gen()
base = self._rp.cap_for(layer) + g * self._spec
return self._stg.submit_pool_sideregion(layer, pred, resident, pool_list, base, gen=g)

View File

@ -0,0 +1 @@
"""线性注意力Qwen3-Next gated-delta相关的自定义 kernel。"""

View File

@ -0,0 +1,226 @@
"""逐步状态输出版 gated-delta Metal kernelvendored 自 mlx_lm.models.gated_delta
来源mlx-lm 0.x `mlx_lm/models/gated_delta.py` `_make_gated_delta_kernel`
本机 mlx 0.31.2本文件****在mlx kernel 基础上多加一个输出 `states_out`
在时间维 `for (t=0..T-1)` 的串行递归里每算完一步就把当前 state 写回
`states_out[B, T, Hv, Dv, Dk]`
为什么 bit-exactmlx kernel 的时间维本就是线程内串行递归不是分块并行 t
寄存器里的 `state` 就是 token t+1 步后的状态全程 fp32同序同融合因此
一次 seq=T 调用产出的 `states_out[:, t]`与逐 tokenT seq=1调用链得到的 state
bit 相等这正是 MTP 投机验证后能按接受长度直接提交 replay的数值根基
用途MTP 验证前向里捕获每个 token 后的 ssm 递归态作为 checkpoint替换原先走
`_gated_delta_step_ops` baseline kernel bit-exact的捕获路径
非目标 Metal(GPU) 路径CPU/ops 回退不在此实现部署恒为 Metal
"""
from typing import Optional, Tuple
import mlx.core as mx
import mlx.nn as nn
# 与 mlx 一致g = exp(-exp(A_log) * softplus(a + dt_bias))fp32。
from mlx_lm.models.gated_delta import compute_g
def _make_gated_delta_multistate_kernel(has_mask=False, vectorized=False):
"""构造逐步状态输出版 kernel除多写 states_out 外与mlx kernel 完全一致。"""
if not mx.metal.is_available():
return None
mask_source = "mask[b_idx * T + t]" if has_mask else "true"
# g 索引方式:标量门控 [B,T,Hv] vs 向量门控 [B,T,Hv,Dk],与 mlx 对齐。
if vectorized:
g_comment = "// g: [B, T, Hv, Dk]"
g_setup = "auto g_ = g + (b_idx * T * Hv + hv_idx) * Dk;"
g_access = "g_[s_idx]"
g_advance = "g_ += Hv * Dk;"
else:
g_comment = "// g: [B, T, Hv]"
g_setup = "auto g_ = g + b_idx * T * Hv;"
g_access = "g_[hv_idx]"
g_advance = "g_ += Hv;"
# 与 mlx 版唯一的差异:时间循环每步末把 state 写入 states_out见下方标注的新增块
source = f"""
auto n = thread_position_in_grid.z;
auto b_idx = n / Hv;
auto hv_idx = n % Hv;
auto hk_idx = hv_idx / (Hv / Hk);
constexpr int n_per_t = Dk / 32;
// q, k: [B, T, Hk, Dk]
auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk;
auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk;
// v, y: [B, T, Hv, Dv]
auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv;
y += b_idx * T * Hv * Dv + hv_idx * Dv;
auto dk_idx = thread_position_in_threadgroup.x;
auto dv_idx = thread_position_in_grid.y;
// state_in, state_out: [B, Hv, Dv, Dk]
auto i_state = state_in + (n * Dv + dv_idx) * Dk;
auto o_state = state_out + (n * Dv + dv_idx) * Dk;
float state[n_per_t];
for (int i = 0; i < n_per_t; ++i) {{
auto s_idx = n_per_t * dk_idx + i;
state[i] = static_cast<float>(i_state[s_idx]);
}}
{g_comment}
{g_setup}
auto beta_ = beta + b_idx * T * Hv;
for (int t = 0; t < T; ++t) {{
if ({mask_source}) {{
float kv_mem = 0.0f;
for (int i = 0; i < n_per_t; ++i) {{
auto s_idx = n_per_t * dk_idx + i;
state[i] = state[i] * {g_access};
kv_mem += state[i] * k_[s_idx];
}}
kv_mem = simd_sum(kv_mem);
auto delta = (v_[dv_idx] - kv_mem) * beta_[hv_idx];
float out = 0.0f;
for (int i = 0; i < n_per_t; ++i) {{
auto s_idx = n_per_t * dk_idx + i;
state[i] = state[i] + k_[s_idx] * delta;
out += state[i] * q_[s_idx];
}}
out = simd_sum(out);
if (thread_index_in_simdgroup == 0) {{
y[dv_idx] = static_cast<InT>(out);
}}
}} else {{
y[dv_idx] = static_cast<InT>(0);
}}
// ===== 新增每步把当前 state 写回 states_out[B, T, Hv, Dv, Dk] =====
// masked state 不变上面 else 分支不更新 state写回的就是上一步的状态
// 与逐 token 单步链语义一致
auto ms_ = states_out
+ ((((b_idx * T + t) * Hv + hv_idx) * Dv + dv_idx) * Dk);
for (int i = 0; i < n_per_t; ++i) {{
auto s_idx = n_per_t * dk_idx + i;
ms_[s_idx] = static_cast<StT>(state[i]);
}}
// ===================================================================
// Increment data pointers to next time step
q_ += Hk * Dk;
k_ += Hk * Dk;
v_ += Hv * Dv;
y += Hv * Dv;
{g_advance}
beta_ += Hv;
}}
for (int i = 0; i < n_per_t; ++i) {{
auto s_idx = n_per_t * dk_idx + i;
o_state[s_idx] = static_cast<StT>(state[i]);
}}
"""
inputs = ["q", "k", "v", "g", "beta", "state_in", "T"]
if has_mask:
inputs.append("mask")
suffix = ""
if vectorized:
suffix += "_vec"
if has_mask:
suffix += "_mask"
return mx.fast.metal_kernel(
name=f"gated_delta_step_multistate{suffix}",
input_names=inputs,
output_names=["y", "state_out", "states_out"],
source=source,
)
_kernel = _make_gated_delta_multistate_kernel(has_mask=False, vectorized=False)
_kernel_masked = _make_gated_delta_multistate_kernel(has_mask=True, vectorized=False)
_kernel_vec = _make_gated_delta_multistate_kernel(has_mask=False, vectorized=True)
_kernel_vec_masked = _make_gated_delta_multistate_kernel(has_mask=True, vectorized=True)
def gated_delta_multistate_kernel(
q: mx.array,
k: mx.array,
v: mx.array,
g: mx.array,
beta: mx.array,
state: mx.array,
mask: Optional[mx.array] = None,
) -> Tuple[mx.array, mx.array, mx.array]:
"""逐步状态输出版 kernel 调用,返回 (y, final_state, states_out)。
形状 mlx `gated_delta_kernel` 一致不在 Python 侧做 GQA repeat
- q, k: [B, T, Hk, Dk]
- v: [B, T, Hv, Dv]
- g: [B, T, Hv]标量 [B, T, Hv, Dk]向量
- beta: [B, T, Hv]
- state:[B, Hv, Dv, Dk]
返回
- y: [B, T, Hv, Dv]
- final_state:[B, Hv, Dv, Dk]== states_out[:, -1]
- states_out: [B, T, Hv, Dv, Dk]每个 token 处理完后的递归态
"""
B, T, Hk, Dk = k.shape
Hv, Dv = v.shape[2:]
input_type = q.dtype
state_type = state.dtype
if g.ndim == 4:
kernel = _kernel_vec_masked if mask is not None else _kernel_vec
else:
kernel = _kernel_masked if mask is not None else _kernel
inputs = [q, k, v, g, beta, state, T]
if mask is not None:
inputs.append(mask)
return kernel(
inputs=inputs,
template=[
("InT", input_type),
("StT", state_type),
("Dk", Dk),
("Dv", Dv),
("Hk", Hk),
("Hv", Hv),
],
grid=(32, Dv, B * Hv),
threadgroup=(32, 4, 1),
output_shapes=[(B, T, Hv, Dv), state.shape, (B, T, Hv, Dv, Dk)],
output_dtypes=[input_type, state_type, state_type],
)
def gated_delta_update_multistate(
q: mx.array,
k: mx.array,
v: mx.array,
a: mx.array,
b: mx.array,
A_log: mx.array,
dt_bias: mx.array,
state: Optional[mx.array] = None,
mask: Optional[mx.array] = None,
) -> Tuple[mx.array, mx.array, mx.array]:
"""对齐 mlx `gated_delta_update` 的签名,多返回逐步状态 states_out。
内部 beta/g 计算与 mlx 完全一致sigmoid(b) / compute_g保证数值路径相同
仅走 Metal kernel 路径multistate GPU 实现
"""
beta = mx.sigmoid(b)
g = compute_g(A_log, a, dt_bias)
if state is None:
B, _, Hk, Dk = q.shape
Hv, Dv = v.shape[-2:]
state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32)
return gated_delta_multistate_kernel(q, k, v, g, beta, state, mask)

75
mlx_streaming/core/mem.py Normal file
View File

@ -0,0 +1,75 @@
"""内存度量统一口径。macOS 上 ru_maxrss 单位是字节。"""
import resource
from dataclasses import dataclass
import mlx.core as mx
def rss_bytes() -> int:
# macOS: ru_maxrss 已是字节Linux 是 KB本项目目标是 macOS
return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
def _call(name: str) -> int:
fn = getattr(mx, name, None)
if fn is None:
fn = getattr(getattr(mx, "metal", object()), name, None)
try:
return int(fn()) if fn else 0
except Exception:
return 0
@dataclass
class MemSnapshot:
rss_bytes: int
mlx_active_bytes: int
mlx_peak_bytes: int
def snapshot() -> MemSnapshot:
return MemSnapshot(
rss_bytes=rss_bytes(),
mlx_active_bytes=_call("get_active_memory"),
mlx_peak_bytes=_call("get_peak_memory"),
)
def clear_cache() -> None:
fn = getattr(mx, "clear_cache", None) or getattr(getattr(mx, "metal", object()), "clear_cache", None)
if fn:
fn()
def reset_peak() -> None:
fn = getattr(mx, "reset_peak_memory", None) or getattr(getattr(mx, "metal", object()), "reset_peak_memory", None)
if fn:
fn()
def _set_limit(name: str, nbytes: int) -> bool:
"""调用 mx.set_*_limit(字节);存在且成功返回 True。"""
fn = getattr(mx, name, None) or getattr(getattr(mx, "metal", object()), name, None)
if fn is None:
return False
try:
fn(int(nbytes))
return True
except Exception:
return False
def setup_memory_hygiene(cache_gb: float = 2.0, wired_gb: float = 0.0) -> dict:
"""长期运行的内存防御:bound MLX 缓冲缓存 + 可选 wire 工作集防 macOS 压缩器。
- cache_gb>0:set_cache_limit,封顶 MLX 可回收缓冲,防长会话里缓存膨胀把常驻推过墙
- wired_gb>0:set_wired_limit,把这么多 GB GPU 缓冲钉为常驻(wired),macOS 不再
压缩/换出这些页 长跑延迟稳定务必 < 系统建议工作集(本机 26.8GB),否则饿死系统
返回实际生效项,便于启动日志记录
"""
applied = {}
if cache_gb and cache_gb > 0:
applied["cache_limit_gb"] = cache_gb if _set_limit("set_cache_limit", cache_gb * 1e9) else None
if wired_gb and wired_gb > 0:
applied["wired_limit_gb"] = wired_gb if _set_limit("set_wired_limit", wired_gb * 1e9) else None
return applied

View File

@ -0,0 +1 @@
"""MoE 模块:门控选专家(gate) / 专家计算(compute) / 自定义算子(custom_kernel) / 热路径块(block)。"""

View File

@ -0,0 +1,431 @@
"""MoE 热路径块:路由 → 选专家 → 取专家权重 → 计算 → 加权合并。
包含两种块
- `StreamingMoeBlock`包住原生 MoE 专家权重常驻switch_mlp只激活选中专家
- `FileStreamingMoeBlock`专家权重从磁盘按需加载流式是低内存推理的核心热路径
集成常驻池 acquirenative-fused-prefetch 预取共享专家叠加等逻辑
"""
import os
import time
import mlx.core as mx
from mlx_streaming import config
from mlx_streaming.core.moe import native_moe
from mlx_streaming.core import route_trace
from mlx_streaming.core.profiling import (
PROF, WINDOW_PROF, PREDICT_RECALL_PROF, MISS_ATTRIB, note_miss_attrib,
note_tprof, TPROF_ON, _PROF_ON, _tick, note_union, UNION_ON)
# decode/verify 热路径判据:seq 短(单 token decode=1、MTP verify=K≤几),与 prefill 长 seq 区分。
_DECODE_SEQ_MAX = 8
from mlx_streaming.core.moe.gate import _effective_top_k
from mlx_streaming.core.cache import autopin
from mlx_streaming.core.moe.compute import (
streaming_switch_glu_forward, PersistentSubGLU)
class StreamingMoeBlock:
"""包住原 Qwen3MoeSparseMoeBlock路由器常驻专家计算改为只算选中专家。
若提供 storeLruExpertStore可进一步把专家权重从磁盘按需加载否则直接在
常驻的 switch_mlp 上做 uniq 切片计算仍只激活少数专家
"""
def __init__(self, orig_block, layer_idx: int, store=None):
self.gate = orig_block.gate # 路由器常驻(很小)
self.top_k = orig_block.top_k
self.norm_topk_prob = orig_block.norm_topk_prob
self.switch_mlp = orig_block.switch_mlp
self.store = store
self.layer_idx = layer_idx
def __call__(self, x: mx.array) -> mx.array:
gates = mx.softmax(self.gate(x), axis=-1, precise=True)
k = _effective_top_k(self.top_k)
inds = mx.argpartition(gates, kth=-k, axis=-1)[..., -k:]
scores = mx.take_along_axis(gates, inds, axis=-1)
if self.norm_topk_prob:
scores = scores / mx.sum(scores, axis=-1, keepdims=True)
mx.eval(inds) # 先物化路由结果,才能按结果取专家
y = streaming_switch_glu_forward(self.switch_mlp, x, inds)
return (y * scores[..., None]).sum(axis=-2)
class FileStreamingMoeBlock:
"""文件后端流式 MoE 块:路由器常驻,专家权重从磁盘按需加载(不持有堆叠 switch_mlp"""
def __init__(self, gate, top_k, norm_topk_prob, store, layer_idx,
hidden, moe_inter, group_size, bits,
proj_bits: dict | None = None,
shared_expert=None, shared_expert_gate=None):
self.gate = gate
self.top_k = top_k
self.norm_topk_prob = norm_topk_prob
self.store = store
self.layer_idx = layer_idx
self.hidden = hidden
self.moe_inter = moe_inter
self.group_size = group_size
self.bits = bits
# Qwen3-Next 等带共享专家的模型:共享专家恒激活、必须常驻(不流式),
# 输出叠加 sigmoid(shared_expert_gate(x)) * shared_expert(x)。None 则退化为纯路由(如 Qwen3-MoE
self.shared_expert = shared_expert
self.shared_expert_gate = shared_expert_gate
# 可选全流式 blob 源STREAM_BLOB=1 时由 model_builder 注入),见 core/blob_loader.py。
self._blob = None
# 持久化子模块:跨 token 复用,避免每调用重建 QSL。
# proj_bits 非空时走混合精度(逐 proj 不同 bit
self._sub = PersistentSubGLU(hidden, moe_inter, group_size, bits,
proj_bits=proj_bits, layer_idx=layer_idx)
def __call__(self, x: mx.array) -> mx.array:
if _PROF_ON:
return self._call_prof(x)
gates = mx.softmax(self.gate(x), axis=-1, precise=True)
k = _effective_top_k(self.top_k)
inds = mx.argpartition(gates, kth=-k, axis=-1)[..., -k:]
scores = mx.take_along_axis(gates, inds, axis=-1)
if self.norm_topk_prob:
scores = scores / mx.sum(scores, axis=-1, keepdims=True)
if config.predict_recall_prof():
ps = getattr(self, "_predicted_set", None)
if ps is not None:
act = {int(i) for i in inds.reshape(-1).tolist()}
PREDICT_RECALL_PROF["hit"] += len(ps & act)
PREDICT_RECALL_PROF["routed"] += len(act)
PREDICT_RECALL_PROF["n"] += 1
if config.native_moe():
native_y = self._try_native_forward(x, inds, scores, gates.shape[-1])
if native_y is not None:
y = native_y
if self.shared_expert is not None:
y = y + mx.sigmoid(self.shared_expert_gate(x)) * self.shared_expert(x)
return y
if self._blob is not None and config.stream_blob():
# 全流式 blob 路径:每层按需并行读专家 → 复用 _sub.forward(MLX quantized_matmul)。
flat = [int(i) for i in inds.reshape(-1).tolist()]
pool_arrays, slots = self._blob.acquire(self.layer_idx, flat)
local = mx.array(slots, dtype=inds.dtype).reshape(inds.shape)
y = self._sub.forward(pool_arrays, len(set(flat)), x, local)
y = (y * scores[..., None]).sum(axis=-2)
if self.shared_expert is not None:
y = y + mx.sigmoid(self.shared_expert_gate(x)) * self.shared_expert(x)
return y
if config.probe_perlayer_sync():
# 诊断:每层强制一次 host 同步(模拟预测 barrier不做任何预取。
# 若单这个就拖慢 ~30% → 税是 eval barrierC++ 读统一内存也救不了。
_ = int(inds.reshape(-1)[:1].item())
if (config.stream_blob_bg()
and getattr(self.store, "_bg", None) is not None):
# 后台预取已物化的专家在此主线程、acquire 前)写进常驻池槽 → 转 miss 为 hit。
if config.window_prof():
t0 = getattr(self, "_submit_t", None)
if t0 is not None:
WINDOW_PROF["sum_s"] += time.perf_counter() - t0
WINDOW_PROF["n"] += 1
self.store.promote_prefetched(self.layer_idx)
# native-fused-prefetch 的 promote把回调预读好的专家写进池槽下移到各 acquire 分支前:
# host/verify 路径已算出真实路由 uniq_set可"只 promote 命中本层路由的专家",零额外同步地
# 丢弃假阳性 → 省掉无用 scatter + 池污染(这是 promote -3.2% 开销的主因)。
_stg_mgr = getattr(self.store, "_staging", None)
# 双源:投机专家留池侧区由 demand_dual 取,不 promote、不驱逐 → 关 promote。
_do_promote = (_stg_mgr is not None and not config.native_no_promote()
and not config.zerocopy_dual_source())
if config.route_trace_enabled():
flat_trace = [int(i) for i in inds.reshape(-1).tolist()]
resident = (self.store.resident_experts(self.layer_idx)
if hasattr(self.store, "resident_experts") else set())
rank = (self.store.resident_lru_scores(self.layer_idx)
if hasattr(self.store, "resident_lru_scores") else {})
route_trace.record(
self.layer_idx, flat_trace, set(flat_trace) - resident, resident, rank)
# native-fused-prefetch在 seq 分支**之前**提交,使 decode(seq=1) 与 MTP verify(seq=K)
# 两条路径都触发预取。dummy 折进 inds(加 0)GPU 路径靠 acquire_gpu 的 n_miss eval、
# host 路径靠后面的 .tolist() eval都会触发完成回调里的 pread。verify 时 x 为 K 个
# tokenverify_in=[x, d_1..d_{K-1}]),预测的是"下一层这 K 个 token 的专家并集"recall≈0.96 的口径)。
# 双源双缓冲:每块开头先推进/检测前向边界(必须在本块预取提交之前,保证同前向内 fill_gen/read_gen 恒定)。
if config.zerocopy_dual_source() and getattr(self, "_vpool", None) is not None:
self._vpool.begin_forward(self.layer_idx)
if (getattr(self.store, "_staging", None) is not None
and not config.native_no_submit()):
_dummy = self._native_fused_prefetch(x)
if _dummy is not None:
inds = inds + (_dummy.reshape(()).astype(inds.dtype) * 0)
layer_cap = self.store.cap_for(self.layer_idx)
# verify(小 seq)可走 GPU 重映射;prefill(大 seq)唯一专家可能超 cap,须留在 host 路径(有超容量 fetch 回退)。
if config.zerocopy_dual_source() and getattr(self, "_vpool", None) is not None:
# 双源两级缓存本就是为 MTP verify 建的:verify 走 dual 路径才能读到侧区,否则落 host
# 路径会白填侧区(预取填了却不读)→ 命中骤降、读盘翻倍。
#
# 但判据必须按 demand_dual 真实区**可安放容量**算,不能按 cap+侧区行:侧区行是预取只读的,
# miss 只能落真实区的 cap 槽,其中 pinned(AUTOPIN)永不可驱逐。设本次唯一专家数 U、
# 真实区 cap=C、pinned=P,可用槽 = C P (本次命中的非 pinned 常驻数),推得
# 「U ≤ C P」是安放得下的充分条件。超了 C++ 分不到槽会把专家落 0 号槽 → 逐位错算
# (prefill chunk 大时曾整段 system+tools 都算错)。U 无同步上界即 inds.size = seq×top_k。
_verify_gpu = (x.shape[1] * k <= self._dual_placeable(layer_cap))
else:
_verify_gpu = (config.verify_gpu_remap() and x.shape[1] * k <= layer_cap)
if (config.resident_pool_enabled()
and config.gpu_remap_enabled()
and (x.shape[1] == 1 or _verify_gpu)):
# decode 热路径:GPU 侧 slot 重映射,命中层零 host 往返(消除每层 .tolist 同步),
# 仅一次 miss 标志同步;真 miss 才回退读盘。decode top-k(≤cap)恒装得下。
# 专家总数 = gate 输出维(gates 末维),无需依赖 gate.weight。
if UNION_ON:
# GPU 路径本不 .tolist 真实路由;仅 UNION_PROF=1 时付一次同步取并集大小。
note_union(x.shape[1], len({int(i) for i in inds.reshape(-1).tolist()}),
self.layer_idx)
# decode/verify GPU 重映射路径:host 无现成真实路由。开关开时用 GPU membership 仅对
# 预读候选现算 used(drain≤budget)过滤假阳性;关时退回整批写入(回退/对照基线)。
if _do_promote:
if config.gpu_remap_promote_filter():
_stg_mgr.promote(self.layer_idx, self.store,
route_inds=inds, num_experts=gates.shape[-1])
else:
_stg_mgr.promote(self.layer_idx, self.store)
if config.miss_attrib():
# 诊断专用:GPU remap 路径本不 .tolist 真实路由(那是它消栅栏的关键),
# 仅 MISS_ATTRIB=1 时付一次 .tolist 取真实路由,统计 decode 热路径 A/B 构成。
_uniq = {int(i) for i in inds.reshape(-1).tolist()}
_ps = getattr(self, "_predicted_set", None)
_res = (self.store.resident_experts(self.layer_idx)
if hasattr(self.store, "resident_experts") else set())
_rdy = _stg_mgr.last_ready.get(self.layer_idx) if _stg_mgr else None
note_miss_attrib(_uniq, _ps, _res, x.shape[1] <= _DECODE_SEQ_MAX, _rdy)
if config.zerocopy_dual_source():
# 双源双缓冲:统一走 VirtualPool.acquire呈现「所有专家都在」视角
# 内部真实区表 侧区(读代) 单次 gather返回 (pool, local, n_experts))。
pool_arrays, local, n_experts = self._vpool.acquire(
self.layer_idx, inds, gates.shape[-1],
seq_len=x.shape[1], layer_cap=layer_cap)
y = self._sub.forward(pool_arrays, n_experts, x, local)
else:
pool_arrays, local = self.store.acquire_gpu(
self.layer_idx, inds, gates.shape[-1])
if config.stg_verify():
# 诊断:消费侧字节校验(默认 off定位混合 ahead 下池槽损坏。
self.store._resident.verify_acquire_bytes(
self.layer_idx, inds, _stg_mgr)
y = self._sub.forward(pool_arrays, layer_cap, x, local)
else:
# prefill/大批量(seq>1)或显式关闭 GPU_REMAP:走 host 路径。
# 只在此处对 inds 做一次 .tolist() 同步uniq/local 全在 Python 里算。
flat = [int(i) for i in inds.reshape(-1).tolist()]
autopin.note(self.layer_idx, flat) # AUTOPIN 热度:搭已物化的 flat,零额外同步(AUTOPIN=0 立即返回)
uniq_set = set(flat)
if UNION_ON:
note_union(x.shape[1], len(uniq_set), self.layer_idx) # 本层路由专家并集(零额外同步,uniq_set 已算)
# promote 只写"预读好 ∩ 本层真实路由"的专家:复用已算的 uniq_set零额外同步
# 假阳性不进池 → 省 scatter、不污染池acquire 前完成,命中转化生效)。
if _do_promote:
_stg_mgr.promote(self.layer_idx, self.store, used=uniq_set)
if config.miss_attrib():
# 在 promote 之后、acquire(读盘)之前统计:本层真实路由的命中构成。
# miss 拆成 A(预测到却没进池budget丢/时序/驱逐) 与 B(没预测到:召回缺口)。
_ps = getattr(self, "_predicted_set", None)
_res = (self.store.resident_experts(self.layer_idx)
if hasattr(self.store, "resident_experts") else set())
_rdy = _stg_mgr.last_ready.get(self.layer_idx) if _stg_mgr else None
note_miss_attrib(uniq_set, _ps, _res, x.shape[1] <= _DECODE_SEQ_MAX, _rdy)
# 池只能同时容纳 ≤该层容量 个唯一专家prefill 唯一数超容量时回退 stack。
# dual 模式_vpool 存在)统一走 VirtualPool.acquire_host 收口 host+fetch
# 非 dual_vpool 为 None保留内联逻辑非目标路径避免额外依赖
if (config.resident_pool_enabled()
and getattr(self, "_vpool", None) is not None):
pool_arrays, local, n_experts = self._vpool.acquire_host(
self.layer_idx, flat, inds.shape, inds.dtype, layer_cap)
y = self._sub.forward(pool_arrays, n_experts, x, local)
elif (config.resident_pool_enabled()
and len(uniq_set) <= layer_cap):
pool_arrays, slots = self.store.acquire(self.layer_idx, flat)
local = mx.array(slots, dtype=inds.dtype).reshape(inds.shape)
y = self._sub.forward(pool_arrays, layer_cap, x, local)
else:
uniq_sorted = sorted(uniq_set)
remap = {g: i for i, g in enumerate(uniq_sorted)}
local = mx.array([remap[i] for i in flat], dtype=inds.dtype).reshape(inds.shape)
fetched = self.store.fetch(self.layer_idx, uniq_sorted)
y = self._sub.forward(fetched, len(uniq_sorted), x, local)
y = (y * scores[..., None]).sum(axis=-2)
if self.shared_expert is not None: # 共享专家(常驻)叠加
y = y + mx.sigmoid(self.shared_expert_gate(x)) * self.shared_expert(x)
return y
def _dual_placeable(self, layer_cap: int) -> int:
"""demand_dual 真实区本层能安放的唯一专家上限 = cap pinned。
pinned 集合只在启动 AUTOPIN 预热时写入生成期恒定,故首次查到后缓存
(decode 每层每步都要判,不能每次都进 C++ 取锁)
"""
n = getattr(self, "_pinned_n", None)
if n is None:
try:
import mlx_streaming.native_moe_ext as _N
n = int(_N.real_pinned_count(int(self.layer_idx)))
except Exception: # noqa: BLE001 取不到就按最保守的「全钉死」算,逼回 host 路径
n = int(layer_cap)
self._pinned_n = n
return max(1, int(layer_cap) - n)
def _native_fused_prefetch(self, x: mx.array):
"""搭车式预取:用下 AHEAD 层 gate 对 x 算预测 inds(lazy),挂 GPU 完成回调,
C++ (GPU 算完后) id + pread 预热下层专家字节返回 dummy 张量(需被 eval 才触发)
全程不加主线程 host 同步靠折进 inds acquire_gpu n_miss eval
"""
model = getattr(self, "_prefetch_model_ref", None)
bl = getattr(self.store, "_blob_loader", None)
if model is None or bl is None:
return None
layers = getattr(model, "layers", [])
vp = getattr(self, "_vpool", None)
if vp is not None:
tgt = vp.target_for(self.layer_idx) # per-layer cutoff ahead
if not (self.layer_idx < tgt < len(layers)): # 0/越界/无前瞻 → 跳过
return None
else:
ahead = max(1, config.cross_layer_ahead(default=1))
tgt = self.layer_idx + ahead
if not (0 <= tgt < len(layers)):
return None
tmlp = getattr(layers[tgt], "mlp", None)
if not isinstance(tmlp, FileStreamingMoeBlock):
return None
try:
from mlx_streaming import native_moe_ext as _N
except Exception:
return None
# 方案B预测"宽集合"(top-N, N=predict_width)只为高 recall不占内存仅一次 gate argpartition
# 真正占 staging 的是 C++ 回调过滤常驻后、按分截到 staging budget 的"缺口"子集。
_tp0 = time.perf_counter() if TPROF_ON else 0.0
predict_width = config.cross_layer_predict_width()
# 默认(PREDICT_USE_X=1):用本层 MoE 输入 x 喂目标层 gate。x 含了本层 attention离 L+1
# 近半层 → 比未归一化输入更新鲜,实测 recall 0.812→0.847(+3.6pp)、还省一次 norm。
# PREDICT_USE_X=0 回退旧路径:目标层 post_attention_layernorm 作用在本层未归一化输入上。
h = getattr(self, "_unnormed_input", None)
if config.predict_use_x():
g = tmlp.gate(x)
elif h is None:
h = x # 退化:无未归一化输入时用 xnorm 偏差,覆盖会差)
g = tmlp.gate(h)
else:
g = tmlp.gate(layers[tgt].post_attention_layernorm(h))
# 按 seq 维聚合成"该批 K 个 token 的专家并集"近似。聚合方式 PREDICT_AGG
# - max默认任一 token 强烈想要即入选mean各 token 平均偏好;
# - union每 token 各取 top-kk 再并集(与真实路由"并集"结构最一致,候选更多,由 handler 去重+截断)。
_agg = config.predict_agg()
kk = min(g.shape[-1], predict_width)
if g.ndim == 3 and _agg == "union":
# union 候选数 = seq × union_k。每 token 只取 top-union_k独立于 predict_width
# 调小 union_k 缩候选 → 减少缺口、缓解 staging/读盘洪水(每 token 真实只路由 top_k 个)。
uk = min(g.shape[-1], config.predict_union_k())
pred = mx.argpartition(g, kth=-uk, axis=-1)[..., -uk:].reshape(-1).astype(mx.uint32)
else:
if g.ndim == 3:
g = g.mean(axis=1) if _agg == "mean" else g.max(axis=1)
# 取 top-kkargpartitionO(E) 不全排序)。实测:加宽 kk 不提升命中——staging 更多 = 后台
# pread 更多、到位更晚 + 抢带宽 → 反而更低;故 width=16 即最优cap 截断极少触发、无需排序。
pred = mx.argpartition(g, kth=-kk, axis=-1)[..., -kk:].reshape(-1).astype(mx.uint32)
if config.predict_recall_prof() or config.miss_attrib():
try:
tmlp._predicted_set = {int(e) for e in pred.tolist()}
except Exception:
pass
if TPROF_ON:
note_tprof("predict_s", time.perf_counter() - _tp0, count_key="predict_n")
stg = getattr(self.store, "_staging", None)
if stg is not None:
resident = (self.store.resident_experts(tgt)
if hasattr(self.store, "resident_experts") else None)
if config.zerocopy_dual_source():
# 双源双缓冲:经 VirtualPool 向填代 submit,C++ 回调把预读段散写进该代侧区行。
rp = self.store._resident
if tgt not in rp._pools:
return None # 目标层池未建(首 token 预热) → 跳过本次
segs = stg.src._segs # (proj, tensor, dt, shape, nb),与池 key 顺序一致
pool_list = [rp._pools[tgt][f"{p}.{t}"] for p, t, *_ in segs]
import os as _os
if _os.environ.get("POOL_PTR_TRACE") is not None and tgt == int(_os.environ["POOL_PTR_TRACE"]):
try:
import mlx_streaming.native_moe_ext as _N
_fk = f"{segs[0][0]}.{segs[0][1]}"
print(f"[POOL_PTR] submit layer={tgt} key={_fk} obj_id={id(rp._pools[tgt][_fk])} "
f"ptr={hex(_N.array_data_ptr(rp._pools[tgt][_fk]))}", flush=True)
except Exception as _e:
print(f"[POOL_PTR] err {_e}", flush=True)
return self._vpool.prefetch(tgt, pred, resident, pool_list)
# miss→hit:回调按目标层常驻快照过滤,只把缺口 pread 进 staging≤budget 行promote 时写池。
if TPROF_ON:
_ts0 = time.perf_counter()
_r = stg.submit(tgt, pred, resident)
note_tprof("submit_s", time.perf_counter() - _ts0, count_key="submit_n")
return _r
return stg.submit(tgt, pred, resident)
# 仅预热字节page cache的轻量版。
path = os.path.join(bl.dir, f"layer{tgt:02d}.blob")
return _N.prefetch_on_complete(pred, path, int(bl.stride), True)
def _try_native_forward(self, x: mx.array, inds: mx.array, scores: mx.array, num_experts: int):
# native 第一版只支持三投影同 bit 的专家;其他情况保持现有 MLX 回退。
proj_bits = getattr(self._sub, "proj_bits", {})
bits_set = {int(proj_bits.get(name, self.bits)) for name in ("gate_proj", "up_proj", "down_proj")}
if len(bits_set) != 1:
return None
bits = bits_set.pop()
if not native_moe.can_native_moe(
self.layer_idx, self.hidden, self.moe_inter,
self.group_size, bits, num_experts):
return None
t_sync = time.perf_counter()
mx.eval(inds)
flat = [int(i) for i in inds.reshape(-1).tolist()]
native_moe.note_route_sync(time.perf_counter() - t_sync)
return native_moe.try_native_moe(
self.layer_idx, flat, x, scores,
self.hidden, self.moe_inter, self.group_size, bits, num_experts)
def _call_prof(self, x: mx.array) -> mx.array:
t = time.perf_counter()
gates = mx.softmax(self.gate(x), axis=-1, precise=True)
k = _effective_top_k(self.top_k)
inds = mx.argpartition(gates, kth=-k, axis=-1)[..., -k:]
scores = mx.take_along_axis(gates, inds, axis=-1)
if self.norm_topk_prob:
scores = scores / mx.sum(scores, axis=-1, keepdims=True)
mx.eval(inds, scores)
_tick("route", t); t = time.perf_counter()
flat = [int(i) for i in inds.reshape(-1).tolist()]
uniq_set = set(flat)
layer_cap = self.store.cap_for(self.layer_idx)
use_pool = (config.resident_pool_enabled()
and len(uniq_set) <= layer_cap)
if not use_pool:
uniq_sorted = sorted(uniq_set)
remap = {g: i for i, g in enumerate(uniq_sorted)}
local = mx.array([remap[i] for i in flat], dtype=inds.dtype).reshape(inds.shape)
mx.eval(local)
_tick("pyremap", t); t = time.perf_counter()
if use_pool:
pool_arrays, slots = self.store.acquire(self.layer_idx, flat)
local = mx.array(slots, dtype=inds.dtype).reshape(inds.shape)
mx.eval(local)
fetched, n_experts = pool_arrays, layer_cap
else:
fetched = self.store.fetch(self.layer_idx, uniq_sorted)
n_experts = len(uniq_sorted)
mx.eval(list(fetched.values()))
_tick("fetch", t); t = time.perf_counter()
y = self._sub.forward(fetched, n_experts, x, local)
mx.eval(y)
_tick("matmul", t); t = time.perf_counter()
y = (y * scores[..., None]).sum(axis=-2)
if self.shared_expert is not None:
y = y + mx.sigmoid(self.shared_expert_gate(x)) * self.shared_expert(x)
mx.eval(y)
_tick("combine", t)
PROF["n_calls"] += 1
return y

View File

@ -0,0 +1,208 @@
"""MoE 专家计算:把选中专家切片成小 SwitchGLU 并前向(数值等价全量 SwitchGLU
设计要点
- `_unique_and_local` top-k 路由的全局专家 id 压成唯一集合 + 本地下标
只在实际激活的少数专家上 gather 计算
- `PersistentSubGLU`跨调用复用一套 SwitchGLU + 3×QuantizedSwitchLinear 对象
解码 batch=1 时唯一专家数恒为 top_k只首次构造之后原地 update 权重省重建开销
"""
from typing import Tuple
import mlx.core as mx
import mlx.nn as nn
from mlx_lm.models.switch_layers import SwitchLinear, QuantizedSwitchLinear, SwitchGLU
from mlx_streaming import config
from mlx_streaming.core.moe.custom_kernel import (
_custom_qproj_enabled, _custom_fused_moe_enabled, _custom_qproj_targets,
_custom_qlinear_indexed, _custom_fused_moe_indexed)
def _unique_and_local(inds: mx.array) -> Tuple[mx.array, mx.array]:
"""返回 (uniq, local)uniq 是排序后的唯一全局专家 idlocal 是与 inds 同形状的本地下标。"""
flat = [int(i) for i in inds.reshape(-1).tolist()]
uniq_sorted = sorted(set(flat))
remap = {g: i for i, g in enumerate(uniq_sorted)}
local = mx.array([remap[i] for i in flat], dtype=inds.dtype).reshape(inds.shape)
uniq = mx.array(uniq_sorted, dtype=inds.dtype)
return uniq, local
def _slice_switch_linear(lin, uniq: mx.array):
"""把一个 SwitchLinear/QuantizedSwitchLinear 沿专家维切到 uniq返回新的小 linear。"""
n = int(uniq.shape[0])
has_bias = "bias" in lin
if isinstance(lin, QuantizedSwitchLinear):
new = QuantizedSwitchLinear(
lin.input_dims, lin.output_dims, n, bias=has_bias,
group_size=lin.group_size, bits=lin.bits, mode=lin.mode,
)
else:
new = SwitchLinear(lin.input_dims, lin.output_dims, n, bias=has_bias)
sliced = {}
for name, p in lin.parameters().items():
if isinstance(p, mx.array) and p.ndim >= 1 and p.shape[0] == lin.num_experts:
sliced[name] = p[uniq] # 第 0 维是专家维,按 uniq 取子集
else:
sliced[name] = p
new.update(sliced)
return new
def streaming_switch_glu_forward(glu: SwitchGLU, x: mx.array, inds: mx.array) -> mx.array:
"""只在 inds 涉及到的唯一专家上做 SwitchGLU 等价计算。"""
uniq, local = _unique_and_local(inds)
sub = SwitchGLU(glu.gate_proj.input_dims, glu.gate_proj.output_dims, int(uniq.shape[0]))
sub.gate_proj = _slice_switch_linear(glu.gate_proj, uniq)
sub.up_proj = _slice_switch_linear(glu.up_proj, uniq)
sub.down_proj = _slice_switch_linear(glu.down_proj, uniq)
sub.activation = glu.activation
return sub(x, local)
def _build_qsl(prefix: str, fetched: dict, in_dims: int, out_dims: int, n: int,
group_size: int, bits: int):
"""用 fetched 里以 prefix.* 为键的堆叠参数构造一个 QuantizedSwitchLinear。"""
lin = QuantizedSwitchLinear(in_dims, out_dims, n, bias=False,
group_size=group_size, bits=bits, mode="affine")
sub = {name.split(".", 1)[1]: v for name, v in fetched.items()
if name.startswith(prefix + ".")}
lin.update(sub)
return lin
def streaming_switch_glu_forward_from_store(store, layer, x, inds, hidden, moe_inter,
group_size, bits):
"""文件后端版:从 store 只取选中专家、构造小 SwitchGLU 计算(等价于全量 SwitchGLU
无状态版本每次新建 sub保留给单测使用端到端走 PersistentSubGLU
"""
uniq, local = _unique_and_local(inds)
fetched = store.fetch(layer, [int(i) for i in uniq.tolist()])
n = int(uniq.shape[0])
sub = SwitchGLU(hidden, moe_inter, n)
sub.gate_proj = _build_qsl("gate_proj", fetched, hidden, moe_inter, n, group_size, bits)
sub.up_proj = _build_qsl("up_proj", fetched, hidden, moe_inter, n, group_size, bits)
sub.down_proj = _build_qsl("down_proj", fetched, moe_inter, hidden, n, group_size, bits)
return sub(x, local)
def _update_qsl(lin, prefix: str, fetched: dict):
"""原地更新一个 QuantizedSwitchLinear 的参数(只换数组引用,不重建对象)。"""
sub = {name.split(".", 1)[1]: v for name, v in fetched.items()
if name.startswith(prefix + ".")}
lin.update(sub)
class PersistentSubGLU:
"""按专家数 n 缓存一套 SwitchGLU + 3×QuantizedSwitchLinear跨调用复用。
解码 batch=1 时每层唯一专家数恒为 top_kn 不变 只在首次/n 变化时构造一次
之后每个 token 仅原地 update 三组权重省掉每调用重建对象与初始化随机量化权重的开销
"""
def __init__(self, hidden: int, moe_inter: int, group_size: int, bits: int,
proj_bits: dict | None = None, layer_idx: int | None = None,
quant_mode: str = "affine", swiglu_limit: float = 0.0):
self.hidden = hidden
self.moe_inter = moe_inter
self.group_size = group_size
self.bits = bits
self.layer_idx = layer_idx
# 混合精度:每个 proj 用各自 bitNone 时三 proj 统一用 bits。
self.proj_bits = proj_bits or {
"gate_proj": bits, "up_proj": bits, "down_proj": bits}
# quant_mode量化模式affine / mxfp4swiglu_limit>0 时走 DeepSeek 手动 clip 路径。
self.quant_mode = quant_mode
self.swiglu_limit = swiglu_limit
self._glu = None
self._n = None
def _ensure(self, n: int):
if self._glu is not None and self._n == n:
return
pb = self.proj_bits
glu = SwitchGLU(self.hidden, self.moe_inter, n)
glu.gate_proj = QuantizedSwitchLinear(
self.hidden, self.moe_inter, n, bias=False,
group_size=self.group_size, bits=pb["gate_proj"], mode=self.quant_mode)
glu.up_proj = QuantizedSwitchLinear(
self.hidden, self.moe_inter, n, bias=False,
group_size=self.group_size, bits=pb["up_proj"], mode=self.quant_mode)
glu.down_proj = QuantizedSwitchLinear(
self.moe_inter, self.hidden, n, bias=False,
group_size=self.group_size, bits=pb["down_proj"], mode=self.quant_mode)
self._glu = glu
self._n = n
def forward(self, fetched: dict, n: int, x: mx.array, local: mx.array) -> mx.array:
self._ensure(n)
_update_qsl(self._glu.gate_proj, "gate_proj", fetched)
_update_qsl(self._glu.up_proj, "up_proj", fetched)
_update_qsl(self._glu.down_proj, "down_proj", fetched)
if self.swiglu_limit > 0:
# DeepSeek 路径:手动 gate/up/clip/silu/down不能用 SwiGLU 融合激活
# (镜像 deepseek_v4.Expertsup 双侧 clip、gate 上侧 minimum
x_exp = mx.expand_dims(x, (-2, -3)) # (..., 1, 1, D)
gate = self._glu.gate_proj(x_exp, local)
up = self._glu.up_proj(x_exp, local)
up = mx.clip(up, -self.swiglu_limit, self.swiglu_limit)
gate = mx.minimum(gate, self.swiglu_limit)
h = nn.silu(gate) * up
y = self._glu.down_proj(h, local)
return y.squeeze(-2)
max_fused_seq = config.custom_fused_moe_max_seq()
if (x.shape[1] <= max_fused_seq
and _custom_fused_moe_enabled(self.layer_idx or -1, self.proj_bits)):
return self._custom_fused_forward(x, local)
max_seq = config.custom_qproj_max_seq()
if (x.shape[1] <= max_seq
and _custom_qproj_enabled(self.layer_idx or -1, self.proj_bits["gate_proj"])):
return self._custom_gate_up_forward(x, local)
return self._glu(x, local)
def _custom_fused_forward(self, x: mx.array, local: mx.array) -> mx.array:
"""完整替换 gate/up/SwiGLU/down保持 SwitchGLU 的 [B,S,K,H] 输出契约。"""
shape = local.shape
k = int(shape[-1])
x_flat = mx.broadcast_to(mx.expand_dims(x, -2), x.shape[:-1] + (k, x.shape[-1]))
x_flat = x_flat.reshape(-1, self.hidden).astype(mx.float32)
idx = local.reshape(-1).astype(mx.uint32)
y = _custom_fused_moe_indexed(
x_flat, idx,
self._glu.gate_proj["weight"], self._glu.gate_proj["scales"], self._glu.gate_proj["biases"],
self._glu.up_proj["weight"], self._glu.up_proj["scales"], self._glu.up_proj["biases"],
self._glu.down_proj["weight"], self._glu.down_proj["scales"], self._glu.down_proj["biases"],
self.hidden, self.moe_inter, self.group_size, self.proj_bits["gate_proj"],
)
return y.reshape(shape + (self.hidden,))
def _custom_gate_up_forward(self, x: mx.array, local: mx.array) -> mx.array:
"""只替换 gate/up projectiondown 仍走 MLX QuantizedSwitchLinear。"""
shape = local.shape
k = int(shape[-1])
x_flat = mx.broadcast_to(mx.expand_dims(x, -2), x.shape[:-1] + (k, x.shape[-1]))
x_flat = x_flat.reshape(-1, self.hidden).astype(mx.float32)
idx = local.reshape(-1).astype(mx.uint32)
tile = config.custom_qproj_tile()
up = _custom_qlinear_indexed(
x_flat, idx,
self._glu.up_proj["weight"], self._glu.up_proj["scales"], self._glu.up_proj["biases"],
self.moe_inter, self.hidden, self.group_size, self.proj_bits["up_proj"], tile)
gate = _custom_qlinear_indexed(
x_flat, idx,
self._glu.gate_proj["weight"], self._glu.gate_proj["scales"], self._glu.gate_proj["biases"],
self.moe_inter, self.hidden, self.group_size, self.proj_bits["gate_proj"], tile)
up = up.reshape(shape + (1, self.moe_inter))
gate = gate.reshape(shape + (1, self.moe_inter))
a = self._glu.activation(up, gate)
if "down" in _custom_qproj_targets():
a_flat = a.squeeze(-2).reshape(-1, self.moe_inter).astype(mx.float32)
y = _custom_qlinear_indexed(
a_flat, idx,
self._glu.down_proj["weight"], self._glu.down_proj["scales"], self._glu.down_proj["biases"],
self.hidden, self.moe_inter, self.group_size, self.proj_bits["down_proj"], tile)
return y.reshape(shape + (self.hidden,))
y = self._glu.down_proj(a, local)
return y.squeeze(-2)

View File

@ -0,0 +1,266 @@
"""自定义 Metal 算子indexed 量化线性 + fused MoE expert实验性加速路径
这些 kernel `mx.fast.metal_kernel` `indices[pair]` 选择专家权重
直接在 packed 量化权重上算 SwiGLU/down避免反量化中间张量仅在对应
CUSTOM_* 环境开关命中且 bit 宽匹配时启用默认关闭不改变模型数值
"""
from typing import Tuple # noqa: F401 (保留以兼容历史 import)
import mlx.core as mx
from mlx_streaming import config
from mlx_streaming.config import parse_layers_env as _parse_layers_env
# kernel 编译缓存:按 (tile) / (hidden,moe_inter,...) 复用已编译 kernel。
_CUSTOM_QKERNEL_CACHE = {}
_CUSTOM_FUSED_MOE_CACHE = {}
def _custom_qproj_enabled(layer_idx: int, bits: int) -> bool:
if not config.custom_qproj():
return False
if bits != config.custom_qproj_bits():
return False
layers = _parse_layers_env("CUSTOM_QPROJ_LAYERS")
return layers is None or int(layer_idx) in layers
def _custom_fused_moe_enabled(layer_idx: int, proj_bits: dict) -> bool:
if not config.custom_fused_moe():
return False
bits = config.custom_fused_moe_bits()
if any(int(proj_bits[name]) != bits for name in ("gate_proj", "up_proj", "down_proj")):
return False
layers = _parse_layers_env("CUSTOM_FUSED_MOE_LAYERS")
return layers is None or int(layer_idx) in layers
def _custom_qproj_targets() -> set[str]:
spec = config.custom_qproj_targets()
return {x.strip() for x in spec.split(",") if x.strip()}
def _custom_qlinear_indexed(x: mx.array, indices: mx.array, weight: mx.array,
scales: mx.array, biases: mx.array,
out_dim: int, in_dim: int, group_size: int,
bits: int, tile: int = 4) -> mx.array:
"""indexed custom qlinear: x[p,in] 使用 indices[p] 选择专家权重,输出 [p,out]。"""
key = tile
kernel = _CUSTOM_QKERNEL_CACHE.get(key)
if kernel is None:
source = r"""
constexpr int lanes_per_row = 256 / rows_per_group;
constexpr int block_size = rows_per_group * lanes_per_row;
uint tid = thread_position_in_threadgroup.x;
uint group_id = thread_position_in_grid.x / block_size;
uint local_row = tid / lanes_per_row;
uint lane = tid % lanes_per_row;
uint global_row = group_id * rows_per_group + local_row;
if (global_row >= pairs * out_dim) return;
uint pair = global_row / out_dim;
uint out_row = global_row % out_dim;
uint expert = indices[pair];
constexpr uint mask = (1u << bits) - 1u;
constexpr int words_per_row = (in_dim * bits) / 32;
constexpr int groups_per_row = in_dim / group_size;
threadgroup float partial[block_size];
float acc = 0.0f;
for (int col = int(lane); col < in_dim; col += lanes_per_row) {
int bit_offset = col * bits;
int word_idx = bit_offset / 32;
int shift = bit_offset % 32;
uint base = (expert * out_dim + out_row) * words_per_row;
uint word = weight[base + word_idx];
uint q = (word >> shift);
if (shift + bits > 32) {
uint next_word = weight[base + word_idx + 1];
q |= (next_word << (32 - shift));
}
q = q & mask;
int g = col / group_size;
uint sb = (expert * out_dim + out_row) * groups_per_row + g;
float wv = float(q) * scales[sb] + biases[sb];
acc += wv * x[pair * in_dim + col];
}
partial[tid] = acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (lane < stride) {
partial[tid] += partial[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (lane == 0) {
y[global_row] = partial[0 + (local_row * lanes_per_row)];
}
"""
kernel = mx.fast.metal_kernel(
name=f"custom_qlinear_indexed_tile{tile}",
input_names=["x", "indices", "weight", "scales", "biases"],
output_names=["y"],
source=source,
)
_CUSTOM_QKERNEL_CACHE[key] = kernel
pairs = int(indices.size)
grid_groups = (pairs * out_dim + tile - 1) // tile
(y,) = kernel(
inputs=[x, indices.astype(mx.uint32), weight, scales, biases],
output_shapes=[(pairs, out_dim)],
output_dtypes=[mx.float32],
grid=(grid_groups * 256, 1, 1),
threadgroup=(256, 1, 1),
template=[
("pairs", pairs),
("out_dim", out_dim),
("in_dim", in_dim),
("group_size", group_size),
("bits", bits),
("rows_per_group", tile),
],
)
return y
def _custom_fused_moe_indexed(x: mx.array, indices: mx.array,
gate_w: mx.array, gate_s: mx.array, gate_b: mx.array,
up_w: mx.array, up_s: mx.array, up_b: mx.array,
down_w: mx.array, down_s: mx.array, down_b: mx.array,
hidden: int, moe_inter: int, group_size: int,
bits: int) -> mx.array:
"""fused MoE expert: x[p] 用 indices[p] 选择专家,输出 down 后的 hidden。"""
lanes_per_row = config.custom_fused_moe_lanes()
block_size = config.custom_fused_moe_block()
key = (hidden, moe_inter, group_size, bits, lanes_per_row, block_size)
kernel = _CUSTOM_FUSED_MOE_CACHE.get(key)
if kernel is None:
source = r"""
uint tid = thread_position_in_threadgroup.x;
uint pair = thread_position_in_grid.x / block_size;
if (pair >= pairs) return;
uint expert = indices[pair];
constexpr uint rows_per_step = block_size / lanes_per_row;
uint local_row = tid / lanes_per_row;
uint row_lane = tid % lanes_per_row;
constexpr uint mask = (1u << bits) - 1u;
constexpr int gu_words_per_row = (hidden * bits) / 32;
constexpr int gu_groups_per_row = hidden / group_size;
constexpr int down_words_per_row = (moe_inter * bits) / 32;
constexpr int down_groups_per_row = moe_inter / group_size;
threadgroup float act[1024];
threadgroup float gate_part[block_size];
threadgroup float up_part[block_size];
// 多个 lane 协作计算同一 row避免每个 dot 完全串行
for (uint row_base = 0; row_base < moe_inter; row_base += rows_per_step) {
uint row = row_base + local_row;
float gate_acc = 0.0f;
float up_acc = 0.0f;
if (row < moe_inter) {
for (uint col = row_lane; col < hidden; col += lanes_per_row) {
int bit_offset = int(col * bits);
int word_idx = bit_offset / 32;
int shift = bit_offset % 32;
uint gu_base = (expert * moe_inter + row) * gu_words_per_row;
uint qg = gate_w[gu_base + word_idx] >> shift;
uint qu = up_w[gu_base + word_idx] >> shift;
if (shift + bits > 32) {
qg |= gate_w[gu_base + word_idx + 1] << (32 - shift);
qu |= up_w[gu_base + word_idx + 1] << (32 - shift);
}
qg &= mask;
qu &= mask;
uint g = col / group_size;
uint sb = (expert * moe_inter + row) * gu_groups_per_row + g;
float xv = x[pair * hidden + col];
gate_acc += (float(qg) * gate_s[sb] + gate_b[sb]) * xv;
up_acc += (float(qu) * up_s[sb] + up_b[sb]) * xv;
}
}
gate_part[tid] = gate_acc;
up_part[tid] = up_acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) {
gate_part[tid] += gate_part[tid + stride];
up_part[tid] += up_part[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < moe_inter) {
float gate_v = gate_part[tid];
float up_v = up_part[tid];
float sig = 1.0f / (1.0f + exp(-gate_v));
act[row] = gate_v * sig * up_v;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
// down projection 输出最终 hidden
for (uint row_base = 0; row_base < hidden; row_base += rows_per_step) {
uint row = row_base + local_row;
float acc = 0.0f;
if (row < hidden) {
for (uint col = row_lane; col < moe_inter; col += lanes_per_row) {
int bit_offset = int(col * bits);
int word_idx = bit_offset / 32;
int shift = bit_offset % 32;
uint base = (expert * hidden + row) * down_words_per_row;
uint q = down_w[base + word_idx] >> shift;
if (shift + bits > 32) {
q |= down_w[base + word_idx + 1] << (32 - shift);
}
q &= mask;
uint g = col / group_size;
uint sb = (expert * hidden + row) * down_groups_per_row + g;
acc += (float(q) * down_s[sb] + down_b[sb]) * act[col];
}
}
gate_part[tid] = acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) {
gate_part[tid] += gate_part[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < hidden) {
y[pair * hidden + row] = gate_part[tid];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
"""
kernel = mx.fast.metal_kernel(
name=f"custom_fused_moe_h{hidden}_i{moe_inter}_b{bits}_l{lanes_per_row}",
input_names=[
"x", "indices",
"gate_w", "gate_s", "gate_b",
"up_w", "up_s", "up_b",
"down_w", "down_s", "down_b",
],
output_names=["y"],
source=source,
)
_CUSTOM_FUSED_MOE_CACHE[key] = kernel
pairs = int(indices.size)
(y,) = kernel(
inputs=[
x, indices.astype(mx.uint32),
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b,
],
output_shapes=[(pairs, hidden)],
output_dtypes=[mx.float32],
grid=(pairs * block_size, 1, 1),
threadgroup=(block_size, 1, 1),
template=[
("pairs", pairs),
("hidden", hidden),
("moe_inter", moe_inter),
("group_size", group_size),
("bits", bits),
("lanes_per_row", lanes_per_row),
("block_size", block_size),
],
)
return y

View File

@ -0,0 +1,32 @@
"""MoE 门控与专家选择top-k 激活数开关 + 跨层专家预测gate 前向 + argpartition"""
import mlx.core as mx
from mlx_streaming import config
def _effective_top_k(default_k: int) -> int:
"""实验开关:降低每 token 激活专家数以测速度/质量取舍。默认不改变模型。"""
override = config.moe_topk_override()
if not override:
return default_k
return max(1, min(default_k, int(override)))
def _predict_layer_experts(norm, gate, top_k: int, x: mx.array, mult: int) -> "tuple[dict[int, float], int]":
"""用 gate(norm(x)) 预测专家集,返回 ({expert_id: 最大 softmax 分数}, num_experts)。
norm/gate 必须取自**目标层**被预测的那层以匹配 probe 验证的
gate_L(post_attention_layernorm_L(h)) 配置recall_miss0.95
"""
gates = mx.softmax(gate(norm(x)), axis=-1, precise=True)
num_experts = gates.shape[-1]
k = min(num_experts, top_k * mult)
inds = mx.argpartition(gates, kth=-k, axis=-1)[..., -k:]
vals = mx.take_along_axis(gates, inds, axis=-1)
mx.eval(inds, vals)
best: "dict[int, float]" = {}
for e, s in zip(inds.reshape(-1).tolist(), vals.reshape(-1).tolist()):
e, s = int(e), float(s)
if s > best.get(e, -1.0):
best[e] = s
return best, num_experts

View File

@ -0,0 +1,504 @@
"""CPython native MoE backend 的 Python 包装层。"""
import importlib
import json
import os
import time
from collections import OrderedDict
from functools import lru_cache
import mlx.core as mx
import numpy as np
from mlx_streaming import config
_EXPERT_STAGE_CACHE: "OrderedDict[tuple, tuple[mx.array, mx.array, mx.array]]" = OrderedDict()
_BUNDLE_STAGE_CACHE: "OrderedDict[tuple, tuple[mx.array, ...]]" = OrderedDict()
_SLOT_POOLS: dict[tuple, "NativeComputeSlotPool"] = {}
_STAGE_STATS = {
"expert_hits": 0,
"expert_misses": 0,
"bundle_hits": 0,
"bundle_misses": 0,
"evictions": 0,
"can_checks": 0,
"can_true": 0,
"can_false": 0,
"calls": 0,
"route_sync_s": 0.0,
"stage_s": 0.0,
"enqueue_s": 0.0,
}
class NativeComputeSlotPool:
"""按层维护 compute-buffer 常驻 slothit 时复用 [cap,...] MLX arrays。"""
def __init__(self, compute_dir: str, layer: int, hidden: int, inter: int,
group: int, bits: int, num_experts: int, cap: int):
self.compute_dir = compute_dir
self.layer = int(layer)
self.hidden = int(hidden)
self.inter = int(inter)
self.group = int(group)
self.bits = int(bits)
self.num_experts = int(num_experts)
self.cap = int(cap)
self.slot_of: "OrderedDict[int, int]" = OrderedDict()
self.expert_of: list[int | None] = [None] * self.cap
self.slot_arrays = [self._empty_slot() for _ in range(self.cap)]
self.pool_arrays: tuple[mx.array, ...] | None = None
self.dirty = True
self.hits = 0
self.misses = 0
self.evictions = 0
self.rebuilds = 0
def _empty_slot(self) -> tuple[mx.array, ...]:
gu_words = self.hidden * self.bits // 32
gu_groups = self.hidden // self.group
down_words = self.inter * self.bits // 32
down_groups = self.inter // self.group
return (
mx.zeros((self.inter, gu_words), dtype=mx.uint32),
mx.zeros((self.inter, gu_groups), dtype=mx.uint16),
mx.zeros((self.inter, gu_groups), dtype=mx.uint16),
mx.zeros((self.inter, gu_words), dtype=mx.uint32),
mx.zeros((self.inter, gu_groups), dtype=mx.uint16),
mx.zeros((self.inter, gu_groups), dtype=mx.uint16),
mx.zeros((self.hidden, down_words), dtype=mx.uint32),
mx.zeros((self.hidden, down_groups), dtype=mx.uint16),
mx.zeros((self.hidden, down_groups), dtype=mx.uint16),
)
def _load_slot(self, expert: int) -> tuple[mx.array, ...]:
gate = _stage_one_projection_uncached(
self.compute_dir, self.layer, "gate_proj", expert,
self.inter, self.hidden, self.group, self.bits, self.num_experts)
up = _stage_one_projection_uncached(
self.compute_dir, self.layer, "up_proj", expert,
self.inter, self.hidden, self.group, self.bits, self.num_experts)
down = _stage_one_projection_uncached(
self.compute_dir, self.layer, "down_proj", expert,
self.hidden, self.inter, self.group, self.bits, self.num_experts)
return (*gate, *up, *down)
def _assign_slot(self, expert: int) -> int:
cached = self.slot_of.get(expert)
if cached is not None:
self.slot_of.move_to_end(expert)
self.hits += 1
return cached
self.misses += 1
if len(self.slot_of) < self.cap:
slot = next(i for i, e in enumerate(self.expert_of) if e is None)
else:
old_expert, slot = self.slot_of.popitem(last=False)
self.expert_of[slot] = None
self.evictions += 1
self.slot_arrays[slot] = self._load_slot(expert)
self.expert_of[slot] = expert
self.slot_of[expert] = slot
self.dirty = True
return slot
def _rebuild_pool_arrays(self) -> None:
self.pool_arrays = tuple(
mx.stack([slot[i] for slot in self.slot_arrays], axis=0)
for i in range(9)
)
self.rebuilds += 1
self.dirty = False
def acquire(self, expert_ids: list[int]) -> tuple[list[int], tuple[mx.array, ...]]:
local = [self._assign_slot(int(e)) for e in expert_ids]
if self.pool_arrays is None or self.dirty:
self._rebuild_pool_arrays()
assert self.pool_arrays is not None
return local, self.pool_arrays
def stats(self) -> dict:
return {
"cap": self.cap,
"entries": len(self.slot_of),
"hits": self.hits,
"misses": self.misses,
"evictions": self.evictions,
"rebuilds": self.rebuilds,
}
def _compute_dir() -> str:
return config.compute_buffer_dir() or os.path.join(
config.expert_dir(default=""), "compute_buffers")
@lru_cache(maxsize=1)
def _load_ext():
try:
return importlib.import_module("mlx_streaming.native_moe_ext")
except Exception:
return None
@lru_cache(maxsize=256)
def _has_layer_buffers(compute_dir: str, layer: int) -> bool:
for proj in ("gate_proj", "up_proj", "down_proj"):
for name in ("weight", "scales", "biases"):
path = os.path.join(compute_dir, f"layer{layer:02d}.{proj}.{name}.bin")
if not os.path.exists(path):
return False
return True
@lru_cache(maxsize=1024)
def _buffer_index_matches(compute_dir: str, layer: int, proj: str,
out_dim: int, in_dim: int, group: int, bits: int,
num_experts: int) -> bool:
path = os.path.join(compute_dir, f"layer{layer:02d}.{proj}.index.json")
if not os.path.exists(path):
return False
try:
with open(path) as f:
index = json.load(f)
except Exception:
return False
tensors = index.get("tensors", {})
weight_shape = tensors.get("weight", {}).get("shape_per_expert")
scale_shape = tensors.get("scales", {}).get("shape_per_expert")
bias_shape = tensors.get("biases", {}).get("shape_per_expert")
expected_weight = [out_dim, in_dim * bits // 32]
expected_grouped = [out_dim, in_dim // group]
return (
int(index.get("num_experts", -1)) == int(num_experts)
and weight_shape == expected_weight
and scale_shape == expected_grouped
and bias_shape == expected_grouped
)
@lru_cache(maxsize=512)
def _has_matching_compute_buffers(compute_dir: str, layer: int,
hidden: int, inter: int, group: int,
bits: int, num_experts: int) -> bool:
if not _has_layer_buffers(compute_dir, layer):
return False
# index.json 是防呆线bit/group 或目录错配时宁可回退 MLX不冒险读错 mmap shape。
return (
_buffer_index_matches(
compute_dir, layer, "gate_proj", inter, hidden, group, bits, num_experts)
and _buffer_index_matches(
compute_dir, layer, "up_proj", inter, hidden, group, bits, num_experts)
and _buffer_index_matches(
compute_dir, layer, "down_proj", hidden, inter, group, bits, num_experts)
)
@lru_cache(maxsize=256)
def _projection_memmaps(compute_dir: str, layer: int, proj: str,
out_dim: int, in_dim: int, group: int, bits: int,
num_experts: int):
words = in_dim * bits // 32
groups = in_dim // group
base = os.path.join(compute_dir, f"layer{layer:02d}.{proj}")
weight = np.memmap(
base + ".weight.bin",
dtype=np.uint32,
mode="r",
shape=(num_experts, out_dim, words),
)
scales = np.memmap(
base + ".scales.bin",
dtype=np.uint16,
mode="r",
shape=(num_experts, out_dim, groups),
)
biases = np.memmap(
base + ".biases.bin",
dtype=np.uint16,
mode="r",
shape=(num_experts, out_dim, groups),
)
return weight, scales, biases
def _stage_projection(compute_dir: str, layer: int, proj: str,
expert_ids: list[int], out_dim: int, in_dim: int,
group: int, bits: int, num_experts: int):
if config.native_moe_stage_cache():
staged = [
_stage_one_projection(
compute_dir, layer, proj, int(e), out_dim, in_dim, group, bits, num_experts)
for e in expert_ids
]
return tuple(mx.stack([entry[i] for entry in staged], axis=0) for i in range(3))
return _stage_projection_uncached(
compute_dir, layer, proj, expert_ids, out_dim, in_dim, group, bits, num_experts)
def _stage_projection_uncached(compute_dir: str, layer: int, proj: str,
expert_ids: list[int], out_dim: int, in_dim: int,
group: int, bits: int, num_experts: int):
weight, scales, biases = _projection_memmaps(
compute_dir, layer, proj, out_dim, in_dim, group, bits, num_experts)
ids = np.asarray(expert_ids, dtype=np.int64)
# advanced indexing 已经产出 numpy 临时数组mx.array 再接管成 MLX-managed staging arrays。
return (
mx.array(np.asarray(weight[ids])),
mx.array(np.asarray(scales[ids])),
mx.array(np.asarray(biases[ids])),
)
def _stage_one_projection_uncached(compute_dir: str, layer: int, proj: str, expert: int,
out_dim: int, in_dim: int, group: int, bits: int,
num_experts: int):
weight, scales, biases = _projection_memmaps(
compute_dir, layer, proj, out_dim, in_dim, group, bits, num_experts)
return (
mx.array(np.asarray(weight[int(expert)])),
mx.array(np.asarray(scales[int(expert)])),
mx.array(np.asarray(biases[int(expert)])),
)
def _stage_one_projection(compute_dir: str, layer: int, proj: str, expert: int,
out_dim: int, in_dim: int, group: int, bits: int,
num_experts: int):
key = (compute_dir, int(layer), proj, int(expert), out_dim, in_dim, group, bits, num_experts)
cached = _EXPERT_STAGE_CACHE.get(key)
if cached is not None:
_EXPERT_STAGE_CACHE.move_to_end(key)
_STAGE_STATS["expert_hits"] += 1
return cached
_STAGE_STATS["expert_misses"] += 1
staged = _stage_one_projection_uncached(
compute_dir, layer, proj, expert, out_dim, in_dim, group, bits, num_experts)
_EXPERT_STAGE_CACHE[key] = staged
_EXPERT_STAGE_CACHE.move_to_end(key)
limit = max(0, config.native_moe_stage_cache_experts())
while limit and len(_EXPERT_STAGE_CACHE) > limit:
_EXPERT_STAGE_CACHE.popitem(last=False)
_STAGE_STATS["evictions"] += 1
return staged
def _slot_pool(compute_dir: str, layer: int, hidden: int, inter: int, group: int,
bits: int, num_experts: int) -> NativeComputeSlotPool:
cap = max(1, config.native_moe_slot_cap())
key = (compute_dir, int(layer), hidden, inter, group, bits, num_experts, cap)
pool = _SLOT_POOLS.get(key)
if pool is None:
pool = NativeComputeSlotPool(
compute_dir, int(layer), hidden, inter, group, bits, num_experts, cap)
_SLOT_POOLS[key] = pool
return pool
def _stage_bundle(compute_dir: str, layer: int, expert_ids: list[int],
hidden: int, inter: int, group: int, bits: int,
num_experts: int):
key = (
compute_dir, int(layer), tuple(int(e) for e in expert_ids),
hidden, inter, group, bits, num_experts,
)
if config.native_moe_stage_bundle_cache():
cached = _BUNDLE_STAGE_CACHE.get(key)
if cached is not None:
_BUNDLE_STAGE_CACHE.move_to_end(key)
_STAGE_STATS["bundle_hits"] += 1
return cached
_STAGE_STATS["bundle_misses"] += 1
gate = _stage_projection(
compute_dir, int(layer), "gate_proj", expert_ids, inter, hidden, group, bits, num_experts)
up = _stage_projection(
compute_dir, int(layer), "up_proj", expert_ids, inter, hidden, group, bits, num_experts)
down = _stage_projection(
compute_dir, int(layer), "down_proj", expert_ids, hidden, inter, group, bits, num_experts)
staged = (*gate, *up, *down)
if config.native_moe_stage_bundle_cache():
_BUNDLE_STAGE_CACHE[key] = staged
_BUNDLE_STAGE_CACHE.move_to_end(key)
limit = max(0, config.native_moe_stage_cache_bundles())
while limit and len(_BUNDLE_STAGE_CACHE) > limit:
_BUNDLE_STAGE_CACHE.popitem(last=False)
_STAGE_STATS["evictions"] += 1
return staged
def prefetch_native_moe_stage(layer: int, expert_ids: list[int],
hidden: int, inter: int, group: int, bits: int,
num_experts: int) -> bool:
"""提前把预测专家拷入 MLX-managed staging cache供后续 fused native op 复用。"""
if not config.native_moe():
return False
if not config.native_moe_stage_prefetch():
return False
if config.native_moe_synthetic():
return False
compute_dir = _compute_dir()
if not compute_dir or not _has_matching_compute_buffers(
compute_dir, layer, hidden, inter, group, bits, num_experts):
return False
try:
uniq = list(dict.fromkeys(int(e) for e in expert_ids))
if not uniq:
return False
for proj, out_dim, in_dim in (
("gate_proj", inter, hidden),
("up_proj", inter, hidden),
("down_proj", hidden, inter),
):
for e in uniq:
_stage_one_projection(
compute_dir, int(layer), proj, e, out_dim, in_dim, group, bits, num_experts)
return True
except Exception:
if config.native_moe_raise():
raise
return False
def stage_cache_stats() -> dict:
slot_stats = {
"pools": len(_SLOT_POOLS),
"entries": sum(len(pool.slot_of) for pool in _SLOT_POOLS.values()),
"hits": sum(pool.hits for pool in _SLOT_POOLS.values()),
"misses": sum(pool.misses for pool in _SLOT_POOLS.values()),
"evictions": sum(pool.evictions for pool in _SLOT_POOLS.values()),
"rebuilds": sum(pool.rebuilds for pool in _SLOT_POOLS.values()),
}
return {
**_STAGE_STATS,
"route_sync_s": round(_STAGE_STATS["route_sync_s"], 6),
"stage_s": round(_STAGE_STATS["stage_s"], 6),
"enqueue_s": round(_STAGE_STATS["enqueue_s"], 6),
"expert_entries": len(_EXPERT_STAGE_CACHE),
"bundle_entries": len(_BUNDLE_STAGE_CACHE),
"slot_pool": slot_stats,
}
def note_route_sync(seconds: float) -> None:
_STAGE_STATS["route_sync_s"] += float(seconds)
def can_native_moe(layer: int, hidden: int, inter: int, group: int, bits: int,
num_experts: int) -> bool:
"""只做缓存化的轻量可用性判断;失败时避免上层触发 route host 同步。"""
_STAGE_STATS["can_checks"] += 1
if not config.native_moe():
_STAGE_STATS["can_false"] += 1
return False
if not config.native_moe_mlx_op():
_STAGE_STATS["can_false"] += 1
return False
if config.native_moe_synthetic():
ok = _load_ext() is not None
_STAGE_STATS["can_true" if ok else "can_false"] += 1
return ok
if (hidden * bits) % 32 != 0 or (inter * bits) % 32 != 0:
_STAGE_STATS["can_false"] += 1
return False
compute_dir = _compute_dir()
if not compute_dir:
_STAGE_STATS["can_false"] += 1
return False
ok = (
_load_ext() is not None
and _has_matching_compute_buffers(
compute_dir, int(layer), hidden, inter, group, bits, num_experts)
)
_STAGE_STATS["can_true" if ok else "can_false"] += 1
return ok
def try_native_moe(layer: int, expert_ids: list[int], x: mx.array, scores: mx.array,
hidden: int, inter: int, group: int, bits: int,
num_experts: int) -> "mx.array | None":
"""尝试 native fused MoE失败返回 None 让调用方回退 MLX 路径。"""
if not config.native_moe():
return None
if not config.native_moe_mlx_op():
return None
if (hidden * bits) % 32 != 0 or (inter * bits) % 32 != 0:
return None
if not expert_ids:
return None
if not can_native_moe(layer, hidden, inter, group, bits, num_experts):
return None
compute_dir = _compute_dir()
ext = _load_ext()
if ext is None:
return None
synthetic = config.native_moe_synthetic()
if (not synthetic and not _has_matching_compute_buffers(
compute_dir, layer, hidden, inter, group, bits, num_experts)):
return None
try:
expert_arr = mx.array([int(e) for e in expert_ids], dtype=mx.uint32)
if synthetic:
t_enqueue = time.perf_counter()
y = ext.fused_moe(
x.astype(mx.float32),
expert_arr,
scores.astype(mx.float32),
compute_dir,
int(layer),
hidden,
inter,
group,
bits,
num_experts,
True,
)
_STAGE_STATS["enqueue_s"] += time.perf_counter() - t_enqueue
else:
t_stage = time.perf_counter()
if config.native_moe_slot_pool():
pool = _slot_pool(
compute_dir, int(layer), hidden, inter, group, bits, num_experts)
local_slots, pool_arrays = pool.acquire(expert_ids)
local_arr = mx.array(local_slots, dtype=mx.uint32)
(
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b,
) = pool_arrays
else:
local_arr = None
(
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b,
) = _stage_bundle(
compute_dir, int(layer), expert_ids,
hidden, inter, group, bits, num_experts)
_STAGE_STATS["stage_s"] += time.perf_counter() - t_stage
t_enqueue = time.perf_counter()
if local_arr is not None:
y = ext.fused_moe_slots(
x.astype(mx.float32),
local_arr,
scores.astype(mx.float32),
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b,
hidden, inter, group, bits,
)
else:
y = ext.fused_moe_staged(
x.astype(mx.float32),
scores.astype(mx.float32),
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b,
hidden, inter, group, bits,
)
_STAGE_STATS["enqueue_s"] += time.perf_counter() - t_enqueue
_STAGE_STATS["calls"] += 1
return y
except Exception:
if config.native_moe_raise():
raise
return None

View File

@ -0,0 +1 @@
"""专家预取模块:后台物化(bg_prefetch) / native staging(native_staging) / 跨层预测预取(cross_layer) / 模型 patch(patch)。"""

View File

@ -0,0 +1,133 @@
"""后台专家预取器:在独立 MLX stream 上物化预测专家(私有 array交接给主线程。
gating 测试验证可行probe_multistream_gate / probe_multistream_handoff
- 后台线程 `with mx.stream(s2):` 物化 + mx.eval可与主线程计算重叠不崩
- 只物化私有array绝不写主线程共享池 stream 写共享张量会报 no Stream
- 主线程只消费 eval的交接 array stream 读安全
"""
import queue
import threading
from collections import OrderedDict
import mlx.core as mx
from mlx_streaming import config
class BackgroundExpertPrefetcher:
def __init__(self, blob_source, window: int = 3, native: bool = False):
self._src = blob_source
self._stream = mx.new_stream(mx.default_device())
self._q: "queue.Queue" = queue.Queue()
self._ready: "OrderedDict[tuple, dict]" = OrderedDict()
self._ready_layers: "list[int]" = []
self._window = window
self._native = native
self._lock = threading.Lock()
self._stop = False
self.submitted = 0
self.materialized = 0
self.taken = 0
# 就绪率诊断promote 时已物化好的(ready_on_time) vs 提交了但还在飞行中(not_ready)。
# 用于量化 attention/GDN 窗口是否够长隐藏 I/O。
self.ready_on_time = 0
self.not_ready = 0
self.materialize_s = 0.0 # bg 线程物化(load+eval)累计墙钟,验证"串行加性"假设
self._native_mat = config.native_materialize()
self._inflight: "dict[int, set]" = {}
self._t = threading.Thread(target=self._loop, daemon=True)
self._t.start()
def submit(self, layer: int, expert_ids) -> None:
ids = [int(e) for e in expert_ids]
if not ids:
return
with self._lock:
self.submitted += len(ids)
self._inflight.setdefault(int(layer), set()).update(ids)
self._q.put((int(layer), ids))
def _loop(self) -> None:
while not self._stop:
try:
layer, ids = self._q.get(timeout=0.1)
except queue.Empty:
continue
try:
import time as _t
_t0 = _t.perf_counter()
with mx.stream(self._stream):
# view_bf16=False后台不做 .view(bfloat16)(会报 no Stream(gpu,0)
# scales/biases 留 uint16主线程 take 时再 view。
# NATIVE_MATERIALIZE=1用 C++ blob_load 把拷贝挪出 GIL减少对主线程争用
if self._native_mat:
experts = self._src.load_experts_native(layer, ids, view_bf16=False)
else:
experts = self._src.load_experts(layer, ids, view_bf16=False)
mx.eval([v for d in experts.values() for v in d.values()])
self.materialize_s += _t.perf_counter() - _t0
with self._lock:
pend = self._inflight.get(layer)
for e, d in experts.items():
self._ready[(layer, int(e))] = d
self.materialized += 1
if pend is not None:
pend.discard(int(e))
if layer not in self._ready_layers:
self._ready_layers.append(layer)
while len(self._ready_layers) > self._window:
old = self._ready_layers.pop(0)
for k in [k for k in self._ready if k[0] == old]:
del self._ready[k]
except Exception:
# 后台失败不影响主线程:主线程会优雅回退到同步 demand 路径。
pass
@staticmethod
def _view_bf16(d: dict) -> dict:
"""主线程消费时把 scales/biases 从 uint16 位重解释回 bfloat16。"""
out = {}
for k, v in d.items():
out[k] = v.view(mx.bfloat16) if (k.endswith(".scales") or k.endswith(".biases")) else v
return out
def take_ready(self, layer: int, e: int) -> "dict | None":
with self._lock:
d = self._ready.pop((int(layer), int(e)), None)
if d is None:
return None
self.taken += 1
return self._view_bf16(d)
def take_ready_layer(self, layer: int) -> "dict[int, dict]":
"""取走某层所有就绪专家:{expert_id: {proj.tensor: mx.array}}(主线程调用)。"""
raw = {}
with self._lock:
for key in [k for k in self._ready if k[0] == int(layer)]:
raw[key[1]] = self._ready.pop(key)
self.taken += 1
return {e: self._view_bf16(d) for e, d in raw.items()}
def ready_count(self, layer: int) -> int:
with self._lock:
return sum(1 for k in self._ready if k[0] == int(layer))
def note_promote(self, layer: int, ready_now: int) -> None:
"""promote_prefetched 调用:记录本次 promote 时已就绪数 vs 仍在飞行中(窗口没盖住)。"""
with self._lock:
self.ready_on_time += int(ready_now)
pend = self._inflight.get(int(layer))
if pend:
self.not_ready += len(pend)
pend.clear()
def stats(self) -> dict:
with self._lock:
return {"submitted": self.submitted, "materialized": self.materialized,
"taken": self.taken, "ready": len(self._ready),
"ready_on_time": self.ready_on_time, "not_ready": self.not_ready,
"materialize_s": round(self.materialize_s, 3)}
def close(self) -> None:
self._stop = True
self._t.join(timeout=2)

View File

@ -0,0 +1,181 @@
"""跨层专家预取:在 attention/GDN 之前用本层输入预测目标层 MoE 路由并提前取专家。
通过 monkeypatch `Qwen3NextDecoderLayer.__call__`在真正进入 MoE 计算前用
gate_L(post_attention_layernorm_L(h)) 预测目标层专家对齐 recall0.95 的探针口径
top-(top_k*mult) 专家在当前层计算窗口内异步预取藏住读盘/物化延迟
支持多条预取后端按环境开关择一STREAM_BLOB_BG后台物化进池STREAM_BLOB
blob 字节预读STREAM_BLOB_LOADERblob-loader 字节预热native stage 预取
"""
import time
import mlx.core as mx
from mlx_streaming import config
from mlx_streaming.core.moe import native_moe
from mlx_streaming.core.moe.gate import _predict_layer_experts
from mlx_streaming.core.moe.block import FileStreamingMoeBlock
# 是否已对 decoder layer 打过预取补丁(全局只打一次)。
_CROSS_LAYER_PREFETCH_PATCHED = False
# 每次「整模型前向」内的全局预取预算(按层序回绕重置),避免无界预取撑爆窗口。
_FILE_PREFETCH_REMAINING = 0
_FILE_PREFETCH_LAST_LAYER = -1
_STAGE_PREFETCH_REMAINING = 0
_STAGE_PREFETCH_LAST_LAYER = -1
def _cross_layer_prefetch_mult() -> int:
return config.cross_layer_mult()
def _submit_missing_prefetch(store, layer: int, predicted) -> "list[int]":
"""只把"预测∩非常驻"的专家提交给后台预取器,避免重载常驻、撑爆 attention 窗口。
返回实际提交的缺失专家列表供测试/统计
"""
bg = getattr(store, "_bg", None)
if bg is None:
return []
resident = store.resident_experts(layer) if hasattr(store, "resident_experts") else set()
missing = [int(e) for e in predicted if int(e) not in resident]
if missing:
bg.submit(int(layer), missing)
return missing
def _prefetch_budget(layer_idx: int, kind: str) -> int:
"""按一次层序前向重置全局预算layer 序号回绕说明进入新 forward。"""
global _FILE_PREFETCH_REMAINING, _FILE_PREFETCH_LAST_LAYER
global _STAGE_PREFETCH_REMAINING, _STAGE_PREFETCH_LAST_LAYER
if kind == "file":
total = config.file_prefetch_global_budget()
if layer_idx <= _FILE_PREFETCH_LAST_LAYER:
_FILE_PREFETCH_REMAINING = total
_FILE_PREFETCH_LAST_LAYER = layer_idx
return max(0, _FILE_PREFETCH_REMAINING)
total = config.stage_prefetch_global_budget()
if layer_idx <= _STAGE_PREFETCH_LAST_LAYER:
_STAGE_PREFETCH_REMAINING = total
_STAGE_PREFETCH_LAST_LAYER = layer_idx
return max(0, _STAGE_PREFETCH_REMAINING)
def _consume_prefetch_budget(n: int, kind: str) -> None:
global _FILE_PREFETCH_REMAINING, _STAGE_PREFETCH_REMAINING
if kind == "file":
_FILE_PREFETCH_REMAINING = max(0, _FILE_PREFETCH_REMAINING - n)
else:
_STAGE_PREFETCH_REMAINING = max(0, _STAGE_PREFETCH_REMAINING - n)
def _prefetch_layers(kind: str) -> "set[int] | None":
env = "FILE_PREFETCH_LAYERS" if kind == "file" else "STAGE_PREFETCH_LAYERS"
return config.parse_layers_env(env)
def enable_cross_layer_prefetch():
"""给 Qwen3NextDecoderLayer 加跨层专家预取。
attention/GDN 之前用当前层输入 hidden post_attention_layernorm 近似预测
本层 MoE 路由并把 top-(top_k*mult) 专家预取进 resident pool
"""
global _CROSS_LAYER_PREFETCH_PATCHED
if _CROSS_LAYER_PREFETCH_PATCHED:
return
from mlx_lm.models.qwen3_next import Qwen3NextDecoderLayer
orig_call = Qwen3NextDecoderLayer.__call__
def patched_call(self, x, mask=None, cache=None):
mlp = getattr(self, "mlp", None)
if mlp is not None and getattr(getattr(mlp, "store", None), "_staging", None) is not None:
# 存本层「未归一化」decoder 输入:供 native-fused-prefetch 用目标层 norm 正确预测
# (对齐 0.95-recall 探针MoE forward 拿到的 x 已被本层 norm 过norm 用错会掉点)。
mlp._unnormed_input = x
if config.cross_layer_prefetch() and config.resident_pool_enabled():
ahead = config.cross_layer_ahead(default=0)
target_layer = getattr(self, "_layer_idx", None)
if target_layer is not None:
target_layer += ahead
target_mlp = None
target_decoder = None
if ahead == 0 and isinstance(mlp, FileStreamingMoeBlock):
target_mlp = mlp
target_decoder = self
elif target_layer is not None:
layers = getattr(getattr(self, "_prefetch_model_ref", None), "layers", [])
if 0 <= target_layer < len(layers):
target_decoder = layers[target_layer]
target_mlp = getattr(target_decoder, "mlp", None)
if isinstance(target_mlp, FileStreamingMoeBlock):
if config.probe_predict_only():
# 诊断:只做预测 gate 前向 + eval(同步在预测上),不 .tolist、不预取。
# 用于拆分 32% 税≈S(+9%)→税是 Python 编排(native 可救)≈B(+32%)→税是
# gate前向被 barrier 暴露在关键路径(简单 native 派发救不了)。
g = target_mlp.gate(target_decoder.post_attention_layernorm(x))
kk = min(g.shape[-1], 4)
ip = mx.argpartition(g, kth=-kk, axis=-1)[..., -kk:]
mx.eval(ip)
return orig_call(self, x, mask=mask, cache=cache)
# 关键:用**目标层**的 post_attention_layernorm匹配 probe 验证的
# gate_L(post_norm_L(h_{L-ahead})) 配置recall_miss≈0.95)。
best, num_experts = _predict_layer_experts(
target_decoder.post_attention_layernorm, target_mlp.gate,
target_mlp.top_k, x, _cross_layer_prefetch_mult())
bg = getattr(target_mlp.store, "_bg", None)
if config.stream_blob_bg() and bg is not None:
# 同层(AHEAD=0)预测 + 只预取"预测∩非常驻",藏进 attention/GDN 窗口。
budget = config.stream_blob_bg_budget(
default=target_mlp.top_k * _cross_layer_prefetch_mult())
picked = [e for e, _ in sorted(best.items(), key=lambda kv: kv[1], reverse=True)][:budget]
_submit_missing_prefetch(target_mlp.store, target_mlp.layer_idx, picked)
if config.window_prof():
target_mlp._submit_t = time.perf_counter()
return orig_call(self, x, mask=mask, cache=cache)
if config.stream_blob() and target_mlp._blob is not None:
# 全流式:后台预读下一层预测专家的字节,与当前层计算重叠。
budget = config.stream_blob_prefetch_budget(default=target_mlp.top_k * 2)
picked = [e for e, _ in sorted(best.items(), key=lambda kv: kv[1], reverse=True)][:budget]
target_mlp._blob.prefetch_async(target_mlp.layer_idx, picked)
return orig_call(self, x, mask=mask, cache=cache)
blob_loader = getattr(target_mlp.store, "_blob_loader", None)
if config.stream_blob_loader() and blob_loader is not None:
# blob-loader 路径:后台只预读字节(与当前层计算重叠),
# demand 时主线程从预热字节快速物化进常驻池,不在主线程同步阻塞读盘。
budget = config.stage_prefetch_per_layer_budget(default=12)
picked = [e for e, _ in sorted(best.items(), key=lambda kv: kv[1], reverse=True)][:budget]
blob_loader.prefetch_async(target_mlp.layer_idx, picked)
return orig_call(self, x, mask=mask, cache=cache)
allowed_layers = _prefetch_layers("stage")
if allowed_layers is not None and target_mlp.layer_idx not in allowed_layers:
return orig_call(self, x, mask=mask, cache=cache)
min_score = config.stage_prefetch_min_score()
per_layer = config.stage_prefetch_per_layer_budget(default=4)
remaining = _prefetch_budget(target_mlp.layer_idx, "stage")
budget = max(0, min(per_layer, remaining))
expert_ids = [
e for e, s in sorted(best.items(), key=lambda kv: kv[1], reverse=True)
if s >= min_score
][:budget]
_consume_prefetch_budget(len(expert_ids), "stage")
proj_bits = getattr(target_mlp._sub, "proj_bits", {})
bits_set = {
int(proj_bits.get(name, target_mlp.bits))
for name in ("gate_proj", "up_proj", "down_proj")
}
if len(bits_set) == 1 and native_moe.prefetch_native_moe_stage(
target_mlp.layer_idx,
expert_ids,
target_mlp.hidden,
target_mlp.moe_inter,
target_mlp.group_size,
bits_set.pop(),
num_experts,
):
pass
else:
target_mlp.store.prefetch(target_mlp.layer_idx, expert_ids)
return orig_call(self, x, mask=mask, cache=cache)
Qwen3NextDecoderLayer.__call__ = patched_call
_CROSS_LAYER_PREFETCH_PATCHED = True

View File

@ -0,0 +1,229 @@
"""native-fused-prefetch 的 miss→hit 落地per-layer 环形 staging buffer + 池 promote。
机制预读 IO 全在 GPU 完成回调里主线程零 pread/ .tolist/ npmx.array
- submit(layer, inds_lazy)取该层环里的一块 buffer prefetch_into_staging
GPU 完成回调C++把专家字节 pread 进这块 buffer并按 gen 原子记录 (expertrow)
- promote(layer, store) C++ 记录纯锁 gen 取回 handler 真正写过的那块 buffer
惰性切片直接 _place_expert 写进池槽 eval避免每层 host 同步拖慢 verify 流水线
为何必须 per-layer而非全局共享小环预读在 GPU 完成回调里异步发生 batched verify
整块前向最后一次 eval同一前向里全部 48 层的 submit 都在飞行 handler 可能要到
接近末尾才触发若用全局小环buffer 会在它的 handler 落地/promote 消费**之前**就被后续层的
submit 复用 handler 覆盖成别层数据gen 匹配只保证 buffer 对象保证不了内容新鲜 非确定性
串扰损坏实测全局环 n_mismatch=60~85所以每层各留独立 buffer环大小只用于跨 token 复用
内存 = 层数 × ring × budget × stridering 默认从 4 降到 2STAGING_RING在保持逐位精确的前提下
staging 占用减半48×2×budget×stride
"""
import os
import time
import mlx.core as mx
from mlx_streaming import config
from mlx_streaming.core.profiling import note_tprof, TPROF_ON
# 诊断门控STG_VERIFY开启时 promote 记录每个 (layer, expert) 落池所用 gen供消费侧字节校验
# 定位损坏来自哪一代 staging。默认关对主路径零影响关时连 dict 写入都不发生)。
_STG_VERIFY = os.environ.get("STG_VERIFY") == "1"
def _typed_seg(seg, dt, shape):
"""按段 dtype 重解释uint32 原样 / uint16→bfloat16 / uint8(mxfp4 scales) 原样。"""
if dt == "uint32":
return seg.view(mx.uint32).reshape(shape)
if dt == "uint16":
return seg.view(mx.uint16).reshape(shape).view(mx.bfloat16)
return seg.reshape(shape)
def route_used_subset(cand: "list[int]", route_inds: mx.array, num_experts: int) -> "set[int]":
"""返回 cand 中“被 route_inds 真实路由到”的子集GPU membershipdrain ≤len(cand))。
cand预读候选专家 idhost 已知通常 budget
route_inds本层真实路由 idGPU lazy任意 shape内部 reshape(-1)
num_experts专家总数 mask
cand 为空 返回空集不触发任何 GPU op
仅对预读候选做成员测试不物化全量路由只把 len(cand) 个布尔拉回 host
"""
if not cand:
return set()
route_mask = mx.zeros((num_experts,), dtype=mx.uint8)
route_mask[route_inds.reshape(-1)] = 1 # GPU scatter幂等
hit = route_mask[mx.array(cand, dtype=mx.uint32)] # GPU gather≤len(cand)
hit_list = hit.tolist() # 唯一一次极小 drain
return {int(cand[i]) for i in range(len(cand)) if hit_list[i]}
class NativeStagingManager:
"""per-layer 环形 staging每层 ring 块 buffer 轮转,按 gen 严格匹配回收。
submit ring[layer][rr] (genbuffer)C++ handler gen 原子记 (expertrow)
promote gen 取回**正是 handler 写过**的那块 buffer只要某层的 buffer 在该层下一次ring
submit 之后被复用前 handler 已落地且 promote 的惰性切片已随本前向 eval 被消费即不串扰
"""
def __init__(self, blob_source, budget: int, ring: int | None = None):
self.src = blob_source # BlobExpertSource有 dir / stride / _segs
self.budget = int(budget)
self.stride = int(blob_source.stride)
# 安全下限=2MTP 每步 verify+replay 会对同层各 submit 一次2 个在飞ring=1 实测串扰损坏。
self.ring = max(2, int(ring) if ring is not None else config.staging_ring())
self._ring: "dict[int, list]" = {} # layer -> [mx.array]*ring
self._rr: "dict[int, int]" = {} # layer -> 下一个要写的环索引
self._gen = 1 # 全局单调 generation
self._gen_buf: "dict[int, mx.array]" = {} # gen -> 该次 submit 写的 buffer
self.submitted = 0
self.promoted = 0
# 诊断用:每层最近一次 promote take 到的"已就绪(C++ pread 完成)专家集"。
# miss_attrib 据此把 miss_A 拆成 时序(不在就绪集) / 驱逐(在就绪集但 acquire 前已不在池)。
self.last_ready: "dict[int, set]" = {}
# 诊断用(STG_VERIFY):记录每个 (layer, expert) 最近一次被 promote 写池时所用的 gen
# 供消费侧字节校验定位「损坏来自哪一代 staging」。默认不影响主路径(只是写 dict)。
self.placed_gen: "dict[tuple, int]" = {}
def _bufs(self, layer: int) -> list:
bl = self._ring.get(layer)
if bl is None:
bl = []
for _ in range(self.ring):
a = mx.zeros((self.budget, self.stride), dtype=mx.uint8)
mx.eval(a)
bl.append(a)
self._ring[layer] = bl
self._rr[layer] = 0
return bl
def submit(self, layer: int, inds_lazy: mx.array, resident: "list[int] | None" = None):
"""方案Binds_lazy 为"预测宽集合"(lazy uint32 [N]按门控分降序N 可远大于 budget)。
resident目标层提交时刻的常驻专家快照host C++ 回调按 resident 过滤再按降序
取前 budget "缺口"专家 pread bufferbudget=buffer 行数返回 dummy折进图触发回调
"""
import mlx_streaming.native_moe_ext as _N
layer = int(layer)
bufs = self._bufs(layer)
rr = self._rr[layer]
buf = bufs[rr]
self._rr[layer] = (rr + 1) % self.ring
gen = self._gen
self._gen += 1
self._gen_buf[gen] = buf
if len(self._gen_buf) > self.ring * 64:
for g in sorted(self._gen_buf)[:len(self._gen_buf) - self.ring * 64]:
self._gen_buf.pop(g, None)
path = f"{self.src.dir}/layer{layer:02d}.blob"
self.submitted += 1
res = [int(e) for e in (resident or [])]
# cap=self.budget回调最多往 buffer 写 budget 行覆盖缺口分布p99≈15 → budget=16 够)。
return _N.prefetch_into_staging(
buf, inds_lazy, layer, gen, path, self.stride, res, self.budget,
config.staging_pread_parallel())
def promote(self, layer: int, store, used: "set[int] | None" = None,
route_inds: "mx.array | None" = None,
num_experts: "int | None" = None) -> int:
"""把该层 staging 里已就绪专家写进常驻池主线程、acquire 前)。
used本层真实路由专家集合host 已知时传入 promote "预读好 ∩ used" 的专家
route_inds/num_expertsGPU 重映射路径专用host 没有现成 used
route_used_subset 仅对预读候选做 GPU membership 现算 useddrain budget
三者皆缺 整批写入兜底行为同改动前
"""
import mlx_streaming.native_moe_ext as _N
layer = int(layer)
_pt0 = time.perf_counter() if TPROF_ON else 0.0
flat = _N.prefetch_staging_take(layer)
if TPROF_ON:
note_tprof("take_s", time.perf_counter() - _pt0)
if not flat:
self.last_ready[layer] = set() # 本层无就绪:miss_A 全归"时序(pread 未完成)"
if TPROF_ON:
note_tprof("promote_s", time.perf_counter() - _pt0, count_key="promote_n")
return 0
gen = int(flat[0])
# 记录本层"已就绪(pread 完成)"专家集:无论 buffer 是否匹配pread 都已完成。
self.last_ready[layer] = {int(flat[i]) for i in range(1, len(flat), 2)}
stg = self._gen_buf.pop(gen, None) # 按 handler 记的 gen 取回**正是它写过**的那块 buffer
if stg is None:
if TPROF_ON:
note_tprof("promote_s", time.perf_counter() - _pt0, count_key="promote_n")
return 0 # buffer 已被回收/不匹配 → 跳过(不会错配)
store._resident._ensure_layer(layer) # promote 在 acquire 前,池可能还没建
resident = (store.resident_experts(layer)
if hasattr(store, "resident_experts") else set())
pairs = [(int(flat[i]), int(flat[i + 1])) for i in range(1, len(flat), 2)]
# GPU 重映射路径无现成 used仅对预读候选做 GPU membership 现算(不物化全量路由)。
if used is None and route_inds is not None and num_experts is not None:
_tr0 = time.perf_counter() if TPROF_ON else 0.0
used = route_used_subset([e for e, _ in pairs], route_inds, int(num_experts))
if TPROF_ON:
note_tprof("route_s", time.perf_counter() - _tr0)
# 受保护(不被驱逐)= 本层真实路由(若已知),否则退化为全部预读专家。
protect = set(used) if used is not None else {e for e, _ in pairs}
# 惰性切片直接放入池:正确性由 gen-匹配(切到 handler 真正写的那块 buffer+ per-layer 环
# (该 buffer 在该层 ring 次 submit 内不被覆盖,而 pool scatter 会在本前向读池时 eval保证。
# 不在此 mx.eval那是每层 host 同步,在 verify 大流水线上会显著拖慢热路径)。
# used 已知时只保留命中真实路由的专家:假阳性不写池(省 scatter、不挤掉有用专家
_tpl0 = time.perf_counter() if TPROF_ON else 0.0
pend = [(e, self._slice(stg[row])) for e, row in pairs
if e not in resident and (used is None or e in used)]
placed = 0
for e, d in pend:
try:
store._resident._place_expert(layer, e, d, current=protect)
except ValueError:
break
if _STG_VERIFY:
self.placed_gen[(layer, e)] = gen # 诊断(默认关):记录该专家这次落池用的 gen
placed += 1
if TPROF_ON:
# place_s 含惰性切片构建 + 池 scatter 入图(不含其 GPU 执行);place_experts=实际写入专家数。
note_tprof("place_s", time.perf_counter() - _tpl0,
count_key="place_experts", count=placed)
note_tprof("promote_s", time.perf_counter() - _pt0, count_key="promote_n")
self.promoted += placed
return placed
# ---- zero-copy 双源解码池侧区直接散写C++ 维护侧区 e→物理行 ----
def submit_pool_sideregion(self, layer, pred, resident, pool_list, base_row, gen=0):
"""zero-copy 预取:触发 C++ 把预读专家段散写进池侧区行(填代 gen。pool_list 为 _segs 顺序的 per-key 池数组。
seg_nbytes 取自 blob 段表 pool key 顺序一致spec_slots = 侧区行数=self.budget"""
import mlx_streaming.native_moe_ext as _N
layer = int(layer)
seg_nbytes = [int(nb) for *_, nb in self.src._segs] # 段字节(与 pool key 顺序一致)
path = f"{self.src.dir}/layer{layer:02d}.blob"
self.submitted += 1
res = [int(e) for e in (resident or [])]
return _N.prefetch_pool_sideregion(
pool_list, seg_nbytes, pred, layer, path, self.stride, res,
int(self.budget), int(base_row), gen=int(gen))
def sideregion_contents(self, layer, gen=0):
"""读该层某代 C++ 侧区缓存当前内容 {expert: 物理侧区行}(纯锁,不消费)。"""
import mlx_streaming.native_moe_ext as _N
flat = _N.sideregion_contents(int(layer), int(gen))
return {int(flat[i]): int(flat[i + 1]) for i in range(0, len(flat), 2)}
def sideregion_kv(self, layer, gen=0):
"""读该层某代侧区为 (keys uint32, vals int32) 两个 device mx.arrayC++ 直接建,供快路径合并)。"""
import mlx_streaming.native_moe_ext as _N
return _N.sideregion_kv(int(layer), int(gen))
def sideregion_reset(self):
"""清空 C++ 侧区缓存(换 prompt/重置统计时调用)。"""
import mlx_streaming.native_moe_ext as _N
_N.sideregion_reset()
def _slice(self, row: mx.array) -> dict:
# 按段 dtype 通用分派(支持 v1 affine 与 v2 mxfp4
# - uint32weight原样 reshape
# - uint16affine 的 scales/biasesview(bfloat16)
# - uint8 mxfp4 的 scales保持原样绝不 view bf16。
out = {}
off = 0
for proj, tensor, dt, shape, nb in self.src._segs:
seg = row[off:off + nb]
out[f"{proj}.{tensor}"] = _typed_seg(seg, dt, shape)
off += nb
return out

View File

@ -0,0 +1,56 @@
"""模型 patch把原生 MoE 块替换成流式块(文件后端 / 常驻切片版),并按需挂跨层预取。"""
from mlx_streaming import config
from mlx_streaming.core.moe.block import FileStreamingMoeBlock, StreamingMoeBlock
from mlx_streaming.core.prefetch.cross_layer import enable_cross_layer_prefetch
def patch_model_filebacked(model, store, hidden, moe_inter, group_size, bits,
proj_bits: dict | None = None,
layer_proj_bits: dict | None = None):
"""把每个 MoE 块替换为 FileStreamingMoeBlock并丢弃常驻的堆叠 switch_mlp。
storeFileExpertStore所有 MoE 层共用 (layer,expert) 缓存
proj_bits非空时走混合精度 proj 不同 bit专家须为对应混合重量化产出
layer_proj_bits{绝对层号: {proj:bits}}非空时逐层用各自 proj_bits优先于 proj_bits
对应 requantize_dir_layered 产出各层 QSL 用该层 bit与流式存盘文件一一对应
返回被替换的层数被替换后原 switch_mlp 不再被引用惰性权重不会被物化
"""
patched = 0
for i, layer in enumerate(model.layers):
layer._layer_idx = i
# 用 object.__setattr__ 存反向 model 引用,避免进 nn.Module(=dict)的子模块树:
# 否则 model→layer→model 成环,wired_limit 等 tree_reduce 遍历会无限递归(RecursionError)。
object.__setattr__(layer, "_prefetch_model_ref", model)
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
pb = layer_proj_bits.get(i, proj_bits) if layer_proj_bits else proj_bits
# 捕获共享专家引用使其常驻Qwen3-Next 有Qwen3-MoE 无)
layer.mlp = FileStreamingMoeBlock(
gate=mlp.gate, top_k=mlp.top_k, norm_topk_prob=mlp.norm_topk_prob,
store=store, layer_idx=i, hidden=hidden, moe_inter=moe_inter,
group_size=group_size, bits=bits, proj_bits=pb,
shared_expert=getattr(mlp, "shared_expert", None),
shared_expert_gate=getattr(mlp, "shared_expert_gate", None),
)
# 同理:不进 dict,避免 mlp→model→...→mlp 成环。供 native-fused-prefetch 取下层 gate。
object.__setattr__(layer.mlp, "_prefetch_model_ref", model)
patched += 1
if config.cross_layer_prefetch() or getattr(store, "_staging", None) is not None:
enable_cross_layer_prefetch()
return patched
def patch_model(model, store_factory=None):
"""把模型里每个 MoE 块(含 switch_mlp 与 gate 的块)替换成 StreamingMoeBlock。
store_factory(layer_idx)->LruExpertStore可为 None仅做 uniq 切片不接磁盘后端
返回被替换的层数便于校验确实命中了 MoE
"""
patched = 0
for i, layer in enumerate(model.layers):
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
store = store_factory(i) if store_factory is not None else None
layer.mlp = StreamingMoeBlock(mlp, layer_idx=i, store=store)
patched += 1
return patched

View File

@ -0,0 +1,131 @@
"""热路径计时与诊断埋点(从 streaming_moe 抽出,集中管理)。
- PROF细粒度分段计时STREAM_PROF=1probe 读取
- WINDOW_PROF同层 submitpromote attention/GDN 窗口时间WINDOW_PROF=1
- PREDICT_RECALL_PROF运行时预测对实际路由的覆盖率分诊预取覆盖问题PREDICT_RECALL_PROF=1
这些 dict 为共享可变对象各模块 import 后原地累加外部run_mtp_spec读同一对象
"""
import time
import mlx.core as mx
from mlx_streaming import config
_PROF_ON = config.stream_prof()
PROF = {"route": 0.0, "pyremap": 0.0, "fetch": 0.0, "matmul": 0.0, "combine": 0.0,
"n_calls": 0}
WINDOW_PROF = {"sum_s": 0.0, "n": 0}
PREDICT_RECALL_PROF = {"hit": 0, "routed": 0, "n": 0}
# 主动预取路径 host 墙钟探针PREFETCH_TPROF=1量主线程"不可与 GPU 重叠"的 CPU 时间。
# 注意MLX 惰性——gate matmul / pool scatter 的 GPU 计算落在前向末尾统一 eval本探针看不到
# 那部分用消融吞吐口径(NATIVE_NO_SUBMIT/NATIVE_NO_PROMOTE 差值)测。本探针专测 host 段:
# - predict_sblock._native_fused_prefetch 里建预测图(gate(x)+argpartition[+predicted_set tolist])
# - submit_s stg.submit(prefetch_into_staging 的 host 调用,注册 GPU 完成回调)
# - promote_spromote 总时长,再细分:
# · take_s prefetch_staging_takeC++ 锁读已就绪记录)
# · route_sroute_used_subset 的成员测试 .tolist()GPU→host 同步drain≤budget
# · place_s_place_expert 循环(建惰性切片 + 池 scatter 入图,不含其 GPU 执行)
# *_n 为各段调用次数(按"层·前向"计)。关时各调用点直接跳过,零开销。
PREFETCH_TPROF = {
"predict_s": 0.0, "predict_n": 0,
"submit_s": 0.0, "submit_n": 0,
"promote_s": 0.0, "promote_n": 0,
"take_s": 0.0, "route_s": 0.0, "place_s": 0.0, "place_experts": 0,
}
TPROF_ON = config.prefetch_tprof()
def note_tprof(seg: str, dt: float, *, count_key: "str | None" = None, count: int = 1) -> None:
"""累加预取某段 host 墙钟。dt 为秒;可选 count_key 同步累加调用/专家计数。"""
PREFETCH_TPROF[seg] += dt
if count_key is not None:
PREFETCH_TPROF[count_key] += count
def tprof_reset() -> None:
for k in PREFETCH_TPROF:
PREFETCH_TPROF[k] = 0 if isinstance(PREFETCH_TPROF[k], int) else 0.0
# miss 归因MISS_ATTRIB=1promote 之后 / acquire 之前统计本层真实路由的命中构成):
# - resident_hitacquire 前已驻留LRU 历史 + 本次 promote 命中)
# - miss_A_predictedmiss 且"在预测集里"——预测对了但没进池budget 丢/时序没到/被驱逐)
# - miss_B_unpredictedmiss 且"不在预测集里"——预测器召回缺口(根本没预测到)
# dec_* 为 decode/verify 热路径专用桶seq 短,与 prefill 长 seq 分离):原 miss_attrib
# 仅埋在 host 分支、被 prefill 样本主导,无法反映 decode 热路径的真实 A/B 构成。
MISS_ATTRIB = {"routed": 0, "resident_hit": 0, "miss_A_predicted": 0,
"miss_B_unpredicted": 0, "n": 0,
"dec_routed": 0, "dec_resident_hit": 0, "dec_miss_A": 0,
"dec_miss_B": 0, "dec_n": 0,
# miss_A 细分decode 桶):时序=就绪集里没有它pread 没完成);
# 驱逐=就绪过pread 完成)但 acquire 前已不在池(被驱逐/没写进)。
"dec_miss_A_timing": 0, "dec_miss_A_evicted": 0}
def note_miss_attrib(uniq, predicted, resident, is_decode: bool, ready=None) -> None:
"""统计本层真实路由 uniq 的命中构成,分 prefill / decode 两桶累加。
uniq本层真实路由专家集合predicted本层预测集residentacquire 前驻留集
ready本层最近 promote take 到的"已就绪(pread 完成)"专家集用于细分 miss_A
host 路径与 GPU remap 路径共用确保两条路径口径一致
"""
pset = predicted or set()
rset = ready or set()
for e in uniq:
MISS_ATTRIB["routed"] += 1
if is_decode:
MISS_ATTRIB["dec_routed"] += 1
if e in resident:
MISS_ATTRIB["resident_hit"] += 1
if is_decode:
MISS_ATTRIB["dec_resident_hit"] += 1
elif e in pset:
MISS_ATTRIB["miss_A_predicted"] += 1
if is_decode:
MISS_ATTRIB["dec_miss_A"] += 1
# 就绪集里有它=pread 完成过却仍 miss → 驱逐;否则 pread 没完成 → 时序。
if e in rset:
MISS_ATTRIB["dec_miss_A_evicted"] += 1
else:
MISS_ATTRIB["dec_miss_A_timing"] += 1
else:
MISS_ATTRIB["miss_B_unpredicted"] += 1
if is_decode:
MISS_ATTRIB["dec_miss_B"] += 1
MISS_ATTRIB["n"] += 1
if is_decode:
MISS_ATTRIB["dec_n"] += 1
# 并集专家数探针(UNION_PROF=1):按本次前向的 seq 长度分桶,记每层路由专家"去重并集"大小。
# seq=K(MTP verify,如 K=3)的桶即"K 个 token 的专家并集",决定池 cap 下限;seq=1 为 decode、
# seq=chunk 为分块 prefill。值为 {seq: [sum_union, n_layer_calls]}。默认关、零开销。
UNION_PROF: "dict[int, list]" = {}
# 分层原始样本(UNION_PROF=1){seq: {layer_idx: [本层每次前向的并集大小...]}}。
# 仅聚合 [sum,n] 无法给出 cap 下限所需的 U_max/p99/分层分布,故额外留每层每次前向的原始值。
# 样本量极小(≤层数×token数),仅采集期占用,退出前汇总。
UNION_SAMPLES: "dict[int, dict[int, list]]" = {}
UNION_ON = config.union_prof()
def note_union(seq: int, union_count: int, layer_idx: int = -1) -> None:
seq = int(seq)
union_count = int(union_count)
e = UNION_PROF.setdefault(seq, [0, 0])
e[0] += union_count
e[1] += 1
if layer_idx >= 0:
UNION_SAMPLES.setdefault(seq, {}).setdefault(int(layer_idx), []).append(union_count)
def union_reset() -> None:
UNION_PROF.clear()
UNION_SAMPLES.clear()
def prof_reset():
for k in PROF:
PROF[k] = 0.0
def _tick(seg, t0):
mx.eval # noqa (占位,避免误用)
PROF[seg] += time.perf_counter() - t0

View File

@ -0,0 +1,50 @@
"""MoE 路由 trace 采集工具(仅 probe 使用)。
热路径默认完全关闭开启后会把每次 MoE block layer 和专家集合记录下来
用于离线替换策略模拟开启会触发路由 ids CPU 同步不用于正式性能测试
"""
import json
_enabled = False
_events: list[dict] = []
def enable() -> None:
global _enabled, _events
_enabled = True
_events = []
def disable() -> None:
global _enabled
_enabled = False
def record(layer: int, experts, miss=None, resident=None, resident_rank=None) -> None:
if not _enabled:
return
vals = [int(e) for e in experts]
rec = {"layer": int(layer), "experts": sorted(set(vals))}
if miss is not None:
rec["miss"] = sorted({int(e) for e in miss})
if resident is not None:
rec["resident"] = sorted({int(e) for e in resident})
if resident_rank is not None:
rec["resident_rank"] = [[int(e), float(v)] for e, v in resident_rank.items()]
_events.append(rec)
def dump_jsonl(path: str) -> int:
with open(path, "w") as f:
for rec in _events:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
return len(_events)
def size() -> int:
return len(_events)
def events() -> list[dict]:
return list(_events)

View File

@ -0,0 +1,252 @@
"""共享装配层:把"文件后端流式加载主模型 + 取 hidden + 贪心生成"等跨入口复用的逻辑
集中在这里,而不是藏在某个入口脚本( validate_mtp)
依赖关系:本模块位于 core(mem/expert_store/streaming_moe) mtp 之上,
是把它们粘合成可运行模型的装配层, cli/ 各入口与测试复用
环境变量:
MODEL 主模型路径(MLX 量化)
EXPERT_DIR 拆分/重量化后的 per-expert safetensors 目录
EXPERT_SLOTS 每层常驻池容量
EXPERT_POOL_PROFILE 每层池预算 JSON(无损省内存,可选)
HIDDEN_VARIANT pre_final_norm(默认)| post_final_norm(排错时切换)
"""
import json
import os
import mlx.core as mx
from mlx_lm import load
from mlx_lm.models.base import create_attention_mask, create_ssm_mask
from mlx_streaming import config
from mlx_streaming.core.cache.expert_store import FileExpertStore
from mlx_streaming.core.prefetch.patch import patch_model_filebacked
MODEL = config.model_path()
EXPERT_DIR = config.expert_dir()
EXPERT_SLOTS = config.expert_slots()
# pre_final_norm(默认)| post_final_norm(排错时切换)
HIDDEN_VARIANT = config.hidden_variant()
# 默认 profile 文件名:放在 EXPERT_DIR 下随专家目录一起走,常用跑法自动启用(无损省内存)
DEFAULT_PROFILE_NAME = "pool_profile.json"
def load_pool_profile(expert_dir: str) -> "dict[int, int] | None":
"""解析每层池预算 profile,返回 layer_caps 或 None。
优先级:环境变量 EXPERT_POOL_PROFILE 显式指定路径 > {expert_dir}/pool_profile.json 默认
EXPERT_POOL_PROFILE=none/0/off 显式关闭(回到 uniform capacity)
profile 无损:仅按各层真实工作集分配,命中率/输出/吞吐不变(caps 仍被 capacity 上限钳制)
"""
p = config.expert_pool_profile()
if p.lower() in ("none", "0", "off"):
return None
if not p: # 未显式指定 → 默认找专家目录下的 profile
cand = os.path.join(expert_dir, DEFAULT_PROFILE_NAME)
p = cand if os.path.exists(cand) else ""
if p and os.path.exists(p):
with open(p) as f:
caps = json.load(f).get("layer_caps", {})
return {int(k): int(v) for k, v in caps.items()}
return None
def build_streaming_model():
"""用文件后端流式 patch 加载主模型(32GB 机器装不下 41GB 非流式)。"""
model, tok = load(MODEL, lazy=True)
# 取首个 MoE 维度
dims = None
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
gp = mlp.switch_mlp.gate_proj
dims = {"hidden": gp.input_dims, "moe_inter": gp.output_dims,
"group_size": getattr(gp, "group_size", 64),
"bits": getattr(gp, "bits", 4)}
break
bits, group, proj_bits, layer_proj_bits = (
dims["bits"], dims["group_size"], None, None)
meta_path = os.path.join(EXPERT_DIR, "_split_meta.json")
if os.path.exists(meta_path):
with open(meta_path) as f:
meta = json.load(f)
ed = meta.get("dims", {})
bits = ed.get("bits", bits)
group = ed.get("group_size", group)
proj_bits = ed.get("proj_bits")
if "per_layer_proj_bits" in ed:
layer_proj_bits = {int(k): v for k, v in ed["per_layer_proj_bits"].items()}
# 每层池预算 profile(pool_footprint 产出):默认从 {EXPERT_DIR}/pool_profile.json 自动启用,
# 无损省内存(命中率/输出/吞吐不变,仅不再为低占用层预留满 capacity)。
layer_caps = load_pool_profile(EXPERT_DIR)
store = FileExpertStore(EXPERT_DIR, capacity=EXPERT_SLOTS, layer_caps=layer_caps)
if config.zerocopy_dual_source():
# 零拷贝双源双缓冲:常驻池换成侧区模式(预分配 cap+2*spec_slots 行、禁 grow复用原池 loader/cap/profile。
from mlx_streaming.core.cache.resident_pool import ResidentExpertPool
_old = store._resident
# 默认单缓冲(持久 LFU,一份工作集,省一半侧区内存=生产路径);仅显式 legacy(SIDEREGION_LFU=0)用双缓冲。
_spec_gens = 1 if config.sideregion_lfu() else 2
store._resident = ResidentExpertPool(
_old.capacity, loader=_old.loader, layer_caps=_old.layer_caps,
spec_slots=config.pool_spec_slots(),
spec_gens=_spec_gens)
if config.stream_blob_loader():
# blob 接入常驻池 miss-loader复用 GPU-remap 快路径,小 EXPERT_SLOTS 即低内存。
store._blob_loader = _make_blob_source(dims, group, bits)
# 主动预取native-fused-prefetch miss→hitopt-inNATIVE_FUSED_PREFETCH=1
# 经"promote 只写真实路由命中专家"修正后已是净正:易缓存基座上 +15.5% tok/s
# demand 11.86→13.70hit 0.731→0.851,读盘 45%;见 active-prefetch-turnaround-2026-06-17.md
# 默认关只因收益依赖场景(基座可缓存性/是否磁盘受限)且只影响速度不影响质量、多占少量 staging 内存,
# 故作 opt-in 而非默认路径,落地配方见上述报告 §6。
if config.native_fused_prefetch() and getattr(store, "_blob_loader", None) is not None:
try:
import mlx_streaming.native_moe_ext # noqa: F401 确认扩展已编译
from mlx_streaming.core.prefetch.native_staging import NativeStagingManager
_budget = (config.pool_spec_slots() if config.zerocopy_dual_source()
else config.stream_blob_bg_budget(default=16))
store._staging = NativeStagingManager(store._blob_loader, budget=_budget)
except Exception:
store._staging = None # 扩展不可用 → 关闭,不影响主路径
# 零拷贝双源不变量staging 侧区行数必须等于池 spec_slots否则 C++ 会越界写池(静默损坏)。
if config.zerocopy_dual_source() and getattr(store, "_staging", None) is not None:
assert store._staging.budget == store._resident.spec_slots, (
f"零拷贝双源要求 staging.budget({store._staging.budget}) "
f"== 池 spec_slots({store._resident.spec_slots})")
if config.stream_blob_bg():
# 后台预取池预填bg 在独立 stream 物化预测专家promote 写进池槽(需 CROSS_LAYER_PREFETCH=1
from mlx_streaming.core.prefetch.bg_prefetch import BackgroundExpertPrefetcher
src = _make_blob_source(dims, group, bits)
store._blob_loader = src
store._bg = BackgroundExpertPrefetcher(
src, window=config.stream_blob_window())
patch_model_filebacked(model, store, dims["hidden"], dims["moe_inter"],
group, bits, proj_bits=proj_bits,
layer_proj_bits=layer_proj_bits)
# 双源双缓冲:构造一个共享 VirtualPoolgen 跨层全局、每前向 +1挂到每个流式 MoE 块。
if config.zerocopy_dual_source() and getattr(store, "_staging", None) is not None:
from mlx_streaming.core.cache.virtual_pool import VirtualPool
from mlx_streaming.core.moe.block import FileStreamingMoeBlock
# 双源模式仍需 ahead 调度block._native_fused_prefetch 靠 target_for 选目标层,
# 不传调度参数会让 target_for 恒返回 0_num_layers=0→ 预取全跳过、侧区永远空。
_vpool = VirtualPool(store._resident, store._staging, config.pool_spec_slots(),
num_layers=len(model.layers),
cutoff=config.cross_layer_cutoff(),
ahead_lo=config.cross_layer_ahead_lo(),
ahead_hi=config.cross_layer_ahead_hi(),
store=store)
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if isinstance(mlp, FileStreamingMoeBlock):
mlp._vpool = _vpool
# 主动预取(非 zerocopy挂 per-layer ahead 调度器 vpoolcutoff让晚层预读更早发起。
if (config.native_fused_prefetch() and not config.zerocopy_dual_source()
and getattr(store, "_staging", None) is not None):
from mlx_streaming.core.cache.virtual_pool import VirtualPool
from mlx_streaming.core.moe.block import FileStreamingMoeBlock
_sched = VirtualPool(num_layers=len(model.layers),
cutoff=config.cross_layer_cutoff(),
ahead_lo=config.cross_layer_ahead_lo(),
ahead_hi=config.cross_layer_ahead_hi())
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if isinstance(mlp, FileStreamingMoeBlock):
mlp._vpool = _sched
if config.stream_blob():
_attach_blob_source(model, dims, group, bits)
# KV 量化(IsoQuant K4/V3 + SO(4) 旋转):仅作用于 12 个全注意力层,128k KV 3.0→~0.68 GiB。
if config.kv_quant():
from mlx_streaming.core.cache.kv_quant_patch import patch_kv_quant
patch_kv_quant(model,
group_size=config.kv_group_size(),
k_bits=config.kv_k_bits(),
v_bits=config.kv_v_bits(),
rotate=config.kv_rotate(),
seed=config.kv_rot_seed())
# AUTOPIN 预热钉死(池 + blob_loader + staging 就绪后、return 前):按历史路由热度把每层
# top-N 热专家预填进常驻池并 pin 住(不参与任何驱逐),消灭冷启动慢热;usage 缺失则跳过
# (会话中计数仍在累计,落盘供下次启动)。AUTOPIN=0 零行为。
if config.autopin():
from mlx_streaming.core.cache import autopin as _autopin
_ap = _autopin.warm_start_pin(store)
if _ap["pinned"]:
print(f"[AUTOPIN] 预热 pin: {_ap['layers']} 层共 {_ap['pinned']} 专家 "
f"(frac={config.autopin_budget_frac()}, 耗时 {_ap['seconds']}s)"
+ (f", 跳过 {_ap['skipped_layers']}" if _ap["skipped_layers"] else ""),
flush=True)
# 把 store 挂到 model 上,供只拿到 model 的调用方(server /api/stats 读专家池健康度)取用。
# 非 array/dict/list/tuple 走 object 属性,不进 nn.Module 参数树。
model._expert_store = store
return model, tok, store
def _make_blob_source(dims, group, bits):
from mlx_streaming.core.cache.blob_loader import BlobExpertSource
blob_dir = config.blob_dir() or os.path.join(EXPERT_DIR, "blobs")
workers = config.stream_blob_workers()
nocache = config.stream_blob_nocache(default="0")
num_experts = 512
idx_path = os.path.join(blob_dir, "blob_index.json")
if os.path.exists(idx_path):
with open(idx_path) as f:
num_experts = int(json.load(f).get("num_experts", num_experts))
return BlobExpertSource(blob_dir, dims["hidden"], dims["moe_inter"], group, bits,
num_experts=num_experts, workers=workers, nocache=nocache)
def _attach_blob_source(model, dims, group, bits):
"""STREAM_BLOB=1给每个流式 MoE 块注入共享 BlobExpertSource全流式低内存路径"""
from mlx_streaming.core.cache.blob_loader import BlobExpertSource
from mlx_streaming.core.moe.block import FileStreamingMoeBlock
blob_dir = config.blob_dir() or os.path.join(EXPERT_DIR, "blobs")
workers = config.stream_blob_workers()
nocache = config.stream_blob_nocache(default="1")
num_experts = 512
idx_path = os.path.join(blob_dir, "blob_index.json")
if os.path.exists(idx_path):
with open(idx_path) as f:
num_experts = int(json.load(f).get("num_experts", num_experts))
src = BlobExpertSource(blob_dir, dims["hidden"], dims["moe_inter"], group, bits,
num_experts=num_experts, workers=workers, nocache=nocache)
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if isinstance(mlp, FileStreamingMoeBlock):
mlp._blob = src
def capture_prenorm_hidden(model, input_ids: mx.array) -> mx.array:
"""跑主模型层循环但跳过最后的 model.norm,返回 last-layer hidden(norm 前)。
HIDDEN_VARIANT=post_final_norm 时返回 norm 之后(用于消歧排错)
"""
inner = model.model
h = inner.embed_tokens(input_ids)
layers = inner.layers
if not layers:
return h
cache = model.make_cache()
fa_idx = next((i for i, l in enumerate(layers) if not l.is_linear), 0)
ssm_idx = next((i for i, l in enumerate(layers) if l.is_linear), 0)
fa_mask = create_attention_mask(h, cache[fa_idx])
ssm_mask = create_ssm_mask(h, cache[ssm_idx])
for layer, c in zip(layers, cache):
mask = ssm_mask if layer.is_linear else fa_mask
h = layer(h, mask=mask, cache=c)
if HIDDEN_VARIANT == "post_final_norm":
h = inner.norm(h)
return h
def greedy(model, input_ids: mx.array, n: int) -> mx.array:
"""主模型贪心生成 n 个 token,返回拼接后的完整序列(用作自投机参考)。"""
cache = model.make_cache()
cur = input_ids
out = []
for _ in range(n):
logits = model(cur, cache=cache)
nxt = mx.argmax(logits[:, -1, :], axis=-1, keepdims=True)
out.append(nxt)
cur = nxt
mx.eval(nxt)
return mx.concatenate([input_ids] + out, axis=1)

View File

View File

View File

@ -0,0 +1,28 @@
"""最小树 A/B 评测的纯裁决逻辑:中位数 + go/no-go 判定(无副作用,可单测)。"""
def median(xs):
"""样本中位数(偶数个取中间两数均值)。xs 非空。"""
s = sorted(xs)
n = len(s)
mid = n // 2
if n % 2 == 1:
return s[mid]
return (s[mid - 1] + s[mid]) / 2
def verdict_from_delta(delta, exact_all, margin=0.05):
"""按跨 prompt 中位相对提速 delta 与 lossless 门给出裁决。
- exact_all=False "bug"(lossless 硬门一票否决)
- delta > margin "go"
- delta < -margin "no-go"
- 其余(含边界) "even"
"""
if not exact_all:
return "bug"
if delta > margin:
return "go"
if delta < -margin:
return "no-go"
return "even"

View File

@ -0,0 +1,188 @@
"""真实 MTP drafter把 Qwen3NextMTP 包成 mtp_generate 需要的 draft/sync 接口。"""
import mlx.core as mx
from mlx_streaming.mtp.qwen3_next_mtp import mtp_step
from mlx_streaming.mtp.kv_cache import _snapshot, _restore
class MTPDrafter:
"""把 Qwen3NextMTP 包成 mtp_generate 需要的 drafter 接口。"""
def __init__(self, mtp, lm_head):
self.mtp = mtp
self.lm_head = lm_head
self.embed_tokens = mtp.embed_tokens
def make_cache(self):
from mlx_lm.models import cache as kc
return [kc.KVCache()] # MTP 单层全注意力
def draft(self, H_last, x_ids, mtp_cache, K, topk: int = 0):
# 注意:每草稿一次 int(argmax) host 同步反而最快 —— 它让 MLX 把每步 draft 的
# 图保持很小、及时释放;试过"全程 GPU argmax、末尾一次同步"攒大惰性图,draft
# 慢 5×、端到端 24.9→16.3(A/B 证伪,见报告)。保持逐步同步。
# topk>0(探针用):额外返回每个位置 MTP 的 top-k 候选 id(降序),量树形展开的救回上界。
drafts, cands = [], []
h, cur = H_last, x_ids
for _ in range(K):
logits, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg = logits[0].reshape(-1)
if topk > 0:
order = [int(i) for i in mx.argsort(lg)[-topk:].tolist()][::-1] # 降序 top-k
d = order[0]
cands.append(order)
else:
d = int(mx.argmax(lg))
drafts.append(d)
h, cur = mh, mx.array([[d]])
if topk > 0:
return drafts, cands
return drafts
def draft_adaptive(self, H_last, x_ids, mtp_cache, depth_max, tau):
"""置信度门控动态深度:逐位贪婪抽,累计置信度 C=∏p_i 跌破 tau 即停,最多 depth_max 位。
返回可变长度(1..depth_max)的草稿链C 用每位 top-1 softmax 概率连乘;低置信提前收敛
(本步只 verify 到当前深度,省后续位置的专家加载),高置信抽满到 depth_max始终至少抽 1
(depth=1 即退化为普通单 token 解码:verify 只喂 [x]恒接受模型真值 1 )
"""
drafts = []
h, cur = H_last, x_ids
conf = 1.0
for i in range(depth_max):
logits, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg = logits[0].reshape(-1)
d = int(mx.argmax(lg))
drafts.append(d)
h, cur = mh, mx.array([[d]])
if i + 1 < depth_max: # 末位无需再判是否加深
conf *= float(mx.softmax(lg)[d])
if conf < tau:
break
return drafts
def draft_adaptive_tree(self, H_last, x_ids, mtp_cache, depth_max, tau):
"""置信度门控变长 + pos0 分支:返回 (chainA, chainB)。
chainA draft_adaptive(累计置信度门控,变长 n=1..depth_max);chainB 是第1 top-2 分支,
续抽到与 chainA 同深度 n( n>=2 ,浅步 depth=1 无位置可救 chainB=None)用快照隔离
A/B 分叉,B 不带 A 的递归污染供合并路径:动态定深 + 对深链做 pos0 救回
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
top2 = [int(i) for i in mx.argsort(lg)[-2:].tolist()][::-1] # [d1a, d1b] 降序
d1a, d1b = top2[0], top2[1]
snap_after_x = _snapshot(mtp_cache)
# chainA:从 d1a 续抽,累计置信度跌破 tau 即停(与 draft_adaptive 同语义)。
chainA = [d1a]
conf = float(mx.softmax(lg)[d1a])
h, cur = mh1, mx.array([[d1a]])
while len(chainA) < depth_max and conf >= tau:
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg2 = lo[0].reshape(-1)
d = int(mx.argmax(lg2))
chainA.append(d)
conf *= float(mx.softmax(lg2)[d])
h, cur = mh, mx.array([[d]])
n = len(chainA)
chainB = None
if n >= 2: # 浅步无后续位置可救,不抽 B
_restore(mtp_cache, snap_after_x)
chainB = [d1b]
h, cur = mh1, mx.array([[d1b]])
for _ in range(n - 1):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chainB.append(d)
h, cur = mh, mx.array([[d]])
return chainA, chainB
def draft_tree(self, H_last, x_ids, mtp_cache, K, pos1=True):
"""最小树:第1(始终)、第2(pos1=True 时)草稿位置展开 top-2,返回三条链,各长 K。
- chainA = [d1a(top-1), d2a(top-1), d3a] top-1, draft
- chainB = [d1b(top-2), d2b, d3b] 第1 草稿位置 top-2 分支 pos0 救回
- chainC = [d1a, d2c(第2 位次选), d3c] None 第2 草稿位置 top-2 分支 pos1 救回
各分叉点用 mtp_cache 快照隔离,保证任一分支不带其它分支的递归污染:snap_after_x(x 处理后)
分叉 d1a/d1b;snap_after_d1a(d1a 处理后)分叉 d2a/d2c探针实测第2 首选错次选对比例
(~11%)高于第1 (~7%), pos1 分支值得抽;但每步多一次 MTP 续抽,pos1=False 时退化为
chainA/chainB(与旧最小树成本一致),便于 A/B 隔离 pos1 增量
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
top2 = [int(i) for i in mx.argsort(lg)[-2:].tolist()][::-1] # [d1a, d1b] 降序
d1a, d1b = top2[0], top2[1]
snap_after_x = _snapshot(mtp_cache) # x 处理后、第1 位分叉前的 MTP 递归态
def _continue(first, h0, n):
"""首 token first 已定,从隐状态 h0 再贪婪续抽 n 个,返回长度 1+n 的链。"""
chain = [first]
h, cur = h0, mx.array([[first]])
for _ in range(n):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chain.append(d)
h, cur = mh, mx.array([[d]])
return chain
chainC = None
if pos1:
# 从 d1a 续抽第2 位,捕获第2 位 top-2(d2a/d2c),供 chainA 与 chainC 分叉。
lo2, mh_after_d1a = mtp_step(self.mtp, mh1, mx.array([[d1a]]), self.lm_head, mtp_cache[0])
lg2 = lo2[0].reshape(-1)
top2_p1 = [int(i) for i in mx.argsort(lg2)[-2:].tolist()][::-1] # [d2a, d2c] 降序
d2a, d2c = top2_p1[0], top2_p1[1]
snap_after_d1a = _snapshot(mtp_cache) # d1a 处理后、第2 位分叉前的递归态
chainA = [d1a] + _continue(d2a, mh_after_d1a, K - 2) # [d1a, d2a, d3a...]
_restore(mtp_cache, snap_after_d1a) # 回到 d1a 态,C 链从第2 位次选分叉
chainC = [d1a] + _continue(d2c, mh_after_d1a, K - 2) # [d1a, d2c, d3c...]
else:
chainA = _continue(d1a, mh1, K - 1) # 仅全 top-1 链(旧最小树成本)
# chainB:回到 x 态,从第1 位次选 d1b 分叉续抽。
_restore(mtp_cache, snap_after_x)
chainB = _continue(d1b, mh1, K - 1) # [d1b, d2b, d3b...]
return chainA, chainB, chainC
def draft_paths(self, H_last, x_ids, mtp_cache, K, P):
"""完整树 batch-of-paths:位置1 展开 top-P,返回 P 条链(各长 K)。
位置1 MTP top-P 候选 [d1_0..d1_{P-1}];每个候选从 pos1 的共享 MTP 递归态 mh1 分叉,
贪婪续抽 K-1 token 成一条链 mtp_cache 快照保证各链从同一起点分叉互不污染
P=1 退化为普通链;P=2 等价 draft_tree
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
firsts = [int(i) for i in mx.argsort(lg)[-P:].tolist()][::-1] # top-P 降序
snap_pos1 = _snapshot(mtp_cache)
def _continue(first):
chain = [first]
h, cur = mh1, mx.array([[first]])
for _ in range(K - 1):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chain.append(d)
h, cur = mh, mx.array([[d]])
return chain
paths = []
for j, f in enumerate(firsts):
if j > 0:
_restore(mtp_cache, snap_pos1) # 回到 pos1 态,从同一起点分叉
paths.append(_continue(f))
return paths
def sync(self, prev_H, rH, replay_in, mtp_cache):
"""用已接受 token 的真实主模型 hidden 推进 MTP KV cache。
MTP 在位置 i 消费 (H_i, t_{i+1}) 预测 t_{i+2};因此提交 accepted prefix
`[t_{i+1}, ..., t_{i+n}]` ,hidden 序列应为 `[H_i, ..., H_{i+n-1}]`
"""
from mlx_streaming.mtp.qwen3_next_mtp import mtp_advance
h_seq = mx.concatenate([prev_H, rH[:, :-1, :]], axis=1)
H = mtp_advance(self.mtp, h_seq, replay_in, mtp_cache[0])
mx.eval(H)

View File

@ -0,0 +1,506 @@
"""Qwen3-Next MTP 自投机贪婪解码循环。
spec/plan: docs/superpowers/{specs,plans}/2026-06-07-qwen3next-mtp-self-speculation*
核心:每步主模型前向出 hidden -> MTP 自回归抽 K 草稿 -> 主模型并行验证 ->
接受最长命中前缀 -> cache 快照/恢复/重放回滚(统一处理 ArraysCache KVCache),
保证输出与非投机贪婪逐 token 等价
cache 校验/提交机制见 mtp/kv_cache.pydrafter 接口见 mtp/drafter.py
"""
import time
import mlx.core as mx
from mlx_lm.models.base import create_attention_mask, create_ssm_mask
from mlx_streaming import config
from mlx_streaming.mtp.kv_cache import (
enable_qwen3next_speculative_checkpoints, begin_speculative_checkpoints,
_snapshot, _restore, commit_verified_prefix, commit_verified_snapshot,
tile_caches, commit_tree_row)
def _batch_direct_commit_guaranteed(model) -> bool:
"""判断该模型的 batch verify 是否「必定」能直接提交(commit_verified_prefix 恒成功)。
commit 成功的充要条件是每个 cache 都满足:可裁剪(KVCache)或已捕获 per-token checkpoint
- 非线性层 KVCache,可裁剪,必成功;
- 线性层 必须是被 patch 过的 Qwen3NextGatedDeltaNet,verify 前向才会写 `_spec_checkpoints`
全部满足才返回 True此时每步的 `snap_m` 回退快照永远用不上,可安全省略以压低在途峰值
通用/玩具递归层( patch)返回 False,保留 snap_m 走安全 replay 兜底(与原行为逐 token 等价)
"""
import mlx_streaming.mtp.kv_cache as _kv
if not getattr(_kv, "_QWEN3NEXT_CHECKPOINTS_PATCHED", False):
return False
try:
from mlx_lm.models.qwen3_next import Qwen3NextGatedDeltaNet
layers = model.model.layers
except (ImportError, AttributeError):
return False
for l in layers:
if getattr(l, "is_linear", False) and not isinstance(
getattr(l, "linear_attn", None), Qwen3NextGatedDeltaNet):
return False
return True
def forward_with_hidden(model, ids, cache, compute_logits: bool = True):
"""跑主模型层循环 + 最终 norm,返回 (logits(1,L,V), hidden(1,L,H))。
logits 恒由 final-norm 后的 H lm_head 得到(保证与贪婪一致)
返回给 MTP drafter hidden:**默认用 final-norm 之前的 h**这与 MTP 训练/验证
(`capture_prenorm_hidden`)消费的输入一致; norm 后的 hidden 会双重归一化显著压低
草稿接受率(不影响正确性, verify 走主模型)`MTP_HIDDEN=post_norm` 可切回旧行为对照
compute_logits=False:跳过 lm_head(2048×151936 的大投影),只更新 cache 并返回 hidden
供分块 prefill 的中间块用(中间块不需要 logits,只需 cache 因果累积)
"""
inner = model.model
h = inner.embed_tokens(ids)
layers = inner.layers
has_full = any(not l.is_linear for l in layers)
has_linear = any(l.is_linear for l in layers)
fa_idx = next((i for i, l in enumerate(layers) if not l.is_linear), 0)
ssm_idx = next((i for i, l in enumerate(layers) if l.is_linear), 0)
fa_mask = create_attention_mask(h, cache[fa_idx]) if has_full else None
ssm_mask = create_ssm_mask(h, cache[ssm_idx]) if has_linear else None
for layer, c in zip(layers, cache):
mask = ssm_mask if layer.is_linear else fa_mask
h = layer(h, mask=mask, cache=c)
H = inner.norm(h)
hidden = H if config.mtp_hidden() == "post_norm" else h
logits = model.lm_head(H) if compute_logits else None
return logits, hidden
def prefill_chunked(model, ids, cache, chunk: "int | None" = None):
"""分块 prefill:把 ids 按 chunk 切片逐块喂入,增量更新 cache,返回最后一块的 (logits, hidden)。
整段 prefill 一次前向的激活峰值 prompt 长度( MoE 层瞬时物化大量唯一专家 + 长序列激活)
分块后每块只算 chunk token,逐块 eval 释放上一块瞬时图,峰值压回 chunk, decode 同稳态
数值等价:全注意力层逐块因果累积与整段完全等价;线性层(gated-delta)分块在 chunk 边界处有
重结合的末位浮点差异( decode 本身逐 token 同源,可忽略)中间块跳过 lm_head 省大投影
chunk<=0 prompt 不超过一块时,退化为整段 prefill
"""
if chunk is None:
chunk = config.prefill_chunk()
L = ids.shape[1]
if chunk <= 0 or L <= chunk:
return forward_with_hidden(model, ids, cache)
last_start = ((L - 1) // chunk) * chunk
for s in range(0, last_start, chunk):
_, h = forward_with_hidden(model, ids[:, s:s + chunk], cache, compute_logits=False)
mx.eval(h) # 物化本块、写 cache、释放上一块瞬时图
return forward_with_hidden(model, ids[:, last_start:], cache)
def forward_with_hidden_stepwise(model, ids, cache, capture_snapshots: bool = False):
"""逐 token 走主模型解码路径,可在每个 token 后保存 cache 快照。
这是 exact direct-commit 的保守路径:verify baseline 都走 N=1 cache 更新,
接受前缀时直接恢复到对应 token 后的快照,不再 replay accepted prefix
"""
logits_parts, hidden_parts, snaps = [], [], []
for i in range(ids.shape[1]):
logits, H = forward_with_hidden(model, ids[:, i:i + 1], cache)
mx.eval(logits, H)
logits_parts.append(logits)
hidden_parts.append(H)
if capture_snapshots and i < ids.shape[1] - 1:
snaps.append(_snapshot(cache))
logits = mx.concatenate(logits_parts, axis=1)
H = mx.concatenate(hidden_parts, axis=1)
if capture_snapshots:
return logits, H, snaps
return logits, H
# ----------------------------------------------------------------- 接受判定
def accept_prefix(drafts, preds):
"""drafts: MTP 抽的 K 个草稿; preds: 主模型对应位置真实下一 token(len==K)。
返回命中草稿数 matched(最长使 preds[j]==drafts[j] 的前缀长度)
发射 token 由调用方组装为 drafts[:matched] + [重放末位 bonus],统一处理纠正/全命中,
保证贪婪下逐 token 与非投机等价
"""
matched = 0
for d, p in zip(drafts, preds):
if int(p) == int(d):
matched += 1
else:
break
return matched
# --------------------------------------------------------- 完整树形验证(batch-of-paths)
def tree_verify_step(model, drafter, x, H_last, x_ids, mtp_cache, main_cache, K, P, snap_d):
"""一次 batched 前向并行验证 P 条候选路径,提交接受最长的赢家路径。
返回 (new_tokens, x_new, rH, accepted_in, matched)
原理:每条路径 `[x, d_1..d_{K-1}]` 是普通线性序列,拍到 batch 维后线性层/全注意力层都走成熟
batch 前向( row batch=1 等价, benchmarks/test_layerwise.py)batch=P 的计算量加宽了
每层预取窗口(hit_rate),多路径又提升接受长度(accept_len)验证后按赢家 row 提取 checkpoint
/trim 提交回 batch=1 cache,后续解码逐 token 从赢家路径续,与非投机贪婪等价
"""
paths = drafter.draft_paths(H_last, x_ids, mtp_cache, K, P)
verify_in = mx.array([[x] + p[: K - 1] for p in paths]) # (P, K)
snap_m = _snapshot(main_cache) # tile 前的 batch=1 快照(回退用)
tile_caches(main_cache, len(paths))
begin_speculative_checkpoints(main_cache)
vlogits, vH = forward_with_hidden(model, verify_in, main_cache)
mx.eval(vlogits, vH)
best_w, matched, preds = 0, -1, None
for wi, p in enumerate(paths):
preds_i = [int(t) for t in mx.argmax(vlogits[wi], axis=-1)]
m = accept_prefix(p, preds_i)
if m > matched: # 平局取更靠前的路径(top-1 优先)
best_w, matched, preds = wi, m, preds_i
drafts = paths[best_w]
accepted_len = min(matched + 1, K)
committed = commit_tree_row(main_cache, verified_len=K,
accepted_len=accepted_len, row=best_w)
if snap_d is not None:
_restore(mtp_cache, snap_d)
if not committed:
# 理论上 checkpoint 齐全不会走到;稳妥回退:恢复 batch=1 主 cache 后重放赢家 accepted prefix。
_restore(main_cache, snap_m)
accepted_in = mx.array([[x] + drafts[:matched]])
rlogits, rH = forward_with_hidden(model, accepted_in, main_cache)
mx.eval(rlogits)
bonus = int(mx.argmax(rlogits[:, -1, :]))
return drafts[:matched] + [bonus], bonus, rH, accepted_in, matched
rH = vH[best_w:best_w + 1, :accepted_len, :]
accepted_in = verify_in[best_w:best_w + 1, :accepted_len]
if matched == K:
return drafts[:K], drafts[-1], rH, accepted_in, matched
return drafts[:matched] + [preds[matched]], preds[matched], rH, accepted_in, matched
# ----------------------------------------------------------------- 主循环
def mtp_generate(model, drafter, tok, prompt, max_tokens, K=3, ids_mode=False,
profile=False, on_tokens=None, main_cache=None, cached_len=0):
"""贪婪 MTP 自投机。
drafter 需提供 draft(H_last(1,1,H), x_ids(1,1), mtp_cache, K) -> list[int](长度 K);
可选 make_cache()->list sync(rH, replay_in, mtp_cache)
ids_mode=True prompt 已是 ids(1,L) 且返回 token id 列表(测试用)
on_tokens:可选流式钩子,每步用本步真正写入 produced token id 列表回调一次
(prefill 的首 token 作为第一次回调);返回 True 表示请求尽快停止生成
跨轮复用(main_cache/cached_len):传入已持有前 cached_len prompt token main_cache,
则只 prefill `prompt[:, cached_len:]`(等价于一次多 token decode,复用历史 KV/递归态,不重算)
调用方须保证 prompt[:cached_len] cache 中已有 token 完全一致(否则结果错误)
不变式:返回时 main_cache 恰好持有 `prompt + produced[:-1]`(produced[-1] pending,未入 cache),
调用方可据此拼下轮的 cached prefixmain_cache=None 时内部新建(默认,整段 prefill)
"""
enable_qwen3next_speculative_checkpoints()
if main_cache is None:
main_cache = model.make_cache()
cached_len = 0
mtp_cache = drafter.make_cache() if hasattr(drafter, "make_cache") else None
ids = prompt if ids_mode else mx.array([tok.encode(prompt)])
if cached_len < 0 or cached_len >= ids.shape[1]:
raise ValueError(
f"cached_len={cached_len} 必须落在 [0, prompt_len={ids.shape[1]}) 内,至少留 1 个 token 供 prefill")
# prefill(分块):只喂尚未在 cache 中的后缀 ids[:, cached_len:](cached_len==0 即整段);
# 得到第 1 个 pending token x 与其 hidden;分块把激活峰值压到与 decode 同稳态。
logits, H = prefill_chunked(model, ids[:, cached_len:], main_cache)
x = int(mx.argmax(logits[:, -1, :]))
H_last = H[:, -1:, :]
produced = [x]
_stop_from_initial = on_tokens is not None and on_tokens([x])
mx.eval(H_last)
t0 = time.perf_counter()
n_steps = 0
# profile 分段累加(秒):把 sync/finalize 拆开,避免隐藏在 replay 名下。
t_draft = t_verify = t_replay = t_snap = 0.0
t_commit = t_sync = t_finalize = 0.0
direct_commits = fallback_replays = replayed_tokens = 0
# matched 直方图:accept_hist[j] = 恰好命中 j 个草稿(j∈0..K)的步数。用于实测每位置接受率,
# 不做几何分布假设。emitted/step = min(matched+1,K),故上限为 K(本实现 verify 只喂 K-1 草稿)。
accept_hist = [0] * (max(K, config.depth_max() if config.adaptive_depth() else K) + 1)
# top-k 覆盖探针:每位置统计模型真实 token 落在 MTP top-1/2/3 的步数(树形救回上界)。
topk_probe = config.accept_topk() if profile else 0
tk_n = [0] * K
tk_cover1 = [0] * K
tk_cover2 = [0] * K
tk_cover3 = [0] * K
verify_mode = config.mtp_verify_mode()
tree_mode = config.tree_top2()
tree_p1_mode = tree_mode and config.tree_top2_p1() # 第2 位(pos1)救回,附加在 tree_top2 上
tree_verify_mode = config.tree_verify() # 完整树形验证(batch-of-paths),优先级最高
# 置信度门控动态深度:仅在 plain 单链路径生效(与 tree/tree_verify/step 互斥)。
adaptive_mode = config.adaptive_depth() and not tree_mode and not tree_verify_mode \
and verify_mode != "step"
adaptive_rescue_mode = adaptive_mode and config.adaptive_rescue() # 深链叠加 pos0 救回
conf_tau_v = config.conf_tau()
depth_max_v = max(1, config.depth_max())
tree_P = max(1, config.tree_branches())
tree_rescues = 0 # 第1 位 top-2 成功救回(B 链首命中)的步数
tree_rescues_p1 = 0 # 第2 位 top-2 成功救回(C 链第2 token 命中)的步数
# 纯 batch 直接提交路径 + 模型保证 commit 恒成功时,跳过每步一次「全 cache 深拷贝 + eval」
# (snap_m 只用于 tree 救回 / step / replay 回退,这些路径都不满足下方条件)。省 ~72MiB 在途
# 峰值和一次同步栅栏,数值完全不变(bit-exact,只改内存调度)。模型结构恒定,循环外算一次即可。
_skip_snap = (verify_mode != "step") and (not tree_mode) and (not tree_verify_mode) \
and (not adaptive_rescue_mode) and _batch_direct_commit_guaranteed(model)
while len(produced) < max_tokens and not _stop_from_initial:
x0 = x
x_ids = mx.array([[x]])
snap_d = _snapshot(mtp_cache) if mtp_cache else None
prev_H_last = H_last
# 完整树形验证:一次 batched 前向验证 P 条路径,提交赢家。自带 step 尾处理后 continue。
if tree_verify_mode and verify_mode != "step":
if profile:
_tic = time.perf_counter()
new_tokens, x, rH, accepted_in, matched = tree_verify_step(
model, drafter, x, H_last, x_ids, mtp_cache, main_cache, K, tree_P, snap_d)
accept_hist[matched] += 1
direct_commits += 1
if profile:
t_verify += time.perf_counter() - _tic
_tic = time.perf_counter()
_stop = False
_n_before = len(produced)
for t in new_tokens:
produced.append(t)
if len(produced) >= max_tokens:
break
# 只把「本步真正写入 produced」的 token 交给回调:多 token 步命中 max_tokens
# 上限时,截断掉的尾巴不应上报给流式消费者(避免超报被丢弃的 token)。
if on_tokens is not None and on_tokens(produced[_n_before:]):
_stop = True
H_last = rH[:, -1:, :]
if mtp_cache is not None and hasattr(drafter, "sync"):
drafter.sync(prev_H_last, rH, accepted_in, mtp_cache)
if profile:
t_sync += time.perf_counter() - _tic
n_steps += 1
mx.eval(x, H_last)
if _stop:
break
continue
if profile:
_tic = time.perf_counter()
tree_b = None
tree_c = None
step_k = K # 本步验证深度(动态深度时可变)
if tree_mode and verify_mode != "step":
# chainA(全 top-1)、chainB(第1 位 top-2 → pos0 救回)、chainC(第2 位 top-2 → pos1 救回)
drafts, tree_b, tree_c = drafter.draft_tree(
H_last, x_ids, mtp_cache, K, pos1=tree_p1_mode)
draft_cands = None
elif adaptive_mode:
# 置信度门控:返回可变长度(1..depth_max)草稿链,本步深度随置信度自适应。
if adaptive_rescue_mode:
# 合并路径:变长 chainA + pos0 分支 chainB(深链才有),供下方 pos0 救回。
drafts, tree_b = drafter.draft_adaptive_tree(
H_last, x_ids, mtp_cache, depth_max_v, conf_tau_v)
else:
drafts = drafter.draft_adaptive(H_last, x_ids, mtp_cache, depth_max_v, conf_tau_v)
draft_cands = None
step_k = len(drafts)
elif topk_probe > 0:
drafts, draft_cands = drafter.draft(H_last, x_ids, mtp_cache, K, topk=topk_probe)
else:
drafts, draft_cands = drafter.draft(H_last, x_ids, mtp_cache, K), None # 长度 K
verify_in = mx.array([[x] + drafts[: step_k - 1]]) # [x, d_1..d_{step_k-1}]
if profile:
mx.eval(verify_in)
t_draft += time.perf_counter() - _tic
_tic = time.perf_counter()
snap_m = None if _skip_snap else _snapshot(main_cache)
if profile:
t_snap += time.perf_counter() - _tic
_tic = time.perf_counter()
verify_snaps = None
if verify_mode == "step":
vlogits, vH, verify_snaps = forward_with_hidden_stepwise(
model, verify_in, main_cache, capture_snapshots=True)
else:
begin_speculative_checkpoints(main_cache)
vlogits, vH = forward_with_hidden(model, verify_in, main_cache)
mx.eval(vlogits, vH)
preds = [int(t) for t in mx.argmax(vlogits[0], axis=-1)] # 长度 K
matched = accept_prefix(drafts, preds)
# 最小树救回:A 链首草稿被拒(matched==0)且 B 候选=模型真实 token(=preds[0])时,
# 恢复 cache、改验 B 链。preds[0] 只依赖 [x],两链一致;故 d1b==preds[0] 即 B 首命中。
# lossless(2026-07-05 定论,见 reports/tree-top2-rescue-2026-07-05.md):曾疑此救回非
# bit-lossless,实为 demand_dual 的一个 bug——inds(argpartition 切片)是非连续视图,C++
# 按连续读 data() 致 seq≥2 的 token≥1 装错专家,spec/救回全体失真。修 demand_dual 用
# contiguous() 物化后,seq=K verify 与 seq=1 解码逐位等价,spec 与本救回均 bit-lossless。
if tree_b is not None and matched == 0 and tree_b[0] == preds[0]:
_restore(main_cache, snap_m)
begin_speculative_checkpoints(main_cache)
# 重建 verify_in 为 B 链:后续 accepted_in=verify_in[:, :accepted_len] 才会取到 chainB,
# 否则残留 chainA 的错误首草稿会经 sync 污染 MTP cache、拉低后续草稿质量(低估救回收益)。
verify_in = mx.array([[x] + tree_b[:step_k - 1]])
vlogits, vH = forward_with_hidden(model, verify_in, main_cache)
mx.eval(vlogits, vH)
preds = [int(t) for t in mx.argmax(vlogits[0], axis=-1)]
drafts = tree_b
matched = accept_prefix(drafts, preds)
tree_rescues += 1
# 第2 位(pos1)救回:A 链首草稿命中(matched==1)但第2 草稿被拒,且 chainC 的第2 token
# =模型真值 preds[1] 时,改验 chainC=[x, d1a, d2c, ...]。preds[1] 只依赖 [x, d1a],A/C
# 前缀一致故 preds[0]/preds[1] 不变、重验后 matched≥2;d3c 可能续命中把接受长度顶到 3。
# 与 pos0 救回同构(纯重验、只提交模型 argmax token → bit-lossless,并集不放大)。
elif tree_c is not None and matched == 1 and len(tree_c) > 1 and tree_c[1] == preds[1]:
_restore(main_cache, snap_m)
begin_speculative_checkpoints(main_cache)
verify_in = mx.array([[x] + tree_c[:step_k - 1]])
vlogits, vH = forward_with_hidden(model, verify_in, main_cache)
mx.eval(vlogits, vH)
preds = [int(t) for t in mx.argmax(vlogits[0], axis=-1)]
drafts = tree_c
matched = accept_prefix(drafts, preds)
tree_rescues_p1 += 1
accept_hist[matched] += 1
if draft_cands is not None:
# preds[i]=模型在位置 i 的真实下一 token;查它是否在 MTP 该位置 top-1/2/3 候选里。
for i in range(K):
p = preds[i]
c = draft_cands[i]
tk_n[i] += 1
if p == c[0]:
tk_cover1[i] += 1
if p in c[:2]:
tk_cover2[i] += 1
if p in c[:3]:
tk_cover3[i] += 1
if profile:
t_verify += time.perf_counter() - _tic
_tic = time.perf_counter()
accepted_len = min(matched + 1, step_k)
if verify_snaps is not None:
# step 模式:逐 token 解码路径产生的快照可精确 direct commit。
committed = commit_verified_snapshot(main_cache, verify_snaps,
accepted_len, verified_len=K)
else:
# batch 模式:一次并行验证后按接受长度直接提交,零 replay。
# - 可裁剪 cacheKVCachetrim 掉 rejected 后缀即精确。
# - 递归状态 cacheQwen3-Next gated-delta 的 conv/ssmverify 前向里用
# multistate kernel 捕获的 per-token checkpoint 与 baseline kernel 逐 bit 等价
# (见 mtp/kv_cache.py 与 tests/test_gated_delta_multistate.py可精确直提交。
# commit_verified_prefix 一次处理混合 cacheKV 走 trim、Arrays 走 checkpoint 直提交。
# 若任一递归 cache 缺 checkpoint异常情况则返回 False自动回退到下方 replay。
committed = commit_verified_prefix(main_cache, verified_len=step_k,
accepted_len=accepted_len)
if snap_d is not None:
_restore(mtp_cache, snap_d)
accepted_in = verify_in[:, :accepted_len]
if profile:
t_commit += time.perf_counter() - _tic
_tic = time.perf_counter()
if committed:
direct_commits += 1
# 直接使用验证前向产生的 hidden 同步 MTP cache,避免主模型 replay。
rH = vH[:, :accepted_len, :]
if matched == step_k:
new_tokens = drafts[:step_k]
x = drafts[-1]
else:
new_tokens = drafts[:matched] + [preds[matched]]
x = preds[matched]
else:
fallback_replays += 1
if snap_m is None:
# 理论不可达:batch 直接提交路径 checkpoint 齐全,commit 必成功。若真走到这里,
# 说明前提被破坏(如未捕获 checkpoint),此时无快照可回退,直接抛错定位而非静默错算。
raise RuntimeError("batch verify commit failed but snapshot was skipped")
# fallback:回滚到验证前,重放 accepted prefix,保证 ArraysCache 正确。
_restore(main_cache, snap_m)
accepted_in = mx.array([[x0] + drafts[:matched]])
replayed_tokens += accepted_in.shape[1]
rlogits, rH = forward_with_hidden(model, accepted_in, main_cache)
mx.eval(rlogits)
bonus = int(mx.argmax(rlogits[:, -1, :]))
new_tokens = drafts[:matched] + [bonus]
x = bonus
if profile:
t_replay += time.perf_counter() - _tic
_tic = time.perf_counter()
_n_before = len(produced)
for t in new_tokens:
produced.append(t)
if len(produced) >= max_tokens:
break
# 只把「本步真正写入 produced」的 token 交给回调:多 token 步命中 max_tokens
# 上限时,截断掉的尾巴不应上报给流式消费者(避免超报被丢弃的 token)。
_stop = on_tokens is not None and on_tokens(produced[_n_before:])
H_last = rH[:, -1:, :]
if mtp_cache is not None and hasattr(drafter, "sync"):
drafter.sync(prev_H_last, rH, accepted_in, mtp_cache)
if profile:
t_sync += time.perf_counter() - _tic
_tic = time.perf_counter()
n_steps += 1
mx.eval(x, H_last)
if profile:
t_finalize += time.perf_counter() - _tic
if _stop:
break
produced = produced[:max_tokens]
wall = time.perf_counter() - t0
# cache 实际驻留 token 数(以可裁剪 cache 的 offset 为准),供跨轮复用精确对账:
# 末步多 token 跨 max_tokens 会 over-commit(cache 领先于 produced),据此识别并禁用复用。
resident_tokens = None
for c in main_cache:
if getattr(c, "is_trimmable", None) and c.is_trimmable() and hasattr(c, "offset"):
resident_tokens = int(c.offset)
break
stats = {
"steps": n_steps,
"tokens": len(produced),
"resident_tokens": resident_tokens,
"avg_accept_len": round(len(produced) / max(n_steps, 1), 3),
"wall_s": round(wall, 3),
"verify_mode": verify_mode,
"direct_commits": direct_commits,
"fallback_replays": fallback_replays,
"replayed_tokens": replayed_tokens,
"accept_hist": accept_hist, # [恰好命中0,1,...,K 个草稿的步数]
"tree_rescues": tree_rescues, # 最小树第1 位成功救回步数(tree_top2 开时)
"tree_rescues_p1": tree_rescues_p1, # 最小树第2 位成功救回步数(tree_top2 开时)
}
if topk_probe > 0:
stats["topk_probe"] = {
"topk": topk_probe, "n": tk_n,
"cover_top1": tk_cover1, "cover_top2": tk_cover2, "cover_top3": tk_cover3,
}
if profile:
seg = t_draft + t_snap + t_verify + t_commit + t_replay + t_sync + t_finalize
stats.update({
"t_draft_s": round(t_draft, 3),
"t_snap_s": round(t_snap, 3),
"t_verify_s": round(t_verify, 3),
"t_commit_s": round(t_commit, 3),
"t_replay_s": round(t_replay, 3),
"t_sync_s": round(t_sync, 3),
"t_finalize_s": round(t_finalize, 3),
# 「重放免费」上限:草稿+验证(不含快照,因 per-token 检查点设计无需全 cache 快照)
"proj_no_replay_tps": round(len(produced) / max(t_draft + t_verify, 1e-6), 2),
"measured_seg_tps": round(len(produced) / max(seg, 1e-6), 2),
})
if ids_mode:
return produced, stats
return tok.decode(produced), stats

View File

@ -0,0 +1,276 @@
"""MTP 投机解码的 KV/递归状态 cache 校验机制per-token checkpoint + 快照/恢复/提交。
speculative decoding 需要在验证 K 个草稿后只保留 accepted_len token cache
贡献两类 cache 处理方式不同
- 可裁剪 cacheKVCache直接 trim rejected 后缀
- 递归状态 cacheQwen3-Next 线性注意力 ArraysCache必须在 verify 前向中逐 token 记录
`[conv, ssm]` checkpoint 才能精确提交否则回退到快照 + replay accepted prefix
"""
import mlx.core as mx
import mlx.nn as nn
_QWEN3NEXT_CHECKPOINTS_PATCHED = False
_EMPTY_CACHE = object()
def enable_qwen3next_speculative_checkpoints():
"""给 Qwen3-Next 线性注意力层加 verify-time per-token cache checkpoint。
普通前向仍走 mlx-lm 原实现;只有 `begin_speculative_checkpoints()` 标记过的
ArraysCache 会走逐 token gated-delta ops,并把每个 prefix 后的 `[conv, ssm]`
状态写入 `cache._spec_checkpoints`
"""
global _QWEN3NEXT_CHECKPOINTS_PATCHED
if _QWEN3NEXT_CHECKPOINTS_PATCHED:
return
from mlx_lm.models.qwen3_next import Qwen3NextGatedDeltaNet
from mlx_streaming.core.linear_attn.gated_delta_multistate import (
gated_delta_update_multistate,
)
orig_call = Qwen3NextGatedDeltaNet.__call__
def patched_call(self, inputs, mask=None, cache=None):
capture = bool(getattr(cache, "_capture_spec_checkpoints", False))
if not capture:
return orig_call(self, inputs, mask=mask, cache=cache)
B, S, _ = inputs.shape
q, k, v, z, b, a = self.fix_query_key_value_ordering(
self.in_proj_qkvz(inputs), self.in_proj_ba(inputs)
)
if cache is not None and cache[0] is not None:
conv_state = cache[0]
else:
conv_state = mx.zeros(
(B, self.conv_kernel_size - 1, self.conv_dim),
dtype=inputs.dtype,
)
mixed_qkv = mx.concatenate(
[q.reshape(B, S, -1), k.reshape(B, S, -1), v.reshape(B, S, -1)], axis=-1
)
if mask is not None:
mixed_qkv = mx.where(mask[..., None], mixed_qkv, 0)
conv_input = mx.concatenate([conv_state, mixed_qkv], axis=1)
n_keep = self.conv_kernel_size - 1
conv_checkpoints = [
mx.contiguous(conv_input[:, i:i + n_keep, :])
for i in range(1, S + 1)
]
if cache is not None:
cache[0] = conv_checkpoints[-1]
conv_out = nn.silu(self.conv1d(conv_input))
q, k, v = [
t.reshape(B, S, h, d)
for t, h, d in zip(
mx.split(conv_out, [self.key_dim, 2 * self.key_dim], -1),
[self.num_k_heads, self.num_k_heads, self.num_v_heads],
[self.head_k_dim, self.head_k_dim, self.head_v_dim],
)
]
state = cache[1] if cache else None
inv_scale = k.shape[-1] ** -0.5
q = (inv_scale**2) * mx.fast.rms_norm(q, None, 1e-6)
k = inv_scale * mx.fast.rms_norm(k, None, 1e-6)
# 关键:用 multistate kernel 一次前向算出每个 token 处理后的 ssm 递归态。
# 它与 baseline 解码走的mlx gated_delta_kernel 是同一份 kernel、同序、fp32
# 因此 states_out[:, i] 与「逐 token 单步解码」逐 bit 等价(见
# tests/test_gated_delta_multistate.py——这是验证后能直接提交、零 replay 的根基。
# 注意kernel 路径内部处理 GQAhk_idx 映射),不在 Python 侧 repeat q/k
# 与 baseline kernel 路径对齐(旧 ops 路径的 repeat 会引入数值差异)。
out, state, states_out = gated_delta_update_multistate(
q, k, v, a, b, self.A_log, self.dt_bias, state, mask
)
if cache is not None:
cache[1] = state
cache.advance(S)
cache._spec_checkpoints = [
[conv_checkpoints[i], states_out[:, i]]
for i in range(S)
]
cache._capture_spec_checkpoints = False
out = self.norm(out, z)
return self.out_proj(out.reshape(B, S, -1))
Qwen3NextGatedDeltaNet.__call__ = patched_call
_QWEN3NEXT_CHECKPOINTS_PATCHED = True
def begin_speculative_checkpoints(caches):
"""标记 ArraysCache 在下一次 verify forward 中记录 per-token checkpoint。"""
for c in caches:
if not c.is_trimmable():
c._capture_spec_checkpoints = True
c._spec_checkpoints = None
# ----------------------------------------------------------------- cache 快照
def _copy_state(st):
"""深拷贝 cache.state(支持 None / array / list / tuple 嵌套)。"""
if st is None:
return None
if isinstance(st, (list, tuple)):
return type(st)(_copy_state(s) for s in st)
return mx.array(st)
def _iter_arrays(st):
if st is None:
return
if isinstance(st, (list, tuple)):
for s in st:
yield from _iter_arrays(s)
else:
yield st
def _snapshot(caches):
"""对每个 cache 深拷贝 state + meta_state(强制 eval,防 update_and_fetch 原地改写)。"""
snaps = []
arrays = []
for c in caches:
if hasattr(c, "keys") and hasattr(c, "empty") and c.empty():
snaps.append((_EMPTY_CACHE, c.meta_state))
continue
st_copy = _copy_state(c.state)
arrays.extend(_iter_arrays(st_copy))
snaps.append((st_copy, c.meta_state))
if arrays:
mx.eval(arrays)
return snaps
def _restore(caches, snaps):
for c, (st, meta) in zip(caches, snaps):
if st is _EMPTY_CACHE:
if hasattr(c, "keys") and hasattr(c, "values"):
c.keys = None
c.values = None
c.offset = 0
elif hasattr(c, "cache"):
c.cache = [None] * len(c.cache)
c.meta_state = meta
continue
# 必须装入副本:ArraysCache.state setter 是别名赋值(self.cache=v),后续前向
# 的 cache[idx]=new 会原地改写这个 list 的元素,从而污染快照本体。救回路径会对
# 同一 snap_m 恢复两次(救回时一次 + fallback 一次),若不复制,第二次恢复到的将是
# 被中间前向污染的状态 → 递归 cache 损坏、输出发散。
c.state = _copy_state(st)
c.meta_state = meta
def commit_verified_prefix(caches, verified_len: int, accepted_len: int) -> bool:
"""把验证前向产生的 cache 直接提交到 accepted prefix。
vLLM 语义下,验证 K token 后只保留 accepted_len token cache 的贡献
KVCache 这类可裁剪 cache,直接 trim rejected 后缀即可; ArraysCache 这类
递归状态 cache,必须有 per-token checkpoint 才能精确提交,否则返回 False 让调用方
fallback replay
"""
rejected = verified_len - accepted_len
if rejected < 0:
raise ValueError("accepted_len cannot exceed verified_len")
def can_commit(c):
if c.is_trimmable():
return True
checkpoints = getattr(c, "_spec_checkpoints", None)
return checkpoints is not None and len(checkpoints) >= accepted_len
if not all(can_commit(c) for c in caches):
return False
for c in caches:
if c.is_trimmable():
if rejected:
c.trim(rejected)
else:
c.state = c._spec_checkpoints[accepted_len - 1]
if hasattr(c, "_spec_checkpoints"):
c._spec_checkpoints = None
if hasattr(c, "_capture_spec_checkpoints"):
c._capture_spec_checkpoints = False
return True
# ----------------------------------------------------------- 树形验证 batch 工具
def _tile_state(st, P: int):
"""把 cache.state 沿 batch 轴(axis0)复制 P 份((1,...) → (P,...))。"""
if st is None:
return None
if isinstance(st, (list, tuple)):
return type(st)(_tile_state(s, P) for s in st)
return mx.contiguous(mx.repeat(st, P, axis=0))
def _row_state(st, w: int):
"""取 cache.state 的第 w 行((P,...) → (1,...)),保持 batch 维。"""
if st is None:
return None
if isinstance(st, (list, tuple)):
return type(st)(_row_state(s, w) for s in st)
return mx.contiguous(st[w:w + 1])
def tile_caches(caches, P: int):
"""把 batch=1 的 cache 全部平铺到 batch=P,供 batch-of-paths 树验证并行前向。
KVCache state setter 会按数组 shape 复位 offset;ArraysCache offset两者 meta
(lengths/left_padding) 在贪婪解码里恒为 None,tile 后语义不变前提:所有 cache 当前 batch=1
"""
for c in caches:
c.state = _tile_state(c.state, P)
def commit_tree_row(caches, verified_len: int, accepted_len: int, row: int) -> bool:
"""把 batched(batch=P)验证前向的第 `row` 条路径,按 accepted_len 提交回 batch=1 主 cache。
等价于先按接受长度裁剪 rejected 后缀,再抽出赢家路径那一行:
- 可裁剪 cache(KVCache):trim(rejected) 后取第 row KV
- 递归状态 cache(ArraysCache): verify 前向捕获的 per-token checkpoint[accepted_len-1]
的第 row [conv, ssm]要求 batched 前向前调过 begin_speculative_checkpoints
提交后主 cache 变回 batch=1,后续解码逐 token 从赢家路径续
"""
rejected = verified_len - accepted_len
if rejected < 0:
raise ValueError("accepted_len cannot exceed verified_len")
def can_commit(c):
if c.is_trimmable():
return True
cks = getattr(c, "_spec_checkpoints", None)
return cks is not None and len(cks) >= accepted_len
if not all(can_commit(c) for c in caches):
return False
for c in caches:
if c.is_trimmable():
if rejected:
c.trim(rejected)
c.state = _row_state(c.state, row)
else:
ck = c._spec_checkpoints[accepted_len - 1] # [conv(P,...), ssm(P,...)]
c.state = [_row_state(ck[0], row), _row_state(ck[1], row)]
if hasattr(c, "_spec_checkpoints"):
c._spec_checkpoints = None
if hasattr(c, "_capture_spec_checkpoints"):
c._capture_spec_checkpoints = False
return True
def commit_verified_snapshot(caches, snapshots, accepted_len: int,
verified_len: int | None = None) -> bool:
"""把 cache 恢复到 stepwise verify 的 accepted_len 后快照。"""
if verified_len is not None and accepted_len == verified_len:
return True
if accepted_len <= 0 or accepted_len > len(snapshots):
return False
_restore(caches, snapshots[accepted_len - 1])
return True

View File

@ -0,0 +1,95 @@
"""Qwen3-Next MTP(多 token 预测)模块的 MLX 实现,复用 mlx-lm 现成子模块类。
前向( vLLM/sglang/trtllm 一致):
emb = pre_fc_norm_embedding(embed(next_id))
hid = pre_fc_norm_hidden(主模型 last-layer hidden, norm 之前)
x = fc(concat([emb, hid], axis=-1)) # emb 在前
x = layer(x) # 全注意力 + MoE 解码层(内部含残差)
logits = lm_head(norm(x))
"""
from typing import Any, Callable, Optional
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_unflatten
from mlx_lm.models.base import create_attention_mask
from mlx_lm.models.qwen3_next import ModelArgs, Qwen3NextDecoderLayer
class Qwen3NextMTP(nn.Module):
def __init__(self, args: ModelArgs):
super().__init__()
h = args.hidden_size
eps = args.rms_norm_eps
self.embed_tokens = nn.Embedding(args.vocab_size, h)
self.pre_fc_norm_embedding = nn.RMSNorm(h, eps=eps)
self.pre_fc_norm_hidden = nn.RMSNorm(h, eps=eps)
self.fc = nn.Linear(2 * h, h, bias=False)
# layer_idx=3 配合 full_attention_interval=4 -> 全注意力 + MoE
self.layer = Qwen3NextDecoderLayer(args, layer_idx=3)
self.norm = nn.RMSNorm(h, eps=eps)
def __call__(
self,
hidden: mx.array, # (B, L, H) 主模型 last-layer hidden(norm 前)
next_ids: mx.array, # (B, L) 每位置的"下一个" token id
lm_head: Callable[[mx.array], mx.array],
cache: Optional[Any] = None,
return_hidden: bool = False,
):
emb = self.pre_fc_norm_embedding(self.embed_tokens(next_ids))
hid = self.pre_fc_norm_hidden(hidden)
x = self.fc(mx.concatenate([emb, hid], axis=-1))
mask = create_attention_mask(x, cache) if cache is not None else "causal"
x = self.layer(x, mask=mask, cache=cache)
H = self.norm(x)
logits = lm_head(H)
if return_hidden:
return logits, H
return logits
def mtp_step(mtp, hidden, token, lm_head, cache):
"""单步:hidden(1,1,H) + token(1,1) -> (logits(1,V), mtp_hidden(1,1,H))。"""
logits, H = mtp(hidden, token, lm_head, cache=cache, return_hidden=True)
return logits[:, -1, :], H[:, -1:, :]
def mtp_advance(mtp, hidden, token, cache):
"""只推进 MTP cache 并返回 hidden,不计算 lm_head logits。"""
emb = mtp.pre_fc_norm_embedding(mtp.embed_tokens(token))
hid = mtp.pre_fc_norm_hidden(hidden)
x = mtp.fc(mx.concatenate([emb, hid], axis=-1))
mask = create_attention_mask(x, cache) if cache is not None else "causal"
x = mtp.layer(x, mask=mask, cache=cache)
return mtp.norm(x)
def load_mtp(args: ModelArgs, weights_path: str, quantize: bool = True,
bits: int = 4) -> Qwen3NextMTP:
"""加载抽取好的 MTP 权重(已 stack 专家、已 norm +1.0)到模块。
quantize=True 时按主模型约定量化:线性层用 `bits`-bit(默认 4),gate/shared_expert_gate
恒用 8-bitquantize=False 保持 bf16 全精度草稿质量最高接受率最高,但显存更大
提高 bits(48) quantize=False 是直接抬升 MTP 草稿接受率的杠杆
"""
model = Qwen3NextMTP(args)
raw = mx.load(weights_path)
# 去掉 'mtp.' 前缀,映射到模块属性路径;mtp.layers.0.* -> layer.*
renamed = {}
for k, v in raw.items():
nk = k[len("mtp."):] if k.startswith("mtp.") else k
nk = nk.replace("layers.0.", "layer.", 1)
renamed[nk] = v
model.update(tree_unflatten(list(renamed.items())))
if quantize:
def pred(path, m):
if not hasattr(m, "to_quantized"):
return False
if path.endswith("mlp.gate") or path.endswith("shared_expert_gate"):
return {"group_size": 64, "bits": 8}
return True
nn.quantize(model, group_size=64, bits=bits, class_predicate=pred)
mx.eval(model.parameters())
return model

Binary file not shown.

View File

View File

@ -0,0 +1,32 @@
"""专家 blob 字节布局的单一真相源v1 affine / v2 mxfp4
Segment 五元组 (proj, tensor, np_dtype_name, shape, nbytes)与现有
blob_loader._layout 的元组结构一致 blob_loader 可直接消费pack 取子集
"""
from typing import List, Tuple
BLOB_V1_AFFINE = "expert_blob_v1"
BLOB_V2_MXFP4 = "expert_blob_v2_mxfp4"
Segment = Tuple[str, str, str, Tuple[int, int], int] # proj, tensor, dtype, shape, nbytes
def _projs(hidden: int, inter: int) -> Tuple[Tuple[str, int, int], ...]:
return (("gate_proj", inter, hidden), ("up_proj", inter, hidden), ("down_proj", hidden, inter))
def layout_for(fmt: str, hidden: int, inter: int, bits: int, group: int) -> Tuple[List[Segment], int]:
segs: List[Segment] = []
for proj, out_d, in_d in _projs(hidden, inter):
words = in_d * bits // 32
groups = in_d // group
segs.append((proj, "weight", "uint32", (out_d, words), out_d * words * 4))
if fmt == BLOB_V1_AFFINE:
segs.append((proj, "scales", "uint16", (out_d, groups), out_d * groups * 2))
segs.append((proj, "biases", "uint16", (out_d, groups), out_d * groups * 2))
elif fmt == BLOB_V2_MXFP4:
segs.append((proj, "scales", "uint8", (out_d, groups), out_d * groups * 1))
else:
raise ValueError(f"未知 blob format: {fmt}")
stride = sum(s[-1] for s in segs)
return segs, stride

View File

@ -0,0 +1,35 @@
"""自动重试的断点续传下载:应对代理对大文件不稳定的情况。
反复调用 snapshot_downloadhuggingface_hub 会从 .incomplete 续传
每轮带超时直到全部文件就绪或达到最大轮数
"""
import os
import sys
import time
from huggingface_hub import snapshot_download
REPO = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
MAX_TRIES = int(os.environ.get("MAX_TRIES", "200"))
os.environ.setdefault("HF_HUB_DOWNLOAD_TIMEOUT", "30")
def main():
for i in range(1, MAX_TRIES + 1):
try:
p = snapshot_download(
REPO,
allow_patterns=["*.json", "*.txt", "*.safetensors", "tokenizer*", "*.model"],
max_workers=1,
)
print(f"DONE -> {p}", flush=True)
return
except Exception as e:
print(f"[try {i}] 失败,续传重试: {repr(e)[:160]}", flush=True)
time.sleep(3)
print("EXHAUSTED: 达到最大重试轮数仍未完成", flush=True)
sys.exit(1)
if __name__ == "__main__":
main()

View File

@ -0,0 +1,131 @@
"""下载 Qwen3-Next 原版末分片并抽取/整理 MTP 权重为单文件。
MTP 全部权重在 model-00041-of-00041.safetensors(~3.9GB)整理规则:
1) mlp.experts.{e}.{proj}.weight(512 )stack mlp.switch_mlp.{proj}.weight
2) 对所有 RMSNorm 权重 +1.0(Qwen3-Next zero-centered RMSNorm,
mlx-lm sanitize 对主模型的处理一致;MTP sanitize 过滤,故需自行补)
"""
import json
import os
import subprocess
import mlx.core as mx
# 与 mlx-lm qwen3_next.sanitize 一致的 norm 后缀(去掉 model.norm,因 MTP 用 mtp.norm)
_NORM_SUFFIXES = (
".input_layernorm.weight",
".post_attention_layernorm.weight",
".q_norm.weight",
".k_norm.weight",
)
def _is_mtp_norm(key: str) -> bool:
if key.endswith(".pre_fc_norm_hidden.weight") or key.endswith(
".pre_fc_norm_embedding.weight"
):
return True
if key == "mtp.norm.weight":
return True
return any(key.endswith(sfx) for sfx in _NORM_SUFFIXES)
def bump_mtp_norms(weights: dict) -> dict:
out = {}
for k, v in weights.items():
if _is_mtp_norm(k) and v.ndim == 1:
out[k] = v + 1.0
else:
out[k] = v
return out
def stack_mtp_experts(weights: dict, num_experts: int) -> dict:
out = dict(weights)
prefix = "mtp.layers.0.mlp"
for proj in ("gate_proj", "up_proj", "down_proj"):
keys = [f"{prefix}.experts.{e}.{proj}.weight" for e in range(num_experts)]
if keys[0] not in out:
continue
stacked = mx.stack([out.pop(k) for k in keys])
out[f"{prefix}.switch_mlp.{proj}.weight"] = stacked
return out
SHARD = "model-00041-of-00041.safetensors"
REPO = "Qwen/Qwen3-Next-80B-A3B-Instruct"
from mlx_streaming import config as _cfg
SHARD_DIR = os.environ.get("MTP_SHARD_DIR", "/tmp/qn_mtp_shard") # 临时分片,可留 /tmp
OUT_PATH = _cfg.mtp_out() # 默认 models/qn_mtp_weights.safetensors持久
CONFIG = _cfg.qn_config() # 默认 models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit/config.json
_DL_URL = (
f"https://www.modelscope.cn/api/v1/models/{REPO}/repo"
f"?Revision=master&FilePath={SHARD}"
)
def _expected_size() -> int:
"""从 ModelScope 文件列表 API 取 SHARD 的真实字节数。
不能用 HEAD/content-length:API 重定向只回小 JSON,content-length 不可靠
"""
url = (
f"https://www.modelscope.cn/api/v1/models/{REPO}/repo/files"
f"?Revision=master&Recursive=true"
)
out = subprocess.run(
["curl", "-sL", "--max-time", "60", url],
capture_output=True, text=True,
).stdout
data = json.loads(out)
files = data.get("Data", {}).get("Files") or data.get("Data", {}).get("files") or []
for f in files:
name = f.get("Name") or f.get("Path") or f.get("name")
if name and SHARD in str(name):
return int(f.get("Size") or f.get("size") or 0)
return 0
def download_shard() -> str:
os.makedirs(SHARD_DIR, exist_ok=True)
out = os.path.join(SHARD_DIR, SHARD)
expect = _expected_size()
for attempt in range(1, 9):
cur = os.path.getsize(out) if os.path.exists(out) else 0
if expect and cur == expect:
print(f"SKIP shard already complete ({cur}B)")
return out
print(f"download attempt {attempt}: have {cur}B / {expect}B")
subprocess.run(
["curl", "-sL", "--max-time", "3600",
"--speed-limit", "51200", "--speed-time", "30",
_DL_URL, "-o", out],
)
cur = os.path.getsize(out) if os.path.exists(out) else 0
if expect and cur != expect:
raise RuntimeError(f"download failed: {cur}B != {expect}B")
return out
def extract(shard_path: str, config_path: str, out_path: str) -> None:
with open(config_path) as f:
cfg = json.load(f)
num_experts = cfg["num_experts"]
all_w = mx.load(shard_path)
mtp = {k: v for k, v in all_w.items() if k.startswith("mtp.")}
print(f"原始 mtp 张量数: {len(mtp)}")
mtp = stack_mtp_experts(mtp, num_experts)
mtp = bump_mtp_norms(mtp)
mx.eval(list(mtp.values()))
mx.save_safetensors(out_path, mtp)
print(f"已写出 {len(mtp)} 个张量 -> {out_path}")
def main():
shard = download_shard()
extract(shard, CONFIG, OUT_PATH)
if __name__ == "__main__":
main()

View File

@ -0,0 +1,87 @@
"""直接把 per-expert safetensors 打包成「每专家一个连续 blob」(按层一个文件)。
字节布局由 prep/blob_layout.py 统一描述v1 affine / v2 mxfp4 blob_loader 完全一致
_split_meta.json blob_format 选段表每专家依次写各 proj [weight, scales(, biases)]
- weight: uint32v1 scales/biases bf16 uint16 原始 2 字节v2 mxfp4 scales uint8
环境变量EXPERT_DIR( per-expert) / BLOB_DIR(输出) / BITS / GROUP / LAYERS(逗号或 all)
"""
import json
import os
import mlx.core as mx
import numpy as np
from mlx_streaming.prep.blob_layout import layout_for, BLOB_V1_AFFINE, BLOB_V2_MXFP4
EXPERT_DIR = os.environ.get("EXPERT_DIR", "/tmp/qwen3_next_experts_8bit_g128")
OUT = os.environ.get("BLOB_DIR", "/tmp/cb_8bit_blob")
BITS = int(os.environ.get("BITS", "8"))
GROUP = int(os.environ.get("GROUP", "128"))
def _meta():
m = json.load(open(os.path.join(EXPERT_DIR, "_split_meta.json")))
d = m["dims"]
fmt = m.get("blob_format", BLOB_V1_AFFINE)
return (int(d["num_experts"]), int(d["hidden"]), int(d["moe_intermediate"]),
fmt, d.get("quant_mode", "affine"))
def _raw_bytes(arr) -> bytes:
"""uint32 → 4B/元素uint8 → 1B/元素;其余(affine 的 bf16 scales/biases) → uint16 2B/元素。"""
if arr.dtype == mx.uint32:
return np.array(arr, copy=False).tobytes()
if arr.dtype == mx.uint8:
return np.array(arr, copy=False).tobytes()
return np.array(arr.view(mx.uint16), copy=False).tobytes()
def pack_layer(layer, num_experts, segs, stride, fmt, quant_mode) -> None:
out_path = os.path.join(OUT, f"layer{layer:02d}.blob")
with open(out_path, "wb") as f:
for e in range(num_experts):
w = mx.load(os.path.join(EXPERT_DIR, f"layer{layer:02d}_expert{e:03d}.safetensors"))
for proj, tensor, dt, shape, nb in segs:
b = _raw_bytes(w[f"{proj}.{tensor}"])
assert len(b) == nb, f"L{layer} e{e} {proj}.{tensor}: {len(b)}!={nb}"
f.write(b)
assert os.path.getsize(out_path) == stride * num_experts
index = {"format": fmt, "quant_mode": quant_mode, "layer": layer,
"num_experts": num_experts, "stride": stride,
"page_aligned": stride % 16384 == 0,
"segments": [{"proj": p, "tensor": t, "nbytes": n} for p, t, _, _, n in segs]}
with open(os.path.join(OUT, f"layer{layer:02d}.blob.index.json"), "w") as f:
json.dump(index, f, ensure_ascii=False, indent=2)
def _resolve_layers(spec: str, num_experts: int) -> list:
spec = spec.strip()
if spec.lower() == "all":
layers = set()
for name in os.listdir(EXPERT_DIR):
if name.startswith("layer") and name.endswith("_expert000.safetensors"):
layers.add(int(name[5:7]))
return sorted(layers)
return [int(x) for x in spec.split(",") if x.strip()]
def main():
os.makedirs(OUT, exist_ok=True)
num_experts, hidden, inter, fmt, quant_mode = _meta()
segs, stride = layout_for(fmt, hidden, inter, BITS, GROUP)
layers = _resolve_layers(os.environ.get("LAYERS", "all"), num_experts)
for i, L in enumerate(layers):
pack_layer(L, num_experts, segs, stride, fmt, quant_mode)
print(f" layer {L} packed ({i+1}/{len(layers)})", flush=True)
summary = {"format": fmt, "quant_mode": quant_mode, "stride": stride,
"page_aligned": stride % 16384 == 0, "num_experts": num_experts,
"layers": layers, "bits": BITS, "group_size": GROUP}
with open(os.path.join(OUT, "blob_index.json"), "w") as f:
json.dump(summary, f, ensure_ascii=False, indent=2)
print(json.dumps({"out": OUT, "n_layers": len(layers), "stride_bytes": stride,
"page_aligned": stride % 16384 == 0, "format": fmt}, ensure_ascii=False))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,83 @@
"""把单层单 projection 打包成连续 compute buffers。
输出三个大 buffer:
- layer43.gate_proj.weight.bin
- layer43.gate_proj.scales.bin
- layer43.gate_proj.biases.bin
这个格式用于验证整层少数大 buffer + custom qlinear kernel 直接按 expert_id 读取
"""
import argparse
import json
import os
from mlx_streaming.prep.pack_expert_ranges import read_safetensors_payloads
def pack_projection(src_dir: str, out_dir: str, layer: int, proj: str, num_experts: int) -> dict:
os.makedirs(out_dir, exist_ok=True)
tensors = ["weight", "scales", "biases"]
handles = {
name: open(os.path.join(out_dir, f"layer{layer:02d}.{proj}.{name}.bin"), "wb")
for name in tensors
}
meta = {
"format": "mlx_streaming_compute_buffer_v1",
"src_dir": src_dir,
"layer": layer,
"proj": proj,
"num_experts": num_experts,
"tensors": {},
}
try:
for expert in range(num_experts):
path = os.path.join(src_dir, f"layer{layer:02d}_expert{expert:03d}.safetensors")
payloads = {rec["key"]: rec for rec in read_safetensors_payloads(path)}
for name in tensors:
key = f"{proj}.{name}"
rec = payloads[key]
f = handles[name]
offset = f.tell()
f.write(rec["payload"])
entry = meta["tensors"].setdefault(name, {
"dtype": rec["dtype"],
"shape_per_expert": rec["shape"],
"nbytes_per_expert": len(rec["payload"]),
"file": f"layer{layer:02d}.{proj}.{name}.bin",
"offsets": [],
})
entry["offsets"].append(offset)
finally:
for f in handles.values():
f.close()
meta_path = os.path.join(out_dir, f"layer{layer:02d}.{proj}.index.json")
with open(meta_path, "w") as f:
json.dump(meta, f, ensure_ascii=False, indent=2)
return meta
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--src", default=os.environ.get("EXPERT_DIR", ""))
ap.add_argument("--out", default=os.environ.get("COMPUTE_BUFFER_DIR", ""))
ap.add_argument("--layer", type=int, default=int(os.environ.get("LAYER", "43")))
ap.add_argument("--proj", default=os.environ.get("PROJ", "gate_proj"))
args = ap.parse_args()
if not args.src:
raise SystemExit("--src / EXPERT_DIR required")
out_dir = args.out or os.path.join(args.src, "compute_buffers")
with open(os.path.join(args.src, "_split_meta.json")) as f:
model_meta = json.load(f)
num_experts = int(model_meta["dims"]["num_experts"])
meta = pack_projection(args.src, out_dir, args.layer, args.proj, num_experts)
print(json.dumps({
"out_dir": out_dir,
"layer": args.layer,
"proj": args.proj,
"num_experts": num_experts,
"files": {k: v["file"] for k, v in meta["tensors"].items()},
}, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,70 @@
"""把 per-expert safetensors 打包成 per-layer bundle。
输入目录:
layer00_expert000.safetensors
layer00_expert001.safetensors
...
输出目录:
layer_bundles/layer00.safetensors
bundle key:
expert000.gate_proj.weight
expert000.gate_proj.scales
...
运行时设置 EXPERT_BUNDLE=1 即可优先读取 bundle
"""
import json
import os
import time
import mlx.core as mx
SRC_DIR = os.environ.get("EXPERT_DIR", "/tmp/qwen3_next_experts")
OUT_DIR = os.environ.get("EXPERT_BUNDLE_DIR", os.path.join(SRC_DIR, "layer_bundles"))
def _pack_layer(src_dir: str, out_dir: str, layer: int, num_experts: int) -> str:
out = {}
for e in range(num_experts):
path = os.path.join(src_dir, f"layer{layer:02d}_expert{e:03d}.safetensors")
rec = mx.load(path)
prefix = f"expert{e:03d}."
for k, v in rec.items():
out[prefix + k] = v
os.makedirs(out_dir, exist_ok=True)
out_path = os.path.join(out_dir, f"layer{layer:02d}.safetensors")
mx.save_safetensors(out_path, out)
return out_path
def main():
t0 = time.perf_counter()
with open(os.path.join(SRC_DIR, "_split_meta.json")) as f:
meta = json.load(f)
layers = [int(x) for x in meta["moe_layers"]]
num_experts = int(meta["dims"]["num_experts"])
os.makedirs(OUT_DIR, exist_ok=True)
for i, layer in enumerate(layers):
out_path = _pack_layer(SRC_DIR, OUT_DIR, layer, num_experts)
print(f" {i + 1}/{len(layers)} layer{layer:02d} -> {out_path} "
f"({round(time.perf_counter() - t0, 1)}s)", flush=True)
bundle_meta = {
"src_dir": SRC_DIR,
"out_dir": OUT_DIR,
"layers": layers,
"num_experts": num_experts,
"format": "expert{EEE}.{tensor_key}",
}
with open(os.path.join(OUT_DIR, "_bundle_meta.json"), "w") as f:
json.dump(bundle_meta, f, ensure_ascii=False, indent=2)
print(json.dumps({
"out_dir": OUT_DIR,
"layers": len(layers),
"elapsed_s": round(time.perf_counter() - t0, 1),
}, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,135 @@
"""把 per-expert safetensors 的原始 payload 打包成 per-layer range pack。
这个格式只用于 de-risk避免运行时解析大量小 safetensors header先验证
expert byte range 读取是否比多个 `mx.load` 更快
"""
import argparse
import json
import os
import struct
import time
from pathlib import Path
ALIGN = 64
def _align(n: int, align: int = ALIGN) -> int:
rem = n % align
return n if rem == 0 else n + (align - rem)
def read_safetensors_payloads(path: str) -> list[dict]:
"""读取 safetensors header 和每个 tensor 的原始 payload bytes。"""
with open(path, "rb") as f:
header_len = struct.unpack("<Q", f.read(8))[0]
header = json.loads(f.read(header_len))
base = 8 + header_len
out = []
for key, meta in header.items():
if key == "__metadata__":
continue
start, end = meta["data_offsets"]
f.seek(base + start)
payload = f.read(end - start)
out.append({
"key": key,
"dtype": meta["dtype"],
"shape": meta["shape"],
"payload": payload,
})
return out
def pack_layer(src_dir: str, out_dir: str, layer: int, num_experts: int) -> dict:
"""打包单层所有专家,返回 index dict。"""
os.makedirs(out_dir, exist_ok=True)
pack_path = os.path.join(out_dir, f"layer{layer:02d}.pack")
index_path = os.path.join(out_dir, f"layer{layer:02d}.index.json")
tsv_path = os.path.join(out_dir, f"layer{layer:02d}.idx")
tensors = []
experts = []
with open(pack_path, "wb") as pack:
for expert in range(num_experts):
src = os.path.join(src_dir, f"layer{layer:02d}_expert{expert:03d}.safetensors")
expert_start = _align(pack.tell())
if expert_start > pack.tell():
pack.write(b"\0" * (expert_start - pack.tell()))
for rec in read_safetensors_payloads(src):
offset = _align(pack.tell())
if offset > pack.tell():
pack.write(b"\0" * (offset - pack.tell()))
payload = rec["payload"]
pack.write(payload)
tensors.append({
"layer": layer,
"expert_id": expert,
"key": rec["key"],
"dtype": rec["dtype"],
"shape": rec["shape"],
"offset": offset,
"nbytes": len(payload),
})
expert_end = pack.tell()
experts.append({
"layer": layer,
"expert_id": expert,
"offset": expert_start,
"nbytes": expert_end - expert_start,
})
index = {
"format": "mlx_streaming_expert_pack_v1",
"alignment": ALIGN,
"layer": layer,
"num_experts": num_experts,
"pack": os.path.basename(pack_path),
"experts": experts,
"tensors": tensors,
}
with open(index_path, "w") as f:
json.dump(index, f, ensure_ascii=False, indent=2)
with open(tsv_path, "w") as f:
f.write("kind\tlayer\texpert_id\tkey\tdtype\tshape\toffset\tnbytes\n")
for rec in experts:
f.write(
f"EXPERT\t{rec['layer']}\t{rec['expert_id']}\t*\t*\t*\t"
f"{rec['offset']}\t{rec['nbytes']}\n"
)
for rec in tensors:
shape = ",".join(str(x) for x in rec["shape"])
f.write(
f"TENSOR\t{rec['layer']}\t{rec['expert_id']}\t{rec['key']}\t"
f"{rec['dtype']}\t{shape}\t{rec['offset']}\t{rec['nbytes']}\n"
)
return index
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--src", default=os.environ.get("EXPERT_DIR", "/tmp/qwen3_next_experts"))
ap.add_argument("--out", default=os.environ.get("EXPERT_PACK_DIR"))
ap.add_argument("--layers", default=os.environ.get("LAYERS", ""))
args = ap.parse_args()
src = args.src
out = args.out or os.path.join(src, "layer_packs")
with open(os.path.join(src, "_split_meta.json")) as f:
meta = json.load(f)
all_layers = [int(x) for x in meta["moe_layers"]]
layers = [int(x) for x in args.layers.split(",") if x.strip()] if args.layers else all_layers
num_experts = int(meta["dims"]["num_experts"])
t0 = time.perf_counter()
for i, layer in enumerate(layers):
pack_layer(src, out, layer, num_experts)
print(f" {i + 1}/{len(layers)} layer{layer:02d} ({round(time.perf_counter() - t0, 1)}s)", flush=True)
with open(os.path.join(out, "_pack_meta.json"), "w") as f:
json.dump({
"src": src,
"out": out,
"layers": layers,
"num_experts": num_experts,
"format": "mlx_streaming_expert_pack_v1",
}, f, ensure_ascii=False, indent=2)
print(json.dumps({"out": out, "layers": len(layers)}, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,92 @@
"""把 compute buffer 重打包成"每专家一个连续 blob"格式(按层一个文件)。
blob 内顺序gate[w,s,b] + up[w,s,b] + down[w,s,b]正好 864KB= 54×16KB 页对齐
读一个专家 = 1 pread(stride, e*stride)而不是当前的 9 次散读
环境变量COMPUTE_BUFFER_DIR()BLOB_DIR(输出)LAYERS(逗号分隔默认 15,25)
"""
import json
import os
import numpy as np
HIDDEN = 2048
INTER = 512
GROUP = 128
BITS = 2
NUM_EXPERTS = 512
SRC = os.environ.get("COMPUTE_BUFFER_DIR", "/tmp/cb_2bit_g128")
OUT = os.environ.get("BLOB_DIR", "/tmp/cb_2bit_blob")
PROJS = (("gate_proj", INTER, HIDDEN), ("up_proj", INTER, HIDDEN), ("down_proj", HIDDEN, INTER))
def _layout():
"""返回 [(proj, tensor, nbytes_per_expert), ...] 和单专家总字节。"""
segs = []
for proj, out_dim, in_dim in PROJS:
words = in_dim * BITS // 32
groups = in_dim // GROUP
segs.append((proj, "weight", out_dim * words * 4))
segs.append((proj, "scales", out_dim * groups * 2))
segs.append((proj, "biases", out_dim * groups * 2))
return segs, sum(s[2] for s in segs)
def repack_layer(layer: int) -> dict:
os.makedirs(OUT, exist_ok=True)
segs, stride = _layout()
# 源 .bin 以 uint8 原始字节映射,按专家字节偏移切片
raw = {}
for proj, _, _ in PROJS:
base = os.path.join(SRC, f"layer{layer:02d}.{proj}")
for tensor in ("weight", "scales", "biases"):
raw[(proj, tensor)] = np.memmap(f"{base}.{tensor}.bin", dtype=np.uint8, mode="r")
out_path = os.path.join(OUT, f"layer{layer:02d}.blob")
with open(out_path, "wb") as f:
for e in range(NUM_EXPERTS):
for proj, tensor, nb in segs:
buf = raw[(proj, tensor)]
f.write(buf[e * nb:(e + 1) * nb].tobytes())
assert os.path.getsize(out_path) == stride * NUM_EXPERTS
index = {
"format": "expert_blob_v1",
"layer": layer,
"num_experts": NUM_EXPERTS,
"stride": stride,
"page_aligned": stride % 16384 == 0,
"segments": [{"proj": p, "tensor": t, "nbytes": n} for p, t, n in segs],
}
with open(os.path.join(OUT, f"layer{layer:02d}.blob.index.json"), "w") as f:
json.dump(index, f, ensure_ascii=False, indent=2)
return index
def _resolve_layers(spec: str) -> list[int]:
"""LAYERS=all → 读源 compute buffer 现有的全部层;否则按逗号解析。"""
spec = spec.strip()
if spec.lower() == "all":
layers = []
for name in os.listdir(SRC):
if name.endswith(".gate_proj.weight.bin") and name.startswith("layer"):
layers.append(int(name[len("layer"):len("layer") + 2]))
return sorted(set(layers))
return [int(x) for x in spec.split(",") if x.strip()]
def main():
layers = _resolve_layers(os.environ.get("LAYERS", "15,25"))
_, stride = _layout()
for L in layers:
repack_layer(L)
# 汇总 index供 loader 校验/发现
summary = {"format": "expert_blob_v1", "stride": stride,
"page_aligned": stride % 16384 == 0, "num_experts": NUM_EXPERTS,
"layers": layers}
with open(os.path.join(OUT, "blob_index.json"), "w") as f:
json.dump(summary, f, ensure_ascii=False, indent=2)
print(json.dumps({"out": OUT, "n_layers": len(layers), "stride_bytes": stride,
"page_aligned": stride % 16384 == 0}, ensure_ascii=False))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,69 @@
"""路线 B 离线工具:把模型里堆叠的 switch_mlp 专家权重按专家拆成 per-expert 小文件。
拆分后每个文件 layer{L:02d}_expert{E:03d}.safetensors 含扁平 dict
gate_proj.weight / gate_proj.scales / gate_proj.biases
up_proj.weight / up_proj.scales / up_proj.biases
down_proj.weight / down_proj.scales / down_proj.biases
非量化模型则只有 .weight可能还有 .bias
一次拆分逐专家物化全程低内存
"""
import os
import sys
import json
import mlx.core as mx
PROJ_NAMES = ["gate_proj", "up_proj", "down_proj"]
def split_switch_glu(switch_glu, out_dir: str, layer: int) -> int:
"""把一个 SwitchGLU 的三组 SwitchLinear 沿专家维拆成 per-expert 文件。返回专家数。"""
os.makedirs(out_dir, exist_ok=True)
E = switch_glu.gate_proj.num_experts
for e in range(E):
d = {}
for proj_name in PROJ_NAMES:
proj = getattr(switch_glu, proj_name)
for pname, p in proj.parameters().items():
if isinstance(p, mx.array) and p.ndim >= 1 and p.shape[0] == E:
d[f"{proj_name}.{pname}"] = p[e]
mx.eval(d) # 只物化这一个专家
path = os.path.join(out_dir, f"layer{layer:02d}_expert{e:03d}.safetensors")
mx.save_safetensors(path, d)
return E
def split_model(model_path: str, out_dir: str) -> dict:
"""加载模型lazy并把所有 MoE 层的专家拆到 out_dir。返回 {dims/统计}。"""
from mlx_lm import load
model, _ = load(model_path, lazy=True)
os.makedirs(out_dir, exist_ok=True)
moe_layers = []
dims = None
for l, layer in enumerate(model.layers):
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
sm = mlp.switch_mlp
split_switch_glu(sm, out_dir, l)
moe_layers.append(l)
if dims is None:
gp = sm.gate_proj
dims = {
"hidden": gp.input_dims,
"moe_intermediate": gp.output_dims,
"num_experts": gp.num_experts,
"group_size": getattr(gp, "group_size", None),
"bits": getattr(gp, "bits", None),
}
meta = {"out_dir": out_dir, "moe_layers": moe_layers, "dims": dims}
with open(os.path.join(out_dir, "_split_meta.json"), "w") as f:
json.dump(meta, f, ensure_ascii=False, indent=2)
return meta
if __name__ == "__main__":
mp = sys.argv[1]
od = sys.argv[2] if len(sys.argv) > 2 else "/tmp/mlx_qwen3_experts"
m = split_model(mp, od)
print(json.dumps(m, ensure_ascii=False, indent=2))

View File

@ -0,0 +1 @@
"""生产推理入口:基线/流式/投机解码等 run_* 启动脚本。"""

View File

@ -0,0 +1,42 @@
"""全量驻留基线:正常加载 + 生成,记录 RSS/峰值/decode 速度。"""
import os
import time
import json
import mlx.core as mx
from mlx_lm import load, generate
from mlx_streaming.core.mem import snapshot, reset_peak
MODEL = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "128"))
def main():
reset_peak()
t0 = time.perf_counter()
model, tok = load(MODEL) # 默认全量加载
mx.eval(model.parameters()) # 强制全部驻留,作为对照上界
load_done = snapshot()
t1 = time.perf_counter()
text = generate(model, tok, prompt=PROMPT, max_tokens=MAXTOK, verbose=False)
t2 = time.perf_counter()
after = snapshot()
out = {
"mode": "baseline_resident",
"model": MODEL,
"load_s": round(t1 - t0, 2),
"gen_s": round(t2 - t1, 2),
"tok_per_s": round(MAXTOK / (t2 - t1), 2),
"rss_gb_after_load": round(load_done.rss_bytes / 1e9, 2),
"rss_gb_after_gen": round(after.rss_bytes / 1e9, 2),
"mlx_peak_gb": round(after.mlx_peak_bytes / 1e9, 2),
}
print(json.dumps(out, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,69 @@
"""路线 Alazy+ 若支持则 mmap加载可选设常驻/缓存上限,量内存与速度。
环境探针结论mlx-lm 0.31.3load 支持 lazy但不支持 use_mmap
本脚本用 _filtered_load_kwargs 只传该版本真正支持的参数缺失的自动剔除
"""
import os
import time
import json
import inspect
import mlx.core as mx
from mlx_lm import load, generate
from mlx_streaming.core.mem import snapshot, reset_peak, clear_cache
MODEL = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "128"))
WIRED_GB = os.environ.get("WIRED_GB") # 例如 "6";空=不设
CACHE_GB = os.environ.get("CACHE_GB") # 例如 "1";空=不设
def _maybe(fn_name, val_bytes):
fn = getattr(mx, fn_name, None) or getattr(getattr(mx, "metal", object()), fn_name, None)
if fn and val_bytes is not None:
prev = fn(int(val_bytes))
print(f"{fn_name}({int(val_bytes)}) 旧值={prev}")
def _filtered_load_kwargs():
sig = inspect.signature(load)
want = {"lazy": True, "use_mmap": True}
return {k: v for k, v in want.items() if k in sig.parameters}
def main():
if WIRED_GB:
_maybe("set_wired_limit", float(WIRED_GB) * 1e9)
if CACHE_GB:
_maybe("set_cache_limit", float(CACHE_GB) * 1e9)
reset_peak()
kw = _filtered_load_kwargs()
print("load kwargs:", kw)
t0 = time.perf_counter()
model, tok = load(MODEL, **kw) # 不强制 eval 全部参数!
t1 = time.perf_counter()
after_load = snapshot() # 关键:此时 RSS 应远低于基线(专家还没换入)
text = generate(model, tok, prompt=PROMPT, max_tokens=MAXTOK, verbose=False)
t2 = time.perf_counter()
clear_cache()
after_gen = snapshot()
out = {
"mode": "mmap_lazy", "model": MODEL,
"wired_gb": WIRED_GB, "cache_gb": CACHE_GB,
"load_s": round(t1 - t0, 2), "gen_s": round(t2 - t1, 2),
"tok_per_s": round(MAXTOK / (t2 - t1), 2),
"rss_gb_after_load": round(after_load.rss_bytes / 1e9, 2),
"rss_gb_after_gen": round(after_gen.rss_bytes / 1e9, 2),
"mlx_active_gb_after_gen": round(after_gen.mlx_active_bytes / 1e9, 2),
"mlx_peak_gb": round(after_gen.mlx_peak_bytes / 1e9, 2),
}
print(json.dumps(out, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,409 @@
"""真实 80B MTP 自投机基准 + 与非投机贪婪逐 token 一致性校验。
环境变量:K / MAXTOK / PROMPT / QN_CONFIG / MTP_OUT(其余主模型路径见 validate_mtp)
"""
import json
import os
import statistics
import time
import mlx.core as mx
from mlx_lm.models.qwen3_next import ModelArgs
from mlx_streaming.core.moe import native_moe
from mlx_streaming.core.mem import snapshot, reset_peak
from mlx_streaming.mtp.drafter import MTPDrafter
from mlx_streaming.mtp.generate import forward_with_hidden, mtp_generate, prefill_chunked
from mlx_streaming.mtp.qwen3_next_mtp import load_mtp
from mlx_streaming.model_builder import build_streaming_model
from mlx_streaming import config as _cfg
QN_CONFIG = _cfg.qn_config()
MTP_OUT = _cfg.mtp_out()
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "96"))
K = int(os.environ.get("K", "3"))
PIN_HOT = int(os.environ.get("PIN_HOT", "0"))
PIN_CAL_TOK = int(os.environ.get("PIN_CAL_TOK", "32"))
# 稳态测速:warmup 跑满 MAXTOK(热 baseline+spec 两条路径的 Metal kernel 与常驻专家池/预取),
# 再各重复 REPEAT 次取中位数,避免冷启动编译/补池污染绝对 tok/s。WARMUP_TOK=0 关 warmup。
WARMUP_TOK = int(os.environ.get("WARMUP_TOK", str(MAXTOK)))
REPEAT = int(os.environ.get("REPEAT", "3"))
def _spec_once(model, drafter, tok, n, k):
ids, stats = mtp_generate(model, drafter, tok,
mx.array([tok.encode(PROMPT)]),
n, K=k, ids_mode=True, profile=True)
tps = round(stats["tokens"] / stats["wall_s"], 2)
return ids, stats, tps
def _baseline_greedy(model, tok, prompt, n):
cache = model.make_cache()
ids = mx.array([tok.encode(prompt)])
t0 = time.perf_counter()
# prefill 分块:把整段 prompt 的激活峰值压到与 decode 同稳态(见 config.prefill_chunk)。
logits, _ = prefill_chunked(model, ids, cache)
_dump_margin = bool(os.environ.get("DUMP_MARGIN"))
out = []
for _ in range(n):
lg = logits[:, -1, :]
nxt = int(mx.argmax(lg))
out.append(nxt)
if _dump_margin:
# 诊断:每步 top-2 logit 与差值,用于判定路径间发散是「FP 近平局」还是「真错」。
v = lg.reshape(-1)
top2 = mx.argpartition(-v, 2)[:2]
t2 = [int(x) for x in top2.tolist()]
vv = {t: float(v[t]) for t in t2}
st = sorted(t2, key=lambda t: -vv[t])
print(f"MARGIN step={len(out)-1} top1={st[0]}({vv[st[0]]:.5f}) "
f"top2={st[1]}({vv[st[1]]:.5f}) gap={vv[st[0]]-vv[st[1]]:.6f}", flush=True)
cur = mx.array([[nxt]])
mx.eval(cur)
if len(out) >= n:
break
logits, _ = forward_with_hidden(model, cur, cache)
return out, round(n / (time.perf_counter() - t0), 2)
def main():
reset_peak()
model, tok, store = build_streaming_model()
with open(QN_CONFIG) as f:
args = ModelArgs.from_dict(json.load(f))
mtp = load_mtp(args, MTP_OUT, quantize=True)
mtp.embed_tokens = model.model.embed_tokens # 共享主模型 embedding
drafter = MTPDrafter(mtp, model.lm_head)
# warmup:同时跑 baseline + spec 两条路径,编译 Metal kernel(含 multistate/batch verify)
# 并把专家常驻池/预取热起来,确保后续测的是稳态而非冷启动。默认 warmup=MAXTOK 跑满全长。
if WARMUP_TOK > 0:
_baseline_greedy(model, tok, PROMPT, WARMUP_TOK)
_spec_once(model, drafter, tok, WARMUP_TOK, K)
# ---- baseline 稳态:重复 REPEAT 次取中位数;最后一次清零统计供命中率口径 ----
base = None
base_tps_runs = []
for r in range(REPEAT):
if r == REPEAT - 1:
store.reset_stats()
base, bt = _baseline_greedy(model, tok, PROMPT, MAXTOK)
base_tps_runs.append(bt)
base_tps = statistics.median(base_tps_runs)
base_miss, base_hit = store.misses, store.hits
base_prefetch_loads = store._resident.prefetch_loads
base_prefetch_hits = store._resident.prefetch_hits
# 可选:baseline 之后、spec 之前校准每层热专家,并预取/钉入 resident pool。
# 这样 disk_load_ratio 的分母仍是未 pin baseline能直接衡量 pin 是否压低 spec miss。
if PIN_HOT > 0:
store.record = True
_baseline_greedy(model, tok, PROMPT, PIN_CAL_TOK)
for li in store.recorded_layers():
store.pin(li, store.hot(li, PIN_HOT))
store.record = False
store.reset_stats()
# 诊断:开 handler 触发时刻探针(仅覆盖 spec 阶段,enable 会清零旧日志)。
_hprof = bool(os.environ.get("STAGING_HPROF"))
if _hprof:
from mlx_streaming import native_moe_ext as _Nhp
_Nhp.staging_hprof_enable(True)
# ---- spec 稳态:重复 REPEAT 次取中位数;最后一次清零统计 + reset_peak + 保留其 stats ----
ids, stats, spec_tps_runs = None, None, []
for r in range(REPEAT):
if r == REPEAT - 1:
store.reset_stats()
reset_peak()
from mlx_streaming.core.profiling import tprof_reset, union_reset
tprof_reset() # 探针只统计最终测量轮(与命中率/内存口径一致)
union_reset() # 并集专家数也只统计最终轮
ids, stats, st = _spec_once(model, drafter, tok, MAXTOK, K)
spec_tps_runs.append(st)
spec_tps = statistics.median(spec_tps_runs)
spec_miss, spec_hit = store.misses, store.hits
spec_prefetch_loads = store._resident.prefetch_loads
spec_prefetch_hits = store._resident.prefetch_hits
after = snapshot()
proj = stats.get("proj_no_replay_tps", 0.0)
result = {
"K": K,
"max_tokens": MAXTOK,
"warmup_tok": WARMUP_TOK,
"repeat": REPEAT,
"exact_match": ids == base,
"n_mismatch": sum(1 for a, b in zip(ids, base) if a != b),
"avg_accept_len": stats["avg_accept_len"],
"steps": stats["steps"],
"verify_mode": stats.get("verify_mode"),
"direct_commits": stats.get("direct_commits"),
"fallback_replays": stats.get("fallback_replays"),
"replayed_tokens": stats.get("replayed_tokens"),
"spec_tok_per_s": spec_tps,
"baseline_tok_per_s": base_tps,
"speedup": round(spec_tps / max(base_tps, 1e-6), 2),
"spec_tps_runs": spec_tps_runs,
"baseline_tps_runs": base_tps_runs,
"spec_tps_minmax": [min(spec_tps_runs), max(spec_tps_runs)],
# 分段计时与「重放免费」投机上限
"t_draft_s": stats.get("t_draft_s"),
"t_snap_s": stats.get("t_snap_s"),
"t_verify_s": stats.get("t_verify_s"),
"t_commit_s": stats.get("t_commit_s"),
"t_replay_s": stats.get("t_replay_s"),
"t_sync_s": stats.get("t_sync_s"),
"t_finalize_s": stats.get("t_finalize_s"),
"proj_no_replay_tps": proj,
"proj_no_replay_speedup": round(proj / max(base_tps, 1e-6), 2),
"baseline_disk_loads": base_miss,
"baseline_prefetch_loads": base_prefetch_loads,
"baseline_prefetch_hits": base_prefetch_hits,
"baseline_hit_rate": round(base_hit / max(base_hit + base_miss, 1), 3),
"spec_disk_loads": spec_miss,
"spec_prefetch_loads": spec_prefetch_loads,
"spec_prefetch_hits": spec_prefetch_hits,
"spec_hit_rate": round(spec_hit / max(spec_hit + spec_miss, 1), 3),
"disk_load_ratio": round(spec_miss / max(base_miss, 1), 2),
# 双源 acquire 分路计数:n_miss==0 走全 GPU 快路径;任一路由 miss 则整层落 host 慢路径
# (.tolist 全批同步 + demand 读盘)。fallback 占比高 → 即使 hit 高,慢路径仍按"层"频繁触发。
"gpu_fastpath": getattr(store._resident, "gpu_fastpath", None),
"gpu_fallback": getattr(store._resident, "gpu_fallback", None),
"pin_hot": PIN_HOT,
"pin_cal_tok": PIN_CAL_TOK,
"pinned_experts": store.pinned_count(),
"expert_slots": store.capacity,
"mlx_active_gb": round(after.mlx_active_bytes / 1e9, 2),
"mlx_peak_gb": round(after.mlx_peak_bytes / 1e9, 2),
"rss_gb": round(after.rss_bytes / 1e9, 2),
# 内存分块(清缓冲后的真实常驻):A 权重 / B 专家池 / C staging / D MTP / F 激活
"mem_breakdown": _mem_breakdown(model, store, mtp),
"prefill_chunk": _cfg.prefill_chunk(),
"native_stage_cache": native_moe.stage_cache_stats(),
"bg_stats": (store._bg.stats() if getattr(store, "_bg", None) is not None else None),
"window_prof": _window_prof(),
"predict_recall": _predict_recall(),
"miss_attrib": _miss_attrib(),
"prefetch_tprof": _prefetch_tprof(stats.get("wall_s")),
"union_experts": _union_prof(),
}
# 噪声地板测量口径:DUMP_IDS=1 时把 baseline greedy 与 spec 的完整 token 序列打进日志,
# 供跨进程 run-to-run 逐位对比(默认关闭,不污染常规输出)。
if os.environ.get("DUMP_IDS"):
print("DUMP_BASE_IDS " + json.dumps(list(base)))
print("DUMP_SPEC_IDS " + json.dumps(list(ids)))
# 字节真值校验自证口径:开 STG_VERIFY 时,把两处校验器的累计计数
# (ok/bad/calls)打进日志。关键:让「0 BAD」可判真伪——若 calls==0 说明本配置根本
# 没触发该校验器(如 STG_VERIFY 在 zerocopy_dual 路径不接线),此时 0 BAD 是空结论。
if os.environ.get("STG_VERIFY"):
from mlx_streaming.core.cache import resident_pool as _rp_mod
from mlx_streaming.core.cache import virtual_pool as _vp_mod
_vsum = {
"STG_VERIFY.resident(verify_acquire_bytes)": dict(_rp_mod._stg_verify_state),
"STG_VERIFY.virtual(_verify_native_bytes)": dict(_vp_mod._stg_verify_state),
}
print("VERIFY_SUMMARY " + json.dumps(_vsum, ensure_ascii=False))
print(json.dumps(result, ensure_ascii=False, indent=2))
if _hprof:
# dump (gen, layer, t_fire) 原始日志,供离线分析回调触发时刻分布。
_flat = _Nhp.staging_hprof_get() # 扁平 [gen,layer,t, ...]
_path = os.environ.get("STAGING_HPROF_OUT", "/tmp/ab/hprof.jsonl")
os.makedirs(os.path.dirname(_path), exist_ok=True)
_n = 0
with open(_path, "w") as _f:
for i in range(0, len(_flat), 3):
_f.write(json.dumps([int(_flat[i]), int(_flat[i + 1]), float(_flat[i + 2])]) + "\n")
_n += 1
print(f"[hprof] wrote {_n} handler records to {_path}")
def _tree_nbytes(obj):
from mlx.utils import tree_flatten
return sum(v.nbytes for _, v in tree_flatten(obj) if isinstance(v, mx.array))
def _mem_breakdown(model, store, mtp):
"""把 decode 稳态内存拆成各块(清掉 MLX 可回收缓冲后量真实常驻)。
A 常驻非专家权重 / B 专家常驻池 / C staging 预取 / D MTP drafter(扣共享 embedding)/
F 激活+临时(= active底 ABCD)返回 GiB 字典
"""
GIB = 1024 ** 3
A = _tree_nbytes(model.parameters())
rp = store._resident
B = sum(v.nbytes for pool in rp._pools.values() for v in pool.values())
C = 0
stg = getattr(store, "_staging", None)
if stg is not None:
from mlx.utils import tree_flatten
C = sum(v.nbytes for _, v in tree_flatten(getattr(stg, "__dict__", {}))
if isinstance(v, mx.array))
# D:MTP drafter 全部权重减去与主模型共享的 embedding(run 里 mtp.embed_tokens = 主模型的)
D = _tree_nbytes(mtp.parameters())
emb = getattr(getattr(mtp, "embed_tokens", None), "weight", None)
shared = 0
if isinstance(emb, mx.array):
# 共享 embedding 在 A 已计入;若 mtp 的 embed 与主模型同一对象则从 D 扣除避免重复
shared = _tree_nbytes(mtp.embed_tokens) if hasattr(mtp, "embed_tokens") else 0
D = max(0, D - shared)
mx.clear_cache()
active = mx.get_active_memory()
peak = mx.get_peak_memory()
F = max(0, active - (A + B + C + D))
pool_rows = sum(rp.allocated_slots(l) for l in rp._pools)
return {
"A_weights_gib": round(A / GIB, 3),
"B_expert_pool_gib": round(B / GIB, 3),
"B_pool_rows": pool_rows,
"C_staging_gib": round(C / GIB, 3),
"D_mtp_drafter_gib": round(D / GIB, 3),
"F_activation_temp_gib": round(F / GIB, 3),
"active_live_gib": round(active / GIB, 3),
"peak_gib": round(peak / GIB, 3),
}
def _union_prof():
"""按 seq 分桶汇报每层路由专家并集大小(avg/max/p99/分层);seq=K 桶即 MTP verify 的专家并集。
UNION_PROF=1 才有数据返回 {seq: {avg_union, max_union, p99_union, per_layer...}};
verify = seq==K(verify_in=[x, d_1..d_{K-1}] K token, mtp/generate.py)
U_max 即该桶各层并集的全局最大值,是池 cap 下限的正确性依据(cap 必须 U_max 才保证不溢出)
"""
from mlx_streaming.core.profiling import UNION_PROF as U, UNION_SAMPLES as S
if not U:
return None
def _p(vals, q):
# 最近秩法分位数(vals 已升序):q∈[0,1]。样本少时直接给保守上界。
if not vals:
return None
idx = max(0, min(len(vals) - 1, int(round(q * (len(vals) - 1)))))
return vals[idx]
out = {}
for seq in sorted(U):
s, n = U[seq]
rec = {"avg_union": round(s / max(1, n), 2), "n_layer_calls": n}
# 由分层原始样本汇总出该 seq 桶的 U_max / p99cap 下限的正确性依据)。
layer_map = S.get(seq)
if layer_map:
flat = sorted(v for lst in layer_map.values() for v in lst)
rec["max_union"] = flat[-1]
rec["p99_union"] = _p(flat, 0.99)
rec["p50_union"] = _p(flat, 0.50)
rec["n_samples"] = len(flat)
# 分层分布:每层的 max / avg / 样本数(按层号排序)。
per_layer = {}
for li in sorted(layer_map):
lst = layer_map[li]
per_layer[li] = {
"max": max(lst),
"avg": round(sum(lst) / len(lst), 2),
"n": len(lst),
}
rec["per_layer"] = per_layer
rec["max_layer_idx"] = max(per_layer, key=lambda li: per_layer[li]["max"])
out[f"seq{seq}"] = rec
# 便捷:verify 桶 = seq==K(verify_in=[x, d_1..d_{K-1}] 恰 K 个 token)
verify_seq = K
verify = out.get(f"seq{verify_seq}")
return {"by_seq": out, "verify_seq": verify_seq,
"verify_avg_union": (verify or {}).get("avg_union"),
"verify_max_union": (verify or {}).get("max_union"),
"verify_p99_union": (verify or {}).get("p99_union")}
def _miss_attrib():
from mlx_streaming.core.profiling import MISS_ATTRIB as M
if not M["n"]:
return None
routed = max(1, M["routed"])
miss = M["miss_A_predicted"] + M["miss_B_unpredicted"]
out = {
"routed": M["routed"],
"hit_rate": round(M["resident_hit"] / routed, 4),
"miss_rate": round(miss / routed, 4),
# A预测到却没进池budget/时序/驱逐B没预测到召回缺口
"miss_A_predicted": M["miss_A_predicted"],
"miss_B_unpredicted": M["miss_B_unpredicted"],
"A_share_of_miss": round(M["miss_A_predicted"] / max(1, miss), 4),
"B_share_of_miss": round(M["miss_B_unpredicted"] / max(1, miss), 4),
}
if M["dec_n"]:
# decode/verify 热路径专用(与上面 prefill-主导的全量分开看):
dr = max(1, M["dec_routed"])
dmiss = M["dec_miss_A"] + M["dec_miss_B"]
out["decode"] = {
"routed": M["dec_routed"],
"hit_rate": round(M["dec_resident_hit"] / dr, 4),
"miss_rate": round(dmiss / dr, 4),
"miss_A_predicted": M["dec_miss_A"],
"miss_B_unpredicted": M["dec_miss_B"],
"A_share_of_miss": round(M["dec_miss_A"] / max(1, dmiss), 4),
"B_share_of_miss": round(M["dec_miss_B"] / max(1, dmiss), 4),
# miss_A 细分:时序(pread 没完成) vs 驱逐(就绪过但 acquire 前不在池)
"miss_A_timing": M["dec_miss_A_timing"],
"miss_A_evicted": M["dec_miss_A_evicted"],
"A_timing_share": round(M["dec_miss_A_timing"] / max(1, M["dec_miss_A"]), 4),
"A_evicted_share": round(M["dec_miss_A_evicted"] / max(1, M["dec_miss_A"]), 4),
}
return out
def _predict_recall():
from mlx_streaming.core.profiling import PREDICT_RECALL_PROF as P
if not P["n"]:
return None
return {"recall": round(P["hit"] / max(1, P["routed"]), 4),
"avg_routed": round(P["routed"] / P["n"], 2), "n": P["n"]}
def _window_prof():
from mlx_streaming.core.profiling import WINDOW_PROF
if not WINDOW_PROF["n"]:
return None
return {"avg_ms": round(WINDOW_PROF["sum_s"] / WINDOW_PROF["n"] * 1000, 4),
"n": WINDOW_PROF["n"]}
def _prefetch_tprof(wall_s=None):
"""预取 host 墙钟探针汇总(PREFETCH_TPROF=1 才有数据)。
汇报每段:总秒数占最终测量轮 wall 的百分比单次调用均值(ms)
注意:这是"主线程不可重叠的 host 时间";gate matmul / pool scatter GPU 执行落在
前向末尾统一 eval,不在此口径,需用消融(NATIVE_NO_SUBMIT/NATIVE_NO_PROMOTE)差值测
"""
from mlx_streaming.core.profiling import PREFETCH_TPROF as T
if not (T["predict_n"] or T["promote_n"] or T["submit_n"]):
return None
def _seg(s, n):
d = {"total_s": round(T[s], 4)}
if wall_s:
d["pct_wall"] = round(T[s] / wall_s * 100, 2)
if n:
d["avg_ms"] = round(T[s] / n * 1000, 4)
return d
host_total = T["predict_s"] + T["submit_s"] + T["promote_s"]
out = {
"predict": _seg("predict_s", T["predict_n"]),
"submit": _seg("submit_s", T["submit_n"]),
"promote": _seg("promote_s", T["promote_n"]),
# promote 内部细分(take 锁读 / route 成员同步 / place 切片+scatter 入图)
"promote_take": _seg("take_s", T["promote_n"]),
"promote_route": _seg("route_s", T["promote_n"]),
"promote_place": _seg("place_s", T["promote_n"]),
"host_total_s": round(host_total, 4),
"host_total_pct_wall": (round(host_total / wall_s * 100, 2) if wall_s else None),
"counts": {"predict_n": T["predict_n"], "submit_n": T["submit_n"],
"promote_n": T["promote_n"], "place_experts": T["place_experts"]},
}
return out
if __name__ == "__main__":
main()

View File

@ -0,0 +1,116 @@
"""投机解码 + 流式专家 验证target=流式 Qwen3-30Bdraft=常驻 Qwen3-0.6B。
独立 draft 投机 I/O 流式 MoE 上到底有没有加速对比不投机 vs 不同
num_draft_tokens decode tok/s 与草稿接受率
环境变量
MODEL/EXPERT_DIR/EXPERT_BITS/EXPERT_GROUP/EXPERT_SLOTS run_streaming
DRAFT draft 模型路径默认 /tmp/qwen3_draft
NDRAFTS 要扫的 num_draft_tokens 列表逗号分隔默认 "2,3,4"
MAXTOK/PROMPT
"""
import os
import time
import json
import mlx.core as mx
from mlx_lm import load, stream_generate
from mlx_streaming.core.mem import snapshot, reset_peak
from mlx_streaming.core.cache.expert_store import FileExpertStore
from mlx_streaming.core.prefetch.patch import patch_model_filebacked
MODEL = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
EXPERT_DIR = os.environ.get("EXPERT_DIR", "/tmp/mlx_qwen3_experts_2bit")
EXPERT_BITS = int(os.environ.get("EXPERT_BITS", "2"))
EXPERT_GROUP = int(os.environ.get("EXPERT_GROUP", "64"))
EXPERT_SLOTS = int(os.environ.get("EXPERT_SLOTS", "64"))
DRAFT = os.environ.get("DRAFT", "/tmp/qwen3_draft")
NDRAFTS = [int(x) for x in os.environ.get("NDRAFTS", "2,3,4").split(",")]
MAXTOK = int(os.environ.get("MAXTOK", "128"))
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
def _first_moe_dims(model):
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
gp = mlp.switch_mlp.gate_proj
return {"hidden": gp.input_dims, "moe_inter": gp.output_dims}
raise RuntimeError("无 MoE 层")
def _build_target():
model, tok = load(MODEL, lazy=True)
dims = _first_moe_dims(model)
# 混合精度:从专家目录 meta 读 proj_bits/bits/group有则优先于环境变量默认
bits, group, proj_bits, layer_proj_bits = EXPERT_BITS, EXPERT_GROUP, None, None
meta_path = os.path.join(EXPERT_DIR, "_split_meta.json")
if os.path.exists(meta_path):
with open(meta_path) as f:
ed = json.load(f).get("dims", {})
bits = ed.get("bits", bits)
group = ed.get("group_size", group)
proj_bits = ed.get("proj_bits")
if "per_layer_proj_bits" in ed:
layer_proj_bits = {int(k): v for k, v in ed["per_layer_proj_bits"].items()}
store = FileExpertStore(EXPERT_DIR, capacity=EXPERT_SLOTS)
patch_model_filebacked(model, store, dims["hidden"], dims["moe_inter"],
group, bits, proj_bits=proj_bits,
layer_proj_bits=layer_proj_bits)
return model, tok, store, {"bits": bits, "group": group, "proj_bits": proj_bits,
"layered": layer_proj_bits is not None}
def _run(model, tok, draft_model, num_draft, store):
store.reset_stats() if hasattr(store, "reset_stats") else None
n_tok = 0
n_draft_acc = 0
last_tps = 0.0
kwargs = {}
if draft_model is not None:
kwargs["num_draft_tokens"] = num_draft
t0 = time.perf_counter()
for r in stream_generate(model, tok, prompt=PROMPT, max_tokens=MAXTOK,
draft_model=draft_model, **kwargs):
n_tok = r.generation_tokens
last_tps = r.generation_tps
if getattr(r, "from_draft", False):
n_draft_acc += 1
wall = time.perf_counter() - t0
return {
"num_draft_tokens": num_draft if draft_model is not None else None,
"gen_tokens": n_tok,
"tok_per_s": round(last_tps, 2),
"wall_s": round(wall, 2),
"draft_accepted_frac": round(n_draft_acc / n_tok, 3) if n_tok else 0.0,
"expert_hit_rate": round(store.hit_rate(), 3),
}
def main():
reset_peak()
model, tok, store, qinfo = _build_target()
draft_model, _ = load(DRAFT) # 常驻小 draft
mx.eval(draft_model.parameters())
results = []
# 1) 不投机基线
results.append({"mode": "no_spec", **_run(model, tok, None, 0, store)})
# 2) 投机:扫 num_draft_tokens
for nd in NDRAFTS:
results.append({"mode": "spec", **_run(model, tok, draft_model, nd, store)})
after = snapshot()
print(json.dumps({
"target": MODEL, "draft": DRAFT,
"expert_bits": qinfo["bits"], "expert_proj_bits": qinfo["proj_bits"],
"expert_slots": EXPERT_SLOTS,
"rss_gb": round(after.rss_bytes / 1e9, 2),
"mlx_peak_gb": round(after.mlx_peak_bytes / 1e9, 2),
"results": results,
}, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,180 @@
"""路线 B 端到端lazy 加载 + 离线拆分专家 + 文件后端流式 + generate量内存/速度/命中率。
环境变量
MODEL 模型 repo 或本地路径默认 models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit
EXPERT_DIR per-expert 拆分输出目录默认 /tmp/mlx_qwen3_experts
EXPERT_SLOTS 每层 LRU 专家槽数默认 8=top_kworst-case 常驻槽数×MoE层数
PROMPT/MAXTOK 提示与生成 token
WIRED_GB 可选set_wired_limit 上限GB
CACHE_GB 可选set_cache_limit 上限GB控制 MLX 缓冲复用上界
"""
import os
import time
import json
import statistics
import mlx.core as mx
from mlx_lm import load, generate
from mlx_streaming import config
from mlx_streaming.core.mem import snapshot, reset_peak, clear_cache
from mlx_streaming.core.cache.expert_store import FileExpertStore
from mlx_streaming.prep.split_experts import split_model
from mlx_streaming.core.prefetch.patch import patch_model_filebacked
from mlx_streaming.model_builder import load_pool_profile
MODEL = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
EXPERT_DIR = os.environ.get("EXPERT_DIR", "/tmp/mlx_qwen3_experts")
EXPERT_SLOTS = int(os.environ.get("EXPERT_SLOTS", "8")) # 每层槽数
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "128"))
# 稳态测速:warmup 跑满 MAXTOK(把整段会路由到的专家都装进常驻池 + 编译 Metal kernel),
# 再重复 REPEAT 次取中位数,避免冷启动/补池污染绝对 tok/s。WARMUP_TOK=0 关 warmup。
WARMUP_TOK = int(os.environ.get("WARMUP_TOK", str(MAXTOK)))
REPEAT = int(os.environ.get("REPEAT", "3"))
WIRED_GB = os.environ.get("WIRED_GB")
CACHE_GB = os.environ.get("CACHE_GB")
CLEAR_ON_EVICT = os.environ.get("CLEAR_ON_EVICT", "0") == "1"
PIN_HOT = int(os.environ.get("PIN_HOT", "0")) # ③ 每层钉住的热专家数0=关)
CAL_TOK = int(os.environ.get("CAL_TOK", "32")) # 校准生成 token 数
# 专家量化格式覆盖:当 EXPERT_DIR 指向重量化目录(2/3-bit)时,文件后端 QSL 要用对应 bit/group
EXPERT_BITS = os.environ.get("EXPERT_BITS")
EXPERT_GROUP = os.environ.get("EXPERT_GROUP")
def _first_moe_dims(model):
for layer in model.layers:
mlp = getattr(layer, "mlp", None)
if mlp is not None and hasattr(mlp, "switch_mlp") and hasattr(mlp, "gate"):
gp = mlp.switch_mlp.gate_proj
return {
"hidden": gp.input_dims, "moe_inter": gp.output_dims,
"group_size": getattr(gp, "group_size", 64),
"bits": getattr(gp, "bits", 4),
}
raise RuntimeError("模型里没有找到 MoE 层")
def main():
if WIRED_GB:
fn = getattr(mx, "set_wired_limit", None)
if fn:
print("set_wired_limit 旧值=", fn(int(float(WIRED_GB) * 1e9)))
if CACHE_GB:
fn = getattr(mx, "set_cache_limit", None)
if fn:
print("set_cache_limit 旧值=", fn(int(float(CACHE_GB) * 1e9)))
# 1. 拆分专家到磁盘(若尚未拆分)
if not os.path.exists(os.path.join(EXPERT_DIR, "_split_meta.json")):
print("拆分专家到", EXPERT_DIR, "...")
t = time.perf_counter()
meta = split_model(MODEL, EXPERT_DIR)
print("拆分完成", round(time.perf_counter() - t, 1), "s; MoE 层=", len(meta["moe_layers"]))
reset_peak()
# 2. lazy 加载(不强制 eval 全部)
t0 = time.perf_counter()
model, tok = load(MODEL, lazy=True)
dims = _first_moe_dims(model)
# 专家目录 meta 是 bit/group/proj_bits 的权威来源(重量化目录都会写);先用它覆盖,
# 再让 EXPERT_BITS/EXPERT_GROUP 环境变量做最高优先级手动覆盖。proj_bits 非空走混合精度。
proj_bits = None
layer_proj_bits = None
meta_path = os.path.join(EXPERT_DIR, "_split_meta.json")
if os.path.exists(meta_path):
with open(meta_path) as f:
ed = json.load(f).get("dims", {})
dims["bits"] = ed.get("bits", dims["bits"])
dims["group_size"] = ed.get("group_size", dims["group_size"])
proj_bits = ed.get("proj_bits")
if "per_layer_proj_bits" in ed: # 逐层混合:键转回 int 层号
layer_proj_bits = {int(k): v for k, v in ed["per_layer_proj_bits"].items()}
if EXPERT_BITS:
dims["bits"] = int(EXPERT_BITS)
if EXPERT_GROUP:
dims["group_size"] = int(EXPERT_GROUP)
# 每层池预算 profile:默认从 {EXPERT_DIR}/pool_profile.json 自动启用(无损省内存)
# EXPERT_POOL_PROFILE 可显式指定路径或设 none 关闭。命中率/输出/吞吐不变。
layer_caps = load_pool_profile(EXPERT_DIR)
# 3. 文件后端 patch丢弃常驻堆叠 switch_mlp
store = FileExpertStore(EXPERT_DIR, capacity=EXPERT_SLOTS, layer_caps=layer_caps,
clear_on_evict=CLEAR_ON_EVICT, record=PIN_HOT > 0)
n = patch_model_filebacked(model, store, dims["hidden"], dims["moe_inter"],
dims["group_size"], dims["bits"],
proj_bits=proj_bits, layer_proj_bits=layer_proj_bits)
t1 = time.perf_counter()
after_patch = snapshot()
# 分块 prefill:把 mlx_lm.generate 内部 prefill 步长压到 config.prefill_chunk()(默认 2),
# 整段 prefill 的激活峰值 ∝prompt 长度 → ∝chunk,使 prefill 与 decode 同稳态。
# PREFILL_CHUNK=0 时不传,回退 mlx_lm 默认 2048。
_ps = config.prefill_chunk()
_gen_kw = {"prefill_step_size": _ps} if _ps > 0 else {}
# 3.5 ③ 热专家常驻:先校准跑一遍统计激活频率,钉住每层最热的 PIN_HOT 个专家
cal_s = 0.0
if PIN_HOT > 0:
tc = time.perf_counter()
generate(model, tok, prompt=PROMPT, max_tokens=CAL_TOK, verbose=False, **_gen_kw)
for li in store.recorded_layers():
store.pin(li, store.hot(li, PIN_HOT))
store.record = False
store.reset_stats()
cal_s = round(time.perf_counter() - tc, 2)
# 3.6 warmup:跑满 MAXTOK 把整段专家装进常驻池并编译 Metal kernel(PIN_HOT 已校准时
# 这里仍补满全长,确保第一次正式测量即稳态)。
if WARMUP_TOK > 0:
generate(model, tok, prompt=PROMPT, max_tokens=WARMUP_TOK, verbose=False, **_gen_kw)
# 4. 生成:重复 REPEAT 次取中位数;最后一次清零专家统计供命中率口径对应稳态。
reset_peak()
text = None
gen_runs = []
for r in range(REPEAT):
if r == REPEAT - 1:
store.reset_stats()
tg = time.perf_counter()
text = generate(model, tok, prompt=PROMPT, max_tokens=MAXTOK, verbose=False, **_gen_kw)
gen_runs.append(round(MAXTOK / (time.perf_counter() - tg), 2))
tok_per_s = statistics.median(gen_runs)
t2 = time.perf_counter()
clear_cache()
after_gen = snapshot()
out = {
"mode": "streaming_filebacked", "model": MODEL,
"expert_dir": EXPERT_DIR,
"expert_bits": dims["bits"], "expert_group": dims["group_size"],
"proj_bits": proj_bits, "layered": layer_proj_bits is not None,
"per_layer_slots": EXPERT_SLOTS, "pool_profile": bool(layer_caps),
"wired_gb": WIRED_GB, "cache_gb": CACHE_GB,
"clear_on_evict": CLEAR_ON_EVICT,
"pin_hot": PIN_HOT, "cal_s": cal_s,
"warmup_tok": WARMUP_TOK, "repeat": REPEAT,
"patched_moe_layers": n,
"resident_experts": store.resident_count(),
"pinned_experts": store.pinned_count(),
"load_patch_s": round(t1 - t0, 2),
"tok_per_s": tok_per_s,
"tok_per_s_runs": gen_runs,
"tok_per_s_minmax": [min(gen_runs), max(gen_runs)],
"rss_gb_after_patch": round(after_patch.rss_bytes / 1e9, 2),
"rss_gb_after_gen": round(after_gen.rss_bytes / 1e9, 2),
"mlx_active_gb_after_gen": round(after_gen.mlx_active_bytes / 1e9, 2),
"mlx_peak_gb": round(after_gen.mlx_peak_bytes / 1e9, 2),
"expert_hit_rate": round(store.hit_rate(), 4),
"expert_hits": store.hits, "expert_misses": store.misses,
# GPU remap 取证:整层全命中走 GPU 快路径 vs 有 miss 回退 host 的层调用占比
"gpu_fastpath": store._resident.gpu_fastpath,
"gpu_fallback": store._resident.gpu_fallback,
"gpu_fastpath_frac": round(
store._resident.gpu_fastpath
/ max(1, store._resident.gpu_fastpath + store._resident.gpu_fallback), 4),
"sample": text[:240],
}
print(json.dumps(out, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()

View File

@ -0,0 +1,9 @@
"""OpenAI 兼容的 FastAPI server。
引擎复用 TUI ChatBackend 抽象:SPARKLE_FAKE=1(或测试显式注入)时用 FakeBackend
免模型开发/测试;默认用 MLXBackend 装配真实流式 MoE 引擎
入口:python -m mlx_streaming.server
"""
from mlx_streaming.server.app import create_app
__all__ = ["create_app"]

View File

@ -0,0 +1,95 @@
"""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()

View File

@ -0,0 +1,116 @@
"""极简 admin 页:单文件 HTML + 内联 JS,不引框架/CDN。中文文案,2.5s 轮询 stats。"""
ADMIN_HTML = """<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>Sparkle 引擎控制台</title>
<style>
body { font-family: -apple-system, "PingFang SC", sans-serif; max-width: 640px;
margin: 32px auto; padding: 0 16px; color: #222; }
h1 { font-size: 20px; }
section { border: 1px solid #ddd; border-radius: 8px; padding: 16px;
margin-bottom: 16px; }
.row { margin: 8px 0; }
.label { display: inline-block; width: 180px; color: #555; }
input[type=number] { width: 90px; }
input[type=range] { width: 240px; vertical-align: middle; }
button { padding: 6px 20px; }
#msg { margin-left: 12px; color: #060; }
.hint { color: #888; font-size: 13px; }
</style>
</head>
<body>
<h1>Sparkle 引擎控制台</h1>
<section>
<div class="row"><span class="label">引擎状态</span><span id="loaded">-</span></div>
<div class="row"><span class="label">模型</span><span id="model">-</span></div>
<div class="row"><span class="label">最近一轮吞吐</span><span id="tps">-</span> tok/s</div>
<div class="row"><span class="label">进程内存 RSS</span><span id="rss">-</span></div>
</section>
<section>
<div class="row">
<span class="label">常驻专家池 expert_slots</span>
<input type="range" id="slots" min="32" max="160" step="1" value="32">
<b id="slotsVal">32</b>
</div>
<div class="row hint">预估池内存:<span id="poolMem"></span>(slots × ~160MB(g64));修改后引擎后台重建,期间对话不可用</div>
<div class="row">
<span class="label">投机宽度 k</span>
<input type="number" id="k" min="1" max="8" step="1" value="3">
</div>
<div class="row">
<span class="label">最大生成 max_tokens</span>
<input type="number" id="maxTokens" min="1" max="32768" step="1" value="4096">
</div>
<div class="row">
<button id="apply">应用</button><span id="msg"></span>
</div>
</section>
<script>
const $ = id => document.getElementById(id);
function fmtBytes(b) {
if (b >= 1024 * 1024 * 1024) return (b / 1024 / 1024 / 1024).toFixed(2) + " GB";
return (b / 1024 / 1024).toFixed(0) + " MB";
}
function updateSlots() {
const v = +$("slots").value;
$("slotsVal").textContent = v;
$("poolMem").textContent = " " + fmtBytes(v * 160 * 1024 * 1024) + " ";
}
$("slots").addEventListener("input", updateSlots);
async function refresh() {
try {
const s = await (await fetch("/api/stats")).json();
$("loaded").textContent = s.loaded ? "已加载" : "未加载 / 重建中";
$("model").textContent = s.model || "-";
$("tps").textContent = (s.tok_per_s || 0).toFixed(1);
$("rss").textContent = fmtBytes(s.rss_bytes || 0);
if (document.activeElement !== $("slots")) {
$("slots").value = s.expert_slots;
updateSlots();
}
} catch (e) { /* 服务暂不可达时下轮再试 */ }
}
$("apply").addEventListener("click", async () => {
const body = {
expert_slots: +$("slots").value,
k: +$("k").value,
max_tokens: +$("maxTokens").value
};
try {
const r = await (await fetch("/api/engine/config", {
method: "POST",
headers: {"Content-Type": "application/json"},
body: JSON.stringify(body)
})).json();
$("msg").textContent = r.reloading ? "已应用,引擎正在后台重建…" : "已应用";
} catch (e) {
$("msg").textContent = "应用失败";
}
setTimeout(() => { $("msg").textContent = ""; }, 4000);
});
(async function init() {
try {
const c = await (await fetch("/api/engine/config")).json();
$("slots").value = c.expert_slots;
$("k").value = c.k;
$("maxTokens").value = c.max_tokens;
} catch (e) { /* 用默认值 */ }
updateSlots();
refresh();
setInterval(refresh, 2500);
})();
</script>
</body>
</html>
"""

104
mlx_streaming/server/app.py Normal file
View File

@ -0,0 +1,104 @@
"""FastAPI 应用装配:OpenAI 路由 + 引擎调参 API + 极简 admin 页。"""
from __future__ import annotations
import asyncio
import os
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request
from fastapi.responses import HTMLResponse
from mlx_streaming.server.admin import ADMIN_HTML
from mlx_streaming.server.openai import make_openai_router
from mlx_streaming.server.state import EngineManager, rss_bytes
from mlx_streaming.tui.backend import FakeBackend
def create_app(args, backend=None) -> FastAPI:
"""装配应用。
args: argparse Namespace 风格对象(至少含 model/k/max_tokens/expert_slots,
真后端还需要 expert_dir/mtp_out/qn_config/spec_slots)
backend: 显式注入的后端(测试用);None 时按 SPARKLE_FAKE=1 FakeBackend,
否则 MLXBackend
"""
if backend is None:
if os.environ.get("SPARKLE_FAKE") == "1":
backend = FakeBackend()
else:
from mlx_streaming.tui.backend import MLXBackend
backend = MLXBackend(args)
mgr = EngineManager(args, backend)
@asynccontextmanager
async def lifespan(app):
# 加载放线程且不等待:真模型以分钟计,不能阻塞 uvicorn 启动;
# 就绪前 loaded=False,chat 请求返回 503。
loop = asyncio.get_running_loop()
app.state.load_future = loop.run_in_executor(None, mgr.load_initial)
yield
app = FastAPI(title="sparkle server", lifespan=lifespan)
# 本机工具服务:放开 CORS,让浏览器能直连 /v1/models、/api/* 做探测
from fastapi.middleware.cors import CORSMiddleware
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
app.state.mgr = mgr
app.include_router(make_openai_router(mgr))
@app.get("/api/stats")
async def api_stats():
return {"tok_per_s": mgr.last_tok_per_s,
"rss_bytes": rss_bytes(),
"expert_slots": mgr.expert_slots,
"loaded": mgr.loaded,
# 专家池「分不到槽被迫落 0 号槽」的累计数。>0 = 有前向拿错专家权重算过,
# 输出不可信(调小 PREFILL_CHUNK 或调大 expert_slots)。健康值恒为 0。
"unplaced_experts": mgr.unplaced_experts(),
# 专家池命中/读盘指标:hit_rate 低就是慢的直接原因(每 miss 一个要读 ~3.3MB)
"pool": mgr.pool_metrics(),
"model": str(mgr.args.model)}
@app.get("/api/engine/config")
async def get_engine_config():
return {"expert_slots": mgr.expert_slots,
"k": mgr.k,
"max_tokens": mgr.max_tokens}
@app.post("/api/engine/config")
async def set_engine_config(req: Request):
"""有多少改多少:k/max_tokens 即时生效;expert_slots 变化触发后台重建。"""
body = await req.json()
if "k" in body:
mgr.k = int(body["k"])
if "max_tokens" in body:
mgr.max_tokens = int(body["max_tokens"])
reloading = False
if "expert_slots" in body:
slots = int(body["expert_slots"])
if slots != mgr.expert_slots:
mgr.expert_slots = slots
mgr.reloading = True # 即刻对 chat 可见(503),不等后台任务起跑
reloading = True
async def _bg_rebuild():
async with mgr.lock: # 等当前生成结束再换引擎
await asyncio.get_running_loop().run_in_executor(
None, mgr.rebuild)
asyncio.create_task(_bg_rebuild())
return {"ok": True, "reloading": reloading,
"expert_slots": mgr.expert_slots,
"k": mgr.k, "max_tokens": mgr.max_tokens}
@app.get("/admin", response_class=HTMLResponse)
async def admin_page():
return ADMIN_HTML
return app

View File

@ -0,0 +1,338 @@
"""OpenAI 兼容端点:/v1/models 与 /v1/chat/completions(SSE 流式 + 普通 JSON)。
生成是阻塞调用,放线程池执行;流式用 asyncio.Queue 把后端回调桥接成 SSE 生成器
客户端断开时通过 threading.Event backend on_text 返回 True 中断生成
"""
from __future__ import annotations
import asyncio
import json
import os
import re
import threading
import time
import uuid
from fastapi import APIRouter, Request
from fastapi.responses import JSONResponse, StreamingResponse
from mlx_streaming import config
from mlx_streaming.server.state import EngineManager
def _sse(obj) -> str:
return f"data: {json.dumps(obj, ensure_ascii=False)}\n\n"
def _delta_of(prev: str, cur: str) -> str:
"""后端上报的是累计全文,这里取出本次新增。
按真实公共前缀切,而不是 `cur[len(prev):]`:后者假设 prev 一定是 cur 的前缀,一旦后端
改写了已发出的尾部(历史上流式解码到半个汉字会先发 U+FFFD下一步再换成真字符),
长度相同就切出空串,那个字符永久丢失此处宁可重发一小段,也不丢字
"""
if cur.startswith(prev):
return cur[len(prev):]
n = 0
for a, b in zip(prev, cur):
if a != b:
break
n += 1
return cur[n:]
_MASK_RE = re.compile(r"\[(NAME|ID|TEAM|STATION|CAR|TRAIN):(emp_[a-z0-9]+)\]", re.I)
def _mask_notice(content: str) -> str:
"""工具结果里带脱敏占位符时,就地追加一句原样透传的约束。
实测(2026-08-02):工具返回 `[TEAM:emp_ypsl004]`,模型答"乘务三组"把占位符替换成
了自己编的名字,用户拿到错数据系统提示里已有同样的规则,但它在 4000 token 之前,
而占位符就在眼前;贴着数据再说一次显著更硬,且模型无关
只在真的检出占位符时加,不给普通工具结果增加 token
"""
if not content or not _MASK_RE.search(content):
return content
return (f"{content}\n\n[脱敏提示] 上面形如 [TEAM:emp_xxx]、[NAME:emp_xxx]、[ID:emp_xxx] 的是"
"脱敏占位符,界面会自动还原成真实值。引用时**必须连方括号整体原样复制**,"
"严禁替换成「乘务一组」这类具体名称(你看不到真实值,写具体名称一定是编的),"
"也不要拆掉方括号或向用户解释占位符。")
def _normalize_messages(raw) -> list[dict]:
"""OpenAI messages → chat template 可用消息。
关键:assistant tool_calls role=tool 的消息**原样透传**Qwen chat template
会把它们渲染成模型训练时学过的格式(<tool_call> <tool_response> 包裹)
此前把工具返回改写成普通 user 文本,模型不认为是权威工具结果,会无视真实数据
自行编造报告(2026-07 实测复现)只做 content 分段列表的拍平
"""
out = []
for m in raw:
role = m.get("role", "user")
content = m.get("content")
if isinstance(content, list): # OpenAI 分段 content:[{"type":"text","text":...}]
content = "".join(
p.get("text", "") for p in content
if isinstance(p, dict) and p.get("type") == "text")
if role == "tool":
# 原生透传:模板渲染为 <tool_response>,模型按训练格式 ground 到工具数据
item = {"role": "tool", "content": _mask_notice(content or "")}
if m.get("tool_call_id"):
item["tool_call_id"] = m["tool_call_id"]
if m.get("name"):
item["name"] = m["name"]
out.append(item)
continue
if role == "assistant" and m.get("tool_calls"):
# 原生透传:模板渲染为 <tool_call>。arguments 必须是 JSON 字符串(OpenAI 规范)
tcs = []
for tc in m["tool_calls"]:
fn = tc.get("function", {})
args = fn.get("arguments")
if args is not None and not isinstance(args, str):
args = json.dumps(args, ensure_ascii=False)
tcs.append({"id": tc.get("id"), "type": "function",
"function": {"name": fn.get("name"), "arguments": args}})
out.append({"role": "assistant", "content": content or None,
"tool_calls": tcs})
continue
out.append({"role": role, "content": content or ""})
return out
def _tool_name(t) -> str:
if not isinstance(t, dict):
return ""
fn = t.get("function")
if isinstance(fn, dict) and fn.get("name"):
return str(fn["name"])
return str(t.get("name") or "")
_tools_filter_seen: set = set()
def _filter_tools(tools):
"""按 config.tools_allow() 白名单裁剪工具,全被裁掉时返回 None(本轮不暴露工具)。
config.tools_allow 的说明:这不只是省 token,主要是把 prompt 头部钉稳,让前缀快照
能命中调用方的工具发现是动态的,工具集一变头部就变,每个新会话都要全量 prefill
首次遇到某个保留/丢弃组合时打一行日志,便于确认线上到底发来了什么
"""
allow = config.tools_allow()
if not allow or not tools:
return tools
allow_set = set(allow)
kept, dropped = [], []
for t in tools:
(kept if _tool_name(t) in allow_set else dropped).append(t)
if dropped:
sig = (tuple(_tool_name(t) for t in kept), tuple(_tool_name(t) for t in dropped))
if sig not in _tools_filter_seen:
_tools_filter_seen.add(sig)
print(f"[TOOLS_ALLOW] 保留 {len(kept)}{list(sig[0])},"
f"丢弃 {len(dropped)}{list(sig[1])}"
f"目的是让 prompt 头部固定、前缀快照命中;设 TOOLS_ALLOW= 可关闭。",
flush=True)
return kept or None
def make_openai_router(mgr: EngineManager) -> APIRouter:
router = APIRouter()
def _model_id() -> str:
return os.path.basename(str(mgr.args.model).rstrip("/")) or "sparkle"
def _auth(req: Request):
"""Bearer 任意非空 token 即放行;设了 SPARKLE_API_KEY 则必须匹配。"""
auth = req.headers.get("authorization", "")
tok = auth[7:].strip() if auth.lower().startswith("bearer ") else ""
want = os.environ.get("SPARKLE_API_KEY", "")
ok = (tok == want) if want else bool(tok)
if ok:
return None
return JSONResponse(status_code=401, content={
"error": {"message": "unauthorized", "type": "authentication_error"}})
def _unavailable():
msg = "engine reloading" if mgr.reloading else "engine not loaded"
return JSONResponse(status_code=503, content={"error": msg})
def _prompt_tokens(messages, tools=None) -> int:
"""prompt token 数:真后端用 tokenizer 精确编码;fake 退化为字符数估算。
必须带上 tools:工具定义是由 chat template 渲染进 prompt ,实测 4 个工具就占
1349 token(占总 prompt 三分之一)漏掉它上报的 usage 会显著偏小,而调用方
(中台) usage 判断要不要压缩上下文,少算就会算漏
"""
tok = getattr(mgr.backend, "_tok", None)
if tok is not None:
try:
from mlx_streaming.cli import _encode_chat
return len(_encode_chat(tok, messages, tools=tools))
except Exception: # noqa: BLE001 估算失败不阻塞生成
pass
return sum(len(m.get("content") or "") for m in messages)
def _usage(prompt_tokens: int, completion_tokens: int) -> dict:
return {"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens}
def _chunk(model_id, cmpl_id, created, delta=None, role=False,
finish=None, usage=None):
if usage is not None: # OpenAI 风格:usage chunk 的 choices 为空数组
return {"id": cmpl_id, "object": "chat.completion.chunk",
"created": created, "model": model_id,
"choices": [], "usage": usage}
d = {}
if role:
d["role"] = "assistant"
if delta:
d["content"] = delta
return {"id": cmpl_id, "object": "chat.completion.chunk",
"created": created, "model": model_id,
"choices": [{"index": 0, "delta": d, "finish_reason": finish}]}
@router.get("/v1/models")
async def list_models(req: Request):
err = _auth(req)
if err is not None:
return err
return {"object": "list", "data": [{
"id": _model_id(), "object": "model",
"created": int(time.time()), "owned_by": "sparkle"}]}
@router.post("/v1/chat/completions")
async def chat_completions(req: Request):
err = _auth(req)
if err is not None:
return err
if not mgr.loaded or mgr.reloading:
return _unavailable()
body = await req.json()
# chat_template_kwargs / enable_thinking:接受并忽略。
# tools:渲染进 chat template(Qwen 模板把工具定义注入提示词),模型据此输出
# <tool_call> 文本,客户端兜底解析——server 侧不做原生 tool parser。
raw_tools = body.get("tools") or None
tools = _filter_tools(raw_tools)
# 「模型没调工具就编数据」有两种成因:调用方压根没发工具,或发了而模型无视。
# 两者在 server 侧看不出区别(都只是一次普通请求),故每轮记一行收到的工具名。
print(f"[REQ] 收到工具 {len(raw_tools or [])}"
f"{[_tool_name(t) for t in (raw_tools or [])]} → 渲染 "
f"{len(tools or [])}{[_tool_name(t) for t in (tools or [])]}", flush=True)
# 采样策略(关键):模型 generation_config 要求 do_sample(temp 0.7/top_p 0.8/top_k 20),
# 贪心会诱发幻觉。默认走采样模式(质量优先,非投机);仅显式 temperature=0 走贪心+MTP。
req_temp = body.get("temperature")
sampling = None
if req_temp is None or float(req_temp) > 0:
gc = getattr(mgr.backend, "_gen_cfg", None) or {}
sampling = {
"temperature": float(req_temp) if req_temp is not None
else float(gc.get("temperature", 0.7)),
"top_p": body.get("top_p") or gc.get("top_p", 0.8),
"top_k": body.get("top_k") or gc.get("top_k", 20),
}
messages = _normalize_messages(body.get("messages") or [])
stream = bool(body.get("stream"))
include_usage = bool((body.get("stream_options") or {}).get("include_usage"))
req_max = body.get("max_tokens") or body.get("max_completion_tokens")
model_id = _model_id()
cmpl_id = "chatcmpl-" + uuid.uuid4().hex[:24]
created = int(time.time())
loop = asyncio.get_running_loop()
disconnect = threading.Event() # 客户端断开 → on_text 返回 True 中断生成
async def _watch_disconnect():
while not disconnect.is_set():
if await req.is_disconnected():
disconnect.set()
return
await asyncio.sleep(0.2)
if not stream:
async with mgr.lock: # 引擎单对话:并发请求排队等锁
if not mgr.loaded or mgr.reloading: # 等锁期间可能开始重建
return _unavailable()
mgr.apply_gen_params(req_max)
prompt_tokens = _prompt_tokens(messages, tools)
watcher = asyncio.create_task(_watch_disconnect()) # 客户端断开也中断,防僵尸生成堵锁
try:
res = await loop.run_in_executor(
None, lambda: mgr.backend.generate(
messages, lambda _f, _n: disconnect.is_set(),
tools=tools, sampling=sampling))
except Exception as e: # noqa: BLE001
return JSONResponse(status_code=500, content={"error": str(e)})
finally:
disconnect.set()
watcher.cancel()
mgr.last_tok_per_s = res.tok_per_s
return {"id": cmpl_id, "object": "chat.completion",
"created": created, "model": model_id,
"choices": [{"index": 0, "finish_reason": "stop",
"message": {"role": "assistant",
"content": res.text}}],
"usage": _usage(prompt_tokens, res.n_tokens)}
async def event_stream():
async with mgr.lock:
if not mgr.loaded or mgr.reloading:
yield _sse({"error": "engine reloading"})
yield "data: [DONE]\n\n"
return
mgr.apply_gen_params(req_max)
prompt_tokens = _prompt_tokens(messages, tools)
q = asyncio.Queue()
def on_text(full, _n):
if disconnect.is_set():
return True
loop.call_soon_threadsafe(q.put_nowait, full)
return False
def _run(): # 阻塞生成,放线程池;结束/异常都经由 queue 通知
try:
res = mgr.backend.generate(messages, on_text, tools=tools,
sampling=sampling)
loop.call_soon_threadsafe(q.put_nowait, ("__done__", res))
except Exception as e: # noqa: BLE001
loop.call_soon_threadsafe(q.put_nowait, ("__error__", e))
watcher = asyncio.create_task(_watch_disconnect())
loop.run_in_executor(None, _run)
prev = ""
result = None
yield _sse(_chunk(model_id, cmpl_id, created, role=True))
try:
while True:
item = await q.get()
if isinstance(item, tuple):
if item[0] == "__done__":
result = item[1]
break
raise item[1]
delta = _delta_of(prev, item)
prev = item
if delta:
yield _sse(_chunk(model_id, cmpl_id, created,
delta=delta))
finally:
disconnect.set() # 客户端断开/生成结束都确保后端被中断
watcher.cancel()
if result is not None:
mgr.last_tok_per_s = result.tok_per_s
yield _sse(_chunk(model_id, cmpl_id, created, finish="stop"))
if include_usage:
n = result.n_tokens if result is not None else 0
yield _sse(_chunk(model_id, cmpl_id, created,
usage=_usage(prompt_tokens, n)))
yield "data: [DONE]\n\n"
return StreamingResponse(event_stream(),
media_type="text/event-stream")
return router

View File

@ -0,0 +1,101 @@
"""引擎单例管理:加载状态、生成串行锁、运行期调参与后台重建。"""
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

View File

@ -0,0 +1,9 @@
"""sparkle 全屏 TUI 包。run_tui 延迟 import app(避免无谓加载 textual)。"""
def run_tui(backend, args) -> int:
"""启动全屏 TUI。backend 为 ChatBackend 实现,args 为解析后的命令行参数。"""
from mlx_streaming.tui.app import SparkleApp
SparkleApp(backend, args).run()
return 0

334
mlx_streaming/tui/app.py Normal file
View File

@ -0,0 +1,334 @@
"""sparkle 全屏 TUI(Textual 实现,视觉对标 opencode)。
界面只依赖 ChatBackend 接口加载与每轮生成均在 worker 线程执行,
通过 call_from_thread 把状态/流式文本安全回推到 UI 线程
"""
from __future__ import annotations
import os
import time
from collections import deque
from typing import Optional
from rich.console import Group
from rich.markdown import Markdown as RichMarkdown
from rich.text import Text
from textual import work
from textual.app import App, ComposeResult
from textual.containers import VerticalScroll
from textual.widgets import Input, Static
from mlx_streaming.core.mem import snapshot
from mlx_streaming.tui.backend import ChatBackend, GenResult
from mlx_streaming.tui.banner import LOGO
_ACCENT = "#2dd4bf"
def _fmt_gb(nbytes: int) -> str:
return f"{nbytes / 1e9:.2f} GB"
def _mem_suffix(peak: bool = False) -> str:
"""状态栏内存后缀。显存外置流式 MoE 的核心卖点是低内存,故常驻展示活跃占用;
结束时附带峰值 MLX 统一内存计数器,开销极低取不到(0)时返回空串,不干扰状态栏"""
snap = snapshot()
if snap.mlx_active_bytes <= 0:
return ""
s = f" · 内存 {_fmt_gb(snap.mlx_active_bytes)}"
if peak and snap.mlx_peak_bytes > 0:
s += f"(峰值 {_fmt_gb(snap.mlx_peak_bytes)})"
return s
# 生成中状态栏用滑动窗口算「瞬时」tok/s 的时间窗(秒);越小越灵敏、越大越平滑。
_TPS_WINDOW = 1.0
_HELP = (
"可用命令:\n"
" /help 显示本帮助\n"
" /reset 清空对话历史(保留 system)\n"
" /clear 清空对话区显示\n"
" /exit 退出\n\n"
"快捷键:Enter 发送 · Esc 中断生成 · Ctrl+C 退出"
)
def _short_model(path: str) -> str:
# 只取模型路径末段,避免顶栏/状态栏过长
return os.path.basename(path.rstrip("/")) or path
class ChatMessage(Static):
"""一条对话消息。role ∈ {'user','assistant'};助手消息支持流式更新与 Markdown 收尾。"""
def __init__(self, role: str, text: str = "", *, final: bool = True):
super().__init__(classes=f"msg {role}")
self.role = role
self.text = text
self.final = final
self._refresh_content()
def stream(self, text: str) -> None:
# 流式过程中用纯文本渲染,避免半截 Markdown 抖动
self.text = text
self.final = False
self._refresh_content()
def finalize(self, text: str) -> None:
# 收尾时切换为 Markdown 渲染
self.text = text
self.final = True
self._refresh_content()
def _refresh_content(self) -> None:
if self.role == "user":
header = Text("", style="bold")
body = Text(self.text)
else:
header = Text("⏺ sparkle", style=f"bold {_ACCENT}")
if not self.text:
body = Text("正在思考…", style="dim italic")
elif self.final:
body = RichMarkdown(self.text)
else:
body = Text(self.text)
self.update(Group(header, body))
class SparkleApp(App):
CSS_PATH = "styles.tcss"
BINDINGS = [
("escape", "interrupt", "中断"),
("ctrl+c", "quit", "退出"),
]
def __init__(self, backend: ChatBackend, args):
super().__init__()
self.backend = backend
self.args = args
self._messages: list[dict] = []
# 保留 system 消息,/reset 时不清空这部分
if getattr(args, "system", None):
self._messages.append({"role": "system", "content": args.system})
self._base_len = len(self._messages)
self._busy = False
self._stop = False
self._cur: Optional[ChatMessage] = None
# 首 token 到达时刻与基准 token 数,用于结束时计算「累计解码平均」tok/s
# (排除 prefill)。_gen_t0 为 0.0 表示本轮尚未收到首 token。
self._gen_t0 = 0.0
self._n0 = 0
# 生成中「瞬时」tok/s 的滑动窗口采样:每项为 (时刻, 累计 token 数)。
self._tps_window = _TPS_WINDOW
self._samples: deque[tuple[float, int]] = deque()
def compose(self) -> ComposeResult:
yield Static(self._top_text(), id="top")
yield VerticalScroll(id="chat")
# 提示放到边框标题里,不用文本区的 placeholder:
# 长占位符在部分终端增量重绘时不会被擦除,打字后会残留「后面还有字」,
# 直到全量重绘(Enter/截图/resize)才消失;边框标题在边框上,不受此影响。
yield Input(id="prompt")
yield Static("", id="status")
def on_mount(self) -> None:
inp = self.query_one("#prompt", Input)
inp.border_title = "输入消息 · Enter 发送 · Esc 中断 · /help"
# 加载完成前禁用输入框,避免用户在模型就绪前发消息
inp.disabled = True
self._set_status("加载模型中…")
self._load()
def on_input_changed(self, event: Input.Changed) -> None:
# 兜底:强制整屏重绘,清除个别终端增量重绘遗留的输入残影
# (等价于截图/resize 触发的全量刷新)。
self.refresh()
def _top_text(self) -> Text:
t = Text()
t.append(LOGO, style=f"bold {_ACCENT}")
t.append(f" {_short_model(self.args.model)} · k={self.args.k} · "
f"{self.args.max_tokens} tok", style="dim")
return t
def _set_status(self, text: str) -> None:
self.query_one("#status", Static).update(
f" {_short_model(self.args.model)} · {text}")
@work(thread=True, exclusive=True, group="load")
def _load(self) -> None:
# 在 worker 线程执行阻塞加载,回调需经 call_from_thread 回到 UI 线程
def on_status(msg: str):
# 应用已退出时不再回推,避免 call_from_thread 抛错
if self.is_running:
self.call_from_thread(self._set_status, f"加载中 · {msg}")
try:
self.backend.load(on_status)
except Exception as e: # noqa: BLE001
if self.is_running:
self.call_from_thread(self._on_load_failed, str(e))
return
if self.is_running:
self.call_from_thread(self._on_load_done)
def _enable_input(self) -> None:
"""重新启用并聚焦输入框(加载完成/一轮生成结束后统一调用)。
刚把 disabled False can_focus 状态可能还没刷新,直接 focus 偶发无效
(表现为输入框失焦时占位提示不消失);故延到下一次刷新后再聚焦,确保稳定拿到焦点
"""
inp = self.query_one("#prompt", Input)
inp.disabled = False
self.call_after_refresh(inp.focus)
def _on_load_done(self) -> None:
self._set_status(f"就绪{_mem_suffix()}")
self._enable_input()
def _on_load_failed(self, err: str) -> None:
self._add("assistant", f"模型加载失败:{err}\n\n请检查路径后用 /exit 退出重试。")
self._set_status("加载失败")
def on_input_submitted(self, event: Input.Submitted) -> None:
text = event.value.strip()
event.input.value = ""
if not text:
return
if text.startswith("/"):
self._command(text)
return
if self._busy:
return
self._start(text)
def _command(self, cmd: str) -> None:
if cmd in ("/exit", "/quit"):
self.exit()
elif cmd == "/help":
self._add("assistant", _HELP)
elif cmd in ("/reset", "/clear") and self._busy:
# 生成中改动历史/移除正在流式的 _cur 会让 _on_stream/_on_done 操作已卸载组件,故拒绝
self._add("assistant", "生成中,请等本轮结束或按 Esc 中断后再执行该命令。")
elif cmd == "/reset":
# 只清空 system 之后的历史
del self._messages[self._base_len:]
self._add("assistant", "对话历史已清空。")
elif cmd == "/clear":
self.query_one("#chat", VerticalScroll).remove_children()
else:
self._add("assistant", f"未知命令:{cmd}(/help 查看可用命令)")
def _start(self, user_text: str) -> None:
self._messages.append({"role": "user", "content": user_text})
self._add("user", user_text)
self._cur = self._add("assistant", "", final=False)
self._busy = True
self._stop = False
# 置 0 表示还没收到首 token;真正的计时起点推迟到第一次 _on_stream。
self._gen_t0 = 0.0
self._n0 = 0
self._samples.clear()
self.query_one("#prompt", Input).disabled = True
self._set_status("思考中…")
self._generate(list(self._messages))
@work(thread=True, exclusive=True, group="gen")
def _generate(self, messages: list[dict]) -> None:
# 生成在 worker 线程运行;on_text 返回 self._stop 让后端可提前中断
def on_text(full: str, n_tokens: int) -> bool:
# 应用已退出时,call_from_thread 会抛错;此时直接请求停止,避免 worker 线程未捕获异常。
if not self.is_running:
return True
try:
self.call_from_thread(self._on_stream, full, n_tokens)
except Exception: # noqa: BLE001
return True
return self._stop
try:
result = self.backend.generate(messages, on_text)
except Exception as e: # noqa: BLE001
if self.is_running:
self.call_from_thread(self._on_error, str(e))
return
if self.is_running:
self.call_from_thread(self._on_done, result)
def _record_sample(self, now: float, n_tokens: int) -> None:
"""记录一次采样,并丢弃早于滑动窗口的旧点(保留跨越窗口边界的那一个)。"""
self._samples.append((now, n_tokens))
w = self._tps_window
# 当第二个点仍比窗口还老时,第一个点已冗余,可丢弃;
# 循环后 _samples[0] 恰好是刚跨过窗口边界的采样,窗口长度 ≈ w。
while len(self._samples) >= 2 and now - self._samples[1][0] >= w:
self._samples.popleft()
def _window_tps(self, now: float) -> Optional[float]:
"""按滑动窗口算瞬时 tok/s;样本不足(不够两点或时间差为 0)时返回 None。"""
if len(self._samples) < 2:
return None
t0, n0 = self._samples[0]
dt = now - t0
if dt <= 0:
return None
return (self._samples[-1][1] - n0) / dt
def _on_stream(self, full: str, n_tokens: int) -> None:
if self._cur is not None:
self._cur.stream(full)
self._scroll_end()
now = time.monotonic()
# 首次回调:prefill 刚结束,记下计时起点与基准 token 数,供结束时算累计解码平均。
if self._gen_t0 == 0.0:
self._gen_t0 = now
self._n0 = n_tokens
# 生成中显示滑动窗口「瞬时」速度:一直在动,能反映后期变慢,不被历史平均拖住。
self._record_sample(now, n_tokens)
tps = self._window_tps(now)
mem = _mem_suffix()
if tps is None:
self._set_status(f"思考中 · {n_tokens} tok{mem}")
else:
self._set_status(f"思考中 · {n_tokens} tok · {tps:.1f} tok/s{mem}")
def _on_done(self, result: GenResult) -> None:
if self._cur is not None:
self._cur.finalize(result.text)
self._messages.append({"role": "assistant", "content": result.text})
self._busy = False
self._cur = None
suffix = " · 已中断" if result.stopped else ""
# 与流式状态栏同口径:从首 token 起、按解码 token 数算,避免结束瞬间数字回落。
# 若本轮没触发过流式(_gen_t0 仍为 0),退回后端上报的 tok/s。
dt = time.monotonic() - self._gen_t0
if self._gen_t0 > 0.0 and dt > 0:
tps = (result.n_tokens - self._n0) / dt
else:
tps = result.tok_per_s
self._set_status(
f"就绪 · {result.n_tokens} tok · {tps:.1f} tok/s{_mem_suffix(peak=True)}{suffix}")
self._enable_input()
def _on_error(self, err: str) -> None:
if self._cur is not None:
self._cur.finalize(f"生成出错:{err}")
self._busy = False
self._cur = None
self._set_status("就绪(上一轮出错)")
self._enable_input()
def action_interrupt(self) -> None:
# 仅在生成中时置位中断标志,由 on_text 闭包读取
if self._busy:
self._stop = True
self._set_status("正在中断…")
def _add(self, role: str, text: str, *, final: bool = True) -> ChatMessage:
msg = ChatMessage(role, text, final=final)
self.query_one("#chat", VerticalScroll).mount(msg)
self._scroll_end()
return msg
def _scroll_end(self) -> None:
self.query_one("#chat", VerticalScroll).scroll_end(animate=False)

View File

@ -0,0 +1,436 @@
"""TUI 后端抽象:把界面与 MLX 推理引擎解耦。
界面只依赖 ChatBackend 接口, import 任何 MLX 符号,从而能用 FakeBackend 做无模型测试
load / generate 都是阻塞调用, UI 层放到 worker 线程执行
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Callable, Protocol
from collections import OrderedDict
@dataclass
class GenResult:
"""一轮生成的汇总。"""
text: str # 完整回答(已截断 EOS)
n_tokens: int # 新生成 token 数
tok_per_s: float # 吞吐
stopped: bool # 是否被用户中断
class ChatBackend(Protocol):
"""聊天后端接口。所有方法阻塞,调用方负责放 worker 线程。"""
def load(self, on_status: Callable[[str], None]) -> None:
"""加载模型/权重;通过 on_status(msg) 上报进度。"""
...
def generate(
self,
messages: list[dict],
on_text: Callable[[str, int], bool],
) -> GenResult:
"""跑一轮生成。每步把「累计完整文本, 已生成 token 数」传给 on_text;返回 True 表示请求中断。"""
...
@dataclass
class FakeBackend:
"""测试用假后端:不加载模型,把预设回答按字符流式吐出。
delay > 0 时每字符间 sleep,用于 --demo 模式模拟真实吐字节奏;测试默认 0(不拖慢)
"""
reply: str = "你好,这是一个测试回答。"
status_msgs: list[str] = field(default_factory=lambda: ["加载中(模拟)…"])
delay: float = 0.0
seen_messages: list[list[dict]] = field(default_factory=list)
def load(self, on_status: Callable[[str], None]) -> None:
for m in self.status_msgs:
on_status(m)
def generate(self, messages, on_text, tools=None, sampling=None) -> GenResult:
import time
self.seen_messages.append([dict(m) for m in messages])
acc = ""
for ch in self.reply:
acc += ch
if self.delay:
time.sleep(self.delay)
# 用字符数近似 token 数(假后端无真实分词)
if on_text(acc, len(acc)):
return GenResult(acc, len(acc), 0.0, stopped=True)
return GenResult(acc, len(self.reply), 0.0, stopped=False)
def _common_prefix_len(a, b) -> int:
"""返回两个 token id 序列的最长公共前缀长度。"""
n = min(len(a), len(b))
i = 0
while i < n and a[i] == b[i]:
i += 1
return i
def _reuse_prefix_len(cached_ids, new_ids) -> int:
"""可跨轮复用的前缀长度:仅当旧 cache 的 token 是新序列的**严格前缀**且新序列更长时,
返回该前缀长度(= len(cached_ids));否则返回 0 表示需全量重建
只在严格前缀时复用,是为了永远只延续不回退cacheQwen3-Next 的线性注意力递归态
无法裁剪回任意历史位置; detokenizeretokenize 不一致/reset编辑历史等都会让公共前缀
短于旧长度,此时回退整段重建,绝不基于错位的 cache 续算
"""
if not cached_ids or len(new_ids) <= len(cached_ids):
return 0
c = _common_prefix_len(cached_ids, new_ids)
return c if c == len(cached_ids) else 0
class MLXBackend:
"""真实后端:封装 _build_engine + mtp_generate。MLX 相关 import 全部延迟到方法内。"""
def __init__(self, args):
self.args = args
self._model = None
self._tok = None
self._drafter = None
# 跨轮复用:持久化上一轮的 main_cache 及其对应的 token 序列(prompt + 已入 cache 的生成 token)。
self._main_cache = None
self._cached_ids: list[int] = []
# 跨会话前缀快照:{head_key: (snaps, head_ids)},LRU。agentic 场景每个新会话
# 都重复同一段长 system+tools 前缀,快照后首轮 prefill 从 ~2600 tok 降到几百。
self._head_snaps: "OrderedDict[tuple, tuple]" = OrderedDict()
# (system 文本, 工具名) → 固定头部 token 数,见 _fixed_prefix_len
self._fixed_head_key = None
self._fixed_head_len = 0
# 流式反分词器(单例,每轮 reset)。建实例要铺一张 15 万项的 id→token 表,不能每轮重建;
# 生成由 server 的单飞锁串行,单例安全。
self._detok = None
def load(self, on_status: Callable[[str], None]) -> None:
from mlx_streaming.cli import _build_engine, _warmup
self._model, self._tok, self._drafter = _build_engine(
self.args, on_status=on_status
)
# 预热:把首轮的 kernel 编译 + 专家池填充开销移到加载阶段,避免第一条消息莫名卡很久。
on_status("预热中(编译 kernel + 填专家池)…")
_warmup(self._model, self._tok, self._drafter, self.args)
# 采样参数:读模型自带 generation_config(Qwen3-Next 官方要求 do_sample
# temp=0.7/top_p=0.8/top_k=20;贪心 argmax 会诱发幻觉/编造,见 2026-07 排查)
import json
import os
self._gen_cfg = {"temperature": 0.7, "top_p": 0.8, "top_k": 20}
try:
with open(os.path.join(str(self.args.model), "generation_config.json")) as f:
gc = json.load(f)
if gc.get("do_sample", True):
for k in ("temperature", "top_p", "top_k"):
if k in gc:
self._gen_cfg[k] = gc[k]
except OSError:
pass
def _generate_sampled(self, ids, on_tokens, max_tokens, sampling,
main_cache, cached_len):
"""采样生成(非投机):prefill 后逐 token 前向 + temp/top_p/top_k 采样。
返回 (produced, forwarded):forwarded = 已写入 cache 的生成 token ,
供跨轮复用不变式(resident == len(ids)+forwarded)判定
Qwen3-Next generation_config 要求 do_sample;贪心 argmax agentic
提示下会编造工具结果/幻觉工具名(实测)代价: MTP 加速,decode 4-6 tok/s
"""
import mlx.core as mx
from mlx_lm.sample_utils import make_sampler
from mlx_streaming.mtp.generate import forward_with_hidden, prefill_chunked
sampler = make_sampler(
temp=float(sampling.get("temperature", 0.7)),
top_p=float(sampling.get("top_p") or 0.0),
top_k=int(sampling.get("top_k") or 0))
def _logprobs(logits):
"""末位 logits → float32 归一化 log-probs,这是 mlx_lm sampler 的入参契约。
必须归一化:apply_top_p 里是 `probs = mx.exp(x)` 再与 `1 - top_p` 比累积和,
喂未归一化的 logits 会让 exp() 变成天文数字累积和瞬间越过阈值 top_p 完全失效
(只砍掉概率 <1e-12 的尾巴)必须转 float32:lm_head 输出是 bf16(8 位尾数),
exp/cumsum/categorical bf16 里做会把分布压得面目全非
"""
lg = logits[:, -1, :].astype(mx.float32)
return lg - mx.logsumexp(lg, axis=-1, keepdims=True)
ids_mx = ids if isinstance(ids, mx.array) else mx.array([ids])
logits, _ = prefill_chunked(self._model, ids_mx[:, cached_len:], main_cache)
cur = _logprobs(logits)
produced = []
forwarded = 0
for _ in range(max_tokens):
nxt = int(sampler(cur))
produced.append(nxt)
if on_tokens(produced[-1:]):
break # 停止:当前 token 未写入 cache
logits, _ = forward_with_hidden(self._model, mx.array([[nxt]]), main_cache)
mx.eval(logits)
forwarded += 1
cur = _logprobs(logits)
return produced, forwarded
# ---- 前缀快照磁盘持久化 ----
# 目录:models/prefix_snapshots/{key}.safetensors + {key}.json(sidecar:head ids/模型名)。
# 首个会话 prefill 一次落盘;之后(含引擎重启后)同前缀会话直接读盘恢复,消灭"首次"。
def _snap_dir(self) -> str:
import os
d = os.path.join(os.path.dirname(str(self.args.expert_dir)) or ".", "prefix_snapshots")
os.makedirs(d, exist_ok=True)
return d
@staticmethod
def _snap_key(ids, head: int) -> str:
import hashlib
import numpy as np
return hashlib.sha1(np.array(ids[:head], dtype=np.int32).tobytes()).hexdigest()[:16]
def _snap_save(self, key: str, snaps, ids, head: int) -> None:
import json
import os
import mlx.core as mx
from mlx.utils import tree_flatten
from mlx_streaming import config as _cfg
flat = {}
metas = []
for ci, (st, meta) in enumerate(snaps):
metas.append(list(meta) if isinstance(meta, (tuple, list)) else (meta or ""))
for k, v in tree_flatten(st):
flat[f"c{ci}.{k}"] = v
path = os.path.join(self._snap_dir(), key)
mx.save_safetensors(path + ".safetensors", flat)
with open(path + ".json", "w") as f:
# kv_quant 决定 cache 类型(量化/非量化),快照跨配置复用会 500 或错算,必须校验
json.dump({"head": head, "model": os.path.basename(str(self.args.model)),
"kv_quant": _cfg.kv_quant(),
"metas": metas, "ids": list(ids[:head])}, f)
def _snap_load(self, key: str, ids, head: int):
"""磁盘命中返回 snaps(供 _cache_restore);缺失/模型不符/损坏返回 None。"""
import json
import os
import mlx.core as mx
from mlx.utils import tree_flatten, tree_unflatten
path = os.path.join(self._snap_dir(), key)
try:
with open(path + ".json") as f:
side = json.load(f)
if side.get("head") != head or side.get("ids") != list(ids[:head]):
return None
if side.get("model") != os.path.basename(str(self.args.model)):
return None
# kv_quant 配置必须一致:量化 KV 的快照装进非量化 cache(或反之)会崩/错算
from mlx_streaming import config as _cfg
if bool(side.get("kv_quant", True)) != bool(_cfg.kv_quant()):
return None
flat = mx.load(path + ".safetensors")
metas = side.get("metas") or []
n_cache = max(int(k.split(".", 1)[0][1:]) for k in flat) + 1
snaps = []
for ci in range(n_cache):
sub = {k.split(".", 1)[1]: v for k, v in flat.items()
if k.startswith(f"c{ci}.")}
meta = metas[ci] if ci < len(metas) else ""
if isinstance(meta, list):
meta = tuple(meta)
snaps.append((tree_unflatten(list(sub.items())), meta))
mx.eval([v for st, _ in snaps for v in tree_flatten(st)])
return snaps
except (OSError, ValueError, KeyError, json.JSONDecodeError):
return None
def _fixed_prefix_len(self, messages, tools) -> int:
"""算出 prompt 里「与用户问题无关」的固定头部长度(system + tools 渲染后的 token 数)。
做法是把同样的 system/tools 配一个哑 user 再渲染一次,取两次渲染的公共前缀
分叉点就是 user 内容的起始位置这样得到的边界与提问内容无关,是快照能命中的最大头部:
再多一个 token 就把问题本身包进快照,换个问法即失效;少了则白白重算已知不变的部分
(实测 system+4 工具 = 3972 token,固定切 2048 会让每个新会话多 prefill 1900+ token)
结果按 (system 文本, 工具名集合) 缓存:每轮只在配置变化时付一次 tokenize
"""
from mlx_streaming.cli import _encode_chat
head_msgs = [m for m in messages if m.get("role") == "system"]
if not head_msgs:
return 0
key = ("".join(m.get("content") or "" for m in head_msgs),
tuple(sorted((t.get("function") or {}).get("name", "")
for t in (tools or []))))
if getattr(self, "_fixed_head_key", None) == key:
return self._fixed_head_len
probe = _encode_chat(self._tok, head_msgs + [{"role": "user", "content": "\x00probe"}],
tools=tools)
real = _encode_chat(self._tok, head_msgs + [{"role": "user", "content": "\x01other"}],
tools=tools)
n = 0
for a, b in zip(probe, real):
if a != b:
break
n += 1
self._fixed_head_key, self._fixed_head_len = key, n
return n
def _snapshot_prefix(self, ids, fixed_head: int = 0):
"""跨会话前缀快照。内存命中/磁盘命中:恢复到新 cache 返回 (cache, head);
未命中:prefill 头部深拷贝存内存 + 落盘后返回 (cache, head)
(首个会话不增加前向量,只是把整段 prefill 拆两段); prompt 或关闭时返回 (None, 0)
head min(PREFIX_SNAPSHOT_HEAD 上限, system+tools 实际长度):快照必须停在
用户问题之前才能被后续会话命中,而停得越晚省下的 prefill 越多
"""
import mlx.core as mx
from mlx_streaming import config
from mlx_streaming.mtp.generate import prefill_chunked
from mlx_streaming.mtp.kv_cache import _restore as _cache_restore
from mlx_streaming.mtp.kv_cache import _snapshot as _cache_snapshot
head = config.prefix_snapshot_head()
if fixed_head > 0:
head = min(head, fixed_head)
head = min(head, len(ids) - 1) # 至少留 1 个 token 走前向,否则拿不到 logits
# 门槛按 head 自身算,不按「prompt 比 head 长多少」:快照省下的就是 head 这段,
# 与后面剩多少无关。头部 3972、问题只有 15 token 恰恰是收益最大的情形。
if head < 256:
return None, 0
key = self._snap_key(ids, head)
ent = self._head_snaps.get(key)
if ent is not None:
cache = self._model.make_cache()
_cache_restore(cache, ent)
self._head_snaps.move_to_end(key)
return cache, head
snaps = self._snap_load(key, ids, head)
if snaps is not None:
cache = self._model.make_cache()
_cache_restore(cache, snaps)
self._head_snaps[key] = snaps
self._head_snaps.move_to_end(key)
return cache, head
cache = self._model.make_cache()
prefill_chunked(self._model, mx.array([ids[:head]]), cache)
snaps = _cache_snapshot(cache)
self._head_snaps[key] = snaps
self._head_snaps.move_to_end(key)
while len(self._head_snaps) > config.prefix_snapshot_entries():
self._head_snaps.popitem(last=False)
try:
self._snap_save(key, snaps, ids, head)
except Exception: # noqa: BLE001 落盘失败不影响主路径(内存快照仍生效)
pass
return cache, head
def generate(self, messages, on_text, tools=None, sampling=None) -> GenResult:
import time
import mlx.core as mx
from mlx_streaming.cli import _encode_chat, _eos_set, _truncate_eos
from mlx_streaming.mtp.generate import mtp_generate
tok = self._tok
eos = _eos_set(tok)
ids = _encode_chat(tok, messages, tools=tools)
# 跨轮复用 KV/递归态:旧 cache 是本轮 prompt 的严格前缀时,只 prefill 新增后缀,
# 不重算整段历史(prefill 从 ∝历史长度 降到 ∝新消息长度)。否则全量重建。
cached_len = (_reuse_prefix_len(self._cached_ids, ids)
if self._main_cache is not None else 0)
main_cache = self._main_cache if cached_len else None
# 跨会话前缀快照:同一段长 system+tools 头部的新会话,跳过头部 prefill。
# 首个带该头部的会话在 prefill 头部后深拷贝一份状态存下(不产生额外前向,
# 只是把一个整段 prefill 拆成两段),后续会话直接恢复。
if main_cache is None:
main_cache, cached_len = self._snapshot_prefix(
ids, fixed_head=self._fixed_prefix_len(messages, tools))
if main_cache is None:
main_cache = self._model.make_cache()
cached_len = 0
# 流式反分词:必须逐 token 喂 StreamingDetokenizer,不能每步 tok.decode(全部已产 token)。
# Qwen 是 byte-level BPE,很多汉字/符号跨 2 个 token,整体 decode 到半个字符时会按
# errors="replace" 吐出 U+FFFD;而下游按字符数算增量(delta = full[len(prev):]),下一步
# U+FFFD 被真字符替换后长度不变 → 增量为空、真字符永久丢失,客户端就留下一堆「<E5A086>」。
# detokenizer 的 .text 只包含已完整的字符(未完成字节留在 _unflushed),恒为下一步的真前缀。
detok = self._detok
if detok is None:
detok = self._detok = tok.detokenizer
detok.reset()
produced_all: list[int] = [] # 已上报的正文 token(不含 EOS 及其后)
stopped = {"v": False}
hit_eos = {"v": False}
def on_tokens(new_ids):
for t in new_ids:
if int(t) in eos:
hit_eos["v"] = True
break
detok.add_token(int(t))
produced_all.append(int(t))
if hit_eos["v"]:
detok.finalize() # 冲出尾部残留字节(正常收尾,不会有半字)
if on_text(detok.text, len(produced_all)): # 用户按 Esc 请求中断
stopped["v"] = True
return True
# 命中 EOS:完整回答已生成,提前停止,避免引擎空跑到 max_tokens 让界面
# 长时间卡在「思考中」。EOS 属正常完成,不算中断。
return hit_eos["v"]
t0 = time.perf_counter()
if sampling is not None:
# 采样模式(server 默认,质量优先):非投机逐 token 采样,无 MTP 加速。
produced, forwarded = self._generate_sampled(
ids, on_tokens, self.args.max_tokens, sampling,
main_cache, cached_len)
stats = {"resident_tokens": len(ids) + forwarded}
else:
produced, stats = mtp_generate(
self._model,
self._drafter,
tok,
mx.array([ids]),
self.args.max_tokens,
K=self.args.k,
ids_mode=True,
profile=False,
on_tokens=on_tokens,
main_cache=main_cache,
cached_len=cached_len,
)
dt = time.perf_counter() - t0
# 持久化本轮 cache 供下轮复用。正常情况下 main_cache 恰好持有 `ids + produced[:-1]`
# (produced[-1] 为 pending 未入 cache)。但末步多 token 跨 max_tokens 会 over-commit:
# cache 领先于 produced,无法用已知 token 精确表述——此时禁用复用,下轮全量重建,绝不错算。
resident = stats.get("resident_tokens")
expected = len(ids) + len(produced) - 1
if resident == expected:
self._main_cache = main_cache
self._cached_ids = list(ids) + list(produced[:-1])
else:
self._main_cache = None
self._cached_ids = []
out_ids = _truncate_eos(produced, eos)
text = tok.decode(out_ids)
tps = len(out_ids) / dt if dt > 0 else 0.0
return GenResult(text, len(out_ids), tps, stopped=stopped["v"])

View File

@ -0,0 +1,3 @@
"""sparkle 顶栏 logo。保持单行,便于与模型信息拼在同一行。"""
LOGO = "▚ sparkle"

View File

@ -0,0 +1,41 @@
Screen {
background: #0b0e14;
}
#top {
height: auto;
padding: 1 2;
background: #11151f;
color: #e6e6e6;
border-bottom: solid #2a2f3a;
}
#chat {
padding: 1 2;
height: 1fr;
}
.msg {
margin: 1 0;
padding: 0 1;
height: auto;
}
.msg.assistant {
border-left: solid #2dd4bf;
padding-left: 2;
}
#prompt {
margin: 0 2;
border: round #2dd4bf;
border-title-color: #8b93a7;
background: #11151f;
}
#status {
height: 1;
padding: 0 2;
background: #11151f;
color: #8b93a7;
}

49
models/README.md Normal file
View File

@ -0,0 +1,49 @@
# 模型与专家数据(不入库,请本机下载/生成)
## 为什么要做
目标不是「把 80B 常驻塞进 64G」而是用本机 **64G 统一内存** 做压力验证:把约 **77G 专家权重** 以流式专家池压到约 **33G 常驻**,在 8bit 下跑通 Qwen3-Next-80B-A3B。
验证点对应下游的 **黄金拐点**Qwen3 **30B 8bit A3B** 一类模型,用同一套流式装载思路,能放进 **32G / 16G** 统一内存机器,相对「整模进内存」可 **节省 50% 以上统一内存**。本仓库的 80B @ 64G 是这条路径的上限探针;量产形态落在更小机型上的 30B 拐点模型。
---
## 运行依赖版本
权重本身不绑版本,但本仓库预编译扩展与锁文件对齐如下(换版本可能无法 `import native_moe_ext` 或加载 8bit
| 项 | 要求 / 已验证 |
|----|----------------|
| Python | **3.14.***(已验证 3.14.6;须匹配 `native_moe_ext.cpython-314-darwin.so` |
| mlx | `>=0.31`(锁文件 **0.31.2** |
| mlx-lm | `>=0.31`(锁文件 **0.31.3** |
| numpy | `>=2.0`(锁文件 **2.4.6** |
| fastapi | `>=0.115`(锁文件 **0.140.0** |
| uvicorn | `>=0.30`(锁文件 **0.51.0** |
| textual | `>=0.80`(锁文件 **8.2.8** |
安装:在仓库根目录 `uv sync`(或使用已有 `.venv`)。完整约束见根目录 `pyproject.toml` / `uv.lock`。换 Python 小版本需 `make -C native/ext native_moe_ext` 重编 `.so`
权重体积约 **160GB**,不进 Git。按下面准备后目录应类似
```
models/
├── Qwen3-Next-80B-A3B-Instruct-MLX-8bit/ # ~79GB 主模型
├── experts_8bit_g64/ # ~77GB 专家 blob运行时常驻约 33G
└── qn_mtp_weights.safetensors # ~3.1GB MTP
```
## 下载NAS
全部模型与专家数据从 NAS 共享下载:
- 链接https://ug.link/dxp4800plus-8761/filemgr/share-download/?id=86b004b90f4a477e84c7f65c87943b2f
- 密码:`sparkle`
解压/放置到本仓库 `models/` 下,目录结构与上文一致即可。
> 说明:本仓库当前跑通的是 **80B** 探针;包内若含 **Qwen3-30B-A3B-Instruct-2507-MLX-8bit**,供下游 16G/32G 机型对照,不替代本机 80B 数据准备。
## 可选
`prefix_snapshots/` 会在首次跑通后自动生成,可不预置。

32
native/ext/CMakeLists.txt Normal file
View File

@ -0,0 +1,32 @@
cmake_minimum_required(VERSION 3.22)
project(native_moe_mlx_ext LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
find_package(Python COMPONENTS Interpreter Development.Module REQUIRED)
find_package(nanobind CONFIG REQUIRED)
find_package(MLX CONFIG REQUIRED)
nanobind_add_module(
native_moe_ext
NB_STATIC
STABLE_ABI
LTO
NOMINSIZE
NB_DOMAIN mlx
compute/fused_moe.cpp
io/blob_load.cpp
io/bg_reader.cpp
prefetch/prefetch.cpp
pool/owned_pool.cpp
pool/side_region.cpp
pool/demand.cpp
bindings.cpp
)
target_link_libraries(native_moe_ext PRIVATE mlx)
if(BUILD_SHARED_LIBS)
target_link_options(native_moe_ext PRIVATE -Wl,-rpath,@loader_path)
endif()

14
native/ext/Makefile Normal file
View File

@ -0,0 +1,14 @@
# 生产 MLX 扩展构建cmake 把 native_moe_mlx_ext.cpp 编成
# mlx_streaming/native_moe_ext$(EXT_SUFFIX).sonanobind + MLX Primitive
PYTHON ?= ../../.venv/bin/python
PY_SITE := $(abspath ../../.venv/lib/python3.12/site-packages)
NANOBIND_DIR := $(shell $(PYTHON) -m nanobind --cmake_dir)
MLX_CMAKE_DIR := $(PY_SITE)/mlx/share/cmake/MLX
native_moe_ext:
cmake -S . -B build -DPython_EXECUTABLE=$(abspath $(PYTHON)) -Dnanobind_DIR=$(NANOBIND_DIR) -DMLX_DIR=$(MLX_CMAKE_DIR) -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=$(abspath ../../mlx_streaming)
cmake --build build --target native_moe_ext -j
clean:
rm -rf build
rm -f ../../mlx_streaming/native_moe_ext*.so

153
native/ext/bindings.cpp Normal file
View File

@ -0,0 +1,153 @@
#include "compute/fused_moe.h"
#include "io/blob_load.h"
#include "io/bg_reader.h"
#include "prefetch/prefetch.h"
#include "pool/owned_pool.h"
#include "pool/side_region.h"
#include "pool/demand.h"
using namespace nb::literals;
NB_MODULE(native_moe_ext, m) {
nb::module_::import_("mlx.core");
m.doc() = "Native MLX MoE extension.";
// ---- 融合 MoE 计算核native_fused.cpp----
m.def(
"fused_moe",
&fused_moe,
"x"_a,
"expert_ids"_a,
"scores"_a,
"compute_dir"_a,
"layer"_a,
"hidden"_a,
"inter"_a,
"group"_a,
"bits"_a,
"num_experts"_a,
"synthetic"_a,
nb::kw_only(),
"stream"_a = nb::none());
m.def(
"fused_moe_staged",
&fused_moe_staged,
"x"_a,
"scores"_a,
"gate_w"_a,
"gate_s"_a,
"gate_b"_a,
"up_w"_a,
"up_s"_a,
"up_b"_a,
"down_w"_a,
"down_s"_a,
"down_b"_a,
"hidden"_a,
"inter"_a,
"group"_a,
"bits"_a,
nb::kw_only(),
"stream"_a = nb::none());
// ---- [1] blob 直读 ----
m.def(
"blob_load",
&blob_load,
"path"_a,
"expert_ids"_a,
"stride"_a,
nb::kw_only(),
"stream"_a = nb::none());
// ---- [2] 轻量预取(无 staging仅预热 page cache) ----
m.def(
"prefetch_on_complete",
&prefetch_on_complete,
"expert_ids"_a,
"path"_a,
"stride"_a,
"do_read"_a = true,
nb::kw_only(),
"stream"_a = nb::none());
// ---- [3] staging miss→hit + 完成回调时刻探针(诊断)----
m.def(
"prefetch_into_staging",
&prefetch_into_staging,
"staging"_a,
"expert_ids"_a,
"layer"_a,
"gen"_a,
"path"_a,
"stride"_a,
"resident"_a,
"cap"_a,
"parallel"_a,
nb::kw_only(),
"stream"_a = nb::none());
m.def("prefetch_staging_take", &prefetch_staging_take, "layer"_a);
m.def("staging_hprof_enable", &staging_hprof_enable, "on"_a);
m.def("staging_hprof_now", &staging_hprof_now);
m.def("staging_hprof_get", &staging_hprof_get);
// ---- [4] 段散写侧区缓存zero-copy dual-source 默认路径)----
m.def("prefetch_pool_sideregion", &prefetch_pool_sideregion,
"pool_list"_a, "seg_nbytes"_a, "expert_ids"_a, "layer"_a, "path"_a, "stride"_a,
"resident"_a, "spec_slots"_a, "base_row"_a, "gen"_a = 0, nb::kw_only(),
"stream"_a = nb::none());
m.def("sideregion_contents", &sideregion_contents, "layer"_a, "gen"_a = 0);
m.def("sideregion_kv", &sideregion_kv, "layer"_a, "gen"_a = 0);
m.def("sideregion_reset", &sideregion_reset);
m.def("sideregion_drain", &sideregion_drain,
nb::call_guard<nb::gil_scoped_release>()); // 等待时释放 GIL
// ---- [5] Route 3 owned 池底座C++ 拥有 buffer + 直写,删 MLX scatter----
m.def("pool_owned_zeros", &pool_owned_zeros, "shape"_a, "dtype"_a);
m.def("pool_write_rows", &pool_write_rows, "pool_list"_a, "srcs_flat"_a, "slots"_a);
m.def("pool_write_stacked", &pool_write_stacked, "pool_list"_a, "stacked_list"_a, "slots"_a);
m.def("array_data_ptr", &array_data_ptr, "a"_a);
// ---- [6] 方案B 真实区槽状态 C++ 接管 + demand_dual ----
m.def("real_init", &real_init, "layer"_a, "cap"_a);
m.def("real_region_contents", &real_region_contents, "layer"_a);
m.def("real_region_count", &real_region_count, "layer"_a);
m.def("real_reset", &real_reset);
m.def("real_pin", &real_pin, "layer"_a, "experts"_a);
m.def("real_pinned_count", &real_pinned_count, "layer"_a);
m.def("real_freq_dump", &real_freq_dump);
m.def("demand_dual", &demand_dual, "inds"_a, "pool_list"_a, "seg_nbytes"_a, "layer"_a,
"side_gen"_a, "path"_a, "stride"_a, "cap"_a, "lfu"_a, "decay_interval"_a, nb::kw_only(),
"stream"_a = nb::none());
m.def("demand_last_stats", &demand_last_stats);
m.def("demand_timings", &demand_timings);
m.def("demand_timing_enable", &demand_timing_enable, "on"_a);
m.def("real_debug_place", &real_debug_place, "layer"_a, "experts_flat"_a, "cap"_a, "lfu"_a,
"decay_interval"_a);
// ---- [7] 自由后台读线程 ----
m.def("bg_reader_start", &bg_reader_start, "workers"_a = 1, "low_cap"_a = 0);
m.def("bg_reader_submit", &bg_reader_submit,
"dst"_a, "experts"_a, "rows"_a, "path"_a, "stride"_a, "ticket"_a, "prio"_a = 0);
m.def("bg_reader_ready", &bg_reader_ready, "ticket"_a);
m.def("bg_reader_wait", &bg_reader_wait, "ticket"_a,
nb::call_guard<nb::gil_scoped_release>()); // 阻塞等时释放 GIL
m.def("bg_reader_stop", &bg_reader_stop);
m.def("bg_pread_into_pool", &bg_pread_into_pool,
"dst"_a, "seg_off"_a, "seg_nb"_a, "slot"_a, "expert"_a,
"path"_a, "stride"_a, "ticket"_a, "prio"_a = 0, "nocache"_a = true);
// ---- 融合 MoE 计算核slots 版native_fused.cpp----
m.def(
"fused_moe_slots",
&fused_moe_slots,
"x"_a,
"local_slots"_a,
"scores"_a,
"gate_w"_a,
"gate_s"_a,
"gate_b"_a,
"up_w"_a,
"up_s"_a,
"up_b"_a,
"down_w"_a,
"down_s"_a,
"down_b"_a,
"hidden"_a,
"inter"_a,
"group"_a,
"bits"_a,
nb::kw_only(),
"stream"_a = nb::none());
}

34
native/ext/common.h Normal file
View File

@ -0,0 +1,34 @@
// 共用前导:系统/Metal/MLX/nanobind 头 + 命名空间别名。所有 TU 都包含。
#pragma once
#include <nanobind/nanobind.h>
#include <nanobind/stl/pair.h>
#include <nanobind/stl/string.h>
#include <nanobind/stl/variant.h>
#include <nanobind/stl/vector.h>
#include <Metal/Metal.hpp>
#include <algorithm>
#include <atomic>
#include <cctype>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <fcntl.h>
#include <map>
#include <mutex>
#include <random>
#include <utility>
#include <stdexcept>
#include <string>
#include <sys/mman.h>
#include <sys/stat.h>
#include <unistd.h>
#include <vector>
#include "mlx/backend/metal/device.h"
#include "mlx/mlx.h"
#include "mlx/primitives.h"
namespace nb = nanobind;
namespace mx = mlx::core;

View File

@ -0,0 +1,664 @@
#include "fused_moe.h"
struct FusedParams {
int experts;
int hidden;
int inter;
int group_size;
int bits;
int k;
};
struct MappedProj {
void* w{MAP_FAILED};
void* s{MAP_FAILED};
void* b{MAP_FAILED};
size_t wn{0}, sn{0}, bn{0};
int wfd{-1}, sfd{-1}, bfd{-1};
};
struct ProjectionShape {
int out_dim;
int in_dim;
size_t weight_per;
size_t scale_per;
};
static size_t file_size(int fd) {
struct stat st {};
if (fstat(fd, &st) != 0) {
throw std::runtime_error("fstat failed");
}
return static_cast<size_t>(st.st_size);
}
static void* map_file(const std::string& path, size_t& nbytes, int& fd) {
fd = open(path.c_str(), O_RDONLY);
if (fd < 0) {
throw std::runtime_error("failed to open " + path);
}
nbytes = file_size(fd);
void* p = mmap(nullptr, nbytes, PROT_READ, MAP_PRIVATE, fd, 0);
if (p == MAP_FAILED) {
throw std::runtime_error("failed to mmap " + path);
}
return p;
}
static MappedProj map_proj(const std::string& dir, int layer, const std::string& proj) {
MappedProj m;
std::string layer_name = "layer" + std::string(layer < 10 ? "0" : "") + std::to_string(layer);
std::string base = dir + "/" + layer_name + "." + proj;
m.w = map_file(base + ".weight.bin", m.wn, m.wfd);
m.s = map_file(base + ".scales.bin", m.sn, m.sfd);
m.b = map_file(base + ".biases.bin", m.bn, m.bfd);
return m;
}
static void unmap_proj(MappedProj& m) {
if (m.w != MAP_FAILED) munmap(m.w, m.wn);
if (m.s != MAP_FAILED) munmap(m.s, m.sn);
if (m.b != MAP_FAILED) munmap(m.b, m.bn);
if (m.wfd >= 0) close(m.wfd);
if (m.sfd >= 0) close(m.sfd);
if (m.bfd >= 0) close(m.bfd);
}
static ProjectionShape proj_shape(int out_dim, int in_dim, int group, int bits) {
int words = (in_dim * bits) / 32;
int groups = in_dim / group;
return ProjectionShape{
out_dim,
in_dim,
static_cast<size_t>(out_dim) * words * sizeof(uint32_t),
static_cast<size_t>(out_dim) * groups * sizeof(uint16_t),
};
}
static uint16_t fp32_to_bf16(float x) {
uint32_t u = 0;
std::memcpy(&u, &x, sizeof(u));
return static_cast<uint16_t>(u >> 16);
}
static void fill_synthetic(
MTL::Buffer* wbuf,
MTL::Buffer* sbuf,
MTL::Buffer* bbuf,
uint32_t seed) {
std::mt19937 rng(seed);
auto* w = reinterpret_cast<uint32_t*>(wbuf->contents());
size_t wn = wbuf->length() / sizeof(uint32_t);
for (size_t i = 0; i < wn; ++i) w[i] = rng();
auto* s = reinterpret_cast<uint16_t*>(sbuf->contents());
auto* b = reinterpret_cast<uint16_t*>(bbuf->contents());
size_t sn = sbuf->length() / sizeof(uint16_t);
uint16_t scale = fp32_to_bf16(0.001f);
uint16_t bias = fp32_to_bf16(-0.032f);
for (size_t i = 0; i < sn; ++i) {
s[i] = scale;
b[i] = bias;
}
}
static void copy_projection(
const MappedProj& src,
const ProjectionShape& shape,
const std::vector<int>& ids,
MTL::Buffer* wbuf,
MTL::Buffer* sbuf,
MTL::Buffer* bbuf) {
auto* wd = static_cast<char*>(wbuf->contents());
auto* sd = static_cast<char*>(sbuf->contents());
auto* bd = static_cast<char*>(bbuf->contents());
auto* ws = static_cast<const char*>(src.w);
auto* ss = static_cast<const char*>(src.s);
auto* bs = static_cast<const char*>(src.b);
for (size_t local = 0; local < ids.size(); ++local) {
size_t expert = static_cast<size_t>(ids[local]);
std::memcpy(wd + local * shape.weight_per, ws + expert * shape.weight_per, shape.weight_per);
std::memcpy(sd + local * shape.scale_per, ss + expert * shape.scale_per, shape.scale_per);
std::memcpy(bd + local * shape.scale_per, bs + expert * shape.scale_per, shape.scale_per);
}
}
static std::string metal_source() {
return R"(
#include <metal_stdlib>
using namespace metal;
struct FusedParams { int experts; int hidden; int inter; int group_size; int bits; int k; };
inline float bf16_to_float(ushort v) { uint u = uint(v) << 16; return as_type<float>(u); }
inline float qvalue(const device uint* w, const device ushort* s, const device ushort* b,
uint expert, uint row, uint col, uint out_dim, uint in_dim,
uint group_size, uint bits) {
uint words_per_row = (in_dim * bits) / 32;
uint groups_per_row = in_dim / group_size;
uint bit_offset = col * bits;
uint word_idx = bit_offset / 32;
uint shift = bit_offset % 32;
uint base = (expert * out_dim + row) * words_per_row;
uint q = w[base + word_idx] >> shift;
if (shift + bits > 32) q |= w[base + word_idx + 1] << (32 - shift);
q &= (1u << bits) - 1u;
uint sb = (expert * out_dim + row) * groups_per_row + col / group_size;
return float(q) * bf16_to_float(s[sb]) + bf16_to_float(b[sb]);
}
kernel void synthetic_moe(
const device float* x [[buffer(0)]],
device float* y [[buffer(1)]],
const device float* scores [[buffer(2)]],
constant FusedParams& p [[buffer(3)]],
uint gid [[thread_position_in_grid]]) {
uint token = gid / p.hidden;
uint col = gid % p.hidden;
float acc = 0.0f;
for (int j = 0; j < p.k; ++j) {
acc += scores[token * p.k + j];
}
y[gid] = x[token * p.hidden + col] * acc;
}
kernel void fused_moe_pairs(
const device float* x [[buffer(0)]],
const device uint* gate_w [[buffer(1)]], const device ushort* gate_s [[buffer(2)]], const device ushort* gate_b [[buffer(3)]],
const device uint* up_w [[buffer(4)]], const device ushort* up_s [[buffer(5)]], const device ushort* up_b [[buffer(6)]],
const device uint* down_w [[buffer(7)]], const device ushort* down_s [[buffer(8)]], const device ushort* down_b [[buffer(9)]],
device float* y [[buffer(10)]], constant FusedParams& p [[buffer(11)]],
uint tid [[thread_position_in_threadgroup]], uint gid [[thread_position_in_grid]]) {
constexpr uint block_size = 512;
constexpr uint lanes_per_row = 16;
constexpr uint rows_per_step = block_size / lanes_per_row;
uint expert = gid / block_size;
if (expert >= uint(p.experts)) return;
uint token = expert / p.k;
uint local_row = tid / lanes_per_row;
uint row_lane = tid % lanes_per_row;
threadgroup float act[1024];
threadgroup float gate_part[block_size];
threadgroup float up_part[block_size];
for (uint row_base = 0; row_base < uint(p.inter); row_base += rows_per_step) {
uint row = row_base + local_row;
float gate_acc = 0.0f, up_acc = 0.0f;
if (row < uint(p.inter)) {
for (uint col = row_lane; col < uint(p.hidden); col += lanes_per_row) {
float xv = x[token * p.hidden + col];
gate_acc += qvalue(gate_w, gate_s, gate_b, expert, row, col, p.inter, p.hidden, p.group_size, p.bits) * xv;
up_acc += qvalue(up_w, up_s, up_b, expert, row, col, p.inter, p.hidden, p.group_size, p.bits) * xv;
}
}
gate_part[tid] = gate_acc;
up_part[tid] = up_acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) { gate_part[tid] += gate_part[tid + stride]; up_part[tid] += up_part[tid + stride]; }
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < uint(p.inter)) {
float g = gate_part[tid];
act[row] = g * (1.0f / (1.0f + exp(-g))) * up_part[tid];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
for (uint row_base = 0; row_base < uint(p.hidden); row_base += rows_per_step) {
uint row = row_base + local_row;
float acc = 0.0f;
if (row < uint(p.hidden)) {
for (uint col = row_lane; col < uint(p.inter); col += lanes_per_row) {
acc += qvalue(down_w, down_s, down_b, expert, row, col, p.hidden, p.inter, p.group_size, p.bits) * act[col];
}
}
gate_part[tid] = acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) gate_part[tid] += gate_part[tid + stride];
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < uint(p.hidden)) {
y[expert * p.hidden + row] = gate_part[tid];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
}
kernel void fused_moe_slots_pairs(
const device float* x [[buffer(0)]],
const device uint* gate_w [[buffer(1)]], const device ushort* gate_s [[buffer(2)]], const device ushort* gate_b [[buffer(3)]],
const device uint* up_w [[buffer(4)]], const device ushort* up_s [[buffer(5)]], const device ushort* up_b [[buffer(6)]],
const device uint* down_w [[buffer(7)]], const device ushort* down_s [[buffer(8)]], const device ushort* down_b [[buffer(9)]],
device float* y [[buffer(10)]], constant FusedParams& p [[buffer(11)]],
const device uint* local_slots [[buffer(12)]],
uint tid [[thread_position_in_threadgroup]], uint gid [[thread_position_in_grid]]) {
constexpr uint block_size = 512;
constexpr uint lanes_per_row = 16;
constexpr uint rows_per_step = block_size / lanes_per_row;
uint active_idx = gid / block_size;
if (active_idx >= uint(p.experts)) return;
uint slot = local_slots[active_idx];
uint token = active_idx / p.k;
uint local_row = tid / lanes_per_row;
uint row_lane = tid % lanes_per_row;
threadgroup float act[1024];
threadgroup float gate_part[block_size];
threadgroup float up_part[block_size];
for (uint row_base = 0; row_base < uint(p.inter); row_base += rows_per_step) {
uint row = row_base + local_row;
float gate_acc = 0.0f, up_acc = 0.0f;
if (row < uint(p.inter)) {
for (uint col = row_lane; col < uint(p.hidden); col += lanes_per_row) {
float xv = x[token * p.hidden + col];
gate_acc += qvalue(gate_w, gate_s, gate_b, slot, row, col, p.inter, p.hidden, p.group_size, p.bits) * xv;
up_acc += qvalue(up_w, up_s, up_b, slot, row, col, p.inter, p.hidden, p.group_size, p.bits) * xv;
}
}
gate_part[tid] = gate_acc;
up_part[tid] = up_acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) { gate_part[tid] += gate_part[tid + stride]; up_part[tid] += up_part[tid + stride]; }
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < uint(p.inter)) {
float g = gate_part[tid];
act[row] = g * (1.0f / (1.0f + exp(-g))) * up_part[tid];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
for (uint row_base = 0; row_base < uint(p.hidden); row_base += rows_per_step) {
uint row = row_base + local_row;
float acc = 0.0f;
if (row < uint(p.hidden)) {
for (uint col = row_lane; col < uint(p.inter); col += lanes_per_row) {
acc += qvalue(down_w, down_s, down_b, slot, row, col, p.hidden, p.inter, p.group_size, p.bits) * act[col];
}
}
gate_part[tid] = acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = lanes_per_row / 2; stride > 0; stride >>= 1) {
if (row_lane < stride) gate_part[tid] += gate_part[tid + stride];
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (row_lane == 0 && row < uint(p.hidden)) {
y[active_idx * p.hidden + row] = gate_part[tid];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
}
kernel void reduce_pairs(
const device float* pair_y [[buffer(0)]],
const device float* scores [[buffer(1)]],
device float* out [[buffer(2)]],
constant FusedParams& p [[buffer(3)]],
uint gid [[thread_position_in_grid]]) {
uint token = gid / p.hidden;
uint col = gid % p.hidden;
float acc = 0.0f;
for (int j = 0; j < p.k; ++j) {
uint pair = token * p.k + uint(j);
acc += pair_y[pair * p.hidden + col] * scores[pair];
}
out[gid] = acc;
}
)";
}
class FusedMoePrimitive : public mx::Primitive {
public:
FusedMoePrimitive(
mx::Stream stream,
std::string compute_dir,
int layer,
int hidden,
int inter,
int group,
int bits,
int num_experts,
bool synthetic,
std::vector<int> expert_ids)
: Primitive(stream),
compute_dir_(std::move(compute_dir)),
layer_(layer),
hidden_(hidden),
inter_(inter),
group_(group),
bits_(bits),
num_experts_(num_experts),
synthetic_(synthetic),
expert_ids_(std::move(expert_ids)) {
if (!synthetic_) {
gate_ = map_proj(compute_dir_, layer_, "gate_proj");
up_ = map_proj(compute_dir_, layer_, "up_proj");
down_ = map_proj(compute_dir_, layer_, "down_proj");
}
}
~FusedMoePrimitive() override {
if (!synthetic_) {
unmap_proj(gate_);
unmap_proj(up_);
unmap_proj(down_);
}
}
const char* name() const override { return "FusedMoePrimitive"; }
void eval_cpu(const std::vector<mx::array>&, std::vector<mx::array>&) override {
throw std::runtime_error("FusedMoePrimitive only supports GPU evaluation");
}
void eval_gpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
auto& x = inputs[0];
auto& scores = inputs[1];
auto& out = outputs[0];
int active = static_cast<int>(expert_ids_.size());
int tokens = static_cast<int>(x.size() / hidden_);
int k = active / std::max(1, tokens);
if (active <= 0 || tokens <= 0 || active != static_cast<int>(scores.size())) {
throw std::runtime_error("FusedMoePrimitive shape mismatch");
}
out.set_data(mx::allocator::malloc(out.nbytes()));
auto& d = mx::metal::device(stream().device);
auto& enc = mx::metal::get_command_encoder(stream());
auto lib = d.get_library("native_moe_mlx_ext", []() { return metal_source(); });
if (synthetic_) {
auto synthetic_kernel = d.get_kernel("synthetic_moe", lib);
enc.set_compute_pipeline_state(synthetic_kernel);
enc.set_input_array(x, 0);
enc.set_output_array(out, 1);
enc.set_input_array(scores, 2);
FusedParams params{active, hidden_, inter_, group_, bits_, k};
enc.set_bytes(params, 3);
enc.dispatch_threads(MTL::Size(out.size(), 1, 1), MTL::Size(std::min<size_t>(out.size(), 256), 1, 1));
return;
}
auto fused = d.get_kernel("fused_moe_pairs", lib);
auto reduce = d.get_kernel("reduce_pairs", lib);
ProjectionShape gu = proj_shape(inter_, hidden_, group_, bits_);
ProjectionShape down = proj_shape(hidden_, inter_, group_, bits_);
size_t gu_w_bytes = static_cast<size_t>(active) * gu.weight_per;
size_t gu_s_bytes = static_cast<size_t>(active) * gu.scale_per;
size_t down_w_bytes = static_cast<size_t>(active) * down.weight_per;
size_t down_s_bytes = static_cast<size_t>(active) * down.scale_per;
buffers_.clear();
auto make_buffer = [&](size_t nbytes) -> MTL::Buffer* {
buffers_.push_back(NS::TransferPtr(d.mtl_device()->newBuffer(nbytes, MTL::ResourceStorageModeShared)));
return buffers_.back().get();
};
MTL::Buffer* gw = make_buffer(gu_w_bytes);
MTL::Buffer* gs = make_buffer(gu_s_bytes);
MTL::Buffer* gb = make_buffer(gu_s_bytes);
MTL::Buffer* uw = make_buffer(gu_w_bytes);
MTL::Buffer* us = make_buffer(gu_s_bytes);
MTL::Buffer* ub = make_buffer(gu_s_bytes);
MTL::Buffer* dw = make_buffer(down_w_bytes);
MTL::Buffer* ds = make_buffer(down_s_bytes);
MTL::Buffer* db = make_buffer(down_s_bytes);
if (synthetic_) {
fill_synthetic(gw, gs, gb, 1);
fill_synthetic(uw, us, ub, 2);
fill_synthetic(dw, ds, db, 3);
} else {
copy_projection(gate_, gu, expert_ids_, gw, gs, gb);
copy_projection(up_, gu, expert_ids_, uw, us, ub);
copy_projection(down_, down, expert_ids_, dw, ds, db);
}
auto pair_out = mx::zeros(mx::Shape{active, hidden_}, mx::float32, stream());
pair_out.set_data(mx::allocator::malloc(pair_out.nbytes()));
enc.set_compute_pipeline_state(fused);
enc.set_input_array(x, 0);
enc.set_buffer(gw, 1);
enc.set_buffer(gs, 2);
enc.set_buffer(gb, 3);
enc.set_buffer(uw, 4);
enc.set_buffer(us, 5);
enc.set_buffer(ub, 6);
enc.set_buffer(dw, 7);
enc.set_buffer(ds, 8);
enc.set_buffer(db, 9);
enc.set_output_array(out, 10);
FusedParams params{active, hidden_, inter_, group_, bits_, k};
enc.set_bytes(params, 11);
enc.set_input_array(scores, 12);
enc.add_temporary(x);
enc.add_temporary(scores);
enc.dispatch_threads(MTL::Size(static_cast<size_t>(active) * 512, 1, 1), MTL::Size(512, 1, 1));
}
private:
std::string compute_dir_;
int layer_;
int hidden_;
int inter_;
int group_;
int bits_;
int num_experts_;
bool synthetic_;
std::vector<int> expert_ids_;
MappedProj gate_;
MappedProj up_;
MappedProj down_;
std::vector<NS::SharedPtr<MTL::Buffer>> buffers_;
};
class StagedFusedMoePrimitive : public mx::Primitive {
public:
StagedFusedMoePrimitive(mx::Stream stream, int hidden, int inter, int group, int bits)
: Primitive(stream), hidden_(hidden), inter_(inter), group_(group), bits_(bits) {}
const char* name() const override { return "StagedFusedMoePrimitive"; }
void eval_cpu(const std::vector<mx::array>&, std::vector<mx::array>&) override {
throw std::runtime_error("StagedFusedMoePrimitive only supports GPU evaluation");
}
void eval_gpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
auto& x = inputs[0];
auto& scores = inputs[1];
auto& out = outputs[0];
int tokens = static_cast<int>(x.size() / hidden_);
int active = static_cast<int>(scores.size());
int k = active / std::max(1, tokens);
if (active <= 0 || tokens <= 0 || active != static_cast<int>(scores.size())) {
throw std::runtime_error("StagedFusedMoePrimitive shape mismatch");
}
out.set_data(mx::allocator::malloc(out.nbytes()));
auto& d = mx::metal::device(stream().device);
auto& enc = mx::metal::get_command_encoder(stream());
auto lib = d.get_library("native_moe_mlx_ext", []() { return metal_source(); });
auto fused = d.get_kernel("fused_moe_pairs", lib);
auto reduce = d.get_kernel("reduce_pairs", lib);
auto pair_out = mx::zeros(mx::Shape{active, hidden_}, mx::float32, stream());
pair_out.set_data(mx::allocator::malloc(pair_out.nbytes()));
enc.set_compute_pipeline_state(fused);
enc.set_input_array(x, 0);
enc.set_input_array(inputs[2], 1);
enc.set_input_array(inputs[3], 2);
enc.set_input_array(inputs[4], 3);
enc.set_input_array(inputs[5], 4);
enc.set_input_array(inputs[6], 5);
enc.set_input_array(inputs[7], 6);
enc.set_input_array(inputs[8], 7);
enc.set_input_array(inputs[9], 8);
enc.set_input_array(inputs[10], 9);
enc.set_output_array(pair_out, 10);
FusedParams params{active, hidden_, inter_, group_, bits_, k};
enc.set_bytes(params, 11);
enc.dispatch_threads(MTL::Size(static_cast<size_t>(active) * 512, 1, 1), MTL::Size(512, 1, 1));
enc.set_compute_pipeline_state(reduce);
enc.set_input_array(pair_out, 0);
enc.set_input_array(scores, 1);
enc.set_output_array(out, 2);
enc.set_bytes(params, 3);
enc.add_temporary(pair_out);
enc.dispatch_threads(MTL::Size(out.size(), 1, 1), MTL::Size(std::min<size_t>(out.size(), 256), 1, 1));
}
private:
int hidden_;
int inter_;
int group_;
int bits_;
};
class SlotFusedMoePrimitive : public mx::Primitive {
public:
SlotFusedMoePrimitive(mx::Stream stream, int hidden, int inter, int group, int bits)
: Primitive(stream), hidden_(hidden), inter_(inter), group_(group), bits_(bits) {}
const char* name() const override { return "SlotFusedMoePrimitive"; }
void eval_cpu(const std::vector<mx::array>&, std::vector<mx::array>&) override {
throw std::runtime_error("SlotFusedMoePrimitive only supports GPU evaluation");
}
void eval_gpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
auto& x = inputs[0];
auto& local_slots = inputs[1];
auto& scores = inputs[2];
auto& out = outputs[0];
int tokens = static_cast<int>(x.size() / hidden_);
int active = static_cast<int>(scores.size());
int k = active / std::max(1, tokens);
if (active <= 0 || tokens <= 0 || active != static_cast<int>(local_slots.size())) {
throw std::runtime_error("SlotFusedMoePrimitive shape mismatch");
}
out.set_data(mx::allocator::malloc(out.nbytes()));
auto& d = mx::metal::device(stream().device);
auto& enc = mx::metal::get_command_encoder(stream());
auto lib = d.get_library("native_moe_mlx_ext", []() { return metal_source(); });
auto fused = d.get_kernel("fused_moe_slots_pairs", lib);
auto reduce = d.get_kernel("reduce_pairs", lib);
auto pair_out = mx::zeros(mx::Shape{active, hidden_}, mx::float32, stream());
pair_out.set_data(mx::allocator::malloc(pair_out.nbytes()));
enc.set_compute_pipeline_state(fused);
enc.set_input_array(x, 0);
enc.set_input_array(inputs[3], 1);
enc.set_input_array(inputs[4], 2);
enc.set_input_array(inputs[5], 3);
enc.set_input_array(inputs[6], 4);
enc.set_input_array(inputs[7], 5);
enc.set_input_array(inputs[8], 6);
enc.set_input_array(inputs[9], 7);
enc.set_input_array(inputs[10], 8);
enc.set_input_array(inputs[11], 9);
enc.set_output_array(pair_out, 10);
FusedParams params{active, hidden_, inter_, group_, bits_, k};
enc.set_bytes(params, 11);
enc.set_input_array(local_slots, 12);
enc.dispatch_threads(MTL::Size(static_cast<size_t>(active) * 512, 1, 1), MTL::Size(512, 1, 1));
enc.set_compute_pipeline_state(reduce);
enc.set_input_array(pair_out, 0);
enc.set_input_array(scores, 1);
enc.set_output_array(out, 2);
enc.set_bytes(params, 3);
enc.add_temporary(pair_out);
enc.dispatch_threads(MTL::Size(out.size(), 1, 1), MTL::Size(std::min<size_t>(out.size(), 256), 1, 1));
}
private:
int hidden_;
int inter_;
int group_;
int bits_;
};
mx::array fused_moe(
const mx::array& x,
const mx::array& expert_ids,
const mx::array& scores,
const std::string& compute_dir,
int layer,
int hidden,
int inter,
int group,
int bits,
int num_experts,
bool synthetic,
mx::StreamOrDevice s = {}) {
if (!synthetic) {
throw std::runtime_error(
"real mmap staging in MLX Primitive needs MLX-managed staging buffers; use synthetic for bridge tests");
}
auto ids = expert_ids;
ids.eval();
std::vector<int> expert_vec(ids.size());
const uint32_t* idp = ids.data<uint32_t>();
for (size_t i = 0; i < ids.size(); ++i) expert_vec[i] = static_cast<int>(idp[i]);
auto out_shape = x.shape();
out_shape.back() = hidden;
return mx::array(
out_shape,
mx::float32,
std::make_shared<FusedMoePrimitive>(
mx::to_stream(s),
compute_dir,
layer,
hidden,
inter,
group,
bits,
num_experts,
synthetic,
std::move(expert_vec)),
std::vector<mx::array>{x, scores});
}
mx::array fused_moe_staged(
const mx::array& x,
const mx::array& scores,
const mx::array& gate_w,
const mx::array& gate_s,
const mx::array& gate_b,
const mx::array& up_w,
const mx::array& up_s,
const mx::array& up_b,
const mx::array& down_w,
const mx::array& down_s,
const mx::array& down_b,
int hidden,
int inter,
int group,
int bits,
mx::StreamOrDevice s = {}) {
auto out_shape = x.shape();
out_shape.back() = hidden;
return mx::array(
out_shape,
mx::float32,
std::make_shared<StagedFusedMoePrimitive>(mx::to_stream(s), hidden, inter, group, bits),
std::vector<mx::array>{
x, scores,
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b});
}
mx::array fused_moe_slots(
const mx::array& x,
const mx::array& local_slots,
const mx::array& scores,
const mx::array& gate_w,
const mx::array& gate_s,
const mx::array& gate_b,
const mx::array& up_w,
const mx::array& up_s,
const mx::array& up_b,
const mx::array& down_w,
const mx::array& down_s,
const mx::array& down_b,
int hidden,
int inter,
int group,
int bits,
mx::StreamOrDevice s = {}) {
auto out_shape = x.shape();
out_shape.back() = hidden;
return mx::array(
out_shape,
mx::float32,
std::make_shared<SlotFusedMoePrimitive>(mx::to_stream(s), hidden, inter, group, bits),
std::vector<mx::array>{
x, local_slots, scores,
gate_w, gate_s, gate_b,
up_w, up_s, up_b,
down_w, down_s, down_b});
}

View File

@ -0,0 +1,23 @@
// fused MoE 计算:合成/真实 mmap staging、staged(预切片权重)、slot(local slot 索引)三条工厂。
// 实现见 fused_moe.cpp这里只暴露给 bindings 的自由函数声明(默认参数放在声明处)。
#pragma once
#include "../common.h"
mx::array fused_moe(
const mx::array& x, const mx::array& expert_ids, const mx::array& scores,
const std::string& compute_dir, int layer, int hidden, int inter, int group,
int bits, int num_experts, bool synthetic, mx::StreamOrDevice s);
mx::array fused_moe_staged(
const mx::array& x, const mx::array& scores,
const mx::array& gate_w, const mx::array& gate_s, const mx::array& gate_b,
const mx::array& up_w, const mx::array& up_s, const mx::array& up_b,
const mx::array& down_w, const mx::array& down_s, const mx::array& down_b,
int hidden, int inter, int group, int bits, mx::StreamOrDevice s);
mx::array fused_moe_slots(
const mx::array& x, const mx::array& local_slots, const mx::array& scores,
const mx::array& gate_w, const mx::array& gate_s, const mx::array& gate_b,
const mx::array& up_w, const mx::array& up_s, const mx::array& up_b,
const mx::array& down_w, const mx::array& down_s, const mx::array& down_b,
int hidden, int inter, int group, int bits, mx::StreamOrDevice s);

183
native/ext/io/bg_reader.cpp Normal file
View File

@ -0,0 +1,183 @@
// [7] 自由后台读线程de-risk双队列 + 低优并发上限的后台 pread 线程池。
#include "bg_reader.h"
#include "blob_io.h"
#include <condition_variable>
#include <functional>
#include <queue>
#include <thread>
#include <unordered_map>
#include <unordered_set>
#include <fcntl.h>
#include <unistd.h>
namespace {
struct ReadOp { uint8_t* dst; size_t nbytes; off_t file_off; };
struct BgJob {
std::vector<ReadOp> ops;
std::string path;
long ticket;
std::vector<mx::array> keep; // 持 buffer 引用,保证 dst 指针在读期间存活
std::function<void()> task; // 若设置worker 直接执行它(侧区异步读用),不走 ops/ticket
int prio = 0; // >0=高优(route 读)=0=低优(投机兜底)
bool nocache = true; // demandroute读设 false 走 page cache段偏移非页对齐
// F_NOCACHE 下多 MB 非对齐 pread 会短读 → 池槽脏字节(实测)。
};
// 双队列 + 低优并发上限:高优(route)读永远优先取、且独占大多数 worker
// 低优(投机)读最多占 low_cap_ 个 worker → 给 route 读留出 worker 和 SSD 带宽,
// 实现"高分抢带宽先就绪、低分延后不阻塞关键路径"的真 IO 优先级。
class BgReader {
public:
~BgReader() { stop(); } // 进程退出兜底:避免 joinable 线程析构触发 std::terminate
void start(int workers, int low_cap) {
{ std::lock_guard<std::mutex> lk(m_); low_cap_ = low_cap; }
ensure_started(workers);
}
// 懒启动:侧区预取的 GPU 回调可能在 Python 没调用 bg_reader_start 时就要派任务,
// 故 submit/submit_task 自动确保线程已起(默认 4 worker调用方无需显式 start。
void ensure_started(int workers) {
std::lock_guard<std::mutex> lk(m_);
if (running_) return;
running_ = true;
for (int i = 0; i < workers; ++i) threads_.emplace_back([this] { loop(); });
}
void submit(BgJob job) {
ensure_started(4);
{ std::lock_guard<std::mutex> lk(m_);
if (job.prio > 0) high_q_.push(std::move(job)); else low_q_.push(std::move(job)); }
cv_.notify_all();
}
void submit_task(std::function<void()> fn) {
ensure_started(4);
BgJob job;
job.task = std::move(fn);
{ std::lock_guard<std::mutex> lk(m_); low_q_.push(std::move(job)); } // 通用任务走低优
cv_.notify_all();
}
bool ready(long ticket) {
std::lock_guard<std::mutex> lk(dm_);
return done_.count(ticket) > 0;
}
void wait(long ticket) {
std::unique_lock<std::mutex> lk(dm_);
dcv_.wait(lk, [&] { return done_.count(ticket) > 0; });
}
void stop() {
{ std::lock_guard<std::mutex> lk(m_); running_ = false; }
cv_.notify_all();
for (auto& t : threads_) if (t.joinable()) t.join();
threads_.clear();
{ std::lock_guard<std::mutex> lk(dm_); done_.clear(); }
{ std::lock_guard<std::mutex> lk(m_); active_low_ = 0; }
}
private:
int low_budget() const { return low_cap_ <= 0 ? 1 << 30 : low_cap_; } // <=0 视为不限流
bool can_take_low() const { return !low_q_.empty() && active_low_ < low_budget(); }
void loop() {
std::unordered_map<std::string, int> fds; // 每线程各自缓存 fd
while (true) {
std::unique_lock<std::mutex> lk(m_);
cv_.wait(lk, [this] { return !high_q_.empty() || can_take_low() || !running_; });
if (!running_ && high_q_.empty() && low_q_.empty()) break;
if (!running_ && high_q_.empty() && !can_take_low()) break; // 退出时低优可超额排空
bool is_low = false;
BgJob job;
if (!high_q_.empty()) { // 高优(route)读永远先取
job = std::move(high_q_.front()); high_q_.pop();
} else if (can_take_low()) { // 低优限流:仅在额度内取
job = std::move(low_q_.front()); low_q_.pop(); ++active_low_; is_low = true;
} else {
continue; // 假醒(低优已满额)→ 回去等
}
lk.unlock();
if (job.task) { job.task(); } // 通用任务(侧区异步读):直接执行,无 ticket
else {
int fd;
std::string key = (job.nocache ? "N:" : "C:") + job.path; // 按 nocache 分别缓存 fd
auto it = fds.find(key);
if (it == fds.end()) {
fd = job.nocache ? open_blob_nocache(job.path.c_str()) : ::open(job.path.c_str(), O_RDONLY);
fds[key] = fd;
} else fd = it->second;
if (fd >= 0)
for (auto& op : job.ops) {
ssize_t got = ::pread(fd, op.dst, op.nbytes, op.file_off);
if (got != static_cast<ssize_t>(op.nbytes))
fprintf(stderr, "[bg pread SHORT] got=%zd want=%zu off=%lld\n",
got, op.nbytes, static_cast<long long>(op.file_off));
}
{ std::lock_guard<std::mutex> lk2(dm_); done_.insert(job.ticket); }
dcv_.notify_all();
}
if (is_low) { // 释放低优额度 → 唤醒别的 worker 再取
std::lock_guard<std::mutex> lk2(m_); --active_low_; cv_.notify_all();
}
}
for (auto& kv : fds) if (kv.second >= 0) ::close(kv.second);
}
std::mutex m_, dm_;
std::condition_variable cv_, dcv_;
std::queue<BgJob> high_q_, low_q_;
int active_low_ = 0; // 当前正在执行的低优读数
int low_cap_ = 0; // 低优并发上限(<=0 不限流,保持旧行为)
std::unordered_set<long> done_;
std::vector<std::thread> threads_;
bool running_ = false;
};
BgReader g_bg;
} // namespace
void bg_submit_task(std::function<void()> fn) { g_bg.submit_task(std::move(fn)); }
void bg_reader_start(int workers, int low_cap) { g_bg.start(workers, low_cap); }
long bg_reader_submit(const mx::array& dst, const std::vector<int>& experts,
const std::vector<int>& rows, const std::string& path,
int stride, long ticket, int prio) {
mx::array d = dst;
d.eval();
uint8_t* base = d.data<uint8_t>();
size_t st = static_cast<size_t>(stride);
BgJob job;
job.path = path;
job.ticket = ticket;
job.prio = prio;
job.keep.push_back(d);
for (size_t i = 0; i < experts.size(); ++i)
job.ops.push_back(ReadOp{base + static_cast<size_t>(rows[i]) * st, st,
static_cast<off_t>(static_cast<size_t>(experts[i]) * st)});
g_bg.submit(std::move(job));
return ticket;
}
// 专家各段直写进多个池段张量的 slot 行(消费侧零 MLX 算子)。
long bg_pread_into_pool(
const std::vector<mx::array>& dst,
const std::vector<long>& seg_off,
const std::vector<long>& seg_nb,
long slot, long expert,
const std::string& path, long stride, long ticket, int prio, bool nocache) {
BgJob job;
job.path = path;
job.ticket = ticket;
job.prio = prio;
job.nocache = nocache;
for (size_t i = 0; i < dst.size(); ++i) {
mx::array d = dst[i];
d.eval();
uint8_t* base = d.data<uint8_t>();
job.keep.push_back(d);
job.ops.push_back(ReadOp{
base + static_cast<size_t>(slot) * static_cast<size_t>(seg_nb[i]),
static_cast<size_t>(seg_nb[i]),
static_cast<off_t>(static_cast<size_t>(expert) * static_cast<size_t>(stride)
+ static_cast<size_t>(seg_off[i]))});
}
g_bg.submit(std::move(job));
return ticket;
}
bool bg_reader_ready(long ticket) { return g_bg.ready(ticket); }
void bg_reader_wait(long ticket) { g_bg.wait(ticket); }
void bg_reader_stop() { g_bg.stop(); }

23
native/ext/io/bg_reader.h Normal file
View File

@ -0,0 +1,23 @@
// 自由后台读线程de-risk脱离 GPU 完成回调、零 GILpread 进调用方 MLX buffer。
// 线程全程只碰 dst 原始指针 / 路径 / 整数,绝不接触 Python 对象 → 无需 GIL。
#pragma once
#include "../common.h"
#include <functional>
void bg_reader_start(int workers, int low_cap = 0);
long bg_reader_submit(const mx::array& dst, const std::vector<int>& experts,
const std::vector<int>& rows, const std::string& path,
int stride, long ticket, int prio = 0);
bool bg_reader_ready(long ticket);
void bg_reader_wait(long ticket);
void bg_reader_stop();
long bg_pread_into_pool(
const std::vector<mx::array>& dst,
const std::vector<long>& seg_off,
const std::vector<long>& seg_nb,
long slot, long expert,
const std::string& path, long stride, long ticket, int prio = 0, bool nocache = true);
// 通用后台任务入口:把任意闭包派到后台线程池(低优队列)。侧区/staging 预取的 GPU 完成回调
// 用它把 pread 派离 Metal 回调线程 → 真正与计算重叠。内部接口,不绑 Python。
void bg_submit_task(std::function<void()> fn);

16
native/ext/io/blob_io.h Normal file
View File

@ -0,0 +1,16 @@
// 共用底层 IO 助手:打开 blob 文件。多个模块(侧区/后台读线程等)的 pread 都用它。
#pragma once
#include <fcntl.h>
#include <unistd.h>
#ifndef F_NOCACHE
#define F_NOCACHE 48 // macOS提示内核读过的页不留 page cache
#endif
// 打开 blob 并设 F_NOCACHE与 demand 侧 blob_loader 对齐,避免预取读把 page cache 灌满 →
// 在内存受限机上累积压力触发"双稳态"慢挡翻转(实测 zerocopy 每轮翻慢挡的根因)。
static inline int open_blob_nocache(const char* path) {
int fd = ::open(path, O_RDONLY);
if (fd >= 0) ::fcntl(fd, F_NOCACHE, 1);
return fd;
}

View File

@ -0,0 +1,58 @@
// [1] blob 直读pread 专家字节进 MLX buffer惰性图节点
// 把一组专家的 blob 字节直接 pread 进 MLX 自有 buffer无 kernel、无额外拷贝
// load 作为惰性图节点:在批量 eval 中执行,避免 Python 侧 per-expert mx.eval 同步。
#include "blob_load.h"
#include <fcntl.h>
#include <unistd.h>
class BlobLoadPrimitive : public mx::Primitive {
public:
BlobLoadPrimitive(mx::Stream stream, std::string path, size_t stride, std::vector<int> experts)
: Primitive(stream), path_(std::move(path)), stride_(stride), experts_(std::move(experts)) {}
const char* name() const override { return "BlobLoadPrimitive"; }
void eval_cpu(const std::vector<mx::array>&, std::vector<mx::array>& outputs) override {
load(outputs[0]);
}
void eval_gpu(const std::vector<mx::array>&, std::vector<mx::array>& outputs) override {
load(outputs[0]);
}
private:
void load(mx::array& out) {
out.set_data(mx::allocator::malloc(out.nbytes()));
uint8_t* dst = out.data<uint8_t>();
int fd = ::open(path_.c_str(), O_RDONLY);
if (fd < 0) throw std::runtime_error("blob open failed: " + path_);
for (size_t i = 0; i < experts_.size(); ++i) {
size_t off = static_cast<size_t>(experts_[i]) * stride_;
ssize_t n = ::pread(fd, dst + i * stride_, stride_, static_cast<off_t>(off));
if (n != static_cast<ssize_t>(stride_)) {
::close(fd);
throw std::runtime_error("blob pread short read");
}
}
::close(fd);
}
std::string path_;
size_t stride_;
std::vector<int> experts_;
};
mx::array blob_load(
const std::string& path,
const mx::array& expert_ids,
int stride,
mx::StreamOrDevice s = {}) {
auto ids = expert_ids;
ids.eval();
std::vector<int> ev(ids.size());
const uint32_t* p = ids.data<uint32_t>();
for (size_t i = 0; i < ids.size(); ++i) ev[i] = static_cast<int>(p[i]);
int n = static_cast<int>(ev.size());
return mx::array(
mx::Shape{n, stride},
mx::uint8,
std::make_shared<BlobLoadPrimitive>(mx::to_stream(s), path, static_cast<size_t>(stride), std::move(ev)),
std::vector<mx::array>{});
}

View File

@ -0,0 +1,7 @@
// blob 直读:把一组专家 blob 字节 pread 进新建的 MLX uint8[n,stride] 数组(惰性图节点)。
#pragma once
#include "../common.h"
mx::array blob_load(
const std::string& path, const mx::array& expert_ids, int stride,
mx::StreamOrDevice s);

331
native/ext/pool/demand.cpp Normal file
View File

@ -0,0 +1,331 @@
// [6] Phase 2 方案B真实区槽状态 C++ 全接管1 次同步版)。
// 复刻 Python ResidentExpertPool 的 _slot_of/_free/_freq + _choose_victim/_alloc_slot 语义:
// - free 优先(front popfree 初始 [0,cap))free 空 → LFU 驱逐,受害者槽直接复用(不回 free)。
// - _choose_victimcandidates=插入序中 ∉current 者victim=min(freq, 候选下标),驱逐不删 freq。
// - dual 语义:真实区命中不移动插入序;仅新放入 miss 追加序尾。
#include "demand.h"
#include "side_region.h"
#include "../io/bg_reader.h"
#include <chrono>
#include <cstdio>
#include <cstdlib>
#include <unordered_map>
#include <unordered_set>
struct RealLayer {
std::vector<int> order; // 插入序(LRU tie-break),与 e2r 同步维护
std::unordered_map<int, int> e2r; // expert -> slot [0,cap)
std::vector<int> free_rows; // 空闲槽(front pop仿 free.pop(0))
std::unordered_map<int, uint32_t> freq; // LFU 频次(驱逐不删,与 Python 一致)
std::unordered_set<int> pinned; // AUTOPIN 钉死集:驱逐永不选(见 real_pin)
int cap = 0;
long access = 0; // 累计访问(decay 用)
bool inited = false;
};
static std::mutex g_real_mutex;
static std::map<int, RealLayer> g_real;
// demand 统计:累计 + 本次(供 Python 更新 rp.hits/misses/gpu_fastpath/gpu_fallback)。
static std::mutex g_dstat_mutex;
// 本次 [hitpos, misspos, loads, fallback01, unplaced]
static long g_d_last[5] = {0, 0, 0, 0, 0};
static std::atomic<long> g_demand_ticket{1000000000}; // demand 并行 pread 用的独立 ticket 段
// 诊断计时(DEMAND_TIMING=1):累计各段主线程微秒,供定位结构性开销。默认关。
static bool g_dt_on = false;
static double g_dt[6] = {0, 0, 0, 0, 0, 0}; // [inds_eval, pool_eval, side_snap, real_lock, core, build]
static inline double dt_now_us() {
return std::chrono::duration<double, std::micro>(
std::chrono::steady_clock::now().time_since_epoch()).count();
}
std::vector<double> demand_timings() { return {g_dt[0], g_dt[1], g_dt[2], g_dt[3], g_dt[4], g_dt[5]}; }
void demand_timing_enable(bool on) {
g_dt_on = on;
for (int i = 0; i < 6; ++i) g_dt[i] = 0;
}
static void real_ensure_locked(RealLayer& c, int cap) {
if (c.inited) return;
c.cap = cap;
c.free_rows.clear();
for (int r = 0; r < cap; ++r) c.free_rows.push_back(r); // free 初始 [0,cap)
c.inited = true;
}
void real_init(int layer, int cap) {
std::lock_guard<std::mutex> lk(g_real_mutex);
real_ensure_locked(g_real[layer], cap);
}
std::vector<int> real_region_contents(int layer) {
std::lock_guard<std::mutex> lk(g_real_mutex);
std::vector<int> out;
auto it = g_real.find(layer);
if (it != g_real.end())
for (auto& p : it->second.e2r) { out.push_back(p.first); out.push_back(p.second); }
return out;
}
int real_region_count(int layer) {
std::lock_guard<std::mutex> lk(g_real_mutex);
auto it = g_real.find(layer);
return it == g_real.end() ? 0 : static_cast<int>(it->second.e2r.size());
}
int real_pinned_count(int layer) {
std::lock_guard<std::mutex> lk(g_real_mutex);
auto it = g_real.find(layer);
return it == g_real.end() ? 0 : static_cast<int>(it->second.pinned.size());
}
void real_reset() {
std::lock_guard<std::mutex> lk(g_real_mutex);
g_real.clear();
}
// 复刻 _choose_victim遍历插入序(order)选 ∉current 且 freq 最小者,并列取最早(候选下标最小)。
// pinned(AUTOPIN 钉死)与 current 同样永不选。返回 expert id-1 表示无可驱逐。调用方须持 g_real_mutex。
static int choose_victim_locked(RealLayer& c, const std::unordered_set<int>& current) {
int victim = -1;
uint32_t best = 0;
for (int e : c.order) {
if (!c.e2r.count(e) || current.count(e) || c.pinned.count(e)) continue;
uint32_t f = c.freq.count(e) ? c.freq[e] : 0;
if (victim < 0 || f < best) { victim = e; best = f; } // 并列不更新 → 保留更早者
}
return victim;
}
// AUTOPIN 预热钉死:把历史热专家注册进真实区并标记 pinned(此后驱逐永不选它)。
// 已驻 → 仅补 pinned 标记、返回原槽;未驻 → free 头取槽登记(池行字节由调用方负责写入)。
// 返回与 experts 平行的槽位;空闲耗尽分不到槽给 -1(调用方跳过该专家)。幂等。
// 前置:调用方须先 real_init(否则 free_rows 为空、全返 -1)。
std::vector<int> real_pin(int layer, const std::vector<int>& experts) {
std::lock_guard<std::mutex> lk(g_real_mutex);
RealLayer& c = g_real[layer];
std::vector<int> out;
out.reserve(experts.size());
for (int e : experts) {
auto it = c.e2r.find(e);
if (it != c.e2r.end()) { c.pinned.insert(e); out.push_back(it->second); continue; }
if (c.free_rows.empty()) { out.push_back(-1); continue; }
int slot = c.free_rows.front();
c.free_rows.erase(c.free_rows.begin()); // free.pop(0)
c.e2r[e] = slot;
c.order.push_back(e);
c.pinned.insert(e);
out.push_back(slot);
}
return out;
}
// AUTOPIN 热度持久化:导出各层 LFU 累计频次,扁平 [layer, expert, count, ...]。
// 仅 lfu 策略下 demand 路径计频;驱逐不删 freq故即该进程全生命周期的路由热度
// 由 Python 侧按增量差分合并进 usage 文件(dual decode 无 Python 物化点,这是唯一计频来源)。
std::vector<long> real_freq_dump() {
std::lock_guard<std::mutex> lk(g_real_mutex);
std::vector<long> out;
for (auto& kv : g_real)
for (auto& p : kv.second.freq) {
out.push_back(static_cast<long>(kv.first));
out.push_back(static_cast<long>(p.first));
out.push_back(static_cast<long>(p.second));
}
return out;
}
// 复刻 _alloc_slote 为 miss不会已在 e2rfree 优先,否则 LFU 驱逐复用受害者槽。
// 返回 slot-1 表示无可驱逐(超容量,不该在 dual 发生)。调用方须持 g_real_mutex。
static int alloc_slot_locked(RealLayer& c, int e, const std::unordered_set<int>& current) {
auto it = c.e2r.find(e);
if (it != c.e2r.end()) return it->second;
int slot;
if (!c.free_rows.empty()) {
slot = c.free_rows.front();
c.free_rows.erase(c.free_rows.begin()); // free.pop(0)
} else {
int victim = choose_victim_locked(c, current);
if (victim < 0) return -1;
slot = c.e2r[victim];
c.e2r.erase(victim);
for (auto oit = c.order.begin(); oit != c.order.end(); ++oit)
if (*oit == victim) { c.order.erase(oit); break; }
}
c.e2r[e] = slot;
c.order.push_back(e);
return slot;
}
// LFU 频次 bumpcanonical本次全部唯一专家各 +1+ decay。调用方须持 g_real_mutex。
static void note_access_locked(RealLayer& c, const std::vector<int>& uniq_access,
bool lfu, int decay_interval) {
if (!lfu) return;
for (int e : uniq_access) c.freq[e] += 1;
c.access += static_cast<long>(uniq_access.size());
if (decay_interval > 0 && c.access >= decay_interval) {
for (auto it = c.freq.begin(); it != c.freq.end();) {
it->second /= 2;
if (it->second == 0) it = c.freq.erase(it);
else ++it;
}
c.access = 0;
}
}
// 核心状态机(纯 CPU、无 I/O给定 host inds(ip[n]) + 侧区快照(side),算 local、分配 miss 槽,
// 把需要落盘的新放入 (expert, slot) 追加到 placements字节由调用方在锁外并行 pread 落池)。
// 返回 local(int32 vector)。stats: [hitpos, misspos, loads(=placements 数), unplaced]。
static std::vector<int32_t> demand_core_locked(
RealLayer& c, const uint32_t* ip, size_t n, const std::unordered_map<int, int>& side,
bool lfu, int decay_interval, std::vector<std::pair<int, int>>& placements, long stats[4]) {
std::vector<int32_t> local(n, -1);
std::vector<int> uniq_miss, access_order;
std::unordered_set<int> miss_seen, access_seen;
int hitpos = 0;
// pass1算命中(侧区覆盖真实区)、收集 miss(首见序)、收集唯一访问(freq)。
for (size_t i = 0; i < n; ++i) {
int e = static_cast<int>(ip[i]);
if (access_seen.insert(e).second) access_order.push_back(e);
auto sit = side.find(e);
if (sit != side.end()) { local[i] = sit->second; ++hitpos; continue; }
auto rit = c.e2r.find(e);
if (rit != c.e2r.end()) { local[i] = rit->second; ++hitpos; continue; }
if (miss_seen.insert(e).second) uniq_miss.push_back(e);
}
note_access_locked(c, access_order, lfu, decay_interval);
// pass2miss 分配槽不落盘。current = 本前向全部唯一路由专家(命中+miss):绝不驱逐本前向要读的
// 任何专家的槽(否则真实区命中专家的槽被 miss 复写 → 脏字节。比 Python 仅护 miss 更严格、更正确)。
std::unordered_map<int, int> new_slot;
long unplaced = 0;
for (int e : uniq_miss) {
int slot = alloc_slot_locked(c, e, access_seen);
if (slot < 0) {
// 超容量:可驱逐槽(= cap pinned 本前向命中数)不够安放全部 miss。只能落 0 号槽,
// 即拿别的专家的权重参与计算 → 该前向逐位错算。调用方须靠 stats[3] 拦截并回退到
// host/fetch 路径;这里只保证「绝不静默」。
new_slot[e] = 0;
++unplaced;
continue;
}
new_slot[e] = slot;
placements.emplace_back(e, slot);
}
// pass3回填 miss 位置。
int misspos = 0;
for (size_t i = 0; i < n; ++i) {
if (local[i] < 0) {
auto it = new_slot.find(static_cast<int>(ip[i]));
local[i] = (it != new_slot.end()) ? it->second : 0;
++misspos;
}
}
stats[0] = hitpos; stats[1] = misspos; stats[2] = static_cast<long>(placements.size());
stats[3] = unplaced;
return local;
}
// demand 全接管inds 惰性(内部 eval 一次=1 次同步)side_gen 指定侧区代pool_list 为 _segs 顺序
// 的 per-key 池数组(已 eval、指针稳定)。返回 local(int32, inds.shape)。
mx::array demand_dual(
const mx::array& inds, const std::vector<mx::array>& pool_list,
const std::vector<int>& seg_nbytes, int layer, int side_gen, const std::string& path,
int stride, int cap, bool lfu, int decay_interval, mx::StreamOrDevice s = {}) {
if (seg_nbytes.size() != pool_list.size())
throw std::invalid_argument("demand_dual: seg_nbytes.size() != pool_list.size()");
double t0 = g_dt_on ? dt_now_us() : 0;
// 关键:inds 常是 argpartition(...)[..., -k:] 切片 → 非连续 strided 视图(每行按父数组 stride E
// 偏移)。若直接按连续读 data(),仅首行(token0)正确,后续 token 读到错位内存 → 装错专家、seq≥2 全错
// (投机 verify 全体受害)。contiguous() 强制物化为连续,读取才逐位正确。
mx::array ids = mx::contiguous(inds);
ids.eval(); // 唯一同步
size_t n = ids.size();
const uint32_t* ip = ids.data<uint32_t>();
double t1 = g_dt_on ? dt_now_us() : 0;
std::vector<uint8_t*> ptrs;
ptrs.reserve(pool_list.size());
for (auto& a0 : pool_list) { mx::array a = a0; a.eval(); ptrs.push_back(a.data<uint8_t>()); }
double t2 = g_dt_on ? dt_now_us() : 0;
std::unordered_map<int, int> side = sideregion_snapshot(layer, side_gen); // 侧区快照(该代)
double ta = g_dt_on ? dt_now_us() : 0;
long stats[4];
std::vector<int32_t> local;
std::vector<std::pair<int, int>> placements; // (expert, slot):锁外并行落盘
double tb;
{
std::lock_guard<std::mutex> lk(g_real_mutex);
tb = g_dt_on ? dt_now_us() : 0;
RealLayer& c = g_real[layer];
real_ensure_locked(c, cap);
local = demand_core_locked(c, ip, n, side, lfu, decay_interval, placements, stats);
}
// 锁外:把 miss 专家的字节 pread 落真实区槽。复刻基线并发模型——多 miss 派给 BgReader worker
// 并行 pread高优队列、直写池段行无主线程 tmp/memcpy本层等待其完成并行 → 远快于串行)。
static const bool kSkipIO = []() {
const char* e = std::getenv("DEMAND_SKIP_IO");
return e && e[0] == '1';
}();
if (!placements.empty() && !kSkipIO) {
std::vector<long> seg_off(seg_nbytes.size()), seg_nb(seg_nbytes.size());
long acc = 0;
for (size_t k = 0; k < seg_nbytes.size(); ++k) {
seg_off[k] = acc; seg_nb[k] = seg_nbytes[k]; acc += seg_nbytes[k];
}
std::vector<long> tickets;
tickets.reserve(placements.size());
for (auto& pr : placements) {
long tk = g_demand_ticket.fetch_add(1);
bg_pread_into_pool(pool_list, seg_off, seg_nb, pr.second, pr.first, path,
static_cast<long>(stride), tk, /*prio=*/1,
/*nocache=*/false); // demand route 读走 cache段偏移非页对齐
tickets.push_back(tk);
}
for (long tk : tickets) bg_reader_wait(tk); // 并行 pread 完成 → 池槽字节就绪
}
double t3 = g_dt_on ? dt_now_us() : 0;
{
std::lock_guard<std::mutex> lk(g_dstat_mutex);
g_d_last[0] = stats[0]; g_d_last[1] = stats[1]; g_d_last[2] = stats[2];
g_d_last[3] = (stats[1] == 0) ? 0 : 1;
g_d_last[4] = stats[3];
}
if (stats[3] > 0) {
static std::atomic<long> warned{0};
long k = warned.fetch_add(1);
if (k < 8 || k % 1000 == 0)
fprintf(stderr,
"[DEMAND_DUAL] 错算layer=%d 有 %ld 个专家分不到槽(落 0 号槽)。"
"inds=%zu cap=%d pinned=%d —— 该前向输出不可信,请调小 PREFILL_CHUNK "
"或调大 EXPERT_SLOTS / 调小 AUTOPIN_BUDGET_FRAC。\n",
layer, stats[3], n, cap, real_pinned_count(layer));
}
mx::array out = mx::array(local.data(), ids.shape(), mx::int32);
if (g_dt_on) {
double t4 = dt_now_us();
g_dt[0] += t1 - t0; g_dt[1] += t2 - t1; g_dt[2] += ta - t2;
g_dt[3] += tb - ta; g_dt[4] += t3 - tb; g_dt[5] += t4 - t3;
}
return out;
}
// 取本次 demand 统计 [hitpos, misspos, loads, fallback01, unplaced](主线程串行,安全)。
std::vector<long> demand_last_stats() {
std::lock_guard<std::mutex> lk(g_dstat_mutex);
return {g_d_last[0], g_d_last[1], g_d_last[2], g_d_last[3], g_d_last[4]};
}
// 测试壳:纯状态推进(不 pread/不侧区),把 experts_flat 当一次 demand 的 inds返回 local 槽位。
// 供 LFU 驱逐语义逐步等价单测(与 Python 参考对拍)。
std::vector<int> real_debug_place(int layer, const std::vector<int>& experts_flat, int cap,
bool lfu, int decay_interval) {
std::vector<uint32_t> u(experts_flat.begin(), experts_flat.end());
long stats[4];
std::vector<std::pair<int, int>> placements;
std::lock_guard<std::mutex> lk(g_real_mutex);
RealLayer& c = g_real[layer];
real_ensure_locked(c, cap);
std::unordered_map<int, int> empty_side;
auto local = demand_core_locked(c, u.data(), u.size(), empty_side, lfu, decay_interval,
placements, stats);
return std::vector<int>(local.begin(), local.end());
}

29
native/ext/pool/demand.h Normal file
View File

@ -0,0 +1,29 @@
// Phase 2 方案B真实区槽状态 C++ 全接管1 次同步版)+ demand_dual。
// 复刻 Python ResidentExpertPool 的 free/LFU 语义,成为 dual-source decode 真实区唯一权威。
#pragma once
#include "../common.h"
void real_init(int layer, int cap); // 初始化某层真实区(cap 槽全空闲),幂等
std::vector<int> real_region_contents(int layer); // [expert, slot, ...]
int real_region_count(int layer);
void real_reset();
// AUTOPIN注册 pinned 热专家(驱逐免选),返回平行槽位(-1=分不到);前置 real_init。
std::vector<int> real_pin(int layer, const std::vector<int>& experts);
// 该层 pinned(永不可驱逐)专家数。调用方据此算「可安放的唯一专家上限」= cap pinned
// 用于在超容量前把该前向分流到 host/fetch 路径(超了会落 0 号槽 → 逐位错算)。
int real_pinned_count(int layer);
// AUTOPIN导出各层 LFU 累计频次,扁平 [layer, expert, count, ...]。
std::vector<long> real_freq_dump();
// demand 全接管inds 惰性(内部 eval 一次)side_gen 侧区代pool_list 为 _segs 顺序池数组。
mx::array demand_dual(
const mx::array& inds, const std::vector<mx::array>& pool_list,
const std::vector<int>& seg_nbytes, int layer, int side_gen, const std::string& path,
int stride, int cap, bool lfu, int decay_interval, mx::StreamOrDevice s);
// [hitpos, misspos, loads, fallback01, unplaced]
// unplaced本次分不到槽、被迫落 0 号槽的唯一专家数。>0 即该前向逐位错算(必须为 0)。
std::vector<long> demand_last_stats();
std::vector<double> demand_timings(); // [inds_eval, pool_eval, state, build] us
void demand_timing_enable(bool on);
// 测试壳:纯状态推进(不 pread/不侧区),返回 local 槽位;供 LFU 驱逐等价对拍。
std::vector<int> real_debug_place(int layer, const std::vector<int>& experts_flat, int cap,
bool lfu, int decay_interval);

View File

@ -0,0 +1,89 @@
// [5] Route 3 Phase 1 底座C++ 拥有的池 buffer。
// 用 mx::allocator::malloc 分配 buffer、no-op deleter 建叶子 mx.arrayC++ 经 g_owned_bufs
// 持有 Buffer 句柄,进程内不释放)。因 C++ 独占持有、MLX 只读,该 buffer 永不被 MLX
// donation/迁移spike 已证),故侧区/demand 的后台异步 pread 可安全直写、消费侧读同一块。
#include "owned_pool.h"
static std::mutex g_owned_mutex;
static std::vector<mx::allocator::Buffer> g_owned_bufs;
static mx::Dtype dtype_from_str(const std::string& s) {
if (s == "uint32") return mx::uint32;
if (s == "uint16") return mx::uint16;
if (s == "uint8") return mx::uint8;
if (s == "int32") return mx::int32;
if (s == "int16") return mx::int16;
if (s == "bfloat16") return mx::bfloat16;
if (s == "float16") return mx::float16;
if (s == "float32") return mx::float32;
throw std::runtime_error("pool_owned_zeros: unsupported dtype " + s);
}
mx::array pool_owned_zeros(const std::vector<int>& shape, const std::string& dtype) {
mx::Dtype dt = dtype_from_str(dtype);
mx::Shape shp(shape.begin(), shape.end());
size_t n = 1;
for (int d : shape) n *= static_cast<size_t>(d);
size_t nbytes = n * static_cast<size_t>(mx::size_of(dt));
auto buf = mx::allocator::malloc(nbytes);
std::memset(buf.raw_ptr(), 0, nbytes);
{
std::lock_guard<std::mutex> lk(g_owned_mutex);
g_owned_bufs.push_back(buf); // C++ 持有,保证进程内不被释放
}
return mx::array(buf, shp, dt, [](mx::allocator::Buffer) {});
}
// demand 真实区落池:把一批已加载专家段 memcpy 进 owned 池行(无 MLX scatter → pool buffer 永不重绑)。
// pool_list[i]:第 i 个 key 的池数组srcs_flat 长 K*mkey-major[k0行0,k0行1,...,k1行0,...]
// slots[j]:第 j 个专家的目标物理行。CPU 端直写,调用点须保证此刻池未被 GPU 并发读该行。
void pool_write_rows(const std::vector<mx::array>& pool_list,
const std::vector<mx::array>& srcs_flat,
const std::vector<int>& slots) {
int K = static_cast<int>(pool_list.size());
int m = static_cast<int>(slots.size());
if (m == 0) return;
if (static_cast<int>(srcs_flat.size()) != K * m)
throw std::runtime_error("pool_write_rows: srcs_flat 长度须 == K*m");
for (int i = 0; i < K; ++i) {
mx::array p = pool_list[i];
p.eval();
uint8_t* base = p.data<uint8_t>();
for (int j = 0; j < m; ++j) {
mx::array s = srcs_flat[static_cast<size_t>(i) * m + j];
s.eval();
size_t nb = s.nbytes();
std::memcpy(base + static_cast<size_t>(slots[j]) * nb, s.data<uint8_t>(), nb);
}
}
}
// 同上,但源是每 key 预堆叠的 (m,*shape) 整块 stacked_list[i],按行 memcpy 到 slots[j]。
void pool_write_stacked(const std::vector<mx::array>& pool_list,
const std::vector<mx::array>& stacked_list,
const std::vector<int>& slots) {
int K = static_cast<int>(pool_list.size());
int m = static_cast<int>(slots.size());
if (m == 0) return;
for (int i = 0; i < K; ++i) {
mx::array p = pool_list[i];
p.eval();
mx::array st = stacked_list[i];
st.eval();
uint8_t* base = p.data<uint8_t>();
const uint8_t* src = st.data<uint8_t>();
size_t rownb = st.nbytes() / static_cast<size_t>(m); // stacked 为 (m,*shape) → 每行字节
for (int j = 0; j < m; ++j)
std::memcpy(base + static_cast<size_t>(slots[j]) * rownb,
src + static_cast<size_t>(j) * rownb, rownb);
}
}
// 临时诊断:返回某 mx.array 底层 buffer 的原始数据指针uintptr。用于对拍
// 「C++ 侧区 memcpy 写入的 buffer 指针」与「Python consume/verify 读到的 pool buffer 指针」
// 是否同一块——若不同即证明 MLX 在两者之间重分配了池 bufferraw 写入被落单。
uintptr_t array_data_ptr(const mx::array& a) {
mx::array b = a;
b.eval();
return reinterpret_cast<uintptr_t>(b.data<uint8_t>());
}

View File

@ -0,0 +1,14 @@
// Route 3 Phase 1 底座C++ 拥有的池 buffer + 直写(替代 mx.zeros 建池 + 消费侧 MLX scatter
#pragma once
#include "../common.h"
// C++ 拥有的池 buffer 包成 mx.arrayno-op deleter进程内持有、永不迁移
// 供侧区/demand 后台 pread 安全直写;替代原 mx.zeros 建池 + 消费侧 MLX scatter。
mx::array pool_owned_zeros(const std::vector<int>& shape, const std::string& dtype);
// demand 真实区落池:把已加载专家段直接 memcpy 进 owned 池行(无 MLX scatter保 buffer 稳定)。
void pool_write_rows(const std::vector<mx::array>& pool_list,
const std::vector<mx::array>& srcs_flat, const std::vector<int>& slots);
void pool_write_stacked(const std::vector<mx::array>& pool_list,
const std::vector<mx::array>& stacked_list, const std::vector<int>& slots);
// 诊断mx.array 底层 buffer 原始指针uintptr用于对拍侧区写入 buffer 与 consume 读到 buffer 是否同一块。
uintptr_t array_data_ptr(const mx::array& a);

View File

@ -0,0 +1,389 @@
// [4] 段散写持久侧区缓存zero-copy dual-source 默认路径)。
// 与持久 staging 缓存同构,但目标是“多个结构化 per-key 池数组”:每个池数组形如
// (cap+spec, ...),行 [base_row, base_row+spec) 为侧区。命中缺口时 pread blob 整行,
// 再把该行内按固定顺序拼接的各段 memcpy 进对应 per-key 数组的同一物理行。
#include "side_region.h"
#include "../io/blob_io.h"
#include "../io/bg_reader.h"
#include <condition_variable>
#include <cstdio>
#include <cstdlib>
#include <functional>
#include <thread>
#include <unordered_set>
struct SideLayer {
std::map<int, int> e2r; // expert -> 物理侧区行 [base_row, base_row+spec)
std::vector<int> free_rows;
std::map<int, uint32_t> freq; // expert -> 预测频次(LFU 分数;仅 SIDEREGION_LFU 用)
bool inited = false;
int base = 0; // 侧区起始物理行 base_row
int spec = 0; // 侧区行数 spec_slots
};
static std::mutex g_side_mutex;
static std::map<std::pair<int, int>, SideLayer> g_side; // 键 (layer, gen):双缓冲两代独立
// 侧区 fill 在途计数eval_gpu 提交预取时 +1后台 read_publish 写完字节后 -1。
// 消费方在前向开头 sideregion_drain() 排空上一前向的 fill保证被消费的侧区行字节
// 已完全写好(含 GPU 完成回调滞后的情形)→ 消灭「GPU 消费 kernel 读到半写侧区行」竞态。
static std::atomic<long> g_side_inflight{0};
static std::mutex g_side_drain_mutex;
static std::condition_variable g_side_drain_cv;
static inline void side_inflight_done() {
{
std::lock_guard<std::mutex> lk(g_side_drain_mutex);
g_side_inflight.fetch_sub(1);
}
g_side_drain_cv.notify_all();
}
void sideregion_drain() {
std::unique_lock<std::mutex> lk(g_side_drain_mutex);
g_side_drain_cv.wait(lk, [] { return g_side_inflight.load() == 0; });
}
// 临时诊断:只追踪指定 (layer,row) 的所有账本/字节事件SIDE_TRACE_LAYER/ROW
static inline bool side_trace_hit(int layer, int row) {
const char* le = std::getenv("SIDE_TRACE_LAYER");
const char* re = std::getenv("SIDE_TRACE_ROW");
if (!le || !re) return false;
return layer == atoi(le) && row == atoi(re);
}
// 临时诊断:跨线程全局事件序号 + 线程短 id用于把 reserve(Metal 回调线程)/read_publish(bg 线程)/
// sideregion_kv(主线程) 的事件按发生顺序排出来root-cause 取证SIDE_TRACE_* 命中时才用)。
static std::atomic<uint64_t> g_side_ev{0};
static inline unsigned side_tid() {
return static_cast<unsigned>(
std::hash<std::thread::id>{}(std::this_thread::get_id()) & 0xffff);
}
class PrefetchPoolSideRegionPrimitive : public mx::Primitive {
public:
PrefetchPoolSideRegionPrimitive(mx::Stream s, std::vector<int> seg_nbytes, int layer, int gen,
std::string path, size_t stride, std::vector<int> resident,
int spec_slots, int base_row)
: Primitive(s), seg_(std::move(seg_nbytes)), layer_(layer), gen_(gen), path_(std::move(path)),
stride_(stride), resident_(std::move(resident)), spec_(spec_slots), base_(base_row) {}
const char* name() const override { return "PrefetchPoolSideRegionPrimitive"; }
void eval_cpu(const std::vector<mx::array>& in, std::vector<mx::array>& out) override {
out[0].set_data(mx::allocator::malloc(out[0].nbytes()));
run(in);
}
void eval_gpu(const std::vector<mx::array>& in, std::vector<mx::array>& out) override {
out[0].set_data(mx::allocator::malloc(out[0].nbytes()));
std::vector<uint8_t*> ptrs;
ptrs.reserve(in.size() - 1);
for (size_t i = 1; i < in.size(); ++i) {
mx::array a = in[i]; // 非 const 拷贝才能取可写指针
ptrs.push_back(a.data<uint8_t>());
}
mx::array ids = in[0];
const uint32_t* idp = ids.data<uint32_t>();
size_t n = ids.size();
std::vector<int> seg = seg_;
int layer = layer_;
int gen = gen_;
std::string path = path_;
size_t stride = stride_;
std::vector<int> resident = resident_;
int spec = spec_, base = base_;
auto& enc = mx::metal::get_command_encoder(stream());
MTL::CommandBuffer* cb = enc.get_command_buffer();
// 提交即计在途:在 eval_gpu预取提交时 +1直到后台字节写完才 -1。这样消费方前向开头
// sideregion_drain() 能等到「即使 GPU 完成回调尚未触发」的 fill闭合跨前向的写-读竞态。
g_side_inflight.fetch_add(1);
// in 按值捕获 → 保活 expert_ids 与所有池数组 bufferidp/ptrs 在回调里指针有效。
cb->addCompletedHandler(
[in, ptrs, seg, idp, n, layer, gen, path, stride, resident, spec, base](MTL::CommandBuffer*) {
// 阶段1回调线程、持锁极短读惰性 id此刻已算完、预留侧区行。
// 必须在回调里——id 只有 command buffer 完成后才有效。
auto to_read = reserve(idp, n, layer, gen, resident, spec, base);
if (to_read.empty()) {
side_inflight_done(); // 无缺口可读:立即消账,避免 drain 空等
return;
}
// 诊断门控 SIDEREGION_SYNC=1回调内同步 pread+memcpy+publish不派 bg
// 消除「异步 bg fill 与下一前向 gather 竞态」这一变量,用于 systematic-debugging 取证。
static const bool kSync = []() {
const char* e = std::getenv("SIDEREGION_SYNC");
return e && e[0] == '1';
}();
if (kSync) {
read_publish(ptrs, seg, to_read, path, stride, layer, gen);
side_inflight_done();
return;
}
// 阶段2+3 派给自由后台线程:~40MB pread + memcpy + 发布 e2r 脱离 Metal 回调线程,
// 与主 stream 后续层计算、多层预取互相并发 → 真正恢复 I/O/计算重叠。
// 闭包再次按值捕获 in保活池/ids buffer 到后台读完ptrs 指向其内存)。
bg_submit_task([in, ptrs, seg, to_read, path, stride, layer, gen]() {
read_publish(ptrs, seg, to_read, path, stride, layer, gen);
side_inflight_done(); // 字节写完才消账 → drain 保证被消费行已就绪
});
});
}
private:
void run(const std::vector<mx::array>& in) {
std::vector<uint8_t*> ptrs;
ptrs.reserve(in.size() - 1);
for (size_t i = 1; i < in.size(); ++i) {
mx::array a = in[i];
a.eval(); // 与 PrefetchStagingCachedPrimitive::eval_cpu 一致:先物化输入
ptrs.push_back(a.data<uint8_t>());
}
mx::array ids = in[0];
ids.eval();
// CPU 路径同步执行(测试/无 GPU 时):预留 + 读发布一气呵成。
auto to_read = reserve(ids.data<uint32_t>(), ids.size(), layer_, gen_, resident_, spec_, base_);
read_publish(ptrs, seg_, to_read, path_, stride_, layer_, gen_);
}
// 阶段1过滤常驻/去重 → 淘汰 ∉P 的旧行 → 为缺口预留物理行(出 free暂不入 e2r
// 避免消费者在字节写好前看到 e2r 命中而 gather 到脏字节)。返回 (expert, 预留行)。
static std::vector<std::pair<int, int>> reserve(
const uint32_t* idp, size_t n, int layer, int gen, const std::vector<int>& resident,
int spec, int base) {
const char* lfu_env = std::getenv("SIDEREGION_LFU"); // 每次读,便于测试切换
// 默认开:持久 LFU 单缓冲=生产路径(cli 默认)。仅显式 SIDEREGION_LFU=0 回退旧 legacy 双缓冲。
bool lfu = !lfu_env || lfu_env[0] != '0';
std::unordered_set<int> res(resident.begin(), resident.end());
std::vector<int> P;
std::unordered_set<int> Pset, seen;
for (size_t i = 0; i < n; ++i) {
int e = static_cast<int>(idp[i]);
if (res.count(e) || !seen.insert(e).second) continue;
P.push_back(e);
Pset.insert(e);
}
std::vector<std::pair<int, int>> to_read;
std::lock_guard<std::mutex> lk(g_side_mutex);
SideLayer& c = g_side[{layer, gen}];
if (!c.inited) {
for (int r = 0; r < spec; ++r) c.free_rows.push_back(base + r);
c.base = base;
c.spec = spec;
c.inited = true;
}
if (!lfu) {
// 旧行为:∉P 全弃(一次性预取批)。
for (auto it = c.e2r.begin(); it != c.e2r.end();) {
if (!Pset.count(it->first)) {
c.free_rows.push_back(it->second);
it = c.e2r.erase(it);
} else {
++it;
}
}
for (int e : P) {
if (c.e2r.count(e) || c.free_rows.empty()) continue;
to_read.emplace_back(e, c.free_rows.back());
c.free_rows.pop_back();
}
return to_read;
}
// LFU 持久:∉P 不清;再预测命中已驻专家 freq+1(越常预测越热)。
for (int e : P) {
if (c.e2r.count(e)) c.freq[e] += 1;
}
for (int e : P) {
if (c.e2r.count(e)) continue; // 已驻,跳过(不重读)
int row;
if (!c.free_rows.empty()) {
row = c.free_rows.back();
c.free_rows.pop_back();
if (side_trace_hit(layer, row))
fprintf(stderr, "[SIDE_TRACE ev=%llu tid=%u] L%d gen%d RESERVE_FROM_FREE row=%d expert=%d\n",
(unsigned long long)g_side_ev.fetch_add(1), side_tid(), layer, gen, row, e);
} else {
// free 空:LFU 淘汰 e2r 中 freq 最小且 ∉P 者(tie-break:最小 expert id)。
int victim = -1;
uint32_t best = 0;
for (auto& kv : c.e2r) {
if (Pset.count(kv.first)) continue; // 不淘本步要用的
uint32_t f = c.freq.count(kv.first) ? c.freq[kv.first] : 0;
if (victim < 0 || f < best || (f == best && kv.first < victim)) {
victim = kv.first;
best = f;
}
}
if (victim < 0) continue; // 全是 P 热,无可淘 → 本步不读
row = c.e2r[victim];
c.e2r.erase(victim);
c.freq.erase(victim);
if (side_trace_hit(layer, row))
fprintf(stderr, "[SIDE_TRACE ev=%llu tid=%u] L%d gen%d EVICT_REUSE row=%d victim=%d newExpert=%d\n",
(unsigned long long)g_side_ev.fetch_add(1), side_tid(), layer, gen, row, victim, e);
}
to_read.emplace_back(e, row);
}
// ===== 临时诊断SIDE_AUDIT=1reserve 结束时审计侧区行账本自洽性 =====
if (std::getenv("SIDE_AUDIT")) side_audit(c, layer, gen, "reserve", to_read);
return to_read;
}
// 审计检测「一行被两专家占用」「free 与 e2r 行重叠」「to_read 行仍被 e2r 占用」。
static void side_audit(SideLayer& c, int layer, int gen, const char* where,
const std::vector<std::pair<int, int>>& to_read) {
std::map<int, int> row_owner; // row -> expert
for (auto& p : c.e2r) {
auto it = row_owner.find(p.second);
if (it != row_owner.end())
fprintf(stderr, "[SIDE_AUDIT] %s L%d gen%d DOUBLE_OWN row=%d experts=%d,%d\n",
where, layer, gen, p.second, it->second, p.first);
row_owner[p.second] = p.first;
}
std::unordered_set<int> free_set(c.free_rows.begin(), c.free_rows.end());
for (auto& p : c.e2r)
if (free_set.count(p.second))
fprintf(stderr, "[SIDE_AUDIT] %s L%d gen%d FREE_E2R_OVERLAP row=%d owned_by=%d\n",
where, layer, gen, p.second, p.first);
if (c.free_rows.size() != free_set.size())
fprintf(stderr, "[SIDE_AUDIT] %s L%d gen%d FREE_DUP free_n=%zu uniq=%zu\n",
where, layer, gen, c.free_rows.size(), free_set.size());
for (auto& pr : to_read)
if (row_owner.count(pr.second))
fprintf(stderr, "[SIDE_AUDIT] %s L%d gen%d TOREAD_LIVE_ROW row=%d assigned_to=%d still_owned_by=%d\n",
where, layer, gen, pr.second, pr.first, row_owner[pr.second]);
}
// 阶段2+3pread blob 整行 + 各段 memcpy 直写进对应 per-key 池数组的物理侧区行(不持锁),
// 写完后持锁发布 e2r。ptrs[i] 为第 i 个池 key 数组的 buffer 指针C++ 拥有、地址恒定不迁移,
// 由 Route 3 底座保证——见 pool_owned_zeros故后台异步直写安全、消费侧读同一 buffer。
// 在后台线程跑:与计算并发,且不阻塞消费侧的 sideregion_kv。
static void read_publish(const std::vector<uint8_t*>& ptrs, const std::vector<int>& seg,
const std::vector<std::pair<int, int>>& to_read,
const std::string& path, size_t stride, int layer, int gen) {
// 段在 blob 记录内的偏移(段顺序 = seg 顺序 = ptrs/池 key 顺序)。
std::vector<size_t> seg_off(seg.size(), 0);
size_t acc = 0;
for (size_t i = 0; i < seg.size(); ++i) { seg_off[i] = acc; acc += static_cast<size_t>(seg[i]); }
int fd = open_blob_nocache(path.c_str());
if (fd < 0) { // 失败则把预留行还回 free
std::lock_guard<std::mutex> lk(g_side_mutex);
SideLayer& c = g_side[{layer, gen}];
for (auto& pr : to_read) c.free_rows.push_back(pr.second);
return;
}
std::vector<uint8_t> rec(stride); // 整条 blob 记录临时缓冲
std::vector<std::pair<int, int>> done;
for (auto& pr : to_read) {
int e = pr.first, row = pr.second;
if (::pread(fd, rec.data(), stride, static_cast<off_t>(static_cast<size_t>(e) * stride)) !=
static_cast<ssize_t>(stride)) {
std::lock_guard<std::mutex> lk(g_side_mutex); // 读失败:行还回 free
g_side[{layer, gen}].free_rows.push_back(row);
continue;
}
// 各段 memcpy 进对应 per-key 池数组的物理行 row(n_slots,*shape) 布局 → 行偏移 = row*seg[i])。
for (size_t i = 0; i < seg.size(); ++i)
std::memcpy(ptrs[i] + static_cast<size_t>(row) * static_cast<size_t>(seg[i]),
rec.data() + seg_off[i], static_cast<size_t>(seg[i]));
if (side_trace_hit(layer, row))
fprintf(stderr, "[SIDE_TRACE ev=%llu tid=%u] L%d gen%d WRITEPOOL row=%d expert=%d\n",
(unsigned long long)g_side_ev.fetch_add(1), side_tid(), layer, gen, row, e);
done.emplace_back(e, row);
}
::close(fd);
{
std::lock_guard<std::mutex> lk(g_side_mutex);
SideLayer& c = g_side[{layer, gen}];
for (auto& pr : done) { // 字节就绪后才发布 e2r
c.e2r[pr.first] = pr.second;
if (!c.freq.count(pr.first)) c.freq[pr.first] = 1; // 新专家初始 freq
if (side_trace_hit(layer, pr.second))
fprintf(stderr, "[SIDE_TRACE ev=%llu tid=%u] L%d gen%d PUBLISH row=%d expert=%d\n",
(unsigned long long)g_side_ev.fetch_add(1), side_tid(), layer, gen, pr.second,
pr.first);
}
if (std::getenv("SIDE_AUDIT")) side_audit(c, layer, gen, "publish", {});
}
}
std::vector<int> seg_;
int layer_;
int gen_;
std::string path_;
size_t stride_;
std::vector<int> resident_;
int spec_;
int base_;
};
mx::array prefetch_pool_sideregion(
const std::vector<mx::array>& pool_list, const std::vector<int>& seg_nbytes,
const mx::array& expert_ids, int layer, const std::string& path, int stride,
const std::vector<int>& resident, int spec_slots, int base_row, int gen,
mx::StreamOrDevice s = {}) {
// 防御:段数必须与池数组个数一一对应(每段映射唯一一个 per-key 数组)。
if (seg_nbytes.size() != pool_list.size()) {
throw std::invalid_argument(
"prefetch_pool_sideregion: seg_nbytes.size()=" + std::to_string(seg_nbytes.size()) +
" != pool_list.size()=" + std::to_string(pool_list.size()));
}
// 防御:各段字节之和必须等于 strideblob 记录恰好是各段的拼接)。
size_t seg_sum = 0;
for (int b : seg_nbytes) seg_sum += static_cast<size_t>(b);
if (seg_sum != static_cast<size_t>(stride)) {
throw std::invalid_argument(
"prefetch_pool_sideregion: sum(seg_nbytes)=" + std::to_string(seg_sum) +
" != stride=" + std::to_string(stride));
}
std::vector<mx::array> inputs;
inputs.push_back(expert_ids);
for (auto& a : pool_list) inputs.push_back(a);
return mx::array(
mx::Shape{1}, mx::uint8,
std::make_shared<PrefetchPoolSideRegionPrimitive>(
mx::to_stream(s), seg_nbytes, layer, gen, path, static_cast<size_t>(stride), resident,
spec_slots, base_row),
inputs);
}
std::vector<int> sideregion_contents(int layer, int gen) {
std::lock_guard<std::mutex> lk(g_side_mutex);
std::vector<int> out;
auto it = g_side.find({layer, gen});
if (it != g_side.end())
for (auto& p : it->second.e2r) { out.push_back(p.first); out.push_back(p.second); }
return out;
}
// 侧区 e2r → 两个 device mx.array (keys uint32, vals int32),直接在 C++ 从 map 建连续 buffer。
// 消掉 Python 侧 dict 构建 + list(...)→mx.array 的每层 host 胶水。
std::pair<mx::array, mx::array> sideregion_kv(int layer, int gen) {
std::vector<uint32_t> keys;
std::vector<int32_t> vals;
{
std::lock_guard<std::mutex> lk(g_side_mutex);
auto it = g_side.find({layer, gen});
if (it != g_side.end()) {
keys.reserve(it->second.e2r.size());
vals.reserve(it->second.e2r.size());
for (auto& p : it->second.e2r) {
keys.push_back(static_cast<uint32_t>(p.first));
vals.push_back(static_cast<int32_t>(p.second));
if (side_trace_hit(layer, p.second))
fprintf(stderr, "[SIDE_TRACE ev=%llu tid=%u] L%d gen%d CONSUME_KV row=%d expert=%d\n",
(unsigned long long)g_side_ev.fetch_add(1), side_tid(), layer, gen, p.second, p.first);
}
}
}
int n = static_cast<int>(keys.size());
return {mx::array(keys.data(), mx::Shape{n}, mx::uint32),
mx::array(vals.data(), mx::Shape{n}, mx::int32)};
}
void sideregion_reset() {
std::lock_guard<std::mutex> lk(g_side_mutex);
g_side.clear();
}
std::unordered_map<int, int> sideregion_snapshot(int layer, int gen) {
std::unordered_map<int, int> side;
std::lock_guard<std::mutex> lk(g_side_mutex);
auto it = g_side.find({layer, gen});
if (it != g_side.end()) for (auto& p : it->second.e2r) side[p.first] = p.second;
return side;
}

View File

@ -0,0 +1,21 @@
// 段散写持久侧区缓存zero-copy dual-source 默认路径keep/evict/read-delta
// 命中缺口时 pread blob 行并把各段 memcpy 进对应 per-key 池数组的 (base_row+cache_row) 行。
#pragma once
#include "../common.h"
#include <unordered_map>
mx::array prefetch_pool_sideregion(
const std::vector<mx::array>& pool_list, const std::vector<int>& seg_nbytes,
const mx::array& expert_ids, int layer, const std::string& path, int stride,
const std::vector<int>& resident, int spec_slots, int base_row, int gen,
mx::StreamOrDevice s);
std::vector<int> sideregion_contents(int layer, int gen); // [expert, phys_row, ...]
std::pair<mx::array, mx::array> sideregion_kv(int layer, int gen); // (keys uint32, vals int32) device 数组
void sideregion_reset();
// 排空侧区在途 fill阻塞直到所有已提交的侧区预取字节写完。消费方在前向开头调用
// 保证被消费的侧区行字节已完全写好(闭合异步写-读竞态)。
void sideregion_drain();
// 内部接口(不绑 Python取某层某代侧区 e2r 快照,供 demand_dual 在自己的核心状态机里
// 叠加侧区命中。等价于在 g_side 锁下拷贝 e2r与原 demand_dual 内联循环一字不差)。
std::unordered_map<int, int> sideregion_snapshot(int layer, int gen);

View File

@ -0,0 +1,217 @@
// [2] 轻量预取(无 staging仅预热 page cache) + [3] staging 版 miss→hit + STAGING_HPROF 探针。
#include "prefetch.h"
#include "../io/bg_reader.h"
#include <chrono>
#include <functional>
#include <map>
#include <tuple>
#include <unordered_set>
#include <fcntl.h>
#include <unistd.h>
// ====== [2] 轻量预取(无 staging 模式)GPU 完成回调里读 inds、pread 预热 page cache ======
// 命门验证:回调在 command buffer 完成后触发,此时 inds 已算完,读到的是正确值。
class PrefetchOnCompletePrimitive : public mx::Primitive {
public:
PrefetchOnCompletePrimitive(mx::Stream s, std::string path, size_t stride, bool do_read)
: Primitive(s), path_(std::move(path)), stride_(stride), do_read_(do_read) {}
const char* name() const override { return "PrefetchOnCompletePrimitive"; }
void eval_cpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
outputs[0].set_data(mx::allocator::malloc(outputs[0].nbytes()));
mx::array ids = inputs[0];
ids.eval();
record_and_read(ids.data<uint32_t>(), ids.size(), path_, stride_, do_read_);
}
void eval_gpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
outputs[0].set_data(mx::allocator::malloc(outputs[0].nbytes()));
mx::array ids = inputs[0]; // by-value保活到回调结束
const uint32_t* ptr = ids.data<uint32_t>(); // 指针此刻有效buffer 已分配),值待 GPU 算完
size_t n = ids.size();
std::string path = path_;
size_t stride = stride_;
bool do_read = do_read_;
auto& enc = mx::metal::get_command_encoder(stream());
MTL::CommandBuffer* cb = enc.get_command_buffer();
cb->addCompletedHandler([ids, ptr, n, path, stride, do_read](MTL::CommandBuffer*) {
// 此时 buffer 已完成 → ptr 指向已算好的 inds 值。ids 按值捕获保证 buffer 不被释放。
record_and_read(ptr, n, path, stride, do_read);
});
}
private:
static void record_and_read(const uint32_t* p, size_t n, const std::string& path,
size_t stride, bool do_read) {
int fd = (do_read && !path.empty()) ? ::open(path.c_str(), O_RDONLY) : -1;
static thread_local std::vector<uint8_t> buf;
if (do_read && buf.size() < stride) buf.resize(stride);
for (size_t i = 0; i < n; ++i) {
int e = static_cast<int>(p[i]);
if (fd >= 0) ::pread(fd, buf.data(), stride, static_cast<off_t>(static_cast<size_t>(e) * stride));
}
if (fd >= 0) ::close(fd);
}
std::string path_;
size_t stride_;
bool do_read_;
};
// 返回一个 dummy 输出挂进图里eval 时触发上面的回调)。
mx::array prefetch_on_complete(
const mx::array& expert_ids,
const std::string& path,
int stride,
bool do_read = true,
mx::StreamOrDevice s = {}) {
return mx::array(
mx::Shape{1},
mx::uint8,
std::make_shared<PrefetchOnCompletePrimitive>(
mx::to_stream(s), path, static_cast<size_t>(stride), do_read),
std::vector<mx::array>{expert_ids});
}
// ---- handler 触发时刻探针(STAGING_HPROF):记录每次 staging 完成回调被 Metal 触发的时刻 ----
// 用于实测"完成回调到底扎堆在 eval 尾、还是逐层铺开"。仅诊断用,默认关。
static std::mutex g_hprof_mutex;
static std::vector<std::tuple<long, int, double>> g_hprof_log; // (gen, layer, t_fire_seconds)
static bool g_hprof_on = false;
static inline double hprof_steady_now() {
return std::chrono::duration<double>(
std::chrono::steady_clock::now().time_since_epoch())
.count();
}
// ====== [3] staging 版 miss→hithandler pread 进 per-layer staging + 记录 (expert→row) ======
static std::mutex g_stg_mutex;
// layer -> (gen, [(expert, row)])handler 原子写 gen+映射,主线程按 gen 匹配 buffer 后 take。
static std::map<int, std::pair<long, std::vector<std::pair<int, int>>>> g_stg_ready;
class PrefetchStagingPrimitive : public mx::Primitive {
public:
// 方案Bexpert_ids 是"预测宽集合"(top-N按门控分降序)resident 是目标层当前常驻专家。
// handler 在回调里过滤掉常驻、按降序取前 cap 个缺口 pread 进 stagingcap=buffer 行数)。
// 这样预测可以很宽(高 recall),而 staging 内存只按 cap 预留(小,覆盖缺口分布即可)。
PrefetchStagingPrimitive(mx::Stream s, int layer, long gen, std::string path, size_t stride,
std::vector<int> resident, int cap, bool parallel)
: Primitive(s), layer_(layer), gen_(gen), path_(std::move(path)), stride_(stride),
resident_(std::move(resident)), cap_(cap), parallel_(parallel) {}
const char* name() const override { return "PrefetchStagingPrimitive"; }
void eval_cpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
outputs[0].set_data(mx::allocator::malloc(outputs[0].nbytes()));
mx::array ids = inputs[0]; ids.eval();
mx::array stg = inputs[1]; stg.eval();
fill(ids.data<uint32_t>(), ids.size(), stg.data<uint8_t>(), layer_, gen_, path_, stride_,
resident_, cap_);
}
void eval_gpu(const std::vector<mx::array>& inputs, std::vector<mx::array>& outputs) override {
outputs[0].set_data(mx::allocator::malloc(outputs[0].nbytes()));
mx::array ids = inputs[0];
mx::array stg = inputs[1];
const uint32_t* idp = ids.data<uint32_t>();
uint8_t* sp = stg.data<uint8_t>();
size_t n = ids.size();
int layer = layer_; long gen = gen_; std::string path = path_; size_t stride = stride_;
std::vector<int> resident = resident_; int cap = cap_; bool parallel = parallel_;
auto& enc = mx::metal::get_command_encoder(stream());
MTL::CommandBuffer* cb = enc.get_command_buffer();
cb->addCompletedHandler(
[ids, stg, idp, sp, n, layer, gen, path, stride, resident, cap, parallel](MTL::CommandBuffer*) {
// 探针:记录本回调被 Metal 触发的时刻(= 该层 pread 能开始跑的时刻)。
if (g_hprof_on) {
double t = hprof_steady_now();
std::lock_guard<std::mutex> lk(g_hprof_mutex);
g_hprof_log.emplace_back(gen, layer, t);
}
// ids/stg 按值捕获保活 buffer派后台线程时再拷一份保活到 fill 跑完。
if (parallel) {
bg_submit_task([ids, stg, idp, sp, n, layer, gen, path, stride, resident, cap]() {
fill(idp, n, sp, layer, gen, path, stride, resident, cap);
});
} else {
fill(idp, n, sp, layer, gen, path, stride, resident, cap);
}
});
}
private:
static void fill(const uint32_t* idp, size_t n, uint8_t* stg,
int layer, long gen, const std::string& path, size_t stride,
const std::vector<int>& resident, int cap) {
int fd = ::open(path.c_str(), O_RDONLY);
if (fd < 0) return;
std::unordered_set<int> res(resident.begin(), resident.end());
std::unordered_set<int> done; // 去重:同一缺口只 pread 一次
std::vector<std::pair<int, int>> ready;
ready.reserve(static_cast<size_t>(cap));
int row = 0; // 写入的 staging 行(≤ cap
for (size_t i = 0; i < n && row < cap; ++i) {
int e = static_cast<int>(idp[i]);
if (res.count(e) || !done.insert(e).second) continue; // 已常驻/已取过 → 跳过
if (::pread(fd, stg + static_cast<size_t>(row) * stride, stride,
static_cast<off_t>(static_cast<size_t>(e) * stride))
== static_cast<ssize_t>(stride)) {
ready.emplace_back(e, row);
++row;
}
}
::close(fd);
std::lock_guard<std::mutex> lk(g_stg_mutex);
g_stg_ready[layer] = {gen, std::move(ready)}; // 原子gen 与映射一起写
}
int layer_;
long gen_;
std::string path_;
size_t stride_;
std::vector<int> resident_; // 目标层提交时刻的常驻专家快照(过滤用)
int cap_; // staging buffer 行数上限
bool parallel_; // true: fill 派后台线程池并行false: 回调线程同步(旧行为)
};
mx::array prefetch_into_staging(
const mx::array& staging, const mx::array& expert_ids, int layer, long gen,
const std::string& path, int stride, const std::vector<int>& resident, int cap,
bool parallel, mx::StreamOrDevice s = {}) {
return mx::array(
mx::Shape{1}, mx::uint8,
std::make_shared<PrefetchStagingPrimitive>(
mx::to_stream(s), layer, gen, path, static_cast<size_t>(stride), resident, cap, parallel),
std::vector<mx::array>{expert_ids, staging});
}
// 取走某层就绪记录:[gen, e0,r0,e1,r1,...](首元素是 generation空表示无就绪。
std::vector<long> prefetch_staging_take(int layer) {
std::lock_guard<std::mutex> lk(g_stg_mutex);
std::vector<long> out;
auto it = g_stg_ready.find(layer);
if (it != g_stg_ready.end()) {
out.push_back(it->second.first); // gen
for (auto& p : it->second.second) { out.push_back(p.first); out.push_back(p.second); }
g_stg_ready.erase(it);
}
return out;
}
// ---- handler 触发时刻探针接口 ----
void staging_hprof_enable(bool on) {
std::lock_guard<std::mutex> lk(g_hprof_mutex);
g_hprof_on = on;
g_hprof_log.clear(); // 开启即清零,便于每次采集干净
}
double staging_hprof_now() { return hprof_steady_now(); } // 与日志同一时钟,供 Python 标 eval 边界
// 扁平返回 [gen0,layer0,t0, gen1,layer1,t1, ...](避开 nanobind tuple caster
std::vector<double> staging_hprof_get() {
std::lock_guard<std::mutex> lk(g_hprof_mutex);
std::vector<double> out;
out.reserve(g_hprof_log.size() * 3);
for (auto& r : g_hprof_log) {
out.push_back(static_cast<double>(std::get<0>(r)));
out.push_back(static_cast<double>(std::get<1>(r)));
out.push_back(std::get<2>(r));
}
return out;
}

View File

@ -0,0 +1,23 @@
// GPU 完成回调预取:算完 inds 后在回调里读 id、pread 预热 page cache 或写 per-layer staging。
#pragma once
#include "../common.h"
// 挂 GPU 完成回调:算完 inds 后在回调里读 id 预热 page cache返回 dummy。
// (无 staging 的轻量预取模式:仅把预测专家字节读进 page cache不落池。
mx::array prefetch_on_complete(
const mx::array& expert_ids, const std::string& path, int stride,
bool do_read, mx::StreamOrDevice s);
// miss→hit方案Bexpert_ids 为预测宽集合(降序),回调按 resident 过滤、取前 cap 个缺口
// pread 进 per-layer staging buffer(cap=buffer 行数),并原子记录 (gen,[expert,row])。
mx::array prefetch_into_staging(
const mx::array& staging, const mx::array& expert_ids, int layer, long gen,
const std::string& path, int stride, const std::vector<int>& resident, int cap,
bool parallel, mx::StreamOrDevice s);
// 取走某层就绪记录:[gen, e0,r0,e1,r1,...];空表示无就绪。
std::vector<long> prefetch_staging_take(int layer);
// handler 触发时刻探针(诊断用)enable 开关并清零now 取同时钟当前秒get 取 (gen,layer,t) 日志。
void staging_hprof_enable(bool on);
double staging_hprof_now();
std::vector<double> staging_hprof_get(); // 扁平 [gen,layer,t, ...]

44
pyproject.toml Normal file
View File

@ -0,0 +1,44 @@
[project]
name = "sparkle-engine"
version = "0.1.0"
description = "SparkleApple Silicon 上流式 MoE 推理Qwen3-Next-80B 8bit"
readme = "README.md"
# 仓库自带 native_moe_ext*.cpython-314*.so须与扩展 ABI 一致
requires-python = "==3.14.*"
dependencies = [
"fastapi>=0.115",
"mlx>=0.31",
"mlx-lm>=0.31",
"numpy>=2.0",
"textual>=0.80",
"uvicorn>=0.30",
]
[project.scripts]
sparkle = "mlx_streaming.cli:main"
# 仅重编 native / 跑单测时需要;日常推理用已有 .venv + .so 即可
[dependency-groups]
dev = [
"httpx>=0.27",
"nanobind>=2.12.0",
"pytest>=8",
"pytest-asyncio>=0.23",
]
[[tool.uv.index]]
# 清华 PyPI 镜像
url = "https://pypi.tuna.tsinghua.edu.cn/simple"
default = true
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["mlx_streaming"]
[tool.pytest.ini_options]
addopts = "-q"
testpaths = ["mlx_streaming/tests"]
asyncio_mode = "auto"

1129
uv.lock generated Normal file

File diff suppressed because it is too large Load Diff