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.already_computed = seq_len
|
||||||
req.is_retracted = False
|
req.is_retracted = False
|
||||||
pre_lens.append(pre_len)
|
pre_lens.append(pre_len)
|
||||||
req.extend_logprob_start_len = 0
|
|
||||||
|
|
||||||
extend_input_logprob_token_ids = None
|
|
||||||
|
|
||||||
# Set fields
|
# Set fields
|
||||||
self.input_ids = torch.tensor(
|
self.input_ids = torch.tensor(
|
||||||
@@ -100,8 +97,8 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
|||||||
self.extend_num_tokens = extend_num_tokens
|
self.extend_num_tokens = extend_num_tokens
|
||||||
self.prefix_lens = [len(r.prefix_indices) for r in reqs]
|
self.prefix_lens = [len(r.prefix_indices) for r in reqs]
|
||||||
self.extend_lens = [r.extend_range.length 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_logprob_start_lens = None
|
||||||
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids
|
self.extend_input_logprob_token_ids = None
|
||||||
self.multimodal_inputs = [r.multimodal_inputs for r in reqs]
|
self.multimodal_inputs = [r.multimodal_inputs for r in reqs]
|
||||||
|
|
||||||
# Build sampling info
|
# Build sampling info
|
||||||
|
|||||||
@@ -844,8 +844,6 @@ class Req(ReqDllmMixin):
|
|||||||
# Prefix info
|
# Prefix info
|
||||||
# The indices to kv cache for the shared prefix.
|
# The indices to kv cache for the shared prefix.
|
||||||
self.prefix_indices: torch.Tensor = torch.empty((0,), dtype=torch.int64)
|
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
|
# TODO(ispobock): rename to last_device_node
|
||||||
self.last_node: Any = None
|
self.last_node: Any = None
|
||||||
self.last_host_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:
|
def set_extend_range(self, start: int, end: int) -> None:
|
||||||
self.extend_range = Range(start, end)
|
self.extend_range = Range(start, end)
|
||||||
self._recompute_extend_logprob_start_len()
|
|
||||||
|
|
||||||
def get_fill_ids(self) -> array:
|
def get_fill_ids(self) -> array:
|
||||||
return self.full_untruncated_fill_ids[: self.extend_range.end]
|
return self.full_untruncated_fill_ids[: self.extend_range.end]
|
||||||
@@ -1458,7 +1455,6 @@ class Req(ReqDllmMixin):
|
|||||||
self.input_token_logprobs = None
|
self.input_token_logprobs = None
|
||||||
self.temp_input_top_logprobs_val = None
|
self.temp_input_top_logprobs_val = None
|
||||||
self.temp_input_top_logprobs_idx = None
|
self.temp_input_top_logprobs_idx = None
|
||||||
self.extend_logprob_start_len = 0
|
|
||||||
self.inflight_middle_chunks = 0
|
self.inflight_middle_chunks = 0
|
||||||
self.mamba_pool_idx = None
|
self.mamba_pool_idx = None
|
||||||
self.mamba_ping_pong_track_buffer = 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()}")
|
logger.info(f"{prefix}: {self.time_stats.convert_to_duration()}")
|
||||||
self.has_log_time_stats = True
|
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):
|
def set_finish_with_abort(self, error_msg: str):
|
||||||
if get_tensor_model_parallel_rank() == 0:
|
if get_tensor_model_parallel_rank() == 0:
|
||||||
logger.error(f"{error_msg}, {self.rid=}")
|
logger.error(f"{error_msg}, {self.rid=}")
|
||||||
@@ -1651,6 +1630,25 @@ def retract_all(
|
|||||||
return retracted_reqs
|
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(
|
def _compute_chunked_req_next_prompt_token(
|
||||||
chunked_req: Optional[Req],
|
chunked_req: Optional[Req],
|
||||||
vocab_size: int,
|
vocab_size: int,
|
||||||
@@ -2004,9 +2002,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
req.extend_range = req.extend_range._replace(
|
req.extend_range = req.extend_range._replace(
|
||||||
start=req.extend_range.start + encoder_len
|
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)
|
req.logprob_start_len = max(req.logprob_start_len, encoder_len)
|
||||||
|
|
||||||
def prepare_for_extend(self):
|
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]
|
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]
|
prefix_lens = [len(r.prefix_indices) for r in reqs]
|
||||||
extend_lens = [r.extend_range.length 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)
|
_pin = is_pin_memory_available(self.device)
|
||||||
# Stay on pinned CPU; H2D is deferred to forward stream via
|
# 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)
|
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.
|
# tokens, and logprob_token_ids is for input logprob, so pad the rest of them by 0.
|
||||||
extend_input_logprob_token_ids.extend(
|
extend_input_logprob_token_ids.extend(
|
||||||
[0]
|
[0]
|
||||||
* (
|
* (
|
||||||
req.extend_range.length
|
req.extend_range.length
|
||||||
- req.extend_logprob_start_len
|
- extend_logprob_start_lens[i]
|
||||||
- len(logprob_token_ids)
|
- 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.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.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
|
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids
|
||||||
|
|
||||||
if get_global_server_args().enable_mamba_extra_buffer():
|
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
|
req.extend_range.length if req.extend_range is not None else 0
|
||||||
for req in batch.reqs
|
for req in batch.reqs
|
||||||
]
|
]
|
||||||
batch_result.extend_logprob_start_len_per_req = [
|
batch_result.extend_logprob_start_len_per_req = (
|
||||||
req.extend_logprob_start_len for req in batch.reqs
|
batch.extend_logprob_start_lens
|
||||||
]
|
)
|
||||||
else:
|
else:
|
||||||
batch_result.extend_input_len_per_req = None
|
batch_result.extend_input_len_per_req = None
|
||||||
batch_result.extend_logprob_start_len_per_req = None
|
batch_result.extend_logprob_start_len_per_req = None
|
||||||
|
|||||||
@@ -79,7 +79,6 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
req.priority = priority
|
req.priority = priority
|
||||||
req.prefix_indices = []
|
req.prefix_indices = []
|
||||||
req.full_untruncated_fill_ids = []
|
req.full_untruncated_fill_ids = []
|
||||||
req.extend_logprob_start_len = 0
|
|
||||||
req.output_ids = [0] * output_len
|
req.output_ids = [0] * output_len
|
||||||
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
|
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
|
||||||
req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time)
|
req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time)
|
||||||
|
|||||||
Reference in New Issue
Block a user