Avoid scattered assignment of extend_input_len and fill_len by merging them into Req.extend_range (#27610)
This commit is contained in:
@@ -1420,7 +1420,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
# inserts committed KV into the radix tree. The last output token
|
||||
# hasn't had KV committed yet (output_ids is 1 ahead).
|
||||
req.full_untruncated_fill_ids = req.origin_input_ids + req.output_ids
|
||||
req.fill_len = req.kv_committed_len
|
||||
# Set prefix_indices so downstream consumers (init_next_round_input,
|
||||
# prepare_for_extend) see the correct prefix length. In the agg path
|
||||
# this is done inside init_next_round_input, but decode-disagg needs
|
||||
@@ -1428,7 +1427,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
req.prefix_indices = (
|
||||
prefix_indices if prefix_len > 0 else torch.empty((0,), dtype=torch.int64)
|
||||
)
|
||||
req.set_extend_input_len(req.fill_len - total_prefix_len)
|
||||
req.set_extend_range(total_prefix_len, req.kv_committed_len)
|
||||
|
||||
# Return the transfer destination indices:
|
||||
if self.scheduler.enable_hisparse:
|
||||
@@ -1909,8 +1908,7 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
# only sees committed KV (full array includes one uncommitted
|
||||
# token because init_next_round_input rebuilt it as full).
|
||||
if req.kv_committed_len is not None:
|
||||
req.fill_len = req.kv_committed_len
|
||||
req.set_extend_input_len(req.fill_len - len(req.prefix_indices))
|
||||
req.set_extend_range(len(req.prefix_indices), req.kv_committed_len)
|
||||
else:
|
||||
waiting_queue.append(req)
|
||||
|
||||
|
||||
@@ -57,7 +57,7 @@ class ReqDllmMixin:
|
||||
def _init_fill_ids_for_dllm(self: Req):
|
||||
self.dllm_block_offset = (
|
||||
0
|
||||
if self.fill_len == 0
|
||||
if not self.dllm_initialized
|
||||
else self.dllm_block_offset + self.dllm_config.block_size
|
||||
)
|
||||
self.full_untruncated_fill_ids = (
|
||||
@@ -65,7 +65,7 @@ class ReqDllmMixin:
|
||||
+ self.output_ids
|
||||
+ array("q", [self.dllm_config.mask_id] * self.dllm_config.block_size)
|
||||
)
|
||||
self.fill_len = len(self.full_untruncated_fill_ids)
|
||||
self.dllm_initialized = True
|
||||
|
||||
def _update_block_offset_for_dllm(self):
|
||||
prefix_len = len(self.prefix_indices)
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.utils.common import (
|
||||
Range,
|
||||
ceil_align,
|
||||
flatten_arrays_to_pinned_cpu,
|
||||
is_pin_memory_available,
|
||||
@@ -726,7 +727,8 @@ class Req(ReqDllmMixin):
|
||||
# Kept in sync by _refresh_fill_ids; admission only updates fill_len,
|
||||
# never mutates this array's length.
|
||||
self.full_untruncated_fill_ids = array("q")
|
||||
self.fill_len: int = 0
|
||||
self.extend_range: Optional[Range] = None
|
||||
self.dllm_initialized: bool = False
|
||||
|
||||
self.session = session
|
||||
self.input_embeds = input_embeds
|
||||
@@ -842,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)
|
||||
# Number of tokens to run prefill.
|
||||
self.extend_input_len = 0
|
||||
# The relative logprob_start_len in an extend batch
|
||||
self.extend_logprob_start_len = 0
|
||||
# TODO(ispobock): rename to last_device_node
|
||||
@@ -1096,6 +1096,18 @@ class Req(ReqDllmMixin):
|
||||
# Whether request reached finished condition
|
||||
return self.finished_reason is not None
|
||||
|
||||
@property
|
||||
def fill_len(self) -> int:
|
||||
return self.extend_range.end
|
||||
|
||||
@property
|
||||
def extend_input_len(self) -> int:
|
||||
return self.extend_range.length
|
||||
|
||||
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.fill_len]
|
||||
|
||||
@@ -1447,7 +1459,8 @@ class Req(ReqDllmMixin):
|
||||
self.num_matched_prefix_tokens = 0
|
||||
self.swa_uuid_for_lock = None
|
||||
self.swa_prefix_lock_released = False
|
||||
self.extend_input_len = 0
|
||||
self.extend_range = None
|
||||
self.dllm_initialized = False
|
||||
self.is_retracted = True
|
||||
self.retracted_stain = True
|
||||
self.input_token_logprobs = None
|
||||
@@ -1470,7 +1483,6 @@ class Req(ReqDllmMixin):
|
||||
self.swa_evicted_seqlen = 0
|
||||
self.extend_batch_idx = 0
|
||||
self.decode_batch_idx = 0
|
||||
self.fill_len = 0
|
||||
|
||||
# When using input_embeds, we cannot easily mix the original input embeddings
|
||||
# with the newly generated output token IDs during re-prefill of retracted request.
|
||||
@@ -1520,14 +1532,13 @@ class Req(ReqDllmMixin):
|
||||
logger.info(f"{prefix}: {self.time_stats.convert_to_duration()}")
|
||||
self.has_log_time_stats = True
|
||||
|
||||
def set_extend_input_len(self, extend_input_len: int):
|
||||
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
|
||||
self.extend_input_len = extend_input_len
|
||||
if self.logprob_start_len == -1:
|
||||
logprob_start_len = len(self.full_untruncated_fill_ids)
|
||||
else:
|
||||
@@ -1997,7 +2008,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
if encoder_len == 0:
|
||||
continue
|
||||
if len(req.prefix_indices) < encoder_len:
|
||||
req.extend_input_len -= encoder_len
|
||||
assert len(req.prefix_indices) == 0
|
||||
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
|
||||
)
|
||||
@@ -2382,8 +2396,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
|
||||
for req in running_batch.reqs:
|
||||
req._refresh_fill_ids()
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
req.set_extend_input_len(1)
|
||||
full_len = len(req.full_untruncated_fill_ids)
|
||||
req.set_extend_range(full_len - 1, full_len)
|
||||
|
||||
# Decode tokens of the running portion live in future_map.output_tokens_buf.
|
||||
self.input_ids = None
|
||||
|
||||
@@ -659,8 +659,7 @@ class PrefillAdder:
|
||||
* self.page_size
|
||||
)
|
||||
|
||||
req.set_extend_input_len(trunc_len)
|
||||
req.fill_len = prefix_len + trunc_len
|
||||
req.set_extend_range(prefix_len, prefix_len + trunc_len)
|
||||
|
||||
self.can_run_list.append(req)
|
||||
|
||||
@@ -684,8 +683,7 @@ class PrefillAdder:
|
||||
)
|
||||
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
|
||||
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
||||
self.can_run_list.append(req)
|
||||
|
||||
# Update budget: reserve max_new_tokens only if not truncated
|
||||
@@ -728,8 +726,7 @@ class PrefillAdder:
|
||||
)
|
||||
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
|
||||
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
@@ -839,10 +836,9 @@ class PrefillAdder:
|
||||
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.set_extend_range(
|
||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||
)
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
@@ -857,9 +853,10 @@ class PrefillAdder:
|
||||
# Chunked prefill
|
||||
trunc_len = self.rem_chunk_tokens
|
||||
|
||||
req.set_extend_input_len(trunc_len)
|
||||
assert len(req.prefix_indices) == 0
|
||||
req.fill_len = len(req.prefix_indices) + trunc_len
|
||||
req.set_extend_range(
|
||||
len(req.prefix_indices), len(req.prefix_indices) + trunc_len
|
||||
)
|
||||
self.can_run_list.append(req)
|
||||
self.new_chunked_req = req
|
||||
self._update_prefill_budget(0, trunc_len, 0, req.retracted_stain)
|
||||
@@ -980,10 +977,9 @@ 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.set_extend_range(
|
||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||
)
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
self.can_run_list.append(req)
|
||||
|
||||
self._req_inc_lock_ref(req)
|
||||
@@ -1022,8 +1018,9 @@ class PrefillAdder:
|
||||
return AddReqResult.OTHER
|
||||
|
||||
# Chunked prefill
|
||||
req.set_extend_input_len(trunc_len)
|
||||
req.fill_len = len(req.prefix_indices) + trunc_len
|
||||
req.set_extend_range(
|
||||
len(req.prefix_indices), len(req.prefix_indices) + trunc_len
|
||||
)
|
||||
|
||||
self.can_run_list.append(req)
|
||||
self.new_chunked_req = req
|
||||
|
||||
@@ -3337,7 +3337,8 @@ class Scheduler(
|
||||
# we can use the correct values in output processing.
|
||||
if batch.return_logprob:
|
||||
batch_result.extend_input_len_per_req = [
|
||||
req.extend_input_len for req in batch.reqs
|
||||
req.extend_input_len 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
|
||||
|
||||
@@ -608,9 +608,10 @@ class SchedulerPPMixin:
|
||||
sampling_params=sampling_params,
|
||||
)
|
||||
req.full_untruncated_fill_ids = req.origin_input_ids
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
req.logprob_start_len = -1
|
||||
req.set_extend_input_len(req.fill_len - len(req.prefix_indices))
|
||||
req.set_extend_range(
|
||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||
)
|
||||
|
||||
# Prepare batch
|
||||
batch = ScheduleBatch.init_new(
|
||||
|
||||
@@ -65,6 +65,7 @@ from typing import (
|
||||
Dict,
|
||||
Generic,
|
||||
List,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Protocol,
|
||||
Sequence,
|
||||
@@ -105,6 +106,15 @@ logger = logging.getLogger(__name__)
|
||||
torch_release = pkg_version.parse(torch.__version__).release
|
||||
|
||||
|
||||
class Range(NamedTuple):
|
||||
start: int
|
||||
end: int
|
||||
|
||||
@property
|
||||
def length(self) -> int:
|
||||
return self.end - self.start
|
||||
|
||||
|
||||
def flatten_arrays_to_pinned_cpu(parts: List[array[int]], pin: bool) -> torch.Tensor:
|
||||
"""Flatten array.array('q') buffers into one int64 CPU tensor.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user