This commit is contained in:
commit
4745f264b2
40
.gitignore
vendored
Normal file
40
.gitignore
vendored
Normal 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
187
README.md
Normal file
@ -0,0 +1,187 @@
|
|||||||
|
# Sparkle
|
||||||
|
|
||||||
|
本地 **Qwen3-Next-80B-A3B Instruct 8bit** 流式推理引擎,对接 ccSparkle Agent(模型源:**Sparkle**)。
|
||||||
|
|
||||||
|
在约 64GB 统一内存的 Apple Silicon 上,把约 77GB 专家权重放在 SSD,运行时按需装入 + 预取,从而跑通 80B 8bit(oMLX / 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) |
|
||||||
|
| 性能基线 | 稳态约 9–10 tok/s(贪心+MTP);采样约 4–6 tok/s;常驻约 23–33GB 量级 |
|
||||||
|
|
||||||
|
依赖安装走清华 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` 调比例(0–1)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 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 槽 → 4,64 槽 → 3),超限会被钳制并打日志。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. 日常对话
|
||||||
|
|
||||||
|
1. 打开对话页(`http://localhost:8090`),顶部选 **Sparkle**
|
||||||
|
2. 首次:`配` → Sparkle tab:Base 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 | 32–160 | 常驻专家数,越大命中率越高、越占内存(每槽 ~160MB)。64GB + AUTOPIN:64 槽约 9 tok/s;**96 槽约 10.3 tok/s(推荐)**;≥128 近红线;**160 会卡死,勿用** | 应用后进程级重启(约 90 秒) |
|
||||||
|
| K(投机宽度) | 1–5 | MTP 草稿 token 数,默认 3;只影响速度 | 即时 |
|
||||||
|
| max_tokens | — | 单次回答最大长度 | 即时 |
|
||||||
|
|
||||||
|
日常 64 够用;长文档可试 96–128;系统开始 swap 就降回 64。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. 故障排查
|
||||||
|
|
||||||
|
| 现象 | 原因 | 处理 |
|
||||||
|
|------|------|------|
|
||||||
|
| 503 / engine reloading | 加载或重建中 | 等 60–90 秒后再试 |
|
||||||
|
| `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 在该模型上易幻觉。默认走采样(约 4–6 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` 调长度(默认 2048,0 关闭)。
|
||||||
|
|
||||||
|
**Q: 速度还能再快吗?**
|
||||||
|
A: 96 槽 + AUTOPIN + prefill 分块约 9–10 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:采样约 **4–6 tok/s**;贪心+MTP 约 **9–10 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 性能基线 |
|
||||||
61
benchmarks/reports/8bit-g64-baseline-2026-07-25.md
Normal file
61
benchmarks/reports/8bit-g64-baseline-2026-07-25.md
Normal 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
3
mlx_streaming/.gitignore
vendored
Normal file
@ -0,0 +1,3 @@
|
|||||||
|
.venv/
|
||||||
|
__pycache__/
|
||||||
|
*.pyc
|
||||||
0
mlx_streaming/__init__.py
Normal file
0
mlx_streaming/__init__.py
Normal file
329
mlx_streaming/cli.py
Normal file
329
mlx_streaming/cli.py
Normal 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 跨大跨度词表分散**的合成 prompt——MoE 路由依赖 token 内容,
|
||||||
|
分散的 id 会命中更多专家、更充分地预填专家池;较长 prompt 又能走通多块分块 prefill。
|
||||||
|
合成 id 直接走 ids_mode(不依赖 tokenizer),分块 prefill(chunk=2)保证长 prompt 也不抬高显存峰值。
|
||||||
|
预热失败不致命(直接吞掉异常),不影响后续真实生成。
|
||||||
|
"""
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from mlx_streaming.mtp.generate import mtp_generate
|
||||||
|
|
||||||
|
try:
|
||||||
|
vocab = int(model.model.embed_tokens.weight.shape[0])
|
||||||
|
n = min(64, vocab) # 预热 prompt 长度(兼顾覆盖与耗时)
|
||||||
|
step = max(1, vocab // n) # 在词表内均匀取样,最大化专家覆盖
|
||||||
|
ids = [(1 + i * step) % vocab for i in range(n)]
|
||||||
|
mtp_generate(model, drafter, tok, mx.array([ids]), 8,
|
||||||
|
K=args.k, ids_mode=True)
|
||||||
|
except Exception: # noqa: BLE001 预热仅为压首轮延迟,失败不应中断启动
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _encode_chat(tok, messages, tools=None):
|
||||||
|
"""把多轮对话按聊天模板编码成 token id 列表。tools(OpenAI 格式)交给模板渲染工具定义。"""
|
||||||
|
tmpl = getattr(tok, "chat_template", None)
|
||||||
|
if tmpl:
|
||||||
|
kwargs = {"add_generation_prompt": True}
|
||||||
|
if tools:
|
||||||
|
kwargs["tools"] = tools
|
||||||
|
out = tok.apply_chat_template(messages, **kwargs)
|
||||||
|
# 返回形态随 transformers 版本而变,三种都要兜住:
|
||||||
|
# transformers 4.x → list[int]
|
||||||
|
# transformers 5.x(默认) → BatchEncoding,即 {"input_ids": [[...]], ...}
|
||||||
|
# 部分实现 → str
|
||||||
|
# 只写 list(out) 会在 5.x 下取到字典的**键名** ["input_ids", "attention_mask"],
|
||||||
|
# 于是 prompt 退化成两个字符串:主生成路径(backend.py 的 ids)喂进去全是乱码,
|
||||||
|
# 且 _fixed_prefix_len 恒为 2 → head < 256 → 前缀快照永久失效。
|
||||||
|
# 单元测试全部 monkeypatch 掉了本函数,捕不到这个回归,故在此显式判型。
|
||||||
|
if isinstance(out, str):
|
||||||
|
return tok.encode(out)
|
||||||
|
ids = out if isinstance(out, (list, tuple)) else out["input_ids"]
|
||||||
|
if len(ids) > 0 and isinstance(ids[0], (list, tuple)):
|
||||||
|
ids = ids[0] # 批次维度
|
||||||
|
return [int(t) for t in ids]
|
||||||
|
# 无聊天模板:朴素拼接兜底。
|
||||||
|
text = ""
|
||||||
|
for m in messages:
|
||||||
|
text += f"{m['role']}: {m['content']}\n"
|
||||||
|
text += "assistant: "
|
||||||
|
return tok.encode(text)
|
||||||
|
|
||||||
|
|
||||||
|
def _eos_set(tok):
|
||||||
|
"""收集所有可能的结束符 token id。"""
|
||||||
|
eos = set()
|
||||||
|
ids = getattr(tok, "eos_token_ids", None)
|
||||||
|
if ids:
|
||||||
|
eos |= set(ids)
|
||||||
|
one = getattr(tok, "eos_token_id", None)
|
||||||
|
if one is not None:
|
||||||
|
eos.add(one)
|
||||||
|
return eos
|
||||||
|
|
||||||
|
|
||||||
|
def _truncate_eos(produced, eos):
|
||||||
|
"""遇到第一个结束符即截断(不含结束符本身)。"""
|
||||||
|
for i, t in enumerate(produced):
|
||||||
|
if t in eos:
|
||||||
|
return produced[:i]
|
||||||
|
return produced
|
||||||
|
|
||||||
|
|
||||||
|
_HELP = """可用命令:
|
||||||
|
/exit, /quit 退出
|
||||||
|
/reset 清空对话历史(保留 system)
|
||||||
|
/help 显示本帮助
|
||||||
|
直接输入文本即可对话。"""
|
||||||
|
|
||||||
|
|
||||||
|
def cmd_chat(args):
|
||||||
|
"""默认启动全屏 TUI;--plain 走纯文本 REPL;--demo 用假后端免模型预览界面。"""
|
||||||
|
if getattr(args, "plain", False):
|
||||||
|
return _chat_repl(args)
|
||||||
|
from mlx_streaming.tui import run_tui
|
||||||
|
if getattr(args, "demo", False):
|
||||||
|
# 免模型预览:秒开 TUI,假流式回答,用于验证界面/占位符/状态栏
|
||||||
|
from mlx_streaming.tui.backend import FakeBackend
|
||||||
|
demo = FakeBackend(
|
||||||
|
reply="这是 --demo 测试回答:界面、逐字流式、状态栏(token 数 / tok·s)"
|
||||||
|
"均为模拟,不加载模型。按 Esc 可中断,/help 看命令。",
|
||||||
|
delay=0.03)
|
||||||
|
return run_tui(demo, args)
|
||||||
|
from mlx_streaming.tui.backend import MLXBackend
|
||||||
|
return run_tui(MLXBackend(args), args)
|
||||||
|
|
||||||
|
|
||||||
|
def _chat_repl(args):
|
||||||
|
model, tok, drafter = _build_engine(args)
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from mlx_streaming.mtp.generate import mtp_generate
|
||||||
|
from mlx_streaming.tui.backend import _reuse_prefix_len
|
||||||
|
|
||||||
|
eos = _eos_set(tok)
|
||||||
|
base_messages = []
|
||||||
|
if args.system:
|
||||||
|
base_messages.append({"role": "system", "content": args.system})
|
||||||
|
messages = list(base_messages)
|
||||||
|
# 跨轮复用:持久化上一轮 main_cache 及其对应的 token 序列(与 MLXBackend 同机制)。
|
||||||
|
main_cache = None
|
||||||
|
cached_ids: list[int] = []
|
||||||
|
|
||||||
|
print("\n正在预热(编译 kernel + 填专家池)…", file=sys.stderr, flush=True)
|
||||||
|
_warmup(model, tok, drafter, args)
|
||||||
|
print("模型已就绪。输入 /help 查看命令,/exit 退出。", file=sys.stderr, flush=True)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
user = input("\n你 > ").strip()
|
||||||
|
except (EOFError, KeyboardInterrupt):
|
||||||
|
print()
|
||||||
|
break
|
||||||
|
if not user:
|
||||||
|
continue
|
||||||
|
if user in ("/exit", "/quit"):
|
||||||
|
break
|
||||||
|
if user == "/reset":
|
||||||
|
messages = list(base_messages)
|
||||||
|
main_cache, cached_ids = None, [] # 清历史同时弃用旧 cache
|
||||||
|
print("对话历史已清空。", file=sys.stderr)
|
||||||
|
continue
|
||||||
|
if user == "/help":
|
||||||
|
print(_HELP, file=sys.stderr)
|
||||||
|
continue
|
||||||
|
|
||||||
|
messages.append({"role": "user", "content": user})
|
||||||
|
ids = _encode_chat(tok, messages)
|
||||||
|
|
||||||
|
# 旧 cache 是本轮 prompt 的严格前缀时只 prefill 新增后缀,否则全量重建。
|
||||||
|
cached_len = (_reuse_prefix_len(cached_ids, ids)
|
||||||
|
if main_cache is not None else 0)
|
||||||
|
cur_cache = main_cache if cached_len else model.make_cache()
|
||||||
|
|
||||||
|
# EOS 提前停止:否则会空跑到 max_tokens,且 produced 会含 EOS 后垃圾 token,
|
||||||
|
# 使 cached_ids 与下轮编码前缀断裂、复用永不触发。
|
||||||
|
produced_all: list[int] = []
|
||||||
|
|
||||||
|
def _on_tokens(new_ids):
|
||||||
|
produced_all.extend(new_ids)
|
||||||
|
truncated = _truncate_eos(produced_all, eos)
|
||||||
|
return len(truncated) < len(produced_all)
|
||||||
|
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
produced, stats = mtp_generate(
|
||||||
|
model, drafter, tok, mx.array([ids]),
|
||||||
|
args.max_tokens, K=args.k, ids_mode=True, profile=args.stats,
|
||||||
|
on_tokens=_on_tokens, main_cache=cur_cache, cached_len=cached_len)
|
||||||
|
dt = time.perf_counter() - t0
|
||||||
|
|
||||||
|
# offset 对账:仅无 over-commit 时记录 cache 供下轮复用(见 MLXBackend.generate)。
|
||||||
|
if stats.get("resident_tokens") == len(ids) + len(produced) - 1:
|
||||||
|
main_cache = cur_cache
|
||||||
|
cached_ids = list(ids) + list(produced[:-1])
|
||||||
|
else:
|
||||||
|
main_cache, cached_ids = None, []
|
||||||
|
|
||||||
|
out_ids = _truncate_eos(produced, eos)
|
||||||
|
text = tok.decode(out_ids)
|
||||||
|
print(f"\n助手 > {text}")
|
||||||
|
messages.append({"role": "assistant", "content": text})
|
||||||
|
|
||||||
|
if args.stats:
|
||||||
|
tps = len(out_ids) / dt if dt > 0 else 0.0
|
||||||
|
print(f"[{len(out_ids)} tok, {tps:.1f} tok/s, "
|
||||||
|
f"accept_len={stats.get('avg_accept_len')}]", file=sys.stderr)
|
||||||
|
|
||||||
|
print("再见。", file=sys.stderr)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _add_chat_args(p):
|
||||||
|
p.add_argument("--model", default=config.model_path(), help="主模型路径(MLX 量化)")
|
||||||
|
p.add_argument("--expert-dir", default=config.expert_dir(),
|
||||||
|
help="拆分后的 per-expert 目录")
|
||||||
|
p.add_argument("--mtp-out", default=config.mtp_out(), help="MTP 权重文件")
|
||||||
|
p.add_argument("--qn-config", default=config.qn_config(),
|
||||||
|
help="Qwen3-Next 配置 JSON")
|
||||||
|
p.add_argument("-k", "--k", type=int, default=3, help="MTP 投机宽度(默认 3)")
|
||||||
|
p.add_argument("-n", "--max-tokens", type=int, default=4096,
|
||||||
|
help="每轮最多生成的新 token 数(默认 4096)")
|
||||||
|
p.add_argument("--expert-slots", type=int, default=32,
|
||||||
|
help="常驻专家池容量(默认 32,同时作为侧区行数默认)")
|
||||||
|
p.add_argument("--spec-slots", type=int, default=None,
|
||||||
|
help="侧区行数 POOL_SPEC_SLOTS(默认跟随 --expert-slots)")
|
||||||
|
p.add_argument("--system", default=None, help="可选 system 提示词")
|
||||||
|
p.add_argument("--stats", action="store_true",
|
||||||
|
help="每轮结束在 stderr 打印 token 数 / tok·s / 接受长度")
|
||||||
|
p.add_argument("--plain", action="store_true",
|
||||||
|
help="用纯文本 REPL,不启动全屏 TUI(终端不兼容/调试时用)")
|
||||||
|
p.add_argument("--demo", action="store_true",
|
||||||
|
help="免模型预览:用假后端秒开 TUI,验证界面/流式/状态栏")
|
||||||
|
p.set_defaults(func=cmd_chat)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_parser():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
prog="sparkle",
|
||||||
|
description="sparkle:Apple Silicon 上的流式 MoE + Qwen3-Next MTP 自投机推理")
|
||||||
|
sub = parser.add_subparsers(dest="cmd")
|
||||||
|
chat = sub.add_parser("chat", help="进入交互式多轮对话(MTP 自投机快路径)")
|
||||||
|
_add_chat_args(chat)
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv=None):
|
||||||
|
argv = list(sys.argv[1:] if argv is None else argv)
|
||||||
|
subcmds = {"chat"}
|
||||||
|
# 让 chat 成为默认子命令:不带子命令(或首参是 flag)时自动补上 chat;
|
||||||
|
# 但保留顶层 -h/--help 直接显示总帮助。
|
||||||
|
if not argv or (argv[0] not in subcmds and argv[0] not in ("-h", "--help")):
|
||||||
|
argv = ["chat"] + argv
|
||||||
|
args = _build_parser().parse_args(argv)
|
||||||
|
return args.func(args)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
298
mlx_streaming/config.py
Normal file
298
mlx_streaming/config.py
Normal file
@ -0,0 +1,298 @@
|
|||||||
|
"""集中配置:所有运行时开关从环境变量读取(env 名与默认值与历史一致,便于现有跑法)。
|
||||||
|
|
||||||
|
设计目的:把原先散落在 core/ 各文件的 `os.environ.get` 全部收口到这里,单点可查、
|
||||||
|
避免默认值漂移。读取均为运行时(每次调用读 env),与原行为一致——测试/probe 用
|
||||||
|
monkeypatch env 仍生效,热路径每 token 读 env 的成本与原先相同。
|
||||||
|
|
||||||
|
对「同一 env 名在不同调用方有不同默认」的项(如 CROSS_LAYER_PREFETCH_AHEAD:hook=0、
|
||||||
|
native 预取=1),accessor 暴露 `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")
|
||||||
|
# A2:GPU 重映射路径下 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) # 安全下限=2:MTP 每步 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
|
||||||
0
mlx_streaming/core/__init__.py
Normal file
0
mlx_streaming/core/__init__.py
Normal file
1
mlx_streaming/core/cache/__init__.py
vendored
Normal file
1
mlx_streaming/core/cache/__init__.py
vendored
Normal file
@ -0,0 +1 @@
|
|||||||
|
"""专家缓存模块:常驻池(resident_pool) / LRU+文件后端存储(expert_store) / blob 流式源(blob_loader)。"""
|
||||||
174
mlx_streaming/core/cache/autopin.py
vendored
Normal file
174
mlx_streaming/core/cache/autopin.py
vendored
Normal file
@ -0,0 +1,174 @@
|
|||||||
|
"""AUTOPIN:专家热度持久化 + 启动预热钉死。
|
||||||
|
|
||||||
|
两部分:
|
||||||
|
- **热度计数与持久化**:搭已有的物化点零同步计频——block host 路径的 flat、
|
||||||
|
acquire_gpu 的 LFU piggyback(全命中)与 miss 回退 flat;dual 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_POLICY≠lfu 时两条 GPU 快路径均不计——绝不为计数给
|
||||||
|
decode 热路径新增 GPU→host 同步。
|
||||||
|
"""
|
||||||
|
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
238
mlx_streaming/core/cache/blob_loader.py
vendored
Normal file
@ -0,0 +1,238 @@
|
|||||||
|
"""全流式 blob 专家源:并行 pread + F_NOCACHE + 即时物化为 MLX 量化数组。
|
||||||
|
|
||||||
|
字节布局由 prep/blob_layout.py 统一描述(单一真相源),支持两种格式:
|
||||||
|
- v1 affine(expert_blob_v1):每 proj 段序 [weight, scales, biases];weight uint32,
|
||||||
|
scales/biases 存 uint16 原始位、用时 .view(bfloat16)。
|
||||||
|
- v2 mxfp4(expert_blob_v2_mxfp4):每 proj 段序 [weight, scales](无 biases);
|
||||||
|
weight uint32,scales 为 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 fcntl:F_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=False:scales/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=False:scales/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.array(np.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] uint8(lazy)
|
||||||
|
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→uint32;v1 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: # uint8(mxfp4 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 列表中的下标(= local,reshape 前)。
|
||||||
|
"""
|
||||||
|
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
604
mlx_streaming/core/cache/expert_store.py
vendored
Normal 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 fcntl:F_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-source:decode 走 staging 双源命中计数(不写池的零拷贝命中)
|
||||||
|
self.staging_hits = 0
|
||||||
|
# 可选 blob miss-loader(STREAM_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 loader:blob 路径一次 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/biases→bf16);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:
|
||||||
|
"""批量 pin(AUTOPIN 预热用):语义同 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)。在 prefill→decode
|
||||||
|
边界调用一次即可把这块预算还给系统/池。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
|
||||||
116
mlx_streaming/core/cache/kv_quant_patch.py
vendored
Normal file
116
mlx_streaming/core/cache/kv_quant_patch.py
vendored
Normal 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
223
mlx_streaming/core/cache/quant_kv.py
vendored
Normal 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)。旋转正交,在注意力分数里自动抵消
|
||||||
|
(q、k 同旋转 → 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
|
||||||
635
mlx_streaming/core/cache/resident_pool.py
vendored
Normal file
635
mlx_streaming/core/cache/resident_pool.py
vendored
Normal 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++ 拥有,侧区异步直写落此 buffer(Route 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 buffer(spec/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: # 方案B:C++ 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]:
|
||||||
|
"""批量写一组 miss:stacked 为按 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 pool;pinned 专家不被 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)模式的 pin:C++ 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:
|
||||||
|
# 一次并行读本层所有 miss(8-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
237
mlx_streaming/core/cache/virtual_pool.py
vendored
Normal 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()
|
||||||
|
# 新前向开头排空上一前向提交的侧区 fill:C++-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>0):C++ demand_dual 唯一权威(真实区 ∪ 侧区读代单次 gather);
|
||||||
|
n_experts = layer_cap + spec_gens*spec_slots。
|
||||||
|
- 非 dual GPU-remap:acquire_gpu;n_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 永不可驱逐 → 可安放上限 cap−pinned。
|
||||||
|
# 超了 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++ 接管落池的字节等价铁证;发现不一致即池装错字节。
|
||||||
|
|
||||||
|
注:不以 local→expert 为判据(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-remap):flat 为 host 侧路由 id 列表。
|
||||||
|
|
||||||
|
返回 (pool_arrays, local, n_experts),与 block.py 原 host/fetch 分支逐元素等价:
|
||||||
|
- uniq <= cap:acquire(flat),local 为槽位;n_experts = layer_cap。
|
||||||
|
- uniq > cap:fetch(uniq_sorted),local 为 remap 到 [0,uniq) 连续索引;n_experts = uniq 数。
|
||||||
|
|
||||||
|
native demand(dual)下真实区槽状态归 C++ g_real,Python 的 _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)
|
||||||
1
mlx_streaming/core/linear_attn/__init__.py
Normal file
1
mlx_streaming/core/linear_attn/__init__.py
Normal file
@ -0,0 +1 @@
|
|||||||
|
"""线性注意力(Qwen3-Next gated-delta)相关的自定义 kernel。"""
|
||||||
226
mlx_streaming/core/linear_attn/gated_delta_multistate.py
Normal file
226
mlx_streaming/core/linear_attn/gated_delta_multistate.py
Normal file
@ -0,0 +1,226 @@
|
|||||||
|
"""逐步状态输出版 gated-delta Metal kernel(vendored 自 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-exact:mlx kernel 的时间维本就是线程内串行递归(不是分块并行),第 t 步
|
||||||
|
寄存器里的 `state` 就是「逐 token 走 t+1 步后的状态」,全程 fp32、同序、同融合。因此
|
||||||
|
一次 seq=T 调用产出的 `states_out[:, t]`,与逐 token(T 次 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
75
mlx_streaming/core/mem.py
Normal 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
|
||||||
1
mlx_streaming/core/moe/__init__.py
Normal file
1
mlx_streaming/core/moe/__init__.py
Normal file
@ -0,0 +1 @@
|
|||||||
|
"""MoE 模块:门控选专家(gate) / 专家计算(compute) / 自定义算子(custom_kernel) / 热路径块(block)。"""
|
||||||
431
mlx_streaming/core/moe/block.py
Normal file
431
mlx_streaming/core/moe/block.py
Normal file
@ -0,0 +1,431 @@
|
|||||||
|
"""MoE 热路径块:路由 → 选专家 → 取专家权重 → 计算 → 加权合并。
|
||||||
|
|
||||||
|
包含两种块:
|
||||||
|
- `StreamingMoeBlock`:包住原生 MoE 块,专家权重常驻(switch_mlp),只激活选中专家。
|
||||||
|
- `FileStreamingMoeBlock`:专家权重从磁盘按需加载(流式),是低内存推理的核心热路径,
|
||||||
|
集成常驻池 acquire、native-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:路由器常驻,专家计算改为只算选中专家。
|
||||||
|
|
||||||
|
若提供 store(LruExpertStore),可进一步把专家权重从磁盘按需加载;否则直接在
|
||||||
|
常驻的 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 barrier,C++ 读统一内存也救不了。
|
||||||
|
_ = 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 个
|
||||||
|
# token(verify_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 # 退化:无未归一化输入时用 x(norm 偏差,覆盖会差)
|
||||||
|
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-kk(argpartition,O(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
|
||||||
208
mlx_streaming/core/moe/compute.py
Normal file
208
mlx_streaming/core/moe/compute.py
Normal 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 是排序后的唯一全局专家 id,local 是与 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_k,n 不变 → 只在首次/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 用各自 bit;None 时三 proj 统一用 bits。
|
||||||
|
self.proj_bits = proj_bits or {
|
||||||
|
"gate_proj": bits, "up_proj": bits, "down_proj": bits}
|
||||||
|
# quant_mode:量化模式(affine / mxfp4);swiglu_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.Experts:up 双侧 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 projection;down 仍走 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)
|
||||||
266
mlx_streaming/core/moe/custom_kernel.py
Normal file
266
mlx_streaming/core/moe/custom_kernel.py
Normal 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
|
||||||
32
mlx_streaming/core/moe/gate.py
Normal file
32
mlx_streaming/core/moe/gate.py
Normal 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_miss≈0.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
|
||||||
504
mlx_streaming/core/moe/native_moe.py
Normal file
504
mlx_streaming/core/moe/native_moe.py
Normal 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 常驻 slot,hit 时复用 [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
|
||||||
1
mlx_streaming/core/prefetch/__init__.py
Normal file
1
mlx_streaming/core/prefetch/__init__.py
Normal file
@ -0,0 +1 @@
|
|||||||
|
"""专家预取模块:后台物化(bg_prefetch) / native staging(native_staging) / 跨层预测预取(cross_layer) / 模型 patch(patch)。"""
|
||||||
133
mlx_streaming/core/prefetch/bg_prefetch.py
Normal file
133
mlx_streaming/core/prefetch/bg_prefetch.py
Normal 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)
|
||||||
181
mlx_streaming/core/prefetch/cross_layer.py
Normal file
181
mlx_streaming/core/prefetch/cross_layer.py
Normal file
@ -0,0 +1,181 @@
|
|||||||
|
"""跨层专家预取:在 attention/GDN 之前用本层输入预测目标层 MoE 路由并提前取专家。
|
||||||
|
|
||||||
|
通过 monkeypatch `Qwen3NextDecoderLayer.__call__`,在真正进入 MoE 计算前用
|
||||||
|
gate_L(post_attention_layernorm_L(h)) 预测目标层专家(对齐 recall≈0.95 的探针口径),
|
||||||
|
把 top-(top_k*mult) 专家在当前层计算窗口内异步预取,藏住读盘/物化延迟。
|
||||||
|
|
||||||
|
支持多条预取后端(按环境开关择一):STREAM_BLOB_BG(后台物化进池)、STREAM_BLOB
|
||||||
|
(blob 字节预读)、STREAM_BLOB_LOADER(blob-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
|
||||||
229
mlx_streaming/core/prefetch/native_staging.py
Normal file
229
mlx_streaming/core/prefetch/native_staging.py
Normal file
@ -0,0 +1,229 @@
|
|||||||
|
"""native-fused-prefetch 的 miss→hit 落地:per-layer 环形 staging buffer + 池 promote。
|
||||||
|
|
||||||
|
机制(预读 IO 全在 GPU 完成回调里、主线程零 pread/零 .tolist/零 np→mx.array):
|
||||||
|
- submit(layer, inds_lazy):取该层环里的一块 buffer,调 prefetch_into_staging →
|
||||||
|
GPU 完成回调(C++)把专家字节 pread 进这块 buffer,并按 gen 原子记录 (expert→row)。
|
||||||
|
- 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 × stride。ring 默认从 4 降到 2(STAGING_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 membership,drain ≤len(cand))。
|
||||||
|
|
||||||
|
cand:预读候选专家 id(host 已知,通常 ≤budget)。
|
||||||
|
route_inds:本层真实路由 id(GPU 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]、记 (gen→buffer);C++ handler 按 gen 原子记 (expert→row);
|
||||||
|
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)
|
||||||
|
# 安全下限=2:MTP 每步 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):
|
||||||
|
"""方案B:inds_lazy 为"预测宽集合"(lazy uint32 [N],按门控分降序,N 可远大于 budget)。
|
||||||
|
|
||||||
|
resident:目标层提交时刻的常驻专家快照(host 侧)。C++ 回调按 resident 过滤、再按降序
|
||||||
|
取前 budget 个"缺口"专家 pread 进 buffer(budget=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_experts:GPU 重映射路径专用——host 没有现成 used 时,
|
||||||
|
用 route_used_subset 仅对预读候选做 GPU membership 现算 used(drain ≤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.array(C++ 直接建,供快路径合并)。"""
|
||||||
|
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):
|
||||||
|
# - uint32:weight,原样 reshape;
|
||||||
|
# - uint16:affine 的 scales/biases,view(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
|
||||||
56
mlx_streaming/core/prefetch/patch.py
Normal file
56
mlx_streaming/core/prefetch/patch.py
Normal 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。
|
||||||
|
|
||||||
|
store:FileExpertStore(所有 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
|
||||||
131
mlx_streaming/core/profiling.py
Normal file
131
mlx_streaming/core/profiling.py
Normal file
@ -0,0 +1,131 @@
|
|||||||
|
"""热路径计时与诊断埋点(从 streaming_moe 抽出,集中管理)。
|
||||||
|
|
||||||
|
- PROF:细粒度分段计时(STREAM_PROF=1),probe 读取。
|
||||||
|
- WINDOW_PROF:同层 submit→promote 的 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_s:block._native_fused_prefetch 里建预测图(gate(x)+argpartition[+predicted_set tolist])
|
||||||
|
# - submit_s :stg.submit(prefetch_into_staging 的 host 调用,注册 GPU 完成回调)
|
||||||
|
# - promote_s:promote 总时长,再细分:
|
||||||
|
# · take_s :prefetch_staging_take(C++ 锁读已就绪记录)
|
||||||
|
# · route_s:route_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=1,promote 之后 / acquire 之前统计本层真实路由的命中构成):
|
||||||
|
# - resident_hit:acquire 前已驻留(LRU 历史 + 本次 promote 命中)
|
||||||
|
# - miss_A_predicted:miss 且"在预测集里"——预测对了但没进池(budget 丢/时序没到/被驱逐)
|
||||||
|
# - miss_B_unpredicted:miss 且"不在预测集里"——预测器召回缺口(根本没预测到)
|
||||||
|
# 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:本层预测集;resident:acquire 前驻留集;
|
||||||
|
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
|
||||||
50
mlx_streaming/core/route_trace.py
Normal file
50
mlx_streaming/core/route_trace.py
Normal 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)
|
||||||
252
mlx_streaming/model_builder.py
Normal file
252
mlx_streaming/model_builder.py
Normal 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→hit):opt-in(NATIVE_FUSED_PREFETCH=1)。
|
||||||
|
# 经"promote 只写真实路由命中专家"修正后已是净正:易缓存基座上 +15.5% tok/s
|
||||||
|
# (demand 11.86→13.70,hit 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)
|
||||||
|
# 双源双缓冲:构造一个共享 VirtualPool(gen 跨层全局、每前向 +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 调度器 vpool(cutoff),让晚层预读更早发起。
|
||||||
|
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)
|
||||||
0
mlx_streaming/models/__init__.py
Normal file
0
mlx_streaming/models/__init__.py
Normal file
0
mlx_streaming/mtp/__init__.py
Normal file
0
mlx_streaming/mtp/__init__.py
Normal file
28
mlx_streaming/mtp/bench_verdict.py
Normal file
28
mlx_streaming/mtp/bench_verdict.py
Normal 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"
|
||||||
188
mlx_streaming/mtp/drafter.py
Normal file
188
mlx_streaming/mtp/drafter.py
Normal 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)
|
||||||
506
mlx_streaming/mtp/generate.py
Normal file
506
mlx_streaming/mtp/generate.py
Normal 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.py;drafter 接口见 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 prefix。main_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。
|
||||||
|
# - 可裁剪 cache(KVCache):trim 掉 rejected 后缀即精确。
|
||||||
|
# - 递归状态 cache(Qwen3-Next gated-delta 的 conv/ssm):verify 前向里用
|
||||||
|
# multistate kernel 捕获的 per-token checkpoint 与 baseline kernel 逐 bit 等价
|
||||||
|
# (见 mtp/kv_cache.py 与 tests/test_gated_delta_multistate.py),可精确直提交。
|
||||||
|
# commit_verified_prefix 一次处理混合 cache:KV 走 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
|
||||||
276
mlx_streaming/mtp/kv_cache.py
Normal file
276
mlx_streaming/mtp/kv_cache.py
Normal file
@ -0,0 +1,276 @@
|
|||||||
|
"""MTP 投机解码的 KV/递归状态 cache 校验机制:per-token checkpoint + 快照/恢复/提交。
|
||||||
|
|
||||||
|
speculative decoding 需要在「验证 K 个草稿」后只保留 accepted_len 个 token 对 cache 的
|
||||||
|
贡献。两类 cache 处理方式不同:
|
||||||
|
- 可裁剪 cache(KVCache):直接 trim 掉 rejected 后缀。
|
||||||
|
- 递归状态 cache(Qwen3-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 路径内部处理 GQA(hk_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
|
||||||
95
mlx_streaming/mtp/qwen3_next_mtp.py
Normal file
95
mlx_streaming/mtp/qwen3_next_mtp.py
Normal 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-bit。quantize=False 保持 bf16 全精度——草稿质量最高、接受率最高,但显存更大。
|
||||||
|
提高 bits(4→8)或 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
|
||||||
BIN
mlx_streaming/native_moe_ext.cpython-314-darwin.so
Executable file
BIN
mlx_streaming/native_moe_ext.cpython-314-darwin.so
Executable file
Binary file not shown.
0
mlx_streaming/prep/__init__.py
Normal file
0
mlx_streaming/prep/__init__.py
Normal file
32
mlx_streaming/prep/blob_layout.py
Normal file
32
mlx_streaming/prep/blob_layout.py
Normal 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
|
||||||
35
mlx_streaming/prep/download_retry.py
Normal file
35
mlx_streaming/prep/download_retry.py
Normal file
@ -0,0 +1,35 @@
|
|||||||
|
"""自动重试的断点续传下载:应对代理对大文件不稳定的情况。
|
||||||
|
|
||||||
|
反复调用 snapshot_download(huggingface_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()
|
||||||
131
mlx_streaming/prep/extract_mtp.py
Normal file
131
mlx_streaming/prep/extract_mtp.py
Normal 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()
|
||||||
87
mlx_streaming/prep/pack_blob_from_experts.py
Normal file
87
mlx_streaming/prep/pack_blob_from_experts.py
Normal 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: uint32;v1 的 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()
|
||||||
83
mlx_streaming/prep/pack_compute_buffers.py
Normal file
83
mlx_streaming/prep/pack_compute_buffers.py
Normal 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()
|
||||||
70
mlx_streaming/prep/pack_expert_bundles.py
Normal file
70
mlx_streaming/prep/pack_expert_bundles.py
Normal 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()
|
||||||
135
mlx_streaming/prep/pack_expert_ranges.py
Normal file
135
mlx_streaming/prep/pack_expert_ranges.py
Normal 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()
|
||||||
92
mlx_streaming/prep/repack_expert_blobs.py
Normal file
92
mlx_streaming/prep/repack_expert_blobs.py
Normal 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()
|
||||||
69
mlx_streaming/prep/split_experts.py
Normal file
69
mlx_streaming/prep/split_experts.py
Normal 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))
|
||||||
1
mlx_streaming/runtime/__init__.py
Normal file
1
mlx_streaming/runtime/__init__.py
Normal file
@ -0,0 +1 @@
|
|||||||
|
"""生产推理入口:基线/流式/投机解码等 run_* 启动脚本。"""
|
||||||
42
mlx_streaming/runtime/run_baseline.py
Normal file
42
mlx_streaming/runtime/run_baseline.py
Normal 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()
|
||||||
69
mlx_streaming/runtime/run_mmap.py
Normal file
69
mlx_streaming/runtime/run_mmap.py
Normal file
@ -0,0 +1,69 @@
|
|||||||
|
"""路线 A:lazy(+ 若支持则 mmap)加载,可选设常驻/缓存上限,量内存与速度。
|
||||||
|
|
||||||
|
环境探针结论(mlx-lm 0.31.3):load 支持 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()
|
||||||
409
mlx_streaming/runtime/run_mtp_spec.py
Normal file
409
mlx_streaming/runtime/run_mtp_spec.py
Normal 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底 − A−B−C−D)。返回 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 / p99(cap 下限的正确性依据)。
|
||||||
|
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()
|
||||||
116
mlx_streaming/runtime/run_spec.py
Normal file
116
mlx_streaming/runtime/run_spec.py
Normal file
@ -0,0 +1,116 @@
|
|||||||
|
"""投机解码 + 流式专家 验证:target=流式 Qwen3-30B,draft=常驻 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()
|
||||||
180
mlx_streaming/runtime/run_streaming.py
Normal file
180
mlx_streaming/runtime/run_streaming.py
Normal 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_k),worst-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()
|
||||||
9
mlx_streaming/server/__init__.py
Normal file
9
mlx_streaming/server/__init__.py
Normal 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"]
|
||||||
95
mlx_streaming/server/__main__.py
Normal file
95
mlx_streaming/server/__main__.py
Normal 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()
|
||||||
116
mlx_streaming/server/admin.py
Normal file
116
mlx_streaming/server/admin.py
Normal 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
104
mlx_streaming/server/app.py
Normal 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
|
||||||
338
mlx_streaming/server/openai.py
Normal file
338
mlx_streaming/server/openai.py
Normal 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
|
||||||
101
mlx_streaming/server/state.py
Normal file
101
mlx_streaming/server/state.py
Normal 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
|
||||||
9
mlx_streaming/tui/__init__.py
Normal file
9
mlx_streaming/tui/__init__.py
Normal 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
334
mlx_streaming/tui/app.py
Normal 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)
|
||||||
436
mlx_streaming/tui/backend.py
Normal file
436
mlx_streaming/tui/backend.py
Normal 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 表示需全量重建。
|
||||||
|
|
||||||
|
只在严格前缀时复用,是为了永远「只延续、不回退」cache——Qwen3-Next 的线性注意力递归态
|
||||||
|
无法裁剪回任意历史位置;而 detokenize→retokenize 不一致、/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"])
|
||||||
3
mlx_streaming/tui/banner.py
Normal file
3
mlx_streaming/tui/banner.py
Normal file
@ -0,0 +1,3 @@
|
|||||||
|
"""sparkle 顶栏 logo。保持单行,便于与模型信息拼在同一行。"""
|
||||||
|
|
||||||
|
LOGO = "▚ sparkle"
|
||||||
41
mlx_streaming/tui/styles.tcss
Normal file
41
mlx_streaming/tui/styles.tcss
Normal 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
49
models/README.md
Normal 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
32
native/ext/CMakeLists.txt
Normal 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
14
native/ext/Makefile
Normal file
@ -0,0 +1,14 @@
|
|||||||
|
# 生产 MLX 扩展构建:cmake 把 native_moe_mlx_ext.cpp 编成
|
||||||
|
# mlx_streaming/native_moe_ext$(EXT_SUFFIX).so(nanobind + 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
153
native/ext/bindings.cpp
Normal 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
34
native/ext/common.h
Normal 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;
|
||||||
664
native/ext/compute/fused_moe.cpp
Normal file
664
native/ext/compute/fused_moe.cpp
Normal 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});
|
||||||
|
}
|
||||||
23
native/ext/compute/fused_moe.h
Normal file
23
native/ext/compute/fused_moe.h
Normal 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
183
native/ext/io/bg_reader.cpp
Normal 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; // demand(route)读设 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
23
native/ext/io/bg_reader.h
Normal file
@ -0,0 +1,23 @@
|
|||||||
|
// 自由后台读线程(de-risk):脱离 GPU 完成回调、零 GIL,pread 进调用方 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
16
native/ext/io/blob_io.h
Normal 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;
|
||||||
|
}
|
||||||
58
native/ext/io/blob_load.cpp
Normal file
58
native/ext/io/blob_load.cpp
Normal 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>{});
|
||||||
|
}
|
||||||
7
native/ext/io/blob_load.h
Normal file
7
native/ext/io/blob_load.h
Normal 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
331
native/ext/pool/demand.cpp
Normal file
@ -0,0 +1,331 @@
|
|||||||
|
// [6] Phase 2 方案B:真实区槽状态 C++ 全接管(1 次同步版)。
|
||||||
|
// 复刻 Python ResidentExpertPool 的 _slot_of/_free/_freq + _choose_victim/_alloc_slot 语义:
|
||||||
|
// - free 优先(front pop,free 初始 [0,cap));free 空 → LFU 驱逐,受害者槽直接复用(不回 free)。
|
||||||
|
// - _choose_victim:candidates=插入序中 ∉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_slot(e 为 miss,不会已在 e2r):free 优先,否则 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 频次 bump(canonical:本次全部唯一专家各 +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);
|
||||||
|
// pass2:miss 分配槽(不落盘)。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
29
native/ext/pool/demand.h
Normal 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);
|
||||||
89
native/ext/pool/owned_pool.cpp
Normal file
89
native/ext/pool/owned_pool.cpp
Normal file
@ -0,0 +1,89 @@
|
|||||||
|
// [5] Route 3 Phase 1 底座:C++ 拥有的池 buffer。
|
||||||
|
// 用 mx::allocator::malloc 分配 buffer、no-op deleter 建叶子 mx.array(C++ 经 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*m,key-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 在两者之间重分配了池 buffer,raw 写入被落单。
|
||||||
|
uintptr_t array_data_ptr(const mx::array& a) {
|
||||||
|
mx::array b = a;
|
||||||
|
b.eval();
|
||||||
|
return reinterpret_cast<uintptr_t>(b.data<uint8_t>());
|
||||||
|
}
|
||||||
14
native/ext/pool/owned_pool.h
Normal file
14
native/ext/pool/owned_pool.h
Normal file
@ -0,0 +1,14 @@
|
|||||||
|
// Route 3 Phase 1 底座:C++ 拥有的池 buffer + 直写(替代 mx.zeros 建池 + 消费侧 MLX scatter)。
|
||||||
|
#pragma once
|
||||||
|
#include "../common.h"
|
||||||
|
|
||||||
|
// C++ 拥有的池 buffer 包成 mx.array(no-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);
|
||||||
389
native/ext/pool/side_region.cpp
Normal file
389
native/ext/pool/side_region.cpp
Normal 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 与所有池数组 buffer;idp/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=1):reserve 结束时审计侧区行账本自洽性 =====
|
||||||
|
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+3:pread 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()));
|
||||||
|
}
|
||||||
|
// 防御:各段字节之和必须等于 stride(blob 记录恰好是各段的拼接)。
|
||||||
|
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;
|
||||||
|
}
|
||||||
21
native/ext/pool/side_region.h
Normal file
21
native/ext/pool/side_region.h
Normal 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);
|
||||||
217
native/ext/prefetch/prefetch.cpp
Normal file
217
native/ext/prefetch/prefetch.cpp
Normal 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→hit:handler 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:
|
||||||
|
// 方案B:expert_ids 是"预测宽集合"(top-N,按门控分降序);resident 是目标层当前常驻专家。
|
||||||
|
// handler 在回调里过滤掉常驻、按降序取前 cap 个缺口 pread 进 staging(cap=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;
|
||||||
|
}
|
||||||
23
native/ext/prefetch/prefetch.h
Normal file
23
native/ext/prefetch/prefetch.h
Normal 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(方案B):expert_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
44
pyproject.toml
Normal file
@ -0,0 +1,44 @@
|
|||||||
|
[project]
|
||||||
|
name = "sparkle-engine"
|
||||||
|
version = "0.1.0"
|
||||||
|
description = "Sparkle:Apple 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"
|
||||||
Loading…
Reference in New Issue
Block a user