Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docker/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -84,7 +84,7 @@ RUN pip install nvidia-cudnn-cu12==9.16.0.29
# reinstall numpy 1.x for megatron
RUN pip install "numpy<2"

RUN pip install https://github.com/zhuzilin/sgl-router/releases/download/v0.3.2-1117d05/sglang_router-0.3.2-cp38-abi3-manylinux_2_28_x86_64.whl --force-reinstall
RUN pip install https://github.com/zhuzilin/sgl-router/releases/download/v0.3.2-3e512c2/sglang_router-0.3.2-cp38-abi3-manylinux_2_28_x86_64.whl --force-reinstall
RUN python -c "import sglang_router; assert 'slime' in sglang_router.__version__"

RUN rm -rf /root/.cache/pip /root/flash-attention
Expand Down
67 changes: 43 additions & 24 deletions docker/patch/latest/sglang-top_p.patch
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py
index 70265a424f5..a19278e8828 100644
index de570fb9fd..a19278e8828 100644
--- a/python/sglang/srt/disaggregation/decode.py
+++ b/python/sglang/srt/disaggregation/decode.py
@@ -1488,6 +1488,8 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
@@ -1476,6 +1476,8 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
output_token_logprobs_idx,
output_top_logprobs_val,
output_top_logprobs_idx,
Expand All @@ -11,7 +11,7 @@ index 70265a424f5..a19278e8828 100644
output_topk_p,
output_topk_index,
output_hidden_states,
@@ -1580,6 +1582,11 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
@@ -1569,6 +1571,11 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
: decode_req.req.logprob.top_logprobs_num
].tolist()
)
Expand All @@ -21,10 +21,10 @@ index 70265a424f5..a19278e8828 100644
+ output_top_p_token_ids[:top_p_token_ids_len].tolist()
+ )

if is_slime_profiling_enabled():
apply_prefill_timing_payload(
decode_req.kv_receiver.clear()
decode_req.kv_receiver = None
diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py
index a44685777ea..f21e4dad4a5 100644
index e1d7d9cb7a..121b1e9baa 100644
--- a/python/sglang/srt/disaggregation/utils.py
+++ b/python/sglang/srt/disaggregation/utils.py
@@ -1,5 +1,6 @@
Expand All @@ -33,17 +33,18 @@ index a44685777ea..f21e4dad4a5 100644
+import logging
import os
import random
import time
@@ -42,6 +43,8 @@ PREFILL_TIMING_DEST_ATTRS = (
("fwd_transfer_total_mb", float),
("fwd_prefill_retry_count", int),
)
from collections import deque
@@ -31,7 +32,9 @@ if TYPE_CHECKING:
#########################
# Constants & Enums
#########################
FAKE_BOOTSTRAP_HOST = "2.2.2.2"
+MAX_PD_TOP_P_TOKEN_IDS = 4096
+logger = logging.getLogger(__name__)


class DisaggregationMode(Enum):
@@ -253,6 +256,12 @@ class MetadataBuffers:
@@ -239,6 +242,12 @@ class MetadataBuffers:
self.output_top_logprobs_idx = torch.zeros(
(size, max_top_logprobs_num), dtype=torch.int32, device=device
)
Expand All @@ -56,16 +57,34 @@ index a44685777ea..f21e4dad4a5 100644
# For PD + spec decode
self.output_topk_p = torch.zeros(
(size, 16), dtype=torch.float32, device=device
@@ -277,6 +286,8 @@ class MetadataBuffers:
("output_token_logprobs_idx", self.output_token_logprobs_idx),
("output_top_logprobs_val", self.output_top_logprobs_val),
("output_top_logprobs_idx", self.output_top_logprobs_idx),
+ ("output_top_p_token_ids_len", self.output_top_p_token_ids_len),
+ ("output_top_p_token_ids", self.output_top_p_token_ids),
("output_topk_p", self.output_topk_p),
("output_topk_index", self.output_topk_index),
("output_hidden_states", self.output_hidden_states),
@@ -301,6 +312,8 @@ class MetadataBuffers:
@@ -266,6 +275,8 @@ class MetadataBuffers:
self.output_token_logprobs_idx.data_ptr(),
self.output_top_logprobs_val.data_ptr(),
self.output_top_logprobs_idx.data_ptr(),
+ self.output_top_p_token_ids_len.data_ptr(),
+ self.output_top_p_token_ids.data_ptr(),
self.output_topk_p.data_ptr(),
self.output_topk_index.data_ptr(),
self.output_hidden_states.data_ptr(),
@@ -278,6 +289,8 @@ class MetadataBuffers:
self.output_token_logprobs_idx.nbytes,
self.output_top_logprobs_val.nbytes,
self.output_top_logprobs_idx.nbytes,
+ self.output_top_p_token_ids_len.nbytes,
+ self.output_top_p_token_ids.nbytes,
self.output_topk_p.nbytes,
self.output_topk_index.nbytes,
self.output_hidden_states.nbytes,
@@ -290,6 +303,8 @@ class MetadataBuffers:
self.output_token_logprobs_idx[0].nbytes,
self.output_top_logprobs_val[0].nbytes,
self.output_top_logprobs_idx[0].nbytes,
+ self.output_top_p_token_ids_len[0].nbytes,
+ self.output_top_p_token_ids[0].nbytes,
self.output_topk_p[0].nbytes,
self.output_topk_index[0].nbytes,
self.output_hidden_states[0].nbytes,
@@ -303,6 +318,8 @@ class MetadataBuffers:
self.output_token_logprobs_idx[idx].clone(),
self.output_top_logprobs_val[idx].clone(),
self.output_top_logprobs_idx[idx].clone(),
Expand All @@ -74,15 +93,15 @@ index a44685777ea..f21e4dad4a5 100644
self.output_topk_p[idx].clone(),
self.output_topk_index[idx].clone(),
self.output_hidden_states[idx].clone(),
@@ -318,6 +331,7 @@ class MetadataBuffers:
@@ -320,6 +337,7 @@ class MetadataBuffers:
self.cached_tokens[req.metadata_buffer_index][1] = req.cached_tokens_device
self.cached_tokens[req.metadata_buffer_index][2] = req.cached_tokens_host
self.cached_tokens[req.metadata_buffer_index][3] = req.cached_tokens_storage
+ self.output_top_p_token_ids_len[req.metadata_buffer_index][0] = 0
if req.return_logprob:
if req.logprob.output_token_logprobs_val: # not none or empty list
self.output_token_logprobs_val[req.metadata_buffer_index][0] = (
@@ -344,6 +358,28 @@ class MetadataBuffers:
@@ -346,6 +364,28 @@ class MetadataBuffers:
dtype=torch.int32,
device="cpu",
)
Expand Down
Loading
Loading