[Fix]: Adjust FutureMap's token_id_bufs Size to Prevent ChunkedPrefill's next_token_ids from Overwriting Previous Prefill Requests' next_token_id (#13713)
Signed-off-by: vito.yy <vito.yy@antgroup.com>
This commit is contained in:
@@ -33,14 +33,26 @@ class FutureMap:
|
||||
def __init__(
|
||||
self,
|
||||
max_running_requests: int,
|
||||
chunked_prefill_size: int,
|
||||
context_len: int,
|
||||
device: torch.device,
|
||||
spec_algo: Optional[SpeculativeAlgorithm] = None,
|
||||
):
|
||||
# FIXME: the calculation of future_limit and future_buffer_len maybe too conservative
|
||||
self.future_ct = 0
|
||||
# A factor of 3 is used to avoid collision in the circular buffer.
|
||||
self.future_limit = max_running_requests * 3
|
||||
# A factor of 5 is used to ensure the buffer is large enough.
|
||||
self.future_buffer_len = max_running_requests * 5
|
||||
|
||||
# Circular buffer layout (wraps in this order):
|
||||
# Running decode batch -> Prefill chunk 1 -> ... -> Prefill chunk N
|
||||
# A running decode batch's result will be resolved after all prefill chunks are done.
|
||||
# reserve `max_num_chunks` extra future slots on top of `max_running_requests * 3`.
|
||||
max_num_chunks = (
|
||||
(context_len + chunked_prefill_size - 1) // chunked_prefill_size
|
||||
if chunked_prefill_size
|
||||
else 0
|
||||
)
|
||||
self.future_limit = max_running_requests * 3 + max_num_chunks
|
||||
# Adding 2 * max_running_requests to future_limit ensures the buffer is sufficiently large.
|
||||
self.future_buffer_len = self.future_limit + 2 * max_running_requests
|
||||
self.device = device
|
||||
self.spec_algo = spec_algo
|
||||
self.buf_initialized = False
|
||||
|
||||
@@ -924,7 +924,11 @@ class Scheduler(
|
||||
).stream(self.copy_stream)
|
||||
|
||||
self.future_map = FutureMap(
|
||||
self.max_running_requests, self.device, self.spec_algorithm
|
||||
self.max_running_requests,
|
||||
self.chunked_prefill_size,
|
||||
self.model_config.context_len,
|
||||
self.device,
|
||||
self.spec_algorithm,
|
||||
)
|
||||
self.batch_record_buf = [None] * 2
|
||||
self.batch_record_ct = 0
|
||||
|
||||
Reference in New Issue
Block a user