Remove Req.extend_logprob_start_len field and make it pure (#27625)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user