Remove Req.extend_logprob_start_len field and make it pure (#27625)

This commit is contained in:
fzyzcjy
2026-06-25 09:06:14 +08:00
committed by GitHub
parent 2c3f007a65
commit 2b64fc7a2c
4 changed files with 36 additions and 36 deletions
@@ -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
+31 -27
View File
@@ -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():
+3 -3
View File
@@ -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