Reduce prefill EAGLE memory pressure under CP shared KV

Prefill CP only needs the local hidden shard for DeepSeek NextN draft extend. The change adds a draft shared-KV path that captures target hidden locally, feeds only the CP-local slice into the draft model, and keeps draft KV writes/transfers on the same shared logical-to-physical page mapping as target KV.\n\nDebug logs are gated behind SGLANG_CP_DRAFT_SHARED_KV_DEBUG and cover scheduler pool selection, KV manager buffer registration, local physical writes, prefill sender filtering, transfer pages, and decode commit metadata so ETE runs can prove draft KV is sharded rather than full-concatenated on a prefill rank.\n\nConstraint: Prefill runs CP while decode remains DP, so prefill must avoid full hidden/KV materialization but decode still receives full logical KV pages.\nRejected: Keep draft extend on full hidden state | preserves correctness but wastes prefill memory and defeats CP shared-KV intent.\nRejected: Transfer draft KV with a separate mapping | target and draft pools share req_to_token logical indices, so duplicating mapping adds risk without benefit.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not remove the debug logs until ETE evidence confirms draft MLA/index writes and transfer pages are CP-sharded on all ranks.\nTested: Remote compileall for changed CP draft, transfer, scheduler, NSA index, MLA write, and EAGLE files.\nNot-tested: Full GLM-5 EAGLE ETE with SGLANG_CP_DRAFT_SHARED_KV_DEBUG=1 after this logging addition; local pytest intentionally not run.
This commit is contained in:
laoyao0822
2026-05-13 22:29:18 +08:00
parent 3fc7a5c18c
commit 99b669f8b9
16 changed files with 951 additions and 31 deletions
@@ -0,0 +1,474 @@
# NSA Prefill CPGLM-5 EAGLE/MTP Draft Shared-KV 支持计划
## 目标
本计划在当前 `cp-hicache-host` 分支上,为 GLM-5 / `GlmMoeDsaForCausalLM` 的 EAGLE/MTP draft prefill 恢复 CP shared-KV 语义:
```text
prefill target model: NSA CP + shared KV + narrow logits path
prefill draft model: NSA CP + shared KV + CP-local token compute
prefill -> decode: transfer target KV 与 draft KV 的 shared-KV shards
decode: 仍按当前 decode DP 路径执行,不引入 decode CP
```
直接目标是让 draft model 不再在 prefill 侧按每个 rank 保存完整 request KV / hidden,而是和 target model 一样按 CP owner shard 写 persistent KV,同时避免 EAGLE target hidden capture 破坏 target narrow output collection。
## 当前范围
### 纳入范围
当前只支持 GLM-5 / DSA 路径:
```text
GlmMoeDsaForCausalLM
-> draft architecture remap: DeepseekV3ForCausalLMNextN
-> python/sglang/srt/models/deepseek_nextn.py
-> NSA / MLA / CP shared-KV path
```
代码依据:
- `python/sglang/srt/configs/model_config.py`
- `GlmMoeDsaForCausalLM` draft 被 remap 到 `DeepseekV3ForCausalLMNextN`
- `python/sglang/srt/models/glm4_moe.py`
- `GlmMoeDsaForCausalLM` 继承 `DeepseekV2ForCausalLM`
- `python/sglang/srt/models/deepseek_nextn.py`
- `DeepseekV3ForCausalLMNextN` 当前已有 NSA CP metadata / split / all-gather 相关骨架。
### 不纳入范围
不在本阶段支持 GLM4 MHA/Radix draft
```text
Glm4MoeForCausalLM / Glm4MoeLiteForCausalLM
-> draft architecture remap: Glm4MoeForCausalLMNextN
-> RadixAttention / MHA path
```
原因:GLM4 draft 走 `Glm4MoeForCausalLMNextN`attention 是 MHA/Radix 路径,不是当前 GLM-5 DSA 的 NSA/MLA shared-KV 路径。MHA/Radix 的 KV write 当前主要是直接按 loc 写 `k_cache[indices] = k``v_cache[indices] = v`,需要单独补 owner filtering 与 logical->physical remap,和本阶段目标解耦。
## 当前问题
### 1. EAGLE target 强制 FULL hidden,禁用 target narrow path
文件:
- `python/sglang/srt/speculative/eagle_worker.py`
- `python/sglang/srt/speculative/multi_layer_eagle_worker.py`
- `python/sglang/srt/speculative/eagle_worker_v2.py`
- `python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py`
- `python/sglang/srt/models/deepseek_v2.py`
当前 EAGLE target extend 会执行:
```python
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL
```
这会让 target model 走 full hidden capture。target CP narrow path 的 gate 里又明确要求:
```python
forward_batch.capture_hidden_mode == CaptureHiddenMode.NULL
```
结果是:EAGLE 开启时 target 无法走 `cp_collect_last_token_hidden(...)` narrow path,而是回退到 `cp_all_gather_rerange_output(...)`
### 2. Deepseek nextn draft CP split 发生得偏晚
文件:
- `python/sglang/srt/models/deepseek_nextn.py`
当前 draft forward 近似流程:
```text
full input_ids -> embed full hidden
full target hidden -> hnorm
cat(full embed, full target hidden)
eh_proj full tokens
CP split hidden / positions
decoder local tokens
CP all-gather output hidden
logits processor
```
这有两个问题:
1. `eh_proj` 前仍然处理 full tokensprefill 显存与计算没有按 CP 缩小。
2. draft 输出仍执行 `cp_all_gather_rerange_output(...)`,没有复用 target narrow 思路。
### 3. Draft KV transfer 当前依赖“target/draft indices 总是共享”的隐式假设
文件:
- `python/sglang/srt/disaggregation/prefill.py`
- `python/sglang/srt/disaggregation/decode.py`
- `python/sglang/srt/managers/scheduler.py`
当前 prefill bootstrap queue 会把 draft pool 的 contiguous buffer 追加到 transfer buffer
```text
target token_to_kv_pool contiguous buffers
+ draft token_to_kv_pool contiguous buffers
```
但在 CP shared-KV 下必须确认:
- draft pool 是否也按 physical shared-KV capacity 初始化;
- draft KV write 是否写入 local physical shard
- prefill->decode transfer 是否覆盖所有 draft shard
- decode 侧是否按 logical loc 正确消费 draft KV。
## 设计原则
1. **不新增 decode CP。** 当前仅 prefill 使用 CPdecode 仍按现有 DP/TP/EAGLE 逻辑执行。
2. **不破坏 target narrow logits path。** EAGLE 需要 draft hidden,但不应该因此强制 target all-gather full hidden。
3. **draft 与 target 使用同一套 CP shared-KV ownership 语义。** logical loc 仍对外暴露,persistent KV pool 内部使用 physical shard。
4. **优先复用 NSA/MLA shared-KV 代码。** 不为 GLM-5 重做一套 MHA/Radix shared-KV writer。
5. **新路径默认 gated。** 在远端验证稳定前,通过环境变量打开,避免影响非 CP / 非 EAGLE / 非 GLM-5 路径。
## 建议环境变量
新增运行时开关:
```text
SGLANG_CP_DRAFT_SHARED_KV=0/1
```
语义:
- `0`:保持现有 EAGLE/MTP draft 行为。
- `1`:在满足以下条件时启用 GLM-5 draft CP shared-KV 路径:
- prefill CP enabled
- CP shared KV enabled
- draft architecture 是 `DeepseekV3ForCausalLMNextN`
- forward mode 是 extend/context-parallel extend
- `nsa_cp_metadata` 能够成功构造。
新增调试开关:
```text
SGLANG_CP_DRAFT_SHARED_KV_DEBUG=0/1
```
> P0 说明:上述环境变量用于 P1+ 功能路径 gated。P0 只落 probe log,不读取新增环境变量。
后续功能路径可打印调试日志,用于确认每个 rank 的:
- draft architecture
- token_to_kv_pool class
- logical / physical capacity
- full token count 与 local token count
- target hidden capture mode
- draft hidden input shape
- draft KV write loc 是否 local physical loc
- prefill transfer 中是否包含 draft KV buffers。
## 实施阶段
### P0:诊断日志
目标:先确认当前实际运行路径,不改变行为。P0 只添加无环境变量控制的 probe 日志;如日志显示现有路径不兼容,后续阶段再决定是否回退或 gated。
修改文件:
- `python/sglang/srt/speculative/eagle_worker.py`
- `python/sglang/srt/speculative/multi_layer_eagle_worker.py`
- `python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py`
- `python/sglang/srt/disaggregation/prefill.py`
- `python/sglang/srt/disaggregation/decode.py`
工作项:
1. 在 draft runner 初始化后打印一次 architecture 与 pool 类型。
2. 在 EAGLE target extend 打印 target capture mode 与 logits output hidden shape。
3. 在 draft extend 打印 draft input hidden shape、`input_ids` shape、`seq_lens_cpu`
4. 在 prefill/decode transfer init 打印 target/draft KV buffer 数量与 item size。
5. 在 scheduler disaggregation 初始化时打印 draft KV pool 来源。
验收:
- 不设置任何新增环境变量时行为不变。
- 日志能确认 GLM-5 draft path 是 `DeepseekV3ForCausalLMNextN`,不是 `Glm4MoeForCausalLMNextN`
- 日志能确认 prefill/decode transfer 是否包含 draft KV buffers。
### P0 观测结论
远端 P0 日志确认当前 GLM-5 EAGLE prefill 的实际路径:
- target prefill 使用 `CaptureHiddenMode.FULL`,所有 CP rank 都得到完整 `hidden_states=(prompt_tokens, hidden_size)`,因此 target narrow path 被禁用。
- draft architecture 是 `DeepseekV3ForCausalLMNextN`,属于本计划支持范围。
- draft runner 的 KV pool 已经是 `NSATokenToKVPool` + `CPSharedPagedTokenToKVPoolAllocator`,说明 memory pool 层具备 shared-KV 基础。
- draft extend 初始化阶段仍传入完整 `input_ids` / `spec_info.hidden_states` / `out_cache_loc`,即每个 CP rank 都准备处理完整 prompt tokens。
- `has_nsa_cp_metadata=False` 出现在 `ForwardBatch.init_new(...)` 之后、draft model forward 之前;metadata 当前由 `DeepseekV3ForCausalLMNextN.forward(...)` 内部构造,所以日志点早于 metadata 创建。真正的问题不是 metadata 永远缺失,而是 draft 在 `eh_proj` 前仍处理 full tokens。
因此 P1/P2 的修复重点是:target 输出侧先保留 CP-local draft hiddendraft 模型在 `eh_proj` 前完成 CP split,并在输出侧只收集 last-token hidden/logits,而不是 full hidden。
### P1target narrow logits 与 draft hidden 解耦
目标:EAGLE target 仍为 draft 提供 hidden,但不再通过 `CaptureHiddenMode.FULL` 禁用 target narrow path。
修改文件:
- `python/sglang/srt/speculative/eagle_worker.py`
- `python/sglang/srt/speculative/multi_layer_eagle_worker.py`
- `python/sglang/srt/model_executor/forward_batch_info.py`
- `python/sglang/srt/layers/logits_processor.py`
- `LogitsProcessorOutput` 定义在该文件,新增 `draft_hidden_states` side-channel 字段
- `python/sglang/srt/models/deepseek_v2.py`
设计:
增加一个独立 side-channel 字段,例如:
```python
draft_hidden_states: Optional[torch.Tensor]
```
不要新增 `CaptureHiddenMode` enum 值。原因是 `CaptureHiddenMode` 被 cuda graph / piecewise runner / schedule batch 复用,新增 enum 容易改变 graph capture 分支。
目标数据流:
```text
target model local hidden after final norm
-> for logits: cp_collect_last_token_hidden(...) narrow path
-> for draft: keep CP-local hidden in draft_hidden_states side-channel
```
EAGLE target extend 获取:
```text
logits_output.next_token_logits # narrow logits path
logits_output.draft_hidden_states # CP-local hidden for draft
```
验收:
- EAGLE 开启且 `SGLANG_CP_DRAFT_SHARED_KV=1` 时,target 不再强制 `CaptureHiddenMode.FULL`
- `DeepseekV2Model._should_use_narrow_output_path(...)` 对 target prefill 返回 true。
- draft hidden side-channel shape 是当前 CP rank local token 数,而不是 full prompt token 数。
### P2Deepseek nextn draft 输入改为 CP-local
目标:draft model 在 `eh_proj` 前就按 CP local tokens 运行,并在 prefill 输出侧只收集 last-token hidden/logits。
修改文件:
- `python/sglang/srt/models/deepseek_nextn.py`
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `python/sglang/srt/layers/logits_processor.py`
- `python/sglang/srt/speculative/eagle_worker.py`
- `python/sglang/srt/speculative/multi_layer_eagle_worker.py`
当前顺序:
```text
embed full input
cat full target hidden
eh_proj full tokens
CP split
decoder local tokens
CP all-gather full output hidden
```
目标顺序:
```text
prepare nsa_cp_metadata from full logical token layout
split input_ids / positions with split_list + zigzag_index
use target draft_hidden_states side-channel if already CP-local; otherwise split full target hidden
embed local input
eh_proj local tokens
run draft decoder local tokens
collect only last-token hidden via cp_collect_last_token_hidden for logits/capture
```
需要明确的约束:
- `input_ids``positions``spec_info.hidden_states` 必须使用同一套 CP split;如果 P1 已经提供 CP-local hidden,则不能重复 split hidden。
- `forward_batch.out_cache_loc` 对外仍保持 full logical loc 语义;KV write 时继续通过 shared-KV helper 计算本 rank local logical loc,再 remap 到 physical loc。不要把 `forward_batch.out_cache_loc` 直接改成本地 loc,否则 radix / PD transfer 语义会被破坏。
- `LogitsProcessor` 的 compact hidden 分支需要支持 `CaptureHiddenMode.LAST`,否则 `cp_collect_last_token_hidden(...)``hidden_states` 会被丢弃,`capture_for_decode(...)` 无法拿到 draft hidden。
- 如果 `can_cp_split(...)` 失败,必须 fallback 到现有 full-token 路径并打印 rate-limited fallback log。
验收:
- debug log 中 draft local token 数约等于 full token 数 / cp_size。
- `eh_proj` 输入 shape 使用 local token 数。
- draft prefill 不再执行 full `cp_all_gather_rerange_output(...)`
- `capture_for_decode(...)` 能拿到 `(batch, hidden)` 的 last-token hidden。
- fallback 原因可观测,例如:`no_nsa_cp_metadata``unsupported_arch``non_extend_mode``cp_shared_kv_disabled`
### P3:输出 narrow/local collection 验证与清理
目标:P2 已把 draft 输出 narrow/local collection 纳入实现;P3 只做验证、清理旧 fallback 以及补充 benchmark/profile 证据。
修改文件:
- `python/sglang/srt/models/deepseek_nextn.py`
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `python/sglang/srt/layers/logits_processor.py`
当前:
```python
hidden_states = cp_all_gather_rerange_output(...)
```
目标:
```text
只收集 draft logits / capture_for_decode 需要的 last-token hidden
```
可以复用 target 的思路:
```python
cp_collect_last_token_hidden(hidden_states, forward_batch, cp_size)
```
但要先确认 `EagleDraftInput.capture_hidden_mode = CaptureHiddenMode.LAST` 对 draft logits processor 的需求:
- `next_token_logits` 需要每个 request 的最后 token hidden
- `capture_for_decode(...)` 需要 draft decode 后续使用的 hidden/topk 信息;
- 不需要完整 prompt hidden。
验收:
- draft extend 不再执行 full `cp_all_gather_rerange_output(...)`
- EAGLE draft next token logits 与关闭新路径时保持一致或在可接受数值误差内一致。
- `capture_for_decode(...)` 能正常生成后续 speculative decode 所需状态。
### P4adraft KV transfer / physical shard debug 日志
目标:在 `SGLANG_CP_DRAFT_SHARED_KV_DEBUG=1` 下补齐 draft KV shared-KV 验证闭环,不改变默认行为。
新增日志点:
- scheduler disaggregation inittarget/draft KV pool class、size、page_size、layer 范围;
- prefill/decode KV managertarget/draft contiguous buffer 数量、lens、item_lens、draft split point
- CP shared KV writelogical→physical remap、MLA KV write path、index K/scale write path
- prefill send chunklogical page_indices、state_indices、是否存在 draft pool
- mooncake register/sender/transfer worker:注册 buffer 数、CP filter 后 pages、logical positions、decode dst pages
- decode prealloc/commitdst pages、metadata index、spec metadata shape。
这些日志用于确认:draft KV 没有拼 full KV 到某个 prefill rank,而是复用 CP shared-KV owner page filter,把各 rank owned physical pages 传到 decode full KV pool。
### P4draft KV shared-KV physical write 与 transfer 验证
目标:确认 draft persistent KV 与 target 一样是 CP owner-sharded physical KV。
修改文件:
- `python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py`
- `python/sglang/srt/layers/attention/nsa_backend.py`
- `python/sglang/srt/layers/attention/nsa/nsa_indexer.py`
- `python/sglang/srt/disaggregation/prefill.py`
- `python/sglang/srt/disaggregation/decode.py`
验证重点:
1. draft `NSATokenToKVPool` 是否用 physical capacity 初始化。
2. draft MLA KV write 是否走 `_maybe_filter_shared_mla_kv_write(...)`
3. draft index K/scale write 是否走 `get_cp_shared_kv_local_physical_out_cache_loc(...)`
4. prefill transfer 是否包含 draft physical shards。
5. decode 是否能按 logical loc 使用 transferred draft KV。
验收:
- 每个 CP rank 只写 owner shard 对应的 physical loc。
- draft KV pool memory 不再按 full logical token 数乘以 cp_size 膨胀。
- prefill->decode ETE 可通过,且 EAGLE speculative decode 正常。
## 远端验证计划
所有运行验证在远端容器执行,不在本地执行。
推荐远端环境:
```text
host: g0034 / ubuntu@10.20.32.34
container: sglang-glm5-dev-2
code dir: /sgl-workspace/sglang-tai
```
验证分层:
### 1. 静态检查
```bash
python -m compileall python/sglang/srt/speculative python/sglang/srt/models python/sglang/srt/layers/attention/nsa
```
### 2. targeted pytest
优先跑和 CP split / schedule copy / speculative 相关的小测试。若没有现成测试,需要新增 narrow unit test 覆盖:
- `ScheduleBatch.copy()` 保留 `spec_info`
- CP split 对 `input_ids``positions``hidden_states` 一致;
- draft path fallback reason 可观测。
### 3. GLM-5 prefill CP + decode DP ETE
启动条件:
```bash
SGLANG_CP_DRAFT_SHARED_KV=1
SGLANG_CP_DRAFT_SHARED_KV_DEBUG=1
SGLANG_CP_SHARED_KV_FUSED_MLA_STORE=1
SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE=1
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER=1
```
观察:
- target narrow path 命中;
- draft architecture 是 `DeepseekV3ForCausalLMNextN`
- draft local token 数约等于 full token 数 / cp_size
- draft KV write 使用 local physical loc
- prefill transfer 包含 draft KV
- decode 侧没有 draft KV missing / bootstrap timeout / speculative 状态异常。
### 4. profile 验证
Nsight / 日志关注点:
- target logits 前不再出现 full hidden all-gather
- draft `eh_proj` 和 decoder token 数下降;
- draft 输出不再 full all-gather
- draft KV materialize/write 不再每 rank 写 full request
- prefill 显存下降,尤其 draft KV pool 和 logits/hidden buffer 部分。
## 风险与处理
### 风险 1side-channel hidden 与 logits processor 输出结构耦合
处理:不复用 `hidden_states` 字段表达两种语义,新增明确字段 `draft_hidden_states`。现有 EAGLE 非 CP 路径继续使用原字段。
### 风险 2draft output narrow 后 `capture_for_decode(...)` 缺状态
处理:先在 P3 前单独梳理 `capture_for_decode(...)` 依赖字段,只收集它真正需要的 last-token hidden/topk,不保留 full prompt hidden。
### 风险 3draft KV transfer 依赖 target/draft loc 完全一致
处理:P4 只在日志与 ETE 确认一致后打开默认路径。若不一致,增加显式 draft logical loc mapping,不继续依赖注释里的隐式假设。
### 风险 4GLM4 MHA/Radix 被误启用
处理:`SGLANG_CP_DRAFT_SHARED_KV=1` 下如果 draft architecture 不是 `DeepseekV3ForCausalLMNextN`,直接 fallback 并打印:
```text
CP draft shared KV fallback: unsupported draft architecture
```
## 完成标准
本阶段完成需要同时满足:
1. GLM-5 EAGLE/MTP prefill CP ETE 通过。
2. target EAGLE 开启时仍命中 CP narrow logits path。
3. draft model 在 `eh_proj` 前已经使用 CP-local token / hidden。
4. draft persistent KV 使用 CP shared-KV physical shard,而不是每 rank full KV。
5. prefill->decode draft KV transfer 正常。
6. fallback 路径安全、可观测,非 GLM-5 draft 不误走新路径。