diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 2fad05156..5ba315ea7 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1218,8 +1218,6 @@ class Req(ReqDllmMixin): ) ) - self.set_extend_input_len(input_len - len(self.prefix_indices)) - def _compute_max_prefix_len(self, input_len: int) -> int: # NOTE: the matched length is at most 1 less than the input length to enable logprob computation max_prefix_len = input_len - 1 diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 1b34f0f54..514628a3c 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -659,7 +659,7 @@ class PrefillAdder: * self.page_size ) - req.extend_input_len = trunc_len + req.set_extend_input_len(trunc_len) req.fill_len = prefix_len + trunc_len self.can_run_list.append(req) @@ -679,9 +679,13 @@ class PrefillAdder: return AddReqResult.NO_TOKEN # Truncate input length to available tokens and update request metadata - truncated = req.extend_input_len > _rem_tokens - req.extend_input_len = min(req.extend_input_len, _rem_tokens) - req.fill_len = len(req.prefix_indices) + req.extend_input_len + cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( + req.prefix_indices + ) + truncated = cand_extend_input_len > _rem_tokens + new_len = min(cand_extend_input_len, _rem_tokens) + req.set_extend_input_len(new_len) + req.fill_len = len(req.prefix_indices) + new_len self.can_run_list.append(req) # Update budget: reserve max_new_tokens only if not truncated @@ -719,9 +723,13 @@ class PrefillAdder: return req _rem_tokens = self.rem_chunk_tokens - truncated = req.extend_input_len > _rem_tokens - req.set_extend_input_len(min(req.extend_input_len, _rem_tokens)) - req.fill_len = len(req.prefix_indices) + req.extend_input_len + cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( + req.prefix_indices + ) + truncated = cand_extend_input_len > _rem_tokens + new_len = min(cand_extend_input_len, _rem_tokens) + req.set_extend_input_len(new_len) + req.fill_len = len(req.prefix_indices) + new_len self.can_run_list.append(req) self._update_prefill_budget( 0, @@ -755,11 +763,14 @@ class PrefillAdder: self.tree_cache.dec_lock_ref(last_node) def add_one_req_ignore_eos(self, req: Req): - paged_input = self.ceil_paged_tokens(req.extend_input_len) + cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( + req.prefix_indices + ) + paged_input = self.ceil_paged_tokens(cand_extend_input_len) if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens): return AddReqResult.NO_TOKEN if self.is_hybrid_swa: - if self._swa_budget_for_req(req.extend_input_len) > self.rem_swa_tokens: + if self._swa_budget_for_req(cand_extend_input_len) > self.rem_swa_tokens: return AddReqResult.NO_TOKEN def add_req_state(r, insert_sort=False): @@ -799,7 +810,7 @@ class PrefillAdder: # Skip this logic for swa. The SWA has different memory management, and # this mechanism is underestimating the memory usage. cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens( - req.extend_input_len + cand_extend_input_len ) tokens_freed = 0 for i, (tokens_left, tokens_occupied) in enumerate(self.req_states): @@ -825,13 +836,13 @@ class PrefillAdder: self._add_dllm_req(req, 0) elif ( self.rem_chunk_tokens is None # chunked prefill is disabled - or req.extend_input_len <= self.rem_chunk_tokens # it is the last chunk + or cand_extend_input_len <= self.rem_chunk_tokens # it is the last chunk ): # Non-chunked prefill — the whole sequence is committed this iter. + req.set_extend_input_len( + len(req.full_untruncated_fill_ids) - len(req.prefix_indices) + ) req.fill_len = len(req.full_untruncated_fill_ids) - assert ( - req.fill_len == len(req.prefix_indices) + req.extend_input_len - ), f"{req.fill_len=} {len(req.prefix_indices)=} {req.extend_input_len=}" self.can_run_list.append(req) self._update_prefill_budget( 0, @@ -889,10 +900,13 @@ class PrefillAdder: max(req.sampling_params.max_new_tokens - len(req.output_ids), 0), CLIP_MAX_NEW_TOKENS, ) - total_tokens = req.extend_input_len + max_new + self.page_size + cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( + req.prefix_indices + ) + total_tokens = cand_extend_input_len + max_new + self.page_size # adjusting the input_tokens based on host_hit_length and page_size - real_input_tokens = req.extend_input_len - req.host_hit_length + real_input_tokens = cand_extend_input_len - req.host_hit_length real_input_tokens = self.ceil_paged_tokens(real_input_tokens) prefix_len = len(req.prefix_indices) @@ -901,7 +915,7 @@ class PrefillAdder: if self.is_hybrid_swa: swa_needed = self._swa_budget_for_req( - req.extend_input_len, swa_host_hit_length=req.swa_host_hit_length + cand_extend_input_len, swa_host_hit_length=req.swa_host_hit_length ) if swa_needed >= self.rem_swa_tokens: return AddReqResult.NO_TOKEN @@ -923,7 +937,7 @@ class PrefillAdder: if self.is_hybrid_swa: swa_needed = self._swa_budget_for_req( - req.extend_input_len, swa_host_hit_length=req.swa_host_hit_length + cand_extend_input_len, swa_host_hit_length=req.swa_host_hit_length ) if swa_needed >= self.rem_swa_tokens: return AddReqResult.NO_TOKEN @@ -937,13 +951,12 @@ class PrefillAdder: ) ) req.prefix_indices = torch.cat([req.prefix_indices, new_indices]) - req.set_extend_input_len( - len(req.full_untruncated_fill_ids) - len(req.prefix_indices) - ) prefix_len = len(req.prefix_indices) req.cache_protected_len = prefix_len - input_tokens = self.ceil_paged_tokens(req.extend_input_len) + input_tokens = self.ceil_paged_tokens( + len(req.full_untruncated_fill_ids) - len(req.prefix_indices) + ) if ( self.rem_chunk_tokens is None @@ -967,10 +980,10 @@ class PrefillAdder: self._req_inc_lock_ref(req) elif self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens: # Non-chunked prefill — the whole sequence is committed this iter. + req.set_extend_input_len( + len(req.full_untruncated_fill_ids) - len(req.prefix_indices) + ) req.fill_len = len(req.full_untruncated_fill_ids) - assert ( - req.fill_len == len(req.prefix_indices) + req.extend_input_len - ), f"{req.fill_len=} {len(req.prefix_indices)=} {req.extend_input_len=}" self.can_run_list.append(req) self._req_inc_lock_ref(req) @@ -1052,7 +1065,8 @@ class PrefillAdder: preemptible_reqs = [] min_tokens_to_remove = ( - req.extend_input_len + len(req.full_untruncated_fill_ids) + - len(req.prefix_indices) + min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS) - self.rem_total_tokens ) diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 7dcc1b477..1b5badbb1 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -77,6 +77,8 @@ class TestPrefillAdder(CustomTestCase): req.rid = str(rid) req.priority = priority req.extend_input_len = 0 + 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)