From a3d88a247b1744ff85cb92aa61150318d22e268d Mon Sep 17 00:00:00 2001 From: Qiaolin Yu Date: Tue, 10 Mar 2026 12:50:57 -0700 Subject: [PATCH] Enable piecewise-cuda-graph when logprob_start_len = -1 (#19453) --- python/sglang/srt/managers/schedule_batch.py | 4 ++-- python/sglang/srt/managers/scheduler.py | 2 +- test/registered/sessions/test_session_control.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 215748014..d625f9047 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1157,7 +1157,7 @@ class Req(ReqDllmMixin): # - extend_input_len: Number of tokens that need to be processed in this extend batch self.extend_input_len = extend_input_len if self.logprob_start_len == -1: - logprob_start_len = len(self.fill_ids) - 1 + logprob_start_len = len(self.fill_ids) else: # logprob_start_len should be at least the length of the prefix indices logprob_start_len = max(self.logprob_start_len, len(self.prefix_indices)) @@ -1582,7 +1582,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): len(req.fill_ids), ) if req.logprob_start_len == -1: - logprob_start_len = len(req.origin_input_ids) - 1 + logprob_start_len = len(req.origin_input_ids) else: logprob_start_len = req.logprob_start_len # Apply logprob_start_len diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 774036268..ded080a28 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1776,7 +1776,7 @@ class Scheduler( if recv_req.return_logprob and recv_req.token_ids_logprob is None: # If logprob is required but neither token_ids_logprob nor logprob_start_len is # set, return the logprobs for output tokens by default - req.logprob_start_len = len(req.origin_input_ids) - 1 + req.logprob_start_len = len(req.origin_input_ids) elif req.is_prefill_only: # For prefill-only requests with logprob_start_len == -1, set logprob_start_len # beyond input sequence to skip input logprob computation entirely diff --git a/test/registered/sessions/test_session_control.py b/test/registered/sessions/test_session_control.py index af166c552..2b495be04 100644 --- a/test/registered/sessions/test_session_control.py +++ b/test/registered/sessions/test_session_control.py @@ -44,7 +44,7 @@ class TestSessionControl(unittest.TestCase): timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=[ "--attention-backend", - "flashinfer", + "triton", "--enable-streaming-session", ], )