最佳化真實的 Snake 工作負載
所採用的 opt-in 路徑結合了 MLX 編譯、16 token 的長度桶和一個有界的已分詞問題字首快取。它不改變權重、不量化模型、不快取預測,也不在問題之間複用雙向編碼器的隱藏狀態。
在一次成對的完整迴圈測試中,隨附的最佳化路徑在 2,400 次移動中達到 75.40 次移動/秒,而同一次測試中的 eager 推理是 70.82 次移動/秒:1.065×,約 6.5%。兩者都是零死亡、2 次安全乾預,並在 2,400/2,400 步上執行動作完全一致。預熱並清空快取後,合併的活動 MLX 記憶體增長為 0 位元組。
啟用它
laya-snake --optimize
laya-snake --optimize --max-speed
通用 API 暴露相同的 opt-in 控制項:
import laya_mlx as laya
agent = laya.load(
"aac6fef/laya-multilingual-mlx",
compile=True,
pad_to_multiple=16,
cache_prompts=True,
)
這三個選項預設全部關閉,保留現有的 eager 行為和基準配置。編譯會針對輸入形狀特化;首次使用和新形狀會產生編譯成本。改變權重或模組結構後要構造新的 Agent。填充把序列長度向上取整到所請求的倍數,但不超過已配置的上下文上限。掩碼會排除填充的 token。
cache_prompts=True 為每個 Agent 保留至多 128 個不可變的 PreparedQuestion 字首,包括標記位置。快取鍵包含 tokenizer 身份、特殊 token、問題型別、有序的渲染選項、指令和字首預算。state 在每次 prepare 呼叫中只做一次清洗和分詞,然後與每個問題字首獨立拼接。問題變化會建立或選取合適的字首。每個問題仍然得到一次完整的模型前向。
形狀與準備的消融
這個 demo 每次移動問三個問題:方向、安全路線估計和食物可達性估計。因此它的批大小是 3,而不是 1。在 32 個取樣的真實錄制棋盤上:
- multilingual 的序列長度是 59、61、63 和 64。16 token 的倍數把它們全部放進一個 64 token 的桶。
- 英文的序列長度是 66、68、69 和 70,對映到一個 80 token 的桶。
- 把 multilingual 輸入從 64 填充到 96 會給逐 token 的工作量增加 50%;它不是在單獨短文本 API 固定樣本中暗示的那種 93 到 96 的小幅調整。
候選在每一個相同 state 內按輪換順序執行,且在此之前先訪問過每一個被測形狀一次。表中是同步的 Agent.predict 延遲,包含分詞和輸出轉換,排除規劃器/UI 工作和初始形狀預熱。
| 變體 | Multilingual p50 / p95(ms) | English p50 / p95(ms) |
|---|---|---|
| Eager | 9.12 / 10.21 | 21.83 / 26.73 |
| 僅字首複用 | 8.95 / 9.75 | 21.60 / 25.23 |
| 編譯,實際長度 | 8.67 / 9.66 | 21.27 / 25.99 |
| 編譯,填充到 96 | 10.92 / 11.72 | 25.66 / 29.88 |
| 編譯 + 字首複用 | 8.66 / 9.21 | 21.03 / 23.97 |
| 編譯,工作負載桶 | 8.66 / 9.55 | 21.78 / 26.12 |
| 編譯 + 桶 + 字首複用 | 8.56 / 9.29 | 21.51 / 24.16 |
全部七個候選在每個 checkpoint 的 32/32 個棋盤上都與 eager 的提議方向和執行方向一致。在這些樣本中,展示的四位小數機率和估計的最大差異是 0。這是有限樣本下舍入輸出的一致,不是聲稱內部浮點張量逐位相同。
消融使用有界的字首準備包裝器來篩選設計。下面的完整迴圈測試使用隨附的真實 compile、pad_to_multiple 和 cache_prompts API 實現。它的正確性測試還比較了 state 截斷、變化的標準、掩碼清洗和快取逐出之下的準備 ID 與標記。
隨附的最佳化路徑也通過了完整的真實 checkpoint 驗證矩陣:三個 checkpoint 在 FP32 和 FP16 下各 63/63 個選出答案一致(共 378/378)。校準機率誤差保持在現有容差內。每種配置還額外通過了 10 次有限、確定性的重複呼叫,測得活動記憶體增長 0 位元組。最佳化後的驗證資料。原始 eager 路徑每種配置 100 次重複的結果仍在原始基準報告中。
對英文來說,在這個樣本里用實際序列長度編譯比強制使用更大的桶更好。demo 預設用 multilingual;通用 API 使用者可以在啟用編譯和字首複用的同時把 pad_to_multiple=None。
完整迴圈的成對測試
四個種子,各 600 次移動,候選順序按種子交替。渲染包含 truecolor Rich 組裝和 ANSI 序列化,不含終端模擬器的繪製。每次移動都做一次全新的預測。結果來自一次本地成對執行。
| 種子 | Eager 移動/秒 | 最佳化後移動/秒 | 分數(兩者) | 動作一致數 |
|---|---|---|---|---|
| 101 | 68.60 | 78.04 | 20 | 600 |
| 102 | 70.07 | 78.62 | 24 | 600 |
| 103 | 76.75 | 85.62 | 23 | 600 |
| 104 | 68.50 | 63.15 | 16 | 600 |
最佳化路徑在一個種子上更慢。因此,6.5% 是這次實測執行中的合併改進,而不是對每一個回合或每一臺機器都保證的改進。更早的寬泛速度掃描和這次較晚的成對測試是不同的執行;不能把它們的絕對速率相減來聲稱加速。完整迴圈資料。
同時依據玩法和延遲來選擇模型
兩個 checkpoint 都跑了 20 個成對種子 × 300 次移動,checkpoint 順序按種子交替。每個回合都使用相同的初始 state、食物 RNG 種子、精簡特徵描述和迴圈護盾。視野長度是固定的;這些是 300 次移動後的分數,不是以死亡或填滿棋盤告終的完整對局。這個模型比較不包含終端渲染。
| Checkpoint | 存活 / 回合數 | 移動數 | 中位 / 平均分數 | 推理 p50 / p95(ms) | 干預次數 |
|---|---|---|---|---|---|
| laya | 20 / 20 | 6000 | 7.0 / 6.9 | 23.15 / 28.21 | 0 |
| multilingual | 20 / 20 | 6000 | 10.0 / 9.9 | 9.38 / 14.38 | 2 |
multilingual 取得了更多吃食進展,並在這個工作負載上更快,所以它仍是預設的 demo checkpoint。這個結果評估的是這套特徵輔助策略,不是通用推理質量或無輔助的 Snake 模型。每個回合與推理。
復現
uv run --extra demo python -m experiments.snake_runtime \
--output artifacts/snake/runtime-multilingual.json
uv run --extra demo python -m experiments.snake_runtime \
--model models/hub/laya-mlx --bucket 80 \
--output artifacts/snake/runtime-english.json
uv run --extra demo python -m benchmarks.snake_optimized \
--output artifacts/snake/optimized-paired.json
uv run --extra demo python -m benchmarks.snake_models \
--episodes 20 --steps 300 --output artifacts/snake/models.json
按順序執行 GPU 測量。先下載兩個本地模型目錄。已入庫的源錄製提供精確取樣的棋盤狀態。原始消融:multilingual、English、初始 96 token 試點。
更早的 compact 與 detailed 提示比較 在 64 個 state 上交替提示順序,並在丟棄前 8 次預熱迭代後發現中位數從 11.80 → 9.29 ms。它最初的臨時記錄沒有儲存棋盤快照,所以它是輔助證據,而不是主要的可復現消融。python -m benchmarks.snake_prompt 提供一個可復現的版本,會儲存 state、種子、完整決策和方法。
實現遵循 MLX 的官方編譯指南:使用長壽命的已編譯可呼叫物件和常規的形狀特化。它不會在依賴形狀的 Python 模型程式碼上使用 shapeless=True。當前文件是在 Context7 CLI 請求因網路錯誤失敗後,通過官方網站核對的。