From 99b669f8b91e336c4dc34d739596c424a7d1d9cc Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Wed, 13 May 2026 22:29:18 +0800 Subject: [PATCH] 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. --- ...nsa_prefill_cp_eagle_mtp_shared_kv_plan.md | 474 ++++++++++++++++++ python/sglang/srt/disaggregation/decode.py | 84 ++++ .../srt/disaggregation/mooncake/conn.py | 48 ++ python/sglang/srt/disaggregation/prefill.py | 74 +++ python/sglang/srt/environ.py | 2 + .../srt/layers/attention/nsa/nsa_indexer.py | 10 + .../sglang/srt/layers/attention/nsa/utils.py | 29 ++ python/sglang/srt/layers/logits_processor.py | 10 +- python/sglang/srt/managers/schedule_batch.py | 2 + python/sglang/srt/managers/scheduler.py | 26 + python/sglang/srt/mem_cache/common.py | 20 +- .../srt/model_executor/forward_batch_info.py | 3 + .../attention_forward_methods/forward_mla.py | 24 + python/sglang/srt/models/deepseek_nextn.py | 107 +++- python/sglang/srt/models/deepseek_v2.py | 9 +- python/sglang/srt/speculative/eagle_worker.py | 60 ++- 16 files changed, 951 insertions(+), 31 deletions(-) create mode 100644 docs/advanced_features/nsa_prefill_cp_eagle_mtp_shared_kv_plan.md diff --git a/docs/advanced_features/nsa_prefill_cp_eagle_mtp_shared_kv_plan.md b/docs/advanced_features/nsa_prefill_cp_eagle_mtp_shared_kv_plan.md new file mode 100644 index 000000000..85b42a3bf --- /dev/null +++ b/docs/advanced_features/nsa_prefill_cp_eagle_mtp_shared_kv_plan.md @@ -0,0 +1,474 @@ +# NSA Prefill CP:GLM-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 tokens,prefill 显存与计算没有按 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 使用 CP,decode 仍按现有 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 hidden;draft 模型在 `eh_proj` 前完成 CP split,并在输出侧只收集 last-token hidden/logits,而不是 full hidden。 + +### P1:target 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 数。 + +### P2:Deepseek 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 所需状态。 + +### P4a:draft KV transfer / physical shard debug 日志 + +目标:在 `SGLANG_CP_DRAFT_SHARED_KV_DEBUG=1` 下补齐 draft KV shared-KV 验证闭环,不改变默认行为。 + +新增日志点: + +- scheduler disaggregation init:target/draft KV pool class、size、page_size、layer 范围; +- prefill/decode KV manager:target/draft contiguous buffer 数量、lens、item_lens、draft split point; +- CP shared KV write:logical→physical remap、MLA KV write path、index K/scale write path; +- prefill send chunk:logical page_indices、state_indices、是否存在 draft pool; +- mooncake register/sender/transfer worker:注册 buffer 数、CP filter 后 pages、logical positions、decode dst pages; +- decode prealloc/commit:dst 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。 + +### P4:draft 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 部分。 + +## 风险与处理 + +### 风险 1:side-channel hidden 与 logits processor 输出结构耦合 + +处理:不复用 `hidden_states` 字段表达两种语义,新增明确字段 `draft_hidden_states`。现有 EAGLE 非 CP 路径继续使用原字段。 + +### 风险 2:draft output narrow 后 `capture_for_decode(...)` 缺状态 + +处理:先在 P3 前单独梳理 `capture_for_decode(...)` 依赖字段,只收集它真正需要的 last-token hidden/topk,不保留 full prompt hidden。 + +### 风险 3:draft KV transfer 依赖 target/draft loc 完全一致 + +处理:P4 只在日志与 ETE 确认一致后打开默认路径。若不一致,增加显式 draft logical loc mapping,不继续依赖注释里的隐式假设。 + +### 风险 4:GLM4 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 不误走新路径。 diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 5f6ac05e5..a3349a9a3 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -77,6 +77,42 @@ if TYPE_CHECKING: CLIP_MAX_NEW_TOKEN = envs.SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION.get() +def _cp_draft_shared_kv_debug(message: str, *args) -> None: + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + logger.info("[CP_DRAFT_SHARED_KV] " + message, *args) + + +def _seq_summary(values) -> str: + if values is None: + return "None" + try: + size = len(values) + except TypeError: + return str(values) + if size == 0: + return "size=0" + try: + head = list(values[: min(8, size)]) + except TypeError: + head = list(values)[: min(8, size)] + try: + min_val = min(values) + max_val = max(values) + return f"size={size} min={min_val} max={max_val} head={head}" + except (TypeError, ValueError): + return f"size={size} head={head}" + + +def _pool_summary(pool) -> str: + if pool is None: + return "None" + parts = [pool.__class__.__name__] + for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"): + if hasattr(pool, attr): + parts.append(f"{attr}={getattr(pool, attr)}") + return " ".join(parts) + + def _kv_locs_to_page_indices_cpu( kv_locs: torch.Tensor, page_size: int, @@ -321,16 +357,39 @@ class DecodePreallocQueue: kv_data_ptrs, kv_data_lens, kv_item_lens = ( self.token_to_kv_pool.get_contiguous_buf_infos() ) + target_kv_buffer_count = len(kv_data_ptrs) + draft_kv_data_lens = [] + draft_kv_item_lens = [] + draft_kv_buffer_count = 0 if self.draft_token_to_kv_pool is not None: # We should also transfer draft model kv cache. The indices are # always shared with a target model. draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = ( self.draft_token_to_kv_pool.get_contiguous_buf_infos() ) + draft_kv_buffer_count = len(draft_kv_data_ptrs) kv_data_ptrs += draft_kv_data_ptrs kv_data_lens += draft_kv_data_lens kv_item_lens += draft_kv_item_lens + kv_args.draft_kv_buffer_start = target_kv_buffer_count + kv_args.draft_kv_buffer_count = draft_kv_buffer_count + _cp_draft_shared_kv_debug( + "decode_kv_manager cp_rank=%s target_pool=(%s) draft_pool=(%s) " + "target_bufs=%s draft_bufs=%s total_bufs=%s target_lens=%s " + "draft_lens=%s target_item_lens=%s draft_item_lens=%s", + self.tp_rank, + _pool_summary(self.token_to_kv_pool), + _pool_summary(self.draft_token_to_kv_pool), + target_kv_buffer_count, + draft_kv_buffer_count, + len(kv_data_ptrs), + _seq_summary(kv_data_lens[:target_kv_buffer_count]), + _seq_summary(draft_kv_data_lens), + _seq_summary(kv_item_lens[:target_kv_buffer_count]), + _seq_summary(draft_kv_item_lens), + ) + kv_args.kv_data_ptrs = kv_data_ptrs kv_args.kv_data_lens = kv_data_lens kv_args.kv_item_lens = kv_item_lens @@ -755,6 +814,19 @@ class DecodePreallocQueue: self.req_to_metadata_buffer_idx_allocator.alloc() ) assert decode_req.metadata_buffer_index is not None + _cp_draft_shared_kv_debug( + "decode_prealloc rid=%s room=%s origin_tokens=%s fill_tokens=%s " + "page_size=%s pages=%s state_pages=%s metadata_idx=%s has_draft_pool=%s", + decode_req.req.rid, + decode_req.req.bootstrap_room, + origin_input_len, + len(kv_loc), + page_size, + _seq_summary(page_indices), + _seq_summary(state_indices), + decode_req.metadata_buffer_index, + self.draft_token_to_kv_pool is not None, + ) decode_req.kv_receiver.init( page_indices, decode_req.metadata_buffer_index, state_indices ) @@ -1008,6 +1080,18 @@ class DecodeTransferQueue: decode_req.req.output_topk_index = output_topk_index decode_req.req.hidden_states_tensor = output_hidden_states + _cp_draft_shared_kv_debug( + "decode_transfer_commit rid=%s room=%s metadata_idx=%s cached_tokens=%s " + "topk_p_shape=%s topk_index_shape=%s hidden_shape=%s", + decode_req.req.rid, + decode_req.req.bootstrap_room, + idx, + decode_req.req.cached_tokens, + tuple(output_topk_p.shape) if output_topk_p is not None else None, + tuple(output_topk_index.shape) if output_topk_index is not None else None, + tuple(output_hidden_states.shape) if output_hidden_states is not None else None, + ) + if decode_req.req.return_logprob: decode_req.req.output_token_logprobs_val.append( output_token_logprobs_val[0].item() diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index f3fe7f538..3d3a90021 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -51,6 +51,17 @@ def _cp_shared_debug_log(key: str, message: str, *args, limit: int = 64) -> None logger.info("[CP_SHARED_KV_DEBUG] " + message, *args) +def _cp_draft_shared_kv_debug(message: str, *args, limit: int = 64) -> None: + if not envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + return + key = "draft:" + message.split(" ", 1)[0] + count = _CP_SHARED_DEBUG_COUNTS.get(key, 0) + if count >= limit: + return + _CP_SHARED_DEBUG_COUNTS[key] = count + 1 + logger.info("[CP_DRAFT_SHARED_KV] " + message, *args) + + def _np_summary(arr) -> str: if arr is None: return "None" @@ -246,6 +257,17 @@ class MooncakeKVManager(CommonKVManager): def register_buffer_to_engine(self): # Batch register KV data buffers if self.kv_args.kv_data_ptrs and self.kv_args.kv_data_lens: + _cp_draft_shared_kv_debug( + "register_buffers mode=%s cp_rank=%s total_kv_bufs=%s " + "draft_start=%s draft_count=%s kv_lens=%s kv_item_lens=%s", + self.disaggregation_mode, + self.attn_cp_rank, + len(self.kv_args.kv_data_ptrs), + getattr(self.kv_args, "draft_kv_buffer_start", None), + getattr(self.kv_args, "draft_kv_buffer_count", None), + _np_summary(self.kv_args.kv_data_lens), + _np_summary(self.kv_args.kv_item_lens), + ) self.engine.batch_register( self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens ) @@ -847,6 +869,16 @@ class MooncakeKVManager(CommonKVManager): chunked_dst_kv_indice = req.dst_kv_indices[ kv_chunk.index_slice ] + _cp_draft_shared_kv_debug( + "transfer_pages cp_rank=%s room=%s prefill_pages=%s " + "logical_positions=%s dst_pages=%s is_last=%s", + self.attn_cp_rank, + kv_chunk.room, + _np_summary(kv_chunk.prefill_kv_indices), + _np_summary(kv_chunk.logical_page_positions), + _np_summary(chunked_dst_kv_indice), + kv_chunk.is_last_chunk, + ) if envs.SGLANG_DEBUG_CP_SHARED_KV.get(): _cp_shared_debug_log( "transfer_worker_kv", @@ -1275,6 +1307,22 @@ class MooncakeKVSender(CommonKVSender): _np_summary(state_logical_page_positions), is_last_chunk, ) + _cp_draft_shared_kv_debug( + "sender_filter cp_rank=%s room=%s page_start=%s orig_kv_pages=%s " + "filtered_kv_pages=%s kv_positions=%s orig_state_pages=%s " + "filtered_state_pages=%s state_positions=%s is_last=%s draft_bufs=%s", + self.kv_mgr.attn_cp_rank, + self.bootstrap_room, + chunk_page_start, + _np_summary(orig_kv_indices), + _np_summary(kv_indices), + _np_summary(logical_page_positions), + _np_summary(orig_state_indices), + _np_summary(state_indices), + _np_summary(state_logical_page_positions), + is_last_chunk, + getattr(self.kv_mgr.kv_args, "draft_kv_buffer_count", None), + ) # Special handling for cp elif self.kv_mgr.enable_all_cp_ranks_for_transfer: kv_indices, index_slice = filter_kv_indices_for_cp_rank( diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index d4a5539df..a04489e17 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -27,6 +27,7 @@ from typing import TYPE_CHECKING, List, Optional import torch from sglang.srt.disaggregation.base import KVPoll +from sglang.srt.environ import envs from sglang.srt.disaggregation.common.conn import CommonKVManager from sglang.srt.disaggregation.utils import ( FAKE_BOOTSTRAP_HOST, @@ -61,6 +62,42 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +def _cp_draft_shared_kv_debug(message: str, *args) -> None: + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + logger.info("[CP_DRAFT_SHARED_KV] " + message, *args) + + +def _seq_summary(values) -> str: + if values is None: + return "None" + try: + size = len(values) + except TypeError: + return str(values) + if size == 0: + return "size=0" + try: + head = list(values[: min(8, size)]) + except TypeError: + head = list(values)[: min(8, size)] + try: + min_val = min(values) + max_val = max(values) + return f"size={size} min={min_val} max={max_val} head={head}" + except (TypeError, ValueError): + return f"size={size} head={head}" + + +def _pool_summary(pool) -> str: + if pool is None: + return "None" + parts = [pool.__class__.__name__] + for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"): + if hasattr(pool, attr): + parts.append(f"{attr}={getattr(pool, attr)}") + return " ".join(parts) + + def _kv_locs_to_page_indices_cpu( kv_locs: torch.Tensor, page_size: int, @@ -154,16 +191,39 @@ class PrefillBootstrapQueue: self.token_to_kv_pool.get_contiguous_buf_infos() ) + target_kv_buffer_count = len(kv_data_ptrs) + draft_kv_data_lens = [] + draft_kv_item_lens = [] + draft_kv_buffer_count = 0 if self.draft_token_to_kv_pool is not None: # We should also transfer draft model kv cache. The indices are # always shared with a target model. draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = ( self.draft_token_to_kv_pool.get_contiguous_buf_infos() ) + draft_kv_buffer_count = len(draft_kv_data_ptrs) kv_data_ptrs += draft_kv_data_ptrs kv_data_lens += draft_kv_data_lens kv_item_lens += draft_kv_item_lens + kv_args.draft_kv_buffer_start = target_kv_buffer_count + kv_args.draft_kv_buffer_count = draft_kv_buffer_count + _cp_draft_shared_kv_debug( + "prefill_kv_manager cp_rank=%s target_pool=(%s) draft_pool=(%s) " + "target_bufs=%s draft_bufs=%s total_bufs=%s target_lens=%s " + "draft_lens=%s target_item_lens=%s draft_item_lens=%s", + self.tp_rank, + _pool_summary(self.token_to_kv_pool), + _pool_summary(self.draft_token_to_kv_pool), + target_kv_buffer_count, + draft_kv_buffer_count, + len(kv_data_ptrs), + _seq_summary(kv_data_lens[:target_kv_buffer_count]), + _seq_summary(draft_kv_data_lens), + _seq_summary(kv_item_lens[:target_kv_buffer_count]), + _seq_summary(draft_kv_item_lens), + ) + kv_args.kv_data_ptrs = kv_data_ptrs kv_args.kv_data_lens = kv_data_lens kv_args.kv_item_lens = kv_item_lens @@ -792,4 +852,18 @@ class SchedulerDisaggregationPrefillMixin: f"Skip sending kv chunk for request {req.rid=} {req.bootstrap_room=} because page_indices is empty" ) return + prefill_queue = getattr(self, "disagg_prefill_bootstrap_queue", None) + _cp_draft_shared_kv_debug( + "prefill_send_kv_chunk rid=%s room=%s start_idx=%s end_idx=%s " + "last_chunk=%s page_size=%s pages=%s state_pages=%s has_draft_pool=%s", + req.rid, + req.bootstrap_room, + start_idx, + end_idx, + last_chunk, + page_size, + _seq_summary(page_indices), + _seq_summary(state_indices), + getattr(prefill_queue, "draft_token_to_kv_pool", None) is not None, + ) req.disagg_kv_sender.send(page_indices, state_indices) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index f54c814b4..2e26548de 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -213,6 +213,8 @@ class Envs: SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH = EnvBool(False) SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX = EnvBool(False) SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES = EnvInt(-1) + SGLANG_CP_DRAFT_SHARED_KV = EnvBool(False) + SGLANG_CP_DRAFT_SHARED_KV_DEBUG = EnvBool(False) SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False) SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(False) SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1) diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 3a33fff0d..d98b1d45a 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -24,6 +24,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( filter_owned_logical_locs, get_or_build_shared_paged_buffer_slot_remap, is_current_only_extend_batch, + log_cp_draft_shared_kv_debug, materialize_shared_paged_buffer, tensor_debug_checksum, tensor_debug_summary, @@ -1497,6 +1498,15 @@ class Indexer(MultiPlatformOp): physical_out_loc = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch) if physical_out_loc is None: return False + log_cp_draft_shared_kv_debug( + "index_write", + "index_write layer=%s tokens=%s physical_tokens=%s pool=%s key_shape=%s", + layer_id, + local_out_loc.numel(), + physical_out_loc.numel(), + forward_batch.token_to_kv_pool.__class__.__name__, + tuple(local_key.shape), + ) self._store_index_k_cache( forward_batch=forward_batch, layer_id=layer_id, diff --git a/python/sglang/srt/layers/attention/nsa/utils.py b/python/sglang/srt/layers/attention/nsa/utils.py index 5a1bc6298..fd42d48bc 100644 --- a/python/sglang/srt/layers/attention/nsa/utils.py +++ b/python/sglang/srt/layers/attention/nsa/utils.py @@ -9,6 +9,7 @@ import torch.nn.functional as F import triton import triton.language as tl +from sglang.srt.environ import envs from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -29,6 +30,23 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_CP_DRAFT_SHARED_KV_DEBUG_COUNTS = {} + + +def log_cp_draft_shared_kv_debug( + key: str, + message: str, + *args, + limit: int = 128, +) -> None: + if not envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + return + count = _CP_DRAFT_SHARED_KV_DEBUG_COUNTS.get(key, 0) + if count >= limit: + return + _CP_DRAFT_SHARED_KV_DEBUG_COUNTS[key] = count + 1 + logger.info("[CP_DRAFT_SHARED_KV] " + message, *args) + def log_cp_shared_kv_direct_write_fallback( reason: str, @@ -565,6 +583,17 @@ def get_cp_shared_kv_local_physical_out_cache_loc(forward_batch: "ForwardBatch") physical_out_cache_loc = layout.logical_locs_to_physical( local_out_cache_loc ).contiguous() + log_cp_draft_shared_kv_debug( + "physical_out_loc", + "physical_out_loc cp_rank=%s cp_size=%s page_size=%s tokens=%s " + "physical_tokens=%s pool=%s", + layout.cp_rank, + layout.cp_size, + layout.page_size, + local_out_cache_loc.numel(), + physical_out_cache_loc.numel(), + getattr(forward_batch, "token_to_kv_pool", None).__class__.__name__, + ) forward_batch.cp_local_physical_out_cache_loc = physical_out_cache_loc return physical_out_cache_loc diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index f33d97950..9c6eed656 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -71,6 +71,10 @@ class LogitsProcessorOutput: # Used by speculative decoding (EAGLE) # The last hidden layers hidden_states: Optional[torch.Tensor] = None + # CP-local hidden states for draft prefill. This is intentionally separate + # from `hidden_states`: the logits path may use compact/narrow hidden while + # EAGLE draft still needs the local target hidden to build draft KV. + draft_hidden_states: Optional[torch.Tensor] = None ## Part 2: This part will be assigned in python/sglang/srt/layers/sampler.py::Sampler # he log probs of output tokens, if SGLANG_RETURN_ORIGINAL_LOGPROB = True, will get the log probs before applying temperature. If False, will get the log probs before applying temperature. @@ -314,7 +318,11 @@ class LogitsProcessor(nn.Module): logits = self._get_logits(hidden_states, lm_head, logits_metadata) return LogitsProcessorOutput( next_token_logits=logits, - hidden_states=None, + hidden_states=( + hidden_states + if logits_metadata.capture_hidden_mode.is_last() + else None + ), mm_input_embeds=logits_metadata.mm_input_embeds, ) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index aded74eb0..c058169bf 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2317,6 +2317,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): else CaptureHiddenMode.NULL ) ), + capture_draft_hidden_states=False, extend_input_logprob_token_ids=self.extend_input_logprob_token_ids, is_prefill_only=self.is_prefill_only, dimensions=self.dimensions, @@ -2498,6 +2499,7 @@ class ModelWorkerBatch: # If set, the output of the batch contains the hidden states of the run. capture_hidden_mode: CaptureHiddenMode = None + capture_draft_hidden_states: bool = False hicache_consumer_index: int = -1 # For matryoshka embeddings diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f4534b028..1a7c0b426 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -241,6 +241,22 @@ TEST_RETRACT = envs.SGLANG_TEST_RETRACT.get() TEST_RETRACT_INTERVAL = envs.SGLANG_TEST_RETRACT_INTERVAL.get() TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get() + +def _cp_draft_pool_summary(pool) -> str: + if pool is None: + return "None" + parts = [pool.__class__.__name__] + for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"): + if hasattr(pool, attr): + parts.append(f"{attr}={getattr(pool, attr)}") + return " ".join(parts) + + +def _cp_draft_shared_kv_debug(message: str, *args) -> None: + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + logger.info("[CP_DRAFT_SHARED_KV] " + message, *args) + + _is_npu = is_npu() @@ -941,6 +957,16 @@ class Scheduler( draft_token_to_kv_pool = self.draft_worker.model_runner.token_to_kv_pool model_config = self.draft_worker.model_config + _cp_draft_shared_kv_debug( + "scheduler_disagg_init mode=%s spec_algorithm=%s draft_worker=%s " + "draft_pool=(%s) target_pool=(%s)", + self.disaggregation_mode, + self.spec_algorithm, + self.draft_worker is not None, + _cp_draft_pool_summary(draft_token_to_kv_pool), + _cp_draft_pool_summary(self.token_to_kv_pool_allocator.get_kvcache()), + ) + if ( self.disaggregation_mode == DisaggregationMode.DECODE ): # *2 for the headroom. diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index b535fdae5..c4a42df79 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -369,16 +369,16 @@ def alloc_paged_token_slots_extend( num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size evict_result = evict_from_tree_cache(tree_cache, num_tokens) - logger.info( - "[MemCache-alloc] alloc_paged_token_slots_extend: extend_num_tokens=%d batch_size=%d num_tokens=%d page_size=%d " - "available_size=%d evicted=%d", - extend_num_tokens, - len(seq_lens_cpu), - num_tokens, - allocator.page_size, - allocator.available_size(), - getattr(evict_result, "num_tokens_evicted", 0), - ) + # logger.info( + # "[MemCache-alloc] alloc_paged_token_slots_extend: extend_num_tokens=%d batch_size=%d num_tokens=%d page_size=%d " + # "available_size=%d evicted=%d", + # extend_num_tokens, + # len(seq_lens_cpu), + # num_tokens, + # allocator.page_size, + # allocator.available_size(), + # getattr(evict_result, "num_tokens_evicted", 0), + # ) alloc_extend_compute_owner = getattr( allocator, "alloc_extend_compute_owner", None diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 4a553ee63..32a88b927 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -399,6 +399,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): spec_algorithm: SpeculativeAlgorithm = None mm_input_embeds: Optional[torch.Tensor] = None capture_hidden_mode: CaptureHiddenMode = None + capture_draft_hidden_states: bool = False + draft_hidden_states: Optional[torch.Tensor] = None # For padding padded_static_len: int = -1 # -1 if not padded @@ -484,6 +486,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): spec_algorithm=batch.spec_algorithm, spec_info=batch.spec_info, capture_hidden_mode=batch.capture_hidden_mode, + capture_draft_hidden_states=batch.capture_draft_hidden_states, input_embeds=batch.input_embeds, token_type_ids=batch.token_type_ids, tbo_split_seq_index=batch.tbo_split_seq_index, diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index b5475f982..b57b2ba3a 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -9,6 +9,7 @@ from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.attention.nsa.utils import ( get_cp_shared_kv_local_out_cache_loc, get_cp_shared_kv_local_physical_out_cache_loc, + log_cp_draft_shared_kv_debug, log_cp_shared_kv_direct_write_fallback, nsa_use_prefill_cp, ) @@ -623,6 +624,18 @@ class DeepseekMLAForwardMixin: k_nope=k_nope, k_rope=k_pe, ): + log_cp_draft_shared_kv_debug( + "mla_tai_write", + "mla_write path=tai_fused layer=%s cp_rank=%s cp_size=%s tokens=%s " + "pool=%s k_nope_shape=%s k_pe_shape=%s", + self.attn_mqa.layer_id, + layout.cp_rank, + layout.cp_size, + local_out_cache_loc.numel(), + forward_batch.token_to_kv_pool.__class__.__name__, + tuple(k_nope.shape), + tuple(k_pe.shape), + ) return True physical_out_cache_loc = get_cp_shared_kv_local_physical_out_cache_loc( @@ -630,6 +643,17 @@ class DeepseekMLAForwardMixin: ) if physical_out_cache_loc is None: return False + log_cp_draft_shared_kv_debug( + "mla_torch_write", + "mla_write path=torch layer=%s tokens=%s physical_tokens=%s pool=%s " + "k_nope_shape=%s k_pe_shape=%s", + self.attn_mqa.layer_id, + local_out_cache_loc.numel(), + physical_out_cache_loc.numel(), + forward_batch.token_to_kv_pool.__class__.__name__, + tuple(k_nope.shape), + tuple(k_pe.shape), + ) forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.attn_mqa, physical_out_cache_loc, diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 207c58624..c0d135f50 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -28,6 +28,8 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r from sglang.srt.layers.attention.nsa.utils import ( can_cp_split, cp_all_gather_rerange_output, + cp_collect_last_token_hidden, + cp_split_and_rebuild_1d, cp_split_and_rebuild_data, cp_split_and_rebuild_position, is_nsa_enable_prefill_cp, @@ -130,6 +132,35 @@ class DeepseekModelNextN(nn.Module): else: self.cp_size = None + def _debug_cp_draft_shared_kv(self, message: str): + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + logger.info("[CP_DRAFT_SHARED_KV] %s", message) + + def _get_cp_local_spec_hidden_states( + self, + forward_batch: ForwardBatch, + spec_hidden_states: torch.Tensor, + *, + full_num_tokens: int, + local_num_tokens: int, + ) -> Optional[torch.Tensor]: + if spec_hidden_states is None: + self._debug_cp_draft_shared_kv("fallback reason=missing_spec_hidden") + return None + + if spec_hidden_states.shape[0] == local_num_tokens: + return spec_hidden_states + + if spec_hidden_states.shape[0] == full_num_tokens: + return cp_split_and_rebuild_data(forward_batch, spec_hidden_states) + + self._debug_cp_draft_shared_kv( + "fallback reason=spec_hidden_shape_mismatch " + f"spec_tokens={spec_hidden_states.shape[0]} " + f"full_tokens={full_num_tokens} local_tokens={local_num_tokens}" + ) + return None + def forward( self, input_ids: torch.Tensor, @@ -145,25 +176,68 @@ class DeepseekModelNextN(nn.Module): ), ) - if input_embeds is None: - hidden_states = self.embed_tokens(input_ids) - else: - hidden_states = input_embeds + use_cp = nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp) + use_cp_local_draft = use_cp and envs.SGLANG_CP_DRAFT_SHARED_KV.get() + if use_cp_local_draft: + local_input_ids = cp_split_and_rebuild_1d(forward_batch, input_ids) + local_num_tokens = local_input_ids.shape[0] + local_positions = cp_split_and_rebuild_position(forward_batch, positions) + spec_hidden_states = self._get_cp_local_spec_hidden_states( + forward_batch, + forward_batch.spec_info.hidden_states, + full_num_tokens=input_ids.shape[0], + local_num_tokens=local_num_tokens, + ) + if spec_hidden_states is None: + use_cp_local_draft = False + else: + positions = local_positions + if input_embeds is None: + hidden_states = self.embed_tokens(local_input_ids) + elif input_embeds.shape[0] == local_num_tokens: + hidden_states = input_embeds + elif input_embeds.shape[0] == input_ids.shape[0]: + hidden_states = cp_split_and_rebuild_data( + forward_batch, input_embeds + ) + else: + self._debug_cp_draft_shared_kv( + "fallback reason=input_embeds_shape_mismatch " + f"input_embed_tokens={input_embeds.shape[0]} " + f"full_tokens={input_ids.shape[0]} " + f"local_tokens={local_num_tokens}" + ) + use_cp_local_draft = False + + if not use_cp_local_draft: + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds + spec_hidden_states = forward_batch.spec_info.hidden_states if hidden_states.shape[0] > 0: hidden_states = self.eh_proj( torch.cat( ( self.enorm(hidden_states), - self.hnorm(forward_batch.spec_info.hidden_states), + self.hnorm(spec_hidden_states), ), dim=-1, ) ) - if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp): + if use_cp and not use_cp_local_draft: hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) + + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get() and use_cp_local_draft: + self._debug_cp_draft_shared_kv( + "local_path " + f"full_tokens={input_ids.shape[0]} " + f"local_tokens={hidden_states.shape[0]} " + f"capture_hidden_mode={forward_batch.capture_hidden_mode}" + ) residual = None with get_global_expert_distribution_recorder().disable_this_region(): hidden_states, residual = self.decoder( @@ -180,14 +254,19 @@ class DeepseekModelNextN(nn.Module): else: hidden_states = self.shared_head.norm(hidden_states) - if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp): - # allgather + rerrange - hidden_states = cp_all_gather_rerange_output( - hidden_states, - self.cp_size, - forward_batch, - torch.cuda.current_stream(), - ) + if use_cp: + if use_cp_local_draft: + hidden_states = cp_collect_last_token_hidden( + hidden_states, forward_batch, self.cp_size + ) + else: + # allgather + rerange + hidden_states = cp_all_gather_rerange_output( + hidden_states, + self.cp_size, + forward_batch, + torch.cuda.current_stream(), + ) return hidden_states diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 1d6a5178e..17744eb2d 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2039,6 +2039,9 @@ class DeepseekV2Model(nn.Module): else: hidden_states, _ = self.norm(hidden_states, residual) + if getattr(forward_batch, "capture_draft_hidden_states", False): + forward_batch.draft_hidden_states = hidden_states + if self.pp_group.is_last_rank and nsa_use_prefill_cp(forward_batch): if self._should_use_narrow_output_path(forward_batch): hidden_states = cp_collect_last_token_hidden( @@ -2221,9 +2224,13 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): hidden_states, aux_hidden_states = hidden_states if self.pp_group.is_last_rank: - return self.logits_processor( + logits_output = self.logits_processor( input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states ) + logits_output.draft_hidden_states = getattr( + forward_batch, "draft_hidden_states", None + ) + return logits_output else: return hidden_states diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 59c63c17c..35ce3ea41 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -5,6 +5,7 @@ from typing import List, Optional, Tuple import torch from sglang.srt.distributed import get_tp_group +from sglang.srt.environ import envs from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import ( EAGLEDraftNpuGraphRunner, ) @@ -276,6 +277,36 @@ class EAGLEWorker(TpModelWorker): def draft_model_runner(self): return self.model_runner + def _debug_cp_draft_shared_kv(self, message: str): + if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get(): + logger.info("[CP_DRAFT_SHARED_KV] %s", message) + + def _can_use_cp_draft_shared_kv(self, batch: ScheduleBatch) -> bool: + if not envs.SGLANG_CP_DRAFT_SHARED_KV.get(): + return False + if not (batch.forward_mode.is_extend() or batch.is_extend_in_batch): + self._debug_cp_draft_shared_kv("fallback reason=non_extend_mode") + return False + draft_architectures = getattr( + self.draft_model_runner.model_config.hf_config, "architectures", [] + ) + if "DeepseekV3ForCausalLMNextN" not in (draft_architectures or []): + self._debug_cp_draft_shared_kv( + f"fallback reason=unsupported_arch architectures={draft_architectures}" + ) + return False + if not getattr(self.target_worker.model_runner, "uses_cp_shared_kv", False): + self._debug_cp_draft_shared_kv( + "fallback reason=target_cp_shared_kv_disabled" + ) + return False + if not getattr(self.draft_model_runner, "uses_cp_shared_kv", False): + self._debug_cp_draft_shared_kv( + "fallback reason=draft_cp_shared_kv_disabled" + ) + return False + return True + def forward_batch_generation(self, batch: ScheduleBatch) -> GenerationBatchResult: """Run speculative decoding forward. @@ -298,9 +329,21 @@ class EAGLEWorker(TpModelWorker): with self.draft_tp_context( self.draft_model_runner.tp_group ), speculative_moe_backend_context(), speculative_moe_a2a_backend_context(): + draft_hidden_states = ( + logits_output.draft_hidden_states + if logits_output.draft_hidden_states is not None + else logits_output.hidden_states + ) + if ( + envs.SGLANG_CP_DRAFT_SHARED_KV.get() + and draft_hidden_states is None + ): + self._debug_cp_draft_shared_kv( + "fallback_failed reason=missing_target_hidden" + ) self.forward_draft_extend( batch, - logits_output.hidden_states, + draft_hidden_states, next_token_ids, seq_lens_cpu, logits_output.mm_input_embeds, @@ -367,15 +410,22 @@ class EAGLEWorker(TpModelWorker): batch: The batch to run. States could be modified. Returns: - logits_output: The output of logits. It will contain the full hidden states. + logits_output: The output of logits. It contains full hidden states on the + legacy path, or CP-local draft hidden states in `draft_hidden_states` + when CP draft shared-KV is enabled. next_token_ids: Next token ids generated. seq_lens_cpu: CPU copy of sequence lengths for the draft prefill path. can_run_cuda_graph: Whether the target prefill ran with cuda graph. """ - # Forward with the target model and get hidden states. - # We need the full hidden states to prefill the KV cache of the draft model. + # Forward with the target model and get hidden states for draft prefill. + # CP draft shared-KV keeps this hidden side channel CP-local so target + # logits can still use the narrow output path. model_worker_batch = batch.get_model_worker_batch() - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL + if self._can_use_cp_draft_shared_kv(batch): + model_worker_batch.capture_hidden_mode = CaptureHiddenMode.NULL + model_worker_batch.capture_draft_hidden_states = True + else: + model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL batch_result = self.target_worker.forward_batch_generation(model_worker_batch) logits_output, next_token_ids = ( batch_result.logits_output,