From 2b64fc7a2cbf38c06169bd05cee958d9f917f4c0 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Thu, 25 Jun 2026 09:06:14 +0800 Subject: [PATCH] Remove Req.extend_logprob_start_len field and make it pure (#27625) --- .../decode_schedule_batch_mixin.py | 7 +-- python/sglang/srt/managers/schedule_batch.py | 58 ++++++++++--------- python/sglang/srt/managers/scheduler.py | 6 +- .../unit/managers/test_prefill_adder.py | 1 - 4 files changed, 36 insertions(+), 36 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py index 4ded8df10..4c30337d5 100644 --- a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py +++ b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py @@ -73,9 +73,6 @@ class ScheduleBatchDisaggregationDecodeMixin: req.already_computed = seq_len req.is_retracted = False pre_lens.append(pre_len) - req.extend_logprob_start_len = 0 - - extend_input_logprob_token_ids = None # Set fields self.input_ids = torch.tensor( @@ -100,8 +97,8 @@ class ScheduleBatchDisaggregationDecodeMixin: self.extend_num_tokens = extend_num_tokens self.prefix_lens = [len(r.prefix_indices) for r in reqs] self.extend_lens = [r.extend_range.length for r in reqs] - self.extend_logprob_start_lens = [r.extend_logprob_start_len for r in reqs] - self.extend_input_logprob_token_ids = extend_input_logprob_token_ids + self.extend_logprob_start_lens = None + self.extend_input_logprob_token_ids = None self.multimodal_inputs = [r.multimodal_inputs for r in reqs] # Build sampling info diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index cb67004b7..cb2823dfd 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -844,8 +844,6 @@ class Req(ReqDllmMixin): # Prefix info # The indices to kv cache for the shared prefix. self.prefix_indices: torch.Tensor = torch.empty((0,), dtype=torch.int64) - # The relative logprob_start_len in an extend batch - self.extend_logprob_start_len = 0 # TODO(ispobock): rename to last_device_node self.last_node: Any = None self.last_host_node: Any = None @@ -1098,7 +1096,6 @@ class Req(ReqDllmMixin): def set_extend_range(self, start: int, end: int) -> None: self.extend_range = Range(start, end) - self._recompute_extend_logprob_start_len() def get_fill_ids(self) -> array: return self.full_untruncated_fill_ids[: self.extend_range.end] @@ -1458,7 +1455,6 @@ class Req(ReqDllmMixin): self.input_token_logprobs = None self.temp_input_top_logprobs_val = None self.temp_input_top_logprobs_idx = None - self.extend_logprob_start_len = 0 self.inflight_middle_chunks = 0 self.mamba_pool_idx = None self.mamba_ping_pong_track_buffer = None @@ -1524,23 +1520,6 @@ class Req(ReqDllmMixin): logger.info(f"{prefix}: {self.time_stats.convert_to_duration()}") self.has_log_time_stats = True - def _recompute_extend_logprob_start_len(self): - # Setting extend_input_len and computing the relative logprob_start_len in an extend batch - # - # Key variables: - # - logprob_start_len: Absolute position in full sequence where logprob computation begins - # - extend_logprob_start_len: Relative position within current extend batch where logprob computation begins - # - extend_input_len: Number of tokens that need to be processed in this extend batch - if self.logprob_start_len == -1: - logprob_start_len = len(self.full_untruncated_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)) - self.extend_logprob_start_len = min( - logprob_start_len - len(self.prefix_indices), - self.extend_range.length, - ) - def set_finish_with_abort(self, error_msg: str): if get_tensor_model_parallel_rank() == 0: logger.error(f"{error_msg}, {self.rid=}") @@ -1651,6 +1630,25 @@ def retract_all( return retracted_reqs +def compute_extend_logprob_start_len( + *, + logprob_start_len: int, + prefix_len: int, + extend_len: int, + full_untruncated_fill_len: int, +) -> int: + # Key variables: + # - logprob_start_len: Absolute position in full sequence where logprob computation begins + # - extend_logprob_start_len: Relative position within current extend batch where logprob computation begins + # - extend_input_len: Number of tokens that need to be processed in this extend batch + if logprob_start_len == -1: + resolved_start = full_untruncated_fill_len + else: + # logprob_start_len should be at least the length of the prefix indices + resolved_start = max(logprob_start_len, prefix_len) + return min(resolved_start - prefix_len, extend_len) + + def _compute_chunked_req_next_prompt_token( chunked_req: Optional[Req], vocab_size: int, @@ -2004,9 +2002,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): req.extend_range = req.extend_range._replace( start=req.extend_range.start + encoder_len ) - req.extend_logprob_start_len = max( - 0, req.extend_logprob_start_len - encoder_len - ) req.logprob_start_len = max(req.logprob_start_len, encoder_len) def prepare_for_extend(self): @@ -2024,6 +2019,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): orig_seq_lens = [max(r.extend_range.end, len(r.origin_input_ids)) for r in reqs] prefix_lens = [len(r.prefix_indices) for r in reqs] extend_lens = [r.extend_range.length for r in reqs] + extend_logprob_start_lens = [ + compute_extend_logprob_start_len( + logprob_start_len=r.logprob_start_len, + prefix_len=prefix_lens[i], + extend_len=extend_lens[i], + full_untruncated_fill_len=len(r.full_untruncated_fill_ids), + ) + for i, r in enumerate(reqs) + ] _pin = is_pin_memory_available(self.device) # Stay on pinned CPU; H2D is deferred to forward stream via @@ -2172,13 +2176,13 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ] extend_input_logprob_token_ids.extend(logprob_token_ids) - # We will need req.extend_range.length - req.extend_logprob_start_len number of + # We will need req.extend_range.length - extend_logprob_start_lens[i] number of # tokens, and logprob_token_ids is for input logprob, so pad the rest of them by 0. extend_input_logprob_token_ids.extend( [0] * ( req.extend_range.length - - req.extend_logprob_start_len + - extend_logprob_start_lens[i] - len(logprob_token_ids) ) ) @@ -2237,7 +2241,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.top_logprobs_nums = [r.logprob.top_logprobs_num for r in reqs] self.token_ids_logprobs = [r.logprob.token_ids_logprob for r in reqs] - self.extend_logprob_start_lens = [r.extend_logprob_start_len for r in reqs] + self.extend_logprob_start_lens = extend_logprob_start_lens self.extend_input_logprob_token_ids = extend_input_logprob_token_ids if get_global_server_args().enable_mamba_extra_buffer(): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index dea1887e5..8f4dab9d6 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3340,9 +3340,9 @@ class Scheduler( req.extend_range.length if req.extend_range is not None else 0 for req in batch.reqs ] - batch_result.extend_logprob_start_len_per_req = [ - req.extend_logprob_start_len for req in batch.reqs - ] + batch_result.extend_logprob_start_len_per_req = ( + batch.extend_logprob_start_lens + ) else: batch_result.extend_input_len_per_req = None batch_result.extend_logprob_start_len_per_req = None diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 6701f8415..501a65d5e 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -79,7 +79,6 @@ class TestPrefillAdder(CustomTestCase): req.priority = priority req.prefix_indices = [] req.full_untruncated_fill_ids = [] - req.extend_logprob_start_len = 0 req.output_ids = [0] * output_len req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens) req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time)