Inline extend_range accessors and remove the extend_input_len/fill_len properties (#27611)
This commit is contained in:
@@ -381,9 +381,8 @@ def prepare_inputs_for_correctness_test(bench_args, tokenizer, custom_prompts):
|
||||
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.origin_input_ids))
|
||||
reqs.append(req)
|
||||
|
||||
return input_ids, reqs
|
||||
@@ -395,14 +394,15 @@ def prepare_extend_inputs_for_correctness_test(
|
||||
for i in range(len(reqs)):
|
||||
req: Req = reqs[i]
|
||||
req.full_untruncated_fill_ids.extend(input_ids[i][bench_args.cut_len :])
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
if model_runner is not None:
|
||||
# Use req.req_pool_idx instead of i to handle slot 0 padding correctly
|
||||
req.prefix_indices = model_runner.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : bench_args.cut_len
|
||||
].to(req.prefix_indices.dtype)
|
||||
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)
|
||||
)
|
||||
return reqs
|
||||
|
||||
|
||||
@@ -428,9 +428,8 @@ def prepare_synthetic_inputs_for_latency_test(
|
||||
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.origin_input_ids))
|
||||
reqs.append(req)
|
||||
|
||||
return reqs
|
||||
|
||||
@@ -1085,7 +1085,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
|
||||
@property
|
||||
def num_tokens_pre_allocated(self):
|
||||
return sum(decode_req.req.fill_len for decode_req in self.transfer_queue.queue)
|
||||
return sum(
|
||||
decode_req.req.extend_range.end for decode_req in self.transfer_queue.queue
|
||||
)
|
||||
|
||||
def _need_space_for_single_req(
|
||||
self, retractable_tokens: Optional[int] = None
|
||||
|
||||
@@ -37,7 +37,7 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
req_pool_indices = []
|
||||
|
||||
# Pre-calculate total size
|
||||
total_size = sum(req.extend_input_len for req in reqs)
|
||||
total_size = sum(req.extend_range.length for req in reqs)
|
||||
out_cache_loc = torch.empty(total_size, dtype=torch.int64, device=self.device)
|
||||
|
||||
# Fill the tensor in one pass
|
||||
@@ -47,20 +47,20 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
pre_len = len(req.prefix_indices)
|
||||
|
||||
chunk = self.req_to_token_pool.req_to_token[req.req_pool_idx][
|
||||
pre_len : pre_len + req.extend_input_len
|
||||
pre_len : pre_len + req.extend_range.length
|
||||
]
|
||||
assert (
|
||||
offset + req.extend_input_len <= total_size
|
||||
), f"Exceeds total size: offset={offset}, req.extend_input_len={req.extend_input_len}, total_size={total_size}"
|
||||
out_cache_loc[offset : offset + req.extend_input_len] = chunk
|
||||
offset += req.extend_input_len
|
||||
offset + req.extend_range.length <= total_size
|
||||
), f"Exceeds total size: offset={offset}, req.extend_range.length={req.extend_range.length}, total_size={total_size}"
|
||||
out_cache_loc[offset : offset + req.extend_range.length] = chunk
|
||||
offset += req.extend_range.length
|
||||
|
||||
seq_len = len(req.origin_input_ids) + max(0, len(req.output_ids) - 1)
|
||||
seq_lens.append(seq_len)
|
||||
if len(req.output_ids) == 0:
|
||||
assert (
|
||||
seq_len - pre_len == req.extend_input_len
|
||||
), f"seq_len={seq_len}, pre_len={pre_len}, req.extend_input_len={req.extend_input_len}"
|
||||
seq_len - pre_len == req.extend_range.length
|
||||
), f"seq_len={seq_len}, pre_len={pre_len}, req.extend_range.length={req.extend_range.length}"
|
||||
|
||||
if not req.retracted_stain:
|
||||
# Clamp to avoid double-counting: already_computed is seeded from
|
||||
@@ -99,7 +99,7 @@ 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_input_len 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.multimodal_inputs = [r.multimodal_inputs for r in reqs]
|
||||
|
||||
@@ -934,7 +934,7 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
elif self.enable_overlap:
|
||||
# Delay KV transfer to process_batch_result_disagg_prefill when overlap is enabled to ensure results are resolved
|
||||
self.chunked_req.tmp_end_idx = min(
|
||||
self.chunked_req.fill_len,
|
||||
self.chunked_req.extend_range.end,
|
||||
len(self.chunked_req.origin_input_ids),
|
||||
)
|
||||
else:
|
||||
@@ -970,7 +970,7 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
end_idx = (
|
||||
end_idx
|
||||
if end_idx is not None
|
||||
else min(req.fill_len, len(req.origin_input_ids))
|
||||
else min(req.extend_range.end, len(req.origin_input_ids))
|
||||
)
|
||||
|
||||
if not last_chunk:
|
||||
@@ -1001,7 +1001,7 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
# length here avoids emitting an extra state page when the sampled
|
||||
# token crosses a page boundary, which mismatched src/dst lengths in
|
||||
# group_concurrent_contiguous.
|
||||
seq_len = min(req.fill_len, len(req.origin_input_ids))
|
||||
seq_len = min(req.extend_range.end, len(req.origin_input_ids))
|
||||
|
||||
def _mamba_payload():
|
||||
return [
|
||||
|
||||
@@ -81,7 +81,7 @@ class SchedulerDllmMixin:
|
||||
continue
|
||||
|
||||
req.full_untruncated_fill_ids[
|
||||
req.fill_len - new_tokens : req.fill_len
|
||||
req.extend_range.end - new_tokens : req.extend_range.end
|
||||
] = array("q", next_token_ids)
|
||||
self.metrics_reporter.num_generated_tokens += new_tokens
|
||||
|
||||
|
||||
@@ -217,7 +217,7 @@ class HiSparseCoordinator:
|
||||
req.hisparse_staging = True
|
||||
|
||||
full_kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
].to(dtype=torch.int64, copy=True)
|
||||
device_indices = (
|
||||
self.mem_pool_device.translate_loc_from_full_to_hisparse_device(
|
||||
@@ -308,7 +308,7 @@ class HiSparseCoordinator:
|
||||
|
||||
def alloc_device_buffer(self, req: Req) -> None:
|
||||
if self.is_dsv4_hisparse:
|
||||
allocated_len = req.fill_len
|
||||
allocated_len = req.extend_range.end
|
||||
alloc_size = self.padded_buffer_size
|
||||
else:
|
||||
allocated_len = req.kv_allocated_len
|
||||
@@ -729,7 +729,7 @@ class HiSparseCoordinator:
|
||||
# Wait for any in-flight staging DMA to complete before freeing
|
||||
self.write_staging_stream.synchronize()
|
||||
|
||||
prefill_len = req.fill_len
|
||||
prefill_len = req.extend_range.end
|
||||
allocated_locs = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, :prefill_len
|
||||
]
|
||||
|
||||
@@ -993,8 +993,8 @@ class Req(ReqDllmMixin):
|
||||
# the start index of the sent kv cache
|
||||
# We want to send it chunk by chunk for chunked prefill.
|
||||
# After every chunk forward, we do the following:
|
||||
# kv_send(req.input_ids[req.start_send_idx:req.fill_len])
|
||||
# start_send_idx = req.fill_len
|
||||
# kv_send(req.input_ids[req.start_send_idx:req.extend_range.end])
|
||||
# start_send_idx = req.extend_range.end
|
||||
self.start_send_idx: int = 0
|
||||
|
||||
# For overlap schedule, we delay the kv transfer until `process_batch_result_disagg_prefill` rather than `process_prefill_chunk` in non-overlap
|
||||
@@ -1096,20 +1096,12 @@ 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]
|
||||
return self.full_untruncated_fill_ids[: self.extend_range.end]
|
||||
|
||||
def _refresh_fill_ids(self) -> None:
|
||||
"""Keep full_untruncated_fill_ids == origin_input_ids + output_ids by
|
||||
@@ -1546,7 +1538,7 @@ class Req(ReqDllmMixin):
|
||||
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_input_len,
|
||||
self.extend_range.length,
|
||||
)
|
||||
|
||||
def set_finish_with_abort(self, error_msg: str):
|
||||
@@ -1667,7 +1659,7 @@ def _compute_chunked_req_next_prompt_token(
|
||||
multimodal placeholder (hash) tokens that lie outside the model vocab."""
|
||||
if chunked_req is None:
|
||||
return None
|
||||
fill_len = chunked_req.fill_len
|
||||
fill_len = chunked_req.extend_range.end
|
||||
origin_ids = chunked_req.origin_input_ids
|
||||
if fill_len >= len(origin_ids):
|
||||
return None
|
||||
@@ -1932,17 +1924,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
input_ids[i] = input_ids[i][encoder_len:]
|
||||
encoder_out_cache_loc.append(self.out_cache_loc[pt : pt + encoder_len])
|
||||
decoder_out_cache_loc.append(
|
||||
self.out_cache_loc[pt + encoder_len : pt + req.extend_input_len]
|
||||
self.out_cache_loc[pt + encoder_len : pt + req.extend_range.length]
|
||||
)
|
||||
self.extend_lens[i] -= encoder_len
|
||||
self.extend_num_tokens -= encoder_len
|
||||
else:
|
||||
decoder_out_cache_loc.append(
|
||||
self.out_cache_loc[pt : pt + req.extend_input_len]
|
||||
self.out_cache_loc[pt : pt + req.extend_range.length]
|
||||
)
|
||||
self.prefix_lens[i] -= encoder_len
|
||||
|
||||
pt += req.extend_input_len
|
||||
pt += req.extend_range.length
|
||||
|
||||
# Reassign: ED stripping rebuilds prefill_input_ids_cpu (CPU pinned);
|
||||
# resolve_forward_inputs will H2D this on forward stream. self.input_ids
|
||||
@@ -1977,7 +1969,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
for i, req in enumerate(self.reqs):
|
||||
encoder_len = self.encoder_lens_cpu[i]
|
||||
old_start_len = self.extend_logprob_start_lens[i]
|
||||
old_contribution = req.extend_input_len - old_start_len
|
||||
old_contribution = req.extend_range.length - old_start_len
|
||||
|
||||
if len(req.prefix_indices) < encoder_len:
|
||||
tokens_to_strip = max(0, encoder_len - old_start_len)
|
||||
@@ -2028,10 +2020,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
reqs = self.reqs
|
||||
input_ids = [r.get_fill_ids()[len(r.prefix_indices) :] for r in reqs]
|
||||
extend_num_tokens = sum(len(ids) for ids in input_ids)
|
||||
seq_lens = [r.fill_len for r in reqs]
|
||||
orig_seq_lens = [max(r.fill_len, len(r.origin_input_ids)) for r in reqs]
|
||||
seq_lens = [r.extend_range.end 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]
|
||||
extend_lens = [r.extend_input_len for r in reqs]
|
||||
extend_lens = [r.extend_range.length for r in reqs]
|
||||
|
||||
_pin = is_pin_memory_available(self.device)
|
||||
# Stay on pinned CPU; H2D is deferred to forward stream via
|
||||
@@ -2071,7 +2063,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
mamba_track_seqlens_cpu = []
|
||||
|
||||
for i, (req, seq_len, pre_len) in enumerate(zip(reqs, seq_lens, prefix_lens)):
|
||||
assert seq_len - pre_len == req.extend_input_len
|
||||
assert seq_len - pre_len == req.extend_range.length
|
||||
|
||||
req.extend_batch_idx += 1
|
||||
|
||||
@@ -2084,7 +2076,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# Slice to match extend_input_len — PrefillAdder truncates
|
||||
# fill_len/extend_input_len on chunk overflow but not input_embeds.
|
||||
input_embeds.extend(
|
||||
req.input_embeds[pre_len : pre_len + req.extend_input_len]
|
||||
req.input_embeds[pre_len : pre_len + req.extend_range.length]
|
||||
)
|
||||
|
||||
if req.positional_embed_overrides is not None:
|
||||
@@ -2096,7 +2088,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
req.positional_embed_overrides.positions
|
||||
):
|
||||
extend_pos = pos - pre_len
|
||||
if extend_pos < 0 or extend_pos >= req.extend_input_len:
|
||||
if extend_pos < 0 or extend_pos >= req.extend_range.length:
|
||||
continue # Outside current extend chunk, skip
|
||||
embeds_to_add.append((embed_idx, input_id_pointer + extend_pos))
|
||||
if embeds_to_add:
|
||||
@@ -2165,7 +2157,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# extend_input_logprob_token_id = [4, 0]
|
||||
global_start_idx, global_end_idx = (
|
||||
len(req.prefix_indices),
|
||||
req.fill_len,
|
||||
req.extend_range.end,
|
||||
)
|
||||
if req.logprob_start_len == -1:
|
||||
logprob_start_len = len(req.origin_input_ids)
|
||||
@@ -2180,12 +2172,12 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
]
|
||||
extend_input_logprob_token_ids.extend(logprob_token_ids)
|
||||
|
||||
# We will need req.extend_input_len - req.extend_logprob_start_len number of
|
||||
# We will need req.extend_range.length - req.extend_logprob_start_len 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_input_len
|
||||
req.extend_range.length
|
||||
- req.extend_logprob_start_len
|
||||
- len(logprob_token_ids)
|
||||
)
|
||||
@@ -2295,7 +2287,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# to force the math calculation to retrieve the correct mamba state from h.
|
||||
return i + 1
|
||||
|
||||
mask = req.extend_input_len >= mamba_cache_chunk_size
|
||||
mask = req.extend_range.length >= mamba_cache_chunk_size
|
||||
track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
||||
mamba_track_seqlen = -1
|
||||
if mask:
|
||||
@@ -2306,13 +2298,13 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# otherwise retrieved from h (i.e. unaligned).
|
||||
# We need to pass the non-aligned seqlen to the calculation. Even though
|
||||
# we pass in mamba_track_seqlen, the actual tracked seqlen is mamba_last_track_seqlen.
|
||||
mamba_track_seqlen = len(req.prefix_indices) + req.extend_input_len
|
||||
mamba_track_seqlen = len(req.prefix_indices) + req.extend_range.length
|
||||
|
||||
# mamba_track_seqlen_aligned/mamba_last_track_seqlen is actual tracked seqlen. Used to pass to
|
||||
# mamba radix cache to track which seqlen this mamba state should store at.
|
||||
mamba_track_seqlen_aligned = (
|
||||
len(req.prefix_indices)
|
||||
+ (req.extend_input_len // mamba_cache_chunk_size)
|
||||
+ (req.extend_range.length // mamba_cache_chunk_size)
|
||||
* mamba_cache_chunk_size
|
||||
)
|
||||
|
||||
@@ -2322,7 +2314,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# by _force_track_h()
|
||||
mamba_track_fla_chunk_aligned = (
|
||||
len(req.prefix_indices)
|
||||
+ (req.extend_input_len // mamba_cache_chunk_size)
|
||||
+ (req.extend_range.length // mamba_cache_chunk_size)
|
||||
* mamba_cache_chunk_size
|
||||
)
|
||||
if mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned:
|
||||
|
||||
@@ -693,7 +693,7 @@ class PrefillAdder:
|
||||
else 0
|
||||
)
|
||||
self._update_prefill_budget(
|
||||
0, req.extend_input_len, max_new_tokens, req.retracted_stain
|
||||
0, req.extend_range.length, max_new_tokens, req.retracted_stain
|
||||
)
|
||||
|
||||
# Return based on remaining token availability
|
||||
@@ -730,7 +730,7 @@ class PrefillAdder:
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
req.extend_input_len,
|
||||
req.extend_range.length,
|
||||
(
|
||||
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS)
|
||||
if not truncated
|
||||
@@ -842,7 +842,7 @@ class PrefillAdder:
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
req.extend_input_len,
|
||||
req.extend_range.length,
|
||||
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
|
||||
req.retracted_stain,
|
||||
)
|
||||
|
||||
@@ -1293,7 +1293,7 @@ class Scheduler(
|
||||
request_lengths = []
|
||||
for req in batch.reqs:
|
||||
start = len(req.prefix_indices)
|
||||
end = start + req.extend_input_len
|
||||
end = start + req.extend_range.length
|
||||
fill_ids = req.origin_input_ids + req.output_ids
|
||||
if start == 0:
|
||||
tokens = fill_ids[start:end]
|
||||
@@ -2623,7 +2623,7 @@ class Scheduler(
|
||||
# beyond what is already cached. A parked chunk (add_chunked_req
|
||||
# hybrid-SWA early-return) leaves fill_len == len(prefix_indices),
|
||||
# so there is nothing new to cache and stashing would be a no-op.
|
||||
if self.chunked_req.fill_len > len(self.chunked_req.prefix_indices):
|
||||
if self.chunked_req.extend_range.end > len(self.chunked_req.prefix_indices):
|
||||
self.stash_chunked_request(self.chunked_req)
|
||||
|
||||
# HiSparse has its own prefill-to-decode transition; skip last_batch merge.
|
||||
@@ -2974,7 +2974,7 @@ class Scheduler(
|
||||
self.enable_priority_scheduling,
|
||||
num_pending_tokens=self.load_inquirer._get_num_pending_tokens(
|
||||
chunk_deduct=(
|
||||
self.chunked_req.extend_input_len
|
||||
self.chunked_req.extend_range.length
|
||||
if self.chunked_req is not None
|
||||
else 0
|
||||
),
|
||||
@@ -3337,7 +3337,7 @@ 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 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
|
||||
]
|
||||
batch_result.extend_logprob_start_len_per_req = [
|
||||
|
||||
@@ -624,7 +624,7 @@ class SchedulerPPMixin:
|
||||
self.spec_algorithm,
|
||||
)
|
||||
|
||||
current_seq_len = req.fill_len
|
||||
current_seq_len = req.extend_range.end
|
||||
|
||||
if is_dp_attention_enabled():
|
||||
# For profiling, we only have one request on PP0
|
||||
@@ -695,7 +695,7 @@ class SchedulerPPMixin:
|
||||
# Release KV cache
|
||||
if req.req_pool_idx is not None:
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
]
|
||||
self.token_to_kv_pool_allocator.free(kv_indices)
|
||||
self.req_to_token_pool.free(req)
|
||||
|
||||
@@ -86,7 +86,7 @@ class ChunkCache(BasePrefixCache):
|
||||
|
||||
def cache_unfinished_req(self, req: Req, chunked=False):
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
]
|
||||
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
|
||||
req.prefix_indices = kv_indices.to(dtype=torch.int64, copy=True)
|
||||
|
||||
@@ -627,7 +627,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
|
||||
def _skip_cache_unfinished_req(req: Req) -> None:
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
]
|
||||
|
||||
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
|
||||
|
||||
@@ -488,7 +488,7 @@ class SWARadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
"""Cache request when it is unfinished."""
|
||||
if self.disable:
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
]
|
||||
|
||||
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
|
||||
|
||||
@@ -359,7 +359,7 @@ class StreamingSession(BasePrefixCache):
|
||||
return False
|
||||
if chunked:
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.req_pool_idx, : req.fill_len
|
||||
req.req_pool_idx, : req.extend_range.end
|
||||
]
|
||||
req.prefix_indices = kv_indices.to(dtype=torch.int64, copy=True)
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user