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,