From 83567c30169b9fc6df5643b5491e690aa6b97322 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 01:30:32 +0000 Subject: [PATCH 01/15] feat: support multi-token KV cache pages --- docs/kv_cache_page_size.md | 43 +++++++ lightllm/common/basemodel/attention/fa3/fp.py | 37 +++--- .../common/basemodel/attention/fa3/mla.py | 9 +- .../basemodel/attention/flashinfer/fp.py | 105 ++++++++++++------ .../basemodel/attention/flashinfer/mla.py | 55 ++++++--- lightllm/common/basemodel/basemodel.py | 26 ++++- lightllm/common/basemodel/cuda_graph.py | 2 + .../common/basemodel/prefill_cuda_graph.py | 8 ++ .../basemodel/triton_kernel/fa3_utils.py | 7 +- .../triton_kernel/repack_kv_index.py | 50 +++++++++ .../deepseek2_mem_manager.py | 2 +- .../kv_cache_mem_manager/mem_manager.py | 18 ++- lightllm/common/req_manager.py | 2 + lightllm/server/api_cli.py | 7 ++ lightllm/server/api_start.py | 22 ++++ lightllm/server/core/objs/start_args_type.py | 1 + .../router/dynamic_prompt/radix_cache.py | 53 ++++++--- .../server/router/model_infer/infer_batch.py | 36 ++++-- .../model_infer/mode_backend/base_backend.py | 22 ++-- .../mode_backend/generic_pre_process.py | 52 +++++++-- .../server/router/req_queue/base_queue.py | 5 + .../req_queue/chunked_prefill/beam_impl.py | 13 +-- .../router/req_queue/chunked_prefill/impl.py | 13 +-- .../common/basemodel/test_model_input.py | 6 +- .../common/basemodel/test_model_output.py | 2 +- .../common/basemodel/test_overlap_utils.py | 2 +- .../basemodel/triton_kernel/test_fa3_utils.py | 14 ++- .../triton_kernel/test_repack_kv_index.py | 31 +++++- unit_tests/common/test_req_manager_page.py | 87 +++++++++++++++ .../router/dynamic_prompt/test_radix_cache.py | 29 +++++ .../mode_backend/test_generic_pre_process.py | 1 + .../mtp_speculative/test_eagle_overlap.py | 8 +- 32 files changed, 622 insertions(+), 146 deletions(-) create mode 100644 docs/kv_cache_page_size.md create mode 100644 unit_tests/common/test_req_manager_page.py diff --git a/docs/kv_cache_page_size.md b/docs/kv_cache_page_size.md new file mode 100644 index 0000000000..f01e869838 --- /dev/null +++ b/docs/kv_cache_page_size.md @@ -0,0 +1,43 @@ +# KV Cache 多 Token 页设计 + +## 目标与不变量 + +启动参数 `--page_size N` 控制注意力 KV Cache 的物理页大小,默认值为 `1`。`N > 1` 时保持 token 级的 +KV 存储和请求表,但资源的申请、缓存和释放以完整物理页为单位。 + +实现依赖以下不变量: + +1. `cur_kv_len` 是请求已经写入有效 KV 的逻辑 token 数;`hold_kv_len` 是请求拥有的物理容量,始终满足 + `cur_kv_len <= hold_kv_len` 且 `hold_kv_len % page_size == 0`。 +2. 一个物理页内的 KV 槽连续,页首索引可被 `page_size` 整除。请求表保存已持有页的全部 token 索引, + 包括尚未使用的尾部槽位。 +3. `ModelInput.mem_indexes` / `InferState.mem_index` 只包含本轮真实参与计算的 token,不能包含预留尾部。 +4. Radix Cache 只插入、拆分、命中和淘汰完整页;不足一页的请求尾部在请求结束或暂停时整页回收。 +5. `page_size=1` 继续走原有批量申请和 token 级页表路径,不改变默认行为。 + +页容量计算、整页申请和本轮索引组装由 Prefill/Decode 输入构造层负责;`ReqManager` 只保留请求 ID、 +请求表及其原有生命周期职责,不提供 page allocator 接口。 + +## 生命周期 + +- Prefill:按每个请求的目标 KV 长度将容量向上对齐。新页一次申请并写入请求表,本轮只返回 + `[cur_kv_len, target_kv_len)` 对应的真实索引。 +- Decode:若尾页仍有预留槽位,不再访问 allocator;跨页时申请一个新页。 +- Prefix Cache 命中:命中长度天然是整页,`cur_kv_len` 与 `hold_kv_len` 同时初始化为命中长度。 +- 完成/暂停:完整逻辑页可进入 Radix Cache;重复页和未完成尾页按物理页展开后释放。 +- Attention:FA3/FlashInfer 页表每项由物理 token 页首索引除以 `page_size` 得到,KV buffer 视为 + `[page_count, page_size, ...]`。 + +## 当前兼容范围 + +首版支持普通与分块 Prefill、Decode、动态 Prompt Cache,以及非量化 FA3/FlashInfer(含 MLA Decode)。 +MTP、PD 分离、CPU KV Cache、DP Prompt Cache 拉取、DP Prefill Balance、diverse mode 和混合线性注意力 +拥有额外的 KV 申请或迁移语义;这些组合在启动阶段明确报错,后续应在各自模块接入统一的页所有权接口后再开放。 + +## 边界处理 + +- KV 总容量和请求表宽度分别向下、向上对齐,防止物理页越界。 +- CUDA Graph 的 HOLD request 映射到额外保留的一整个物理页;请求表按 + `[HOLD, HOLD+1, ..., HOLD+page_size-1]` 循环填充,而不是重复同一个 token 索引。 +- Radix 子节点用“首个完整 token 页”作为键,避免不同序列仅首 token 相同造成页级分支冲突。 +- 非法的 `page_size < 1` 以及尚未支持的功能组合在模型加载前失败。 diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 7ba00e5911..bfdd35a07d 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -1,5 +1,6 @@ import dataclasses import torch +import triton from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional, TYPE_CHECKING from lightllm.utils.sgl_utils import flash_attn_with_kvcache, flash_attn_with_kvcache_autotune @@ -18,6 +19,7 @@ class Fa3AttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) + self.page_size = model.args.page_size # 延迟到首次获取 page table 时再初始化,避免 PD 分离模式下的 prefill 节点 # 分配仅供 decode 使用的 buffer,减少显存浪费。 @@ -37,7 +39,7 @@ def _init_page_table_buffers(self): self.page_table_max_batch_size = max(running_max_batch_size, model.graph_max_batch_size) # max_seq_length is max_req_total_len plus the MTP headroom reserved when # the model is initialized. - self.page_table_max_seq_len = model.max_seq_length + self.page_table_max_seq_len = triton.cdiv(model.max_seq_length, self.page_size) buffer_count = 2 if args.enable_decode_microbatch_overlap else 1 workspace_size = self.page_table_max_batch_size * self.page_table_max_seq_len self.page_table_buffers = [ @@ -57,12 +59,13 @@ def get_page_table_view(self, att_batch_size, max_kv_len, microbatch_index): f"FA3 attention batch size {att_batch_size} exceeds page-table capacity " f"{self.page_table_max_batch_size}" ) - if max_kv_len > self.page_table_max_seq_len: + max_page_len = triton.cdiv(max_kv_len, self.page_size) + if max_page_len > self.page_table_max_seq_len: raise RuntimeError( f"FA3 max KV sequence length {max_kv_len} exceeds page-table capacity " f"{self.page_table_max_seq_len}" ) - return self.page_table_buffers[microbatch_index][: att_batch_size * max_kv_len].reshape( - att_batch_size, max_kv_len + return self.page_table_buffers[microbatch_index][: att_batch_size * max_page_len].reshape( + att_batch_size, max_page_len ) def create_att_prefill_state(self, infer_state) -> "Fa3PrefillAttState": @@ -84,14 +87,18 @@ def init_state(self): self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() self.page_table = torch.empty( - (self.infer_state.batch_size, self.infer_state.max_kv_seq_len), + ( + self.infer_state.batch_size, + triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size), + ), dtype=torch.int32, device=self.infer_state.input_ids.device, ) - self.page_table.copy_( - self.infer_state.req_manager.req_to_token_indexs[ - self.infer_state.b_req_idx, : self.infer_state.max_kv_seq_len - ] + page_table_copy( + page_table=self.page_table, + req_to_token_indexs=self.infer_state.req_manager.req_to_token_indexs, + b_req_idx=self.infer_state.b_req_idx, + page_size=self.backend.page_size, ) def prefill_att( @@ -131,8 +138,8 @@ def _nomarl_prefill_att( sm_scale = 1.0 / (Lq ** 0.5) o = flash_attn_with_kvcache( q=q, - k_cache=k.view(k.shape[0], 1, k.shape[1], k.shape[2]), - v_cache=v.view(v.shape[0], 1, v.shape[1], v.shape[2]), + k_cache=k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]), + v_cache=v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]), page_table=self.page_table, cache_seqlens=self.infer_state.b_seq_len, cu_seqlens_q=self.cu_seqlens_q, @@ -225,6 +232,7 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor): att_batch_size = b_att_req_idx.shape[0] model = self.backend.model actual_max_kv_len = self.infer_state.max_kv_seq_len + actual_max_page_len = triton.cdiv(actual_max_kv_len, self.backend.page_size) page_table_width = actual_max_kv_len if model.graph is not None and model.graph.can_run( batch_size=self.infer_state.batch_size, @@ -242,9 +250,10 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor): ) page_table_copy( - page_table=self.page_table[:, :actual_max_kv_len], + page_table=self.page_table[:, :actual_max_page_len], req_to_token_indexs=model.req_manager.req_to_token_indexs, b_req_idx=b_att_req_idx, + page_size=self.backend.page_size, ) def copy_for_decode_cuda_graph(self, new_state: "Fa3DecodeAttState"): @@ -290,8 +299,8 @@ def _normal_decode_att( sm_scale = 1.0 / (Lq ** 0.5) o = flash_attn_with_kvcache_autotune( q=q, - k_cache=k.view(k.shape[0], 1, k.shape[1], k.shape[2]), - v_cache=v.view(v.shape[0], 1, v.shape[1], v.shape[2]), + k_cache=k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]), + v_cache=v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]), page_table=self.page_table, cache_seqlens=self.b_att_seq_len, cu_seqlens_q=self.cu_seqlens_q, diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 6e64ec7c97..7f85bbcc8d 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -1,5 +1,6 @@ import dataclasses import torch +import triton from ..base_att import BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional, TYPE_CHECKING, Tuple from lightllm.utils.sgl_utils import flash_attn_with_kvcache @@ -153,6 +154,7 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor): att_batch_size = b_att_req_idx.shape[0] model = self.backend.model actual_max_kv_len = self.infer_state.max_kv_seq_len + actual_max_page_len = triton.cdiv(actual_max_kv_len, self.backend.page_size) page_table_width = actual_max_kv_len if model.graph is not None and model.graph.can_run( batch_size=self.infer_state.batch_size, @@ -170,9 +172,10 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor): ) page_table_copy( - page_table=self.page_table[:, :actual_max_kv_len], + page_table=self.page_table[:, :actual_max_page_len], req_to_token_indexs=model.req_manager.req_to_token_indexs, b_req_idx=b_att_req_idx, + page_size=self.backend.page_size, ) def copy_for_decode_cuda_graph(self, new_state: "MlaFa3DecodeAttState"): @@ -213,8 +216,8 @@ def _mla_decode_att( kv = k qk_rope_head_dim = 64 kv_lora_rank = kv.shape[-1] - qk_rope_head_dim - k_rope = kv[:, :, -qk_rope_head_dim:].view(-1, 1, 1, qk_rope_head_dim) - kv_nope = kv[:, :, :-qk_rope_head_dim].view(-1, 1, 1, kv_lora_rank) + k_rope = kv[:, :, -qk_rope_head_dim:].view(-1, self.backend.page_size, 1, qk_rope_head_dim) + kv_nope = kv[:, :, :-qk_rope_head_dim].view(-1, self.backend.page_size, 1, kv_lora_rank) k_descale, v_descale = None, None assert att_control.mla_decode softmax_scale = att_control.mla_decode_dict["softmax_scale"] diff --git a/lightllm/common/basemodel/attention/flashinfer/fp.py b/lightllm/common/basemodel/attention/flashinfer/fp.py index ec0b42dca3..ed00d67af5 100644 --- a/lightllm/common/basemodel/attention/flashinfer/fp.py +++ b/lightllm/common/basemodel/attention/flashinfer/fp.py @@ -1,8 +1,9 @@ import dataclasses import torch +import triton from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from lightllm.utils.dist_utils import get_dp_world_size, get_current_device_id -from ...triton_kernel.repack_kv_index import repack_kv_index +from ...triton_kernel.repack_kv_index import repack_kv_index, repack_page_kv_index from .env_utils import set_flashinfer_envs from .utils import should_init_decode_wrapper @@ -14,6 +15,7 @@ class FlashInferAttBackend(BaseAttBackend): def __init__(self, model): set_flashinfer_envs() super().__init__(model=model) + self.page_size = model.args.page_size tp_world_size = get_dp_world_size() self.tp_q_head_num = model.config["num_attention_heads"] // tp_world_size self.tp_kv_head_num = max(model.config["num_key_value_heads"] // tp_world_size, 1) @@ -22,10 +24,14 @@ def __init__(self, model): self.max_seq_length = model.max_seq_length self.kv_indices_buffer = [ torch.empty( - model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id() + model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size), + dtype=torch.int32, + device=get_current_device_id(), ), torch.empty( - model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id() + model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size), + dtype=torch.int32, + device=get_current_device_id(), ), ] self.q_data_type = model.data_type @@ -51,21 +57,34 @@ def init_state(self): device = self.infer_state.input_ids.device q_starts = self.infer_state.b1_cu_q_seq_len.int() - kv_starts = self.infer_state.b1_cu_kv_seq_len.int() - kv_last_page_len = torch.full((batch_size,), 1, dtype=torch.int32, device=device) + kv_starts = self.infer_state.b1_cu_kv_seq_len.int().clone() + b_page_len = triton.cdiv(self.infer_state.b_seq_len, self.backend.page_size) + kv_starts[1:] = b_page_len.cumsum(0) + kv_last_page_len = self.infer_state.b_seq_len - (b_page_len - 1) * self.backend.page_size kv_indices = torch.empty( - batch_size * self.backend.max_seq_length, + batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size), dtype=torch.int32, device=device, ) - repack_kv_index( - self.infer_state.req_manager.req_to_token_indexs, - self.infer_state.b_req_idx, - self.infer_state.b_seq_len, - kv_starts[:-1], - self.infer_state.max_kv_seq_len, - kv_indices, - ) + if self.backend.page_size == 1: + repack_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + self.infer_state.b_seq_len, + kv_starts[:-1], + self.infer_state.max_kv_seq_len, + kv_indices, + ) + else: + repack_page_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + b_page_len, + kv_starts[:-1], + triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size), + kv_indices, + self.backend.page_size, + ) self.prefill_wrapper = flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper( self.backend.get_gpu_workspace_buffer( key_name=self.backend.workspace_buffer_key, @@ -84,7 +103,7 @@ def init_state(self): self.backend.tp_q_head_num, self.backend.tp_kv_head_num, self.backend.head_dim, - 1, + self.backend.page_size, causal=True, pos_encoding_mode="NONE", logits_soft_cap=0.0, @@ -119,7 +138,10 @@ def _nomarl_prefill_att( o_tensor = alloc_func(q.shape, q.dtype, device="cuda") self.prefill_wrapper.run( q, - (k.unsqueeze(1), v.unsqueeze(1)), + ( + k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]), + v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]), + ), out=o_tensor, ) return o_tensor @@ -141,30 +163,42 @@ def init_state(self): self.backend: FlashInferAttBackend = self.backend device = self.infer_state.input_ids.device model = self.backend.model - self.kv_last_page_len_buffer = torch.full((self.infer_state.batch_size,), 1, dtype=torch.int32, device=device) + b_page_len = triton.cdiv(self.infer_state.b_seq_len, self.backend.page_size) + self.kv_last_page_len_buffer = self.infer_state.b_seq_len - (b_page_len - 1) * self.backend.page_size + buffer_len = self.infer_state.batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size) if ( self.infer_state.batch_size <= model.graph_max_batch_size and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch ): - self.kv_indices = self.backend.kv_indices_buffer[self.infer_state.microbatch_index][ - : self.infer_state.batch_size * self.backend.max_seq_length - ] + self.kv_indices = self.backend.kv_indices_buffer[self.infer_state.microbatch_index][:buffer_len] else: self.kv_indices = torch.empty( - self.infer_state.batch_size * self.backend.max_seq_length, + buffer_len, dtype=torch.int32, device=device, ) - repack_kv_index( - self.infer_state.req_manager.req_to_token_indexs, - self.infer_state.b_req_idx, - self.infer_state.b_seq_len, - self.infer_state.b_kv_start_loc, - self.infer_state.max_kv_seq_len, - self.kv_indices, - ) - self.kv_starts = self.infer_state.b1_cu_kv_seq_len.int() + self.kv_starts = self.infer_state.b1_cu_kv_seq_len.int().clone() + self.kv_starts[1:] = b_page_len.cumsum(0) + if self.backend.page_size == 1: + repack_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + self.infer_state.b_seq_len, + self.kv_starts[:-1], + self.infer_state.max_kv_seq_len, + self.kv_indices, + ) + else: + repack_page_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + b_page_len, + self.kv_starts[:-1], + triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size), + self.kv_indices, + self.backend.page_size, + ) if not self._should_init_decode_wrapper(): # 处于 graph replay 回放阶段,不需要特殊初始化 decode wrapper。 return @@ -189,7 +223,7 @@ def init_state(self): self.backend.tp_q_head_num, self.backend.tp_kv_head_num, self.backend.head_dim, - 1, + self.backend.page_size, q_data_type=self.backend.q_data_type, kv_data_type=self.backend.kv_data_type, non_blocking=True, @@ -210,7 +244,7 @@ def _refresh_cuda_graph_decode_plan(self, max_kv_len: int): dtype=torch.int32, device="cpu", ) - * max_kv_len + * triton.cdiv(max_kv_len, self.backend.page_size) ) fast_decode_plan( @@ -221,7 +255,7 @@ def _refresh_cuda_graph_decode_plan(self, max_kv_len: int): num_qo_heads=self.backend.tp_q_head_num, num_kv_heads=self.backend.tp_kv_head_num, head_dim=self.backend.head_dim, - page_size=1, + page_size=self.backend.page_size, q_data_type=self.backend.q_data_type, kv_data_type=self.backend.kv_data_type, non_blocking=True, @@ -259,7 +293,10 @@ def _normal_decode_att( o_tensor = alloc_func(q.shape, q.dtype) self.decode_wrapper.run( q, - (k.unsqueeze(1), v.unsqueeze(1)), + ( + k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]), + v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]), + ), out=o_tensor, ) return o_tensor diff --git a/lightllm/common/basemodel/attention/flashinfer/mla.py b/lightllm/common/basemodel/attention/flashinfer/mla.py index 2da8b4423c..f7d7b9712c 100644 --- a/lightllm/common/basemodel/attention/flashinfer/mla.py +++ b/lightllm/common/basemodel/attention/flashinfer/mla.py @@ -1,8 +1,9 @@ import dataclasses import torch +import triton from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from lightllm.utils.dist_utils import get_dp_world_size, get_current_device_id -from ...triton_kernel.repack_kv_index import repack_kv_index +from ...triton_kernel.repack_kv_index import repack_kv_index, repack_page_kv_index from ...triton_kernel.flashinfer_mla_plan import fill_mla_decode_plan_for_cuda_graph from typing import Tuple from .env_utils import set_flashinfer_envs @@ -16,6 +17,7 @@ class MlaFlashInferAttBackend(BaseAttBackend): def __init__(self, model): set_flashinfer_envs() super().__init__(model=model) + self.page_size = model.args.page_size num_heads = model.config["num_attention_heads"] self.tp_q_head_num = num_heads // get_dp_world_size() self.qk_nope_head_dim = model.qk_nope_head_dim @@ -28,10 +30,14 @@ def __init__(self, model): self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) self.kv_indices_buffer = [ torch.empty( - model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id() + model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size), + dtype=torch.int32, + device=get_current_device_id(), ), torch.empty( - model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id() + model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size), + dtype=torch.int32, + device=get_current_device_id(), ), ] @@ -135,29 +141,42 @@ def init_state(self): device = self.infer_state.input_ids.device batch_size = self.infer_state.batch_size - self.kv_starts = self.infer_state.b1_cu_kv_seq_len + self.kv_starts = self.infer_state.b1_cu_kv_seq_len.clone() + b_page_len = triton.cdiv(self.infer_state.b_seq_len, self.backend.page_size) + self.kv_starts[1:] = b_page_len.cumsum(0) self.q_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device="cuda") self.q_indptr_host = torch.arange(batch_size + 1, dtype=torch.int32, device="cpu") if batch_size <= model.graph_max_batch_size and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch: self.kv_indices = self.backend.kv_indices_buffer[self.infer_state.microbatch_index][ - : batch_size * self.backend.max_seq_length + : batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size) ] else: self.kv_indices = torch.empty( - batch_size * self.backend.max_seq_length, + batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size), dtype=torch.int32, device=device, ) - repack_kv_index( - self.infer_state.req_manager.req_to_token_indexs, - self.infer_state.b_req_idx, - self.infer_state.b_seq_len, - self.infer_state.b_kv_start_loc, - self.infer_state.max_kv_seq_len, - self.kv_indices, - ) + if self.backend.page_size == 1: + repack_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + self.infer_state.b_seq_len, + self.kv_starts[:-1], + self.infer_state.max_kv_seq_len, + self.kv_indices, + ) + else: + repack_page_kv_index( + self.infer_state.req_manager.req_to_token_indexs, + self.infer_state.b_req_idx, + b_page_len, + self.kv_starts[:-1], + triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size), + self.kv_indices, + self.backend.page_size, + ) if not self._should_init_decode_wrapper(): return @@ -183,7 +202,7 @@ def init_state(self): self.backend.tp_q_head_num, self.backend.kv_lora_rank, self.backend.qk_rope_head_dim, - 1, + self.backend.page_size, False, # causal self.backend.softmax_scale, self.backend.q_data_type, @@ -204,7 +223,7 @@ def _refresh_cuda_graph_decode_plan(self, max_kv_len: int): self.kv_starts, self.infer_state.batch_size, self.backend.tp_q_head_num, - max_kv_len, + triton.cdiv(max_kv_len, self.backend.page_size), ) def decode_att( @@ -247,8 +266,8 @@ def _mla_decode_att( self.decode_wrapper.run( q_nope, q_rope, - k[:, :, :-qk_rope_head_dim], - k[:, :, -qk_rope_head_dim:], + k[:, :, :-qk_rope_head_dim].view(-1, self.backend.page_size, 1, k.shape[-1] - qk_rope_head_dim), + k[:, :, -qk_rope_head_dim:].view(-1, self.backend.page_size, 1, qk_rope_head_dim), out=o_tensor, return_lse=False, ) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index e80f2b552f..3914cc6b2e 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -123,6 +123,14 @@ def __init__(self, kvargs): # 因为类似 qwen3.5 的linear 架构的模型,其 req_manager 会存储运行时使用的大量 linear state # 这可能会占用大量的显存,所以,req_manger 中保存的 mem_manger 是mem manager 初始化后再赋值 self.req_manager.mem_manager = self.mem_manager + hold_row = self.req_manager.req_to_token_indexs[self.req_manager.HOLD_REQUEST_ID] + hold_page = torch.arange( + self.mem_manager.HOLD_TOKEN_MEMINDEX, + self.mem_manager.HOLD_TOKEN_MEMINDEX + self.mem_manager.page_size, + dtype=hold_row.dtype, + device=hold_row.device, + ) + hold_row.view(-1, self.mem_manager.page_size).copy_(hold_page) self._check_mem_size() self._init_infer_layer() self._init_some_value() @@ -457,7 +465,7 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.mem_indexes, (0, padded_batch_size), mode="constant", - value=self.mem_manager.HOLD_TOKEN_MEMINDEX, + value=self.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % self.mem_manager.page_size), ) new_model_input.multimodal_params = new_model_input.multimodal_params + [ {"images": [], "audios": []} for _ in range(padded_batch_size) @@ -496,11 +504,17 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.max_kv_seq_len = max(padded_token_num, model_input.max_kv_seq_len) new_model_input.max_cache_len = max(0, model_input.max_cache_len) new_model_input.input_ids = F.pad(new_model_input.input_ids, (0, padded_token_num), mode="constant", value=1) - new_model_input.mem_indexes = F.pad( - new_model_input.mem_indexes, - (0, padded_token_num), - mode="constant", - value=self.mem_manager.HOLD_TOKEN_MEMINDEX, + hold_mem_indexes = ( + self.mem_manager.HOLD_TOKEN_MEMINDEX + + torch.arange( + padded_token_num, + dtype=new_model_input.mem_indexes.dtype, + device=new_model_input.mem_indexes.device, + ) + % self.mem_manager.page_size + ) + new_model_input.mem_indexes = torch.cat( + (new_model_input.mem_indexes, hold_mem_indexes), ) new_model_input.b_req_idx = F.pad( new_model_input.b_req_idx, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 5849cccf54..6c56bf9c13 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -256,6 +256,7 @@ def warmup(self, model): max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() + mem_indexes.fill_(model.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % model.mem_manager.page_size)) b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) @@ -317,6 +318,7 @@ def warmup_overlap(self, model): max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() + mem_indexes.fill_(model.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % model.mem_manager.page_size)) b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index bf6039a48f..6a40920cc9 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -195,6 +195,10 @@ def warmup(self, model): total_token_num = handle_token_num input_ids = torch.tensor([1 for _ in range(total_token_num)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() + mem_indexes.copy_( + model.mem_manager.HOLD_TOKEN_MEMINDEX + + torch.arange(total_token_num, dtype=torch.int32, device="cuda") % model.mem_manager.page_size + ) b_req_idx = torch.tensor([model.req_manager.HOLD_REQUEST_ID], dtype=torch.int32, device="cuda") b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) @@ -256,6 +260,10 @@ def warmup_overlap(self, model): total_token_num = handle_token_num input_ids = torch.tensor([1 for _ in range(total_token_num)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() + mem_indexes.copy_( + model.mem_manager.HOLD_TOKEN_MEMINDEX + + torch.arange(total_token_num, dtype=torch.int32, device="cuda") % model.mem_manager.page_size + ) b_req_idx = torch.tensor([model.req_manager.HOLD_REQUEST_ID], dtype=torch.int32, device="cuda") b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) diff --git a/lightllm/common/basemodel/triton_kernel/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index 3d04558273..80520a6591 100644 --- a/lightllm/common/basemodel/triton_kernel/fa3_utils.py +++ b/lightllm/common/basemodel/triton_kernel/fa3_utils.py @@ -18,6 +18,7 @@ def page_table_copy_kernel( page_table_stride_1, req_to_token_stride_0, req_to_token_stride_1, + PAGE_SIZE: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): cur_batch = tl.program_id(axis=0) @@ -27,10 +28,10 @@ def page_table_copy_kernel( offs = cur_block * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offs < max_seq_len_k - input_pos = cur_req_idx * req_to_token_stride_0 + offs * req_to_token_stride_1 + input_pos = cur_req_idx * req_to_token_stride_0 + offs * PAGE_SIZE * req_to_token_stride_1 output_pos = cur_batch * page_table_stride_0 + offs * page_table_stride_1 - mem_index = tl.load(req_to_token_indexs_ptr + input_pos, mask=mask) + mem_index = tl.load(req_to_token_indexs_ptr + input_pos, mask=mask) // PAGE_SIZE tl.store(page_table_ptr + output_pos, mem_index, mask=mask) @@ -38,6 +39,7 @@ def page_table_copy( page_table, # destination tensor [batch, seq] req_to_token_indexs, # source tensor [batch, seq] b_req_idx, # request index to copy from + page_size: int = 1, ): assert page_table.dim() == 2, "page_table should be 2D" assert req_to_token_indexs.dim() == 2, "req_to_token_indexs should be 2D" @@ -58,6 +60,7 @@ def page_table_copy( page_table_stride_1=page_table.stride(1), req_to_token_stride_0=req_to_token_indexs.stride(0), req_to_token_stride_1=req_to_token_indexs.stride(1), + PAGE_SIZE=page_size, BLOCK_SIZE=BLOCK_SIZE, ) diff --git a/lightllm/common/basemodel/triton_kernel/repack_kv_index.py b/lightllm/common/basemodel/triton_kernel/repack_kv_index.py index c218d15e08..2e7348293e 100644 --- a/lightllm/common/basemodel/triton_kernel/repack_kv_index.py +++ b/lightllm/common/basemodel/triton_kernel/repack_kv_index.py @@ -56,6 +56,56 @@ def repack_kv_index(kv_index, req_index, seq_len, start_loc, max_seq_len, out_kv return +@triton.jit +def _fwd_kernel_repack_page_kv_index( + kv_index, + req_index, + out_kv_index, + page_len, + start_loc, + kv_stride_h, + PAGE_SIZE: tl.constexpr, + SEQ_BLOCK: tl.constexpr, +): + cur_batch = tl.program_id(0) + start_page = tl.program_id(1) + cur_page_len = tl.load(page_len + cur_batch) + cur_req_idx = tl.load(req_index + cur_batch) + cur_start_loc = tl.load(start_loc + cur_batch) + + page_offsets = start_page * SEQ_BLOCK + tl.arange(0, SEQ_BLOCK) + token_offsets = page_offsets * PAGE_SIZE + token_index = tl.load( + kv_index + kv_stride_h * cur_req_idx + token_offsets, + mask=page_offsets < cur_page_len, + other=0, + ) + tl.store( + out_kv_index + cur_start_loc + page_offsets, + token_index // PAGE_SIZE, + mask=page_offsets < cur_page_len, + ) + + +@torch.no_grad() +def repack_page_kv_index(kv_index, req_index, page_len, start_loc, max_page_len, out_kv_index, page_size): + """Pack one physical page id per logical request page.""" + batch_size = req_index.shape[0] + block = 64 + _fwd_kernel_repack_page_kv_index[(batch_size, triton.cdiv(max_page_len, block))]( + kv_index, + req_index, + out_kv_index, + page_len, + start_loc, + kv_index.stride(0), + PAGE_SIZE=page_size, + SEQ_BLOCK=block, + num_warps=8, + num_stages=1, + ) + + def repack_kv_ref(req_to_token_indexs, b_req_idx, b_seq_len, b_start_loc, output): for b, sl, start in zip(b_req_idx, b_seq_len, b_start_loc): output[start : start + sl] = req_to_token_indexs[b][:sl] diff --git a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py index 9eb02b963c..d2b14de1b1 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py @@ -26,7 +26,7 @@ def get_cell_size(self): return self.head_num * self.head_dim * self.layer_num * torch._utils._element_size(self.dtype) def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): - self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device="cuda") + self.kv_buffer = torch.empty((layer_num, size + self.page_size, head_num, head_dim), dtype=dtype, device="cuda") def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: self.kv_move_buffer = torch.empty( diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 658d3e899c..3e3ab6ed17 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -36,6 +36,12 @@ def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False # profile the max total token num if the size is None self.profile_size(mem_fraction) + # A physical KV page must never straddle the allocator boundary. The + # unused remainder is intentionally dropped so every managed page has + # a stable ``mem_index // page_size`` id. + self.page_size = get_env_start_args().page_size + self.size = self.size // self.page_size * self.page_size + self.allocator = KvCacheAllocator(self.size) self._init_buffers( @@ -83,7 +89,9 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): # 分配,内部实际也没有管理,这个token是预留来对一些特殊的运行模式,如多dp下,overlap microbatch # 等模式下 padding 一些请求,使推理过程可以正常运行采用的,其索引值为size,存储在HOLD_TOKEN_MEMINDEX # 成员变量中,其与 req_manager 中的HOLD_REQUEST_ID具有类似的作用和意义。 - self.kv_buffer = torch.empty((layer_num, size + 1, 2 * head_num, head_dim), dtype=dtype, device="cuda") + self.kv_buffer = torch.empty( + (layer_num, size + self.page_size, 2 * head_num, head_dim), dtype=dtype, device="cuda" + ) def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir) @@ -175,17 +183,17 @@ def resize_mem(self, new_size): """ just for test code """ - size = new_size + size = new_size // self.page_size * self.page_size dtype = self.dtype head_num = self.head_num head_dim = self.head_dim layer_num = self.layer_num - self.size = new_size - self.allocator.resize(new_size) - self.HOLD_TOKEN_MEMINDEX = self.size + self.size = size + self.allocator.resize(size) self._free_buffers() self._init_buffers(size, dtype, head_num, head_dim, layer_num) + self.HOLD_TOKEN_MEMINDEX = self.size return def get_index_kv_buffer(self, index): diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 3de7de8f12..5b489f5452 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -67,6 +67,8 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager: MemoryMana # 的那个batch size 进行运行,所有 padding 的请求都会使用预留的这个请求管理 id 进行处理 # 这样让 DP 的实现更为简化一些。 self.req_list = _ReqLinkedList(max_request_num) + page_size = get_env_start_args().page_size + max_sequence_length = (max_sequence_length + page_size - 1) // page_size * page_size self.req_to_token_indexs = torch.zeros( (max_request_num + 1, max_sequence_length), dtype=torch.int32, device="cuda" ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index d72e63d724..45037882f3 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -444,6 +444,13 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: for hybrid linear-attention models, the second value selects the linear-attention backend (currently triton only); when omitted, it defaults to auto""", ) + parser.add_argument( + "--page_size", + type=int, + default=1, + help="""KV cache page size in tokens. Values greater than 1 make each request reserve + page-aligned contiguous KV slots and make paged attention/cache reuse operate on full pages.""", + ) parser.add_argument( "--vit_att_backend", type=str, diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 37fe837ad1..c0a4daca42 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -163,6 +163,28 @@ def _launch_subprocesses(args: StartArgs): f"{sorted(allowed_ep_decode_att_backends)}; flashinfer is not supported." ) + if args.page_size < 1: + raise ValueError(f"--page_size must be >= 1, got {args.page_size}") + + if args.page_size > 1: + unsupported_options = { + "MTP": args.mtp_mode is not None, + "PD split mode": args.run_mode in ("prefill", "decode"), + "CPU KV cache": args.enable_cpu_cache, + "DP prompt-cache fetch": args.enable_dp_prompt_cache_fetch, + "DP prefill balance": args.enable_dp_prefill_balance, + "diverse mode": args.diverse_mode, + "hybrid linear attention": is_linear_att_mixed_model(args.model_dir), + } + enabled_unsupported = [name for name, enabled in unsupported_options.items() if enabled] + if enabled_unsupported: + raise ValueError( + f"--page_size > 1 does not yet support {', '.join(enabled_unsupported)}; " + "use --page_size 1 for this configuration" + ) + if args.llm_kv_type != "None": + raise ValueError("--page_size > 1 currently supports only the unquantized LLM KV cache") + # mtp params check if args.mtp_mode is not None: if args.mtp_draft_model_dir is None: diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 5ccaa5a401..741868ee3e 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -166,6 +166,7 @@ class StartArgs: llm_decode_att_backend: List[str] = field( default_factory=lambda: ["auto"], metadata={"choices": ["auto", "triton", "fa3", "flashinfer"]} ) + page_size: int = field(default=1) vit_att_backend: List[str] = field( default_factory=lambda: ["auto"], metadata={"choices": ["auto", "triton", "fa3", "sdpa", "xformers"]} ) diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index c103a61473..030f75baab 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -20,7 +20,8 @@ def generate_time_id(self): class TreeNode: - def __init__(self): + def __init__(self, page_size: int = 1): + self.page_size = page_size self.children: Dict[int, TreeNode] = {} # 这里的键 为 token_id_key 的第一个元素 self.parent: TreeNode = None self.token_id_key: torch.Tensor = None @@ -34,14 +35,21 @@ def __init__(self): def get_compare_key(self): return (0 if self.ref_counter == 0 else 1, len(self.children), self.time_id) + def get_child_key(self, token_ids: torch.Tensor): + first_page = token_ids[: self.page_size] + if self.page_size == 1: + return first_page.item() + return tuple(first_page.tolist()) + def split_node(self, prefix_len): - split_parent_node = TreeNode() + assert prefix_len > 0 and prefix_len % self.page_size == 0 + split_parent_node = TreeNode(page_size=self.page_size) split_parent_node.parent = self.parent - split_parent_node.parent.children[self.token_id_key[0].item()] = split_parent_node + split_parent_node.parent.children[self.get_child_key(self.token_id_key)] = split_parent_node split_parent_node.token_id_key = self.token_id_key[0:prefix_len] split_parent_node.token_mem_index_value = self.token_mem_index_value[0:prefix_len] split_parent_node.children = {} - split_parent_node.children[self.token_id_key[prefix_len].item()] = self + split_parent_node.children[self.get_child_key(self.token_id_key[prefix_len:])] = self split_parent_node.ref_counter = self.ref_counter new_len = len(split_parent_node.token_mem_index_value) @@ -57,12 +65,12 @@ def split_node(self, prefix_len): return split_parent_node def add_and_return_new_child(self, token_id_key, token_mem_index_value): - child = TreeNode() + child = TreeNode(page_size=self.page_size) child.token_id_key = token_id_key child.token_mem_index_value = token_mem_index_value - first_token_key = child.token_id_key[0].item() - assert first_token_key not in self.children.keys() - self.children[first_token_key] = child + child_key = child.get_child_key(child.token_id_key) + assert child_key not in self.children.keys() + self.children[child_key] = child child.parent = self new_len = len(child.token_mem_index_value) @@ -71,7 +79,7 @@ def add_and_return_new_child(self, token_id_key, token_mem_index_value): return child def remove_child(self, child_node: "TreeNode"): - del self.children[child_node.token_id_key[0].item()] + del self.children[child_node.get_child_key(child_node.token_id_key)] child_node.parent = None return @@ -103,15 +111,18 @@ class RadixCache: unique_name 主要用于解决单机,多实列部署时的shm冲突 """ - def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None): + def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None, page_size: int = 1): from lightllm.common.kv_cache_mem_manager import MemoryManager self.total_token_num = total_token_num self.mem_manager: MemoryManager = mem_manager self._key_dtype = torch.int64 self._value_dtype = torch.int64 + if page_size < 1: + raise ValueError(f"page_size must be >= 1, got {page_size}") + self.page_size = page_size - self.root_node = TreeNode() + self.root_node = TreeNode(page_size=page_size) self.root_node.token_id_key = torch.zeros((0,), device="cpu", dtype=self._key_dtype) self.root_node.token_mem_index_value = torch.zeros((0,), device="cpu", dtype=self._value_dtype) self.root_node.ref_counter = 1 # 初始化为 1 保证永远不会被 evict 掉 @@ -131,8 +142,11 @@ def insert(self, key, value=None) -> Tuple[int, Optional[TreeNode]]: value = key assert len(key) == len(value) # and len(key) >= 1 - if len(key) == 0: + aligned_len = len(key) // self.page_size * self.page_size + if aligned_len == 0: return 0, None + key = key[:aligned_len] + value = value[:aligned_len] return self._insert_helper(self.root_node, key, value) def _insert_helper(self, node: TreeNode, key, value) -> Tuple[int, Optional[TreeNode]]: @@ -172,10 +186,12 @@ def _insert_helper_no_recursion( if node.is_leaf(): self.evict_tree_set.discard(node) - first_key_id = key[0].item() + first_key_id = node.get_child_key(key) if first_key_id in node.children.keys(): child: TreeNode = node.children[first_key_id] prefix_len = match(key, child.token_id_key) + prefix_len = prefix_len // self.page_size * self.page_size + assert prefix_len > 0 if prefix_len == len(key): if prefix_len == len(child.token_id_key): if child.is_leaf(): @@ -232,7 +248,10 @@ def _insert_helper_no_recursion( return 0, new_node def match_prefix(self, key, update_refs=False): - assert len(key) != 0 + aligned_len = len(key) // self.page_size * self.page_size + if aligned_len == 0: + return None, 0, None + key = key[:aligned_len] ans_value_list = [] tree_node = self._match_prefix_helper(self.root_node, key, ans_value_list, update_refs=update_refs) if tree_node != self.root_node: @@ -291,12 +310,14 @@ def _match_prefix_helper_no_recursion( if len(key) == 0: return node - first_key_id = key[0].item() + first_key_id = node.get_child_key(key) if first_key_id not in node.children.keys(): return node else: child = node.children[first_key_id] prefix_len = match(key, child.token_id_key) + prefix_len = prefix_len // self.page_size * self.page_size + assert prefix_len > 0 if prefix_len == len(child.token_id_key): ans_value_list.append(child.token_mem_index_value) return (child, key[prefix_len:]) @@ -374,7 +395,7 @@ def _try_merge(self, child_node: TreeNode) -> Optional[TreeNode]: child_node.time_id = max(parent_node.time_id, child_node.time_id) grandparent_node = parent_node.parent - key_in_grandparent = parent_node.token_id_key[0].item() + key_in_grandparent = parent_node.get_child_key(parent_node.token_id_key) grandparent_node.children[key_in_grandparent] = child_node child_node.parent = grandparent_node diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 67547f1996..d037d7356f 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -125,7 +125,8 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: def free_a_req_mem(self, free_token_index: List, req: "InferReq"): if self.radix_cache is None: - free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) + owned_kv_len = req.cur_kv_len if self.args.page_size == 1 else req.hold_kv_len + free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0:owned_kv_len]) else: if not self.is_linear_att_mixed_model: self._full_att_free_req(free_token_index=free_token_index, req=req) @@ -133,18 +134,23 @@ def free_a_req_mem(self, free_token_index: List, req: "InferReq"): self._linear_att_free_req(free_token_index=free_token_index, req=req) assert len(req.linear_att_len_to_big_page_id) == 0 req.cur_kv_len = 0 + req.hold_kv_len = 0 req.shm_req.shm_cur_kv_len = req.cur_kv_len return def _full_att_free_req(self, free_token_index: List, req: "InferReq"): + page_size = self.args.page_size + cache_kv_len = req.cur_kv_len // page_size * page_size input_token_ids = req.get_input_token_ids() - key = torch.tensor(input_token_ids[0 : req.cur_kv_len], dtype=torch.int64, device="cpu") + key = torch.tensor(input_token_ids[0:cache_kv_len], dtype=torch.int64, device="cpu") # .cpu() 是 流内阻塞操作 - value = self.req_manager.req_to_token_indexs[req.req_idx][: req.cur_kv_len].detach().cpu() + value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_kv_len].detach().cpu() prefix_len, _ = self.radix_cache.insert(key, value) old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][old_prefix_len:prefix_len]) + if page_size > 1: + free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][cache_kv_len : req.hold_kv_len]) if req.shared_kv_node is not None: assert req.shared_kv_node.node_prefix_total_len <= prefix_len self.radix_cache.dec_node_ref_counter(req.shared_kv_node) @@ -530,6 +536,9 @@ def __init__( self.shm_index = shm_index self.multimodal_params = multimodal_params self.vocab_size = vocab_size + # cur_kv_len is the logical length already written; hold_kv_len is the + # page-aligned physical capacity owned by this request. + self.hold_kv_len = 0 # 请求需要被暂停 self.wait_pause = False @@ -640,6 +649,7 @@ def _match_radix_cache(self): # 从 cpu 到 gpu 是流内阻塞操作 g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 + self.hold_kv_len = self.cur_kv_len self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 self.shm_req.shm_cur_kv_len = self.cur_kv_len @@ -909,15 +919,27 @@ def prefill_need_token_num(self, is_chuncked_prefill: bool): else: input_token_ids = self.get_input_token_ids() - seq_len = len(input_token_ids) - input_token_len = seq_len - self.cur_kv_len - return input_token_len + return len(input_token_ids) - self.cur_kv_len + + def prefill_kv_alloc_need(self, is_chuncked_prefill: bool) -> int: + if is_chuncked_prefill: + target_kv_len = len(self.get_chuncked_input_token_ids()) + else: + target_kv_len = len(self.get_input_token_ids()) + return self._kv_cache_alloc_need(target_kv_len) def decode_need_token_num(self) -> int: raise NotImplementedError("error") def _normal_decode_need_token_num(self) -> int: - return 1 + return self._kv_cache_alloc_need(self.cur_kv_len + 1) + + def _kv_cache_alloc_need(self, target_kv_len: int) -> int: + page_size = self.args.page_size + if page_size == 1: + return target_kv_len - self.cur_kv_len + target_hold_len = (target_kv_len + page_size - 1) // page_size * page_size + return target_hold_len - self.hold_kv_len def _mtp_decode_need_token_num(self) -> int: return (1 + self.mtp_step) * 2 diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 28f2abf74b..37a4551f24 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -176,6 +176,7 @@ def init_model(self, kvargs): total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, mem_manager=self.model.mem_manager, + page_size=self.args.page_size, ) if "prompt_cache_kv_buffer" in model_cfg: @@ -733,12 +734,13 @@ def _get_classed_reqs( continue token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill) + alloc_token_num = req_obj.prefill_kv_alloc_need(is_chuncked_prefill=not self.disable_chunked_prefill) if prefill_tokens + token_num > self.batch_max_tokens: continue - if token_num <= can_alloc_token_num: + if alloc_token_num <= can_alloc_token_num: prefill_tokens += token_num prefill_reqs.append(req_obj) - can_alloc_token_num -= token_num + can_alloc_token_num -= alloc_token_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True @@ -951,14 +953,20 @@ def preload_prompt_cache_kv_buffer(self, model_cfg): ) prompt_cache_kv_buffer = torch.load(prompt_cache_kv_buffer_path, weights_only=True, map_location="cpu") intact_kv_len = len(model_cfg["prompt_cache_token_ids"]) - intact_kv_index = self.radix_cache.mem_manager.alloc(intact_kv_len) - self.radix_cache.mem_manager.load_index_kv_buffer(intact_kv_index, prompt_cache_kv_buffer) + page_size = self.args.page_size + cache_kv_len = intact_kv_len // page_size * page_size + intact_hold_len = (intact_kv_len + page_size - 1) // page_size * page_size + intact_kv_index = self.radix_cache.mem_manager.alloc(intact_hold_len) + self.radix_cache.mem_manager.load_index_kv_buffer(intact_kv_index[:intact_kv_len], prompt_cache_kv_buffer) self.radix_cache.insert( - torch.tensor(model_cfg["prompt_cache_token_ids"], dtype=torch.int64, device="cpu"), - intact_kv_index, + torch.tensor(model_cfg["prompt_cache_token_ids"][:cache_kv_len], dtype=torch.int64, device="cpu"), + intact_kv_index[:cache_kv_len], ) + if intact_hold_len > cache_kv_len: + self.radix_cache.mem_manager.free(intact_kv_index[cache_kv_len:intact_hold_len]) self.radix_cache.match_prefix( - torch.tensor(model_cfg["prompt_cache_token_ids"], dtype=torch.int64, device="cpu"), update_refs=True + torch.tensor(model_cfg["prompt_cache_token_ids"][:cache_kv_len], dtype=torch.int64, device="cpu"), + update_refs=True, ) def init_rank_infos(self): diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 22731439c4..6833dda0a4 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -7,6 +7,42 @@ INT64_MAX = torch.iinfo(torch.int64).max +def _alloc_kv_mem_indexes(reqs: List[InferReq], target_kv_lens: List[int], token_num: int) -> torch.Tensor: + """Allocate KV slots for this step; page ownership stays in the caller.""" + req_manager = g_infer_context.req_manager + mem_manager = req_manager.mem_manager + page_size = g_infer_context.args.page_size + + if page_size == 1: + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_num) + return mem_manager.alloc(token_num) + + assert len(reqs) == len(target_kv_lens) + alloc_need = sum(req._kv_cache_alloc_need(target_len) for req, target_len in zip(reqs, target_kv_lens)) + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(alloc_need) + + for req, target_len in zip(reqs, target_kv_lens): + assert req.cur_kv_len <= target_len and req.cur_kv_len <= req.hold_kv_len + assert req.hold_kv_len % page_size == 0 + new_hold_kv_len = (target_len + page_size - 1) // page_size * page_size + if new_hold_kv_len > req.hold_kv_len: + new_page_indexes = mem_manager.alloc(new_hold_kv_len - req.hold_kv_len) + req_manager.req_to_token_indexs[req.req_idx, req.hold_kv_len : new_hold_kv_len] = new_page_indexes + req.hold_kv_len = new_hold_kv_len + + # The request table is the source of truth. Return only the logical slots + # consumed by this step; the reserved tail remains owned by the request. + step_indexes = [ + req_manager.req_to_token_indexs[req.req_idx, req.cur_kv_len : target_len] + for req, target_len in zip(reqs, target_kv_lens) + ] + if step_indexes: + return torch.cat(step_indexes) + return torch.empty((0,), dtype=torch.int32, device=req_manager.req_to_token_indexs.device) + + def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> Tuple[ModelInput, List[InferReq]]: run_reqs = [] total_token_num = 0 @@ -66,9 +102,8 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len # dynamic prompt cache 准备 token - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(input_ids.shape[0]) - mem_indexes = g_infer_context.req_manager.mem_manager.alloc(input_ids.shape[0]) + target_kv_lens = [int(seq_len) for seq_len in b_seq_len] + mem_indexes = _alloc_kv_mem_indexes(run_reqs, target_kv_lens, input_ids.shape[0]) model_input = ModelInput( batch_size=b_seq_len.shape[0], @@ -77,7 +112,8 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> max_kv_seq_len=max_kv_seq_len, max_cache_len=max_cache_len, input_ids=input_ids, - mem_indexes_cpu=mem_indexes, + mem_indexes=mem_indexes if mem_indexes.is_cuda else None, + mem_indexes_cpu=mem_indexes if not mem_indexes.is_cuda else None, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -141,9 +177,8 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In ) # dynamic prompt cache 准备 token - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(b_seq_len.shape[0]) - mem_indexes = g_infer_context.req_manager.mem_manager.alloc(b_seq_len.shape[0]) + target_kv_lens = [req.cur_kv_len + 1 for req in req_objs] + mem_indexes = _alloc_kv_mem_indexes(req_objs, target_kv_lens, b_seq_len.shape[0]) model_input = ModelInput( batch_size=b_seq_len.shape[0], @@ -151,7 +186,8 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In max_q_seq_len=max_q_seq_len, max_kv_seq_len=max_kv_seq_len, input_ids=None, - mem_indexes_cpu=mem_indexes, + mem_indexes=mem_indexes if mem_indexes.is_cuda else None, + mem_indexes_cpu=mem_indexes if not mem_indexes.is_cuda else None, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 9af7afd1b4..f5d33fa35a 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -29,6 +29,11 @@ def __init__(self, args: StartArgs, router, dp_index, dp_size_in_node) -> None: self.waiting_req_list: List[Req] = [] # List of queued requests self.router_token_ratio = args.router_token_ratio # ratio to determine whether the router is busy + def add_kv_page_reservation(self, token_num: int, req_num: int) -> int: + """Conservatively include each running request's incomplete tail page.""" + page_size = self.args.page_size + return token_num + req_num * (page_size - 1) + def free_aborted_req_cpu_cache_pages(self, req: Req): if self.args.enable_cpu_cache: self.router.cpu_cache_client.lock.acquire_sleep1ms() diff --git a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py index 962c9c31cb..42476333e0 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py @@ -45,7 +45,8 @@ def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy, new # prefill token 计算, 因为对beam的prefill计算过程是共享的,所以只计算一个请求对应的token数量 new_batch_first_router_need_tokens += req.get_first_router_need_tokens() - ok_token_num = need_max_token_num < self.max_total_tokens + estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) + ok_token_num = estimated_need_token_num < self.max_total_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size @@ -56,9 +57,9 @@ def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy, new ) if ok_token_num and ok_req_num and ok_prefill: - self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) + self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( - need_max_token_num / self.max_total_tokens, + estimated_need_token_num / self.max_total_tokens, self.dp_index, ) return True, new_batch_first_router_need_tokens @@ -144,7 +145,5 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): assert cur_input_len - req.input_len >= 0 cumsum_len += cur_input_len - req.input_len # 减去共享的部分 need_max_token_num = max(need_max_token_num, cumsum_len + index * cur_ouput_len) - return ( - need_max_token_num, - need_max_token_num / self.max_total_tokens, - ) + estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) + return (estimated_need_token_num, estimated_need_token_num / self.max_total_tokens) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py index f8fe510989..799af7f97a 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py @@ -33,7 +33,8 @@ def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens size_array = np.arange(1, len(self.cache_len_list) + 1, 1) need_max_token_num = (left_out_len_array * size_array + cum_run_len_array).max() - ok_token_num = need_max_token_num < self.max_total_tokens + estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) + ok_token_num = estimated_need_token_num < self.max_total_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size @@ -45,9 +46,9 @@ def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens ) if ok_token_num and ok_req_num and ok_prefill: - self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) + self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( - need_max_token_num / self.max_total_tokens, + estimated_need_token_num / self.max_total_tokens, self.dp_index, ) return True, new_batch_first_router_need_tokens @@ -109,7 +110,5 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): else: need_max_token_num = 0 - return ( - need_max_token_num, - need_max_token_num / self.max_total_tokens, - ) + estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) + return (estimated_need_token_num, estimated_need_token_num / self.max_total_tokens) diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 056513bdde..9fd44487c6 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -126,7 +126,7 @@ def test_padded_prefill_adds_non_decode_request_marker(): multimodal_params=[{"images": [], "audios": []}], ) model = SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=-1), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=-1, page_size=1), req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1), ) @@ -161,7 +161,7 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): mtp_draft_input_hiddens=torch.empty((0, 4), dtype=torch.float32), ) model = SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), req_manager=SimpleNamespace(HOLD_REQUEST_ID=88), ) @@ -199,7 +199,7 @@ def test_padded_decode_builds_internal_request_from_empty_token_tensor(): multimodal_params=[], ) model = SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), req_manager=SimpleNamespace(HOLD_REQUEST_ID=88), ) diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 6f9477e294..7a5f565aef 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -118,7 +118,7 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode) model.tp_world_size_ = tp_world_size - model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=99) + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) pad_batch_sizes = [] diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py index b286998e04..ec5dad5d34 100644 --- a/unit_tests/common/basemodel/test_overlap_utils.py +++ b/unit_tests/common/basemodel/test_overlap_utils.py @@ -186,7 +186,7 @@ def test_overlap_decode_cuda_pads_empty_side_and_unpads_outputs(monkeypatch): model.tp_world_size_ = 2 model.graph = None model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) - model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=77) + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=77, page_size=1) infer_batch_sizes = [] def fake_create_inferstate(model_input, microbatch_index): diff --git a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py index c1ef686f8e..55e7d0724a 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py @@ -4,7 +4,19 @@ if not torch.cuda.is_available(): pytest.skip("requires CUDA", allow_module_level=True) -from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params +from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params, page_table_copy + + +def test_page_table_copy_uses_page_bases(): + req_to_token_indexs = torch.tensor([list(range(40, 52)), list(range(80, 92))], dtype=torch.int32, device="cuda") + page_table = torch.empty((2, 3), dtype=torch.int32, device="cuda") + page_table_copy( + page_table=page_table, + req_to_token_indexs=req_to_token_indexs, + b_req_idx=torch.tensor([1, 0], dtype=torch.int32, device="cuda"), + page_size=4, + ) + assert page_table.cpu().tolist() == [[20, 21, 22], [10, 11, 12]] def _reference_dynamic_spec_fa3_decode_params(b_req_idx, b_seq_len, b_mark_mtp_shared_group, hold_req_id): diff --git a/unit_tests/common/basemodel/triton_kernel/test_repack_kv_index.py b/unit_tests/common/basemodel/triton_kernel/test_repack_kv_index.py index b5184d3caa..bb32a9ede6 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_repack_kv_index.py +++ b/unit_tests/common/basemodel/triton_kernel/test_repack_kv_index.py @@ -1,7 +1,7 @@ import torch import pytest from lightllm.utils.log_utils import init_logger -from lightllm.common.basemodel.triton_kernel.repack_kv_index import repack_kv_index +from lightllm.common.basemodel.triton_kernel.repack_kv_index import repack_kv_index, repack_page_kv_index logger = init_logger(__name__) @@ -41,3 +41,32 @@ def repack_kv_ref(req_to_token_indexs, b_req_idx, b_seq_len, b_start_loc, output repack_kv_ref(req_to_token_indexs, b_req_idx, b_seq_len, b_start_loc, ref) repack_kv_index(req_to_token_indexs, b_req_idx, b_seq_len, b_start_loc, MAX_SEQ_LEN, output) assert torch.allclose(output.float(), ref.float()) + + +def test_repack_page_kv_index(): + page_size = 4 + req_to_token_indexs = torch.tensor( + [ + list(range(40, 52)), + list(range(80, 92)), + list(range(120, 132)), + ], + dtype=torch.int32, + device="cuda", + ) + req_indexes = torch.tensor([2, 0, 1], dtype=torch.int32, device="cuda") + page_lens = torch.tensor([2, 1, 3], dtype=torch.int32, device="cuda") + starts = torch.tensor([0, 2, 3], dtype=torch.int32, device="cuda") + output = torch.empty((6,), dtype=torch.int32, device="cuda") + + repack_page_kv_index( + req_to_token_indexs, + req_indexes, + page_lens, + starts, + max_page_len=3, + out_kv_index=output, + page_size=page_size, + ) + + assert output.cpu().tolist() == [30, 31, 10, 20, 21, 22] diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py new file mode 100644 index 0000000000..b53cd78318 --- /dev/null +++ b/unit_tests/common/test_req_manager_page.py @@ -0,0 +1,87 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.server.router.model_infer.mode_backend import generic_pre_process + + +class _FakeMemManager: + page_size = 4 + + def __init__(self): + self.next_index = 0 + self.alloc_sizes = [] + + def alloc(self, size): + self.alloc_sizes.append(size) + result = torch.arange(self.next_index, self.next_index + size, dtype=torch.int32) + self.next_index += size + return result + + +def _make_context(monkeypatch): + mem_manager = _FakeMemManager() + req_manager = SimpleNamespace( + mem_manager=mem_manager, + req_to_token_indexs=torch.full((2, 16), -1, dtype=torch.int32), + ) + context = SimpleNamespace( + args=SimpleNamespace(page_size=4), + req_manager=req_manager, + radix_cache=None, + ) + monkeypatch.setattr(generic_pre_process, "g_infer_context", context) + return context + + +def _make_req(req_idx): + req = SimpleNamespace(req_idx=req_idx, cur_kv_len=0, hold_kv_len=0) + req._kv_cache_alloc_need = lambda target_len: (target_len + 3) // 4 * 4 - req.hold_kv_len + return req + + +def test_request_reuses_reserved_page_tail_before_allocating_next_page(monkeypatch): + context = _make_context(monkeypatch) + req = _make_req(0) + + indexes = generic_pre_process._alloc_kv_mem_indexes([req], [3], 3) + assert indexes.tolist() == [0, 1, 2] + assert req.hold_kv_len == 4 + assert context.req_manager.mem_manager.alloc_sizes == [4] + assert context.req_manager.req_to_token_indexs[0, :4].tolist() == [0, 1, 2, 3] + + req.cur_kv_len = 3 + indexes = generic_pre_process._alloc_kv_mem_indexes([req], [4], 1) + assert indexes.tolist() == [3] + assert context.req_manager.mem_manager.alloc_sizes == [4] + + req.cur_kv_len = 4 + indexes = generic_pre_process._alloc_kv_mem_indexes([req], [6], 2) + assert indexes.tolist() == [4, 5] + assert req.hold_kv_len == 8 + assert context.req_manager.mem_manager.alloc_sizes == [4, 4] + assert context.req_manager.req_to_token_indexs[0, :8].tolist() == list(range(8)) + + +def test_step_indexes_follow_request_order(monkeypatch): + _make_context(monkeypatch) + req0 = _make_req(0) + req1 = _make_req(1) + + indexes = generic_pre_process._alloc_kv_mem_indexes([req0, req1], [2, 3], 5) + + assert indexes.tolist() == [0, 1, 4, 5, 6] + assert req0.hold_kv_len == req1.hold_kv_len == 4 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_step_indexes_are_assembled_from_gpu_request_table(monkeypatch): + context = _make_context(monkeypatch) + context.req_manager.req_to_token_indexs = context.req_manager.req_to_token_indexs.cuda() + req = _make_req(0) + + indexes = generic_pre_process._alloc_kv_mem_indexes([req], [3], 3) + + assert indexes.is_cuda + assert indexes.cpu().tolist() == [0, 1, 2] diff --git a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py index dfeda0b6f7..594005236f 100644 --- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py +++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py @@ -257,5 +257,34 @@ def test_case10(): assert tree.root_node.ref_counter == 1 +def test_page_aligned_insert_and_match(): + tree = RadixCache("paged_radix_test", 100, 99, page_size=4) + values = torch.arange(100, 110, dtype=torch.int64) + + prefix_len, _ = tree.insert(torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]), values) + assert prefix_len == 0 + assert tree.get_tree_total_tokens_num() == 8 + + # A mismatch inside a page cannot produce a partial-page cache hit. + node, matched_len, matched_values = tree.match_prefix(torch.tensor([1, 2, 0, 4, 5, 6, 7, 8])) + assert node is None + assert matched_len == 0 + assert matched_values is None + + node, matched_len, matched_values = tree.match_prefix(torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 11])) + assert node is not None + assert matched_len == 8 + assert matched_values.tolist() == list(range(100, 108)) + + # The second sequence shares exactly one page and then branches by its + # complete second-page key. + prefix_len, _ = tree.insert( + torch.tensor([1, 2, 3, 4, 5, 6, 0, 8]), + torch.arange(200, 208, dtype=torch.int64), + ) + assert prefix_len == 4 + assert tree.get_tree_total_tokens_num() == 12 + + if __name__ == "__main__": pytest.main() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 62296634a9..1e789d4420 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -10,6 +10,7 @@ def _patch_empty_input_context(monkeypatch): alloc=lambda size: torch.empty((size,), dtype=torch.int32), ) infer_context = SimpleNamespace( + args=SimpleNamespace(page_size=1), req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager), radix_cache=None, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 2e24262e7e..c046c63c51 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -112,7 +112,7 @@ def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): draft_models=[draft_model], model=SimpleNamespace( req_manager=SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), ) ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), @@ -165,7 +165,7 @@ def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): draft_models=[draft_model], model=SimpleNamespace( req_manager=SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), ) ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), @@ -207,7 +207,7 @@ def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): draft_models=[draft_model], model=SimpleNamespace( req_manager=SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), ) ), _gen_argmax_token_ids_and_prob=lambda output: ( @@ -314,7 +314,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): draft_models=[draft_model], model=SimpleNamespace( req_manager=SimpleNamespace( - mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1), ) ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), From 1be6cb02a385ba8dee52b851953bdf1f06b9b4ae Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 02:55:15 +0000 Subject: [PATCH 02/15] refactor: preallocate KV cache during scheduling --- docs/kv_cache_page_size.md | 14 +- lightllm/common/basemodel/basemodel.py | 152 +++++++++++------- lightllm/common/basemodel/batch_objs.py | 10 +- .../triton_kernel/copy_kv_index_to_req.py | 95 ++++++++++- .../triton_kernel/dynamic_mtp_utils.py | 3 +- lightllm/server/core/objs/req.py | 7 +- .../server/router/model_infer/infer_batch.py | 12 +- .../model_infer/mode_backend/base_backend.py | 19 +++ .../mode_backend/chunked_prefill/impl.py | 13 +- .../mode_backend/dp_backend/impl.py | 25 +-- .../mode_backend/generic_pre_process.py | 50 +----- .../model_infer/mtp_speculative/engine.py | 21 ++- .../common/basemodel/test_model_output.py | 2 +- .../common/basemodel/test_overlap_utils.py | 5 +- .../test_select_kv_index_from_req.py | 77 +++++++++ unit_tests/common/test_req_manager_page.py | 105 +++++++++--- unit_tests/server/core/objs/test_req.py | 13 ++ .../mode_backend/test_generic_pre_process.py | 14 +- .../mtp_speculative/test_planner.py | 1 + 19 files changed, 452 insertions(+), 186 deletions(-) create mode 100644 unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py diff --git a/docs/kv_cache_page_size.md b/docs/kv_cache_page_size.md index f01e869838..9ff7089b59 100644 --- a/docs/kv_cache_page_size.md +++ b/docs/kv_cache_page_size.md @@ -11,18 +11,20 @@ KV 存储和请求表,但资源的申请、缓存和释放以完整物理页 `cur_kv_len <= hold_kv_len` 且 `hold_kv_len % page_size == 0`。 2. 一个物理页内的 KV 槽连续,页首索引可被 `page_size` 整除。请求表保存已持有页的全部 token 索引, 包括尚未使用的尾部槽位。 -3. `ModelInput.mem_indexes` / `InferState.mem_index` 只包含本轮真实参与计算的 token,不能包含预留尾部。 +3. `req_to_token_indexs` 保存请求拥有的完整页;模型创建 `InferState` 时再从中选出本轮真实参与计算的 + token,`InferState.mem_index` 不包含预留尾部。 4. Radix Cache 只插入、拆分、命中和淘汰完整页;不足一页的请求尾部在请求结束或暂停时整页回收。 -5. `page_size=1` 继续走原有批量申请和 token 级页表路径,不改变默认行为。 +5. `page_size=1` 与多 token 页使用相同的调度期预留和模型执行期索引选择路径。 -页容量计算、整页申请和本轮索引组装由 Prefill/Decode 输入构造层负责;`ReqManager` 只保留请求 ID、 -请求表及其原有生命周期职责,不提供 page allocator 接口。 +页容量计算和整页申请由调度层在请求获准进入 Prefill/Decode batch 时完成;输入构造层只构造本轮输入, +真实推理索引由模型执行层从请求表选取。 +`ReqManager` 只保留请求 ID、请求表及其原有生命周期职责,不提供 page allocator 接口。 ## 生命周期 -- Prefill:按每个请求的目标 KV 长度将容量向上对齐。新页一次申请并写入请求表,本轮只返回 +- Prefill:请求通过调度后,按目标 KV 长度将容量向上对齐,新页一次申请并完整写入请求表。模型执行时再选择 `[cur_kv_len, target_kv_len)` 对应的真实索引。 -- Decode:若尾页仍有预留槽位,不再访问 allocator;跨页时申请一个新页。 +- Decode:请求通过调度后,若尾页仍有预留槽位则不再访问 allocator;容量不足时先补齐完整页再执行。 - Prefix Cache 命中:命中长度天然是整页,`cur_kv_len` 与 `hold_kv_len` 同时初始化为命中长度。 - 完成/暂停:完整逻辑页可进入 Radix Cache;重复页和未完成尾页按物理页展开后释放。 - Attention:FA3/FlashInfer 页表每项由物理 token 页首索引除以 `page_size` 得到,KV buffer 视为 diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 3914cc6b2e..2c2cda2bca 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -17,7 +17,11 @@ from lightllm.common.req_manager import ReqManager from lightllm.common.infer_utils import init_req_to_token_indexes from lightllm.common.build_utils import repair_config -from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import copy_kv_index_to_req +from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import ( + copy_kv_index_to_req, + select_kv_index_from_req, + select_kv_index_from_req_prefill, +) from lightllm.common.basemodel.layer_infer.cache_tensor_manager import g_cache_manager from lightllm.common.basemodel.cuda_graph import CudaGraph from lightllm.common.basemodel.prefill_cuda_graph import PrefillCudaGraph @@ -379,6 +383,7 @@ def _init_hidden_collector(self): @torch.no_grad() def forward(self, model_input: ModelInput): model_input.to_cuda() + self._select_page_mem_indexes(model_input) assert model_input.mem_indexes.is_cuda if model_input.is_prefill: @@ -386,6 +391,27 @@ def forward(self, model_input: ModelInput): else: return self._decode(model_input) + def _select_page_mem_indexes(self, model_input: ModelInput): + if not model_input.mem_indexes_from_req_table or model_input.mem_indexes is not None: + return + if model_input.is_prefill: + model_input.mem_indexes = select_kv_index_from_req_prefill( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + b_start_loc=model_input.b_prefill_start_loc, + max_q_seq_len=model_input.max_q_seq_len, + token_num=model_input.input_ids.shape[0], + ) + else: + model_input.mem_indexes = select_kv_index_from_req( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + ) + return + def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() infer_state.hidden_collector = self.hidden_collector_prototype.new_instance() @@ -603,15 +629,16 @@ def _prefill( ) infer_state = self._create_inferstate(model_input) - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state.b_req_idx, - b_seq_len=infer_state.b_seq_len, - b_ready_cache_len=infer_state.b_ready_cache_len, - b_start_loc=model_input.b_prefill_start_loc, - alloc_mem_index=infer_state.mem_index, - max_q_seq_len=infer_state.max_q_seq_len, - ) + if not model_input.mem_indexes_from_req_table: + init_req_to_token_indexes( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=infer_state.b_req_idx, + b_seq_len=infer_state.b_seq_len, + b_ready_cache_len=infer_state.b_ready_cache_len, + b_start_loc=model_input.b_prefill_start_loc, + alloc_mem_index=infer_state.mem_index, + max_q_seq_len=infer_state.max_q_seq_len, + ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -671,12 +698,13 @@ def _decode( # attention backend 会根据该标记准备 CUDA Graph capture 专用状态, # 因此必须在 init_att_state 之前完成赋值。 infer_state.is_cuda_graph = need_capture - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state.b_req_idx, - infer_state.b_seq_len, - infer_state.mem_index, - ) + if not model_input.mem_indexes_from_req_table: + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state.b_req_idx, + infer_state.b_seq_len, + infer_state.mem_index, + ) infer_state.init_some_extra_state(self) infer_state.init_att_state() @@ -795,6 +823,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod for model_input in (model_input0, model_input1): model_input.to_cuda() + self._select_page_mem_indexes(model_input) if self.args.enable_prefill_decode_mixed and model_input.input_ids.shape[0] > 0: gather_token_prefill_decode_mixed( input_ids=model_input.input_ids, @@ -832,28 +861,30 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input ) infer_state0 = self._create_inferstate(model_input0, 0) - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state0.b_req_idx, - b_seq_len=infer_state0.b_seq_len, - b_ready_cache_len=infer_state0.b_ready_cache_len, - b_start_loc=model_input0.b_prefill_start_loc, - alloc_mem_index=infer_state0.mem_index, - max_q_seq_len=infer_state0.max_q_seq_len, - ) + if not model_input0.mem_indexes_from_req_table: + init_req_to_token_indexes( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=infer_state0.b_req_idx, + b_seq_len=infer_state0.b_seq_len, + b_ready_cache_len=infer_state0.b_ready_cache_len, + b_start_loc=model_input0.b_prefill_start_loc, + alloc_mem_index=infer_state0.mem_index, + max_q_seq_len=infer_state0.max_q_seq_len, + ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(model_input1, 1) - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state1.b_req_idx, - b_seq_len=infer_state1.b_seq_len, - b_ready_cache_len=infer_state1.b_ready_cache_len, - b_start_loc=model_input1.b_prefill_start_loc, - alloc_mem_index=infer_state1.mem_index, - max_q_seq_len=infer_state1.max_q_seq_len, - ) + if not model_input1.mem_indexes_from_req_table: + init_req_to_token_indexes( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=infer_state1.b_req_idx, + b_seq_len=infer_state1.b_seq_len, + b_ready_cache_len=infer_state1.b_ready_cache_len, + b_start_loc=model_input1.b_prefill_start_loc, + alloc_mem_index=infer_state1.mem_index, + max_q_seq_len=infer_state1.max_q_seq_len, + ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -885,6 +916,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode for model_input in (model_input0, model_input1): model_input.to_cuda() + self._select_page_mem_indexes(model_input) if model_input.input_ids is None: if model_input.batch_size > 0: model_input.input_ids = gather_token( @@ -918,23 +950,25 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 padded_model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) infer_state0 = self._create_inferstate(padded_model_input0, 0) infer_state0.is_cuda_graph = need_capture - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state0.b_req_idx, - infer_state0.b_seq_len, - infer_state0.mem_index, - ) + if not padded_model_input0.mem_indexes_from_req_table: + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state0.b_req_idx, + infer_state0.b_seq_len, + infer_state0.mem_index, + ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state1.is_cuda_graph = need_capture - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state1.b_req_idx, - infer_state1.b_seq_len, - infer_state1.mem_index, - ) + if not padded_model_input1.mem_indexes_from_req_table: + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state1.b_req_idx, + infer_state1.b_seq_len, + infer_state1.mem_index, + ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -956,22 +990,24 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) infer_state0 = self._create_inferstate(model_input0, 0) - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state0.b_req_idx, - infer_state0.b_seq_len, - infer_state0.mem_index, - ) + if not model_input0.mem_indexes_from_req_table: + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state0.b_req_idx, + infer_state0.b_seq_len, + infer_state0.mem_index, + ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(model_input1, 1) - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state1.b_req_idx, - infer_state1.b_seq_len, - infer_state1.mem_index, - ) + if not model_input1.mem_indexes_from_req_table: + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state1.b_req_idx, + infer_state1.b_seq_len, + infer_state1.mem_index, + ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index ae645d4b7b..f783604ba9 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -33,6 +33,9 @@ class ModelInput: # radix node;该 id 只用于重建 diverse attention 的 b_mark_shared_group。 b_shared_radix_node_id: torch.Tensor = None mem_indexes: torch.Tensor = None + # The scheduler reserves KV capacity directly in req_to_token_indexs; + # BaseModel resolves the logical indexes when execution starts. + mem_indexes_from_req_table: bool = False is_prefill: bool = False b_ready_cache_len: torch.Tensor = None # Request/row-aligned MRoPE position offset. It is decode-only; prefill @@ -59,7 +62,7 @@ def to_cuda(self): self.check_input() # Prefill 和 decode 都必须提供的公共张量。 - if self.mem_indexes is None: + if self.mem_indexes is None and self.mem_indexes_cpu is not None: self.mem_indexes = self.mem_indexes_cpu.cuda(non_blocking=True) self.b_req_idx = self.b_req_idx.cuda(non_blocking=True) self.b_seq_len = self.b_seq_len.cuda(non_blocking=True) @@ -94,7 +97,7 @@ def check_input(self): assert self.b_mtp_index is not None assert self.b_seq_len is not None assert self.multimodal_params is not None - assert self.mem_indexes is not None or self.mem_indexes_cpu is not None + assert self.mem_indexes_from_req_table or self.mem_indexes is not None or self.mem_indexes_cpu is not None assert self.b_req_idx.shape == (self.batch_size,) assert self.b_mtp_index.shape == self.b_req_idx.shape @@ -124,7 +127,8 @@ def check_input(self): assert self.b_shared_radix_node_id.shape == self.b_req_idx.shape mem_indexes = self.mem_indexes if self.mem_indexes is not None else self.mem_indexes_cpu - assert mem_indexes.ndim == 1 + if mem_indexes is not None: + assert mem_indexes.ndim == 1 @dataclass diff --git a/lightllm/common/basemodel/triton_kernel/copy_kv_index_to_req.py b/lightllm/common/basemodel/triton_kernel/copy_kv_index_to_req.py index 8dc7a487e5..efe563d764 100644 --- a/lightllm/common/basemodel/triton_kernel/copy_kv_index_to_req.py +++ b/lightllm/common/basemodel/triton_kernel/copy_kv_index_to_req.py @@ -37,6 +37,39 @@ def copy_kv_index_to_req(req_to_token_indexs, b_req_idx, b_seq_len, memindex): return +@triton.jit +def _fwd_kernel_select_kv_index_from_req( + req_to_token_indexs, + b_req_idx, + b_seq_len, + out_memindex, + stride_req_to_token_b, + stride_req_to_token_s, +): + cur_index = tl.program_id(0) + cur_req_idx = tl.load(b_req_idx + cur_index) + cur_seq_len = tl.load(b_seq_len + cur_index) + src_offset = req_to_token_indexs + cur_req_idx * stride_req_to_token_b + (cur_seq_len - 1) * stride_req_to_token_s + tl.store(out_memindex + cur_index, tl.load(src_offset)) + + +@torch.no_grad() +def select_kv_index_from_req(req_to_token_indexs, b_req_idx, b_seq_len): + """Select the one logical KV slot consumed by each decode row.""" + out_memindex = torch.empty_like(b_req_idx) + _fwd_kernel_select_kv_index_from_req[(b_req_idx.shape[0],)]( + req_to_token_indexs, + b_req_idx, + b_seq_len, + out_memindex, + req_to_token_indexs.stride(0), + req_to_token_indexs.stride(1), + num_warps=1, + num_stages=1, + ) + return out_memindex + + @triton.jit def _fwd_kernel_copy_kv_index_to_req_prefill( req_to_token_indexs, @@ -108,4 +141,64 @@ def copy_kv_index_to_req_prefill( num_warps=num_warps, num_stages=1, ) - return + + +@triton.jit +def _fwd_kernel_select_kv_index_from_req_prefill( + req_to_token_indexs, + b_req_idx, + b_seq_len, + b_ready_cache_len, + b_start_loc, + out_memindex, + stride_req_to_token_b, + stride_req_to_token_s, + BLOCK: tl.constexpr, +): + block_index = tl.program_id(0) + batch_index = tl.program_id(1) + cur_req_idx = tl.load(b_req_idx + batch_index) + cur_seq_len = tl.load(b_seq_len + batch_index) + cur_ready_cache_len = tl.load(b_ready_cache_len + batch_index) + cur_start_loc = tl.load(b_start_loc + batch_index) + copy_len = cur_seq_len - cur_ready_cache_len + + block_range = block_index * BLOCK + tl.arange(0, BLOCK) + block_mask = block_range < copy_len + src_offset = ( + req_to_token_indexs + + cur_req_idx * stride_req_to_token_b + + (cur_ready_cache_len + block_range) * stride_req_to_token_s + ) + memindex = tl.load(src_offset, mask=block_mask) + tl.store(out_memindex + cur_start_loc + block_range, memindex, mask=block_mask) + + +@torch.no_grad() +def select_kv_index_from_req_prefill( + req_to_token_indexs, + b_req_idx, + b_seq_len, + b_ready_cache_len, + b_start_loc, + max_q_seq_len, + token_num, +): + """Select the logical KV slots consumed by the current prefill step.""" + out_memindex = torch.empty((token_num,), dtype=torch.int32, device=req_to_token_indexs.device) + block, num_warps = get_triton_config(max_q_seq_len) + grid = (triton.cdiv(max_q_seq_len, block), b_req_idx.shape[0]) + _fwd_kernel_select_kv_index_from_req_prefill[grid]( + req_to_token_indexs, + b_req_idx, + b_seq_len, + b_ready_cache_len, + b_start_loc, + out_memindex, + req_to_token_indexs.stride(0), + req_to_token_indexs.stride(1), + BLOCK=block, + num_warps=num_warps, + num_stages=1, + ) + return out_memindex diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index c93939fb8b..a87fd08400 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -324,7 +324,8 @@ def prepare_dynamic_mtp_model_input( # All compaction work stays on the current CUDA stream and needs no host sync. model_input.to_cuda() - assert model_input.mem_indexes.shape[0] == dynamic_batch_size + if not model_input.mem_indexes_from_req_table: + assert model_input.mem_indexes.shape[0] == dynamic_batch_size selected_row_mask = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 87f54fd9e7..df35f90ab6 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -504,11 +504,8 @@ def get_decode_need_tokens(self): # 当开启 mtp 模式以后,每一次 decode 需要的 token 数量会增加 need_tokens = min(self.input_len + self.shm_cur_output_len - self.shm_cur_kv_len, self.chunked_prefill_size) if need_tokens == 1 and self._mtp_step > 0: - # self._mtp_step > 0 时,说明开启了mtp 模式,每次decode需要额外的mem token 资源 - # "vanilla_with_att" 模式需要的 mem 用量为 self._mtp_step + 1 - # "eagle_with_att" 模式需要的 mem 用量为 (self._mtp_step + 1)* 2 - # 为了简化统一 返回 (self._mtp_step + 1)* 2 - need_tokens = (self._mtp_step + 1) * 2 + # target verify 及后续 MTP 操作统一预留三倍窗口。 + need_tokens = (self._mtp_step + 1) * 3 return need_tokens diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index d037d7356f..2a7a6039c7 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -125,8 +125,7 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: def free_a_req_mem(self, free_token_index: List, req: "InferReq"): if self.radix_cache is None: - owned_kv_len = req.cur_kv_len if self.args.page_size == 1 else req.hold_kv_len - free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0:owned_kv_len]) + free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.hold_kv_len]) else: if not self.is_linear_att_mixed_model: self._full_att_free_req(free_token_index=free_token_index, req=req) @@ -149,8 +148,7 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): prefix_len, _ = self.radix_cache.insert(key, value) old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][old_prefix_len:prefix_len]) - if page_size > 1: - free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][cache_kv_len : req.hold_kv_len]) + free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][cache_kv_len : req.hold_kv_len]) if req.shared_kv_node is not None: assert req.shared_kv_node.node_prefix_total_len <= prefix_len self.radix_cache.dec_node_ref_counter(req.shared_kv_node) @@ -936,13 +934,11 @@ def _normal_decode_need_token_num(self) -> int: def _kv_cache_alloc_need(self, target_kv_len: int) -> int: page_size = self.args.page_size - if page_size == 1: - return target_kv_len - self.cur_kv_len target_hold_len = (target_kv_len + page_size - 1) // page_size * page_size - return target_hold_len - self.hold_kv_len + return max(target_hold_len - self.hold_kv_len, 0) def _mtp_decode_need_token_num(self) -> int: - return (1 + self.mtp_step) * 2 + return self._kv_cache_alloc_need(self.cur_kv_len + 3 * (1 + self.mtp_step)) class InferReqUpdatePack: diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 37a4551f24..dc6e1e579f 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -630,6 +630,23 @@ def remaining_prefill_tokens(req: InferReq) -> int: return sorted_reqs # 一些可以复用的通用功能函数 + def _alloc_req_kv_mem(self, req_obj: InferReq, alloc_token_num: int): + if alloc_token_num == 0: + return + + assert alloc_token_num > 0 and alloc_token_num % self.args.page_size == 0 + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(alloc_token_num) + + old_hold_kv_len = req_obj.hold_kv_len + new_hold_kv_len = old_hold_kv_len + alloc_token_num + mem_indexes = g_infer_context.req_manager.mem_manager.alloc(alloc_token_num) + g_infer_context.req_manager.req_to_token_indexs[ + req_obj.req_idx, old_hold_kv_len:new_hold_kv_len + ] = mem_indexes + req_obj.hold_kv_len = new_hold_kv_len + return + def _get_classed_reqs( self, req_ids: List[int] = None, @@ -720,6 +737,7 @@ def _get_classed_reqs( if is_decode: token_num = req_obj.decode_need_token_num() if token_num <= can_alloc_token_num: + self._alloc_req_kv_mem(req_obj, token_num) decode_reqs.append(req_obj) can_alloc_token_num -= token_num else: @@ -738,6 +756,7 @@ def _get_classed_reqs( if prefill_tokens + token_num > self.batch_max_tokens: continue if alloc_token_num <= can_alloc_token_num: + self._alloc_req_kv_mem(req_obj, alloc_token_num) prefill_tokens += token_num prefill_reqs.append(req_obj) can_alloc_token_num -= alloc_token_num diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 4d09476849..9322ec9d0d 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -377,12 +377,13 @@ def decode_mtp( extra_post_req_handle_func=self.extra_post_req_handle_func, ) - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu == 0, - ), - ) + if not model_input.mem_indexes_from_req_table: + proposal.extra_mem_indexes_cpu.append( + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, + ), + ) mtp_utils.free_mem_indexes( backend=self, extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 9a81927bc1..f9693f6cc1 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -604,12 +604,13 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): verify_run_reqs=run_reqs, ) - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu == 0, - ), - ) + if not model_input.mem_indexes_from_req_table: + proposal.extra_mem_indexes_cpu.append( + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, + ), + ) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( @@ -889,18 +890,20 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf req_num=req_num, accept_lengths_cpu=mtp_accept_len_cpu, ) - proposal.extra_mem_indexes_cpu.extend( - ( + if not model_input0.mem_indexes_from_req_table: + proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=model_input0.mem_indexes_cpu, free_mask_cpu=accepted_index_cpu0 == 0, - ), + ) + ) + if not model_input1.mem_indexes_from_req_table: + proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=model_input1.mem_indexes_cpu, free_mask_cpu=accepted_index_cpu1 == 0, - ), + ) ) - ) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 6833dda0a4..a9d2ff6042 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -7,42 +7,6 @@ INT64_MAX = torch.iinfo(torch.int64).max -def _alloc_kv_mem_indexes(reqs: List[InferReq], target_kv_lens: List[int], token_num: int) -> torch.Tensor: - """Allocate KV slots for this step; page ownership stays in the caller.""" - req_manager = g_infer_context.req_manager - mem_manager = req_manager.mem_manager - page_size = g_infer_context.args.page_size - - if page_size == 1: - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_num) - return mem_manager.alloc(token_num) - - assert len(reqs) == len(target_kv_lens) - alloc_need = sum(req._kv_cache_alloc_need(target_len) for req, target_len in zip(reqs, target_kv_lens)) - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(alloc_need) - - for req, target_len in zip(reqs, target_kv_lens): - assert req.cur_kv_len <= target_len and req.cur_kv_len <= req.hold_kv_len - assert req.hold_kv_len % page_size == 0 - new_hold_kv_len = (target_len + page_size - 1) // page_size * page_size - if new_hold_kv_len > req.hold_kv_len: - new_page_indexes = mem_manager.alloc(new_hold_kv_len - req.hold_kv_len) - req_manager.req_to_token_indexs[req.req_idx, req.hold_kv_len : new_hold_kv_len] = new_page_indexes - req.hold_kv_len = new_hold_kv_len - - # The request table is the source of truth. Return only the logical slots - # consumed by this step; the reserved tail remains owned by the request. - step_indexes = [ - req_manager.req_to_token_indexs[req.req_idx, req.cur_kv_len : target_len] - for req, target_len in zip(reqs, target_kv_lens) - ] - if step_indexes: - return torch.cat(step_indexes) - return torch.empty((0,), dtype=torch.int32, device=req_manager.req_to_token_indexs.device) - - def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> Tuple[ModelInput, List[InferReq]]: run_reqs = [] total_token_num = 0 @@ -101,10 +65,6 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> b_q_seq_len = torch.tensor(b_q_seq_len, dtype=torch.int32, device="cpu") b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len - # dynamic prompt cache 准备 token - target_kv_lens = [int(seq_len) for seq_len in b_seq_len] - mem_indexes = _alloc_kv_mem_indexes(run_reqs, target_kv_lens, input_ids.shape[0]) - model_input = ModelInput( batch_size=b_seq_len.shape[0], total_token_num=total_token_num, @@ -112,8 +72,7 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> max_kv_seq_len=max_kv_seq_len, max_cache_len=max_cache_len, input_ids=input_ids, - mem_indexes=mem_indexes if mem_indexes.is_cuda else None, - mem_indexes_cpu=mem_indexes if not mem_indexes.is_cuda else None, + mem_indexes_from_req_table=True, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -176,18 +135,13 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In device="cpu", ) - # dynamic prompt cache 准备 token - target_kv_lens = [req.cur_kv_len + 1 for req in req_objs] - mem_indexes = _alloc_kv_mem_indexes(req_objs, target_kv_lens, b_seq_len.shape[0]) - model_input = ModelInput( batch_size=b_seq_len.shape[0], total_token_num=total_token_num, max_q_seq_len=max_q_seq_len, max_kv_seq_len=max_kv_seq_len, input_ids=None, - mem_indexes=mem_indexes if mem_indexes.is_cuda else None, - mem_indexes_cpu=mem_indexes if not mem_indexes.is_cuda else None, + mem_indexes_from_req_table=True, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 9d59afd8a7..7e3f55f488 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -83,18 +83,15 @@ def prepare_decode_model_input( from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import prepare_dynamic_mtp_model_input from lightllm.server.router.model_infer.infer_batch import g_infer_context - # mem_indexes 是本轮 decode 新申请、尚未绑定请求和 token 位置的 KV slot。 - # 动态 verify 只需要保留 dynamic_batch_size 个任意 slot,因此 CPU 和已存在的 - # GPU 索引都可以直接截取前缀,无需等待 selected_row_mask_cpu。多申请的 CPU - # 尾部索引在这里立即归还;后续 forward 会根据压缩后的 b_req_idx/b_seq_len - # 建立保留 slot 与实际请求位置之间的映射。该操作需要放在下方动态输入构建 - # 之前,避免其内部 to_cuda 将原始完整 batch 的 mem indexes 全量复制到 GPU。 - unused_mem_indexes_cpu = model_input.mem_indexes_cpu[plan.dynamic_batch_size :] - model_input.mem_indexes_cpu = model_input.mem_indexes_cpu[: plan.dynamic_batch_size] - if model_input.mem_indexes is not None: - model_input.mem_indexes = model_input.mem_indexes[: plan.dynamic_batch_size] - if unused_mem_indexes_cpu.numel() > 0: - g_infer_context.req_manager.mem_manager.free(unused_mem_indexes_cpu) + if not model_input.mem_indexes_from_req_table: + # 兼容外部构造的旧式 ModelInput:尚未绑定请求位置的临时 KV slot + # 可以随动态 verify 行一起裁剪并立即释放。 + unused_mem_indexes_cpu = model_input.mem_indexes_cpu[plan.dynamic_batch_size :] + model_input.mem_indexes_cpu = model_input.mem_indexes_cpu[: plan.dynamic_batch_size] + if model_input.mem_indexes is not None: + model_input.mem_indexes = model_input.mem_indexes[: plan.dynamic_batch_size] + if unused_mem_indexes_cpu.numel() > 0: + g_infer_context.req_manager.mem_manager.free(unused_mem_indexes_cpu) model_input, selected_row_mask = prepare_dynamic_mtp_model_input( model_input=model_input, diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 7a5f565aef..c604f70663 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -116,7 +116,7 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): expected_batch_size, ) in execution_configs: model = TpPartBaseModel.__new__(TpPartBaseModel) - model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode) + model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode, page_size=1) model.tp_world_size_ = tp_world_size model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=99, page_size=1) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py index ec5dad5d34..5e93f5960e 100644 --- a/unit_tests/common/basemodel/test_overlap_utils.py +++ b/unit_tests/common/basemodel/test_overlap_utils.py @@ -92,7 +92,7 @@ def fake_gather(**kwargs): monkeypatch.setattr(basemodel, "gather_token_prefill_decode_mixed", fake_gather) model = TpPartBaseModel.__new__(TpPartBaseModel) - model.args = SimpleNamespace(enable_prefill_decode_mixed=True) + model.args = SimpleNamespace(enable_prefill_decode_mixed=True, page_size=1) model.req_manager = SimpleNamespace( req_sampling_params_manager=SimpleNamespace(req_to_next_token_ids=object()), ) @@ -138,6 +138,7 @@ def fake_gather(**kwargs): monkeypatch.setattr(basemodel, "gather_token", fake_gather) model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(page_size=1) model.req_manager = SimpleNamespace( req_sampling_params_manager=SimpleNamespace(req_to_next_token_ids=object()), ) @@ -182,7 +183,7 @@ def test_overlap_decode_cuda_pads_empty_side_and_unpads_outputs(monkeypatch): model_input1.to_cuda() model = TpPartBaseModel.__new__(TpPartBaseModel) - model.args = SimpleNamespace(enable_tpsp_mix_mode=True) + model.args = SimpleNamespace(enable_tpsp_mix_mode=True, page_size=1) model.tp_world_size_ = 2 model.graph = None model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) diff --git a/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py b/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py new file mode 100644 index 0000000000..2b0affcc4b --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py @@ -0,0 +1,77 @@ +import pytest +import torch +from types import SimpleNamespace + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import ( + select_kv_index_from_req, + select_kv_index_from_req_prefill, +) + + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") + + +def test_select_kv_index_from_req_for_decode(): + req_to_token_indexs = torch.tensor( + [ + [10, 11, 12, 13, 14, 15, 16, 17], + [20, 21, 22, 23, 24, 25, 26, 27], + ], + dtype=torch.int32, + device="cuda", + ) + + mem_indexes = select_kv_index_from_req( + req_to_token_indexs=req_to_token_indexs, + b_req_idx=torch.tensor([1, 0], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 6], dtype=torch.int32, device="cuda"), + ) + + assert mem_indexes.cpu().tolist() == [22, 15] + + +def test_select_kv_index_from_req_for_prefill(): + req_to_token_indexs = torch.tensor( + [ + [10, 11, 12, 13, 14, 15, 16, 17], + [20, 21, 22, 23, 24, 25, 26, 27], + ], + dtype=torch.int32, + device="cuda", + ) + + mem_indexes = select_kv_index_from_req_prefill( + req_to_token_indexs=req_to_token_indexs, + b_req_idx=torch.tensor([1, 0], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([4, 7], dtype=torch.int32, device="cuda"), + b_ready_cache_len=torch.tensor([2, 3], dtype=torch.int32, device="cuda"), + b_start_loc=torch.tensor([0, 2], dtype=torch.int32, device="cuda"), + max_q_seq_len=4, + token_num=6, + ) + + assert mem_indexes.cpu().tolist() == [22, 23, 13, 14, 15, 16] + + +def test_model_selects_reserved_indexes_when_execution_starts(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(page_size=4) + model.req_manager = SimpleNamespace( + req_to_token_indexs=torch.tensor( + [[10, 11, 12, 13], [20, 21, 22, 23]], + dtype=torch.int32, + device="cuda", + ) + ) + model_input = SimpleNamespace( + is_prefill=False, + mem_indexes_from_req_table=True, + b_req_idx=torch.tensor([1, 0], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4], dtype=torch.int32, device="cuda"), + mem_indexes=None, + ) + + model._select_page_mem_indexes(model_input) + + assert model_input.mem_indexes.cpu().tolist() == [22, 13] diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py index b53cd78318..50097d40f9 100644 --- a/unit_tests/common/test_req_manager_page.py +++ b/unit_tests/common/test_req_manager_page.py @@ -1,9 +1,10 @@ from types import SimpleNamespace -import pytest import torch from lightllm.server.router.model_infer.mode_backend import generic_pre_process +from lightllm.server.router.model_infer.infer_batch import InferReq, InferenceContext +from lightllm.server.router.model_infer.mode_backend import base_backend class _FakeMemManager: @@ -32,7 +33,10 @@ def _make_context(monkeypatch): radix_cache=None, ) monkeypatch.setattr(generic_pre_process, "g_infer_context", context) - return context + monkeypatch.setattr(base_backend, "g_infer_context", context) + backend = base_backend.ModeBackend.__new__(base_backend.ModeBackend) + backend.args = context.args + return context, backend def _make_req(req_idx): @@ -42,46 +46,107 @@ def _make_req(req_idx): def test_request_reuses_reserved_page_tail_before_allocating_next_page(monkeypatch): - context = _make_context(monkeypatch) + context, backend = _make_context(monkeypatch) req = _make_req(0) - indexes = generic_pre_process._alloc_kv_mem_indexes([req], [3], 3) - assert indexes.tolist() == [0, 1, 2] + backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(3)) assert req.hold_kv_len == 4 assert context.req_manager.mem_manager.alloc_sizes == [4] assert context.req_manager.req_to_token_indexs[0, :4].tolist() == [0, 1, 2, 3] req.cur_kv_len = 3 - indexes = generic_pre_process._alloc_kv_mem_indexes([req], [4], 1) - assert indexes.tolist() == [3] + backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(4)) assert context.req_manager.mem_manager.alloc_sizes == [4] req.cur_kv_len = 4 - indexes = generic_pre_process._alloc_kv_mem_indexes([req], [6], 2) - assert indexes.tolist() == [4, 5] + backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(6)) assert req.hold_kv_len == 8 assert context.req_manager.mem_manager.alloc_sizes == [4, 4] assert context.req_manager.req_to_token_indexs[0, :8].tolist() == list(range(8)) -def test_step_indexes_follow_request_order(monkeypatch): - _make_context(monkeypatch) +def test_reservation_fills_each_request_table_row(monkeypatch): + _, backend = _make_context(monkeypatch) req0 = _make_req(0) req1 = _make_req(1) - indexes = generic_pre_process._alloc_kv_mem_indexes([req0, req1], [2, 3], 5) + backend._alloc_req_kv_mem(req0, req0._kv_cache_alloc_need(2)) + backend._alloc_req_kv_mem(req1, req1._kv_cache_alloc_need(3)) - assert indexes.tolist() == [0, 1, 4, 5, 6] + assert generic_pre_process.g_infer_context.req_manager.req_to_token_indexs[0, :4].tolist() == [0, 1, 2, 3] + assert generic_pre_process.g_infer_context.req_manager.req_to_token_indexs[1, :4].tolist() == [4, 5, 6, 7] assert req0.hold_kv_len == req1.hold_kv_len == 4 -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -def test_step_indexes_are_assembled_from_gpu_request_table(monkeypatch): - context = _make_context(monkeypatch) - context.req_manager.req_to_token_indexs = context.req_manager.req_to_token_indexs.cuda() +def test_decode_reserves_mtp_headroom(monkeypatch): + context, backend = _make_context(monkeypatch) req = _make_req(0) + req.cur_kv_len = 3 + req.hold_kv_len = 4 + req.mtp_step = 2 + req.multimodal_params = {"images": [], "audios": []} + req.shared_kv_node = None + req.get_cur_total_len = lambda: 4 + req.get_radix_cache_shared_len = lambda: 0 + req.args = context.args + + context.req_manager.req_to_token_indexs[0, :4] = torch.arange(4, dtype=torch.int32) + context.req_manager.mem_manager.next_index = 4 + alloc_token_num = InferReq._mtp_decode_need_token_num(req) + backend._alloc_req_kv_mem(req, alloc_token_num) + + model_input, run_reqs = generic_pre_process.prepare_decode_inputs([req]) + + assert run_reqs == [req, req, req] + assert model_input.b_seq_len.tolist() == [4, 5, 6] + assert model_input.mem_indexes_cpu is None + assert req.hold_kv_len == 12 + assert context.req_manager.mem_manager.alloc_sizes == [8] + assert context.req_manager.req_to_token_indexs[0, :12].tolist() == list(range(12)) + + +def test_page_size_one_uses_the_same_scheduler_preallocation(monkeypatch): + context, backend = _make_context(monkeypatch) + context.args.page_size = 1 + req = _make_req(0) + req.cur_kv_len = 3 + req.hold_kv_len = 3 + req.mtp_step = 2 + req.multimodal_params = {"images": [], "audios": []} + req.shared_kv_node = None + req.get_cur_total_len = lambda: 4 + req.get_radix_cache_shared_len = lambda: 0 + req.args = context.args + req._kv_cache_alloc_need = lambda target_len: InferReq._kv_cache_alloc_need(req, target_len) + + context.req_manager.req_to_token_indexs[0, :3] = torch.arange(3, dtype=torch.int32) + context.req_manager.mem_manager.next_index = 3 + alloc_token_num = InferReq._mtp_decode_need_token_num(req) + backend._alloc_req_kv_mem(req, alloc_token_num) + + model_input, _ = generic_pre_process.prepare_decode_inputs([req]) + + assert model_input.mem_indexes_cpu is None + assert model_input.mem_indexes_from_req_table is True + assert req.hold_kv_len == 12 + assert context.req_manager.mem_manager.alloc_sizes == [9] + assert context.req_manager.req_to_token_indexs[0, :12].tolist() == list(range(12)) + + +def test_page_size_one_frees_all_preallocated_indexes(): + infer_context = InferenceContext.__new__(InferenceContext) + infer_context.args = SimpleNamespace(page_size=1) + infer_context.radix_cache = None + infer_context.req_manager = SimpleNamespace(req_to_token_indexs=torch.arange(16, dtype=torch.int32)[None, :]) + req = SimpleNamespace( + req_idx=0, + cur_kv_len=3, + hold_kv_len=12, + shm_req=SimpleNamespace(shm_cur_kv_len=3), + ) + free_token_indexes = [] - indexes = generic_pre_process._alloc_kv_mem_indexes([req], [3], 3) + infer_context.free_a_req_mem(free_token_indexes, req) - assert indexes.is_cuda - assert indexes.cpu().tolist() == [0, 1, 2] + assert free_token_indexes[0].tolist() == list(range(12)) + assert req.cur_kv_len == req.hold_kv_len == 0 diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 51f8d1c82c..74d65d1817 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -1,5 +1,6 @@ import pytest import easydict +from types import SimpleNamespace from lightllm.server.core.objs.req import Req, ChunkedPrefillReq, SamplingParams from lightllm.server.core.objs.token_metadata import ReqFinalTokenMetadata from lightllm.utils.envs_utils import set_env_start_args @@ -54,6 +55,18 @@ def test_get_used_tokens(req): assert req.get_used_tokens() == 5 +def test_mtp_decode_reserves_three_windows(): + req = SimpleNamespace( + input_len=4, + shm_cur_output_len=0, + shm_cur_kv_len=3, + chunked_prefill_size=128, + _mtp_step=2, + ) + + assert ChunkedPrefillReq.get_decode_need_tokens(req) == 9 + + def test_final_token_metadata_read_returns_actual_prompt_tokens(req): req.sample_params.prompt_logprobs = 0 req.shm_logprobs.arr["logprob"][1] = -0.5 diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 1e789d4420..05745016cf 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -11,7 +11,11 @@ def _patch_empty_input_context(monkeypatch): ) infer_context = SimpleNamespace( args=SimpleNamespace(page_size=1), - req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager), + req_manager=SimpleNamespace( + HOLD_REQUEST_ID=-1, + mem_manager=mem_manager, + req_to_token_indexs=torch.zeros((16, 32), dtype=torch.int32), + ), radix_cache=None, ) monkeypatch.setattr(generic_pre_process, "g_infer_context", infer_context) @@ -27,6 +31,7 @@ def _make_prefill_req(req_idx: int, token_num: int): return SimpleNamespace( req_idx=req_idx, cur_kv_len=0, + hold_kv_len=token_num, multimodal_params={"images": [], "audios": []}, get_chuncked_input_token_ids=lambda: input_token_ids, get_input_token_ids=lambda: input_token_ids, @@ -38,6 +43,7 @@ def _make_decode_req(req_idx: int): return SimpleNamespace( req_idx=req_idx, cur_kv_len=3, + hold_kv_len=4, mtp_step=0, multimodal_params={"images": [], "audios": []}, shared_kv_node=None, @@ -54,7 +60,7 @@ def test_prepare_prefill_inputs_allows_empty_batch(monkeypatch): assert run_reqs == [] assert model_input.batch_size == 0 assert model_input.input_ids.shape == (0,) - assert model_input.mem_indexes_cpu.shape == (0,) + assert model_input.mem_indexes_cpu is None assert model_input.b_req_idx.shape == (0,) assert model_input.b_prefill_start_loc.shape == (0,) assert model_input.b_prefill_has_output_cpu == [] @@ -70,7 +76,7 @@ def test_prepare_decode_inputs_allows_empty_batch(monkeypatch): assert run_reqs == [] assert model_input.batch_size == 0 assert model_input.input_ids is None - assert model_input.mem_indexes_cpu.shape == (0,) + assert model_input.mem_indexes_cpu is None assert model_input.b_req_idx.shape == (0,) assert model_input.b_position_delta.shape == (0,) assert model_input.b_shared_seq_len.shape == (0,) @@ -170,4 +176,4 @@ def test_overlap_decode_preserves_empty_microbatch(monkeypatch): assert run_reqs1 == [] assert model_input1.batch_size == 0 assert model_input1.b_req_idx.shape == (0,) - assert model_input1.mem_indexes_cpu.shape == (0,) + assert model_input1.mem_indexes_cpu is None diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 7baf34061b..1977606198 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -704,6 +704,7 @@ def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypa batch_size=4, mem_indexes=torch.tensor([20, 21, 22, 23], dtype=torch.int32), mem_indexes_cpu=torch.tensor([10, 11, 12, 13], dtype=torch.int32), + mem_indexes_from_req_table=False, ) plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=1, pre_draft_step=1) From edda3b6e46dd847cffab944209dd6b4adc0b3529 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 03:08:51 +0000 Subject: [PATCH 03/15] refactor: derive KV indexes in infer state --- docs/kv_cache_page_size.md | 4 +- lightllm/common/basemodel/basemodel.py | 123 ++---------------- lightllm/common/basemodel/batch_objs.py | 15 --- lightllm/common/basemodel/cuda_graph.py | 7 - .../common/basemodel/prefill_cuda_graph.py | 13 -- .../triton_kernel/dynamic_mtp_utils.py | 2 - .../triton_kernel/select_mtp_rows.py | 11 -- lightllm/models/qwen3_dflash/model.py | 2 +- .../model_infer/mode_backend/base_backend.py | 3 +- .../mode_backend/chunked_prefill/impl.py | 8 -- .../mode_backend/dp_backend/impl.py | 24 ---- .../mode_backend/generic_pre_process.py | 2 - .../dp_overlap_proposers/eagle_no_att.py | 3 - .../dp_overlap_proposers/eagle_with_att.py | 13 +- .../dp_overlap_proposers/vanilla_no_att.py | 3 - .../model_infer/mtp_speculative/engine.py | 11 -- .../mtp_speculative/proposers/dflash.py | 14 +- .../mtp_speculative/proposers/dspark.py | 17 +-- .../mtp_speculative/proposers/eagle_no_att.py | 3 - .../proposers/eagle_with_att.py | 17 +-- .../proposers/vanilla_no_att.py | 3 - .../common/basemodel/test_model_input.py | 7 - .../common/basemodel/test_model_output.py | 5 +- .../common/basemodel/test_overlap_utils.py | 8 +- .../triton_kernel/test_dynamic_mtp_utils.py | 11 -- .../test_select_kv_index_from_req.py | 6 +- unit_tests/common/test_req_manager_page.py | 3 - .../models/test_qwen3_dspark_model_output.py | 2 +- .../test_dp_overlap_spec_engine.py | 1 - .../mode_backend/test_generic_pre_process.py | 3 - .../mtp_speculative/test_dflash.py | 10 +- .../mtp_speculative/test_dspark.py | 10 +- .../mtp_speculative/test_eagle_no_att.py | 4 - .../mtp_speculative/test_eagle_overlap.py | 41 +----- .../mtp_speculative/test_eagle_with_att.py | 12 +- .../mtp_speculative/test_planner.py | 18 +-- .../mtp_speculative/test_vanilla_no_att.py | 9 -- 37 files changed, 38 insertions(+), 410 deletions(-) diff --git a/docs/kv_cache_page_size.md b/docs/kv_cache_page_size.md index 9ff7089b59..a89df89afc 100644 --- a/docs/kv_cache_page_size.md +++ b/docs/kv_cache_page_size.md @@ -11,8 +11,8 @@ KV 存储和请求表,但资源的申请、缓存和释放以完整物理页 `cur_kv_len <= hold_kv_len` 且 `hold_kv_len % page_size == 0`。 2. 一个物理页内的 KV 槽连续,页首索引可被 `page_size` 整除。请求表保存已持有页的全部 token 索引, 包括尚未使用的尾部槽位。 -3. `req_to_token_indexs` 保存请求拥有的完整页;模型创建 `InferState` 时再从中选出本轮真实参与计算的 - token,`InferState.mem_index` 不包含预留尾部。 +3. `req_to_token_indexs` 保存请求拥有的完整页;创建 `InferState` 时根据请求索引和序列位置直接聚合本轮 + 真实参与计算的 token,`ModelInput` 不携带物理 KV 索引,`InferState.mem_index` 不包含预留尾部。 4. Radix Cache 只插入、拆分、命中和淘汰完整页;不足一页的请求尾部在请求结束或暂停时整页回收。 5. `page_size=1` 与多 token 页使用相同的调度期预留和模型执行期索引选择路径。 diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 2c2cda2bca..1081c94d7f 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -15,10 +15,8 @@ from lightllm.common.kv_cache_mem_manager import MemoryManager from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.req_manager import ReqManager -from lightllm.common.infer_utils import init_req_to_token_indexes from lightllm.common.build_utils import repair_config from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import ( - copy_kv_index_to_req, select_kv_index_from_req, select_kv_index_from_req_prefill, ) @@ -383,19 +381,15 @@ def _init_hidden_collector(self): @torch.no_grad() def forward(self, model_input: ModelInput): model_input.to_cuda() - self._select_page_mem_indexes(model_input) - assert model_input.mem_indexes.is_cuda if model_input.is_prefill: return self._prefill(model_input=model_input) else: return self._decode(model_input) - def _select_page_mem_indexes(self, model_input: ModelInput): - if not model_input.mem_indexes_from_req_table or model_input.mem_indexes is not None: - return + def _select_mem_indexes(self, model_input: ModelInput): if model_input.is_prefill: - model_input.mem_indexes = select_kv_index_from_req_prefill( + return select_kv_index_from_req_prefill( req_to_token_indexs=self.req_manager.req_to_token_indexs, b_req_idx=model_input.b_req_idx, b_seq_len=model_input.b_seq_len, @@ -404,13 +398,11 @@ def _select_page_mem_indexes(self, model_input: ModelInput): max_q_seq_len=model_input.max_q_seq_len, token_num=model_input.input_ids.shape[0], ) - else: - model_input.mem_indexes = select_kv_index_from_req( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=model_input.b_req_idx, - b_seq_len=model_input.b_seq_len, - ) - return + return select_kv_index_from_req( + req_to_token_indexs=self.req_manager.req_to_token_indexs, + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + ) def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -442,7 +434,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) infer_state.mem_manager = self.mem_manager infer_state.req_manager = self.req_manager - infer_state.mem_index = model_input.mem_indexes + infer_state.mem_index = self._select_mem_indexes(model_input) infer_state.microbatch_index = microbatch_index infer_state.dist_group = dist_group_manager.get_group(microbatch_index) @@ -487,12 +479,6 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_position_delta = F.pad( new_model_input.b_position_delta, (0, padded_batch_size), mode="constant", value=0 ) - new_model_input.mem_indexes = F.pad( - new_model_input.mem_indexes, - (0, padded_batch_size), - mode="constant", - value=self.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % self.mem_manager.page_size), - ) new_model_input.multimodal_params = new_model_input.multimodal_params + [ {"images": [], "audios": []} for _ in range(padded_batch_size) ] @@ -530,18 +516,6 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.max_kv_seq_len = max(padded_token_num, model_input.max_kv_seq_len) new_model_input.max_cache_len = max(0, model_input.max_cache_len) new_model_input.input_ids = F.pad(new_model_input.input_ids, (0, padded_token_num), mode="constant", value=1) - hold_mem_indexes = ( - self.mem_manager.HOLD_TOKEN_MEMINDEX - + torch.arange( - padded_token_num, - dtype=new_model_input.mem_indexes.dtype, - device=new_model_input.mem_indexes.device, - ) - % self.mem_manager.page_size - ) - new_model_input.mem_indexes = torch.cat( - (new_model_input.mem_indexes, hold_mem_indexes), - ) new_model_input.b_req_idx = F.pad( new_model_input.b_req_idx, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID ) @@ -629,16 +603,6 @@ def _prefill( ) infer_state = self._create_inferstate(model_input) - if not model_input.mem_indexes_from_req_table: - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state.b_req_idx, - b_seq_len=infer_state.b_seq_len, - b_ready_cache_len=infer_state.b_ready_cache_len, - b_start_loc=model_input.b_prefill_start_loc, - alloc_mem_index=infer_state.mem_index, - max_q_seq_len=infer_state.max_q_seq_len, - ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -698,13 +662,6 @@ def _decode( # attention backend 会根据该标记准备 CUDA Graph capture 专用状态, # 因此必须在 init_att_state 之前完成赋值。 infer_state.is_cuda_graph = need_capture - if not model_input.mem_indexes_from_req_table: - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state.b_req_idx, - infer_state.b_seq_len, - infer_state.mem_index, - ) infer_state.init_some_extra_state(self) infer_state.init_att_state() @@ -823,7 +780,6 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod for model_input in (model_input0, model_input1): model_input.to_cuda() - self._select_page_mem_indexes(model_input) if self.args.enable_prefill_decode_mixed and model_input.input_ids.shape[0] > 0: gather_token_prefill_decode_mixed( input_ids=model_input.input_ids, @@ -836,9 +792,6 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod return self._microbatch_overlap_prefill_cuda(model_input0, model_input1) def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input1: ModelInput): - assert model_input0.mem_indexes.is_cuda - assert model_input1.mem_indexes.is_cuda - assert self.args.enable_tpsp_mix_mode origin_handle_token_num0 = model_input0.input_ids.shape[0] origin_handle_token_num1 = model_input1.input_ids.shape[0] @@ -861,30 +814,10 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input ) infer_state0 = self._create_inferstate(model_input0, 0) - if not model_input0.mem_indexes_from_req_table: - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state0.b_req_idx, - b_seq_len=infer_state0.b_seq_len, - b_ready_cache_len=infer_state0.b_ready_cache_len, - b_start_loc=model_input0.b_prefill_start_loc, - alloc_mem_index=infer_state0.mem_index, - max_q_seq_len=infer_state0.max_q_seq_len, - ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(model_input1, 1) - if not model_input1.mem_indexes_from_req_table: - init_req_to_token_indexes( - req_to_token_indexs=self.req_manager.req_to_token_indexs, - b_req_idx=infer_state1.b_req_idx, - b_seq_len=infer_state1.b_seq_len, - b_ready_cache_len=infer_state1.b_ready_cache_len, - b_start_loc=model_input1.b_prefill_start_loc, - alloc_mem_index=infer_state1.mem_index, - max_q_seq_len=infer_state1.max_q_seq_len, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -916,7 +849,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode for model_input in (model_input0, model_input1): model_input.to_cuda() - self._select_page_mem_indexes(model_input) if model_input.input_ids is None: if model_input.batch_size > 0: model_input.input_ids = gather_token( @@ -934,8 +866,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1: ModelInput): assert self.args.enable_tpsp_mix_mode - assert model_input0.mem_indexes.is_cuda - assert model_input1.mem_indexes.is_cuda origin_batch_size0 = model_input0.batch_size origin_batch_size1 = model_input1.batch_size @@ -950,25 +880,11 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 padded_model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) infer_state0 = self._create_inferstate(padded_model_input0, 0) infer_state0.is_cuda_graph = need_capture - if not padded_model_input0.mem_indexes_from_req_table: - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state0.b_req_idx, - infer_state0.b_seq_len, - infer_state0.mem_index, - ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state1.is_cuda_graph = need_capture - if not padded_model_input1.mem_indexes_from_req_table: - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state1.b_req_idx, - infer_state1.b_seq_len, - infer_state1.mem_index, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -990,24 +906,10 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) infer_state0 = self._create_inferstate(model_input0, 0) - if not model_input0.mem_indexes_from_req_table: - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state0.b_req_idx, - infer_state0.b_seq_len, - infer_state0.mem_index, - ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() infer_state1 = self._create_inferstate(model_input1, 1) - if not model_input1.mem_indexes_from_req_table: - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state1.b_req_idx, - infer_state1.b_seq_len, - infer_state1.mem_index, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -1149,6 +1051,7 @@ def _check_max_len_infer(self): dummy_input_ids = torch.ones(self.batch_max_tokens, dtype=torch.int64, device="cuda") b_req_idx = torch.tensor([self.req_manager.alloc()], dtype=torch.int32, device="cuda") mem_indexes = self.mem_manager.alloc(len(dummy_input_ids)).cuda() + self.req_manager.req_to_token_indexs[b_req_idx[0], : len(mem_indexes)] = mem_indexes b_seq_len = torch.ones(1, dtype=torch.int32, device="cuda") b_seq_len[:] = self.batch_max_tokens b_ready_cache_len = torch.zeros(1, dtype=torch.int32, device="cuda") @@ -1163,7 +1066,6 @@ def _check_max_len_infer(self): max_kv_seq_len=self.batch_max_tokens, max_cache_len=0, input_ids=dummy_input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, @@ -1228,6 +1130,7 @@ def _autotune_warmup(self): ) b_req_idx = torch.tensor([self.req_manager.alloc()], dtype=torch.int32, device="cuda") mem_indexes = self.mem_manager.alloc(len(dummy_input_ids)).cuda() + self.req_manager.req_to_token_indexs[b_req_idx[0], : len(mem_indexes)] = mem_indexes b_seq_len = torch.ones(1, dtype=torch.int32, device="cuda") b_seq_len[:] = input_len b_ready_cache_len = torch.zeros(1, dtype=torch.int32, device="cuda") @@ -1242,7 +1145,6 @@ def _autotune_warmup(self): max_kv_seq_len=input_len, max_cache_len=0, input_ids=dummy_input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, @@ -1291,9 +1193,6 @@ def _init_padded_req(self): b_req_idx = torch.tensor( [self.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) - mem_indexes = torch.tensor( - [self.mem_manager.HOLD_TOKEN_MEMINDEX for _ in range(batch_size)], dtype=torch.int32, device="cuda" - ) b_seq_len = torch.ones(batch_size, dtype=torch.int32, device="cuda") b_ready_cache_len = torch.zeros(batch_size, dtype=torch.int32, device="cuda") b_q_seq_len = b_seq_len - b_ready_cache_len @@ -1308,7 +1207,6 @@ def _init_padded_req(self): max_kv_seq_len=prefill_input_len, max_cache_len=0, input_ids=dummy_input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -1327,7 +1225,6 @@ def _init_padded_req(self): del model_input del dummy_input_ids del b_req_idx - del mem_indexes del b_seq_len del b_ready_cache_len del model_output diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index f783604ba9..bf73a10f91 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -32,10 +32,6 @@ class ModelInput: # Decode 逐行携带的 radix node 标识。相同 id 表示请求引用同一个共享 # radix node;该 id 只用于重建 diverse attention 的 b_mark_shared_group。 b_shared_radix_node_id: torch.Tensor = None - mem_indexes: torch.Tensor = None - # The scheduler reserves KV capacity directly in req_to_token_indexs; - # BaseModel resolves the logical indexes when execution starts. - mem_indexes_from_req_table: bool = False is_prefill: bool = False b_ready_cache_len: torch.Tensor = None # Request/row-aligned MRoPE position offset. It is decode-only; prefill @@ -44,8 +40,6 @@ class ModelInput: b_position_delta: torch.Tensor = None b_prefill_start_loc: torch.Tensor = None multimodal_params: list = None - # cpu 变量 - mem_indexes_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 # 的一些变量 # 标记 prefill 请求是否会在本轮产生输出。Prefill 必填(空 batch 使用空 list),decode 不使用。 @@ -62,8 +56,6 @@ def to_cuda(self): self.check_input() # Prefill 和 decode 都必须提供的公共张量。 - if self.mem_indexes is None and self.mem_indexes_cpu is not None: - self.mem_indexes = self.mem_indexes_cpu.cuda(non_blocking=True) self.b_req_idx = self.b_req_idx.cuda(non_blocking=True) self.b_seq_len = self.b_seq_len.cuda(non_blocking=True) self.b_mtp_index = self.b_mtp_index.cuda(non_blocking=True) @@ -97,8 +89,6 @@ def check_input(self): assert self.b_mtp_index is not None assert self.b_seq_len is not None assert self.multimodal_params is not None - assert self.mem_indexes_from_req_table or self.mem_indexes is not None or self.mem_indexes_cpu is not None - assert self.b_req_idx.shape == (self.batch_size,) assert self.b_mtp_index.shape == self.b_req_idx.shape assert self.b_seq_len.shape == self.b_req_idx.shape @@ -126,11 +116,6 @@ def check_input(self): assert self.b_shared_seq_len.shape == self.b_req_idx.shape assert self.b_shared_radix_node_id.shape == self.b_req_idx.shape - mem_indexes = self.mem_indexes if self.mem_indexes is not None else self.mem_indexes_cpu - if mem_indexes is not None: - assert mem_indexes.ndim == 1 - - @dataclass class ModelMtpOutputCollector: """保存一次模型 forward 为 MTP 推理产生的可选输出。""" diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 6c56bf9c13..bbb661af6b 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -255,8 +255,6 @@ def warmup(self, model): total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") - mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() - mem_indexes.fill_(model.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % model.mem_manager.page_size)) b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) @@ -271,7 +269,6 @@ def warmup(self, model): max_q_seq_len=1, max_kv_seq_len=max_len_in_batch, input_ids=input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, @@ -285,7 +282,6 @@ def warmup(self, model): model_output: ModelOutput = model.forward(model_input) del model_output del input_ids - del mem_indexes del b_req_idx del b_seq_len @@ -317,8 +313,6 @@ def warmup_overlap(self, model): total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") - mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() - mem_indexes.fill_(model.mem_manager.HOLD_TOKEN_MEMINDEX + (1 % model.mem_manager.page_size)) b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) @@ -335,7 +329,6 @@ def warmup_overlap(self, model): max_kv_seq_len=max_len_in_batch, input_ids=input_ids, b_mtp_index=b_mtp_index, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_shared_seq_len=b_shared_seq_len, diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 6a40920cc9..28d44fefbf 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -194,11 +194,6 @@ def warmup(self, model): logger.info(f"Capture prefill cudagraph, handle_token_num: {handle_token_num}") total_token_num = handle_token_num input_ids = torch.tensor([1 for _ in range(total_token_num)], dtype=torch.int64, device="cuda") - mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() - mem_indexes.copy_( - model.mem_manager.HOLD_TOKEN_MEMINDEX - + torch.arange(total_token_num, dtype=torch.int32, device="cuda") % model.mem_manager.page_size - ) b_req_idx = torch.tensor([model.req_manager.HOLD_REQUEST_ID], dtype=torch.int32, device="cuda") b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) @@ -214,7 +209,6 @@ def warmup(self, model): max_kv_seq_len=total_token_num, max_cache_len=0, input_ids=input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -229,7 +223,6 @@ def warmup(self, model): model_output: ModelOutput = model.forward(model_input) del model_output del input_ids - del mem_indexes del b_req_idx del b_seq_len @@ -259,11 +252,6 @@ def warmup_overlap(self, model): # dummy prefill, capture the cudagraph total_token_num = handle_token_num input_ids = torch.tensor([1 for _ in range(total_token_num)], dtype=torch.int64, device="cuda") - mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() - mem_indexes.copy_( - model.mem_manager.HOLD_TOKEN_MEMINDEX - + torch.arange(total_token_num, dtype=torch.int32, device="cuda") % model.mem_manager.page_size - ) b_req_idx = torch.tensor([model.req_manager.HOLD_REQUEST_ID], dtype=torch.int32, device="cuda") b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) @@ -279,7 +267,6 @@ def warmup_overlap(self, model): max_kv_seq_len=total_token_num, max_cache_len=0, input_ids=input_ids, - mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index a87fd08400..ad995b2b5d 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -324,8 +324,6 @@ def prepare_dynamic_mtp_model_input( # All compaction work stays on the current CUDA stream and needs no host sync. model_input.to_cuda() - if not model_input.mem_indexes_from_req_table: - assert model_input.mem_indexes.shape[0] == dynamic_batch_size selected_row_mask = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, diff --git a/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py index 5e230ce855..b542eaeee8 100644 --- a/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py +++ b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py @@ -13,7 +13,6 @@ class SelectedMtpRows(NamedTuple): b_req_idx: torch.Tensor b_mtp_index: torch.Tensor b_seq_len: torch.Tensor - mem_indexes: torch.Tensor b_shared_seq_len: torch.Tensor b_shared_radix_node_id: torch.Tensor b_position_delta: torch.Tensor @@ -34,8 +33,6 @@ def _select_accepted_tail_rows_kernel( b_mtp_index_stride, b_seq_len, b_seq_len_stride, - mem_indexes, - mem_indexes_stride, b_shared_seq_len, b_shared_seq_len_stride, b_shared_radix_node_id, @@ -49,7 +46,6 @@ def _select_accepted_tail_rows_kernel( out_b_req_idx, out_b_mtp_index, out_b_seq_len, - out_mem_indexes, out_b_shared_seq_len, out_b_shared_radix_node_id, out_b_position_delta, @@ -64,7 +60,6 @@ def _select_accepted_tail_rows_kernel( tl.store(out_b_req_idx + out_row, tl.load(b_req_idx + src_row * b_req_idx_stride)) tl.store(out_b_mtp_index + out_row, tl.load(b_mtp_index + src_row * b_mtp_index_stride)) tl.store(out_b_seq_len + out_row, tl.load(b_seq_len + src_row * b_seq_len_stride)) - tl.store(out_mem_indexes + out_row, tl.load(mem_indexes + src_row * mem_indexes_stride)) tl.store( out_b_shared_seq_len + out_row, tl.load(b_shared_seq_len + src_row * b_shared_seq_len_stride), @@ -103,7 +98,6 @@ def select_accepted_tail_rows( b_req_idx: torch.Tensor, b_mtp_index: torch.Tensor, b_seq_len: torch.Tensor, - mem_indexes: torch.Tensor, b_shared_seq_len: torch.Tensor, b_shared_radix_node_id: torch.Tensor, b_position_delta: torch.Tensor, @@ -121,7 +115,6 @@ def select_accepted_tail_rows( b_req_idx, b_mtp_index, b_seq_len, - mem_indexes, b_shared_seq_len, b_shared_radix_node_id, b_position_delta, @@ -134,7 +127,6 @@ def select_accepted_tail_rows( b_req_idx=b_req_idx.new_empty((req_num,)), b_mtp_index=b_mtp_index.new_empty((req_num,)), b_seq_len=b_seq_len.new_empty((req_num,)), - mem_indexes=mem_indexes.new_empty((req_num,)), b_shared_seq_len=b_shared_seq_len.new_empty((req_num,)), b_shared_radix_node_id=b_shared_radix_node_id.new_empty((req_num,)), b_position_delta=b_position_delta.new_empty((req_num,)), @@ -159,8 +151,6 @@ def select_accepted_tail_rows( b_mtp_index_stride=b_mtp_index.stride(0), b_seq_len=b_seq_len, b_seq_len_stride=b_seq_len.stride(0), - mem_indexes=mem_indexes, - mem_indexes_stride=mem_indexes.stride(0), b_shared_seq_len=b_shared_seq_len, b_shared_seq_len_stride=b_shared_seq_len.stride(0), b_shared_radix_node_id=b_shared_radix_node_id, @@ -174,7 +164,6 @@ def select_accepted_tail_rows( out_b_req_idx=selected.b_req_idx, out_b_mtp_index=selected.b_mtp_index, out_b_seq_len=selected.b_seq_len, - out_mem_indexes=selected.mem_indexes, out_b_shared_seq_len=selected.b_shared_seq_len, out_b_shared_radix_node_id=selected.b_shared_radix_node_id, out_b_position_delta=selected.b_position_delta, diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index ab0c76c149..e49e216fe4 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -103,7 +103,7 @@ def _decode(self, model_input: ModelInput) -> ModelOutput: infer_state.position_cos = torch.index_select(self._cos_cached, 0, position_ids) infer_state.position_sin = torch.index_select(self._sin_cached, 0, position_ids) infer_state.mem_manager = self.mem_manager - infer_state.mem_index = model_input.mem_indexes + infer_state.mem_index = self._select_mem_indexes(model_input) hidden = self.pre_infer.context_forward(None, infer_state, self.pre_post_weight) for layer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index dc6e1e579f..bb4d2d899c 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -405,6 +405,7 @@ def _capture_prompt_logprobs_if_needed( return mgr = PromptLogprobsCaptureManager.get_instance() + mem_indexes = self.model._select_mem_indexes(model_input) start_loc = 0 for req_obj in run_reqs: @@ -438,7 +439,7 @@ def _capture_prompt_logprobs_if_needed( top_token_ids = torch.nn.functional.pad(top_token_ids, padding, value=-1) top_logprobs = torch.nn.functional.pad(top_logprobs, padding, value=float("-inf")) mgr.capture( - mem_indexes=model_input.mem_indexes[start_loc : start_loc + capture_count], + mem_indexes=mem_indexes[start_loc : start_loc + capture_count], top_token_ids=top_token_ids, top_logprobs=top_logprobs, ) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 9322ec9d0d..d8f93872a7 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -14,7 +14,6 @@ from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState @@ -377,13 +376,6 @@ def decode_mtp( extra_post_req_handle_func=self.extra_post_req_handle_func, ) - if not model_input.mem_indexes_from_req_table: - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu == 0, - ), - ) mtp_utils.free_mem_indexes( backend=self, extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index f9693f6cc1..d4eb050874 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -19,7 +19,6 @@ from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from .control_state import DPControlState @@ -604,14 +603,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): verify_run_reqs=run_reqs, ) - if not model_input.mem_indexes_from_req_table: - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu == 0, - ), - ) - select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, @@ -890,21 +881,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf req_num=req_num, accept_lengths_cpu=mtp_accept_len_cpu, ) - if not model_input0.mem_indexes_from_req_table: - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input0.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu0 == 0, - ) - ) - if not model_input1.mem_indexes_from_req_table: - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input1.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu1 == 0, - ) - ) - select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index a9d2ff6042..7d3f34c4a7 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -72,7 +72,6 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> max_kv_seq_len=max_kv_seq_len, max_cache_len=max_cache_len, input_ids=input_ids, - mem_indexes_from_req_table=True, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -141,7 +140,6 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In max_q_seq_len=max_q_seq_len, max_kv_seq_len=max_kv_seq_len, input_ids=None, - mem_indexes_from_req_table=True, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 077d478fdc..f35f6e5065 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -98,7 +98,6 @@ def propose_next_overlap( b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, b_seq_len=model_input.b_seq_len, - mem_indexes=model_input.mem_indexes, b_shared_seq_len=model_input.b_shared_seq_len, b_shared_radix_node_id=model_input.b_shared_radix_node_id, b_position_delta=model_input.b_position_delta, @@ -109,11 +108,9 @@ def propose_next_overlap( draft_input.b_req_idx = selected_rows.b_req_idx draft_input.b_mtp_index = selected_rows.b_mtp_index draft_input.b_seq_len = selected_rows.b_seq_len - draft_input.mem_indexes = selected_rows.mem_indexes draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id draft_input.b_position_delta = selected_rows.b_position_delta - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_req_num)] draft_inputs.append(draft_input) draft_token_ids_by_batch.append(selected_rows.input_ids) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 6b8c23e8fd..0ab8ad9530 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -4,14 +4,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( get_dp_overlap_req_start_rows, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -185,21 +183,12 @@ def propose_next_overlap( empty_multimodal_params = {"images": [], "audios": []} model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) - for step in range(1, draft_step): - mem_start = (step - 1) * req_num - step_mem_indexes = extra_mem_indexes[mem_start : mem_start + req_num] - mem_offset = 0 for batch_index, model_input in enumerate(model_inputs): - batch_req_num = req_num_by_batch[batch_index] model_input.input_ids = draft_token_ids_by_batch[batch_index] model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - model_input.mem_indexes = step_mem_indexes[mem_offset : mem_offset + batch_req_num] model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len - mem_offset += batch_req_num draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): @@ -221,7 +210,7 @@ def propose_next_overlap( return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 4c3261aa26..d134206d7e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -98,7 +98,6 @@ def propose_next_overlap( b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, b_seq_len=model_input.b_seq_len, - mem_indexes=model_input.mem_indexes, b_shared_seq_len=model_input.b_shared_seq_len, b_shared_radix_node_id=model_input.b_shared_radix_node_id, b_position_delta=model_input.b_position_delta, @@ -109,11 +108,9 @@ def propose_next_overlap( draft_input.b_req_idx = selected_rows.b_req_idx draft_input.b_mtp_index = selected_rows.b_mtp_index draft_input.b_seq_len = selected_rows.b_seq_len - draft_input.mem_indexes = selected_rows.mem_indexes draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id draft_input.b_position_delta = selected_rows.b_position_delta - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_req_num)] draft_inputs.append(draft_input) draft_token_ids_by_batch.append(selected_rows.input_ids) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 7e3f55f488..83a73c7d9e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -81,17 +81,6 @@ def prepare_decode_model_input( return model_input, None from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import prepare_dynamic_mtp_model_input - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - if not model_input.mem_indexes_from_req_table: - # 兼容外部构造的旧式 ModelInput:尚未绑定请求位置的临时 KV slot - # 可以随动态 verify 行一起裁剪并立即释放。 - unused_mem_indexes_cpu = model_input.mem_indexes_cpu[plan.dynamic_batch_size :] - model_input.mem_indexes_cpu = model_input.mem_indexes_cpu[: plan.dynamic_batch_size] - if model_input.mem_indexes is not None: - model_input.mem_indexes = model_input.mem_indexes[: plan.dynamic_batch_size] - if unused_mem_indexes_cpu.numel() > 0: - g_infer_context.req_manager.mem_manager.free(unused_mem_indexes_cpu) model_input, selected_row_mask = prepare_dynamic_mtp_model_input( model_input=model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 23f81f546a..fc2317167b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -5,11 +5,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - BaseSpecProposer, - MtpMemIndexesToFree, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( DFlashSpecProposal, ) @@ -90,9 +86,7 @@ def propose_next( draft_model.forward(verify_draft_input) # 每个请求始终展开完整 block,未被本轮 proposal 返回的 block 尾部仍会 - # 参与 parallel forward。所有临时 KV slot 在 verify 后通过 proposal - # 统一释放。 - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * block_size) + # 参与 parallel forward;KV 位置由请求页表和序列位置确定。 block_input_ids = target_next_token_ids.new_full( (req_num * block_size,), fill_value=draft_model.mask_token_id, @@ -140,8 +134,6 @@ def propose_next( .repeat_interleave(block_size) .contiguous() ) - draft_input.mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] draft_output = draft_model.forward(draft_input) @@ -160,6 +152,6 @@ def propose_next( schedule_scores = block_draft_token_probs[:, :draft_step].float().contiguous() return DFlashSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 5e3ea5694f..03cd6862c5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -5,12 +5,8 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - BaseSpecProposer, - MtpMemIndexesToFree, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( DSparkSpecProposal, ) @@ -92,10 +88,9 @@ def propose_next( verify_draft_input.mtp_draft_input_hiddens = target_model_output.mtp_collector.spec_hidden draft_model.forward(verify_draft_input) - # DSpark 每个请求固定展开一个完整 block,临时 KV 在 target verify 完成 - # 后通过 proposal 统一释放。block 第一行是 accepted-tail anchor,其余行 - # 使用 mask token,由 parallel backbone 一次并行计算。 - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * block_size) + # DSpark 每个请求固定展开一个完整 block。block 第一行是 + # accepted-tail anchor,其余行使用 mask token,由 parallel backbone + # 一次并行计算;KV 位置由请求页表和序列位置确定。 block_input_ids = target_next_token_ids.new_full( (req_num * block_size,), fill_value=draft_model.mask_token_id, @@ -143,8 +138,6 @@ def propose_next( .repeat_interleave(block_size) .contiguous() ) - draft_input.mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] draft_output = draft_model.forward(draft_input) @@ -184,7 +177,7 @@ def propose_next( return DSparkSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, schedule_scores_cpu=schedule_scores_cpu, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index 629eb7fd3e..2c0942fa11 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -62,7 +62,6 @@ def propose_next( b_req_idx=target_model_input.b_req_idx, b_mtp_index=target_model_input.b_mtp_index, b_seq_len=target_model_input.b_seq_len, - mem_indexes=target_model_input.mem_indexes, b_shared_seq_len=target_model_input.b_shared_seq_len, b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, b_position_delta=target_model_input.b_position_delta, @@ -75,11 +74,9 @@ def propose_next( draft_input.b_req_idx = selected_rows.b_req_idx draft_input.b_mtp_index = selected_rows.b_mtp_index draft_input.b_seq_len = selected_rows.b_seq_len - draft_input.mem_indexes = selected_rows.mem_indexes draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id draft_input.b_position_delta = selected_rows.b_position_delta - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] # EAGLE 使用同一个 draft model 递归生成多个 token。每一级将上一级 diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 3d2c0a0e86..518f08e0ce 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -5,11 +5,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - BaseSpecProposer, - MtpMemIndexesToFree, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -96,11 +92,6 @@ def propose_next( schedule_scores=torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None, ) - # 后续递归每步、每请求各写一个临时 KV。proposal 在 verify 完成后 - # 统一释放这些 slot,因此同时保留 CPU 索引用于资源回收。 - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) - # 一次 Triton kernel 合并抽取 accepted-tail 的 hidden、请求索引、 # 序列长度、position delta 和共享 radix 元数据,构造 req_num 行的 # 单 token decode 输入。通用算子同时返回 accepted-tail input ids, @@ -114,7 +105,6 @@ def propose_next( b_req_idx=target_model_input.b_req_idx, b_mtp_index=target_model_input.b_mtp_index, b_seq_len=target_model_input.b_seq_len, - mem_indexes=target_model_input.mem_indexes, b_shared_seq_len=target_model_input.b_shared_seq_len, b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, b_position_delta=target_model_input.b_position_delta, @@ -136,14 +126,11 @@ def propose_next( draft_input.b_position_delta = selected_rows.b_position_delta draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] for step in range(1, draft_step): - mem_start = (step - 1) * req_num draft_input.input_ids = draft_token_ids draft_input.mtp_draft_input_hiddens = draft_hidden - draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + req_num] draft_input.max_kv_seq_len = max_kv_seq_len + step draft_input.total_token_num = req_num * draft_input.max_kv_seq_len draft_output = draft_model.forward(draft_input) @@ -161,7 +148,7 @@ def propose_next( schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index ea6a2d5e11..44973679e6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -56,7 +56,6 @@ def propose_next( b_req_idx=target_model_input.b_req_idx, b_mtp_index=target_model_input.b_mtp_index, b_seq_len=target_model_input.b_seq_len, - mem_indexes=target_model_input.mem_indexes, b_shared_seq_len=target_model_input.b_shared_seq_len, b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, b_position_delta=target_model_input.b_position_delta, @@ -69,11 +68,9 @@ def propose_next( draft_input.b_req_idx = selected_rows.b_req_idx draft_input.b_mtp_index = selected_rows.b_mtp_index draft_input.b_seq_len = selected_rows.b_seq_len - draft_input.mem_indexes = selected_rows.mem_indexes draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id draft_input.b_position_delta = selected_rows.b_position_delta - draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] for step in range(draft_step): diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 9fd44487c6..105bca2c7f 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -17,7 +17,6 @@ def _create_model_input(*, is_prefill=False): b_req_idx=torch.arange(batch_size, dtype=torch.int32), b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), b_seq_len=torch.ones(batch_size, dtype=torch.int32), - mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), is_prefill=is_prefill, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], ) @@ -46,7 +45,6 @@ def test_decode_requires_shared_radix_metadata(): b_mtp_index=torch.zeros(1, dtype=torch.int32), b_seq_len=torch.ones(1, dtype=torch.int32), b_position_delta=torch.zeros(1, dtype=torch.int32), - mem_indexes_cpu=torch.zeros(1, dtype=torch.int32), is_prefill=False, multimodal_params=[{"images": [], "audios": []}], ) @@ -120,7 +118,6 @@ def test_padded_prefill_adds_non_decode_request_marker(): b_is_decode_req=torch.ones(1, dtype=torch.bool), b_ready_cache_len=torch.zeros(1, dtype=torch.int32), b_prefill_start_loc=torch.zeros(1, dtype=torch.int32), - mem_indexes=torch.arange(2, dtype=torch.int32), is_prefill=True, b_prefill_has_output_cpu=[False], multimodal_params=[{"images": [], "audios": []}], @@ -154,7 +151,6 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): b_is_decode_req=torch.empty((0,), dtype=torch.bool), b_ready_cache_len=torch.empty((0,), dtype=torch.int32), b_prefill_start_loc=torch.empty((0,), dtype=torch.int32), - mem_indexes=torch.empty((0,), dtype=torch.int32), is_prefill=True, b_prefill_has_output_cpu=[], multimodal_params=[], @@ -174,7 +170,6 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): assert model_input.batch_size == 0 assert padded_input.batch_size == 1 assert padded_input.input_ids.tolist() == [1] - assert padded_input.mem_indexes.tolist() == [99] assert padded_input.b_req_idx.tolist() == [88] assert padded_input.b_seq_len.tolist() == [1] assert padded_input.b_prefill_has_output_cpu == [False] @@ -194,7 +189,6 @@ def test_padded_decode_builds_internal_request_from_empty_token_tensor(): b_position_delta=torch.empty((0,), dtype=torch.int32), b_shared_seq_len=torch.empty((0,), dtype=torch.int32), b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), - mem_indexes=torch.empty((0,), dtype=torch.int32), is_prefill=False, multimodal_params=[], ) @@ -212,7 +206,6 @@ def test_padded_decode_builds_internal_request_from_empty_token_tensor(): assert model_input.batch_size == 0 assert padded_input.batch_size == 1 assert padded_input.input_ids.tolist() == [1] - assert padded_input.mem_indexes.tolist() == [99] assert padded_input.b_req_idx.tolist() == [88] assert padded_input.b_seq_len.tolist() == [2] assert padded_input.b_shared_seq_len.tolist() == [0] diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index c604f70663..ab250c2217 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -88,7 +88,6 @@ def _create_empty_decode_input(): b_position_delta=torch.empty((0,), dtype=torch.int32), b_shared_seq_len=torch.empty((0,), dtype=torch.int32), b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), - mem_indexes=torch.empty((0,), dtype=torch.int32), is_prefill=False, multimodal_params=[], ) @@ -96,8 +95,6 @@ def _create_empty_decode_input(): @torch.no_grad() def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): - monkeypatch.setattr(basemodel, "copy_kv_index_to_req", lambda *args: None) - execution_configs = ( # eager 普通模式:空 batch 只补一个 dummy request。 (None, False, 1, False, 1), @@ -135,7 +132,7 @@ def create_infer_state(model_input): infer_state = SimpleNamespace( b_req_idx=model_input.b_req_idx, b_seq_len=model_input.b_seq_len, - mem_index=model_input.mem_indexes, + mem_index=torch.empty((model_input.batch_size,), dtype=torch.int32), init_some_extra_state=lambda _: None, is_cuda_graph=False, ) diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py index 5e93f5960e..e34bc34b9c 100644 --- a/unit_tests/common/basemodel/test_overlap_utils.py +++ b/unit_tests/common/basemodel/test_overlap_utils.py @@ -27,7 +27,6 @@ def _make_prefill_input(input_ids: list[int], req_idx: int, is_decode_req: bool) b_is_decode_req=torch.tensor([is_decode_req]), b_ready_cache_len=torch.zeros(1, dtype=torch.int32), b_prefill_start_loc=torch.zeros(1, dtype=torch.int32), - mem_indexes_cpu=torch.arange(token_num, dtype=torch.int32), is_prefill=True, b_prefill_has_output_cpu=[True], multimodal_params=_empty_multimodal_params(1), @@ -48,7 +47,6 @@ def _make_decode_input() -> ModelInput: b_position_delta=torch.arange(batch_size, dtype=torch.int32), b_shared_seq_len=torch.arange(10, 16, dtype=torch.int32), b_shared_radix_node_id=torch.arange(20, 26, dtype=torch.int64), - mem_indexes_cpu=torch.arange(100, 106, dtype=torch.int32), is_prefill=False, multimodal_params=_empty_multimodal_params(batch_size), ) @@ -67,7 +65,6 @@ def _make_empty_decode_input() -> ModelInput: b_position_delta=torch.empty((0,), dtype=torch.int32), b_shared_seq_len=torch.empty((0,), dtype=torch.int32), b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), - mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), is_prefill=False, multimodal_params=[], ) @@ -175,7 +172,6 @@ def test_overlap_decode_cuda_pads_empty_side_and_unpads_outputs(monkeypatch): model_input0.b_position_delta = model_input0.b_position_delta[:1] model_input0.b_shared_seq_len = model_input0.b_shared_seq_len[:1] model_input0.b_shared_radix_node_id = model_input0.b_shared_radix_node_id[:1] - model_input0.mem_indexes_cpu = model_input0.mem_indexes_cpu[:1] model_input0.multimodal_params = model_input0.multimodal_params[:1] model_input0.check_input() model_input1 = _make_empty_decode_input() @@ -195,7 +191,7 @@ def fake_create_inferstate(model_input, microbatch_index): return SimpleNamespace( b_req_idx=model_input.b_req_idx, b_seq_len=model_input.b_seq_len, - mem_index=model_input.mem_indexes, + mem_index=torch.empty((model_input.batch_size,), dtype=torch.int32, device="cuda"), init_some_extra_state=lambda _: None, init_att_state=lambda: None, ) @@ -205,8 +201,6 @@ def fake_create_inferstate(model_input, microbatch_index): ModelOutput(logits=torch.zeros((infer_state0.b_req_idx.shape[0], 1), device="cuda")), ModelOutput(logits=torch.zeros((infer_state1.b_req_idx.shape[0], 1), device="cuda")), ) - monkeypatch.setattr(basemodel, "copy_kv_index_to_req", lambda *args, **kwargs: None) - output0, output1 = model._microbatch_overlap_decode_cuda(model_input0, model_input1) assert infer_batch_sizes == [2, 2] diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index 2cb95f5923..cd136ae713 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -53,8 +53,6 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): b_shared_radix_node_id=torch.tensor( [10, 10, 10, 10, 11, 11, 11, 11, 12, 12, 12, 12], dtype=torch.int64, device="cuda" ), - mem_indexes=torch.arange(8, dtype=torch.int32, device="cuda") + 100, - mem_indexes_cpu=torch.arange(8, dtype=torch.int32, device="cpu") + 100, is_prefill=False, multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), @@ -102,13 +100,6 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): compacted_input.b_shared_radix_node_id.cpu(), torch.tensor([10, 10, 10, 11, 12, 12, 12, 12], dtype=torch.int64), ) - assert torch.equal( - compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) - ) - # CPU/GPU mem indexes are trimmed by SpecEngine.prepare_decode_model_input - # before entering this lower-level row compaction helper. - assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(8, dtype=torch.int32) + 100) - expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) @@ -126,8 +117,6 @@ def test_compaction_preserves_shared_radix_metadata(): b_position_delta=torch.zeros(5, dtype=torch.int32, device="cuda"), b_shared_seq_len=torch.full((5,), 7, dtype=torch.int32, device="cuda"), b_shared_radix_node_id=torch.full((5,), 10, dtype=torch.int64, device="cuda"), - mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), - mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(5)], ) diff --git a/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py b/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py index 2b0affcc4b..0dab062934 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py +++ b/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py @@ -66,12 +66,10 @@ def test_model_selects_reserved_indexes_when_execution_starts(): ) model_input = SimpleNamespace( is_prefill=False, - mem_indexes_from_req_table=True, b_req_idx=torch.tensor([1, 0], dtype=torch.int32, device="cuda"), b_seq_len=torch.tensor([3, 4], dtype=torch.int32, device="cuda"), - mem_indexes=None, ) - model._select_page_mem_indexes(model_input) + mem_indexes = model._select_mem_indexes(model_input) - assert model_input.mem_indexes.cpu().tolist() == [22, 13] + assert mem_indexes.cpu().tolist() == [22, 13] diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py index 50097d40f9..c87ecfcf34 100644 --- a/unit_tests/common/test_req_manager_page.py +++ b/unit_tests/common/test_req_manager_page.py @@ -99,7 +99,6 @@ def test_decode_reserves_mtp_headroom(monkeypatch): assert run_reqs == [req, req, req] assert model_input.b_seq_len.tolist() == [4, 5, 6] - assert model_input.mem_indexes_cpu is None assert req.hold_kv_len == 12 assert context.req_manager.mem_manager.alloc_sizes == [8] assert context.req_manager.req_to_token_indexs[0, :12].tolist() == list(range(12)) @@ -126,8 +125,6 @@ def test_page_size_one_uses_the_same_scheduler_preallocation(monkeypatch): model_input, _ = generic_pre_process.prepare_decode_inputs([req]) - assert model_input.mem_indexes_cpu is None - assert model_input.mem_indexes_from_req_table is True assert req.hold_kv_len == 12 assert context.req_manager.mem_manager.alloc_sizes == [9] assert context.req_manager.req_to_token_indexs[0, :12].tolist() == list(range(12)) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 94724940d3..73dee9952d 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -53,6 +53,7 @@ def test_parallel_block_decode_commits_target_hiddens_directly(model_class): target_hiddens = torch.arange(6, dtype=torch.float32).view(2, 3) mem_indexes = torch.tensor([7, 11]) + model._select_mem_indexes = lambda _: mem_indexes observed_states = [] class PreInfer: @@ -78,7 +79,6 @@ def context_forward(self, hidden, infer_state, layer_weight): model_input = SimpleNamespace( batch_size=2, b_seq_len=torch.tensor([3, 5]), - mem_indexes=mem_indexes, mtp_draft_input_hiddens=target_hiddens, ) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index fcdf4e3295..5a341f9a09 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -238,7 +238,6 @@ def test_dp_decode_mtp_runs_common_engine_for_empty_batch(monkeypatch): batch_size=0, b_req_idx=empty_i32, b_mtp_index=empty_i32, - mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), ) model_output = SimpleNamespace(logits=torch.empty((0, 8), device=device)) calls = [] diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 05745016cf..49931a8628 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -60,7 +60,6 @@ def test_prepare_prefill_inputs_allows_empty_batch(monkeypatch): assert run_reqs == [] assert model_input.batch_size == 0 assert model_input.input_ids.shape == (0,) - assert model_input.mem_indexes_cpu is None assert model_input.b_req_idx.shape == (0,) assert model_input.b_prefill_start_loc.shape == (0,) assert model_input.b_prefill_has_output_cpu == [] @@ -76,7 +75,6 @@ def test_prepare_decode_inputs_allows_empty_batch(monkeypatch): assert run_reqs == [] assert model_input.batch_size == 0 assert model_input.input_ids is None - assert model_input.mem_indexes_cpu is None assert model_input.b_req_idx.shape == (0,) assert model_input.b_position_delta.shape == (0,) assert model_input.b_shared_seq_len.shape == (0,) @@ -176,4 +174,3 @@ def test_overlap_decode_preserves_empty_microbatch(monkeypatch): assert run_reqs1 == [] assert model_input1.batch_size == 0 assert model_input1.b_req_idx.shape == (0,) - assert model_input1.mem_indexes_cpu is None diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py b/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py index ac3a303075..0bb24b462f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py @@ -2,7 +2,6 @@ import torch -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import DFlashSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -60,8 +59,6 @@ def forward(model_input): ), enable_dynmaic_mtp=True, ) - extra_mem_indexes_cpu = torch.arange(100, 106, dtype=torch.int32) - monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_num: extra_mem_indexes_cpu) monkeypatch.setattr(torch.Tensor, "cuda", lambda self, non_blocking=False: self) monkeypatch.setattr( g_pin_mem_manager, @@ -83,8 +80,6 @@ def forward(model_input): b_position_delta=torch.tensor([1, 1, 1, 2, 2], dtype=torch.int32), b_shared_seq_len=torch.tensor([3, 3, 3, 6, 6], dtype=torch.int32), b_shared_radix_node_id=torch.tensor([17, 17, 17, 19, 19], dtype=torch.int64), - mem_indexes=torch.arange(5, dtype=torch.int32), - mem_indexes_cpu=torch.arange(5, dtype=torch.int32), multimodal_params=[{"images": [], "audios": []} for _ in range(5)], mtp_draft_input_hiddens=None, ) @@ -113,13 +108,10 @@ def forward(model_input): assert torch.equal(block_draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) assert torch.equal(block_draft_input.b_position_delta, torch.tensor([1, 1, 1, 2, 2, 2], dtype=torch.int32)) assert block_draft_input.mtp_draft_input_hiddens is None - assert block_draft_input.mem_indexes is extra_mem_indexes_cpu - assert block_draft_input.mem_indexes_cpu is None assert torch.equal(proposal.token_ids, torch.tensor([[30, 31], [40, 41]], dtype=torch.int64)) torch.testing.assert_close( proposal.schedule_scores, flat_draft_token_probs.reshape(2, block_size)[:, :2].float(), ) - assert len(proposal.extra_mem_indexes_cpu) == 1 - assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu + assert proposal.extra_mem_indexes_cpu == [] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py index d341dc7dbd..64fb40f677 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py @@ -2,7 +2,6 @@ import torch -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -67,8 +66,6 @@ def forward(model_input): backend=SimpleNamespace(draft_models=[draft_model]), enable_dynmaic_mtp=True, ) - extra_mem_indexes_cpu = torch.arange(100, 106, dtype=torch.int32) - monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_num: extra_mem_indexes_cpu) monkeypatch.setattr(torch.Tensor, "cuda", lambda self, non_blocking=False: self) monkeypatch.setattr( g_pin_mem_manager, @@ -95,8 +92,6 @@ def forward(model_input): b_position_delta=torch.tensor([1, 1, 1, 2, 2], dtype=torch.int32), b_shared_seq_len=torch.tensor([3, 3, 3, 6, 6], dtype=torch.int32), b_shared_radix_node_id=torch.tensor([17, 17, 17, 19, 19], dtype=torch.int64), - mem_indexes=torch.arange(5, dtype=torch.int32), - mem_indexes_cpu=torch.arange(5, dtype=torch.int32), multimodal_params=[{"images": [], "audios": []} for _ in range(5)], mtp_draft_input_hiddens=None, ) @@ -124,8 +119,6 @@ def forward(model_input): assert torch.equal(block_draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) assert torch.equal(block_draft_input.b_position_delta, torch.tensor([1, 1, 1, 2, 2, 2], dtype=torch.int32)) assert block_draft_input.mtp_draft_input_hiddens is None - assert block_draft_input.mem_indexes is extra_mem_indexes_cpu - assert block_draft_input.mem_indexes_cpu is None assert torch.equal(proposal.token_ids, torch.tensor([[30, 31], [40, 41]], dtype=torch.int64)) torch.testing.assert_close( @@ -134,5 +127,4 @@ def forward(model_input): ) assert torch.equal(proposal.schedule_scores_cpu, proposal.schedule_scores) assert proposal.schedule_scores_cpu is not proposal.schedule_scores - assert len(proposal.extra_mem_indexes_cpu) == 1 - assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu + assert proposal.extra_mem_indexes_cpu == [] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py index 9cb533e4e8..fcea551c0d 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py @@ -33,7 +33,6 @@ def forward(model_input): "draft_hidden": model_input.mtp_draft_input_hiddens.cpu(), "b_req_idx": model_input.b_req_idx.cpu(), "b_seq_len": model_input.b_seq_len.cpu(), - "mem_indexes": model_input.mem_indexes.cpu(), } ) return draft_outputs[len(draft_calls) - 1] @@ -53,8 +52,6 @@ def forward(model_input): b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32, device=device), b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32, device=device), b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32, device=device), - mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32, device=device), - mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device=device), b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32, device=device), b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64, device=device), @@ -82,7 +79,6 @@ def forward(model_input): torch.testing.assert_close(draft_calls[0]["draft_hidden"], torch.tensor([[2.0, 3.0], [8.0, 9.0]])) torch.testing.assert_close(draft_calls[0]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32)) torch.testing.assert_close(draft_calls[0]["b_seq_len"], torch.tensor([11, 21], dtype=torch.int32)) - torch.testing.assert_close(draft_calls[0]["mem_indexes"], torch.tensor([101, 104], dtype=torch.int32)) torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([21, 24])) torch.testing.assert_close( draft_calls[1]["draft_hidden"], diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index c046c63c51..546d45b1e8 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -4,7 +4,6 @@ import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import eagle_with_att from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( get_dp_overlap_req_start_rows, @@ -78,8 +77,6 @@ def _target_input(batch_size, b_mtp_index=None, device="cpu"): b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device=device), b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32, device=device), b_shared_radix_node_id=torch.arange(batch_size, dtype=torch.int64, device=device), - mem_indexes=torch.arange(batch_size, dtype=torch.int32, device=device), - mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), max_kv_seq_len=16, max_cache_len=16, is_prefill=False, @@ -118,11 +115,6 @@ def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - monkeypatch.setattr( - mtp_utils, - "alloc_mem_indexes", - lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) model_input0 = _target_input(batch_size=3) model_input1 = _target_input(batch_size=6) @@ -148,17 +140,10 @@ def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): assert draft_model.decode_batch_sizes == [(3, 6), (1, 2)] assert proposal.token_ids.shape == (3, 2) assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) - assert len(proposal.extra_mem_indexes_cpu) == 1 - assert torch.equal( - proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, - torch.arange(3, dtype=torch.int32), - ) - assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None - assert torch.equal(model_input0.mem_indexes, torch.arange(3, dtype=torch.int32)) - assert torch.equal(model_input1.mem_indexes, torch.arange(6, dtype=torch.int32)) + assert proposal.extra_mem_indexes_cpu == [] -def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): +def test_overlap_eagle_supports_empty_verify_rows(): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -171,11 +156,6 @@ def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - monkeypatch.setattr( - mtp_utils, - "alloc_mem_indexes", - lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) proposal = proposer.propose_next_overlap( target_model_input0=_target_input(batch_size=0), @@ -216,11 +196,6 @@ def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): ), ) proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=True) - monkeypatch.setattr( - mtp_utils, - "alloc_mem_indexes", - lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) proposal = proposer.propose_next_overlap( target_model_input0=_target_input( @@ -320,11 +295,6 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapEagle3Proposer(backend=backend, enable_dynmaic_mtp=False) - monkeypatch.setattr( - mtp_utils, - "alloc_mem_indexes", - lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) model_input0 = _target_input(batch_size=3) model_input1 = _target_input(batch_size=6) @@ -354,12 +324,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): assert draft_model.decode_batch_sizes == [(3, 6), (1, 2)] assert proposal.token_ids.shape == (3, 2) assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) - assert len(proposal.extra_mem_indexes_cpu) == 1 - assert torch.equal( - proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, - torch.arange(3, dtype=torch.int32), - ) - assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None + assert proposal.extra_mem_indexes_cpu == [] def test_eagle3_maps_draft_token_ids_in_proposer(): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py index f50780d53c..c8b0db5af6 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py @@ -4,7 +4,6 @@ import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal @@ -144,7 +143,6 @@ def forward(model_input): "b_req_idx": model_input.b_req_idx.clone(), "b_mtp_index": model_input.b_mtp_index.clone(), "b_seq_len": model_input.b_seq_len.clone(), - "mem_indexes": model_input.mem_indexes.clone(), "max_kv_seq_len": model_input.max_kv_seq_len, "total_token_num": model_input.total_token_num, } @@ -177,17 +175,12 @@ def forward(model_input): b_req_idx=torch.tensor([7, 7, 7, 9, 9, 9], dtype=torch.int32, device=device), b_mtp_index=torch.tensor([0, 1, 2, 0, 1, 2], dtype=torch.int32, device=device), b_seq_len=torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int32, device=device), - mem_indexes=torch.tensor([100, 101, 102, 103, 104, 105], dtype=torch.int32, device=device), - mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104, 105], dtype=torch.int32), b_position_delta=torch.tensor([0, 1, 2, 3, 4, 5], dtype=torch.int32, device=device), b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6, 6], dtype=torch.int32, device=device), b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90, 90], dtype=torch.int64, device=device), multimodal_params=[{"images": [], "audios": []} for _ in range(6)], mtp_draft_input_hiddens=None, ) - extra_mem_indexes_cpu = torch.tensor([200, 201], dtype=torch.int32) - monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_count: extra_mem_indexes_cpu) - proposal = proposer.propose_next( target_model_input=target_input, target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), @@ -200,14 +193,12 @@ def forward(model_input): assert isinstance(proposal, EagleSpecProposal) torch.testing.assert_close(proposal.token_ids, torch.tensor([[32, 40], [34, 41]], device=device)) torch.testing.assert_close(proposal.schedule_scores, torch.tensor([[0.32, 0.40], [0.34, 0.41]], device=device)) - assert len(proposal.extra_mem_indexes_cpu) == 1 - assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu + assert proposal.extra_mem_indexes_cpu == [] assert len(draft_calls) == 2 assert draft_calls[0]["model_input"] is not target_input assert draft_calls[0]["batch_size"] == 6 torch.testing.assert_close(draft_calls[0]["input_ids"], original_input_ids) torch.testing.assert_close(draft_calls[0]["draft_hidden"], target_hidden) - torch.testing.assert_close(draft_calls[0]["mem_indexes"], target_input.mem_indexes) assert draft_calls[1]["batch_size"] == 2 torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([32, 34], device=device)) torch.testing.assert_close( @@ -217,7 +208,6 @@ def forward(model_input): torch.testing.assert_close(draft_calls[1]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32, device=device)) torch.testing.assert_close(draft_calls[1]["b_mtp_index"], torch.zeros(2, dtype=torch.int32, device=device)) torch.testing.assert_close(draft_calls[1]["b_seq_len"], torch.tensor([13, 22], dtype=torch.int32, device=device)) - torch.testing.assert_close(draft_calls[1]["mem_indexes"], extra_mem_indexes_cpu.to(device)) assert draft_calls[1]["max_kv_seq_len"] == 23 assert draft_calls[1]["total_token_num"] == 46 assert target_input.input_ids is original_input_ids diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 1977606198..f029f16229 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -664,21 +664,12 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): assert ready_plan.all_reqs_have_proposals -def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypatch): +def test_dynamic_prepare_delegates_row_compaction(monkeypatch): from lightllm.common.basemodel.triton_kernel import dynamic_mtp_utils - from lightllm.server.router.model_infer import infer_batch as infer_batch_module from lightllm.server.router.model_infer.mtp_speculative import ( engine as engine_module, ) - freed = [] - monkeypatch.setattr( - infer_batch_module, - "g_infer_context", - SimpleNamespace( - req_manager=SimpleNamespace(mem_manager=SimpleNamespace(free=lambda indexes: freed.append(indexes.clone()))) - ), - ) selected_row_mask = torch.tensor([1, 0, 1, 0], dtype=torch.int32) monkeypatch.setattr( dynamic_mtp_utils, @@ -702,9 +693,6 @@ def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypa ) model_input = SimpleNamespace( batch_size=4, - mem_indexes=torch.tensor([20, 21, 22, 23], dtype=torch.int32), - mem_indexes_cpu=torch.tensor([10, 11, 12, 13], dtype=torch.int32), - mem_indexes_from_req_table=False, ) plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=1, pre_draft_step=1) @@ -716,10 +704,6 @@ def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypa assert compacted_input is model_input assert selected_mask_cpu is async_mask - assert model_input.mem_indexes.tolist() == [20, 21] - assert model_input.mem_indexes_cpu.tolist() == [10, 11] - assert len(freed) == 1 - assert freed[0].tolist() == [12, 13] def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py index 8ce9114aab..bedf17a5b0 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -24,7 +24,6 @@ def test_select_accepted_tail_rows_triton_matches_index_select(): b_req_idx = (torch.arange(16, dtype=torch.int32, device=device) + 10)[::2] b_mtp_index = torch.arange(8, dtype=torch.int32, device=device) b_seq_len = (torch.arange(16, dtype=torch.int32, device=device) + 20)[::2] - mem_indexes = (torch.arange(16, dtype=torch.int32, device=device) + 100)[::2] b_shared_seq_len = (torch.arange(16, dtype=torch.int32, device=device) + 30)[::2] b_shared_radix_node_id = (torch.arange(16, dtype=torch.int64, device=device) + 1000)[::2] b_position_delta = (torch.arange(16, dtype=torch.int32, device=device) - 8)[::2] @@ -37,7 +36,6 @@ def test_select_accepted_tail_rows_triton_matches_index_select(): b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, - mem_indexes=mem_indexes, b_shared_seq_len=b_shared_seq_len, b_shared_radix_node_id=b_shared_radix_node_id, b_position_delta=b_position_delta, @@ -48,7 +46,6 @@ def test_select_accepted_tail_rows_triton_matches_index_select(): torch.testing.assert_close(selected.b_req_idx, b_req_idx.index_select(0, expected_rows)) torch.testing.assert_close(selected.b_mtp_index, b_mtp_index.index_select(0, expected_rows)) torch.testing.assert_close(selected.b_seq_len, b_seq_len.index_select(0, expected_rows)) - torch.testing.assert_close(selected.mem_indexes, mem_indexes.index_select(0, expected_rows)) torch.testing.assert_close(selected.b_shared_seq_len, b_shared_seq_len.index_select(0, expected_rows)) torch.testing.assert_close( selected.b_shared_radix_node_id, @@ -82,8 +79,6 @@ def draft_forward(model_input): b_req_idx=empty_i32, b_mtp_index=empty_i32, b_seq_len=empty_i32, - mem_indexes=empty_i32, - mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), b_position_delta=empty_i32, b_shared_seq_len=empty_i32, b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64, device=device), @@ -135,7 +130,6 @@ def forward(model_input): "b_req_idx": model_input.b_req_idx.cpu(), "b_mtp_index": model_input.b_mtp_index.cpu(), "b_seq_len": model_input.b_seq_len.cpu(), - "mem_indexes": model_input.mem_indexes.cpu(), } ) return draft_outputs[step] @@ -154,8 +148,6 @@ def forward(model_input): b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32, device=device), b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32, device=device), b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32, device=device), - mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32, device=device), - mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device=device), b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32, device=device), b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64, device=device), @@ -183,7 +175,6 @@ def forward(model_input): torch.testing.assert_close(draft_calls[0]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32)) torch.testing.assert_close(draft_calls[0]["b_mtp_index"], torch.tensor([1, 1], dtype=torch.int32)) torch.testing.assert_close(draft_calls[0]["b_seq_len"], torch.tensor([11, 21], dtype=torch.int32)) - torch.testing.assert_close(draft_calls[0]["mem_indexes"], torch.tensor([101, 104], dtype=torch.int32)) torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([21, 24])) torch.testing.assert_close( draft_calls[1]["draft_hidden"], From dc1364d51ba025d1f237189caf7d4f0f7c50263f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 03:26:35 +0000 Subject: [PATCH 04/15] fix: assert prompt cache page alignment --- lightllm/server/router/model_infer/infer_batch.py | 1 + 1 file changed, 1 insertion(+) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 2a7a6039c7..8bccdf12b6 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -648,6 +648,7 @@ def _match_radix_cache(self): g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.hold_kv_len = self.cur_kv_len + assert self.hold_kv_len % self.args.page_size == 0 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 self.shm_req.shm_cur_kv_len = self.cur_kv_len From dc74842ebca3766fcf215a77843d675b26385d82 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 05:07:47 +0000 Subject: [PATCH 05/15] refactor: align KV allocation needs to pages --- .../server/router/model_infer/infer_batch.py | 29 ++++----------- .../model_infer/mode_backend/base_backend.py | 22 +++++++---- unit_tests/common/test_req_manager_page.py | 37 ++++++++++++++----- 3 files changed, 49 insertions(+), 39 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 8bccdf12b6..ce10d04980 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -577,10 +577,6 @@ def __init__( # mtp_step 用来记录一个请求 draft模型每步需要生成的token数量 # 正常模式下,这个值为0,在 mtp 模式下,这个值为 draft 模型每步需要生成的token数量 self.mtp_step: int = get_env_start_args().mtp_step - if self.mtp_step > 0: - self.decode_need_token_num = self._mtp_decode_need_token_num - else: - self.decode_need_token_num = self._normal_decode_need_token_num if g_infer_context.is_linear_att_mixed_model: self.get_chuncked_input_token_len = self.get_chuncked_input_token_len_for_linear_att @@ -914,32 +910,21 @@ def _stop_sequences_matched(self, output_len: int): def prefill_need_token_num(self, is_chuncked_prefill: bool): if is_chuncked_prefill: - input_token_ids = self.get_chuncked_input_token_ids() + target_kv_len = self.get_chuncked_input_token_len() else: - input_token_ids = self.get_input_token_ids() - - return len(input_token_ids) - self.cur_kv_len - - def prefill_kv_alloc_need(self, is_chuncked_prefill: bool) -> int: - if is_chuncked_prefill: - target_kv_len = len(self.get_chuncked_input_token_ids()) - else: - target_kv_len = len(self.get_input_token_ids()) + target_kv_len = self.get_cur_total_len() return self._kv_cache_alloc_need(target_kv_len) def decode_need_token_num(self) -> int: - raise NotImplementedError("error") - - def _normal_decode_need_token_num(self) -> int: - return self._kv_cache_alloc_need(self.cur_kv_len + 1) + decode_token_num = 1 if self.mtp_step == 0 else 3 * (1 + self.mtp_step) + return self._kv_cache_alloc_need(self.cur_kv_len + decode_token_num) def _kv_cache_alloc_need(self, target_kv_len: int) -> int: page_size = self.args.page_size target_hold_len = (target_kv_len + page_size - 1) // page_size * page_size - return max(target_hold_len - self.hold_kv_len, 0) - - def _mtp_decode_need_token_num(self) -> int: - return self._kv_cache_alloc_need(self.cur_kv_len + 3 * (1 + self.mtp_step)) + alloc_token_num = max(target_hold_len - self.hold_kv_len, 0) + assert alloc_token_num % page_size == 0 + return alloc_token_num class InferReqUpdatePack: diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index bb4d2d899c..8a28f44652 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -409,7 +409,10 @@ def _capture_prompt_logprobs_if_needed( start_loc = 0 for req_obj in run_reqs: - q_len = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill) + if self.disable_chunked_prefill: + q_len = req_obj.get_cur_total_len() - req_obj.cur_kv_len + else: + q_len = req_obj.get_chuncked_input_token_len() - req_obj.cur_kv_len topk = req_obj.sampling_param.shm_param.prompt_logprobs capture_count = min(q_len, req_obj.shm_req.input_len - req_obj.cur_kv_len - 1) if capture_count > 0 and topk == 0 and self.is_master_in_dp: @@ -736,11 +739,11 @@ def _get_classed_reqs( is_decode = False if is_decode: - token_num = req_obj.decode_need_token_num() - if token_num <= can_alloc_token_num: - self._alloc_req_kv_mem(req_obj, token_num) + alloc_token_num = req_obj.decode_need_token_num() + if alloc_token_num <= can_alloc_token_num: + self._alloc_req_kv_mem(req_obj, alloc_token_num) decode_reqs.append(req_obj) - can_alloc_token_num -= token_num + can_alloc_token_num -= alloc_token_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True @@ -752,10 +755,15 @@ def _get_classed_reqs( if req_obj.is_slave_req(): continue - token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill) - alloc_token_num = req_obj.prefill_kv_alloc_need(is_chuncked_prefill=not self.disable_chunked_prefill) + if self.disable_chunked_prefill: + token_num = req_obj.get_cur_total_len() - req_obj.cur_kv_len + else: + token_num = req_obj.get_chuncked_input_token_len() - req_obj.cur_kv_len if prefill_tokens + token_num > self.batch_max_tokens: continue + alloc_token_num = req_obj.prefill_need_token_num( + is_chuncked_prefill=not self.disable_chunked_prefill + ) if alloc_token_num <= can_alloc_token_num: self._alloc_req_kv_mem(req_obj, alloc_token_num) prefill_tokens += token_num diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py index c87ecfcf34..23e4c1a869 100644 --- a/unit_tests/common/test_req_manager_page.py +++ b/unit_tests/common/test_req_manager_page.py @@ -41,7 +41,7 @@ def _make_context(monkeypatch): def _make_req(req_idx): req = SimpleNamespace(req_idx=req_idx, cur_kv_len=0, hold_kv_len=0) - req._kv_cache_alloc_need = lambda target_len: (target_len + 3) // 4 * 4 - req.hold_kv_len + req._kv_cache_alloc_need = lambda target_len: InferReq._kv_cache_alloc_need(req, target_len) return req @@ -49,17 +49,17 @@ def test_request_reuses_reserved_page_tail_before_allocating_next_page(monkeypat context, backend = _make_context(monkeypatch) req = _make_req(0) - backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(3)) + backend._alloc_req_kv_mem(req, alloc_token_num=4) assert req.hold_kv_len == 4 assert context.req_manager.mem_manager.alloc_sizes == [4] assert context.req_manager.req_to_token_indexs[0, :4].tolist() == [0, 1, 2, 3] req.cur_kv_len = 3 - backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(4)) + backend._alloc_req_kv_mem(req, alloc_token_num=0) assert context.req_manager.mem_manager.alloc_sizes == [4] req.cur_kv_len = 4 - backend._alloc_req_kv_mem(req, req._kv_cache_alloc_need(6)) + backend._alloc_req_kv_mem(req, alloc_token_num=4) assert req.hold_kv_len == 8 assert context.req_manager.mem_manager.alloc_sizes == [4, 4] assert context.req_manager.req_to_token_indexs[0, :8].tolist() == list(range(8)) @@ -70,14 +70,33 @@ def test_reservation_fills_each_request_table_row(monkeypatch): req0 = _make_req(0) req1 = _make_req(1) - backend._alloc_req_kv_mem(req0, req0._kv_cache_alloc_need(2)) - backend._alloc_req_kv_mem(req1, req1._kv_cache_alloc_need(3)) + backend._alloc_req_kv_mem(req0, alloc_token_num=4) + backend._alloc_req_kv_mem(req1, alloc_token_num=4) assert generic_pre_process.g_infer_context.req_manager.req_to_token_indexs[0, :4].tolist() == [0, 1, 2, 3] assert generic_pre_process.g_infer_context.req_manager.req_to_token_indexs[1, :4].tolist() == [4, 5, 6, 7] assert req0.hold_kv_len == req1.hold_kv_len == 4 +def test_need_token_num_returns_page_aligned_allocation(): + req = SimpleNamespace( + args=SimpleNamespace(page_size=4), + cur_kv_len=3, + hold_kv_len=4, + mtp_step=0, + get_chuncked_input_token_len=lambda: 6, + get_cur_total_len=lambda: 10, + ) + req._kv_cache_alloc_need = lambda target_len: InferReq._kv_cache_alloc_need(req, target_len) + + assert InferReq.prefill_need_token_num(req, is_chuncked_prefill=True) == 4 + assert InferReq.prefill_need_token_num(req, is_chuncked_prefill=False) == 8 + assert InferReq.decode_need_token_num(req) == 0 + + req.cur_kv_len = 4 + assert InferReq.decode_need_token_num(req) == 4 + + def test_decode_reserves_mtp_headroom(monkeypatch): context, backend = _make_context(monkeypatch) req = _make_req(0) @@ -92,7 +111,7 @@ def test_decode_reserves_mtp_headroom(monkeypatch): context.req_manager.req_to_token_indexs[0, :4] = torch.arange(4, dtype=torch.int32) context.req_manager.mem_manager.next_index = 4 - alloc_token_num = InferReq._mtp_decode_need_token_num(req) + alloc_token_num = InferReq.decode_need_token_num(req) backend._alloc_req_kv_mem(req, alloc_token_num) model_input, run_reqs = generic_pre_process.prepare_decode_inputs([req]) @@ -116,11 +135,9 @@ def test_page_size_one_uses_the_same_scheduler_preallocation(monkeypatch): req.get_cur_total_len = lambda: 4 req.get_radix_cache_shared_len = lambda: 0 req.args = context.args - req._kv_cache_alloc_need = lambda target_len: InferReq._kv_cache_alloc_need(req, target_len) - context.req_manager.req_to_token_indexs[0, :3] = torch.arange(3, dtype=torch.int32) context.req_manager.mem_manager.next_index = 3 - alloc_token_num = InferReq._mtp_decode_need_token_num(req) + alloc_token_num = InferReq.decode_need_token_num(req) backend._alloc_req_kv_mem(req, alloc_token_num) model_input, _ = generic_pre_process.prepare_decode_inputs([req]) From 585f857f1118c5a850af0bb54f6150744d7429c2 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 06:12:30 +0000 Subject: [PATCH 06/15] refactor: reduce MTP KV reservation window --- lightllm/server/core/objs/req.py | 4 ++-- lightllm/server/router/model_infer/infer_batch.py | 2 +- unit_tests/common/test_req_manager_page.py | 6 +++--- unit_tests/server/core/objs/test_req.py | 4 ++-- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index df35f90ab6..8fb7c1d717 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -504,8 +504,8 @@ def get_decode_need_tokens(self): # 当开启 mtp 模式以后,每一次 decode 需要的 token 数量会增加 need_tokens = min(self.input_len + self.shm_cur_output_len - self.shm_cur_kv_len, self.chunked_prefill_size) if need_tokens == 1 and self._mtp_step > 0: - # target verify 及后续 MTP 操作统一预留三倍窗口。 - need_tokens = (self._mtp_step + 1) * 3 + # target verify 及后续 MTP 操作统一预留两倍窗口。 + need_tokens = (self._mtp_step + 1) * 2 return need_tokens diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index ce10d04980..3b128312ec 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -916,7 +916,7 @@ def prefill_need_token_num(self, is_chuncked_prefill: bool): return self._kv_cache_alloc_need(target_kv_len) def decode_need_token_num(self) -> int: - decode_token_num = 1 if self.mtp_step == 0 else 3 * (1 + self.mtp_step) + decode_token_num = 1 if self.mtp_step == 0 else 2 * (1 + self.mtp_step) return self._kv_cache_alloc_need(self.cur_kv_len + decode_token_num) def _kv_cache_alloc_need(self, target_kv_len: int) -> int: diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py index 23e4c1a869..ed508fbcdd 100644 --- a/unit_tests/common/test_req_manager_page.py +++ b/unit_tests/common/test_req_manager_page.py @@ -142,9 +142,9 @@ def test_page_size_one_uses_the_same_scheduler_preallocation(monkeypatch): model_input, _ = generic_pre_process.prepare_decode_inputs([req]) - assert req.hold_kv_len == 12 - assert context.req_manager.mem_manager.alloc_sizes == [9] - assert context.req_manager.req_to_token_indexs[0, :12].tolist() == list(range(12)) + assert req.hold_kv_len == 9 + assert context.req_manager.mem_manager.alloc_sizes == [6] + assert context.req_manager.req_to_token_indexs[0, :9].tolist() == list(range(9)) def test_page_size_one_frees_all_preallocated_indexes(): diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 74d65d1817..6b85ff3b25 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -55,7 +55,7 @@ def test_get_used_tokens(req): assert req.get_used_tokens() == 5 -def test_mtp_decode_reserves_three_windows(): +def test_mtp_decode_reserves_two_windows(): req = SimpleNamespace( input_len=4, shm_cur_output_len=0, @@ -64,7 +64,7 @@ def test_mtp_decode_reserves_three_windows(): _mtp_step=2, ) - assert ChunkedPrefillReq.get_decode_need_tokens(req) == 9 + assert ChunkedPrefillReq.get_decode_need_tokens(req) == 6 def test_final_token_metadata_read_returns_actual_prompt_tokens(req): From 2392dadf383fe0df8092ea50286bc037b0e6a9f5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 06:44:54 +0000 Subject: [PATCH 07/15] refactor: move token scheduling control to infer backend --- lightllm/server/core/objs/req.py | 23 +-------------- lightllm/server/router/batch.py | 9 ------ .../server/router/req_queue/base_queue.py | 2 -- .../req_queue/chunked_prefill/beam_impl.py | 25 ++++------------ .../router/req_queue/chunked_prefill/impl.py | 27 ++++------------- .../server/router/req_queue/dp_base_queue.py | 4 --- unit_tests/server/core/objs/test_req.py | 29 ++++++++----------- 7 files changed, 24 insertions(+), 95 deletions(-) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 8fb7c1d717..f21ac50b79 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -423,12 +423,6 @@ def get_used_tokens(self): def get_tuple_tokens(self, is_busy, ema_req_out_len): raise NotImplementedError("Subclasses should implement this method") - def get_decode_need_tokens(self): - raise NotImplementedError("Subclasses should implement this method") - - def get_first_router_need_tokens(self): - raise NotImplementedError("Subclasses should implement this method") - def get_output_logprobs_metadata(self, src_index: int, tokenizer=None): token_id = int(self.shm_prompt_ids.arr[src_index]) rank = int(self.shm_logprobs.arr["rank"][src_index]) @@ -494,21 +488,6 @@ def get_tuple_tokens(self, is_busy, ema_req_out_len): - 1 ) b_len = max(0, b_len) + ADDED_OUTPUT_LEN + b_len = (b_len + args.page_size - 1) // args.page_size * args.page_size return (a_len, b_len) - - def get_decode_need_tokens(self): - """ - chunkedprefill 调度模式的实现 - """ - # 当开启 mtp 模式以后,每一次 decode 需要的 token 数量会增加 - need_tokens = min(self.input_len + self.shm_cur_output_len - self.shm_cur_kv_len, self.chunked_prefill_size) - if need_tokens == 1 and self._mtp_step > 0: - # target verify 及后续 MTP 操作统一预留两倍窗口。 - need_tokens = (self._mtp_step + 1) * 2 - - return need_tokens - - def get_first_router_need_tokens(self): - - return min(self.input_len + self.shm_cur_output_len, self.chunked_prefill_size) diff --git a/lightllm/server/router/batch.py b/lightllm/server/router/batch.py index 24d0b9b824..e9d462a256 100644 --- a/lightllm/server/router/batch.py +++ b/lightllm/server/router/batch.py @@ -22,15 +22,6 @@ def input_tokens(self): batch_input_tokens += req.input_len return batch_input_tokens - def get_batch_decode_need_tokens(self): - new_batch_decode_need_tokens = [0 for _ in range(self.dp_size_in_node)] # for chunked prefill - - for req in self.reqs: - req_dp_index = req.sample_params.suggested_dp_index - new_batch_decode_need_tokens[req_dp_index] += req.get_decode_need_tokens() - - return new_batch_decode_need_tokens - def get_req_list_for_dp(self, dp_index: int): if self.dp_size_in_node == 1: return self.reqs diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index f5d33fa35a..645df5d7b8 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -23,8 +23,6 @@ def __init__(self, args: StartArgs, router, dp_index, dp_size_in_node) -> None: # 在极端情况下减少,在非特定模式下,get_fixed_kv_len() 返回的都是 # 0, 不会有任何影响。 self.max_total_tokens = args.max_total_token_num - get_fixed_kv_len() - assert args.batch_max_tokens is not None - self.batch_max_tokens = args.batch_max_tokens self.running_max_req_size = args.running_max_req_size # Maximum number of concurrent requests self.waiting_req_list: List[Req] = [] # List of queued requests self.router_token_ratio = args.router_token_ratio # ratio to determine whether the router is busy diff --git a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py index 42476333e0..c768953e3e 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py @@ -21,7 +21,7 @@ def _init_cache_list(self, current_batch: Batch, is_busy): return # @calculate_time(show=True, min_cost_ms=0.1) - def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy, new_batch_first_router_need_tokens): + def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy): for req in cur_handle_group_reqs: self.cache_len_list.append( (req, req.get_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len)) @@ -43,28 +43,20 @@ def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy, new cumsum_len += cur_input_len - req.input_len # 减去共享的部分 need_max_token_num = max(need_max_token_num, cumsum_len + index * cur_ouput_len) - # prefill token 计算, 因为对beam的prefill计算过程是共享的,所以只计算一个请求对应的token数量 - new_batch_first_router_need_tokens += req.get_first_router_need_tokens() estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) ok_token_num = estimated_need_token_num < self.max_total_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size - # 长短请求模式由 Infer 控制单轮 prefill token 上限。 - ok_prefill = ( - self.args.short_prefill_token_threshold is not None - or new_batch_first_router_need_tokens <= self.batch_max_tokens - ) - - if ok_token_num and ok_req_num and ok_prefill: + if ok_token_num and ok_req_num: self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( estimated_need_token_num / self.max_total_tokens, self.dp_index, ) - return True, new_batch_first_router_need_tokens + return True else: - return False, new_batch_first_router_need_tokens + return False # @calculate_time(show=True, min_cost_ms=10) def generate_new_batch(self, current_batch: Batch): @@ -86,16 +78,13 @@ def generate_new_batch(self, current_batch: Batch): self._init_cache_list(current_batch, is_busy) can_run_list = [] - new_batch_first_router_need_tokens = 0 # 主要是对 prefill 大块计算时候的token数量限制 cur_group_reqs = [] for req in self.waiting_req_list: if self._add_to_group(cur_group_reqs, req): continue - ok_insert, new_batch_first_router_need_tokens = self._can_add_new_group_reqs( - cur_group_reqs, is_busy, new_batch_first_router_need_tokens - ) + ok_insert = self._can_add_new_group_reqs(cur_group_reqs, is_busy) if ok_insert: can_run_list.extend(cur_group_reqs) cur_group_reqs = [req] # 等待判断的组 @@ -104,9 +93,7 @@ def generate_new_batch(self, current_batch: Batch): break if len(cur_group_reqs) != 0: - ok_insert, new_batch_first_router_need_tokens = self._can_add_new_group_reqs( - cur_group_reqs, is_busy, new_batch_first_router_need_tokens - ) + ok_insert = self._can_add_new_group_reqs(cur_group_reqs, is_busy) if ok_insert: can_run_list.extend(cur_group_reqs) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py index 799af7f97a..1f236ff870 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py @@ -5,10 +5,6 @@ class ChunkedPrefillQueue(BaseQueue): - def __init__(self, args, router, dp_index, dp_size_in_node) -> None: - super().__init__(args, router, dp_index, dp_size_in_node) - self.batch_max_tokens = self.batch_max_tokens * 2 - def _init_cache_list(self, current_batch: Batch, is_busy): if current_batch is not None: self.cache_len_list = [ @@ -21,7 +17,7 @@ def _init_cache_list(self, current_batch: Batch, is_busy): return # @calculate_time(show=True, min_cost_ms=0.1) - def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens): + def _can_add_new_req(self, req: Req, is_busy): self.cache_len_list.append( req.get_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) ) # hard to analysis @@ -38,22 +34,15 @@ def _can_add_new_req(self, req: Req, is_busy, new_batch_first_router_need_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size - new_batch_first_router_need_tokens += req.get_first_router_need_tokens() - # 长短请求模式由 Infer 控制单轮 prefill token 上限。 - ok_prefill = ( - self.args.short_prefill_token_threshold is not None - or new_batch_first_router_need_tokens <= self.batch_max_tokens - ) - - if ok_token_num and ok_req_num and ok_prefill: + if ok_token_num and ok_req_num: self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( estimated_need_token_num / self.max_total_tokens, self.dp_index, ) - return True, new_batch_first_router_need_tokens + return True else: - return False, new_batch_first_router_need_tokens + return False # @calculate_time(show=True, min_cost_ms=10) def generate_new_batch(self, current_batch: Batch): @@ -72,10 +61,6 @@ def generate_new_batch(self, current_batch: Batch): is_busy = self.is_busy() - new_batch_first_router_need_tokens = ( - 0 if current_batch is None else current_batch.get_batch_decode_need_tokens()[self.dp_index] - ) - self._init_cache_list(current_batch, is_busy) can_run_list = [] consumed_req_count = 0 @@ -83,9 +68,7 @@ def generate_new_batch(self, current_batch: Batch): waiting_queue = self.waiting_req_list for req in waiting_queue: - ok_insert, new_batch_first_router_need_tokens = self._can_add_new_req( - req, is_busy, new_batch_first_router_need_tokens - ) + ok_insert = self._can_add_new_req(req, is_busy) if ok_insert: consumed_req_count += 1 can_run_list.append(req) diff --git a/lightllm/server/router/req_queue/dp_base_queue.py b/lightllm/server/router/req_queue/dp_base_queue.py index af8f875d4e..c465bee562 100644 --- a/lightllm/server/router/req_queue/dp_base_queue.py +++ b/lightllm/server/router/req_queue/dp_base_queue.py @@ -18,10 +18,6 @@ def __init__(self, args, router, base_queue_class, dp_size_in_node) -> None: self.inner_queues: List[BaseQueue] = [ base_queue_class(args, router, dp_index, dp_size_in_node) for dp_index in range(self.dp_size_in_node) ] - # 在调度这放松,在推理时约束。 - # 避免prefill 模式下的情况下,推理完成了,调度没及时获取信息,导致调度bs 过小 - for queue in self.inner_queues: - queue.batch_max_tokens = int(args.batch_max_tokens * 2) self.dp_balancer = get_dp_balancer(args, dp_size_in_node, self.inner_queues) self.reqs_waiting_for_dp_index: List[List[Req]] = [] return diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 6b85ff3b25..e861cb2240 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -17,6 +17,8 @@ def setup_module_env(): "cpu_cache_token_page_size": 256, "enable_cpu_cache": False, "model_dir": "", + "page_size": 4, + "router_max_wait_tokens": 1, } ) ) @@ -55,18 +57,6 @@ def test_get_used_tokens(req): assert req.get_used_tokens() == 5 -def test_mtp_decode_reserves_two_windows(): - req = SimpleNamespace( - input_len=4, - shm_cur_output_len=0, - shm_cur_kv_len=3, - chunked_prefill_size=128, - _mtp_step=2, - ) - - assert ChunkedPrefillReq.get_decode_need_tokens(req) == 6 - - def test_final_token_metadata_read_returns_actual_prompt_tokens(req): req.sample_params.prompt_logprobs = 0 req.shm_logprobs.arr["logprob"][1] = -0.5 @@ -84,11 +74,16 @@ def test_final_token_metadata_read_returns_actual_prompt_tokens(req): ] -# def test_chunked_req_get_tuple_tokens(): -# chunked_req = ChunkedPrefillReq() -# chunked_req.init(1, [1, 2, 3], {"max_new_tokens": 1}, None, chunked_prefill_size=256) -# result = chunked_req.get_tuple_tokens(False, 10) -# assert isinstance(result, tuple) +def test_chunked_req_get_tuple_tokens_aligns_remaining_len_to_page(): + req = SimpleNamespace( + input_len=10, + shm_cur_output_len=0, + shm_cur_kv_len=0, + chunked_prefill_size=4, + sample_params=SimpleNamespace(ignore_eos=True, max_new_tokens=5), + ) + + assert ChunkedPrefillReq.get_tuple_tokens(req, False, 10) == (11, 28) def test_finish_status(req): From 3e1b0d7ea2abf0a279ff3a3359a6113e93636bdf Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 07:45:02 +0000 Subject: [PATCH 08/15] refactor: remove duplicate KV page reservation --- lightllm/server/router/req_queue/base_queue.py | 5 ----- .../router/req_queue/chunked_prefill/beam_impl.py | 10 ++++------ .../server/router/req_queue/chunked_prefill/impl.py | 10 ++++------ 3 files changed, 8 insertions(+), 17 deletions(-) diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 645df5d7b8..90d4b0cbc0 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -27,11 +27,6 @@ def __init__(self, args: StartArgs, router, dp_index, dp_size_in_node) -> None: self.waiting_req_list: List[Req] = [] # List of queued requests self.router_token_ratio = args.router_token_ratio # ratio to determine whether the router is busy - def add_kv_page_reservation(self, token_num: int, req_num: int) -> int: - """Conservatively include each running request's incomplete tail page.""" - page_size = self.args.page_size - return token_num + req_num * (page_size - 1) - def free_aborted_req_cpu_cache_pages(self, req: Req): if self.args.enable_cpu_cache: self.router.cpu_cache_client.lock.acquire_sleep1ms() diff --git a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py index c768953e3e..57296bf22f 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/beam_impl.py @@ -43,15 +43,14 @@ def _can_add_new_group_reqs(self, cur_handle_group_reqs: List[Req], is_busy): cumsum_len += cur_input_len - req.input_len # 减去共享的部分 need_max_token_num = max(need_max_token_num, cumsum_len + index * cur_ouput_len) - estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) - ok_token_num = estimated_need_token_num < self.max_total_tokens + ok_token_num = need_max_token_num < self.max_total_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size if ok_token_num and ok_req_num: - self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) + self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( - estimated_need_token_num / self.max_total_tokens, + need_max_token_num / self.max_total_tokens, self.dp_index, ) return True @@ -132,5 +131,4 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): assert cur_input_len - req.input_len >= 0 cumsum_len += cur_input_len - req.input_len # 减去共享的部分 need_max_token_num = max(need_max_token_num, cumsum_len + index * cur_ouput_len) - estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) - return (estimated_need_token_num, estimated_need_token_num / self.max_total_tokens) + return (need_max_token_num, need_max_token_num / self.max_total_tokens) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl.py b/lightllm/server/router/req_queue/chunked_prefill/impl.py index 1f236ff870..7723ee51b3 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl.py @@ -29,15 +29,14 @@ def _can_add_new_req(self, req: Req, is_busy): size_array = np.arange(1, len(self.cache_len_list) + 1, 1) need_max_token_num = (left_out_len_array * size_array + cum_run_len_array).max() - estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) - ok_token_num = estimated_need_token_num < self.max_total_tokens + ok_token_num = need_max_token_num < self.max_total_tokens ok_req_num = len(self.cache_len_list) <= self.running_max_req_size if ok_token_num and ok_req_num: - self.router.shared_token_load.set_estimated_peak_token_count(estimated_need_token_num, self.dp_index) + self.router.shared_token_load.set_estimated_peak_token_count(need_max_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( - estimated_need_token_num / self.max_total_tokens, + need_max_token_num / self.max_total_tokens, self.dp_index, ) return True @@ -93,5 +92,4 @@ def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): else: need_max_token_num = 0 - estimated_need_token_num = self.add_kv_page_reservation(need_max_token_num, len(self.cache_len_list)) - return (estimated_need_token_num, estimated_need_token_num / self.max_total_tokens) + return (need_max_token_num, need_max_token_num / self.max_total_tokens) From 329f5b60072a156cd082ab9fc7381a1b0a366a54 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 08:01:20 +0000 Subject: [PATCH 09/15] fix: restrict fixed KV cache to single-token pages --- lightllm/utils/config_utils.py | 1 + .../utils/test_config_utils_fixed_kv.py | 26 +++++++++++++++++++ 2 files changed, 27 insertions(+) create mode 100644 unit_tests/utils/test_config_utils_fixed_kv.py diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 6695f0ec44..00fc5091e0 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -368,6 +368,7 @@ def get_dtype(model_path: str): @lru_cache(maxsize=None) def get_fixed_kv_len(): start_args = get_env_start_args() + assert start_args.page_size == 1, "fixed KV cache only supports page_size == 1" model_cfg = get_config_json(start_args.model_dir) if "prompt_cache_token_ids" in model_cfg: return len(model_cfg["prompt_cache_token_ids"]) diff --git a/unit_tests/utils/test_config_utils_fixed_kv.py b/unit_tests/utils/test_config_utils_fixed_kv.py new file mode 100644 index 0000000000..127cd6eb4c --- /dev/null +++ b/unit_tests/utils/test_config_utils_fixed_kv.py @@ -0,0 +1,26 @@ +from types import SimpleNamespace + +import pytest + +from lightllm.utils import config_utils + + +@pytest.fixture(autouse=True) +def clear_fixed_kv_len_cache(): + config_utils.get_fixed_kv_len.cache_clear() + yield + config_utils.get_fixed_kv_len.cache_clear() + + +def test_fixed_kv_len_supports_page_size_one(monkeypatch): + monkeypatch.setattr(config_utils, "get_env_start_args", lambda: SimpleNamespace(page_size=1, model_dir="model")) + monkeypatch.setattr(config_utils, "get_config_json", lambda _: {"prompt_cache_token_ids": [1, 2, 3]}) + + assert config_utils.get_fixed_kv_len() == 3 + + +def test_fixed_kv_len_rejects_paged_kv(monkeypatch): + monkeypatch.setattr(config_utils, "get_env_start_args", lambda: SimpleNamespace(page_size=2, model_dir="model")) + + with pytest.raises(AssertionError, match="fixed KV cache only supports page_size == 1"): + config_utils.get_fixed_kv_len() From 7d9166ba899f481d76207e3fd05c0f961a17e19d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 28 Aug 2026 09:13:39 +0000 Subject: [PATCH 10/15] refactor: split PD queues by inference stage --- lightllm/server/core/objs/req.py | 25 ++++ lightllm/server/router/req_queue/__init__.py | 9 +- .../chunked_prefill/impl_for_pd_decode.py | 108 ++++++++++++++++++ ...{impl_for_pd.py => impl_for_pd_prefill.py} | 37 +++--- unit_tests/server/core/objs/test_req.py | 11 ++ .../req_queue/test_pd_queue_selection.py | 95 +++++++++++++++ 6 files changed, 263 insertions(+), 22 deletions(-) create mode 100644 lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py rename lightllm/server/router/req_queue/chunked_prefill/{impl_for_pd.py => impl_for_pd_prefill.py} (75%) create mode 100644 unit_tests/server/router/req_queue/test_pd_queue_selection.py diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index f21ac50b79..1e58d883a7 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -3,6 +3,7 @@ import ctypes import asyncio import numpy as np +import triton import time from .sampling_params import SamplingParams from .out_token_circlequeue import CircularQueue @@ -423,6 +424,9 @@ def get_used_tokens(self): def get_tuple_tokens(self, is_busy, ema_req_out_len): raise NotImplementedError("Subclasses should implement this method") + def get_pd_decode_mode_tuple_tokens(self, is_busy, ema_req_out_len): + raise NotImplementedError("Subclasses should implement this method") + def get_output_logprobs_metadata(self, src_index: int, tokenizer=None): token_id = int(self.shm_prompt_ids.arr[src_index]) rank = int(self.shm_logprobs.arr["rank"][src_index]) @@ -491,3 +495,24 @@ def get_tuple_tokens(self, is_busy, ema_req_out_len): b_len = (b_len + args.page_size - 1) // args.page_size * args.page_size return (a_len, b_len) + + def get_pd_decode_mode_tuple_tokens(self, is_busy, ema_req_out_len): + args = get_env_start_args() + has_out_len = self.shm_cur_output_len + if self.sample_params.ignore_eos: + cur_max_new_token_len = self.sample_params.max_new_tokens + elif is_busy: + cur_max_new_token_len = self.sample_params.max_new_tokens + else: + cur_max_new_token_len = min( + self.sample_params.max_new_tokens, + max(int(1.1 * has_out_len), ema_req_out_len), + ) + + # PD decode 节点只运行 decode,不需要考虑 chunked prefill 带来的调度等待时间。 + # 当前占用和预计剩余增长分别按 page_size 对齐,使后续峰值计算直接使用物理 KV 容量。 + a_len = max(self.input_len + has_out_len + 1, self.shm_cur_kv_len + 1) + a_len = triton.cdiv(a_len, args.page_size) * args.page_size + b_len = max(0, cur_max_new_token_len - has_out_len - 1) + ADDED_OUTPUT_LEN + b_len = triton.cdiv(b_len, args.page_size) * args.page_size + return (a_len, b_len) diff --git a/lightllm/server/router/req_queue/__init__.py b/lightllm/server/router/req_queue/__init__.py index c0de01db28..5876833536 100644 --- a/lightllm/server/router/req_queue/__init__.py +++ b/lightllm/server/router/req_queue/__init__.py @@ -1,6 +1,7 @@ from .chunked_prefill.impl import ChunkedPrefillQueue from .chunked_prefill.beam_impl import ChunkedBeamContinuesBatchQueue -from .chunked_prefill.impl_for_pd import PDQueue +from .chunked_prefill.impl_for_pd_prefill import PDPrefillQueue +from .chunked_prefill.impl_for_pd_decode import PDDecodeQueue from .dp_base_queue import DpQueue @@ -11,8 +12,10 @@ def _get_req_queue_class(args, router, dp_size_in_node: int): return ChunkedPrefillQueue if args.first_token_constraint_mode: return ChunkedPrefillQueue - if args.run_mode in ["prefill", "decode"]: - return PDQueue + if args.run_mode == "prefill": + return PDPrefillQueue + if args.run_mode == "decode": + return PDDecodeQueue if args.disable_chunked_prefill: # 虽然也使用chuncked prefill queue 但是由于 args.chunked_prefill_size = args.max_req_total_len diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py new file mode 100644 index 0000000000..361ed54417 --- /dev/null +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py @@ -0,0 +1,108 @@ +import uuid +import numpy as np +import triton +from typing import Tuple +from ...batch import Batch, Req +from lightllm.server.router.req_queue.base_queue import BaseQueue + + +class PDDecodeQueue(BaseQueue): + def __init__(self, args, router, dp_index, dp_size_in_node) -> None: + super().__init__(args, router, dp_index, dp_size_in_node) + + # @calculate_time(show=True, min_cost_ms=0.1) + def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_num: int) -> Tuple[bool, int, int]: + # 新请求尚未进入 decode 阶段,缺少实际输出长度等运行信息,只能按输入长度加最大输出长度 + # 保守估算该请求最多可能占用的 KV 资源。 + req_token_num = req.input_len + req.sample_params.max_new_tokens + req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + estimated_peak_token_num += req_token_num + ok_token_num = estimated_peak_token_num < self.max_total_tokens + batch_req_num += 1 + ok_req_num = batch_req_num <= self.running_max_req_size + + if ok_token_num and ok_req_num: + self.router.shared_token_load.set_estimated_peak_token_count(estimated_peak_token_num, self.dp_index) + self.router.shared_token_load.set_dynamic_max_load( + estimated_peak_token_num / self.max_total_tokens, + self.dp_index, + ) + return True, estimated_peak_token_num, batch_req_num + else: + return False, None, None + + def _caclu_batch_estimated_peak_token_num(self, batch: Batch): + is_busy = self.is_busy() + estimated_peak_token_num = 0 + decoding_req_list = [] + if batch is not None: + for req in batch.reqs: + if req.sample_params.suggested_dp_index == self.dp_index: + if req.is_infer_decode(): + # 请求进入 decode 阶段后,可以结合已经运行的 token 数量和预计剩余输出长度, + # 使用连续批处理峰值算法估算其动态 KV 占用。 + decoding_req_list.append( + req.get_pd_decode_mode_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) + ) + else: + # 尚未进入 decode 阶段的请求没有足够的动态信息,仍按输入长度加最大输出长度 + # 预留其最大 KV 资源。 + req_token_num = req.input_len + req.sample_params.max_new_tokens + req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + estimated_peak_token_num += req_token_num + + if decoding_req_list: + # 按预计剩余输出长度排序,计算每个请求结束时仍存活请求的 KV 占用峰值, + # 再与未进入 decode 阶段请求的保守占用相加,得到整个 batch 的最终峰值 token 估算。 + decoding_req_list.sort(key=lambda x: -x[1]) + left_out_len_array = np.array([e[1] for e in decoding_req_list]) + has_run_len_array = np.array([e[0] for e in decoding_req_list]) + cum_run_len_array = np.cumsum(has_run_len_array) + size_array = np.arange(1, len(decoding_req_list) + 1, 1) + estimated_peak_token_num += (left_out_len_array * size_array + cum_run_len_array).max() + + return estimated_peak_token_num + + # @calculate_time(show=True, min_cost_ms=10) + def generate_new_batch(self, current_batch: Batch): + if len(self.waiting_req_list) == 0: + return None + + # 如果当前已经被调度的请求数量超过了上限,直接不调度新的请求了。 + exist_req_num = self.get_batch_dp_req_size(current_batch) + req_is_full = exist_req_num >= self.running_max_req_size + if req_is_full: + return None + + self.filter_aborted_reqs() + if len(self.waiting_req_list) == 0: + return None + + estimated_peak_token_num = self._caclu_batch_estimated_peak_token_num(current_batch) + batch_req_num = exist_req_num + + can_run_list = [] + consumed_req_count = 0 + + waiting_queue = self.waiting_req_list + + for req in waiting_queue: + ok_insert, estimated_peak_token_num, batch_req_num = self._can_add_new_req( + req=req, estimated_peak_token_num=estimated_peak_token_num, batch_req_num=batch_req_num + ) + if ok_insert: + consumed_req_count += 1 + can_run_list.append(req) + else: + break + new_batch = None + if len(can_run_list) != 0: + new_batch = Batch(uuid.uuid4().int, can_run_list, dp_size_in_node=self.dp_size_in_node) + self.waiting_req_list = self.waiting_req_list[consumed_req_count:] + return new_batch + + def _calcu_batch_token_load_batch_not_none(self, current_batch: Batch): + + estimated_peak_token_num = self._caclu_batch_estimated_peak_token_num(current_batch) + + return (estimated_peak_token_num, estimated_peak_token_num / self.max_total_tokens) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py similarity index 75% rename from lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py rename to lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py index f6e33144ef..c6d8212f54 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py @@ -1,17 +1,27 @@ import uuid -import numpy as np +import triton from typing import Tuple from ...batch import Batch, Req from lightllm.server.router.req_queue.base_queue import BaseQueue +from lightllm.utils.log_utils import init_logger -class PDQueue(BaseQueue): +logger = init_logger(__name__) + + +class PDPrefillQueue(BaseQueue): def __init__(self, args, router, dp_index, dp_size_in_node) -> None: super().__init__(args, router, dp_index, dp_size_in_node) + logger.info( + "PD prefill requests normally generate only one output token; " + "estimate peak KV usage by summing their page-aligned token counts" + ) # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_num: int) -> Tuple[bool, int, int]: - estimated_peak_token_num += req.input_len + req.sample_params.max_new_tokens + req_token_num = req.input_len + req.sample_params.max_new_tokens + req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + estimated_peak_token_num += req_token_num ok_token_num = estimated_peak_token_num < self.max_total_tokens batch_req_num += 1 ok_req_num = batch_req_num <= self.running_max_req_size @@ -27,26 +37,15 @@ def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_nu return False, None, None def _caclu_batch_estimated_peak_token_num(self, batch: Batch): - is_busy = self.is_busy() estimated_peak_token_num = 0 - decoding_req_list = [] if batch is not None: for req in batch.reqs: if req.sample_params.suggested_dp_index == self.dp_index: - if req.is_infer_decode(): - decoding_req_list.append( - req.get_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) - ) - else: - estimated_peak_token_num += req.input_len + req.sample_params.max_new_tokens - - if decoding_req_list: - decoding_req_list.sort(key=lambda x: -x[1]) - left_out_len_array = np.array([e[1] for e in decoding_req_list]) - has_run_len_array = np.array([e[0] for e in decoding_req_list]) - cum_run_len_array = np.cumsum(has_run_len_array) - size_array = np.arange(1, len(decoding_req_list) + 1, 1) - estimated_peak_token_num += (left_out_len_array * size_array + cum_run_len_array).max() + # PD prefill 请求通常只生成一个 token,其 KV 占用不会像 decode 请求一样持续增长, + # 因此将每个请求按 page_size 对齐后的 token 数量直接线性相加即可完成估算。 + req_token_num = req.input_len + req.sample_params.max_new_tokens + req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + estimated_peak_token_num += req_token_num return estimated_peak_token_num diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index e861cb2240..694aae24ca 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -86,6 +86,17 @@ def test_chunked_req_get_tuple_tokens_aligns_remaining_len_to_page(): assert ChunkedPrefillReq.get_tuple_tokens(req, False, 10) == (11, 28) +def test_chunked_req_get_pd_decode_mode_tuple_tokens_uses_decode_estimation(): + req = SimpleNamespace( + input_len=10, + shm_cur_output_len=3, + shm_cur_kv_len=12, + sample_params=SimpleNamespace(ignore_eos=True, max_new_tokens=20), + ) + + assert ChunkedPrefillReq.get_pd_decode_mode_tuple_tokens(req, False, 10) == (16, 32) + + def test_finish_status(req): req.finish_status.set_status(req.finish_status.FINISHED_STOP) assert req.finish_status.is_finished() diff --git a/unit_tests/server/router/req_queue/test_pd_queue_selection.py b/unit_tests/server/router/req_queue/test_pd_queue_selection.py new file mode 100644 index 0000000000..1d8bc9a8dd --- /dev/null +++ b/unit_tests/server/router/req_queue/test_pd_queue_selection.py @@ -0,0 +1,95 @@ +from types import SimpleNamespace + +from lightllm.server.router.req_queue import _get_req_queue_class +from lightllm.server.router.batch import Batch +from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd_decode import PDDecodeQueue +from lightllm.server.router.req_queue.chunked_prefill.impl_for_pd_prefill import PDPrefillQueue + + +def _make_args(run_mode: str): + return SimpleNamespace( + run_mode=run_mode, + diverse_mode=False, + output_constraint_mode="none", + first_token_constraint_mode=False, + disable_chunked_prefill=False, + ) + + +def test_pd_prefill_uses_prefill_queue(): + args = _make_args("prefill") + + assert _get_req_queue_class(args, router=None, dp_size_in_node=1) is PDPrefillQueue + + +def test_pd_decode_uses_decode_queue(): + args = _make_args("decode") + + assert _get_req_queue_class(args, router=None, dp_size_in_node=1) is PDDecodeQueue + + +def test_pd_prefill_peak_tokens_do_not_use_decode_estimation(): + queue = PDPrefillQueue.__new__(PDPrefillQueue) + queue.dp_index = 0 + queue.args = SimpleNamespace(page_size=16) + reqs = [ + SimpleNamespace( + request_id="req-0", + input_len=10, + sample_params=SimpleNamespace(max_new_tokens=1, suggested_dp_index=0), + ), + SimpleNamespace( + request_id="req-1", + input_len=20, + sample_params=SimpleNamespace(max_new_tokens=1, suggested_dp_index=0), + ), + SimpleNamespace( + request_id="req-2", + input_len=100, + sample_params=SimpleNamespace(max_new_tokens=1, suggested_dp_index=1), + ), + ] + batch = Batch(batch_id=1, reqs=reqs, dp_size_in_node=2) + + assert queue._caclu_batch_estimated_peak_token_num(batch) == 48 + + +def test_pd_decode_aligns_non_decode_requests_to_page_size(): + queue = PDDecodeQueue.__new__(PDDecodeQueue) + queue.args = SimpleNamespace(page_size=16) + queue.dp_index = 0 + queue.is_busy = lambda: False + reqs = [ + SimpleNamespace( + request_id="req-0", + input_len=20, + sample_params=SimpleNamespace(max_new_tokens=1, suggested_dp_index=0), + is_infer_decode=lambda: False, + ), + SimpleNamespace( + request_id="req-1", + input_len=100, + sample_params=SimpleNamespace(max_new_tokens=1, suggested_dp_index=1), + is_infer_decode=lambda: False, + ), + ] + batch = Batch(batch_id=1, reqs=reqs, dp_size_in_node=2) + + assert queue._caclu_batch_estimated_peak_token_num(batch) == 32 + + +def test_pd_decode_uses_pd_decode_tuple_estimation(): + queue = PDDecodeQueue.__new__(PDDecodeQueue) + queue.args = SimpleNamespace(page_size=16) + queue.dp_index = 0 + queue.is_busy = lambda: False + queue.router = SimpleNamespace(router_statics=SimpleNamespace(ema_req_out_len=10)) + req = SimpleNamespace( + request_id="req-0", + sample_params=SimpleNamespace(suggested_dp_index=0), + is_infer_decode=lambda: True, + get_pd_decode_mode_tuple_tokens=lambda is_busy, ema_req_out_len: (16, 32), + ) + batch = Batch(batch_id=1, reqs=[req], dp_size_in_node=1) + + assert queue._caclu_batch_estimated_peak_token_num(batch) == 48 From 4f5409b7140da115a34c72c5c9eea53aad908a49 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 31 Aug 2026 02:03:24 +0000 Subject: [PATCH 11/15] fix: reserve request capacity for MTP overlap --- lightllm/server/httpserver/manager.py | 12 ++++++++++-- .../httpserver/test_running_request_lifecycle.py | 9 +++++++++ 2 files changed, 19 insertions(+), 2 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 9e6f77e58d..da5eb98347 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -619,13 +619,21 @@ async def _encode( def get_real_supported_max_req_total_len(self): # 得到系统真正能支持的最大长度,同时收到启动参数中模型支持长度的限制,也收到token容量的限制。 - return min(self.shm_max_total_token_num.get_value() - 36, self.max_req_total_len) + # MTP overlap 模式下,达到最大输出长度时,可能仍有部分 accepted token 已提交但停止状态尚未生效; + # 同时下一轮 overlap 还需要为 target verify 和 draft token 保留 KV 位置。因此按三倍 + # (mtp_step + 1) 预留额外 token,避免请求逻辑长度贴近总容量时发生 KV 申请失败。 + mtp_overlap_token_reserve = 3 * (self.args.mtp_step + 1) + return min( + self.shm_max_total_token_num.get_value() - 36 - mtp_overlap_token_reserve, + self.max_req_total_len, + ) async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params: SamplingParams): if not prompt_ids: raise ValueError("prompt_ids is empty") prompt_tokens = len(prompt_ids) - # 这里 -36 是保留一些不可预知的边界余量,防止系统出错 + # -36 用于保留通用边界余量,MTP overlap 所需的额外 KV 窗口由 + # get_real_supported_max_req_total_len 单独扣除。 real_supported_max_req_total_len = self.get_real_supported_max_req_total_len() if prompt_tokens + sampling_params.max_new_tokens > real_supported_max_req_total_len: diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..dcd4f6160e 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -209,3 +209,12 @@ async def run(): assert manager.run_reqs_count_mark.get_value() == 0 asyncio.run(run()) + + +def test_real_supported_max_req_total_len_reserves_mtp_overlap_tokens(): + manager = HttpServerManager.__new__(HttpServerManager) + manager.shm_max_total_token_num = _ValueMark(1000) + manager.max_req_total_len = 2000 + manager.args = SimpleNamespace(mtp_step=2) + + assert manager.get_real_supported_max_req_total_len() == 955 From e4fca207ed332f81a76f5aefadae50e402c51f24 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 31 Aug 2026 02:10:40 +0000 Subject: [PATCH 12/15] fix: reserve page capacity for request admission --- lightllm/server/httpserver/manager.py | 4 +++- .../server/httpserver/test_running_request_lifecycle.py | 4 ++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index da5eb98347..fe82689f39 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -623,8 +623,10 @@ def get_real_supported_max_req_total_len(self): # 同时下一轮 overlap 还需要为 target verify 和 draft token 保留 KV 位置。因此按三倍 # (mtp_step + 1) 预留额外 token,避免请求逻辑长度贴近总容量时发生 KV 申请失败。 mtp_overlap_token_reserve = 3 * (self.args.mtp_step + 1) + # 调度器会把单请求的 KV 资源向上扩展到 page_size 的整数倍。额外预留一个页面, + # 可以在输入阶段截断物理容量不足的请求,避免请求进入等待队列后始终无法被调度。 return min( - self.shm_max_total_token_num.get_value() - 36 - mtp_overlap_token_reserve, + self.shm_max_total_token_num.get_value() - 36 - mtp_overlap_token_reserve - self.args.page_size, self.max_req_total_len, ) diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index dcd4f6160e..c0f1dd6061 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -215,6 +215,6 @@ def test_real_supported_max_req_total_len_reserves_mtp_overlap_tokens(): manager = HttpServerManager.__new__(HttpServerManager) manager.shm_max_total_token_num = _ValueMark(1000) manager.max_req_total_len = 2000 - manager.args = SimpleNamespace(mtp_step=2) + manager.args = SimpleNamespace(mtp_step=2, page_size=4) - assert manager.get_real_supported_max_req_total_len() == 955 + assert manager.get_real_supported_max_req_total_len() == 951 From 13184d72bd7d2c92a3ff0f5e7827d58cda1eb070 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 31 Aug 2026 02:22:33 +0000 Subject: [PATCH 13/15] refactor: simplify request peak token estimation --- lightllm/server/core/objs/req.py | 51 +------------------ .../chunked_prefill/impl_for_pd_decode.py | 2 +- unit_tests/server/core/objs/test_req.py | 17 +------ .../req_queue/test_pd_queue_selection.py | 4 +- 4 files changed, 6 insertions(+), 68 deletions(-) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 1e58d883a7..55e00f40a7 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -3,7 +3,6 @@ import ctypes import asyncio import numpy as np -import triton import time from .sampling_params import SamplingParams from .out_token_circlequeue import CircularQueue @@ -424,9 +423,6 @@ def get_used_tokens(self): def get_tuple_tokens(self, is_busy, ema_req_out_len): raise NotImplementedError("Subclasses should implement this method") - def get_pd_decode_mode_tuple_tokens(self, is_busy, ema_req_out_len): - raise NotImplementedError("Subclasses should implement this method") - def get_output_logprobs_metadata(self, src_index: int, tokenizer=None): token_id = int(self.shm_prompt_ids.arr[src_index]) rank = int(self.shm_logprobs.arr["rank"][src_index]) @@ -454,26 +450,11 @@ def print_time_log(self, log_info: str): return -# 由于目前加入了很多异步调度的方法,为了缓解异步调度带来的很多 -# 估计不准确的问题,通过加长输出的长度,进行偏向保守一些的调度 -# 理论上不会多估计太多的 token 占用量, 同时得到较高的token显存 -# 使用率 -ADDED_OUTPUT_LEN = 16 - - class ChunkedPrefillReq(Req): _pack_ = 4 def get_tuple_tokens(self, is_busy, ema_req_out_len): args = get_env_start_args() - # chuncked prefill 推理的过程中,存在很多模式的延迟 step 推理的控制, 用于 - # 保证更好的包间数据或者是提升 dp 模式下prefill 的效率,但是在估计 token 显存 - # 占用量的过程中,分chuncked 需要考虑其因为分 chuncked带来的生命期的延长,具体 - # 体现就是在 b_len 的计算中,xxx * (max_waiting_token + 1) 的部分,这部分 - # 就是通过模拟加长其输出token长度,来延长其在估计阶段的生命周期。max_waiting_token - # 的计算是保守的,每次chuncked prefill 延迟的最大步数为两种模式之合,因为 - # 这个并不会导致预估的token占用量大幅增加,所以可以放心使用。 - max_waiting_token = args.router_max_wait_tokens has_out_len = self.shm_cur_output_len if self.sample_params.ignore_eos: cur_max_new_token_len = self.sample_params.max_new_tokens @@ -483,36 +464,6 @@ def get_tuple_tokens(self, is_busy, ema_req_out_len): cur_max_new_token_len = min(self.sample_params.max_new_tokens, max(int(1.1 * has_out_len), ema_req_out_len)) a_len = max(self.input_len + has_out_len + 1, self.shm_cur_kv_len + 1) - b_len = ( - (self.input_len + has_out_len - self.shm_cur_kv_len + self.chunked_prefill_size - 1) - // self.chunked_prefill_size - * (max_waiting_token + 1) - + cur_max_new_token_len - - has_out_len - - 1 - ) - b_len = max(0, b_len) + ADDED_OUTPUT_LEN - b_len = (b_len + args.page_size - 1) // args.page_size * args.page_size + b_len = max(0, cur_max_new_token_len - has_out_len - 1) + args.page_size return (a_len, b_len) - - def get_pd_decode_mode_tuple_tokens(self, is_busy, ema_req_out_len): - args = get_env_start_args() - has_out_len = self.shm_cur_output_len - if self.sample_params.ignore_eos: - cur_max_new_token_len = self.sample_params.max_new_tokens - elif is_busy: - cur_max_new_token_len = self.sample_params.max_new_tokens - else: - cur_max_new_token_len = min( - self.sample_params.max_new_tokens, - max(int(1.1 * has_out_len), ema_req_out_len), - ) - - # PD decode 节点只运行 decode,不需要考虑 chunked prefill 带来的调度等待时间。 - # 当前占用和预计剩余增长分别按 page_size 对齐,使后续峰值计算直接使用物理 KV 容量。 - a_len = max(self.input_len + has_out_len + 1, self.shm_cur_kv_len + 1) - a_len = triton.cdiv(a_len, args.page_size) * args.page_size - b_len = max(0, cur_max_new_token_len - has_out_len - 1) + ADDED_OUTPUT_LEN - b_len = triton.cdiv(b_len, args.page_size) * args.page_size - return (a_len, b_len) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py index 361ed54417..06d8dcf357 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py @@ -42,7 +42,7 @@ def _caclu_batch_estimated_peak_token_num(self, batch: Batch): # 请求进入 decode 阶段后,可以结合已经运行的 token 数量和预计剩余输出长度, # 使用连续批处理峰值算法估算其动态 KV 占用。 decoding_req_list.append( - req.get_pd_decode_mode_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) + req.get_tuple_tokens(is_busy, self.router.router_statics.ema_req_out_len) ) else: # 尚未进入 decode 阶段的请求没有足够的动态信息,仍按输入长度加最大输出长度 diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 694aae24ca..78908087c4 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -18,7 +18,6 @@ def setup_module_env(): "enable_cpu_cache": False, "model_dir": "", "page_size": 4, - "router_max_wait_tokens": 1, } ) ) @@ -74,27 +73,15 @@ def test_final_token_metadata_read_returns_actual_prompt_tokens(req): ] -def test_chunked_req_get_tuple_tokens_aligns_remaining_len_to_page(): +def test_chunked_req_get_tuple_tokens_adds_page_reserve(): req = SimpleNamespace( input_len=10, shm_cur_output_len=0, shm_cur_kv_len=0, - chunked_prefill_size=4, sample_params=SimpleNamespace(ignore_eos=True, max_new_tokens=5), ) - assert ChunkedPrefillReq.get_tuple_tokens(req, False, 10) == (11, 28) - - -def test_chunked_req_get_pd_decode_mode_tuple_tokens_uses_decode_estimation(): - req = SimpleNamespace( - input_len=10, - shm_cur_output_len=3, - shm_cur_kv_len=12, - sample_params=SimpleNamespace(ignore_eos=True, max_new_tokens=20), - ) - - assert ChunkedPrefillReq.get_pd_decode_mode_tuple_tokens(req, False, 10) == (16, 32) + assert ChunkedPrefillReq.get_tuple_tokens(req, False, 10) == (11, 8) def test_finish_status(req): diff --git a/unit_tests/server/router/req_queue/test_pd_queue_selection.py b/unit_tests/server/router/req_queue/test_pd_queue_selection.py index 1d8bc9a8dd..91144167c7 100644 --- a/unit_tests/server/router/req_queue/test_pd_queue_selection.py +++ b/unit_tests/server/router/req_queue/test_pd_queue_selection.py @@ -78,7 +78,7 @@ def test_pd_decode_aligns_non_decode_requests_to_page_size(): assert queue._caclu_batch_estimated_peak_token_num(batch) == 32 -def test_pd_decode_uses_pd_decode_tuple_estimation(): +def test_pd_decode_uses_tuple_estimation(): queue = PDDecodeQueue.__new__(PDDecodeQueue) queue.args = SimpleNamespace(page_size=16) queue.dp_index = 0 @@ -88,7 +88,7 @@ def test_pd_decode_uses_pd_decode_tuple_estimation(): request_id="req-0", sample_params=SimpleNamespace(suggested_dp_index=0), is_infer_decode=lambda: True, - get_pd_decode_mode_tuple_tokens=lambda is_busy, ema_req_out_len: (16, 32), + get_tuple_tokens=lambda is_busy, ema_req_out_len: (16, 32), ) batch = Batch(batch_id=1, reqs=[req], dp_size_in_node=1) From 3cee082635abda694277d4d098cc68d442b2fb8e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 31 Aug 2026 02:46:18 +0000 Subject: [PATCH 14/15] refactor: simplify PD prefill token estimation --- .../req_queue/chunked_prefill/impl_for_pd_prefill.py | 9 ++++----- .../server/router/req_queue/test_pd_queue_selection.py | 2 +- 2 files changed, 5 insertions(+), 6 deletions(-) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py index c6d8212f54..8c80a8296b 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_prefill.py @@ -1,5 +1,4 @@ import uuid -import triton from typing import Tuple from ...batch import Batch, Req from lightllm.server.router.req_queue.base_queue import BaseQueue @@ -14,13 +13,13 @@ def __init__(self, args, router, dp_index, dp_size_in_node) -> None: super().__init__(args, router, dp_index, dp_size_in_node) logger.info( "PD prefill requests normally generate only one output token; " - "estimate peak KV usage by summing their page-aligned token counts" + "estimate peak KV usage by adding one page to each request and summing the token counts" ) # @calculate_time(show=True, min_cost_ms=0.1) def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_num: int) -> Tuple[bool, int, int]: req_token_num = req.input_len + req.sample_params.max_new_tokens - req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + req_token_num += self.args.page_size estimated_peak_token_num += req_token_num ok_token_num = estimated_peak_token_num < self.max_total_tokens batch_req_num += 1 @@ -42,9 +41,9 @@ def _caclu_batch_estimated_peak_token_num(self, batch: Batch): for req in batch.reqs: if req.sample_params.suggested_dp_index == self.dp_index: # PD prefill 请求通常只生成一个 token,其 KV 占用不会像 decode 请求一样持续增长, - # 因此将每个请求按 page_size 对齐后的 token 数量直接线性相加即可完成估算。 + # 因此为每个请求额外增加一个 page_size 后直接线性相加即可完成估算。 req_token_num = req.input_len + req.sample_params.max_new_tokens - req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + req_token_num += self.args.page_size estimated_peak_token_num += req_token_num return estimated_peak_token_num diff --git a/unit_tests/server/router/req_queue/test_pd_queue_selection.py b/unit_tests/server/router/req_queue/test_pd_queue_selection.py index 91144167c7..16e57bac7d 100644 --- a/unit_tests/server/router/req_queue/test_pd_queue_selection.py +++ b/unit_tests/server/router/req_queue/test_pd_queue_selection.py @@ -51,7 +51,7 @@ def test_pd_prefill_peak_tokens_do_not_use_decode_estimation(): ] batch = Batch(batch_id=1, reqs=reqs, dp_size_in_node=2) - assert queue._caclu_batch_estimated_peak_token_num(batch) == 48 + assert queue._caclu_batch_estimated_peak_token_num(batch) == 64 def test_pd_decode_aligns_non_decode_requests_to_page_size(): From 1006160522c837e322af08260b3b53223ebe60ec Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 31 Aug 2026 03:02:58 +0000 Subject: [PATCH 15/15] refactor: simplify PD decode token estimation --- .../router/req_queue/chunked_prefill/impl_for_pd_decode.py | 5 ++--- .../server/router/req_queue/test_pd_queue_selection.py | 4 ++-- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py index 06d8dcf357..d31501f0e1 100644 --- a/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py @@ -1,6 +1,5 @@ import uuid import numpy as np -import triton from typing import Tuple from ...batch import Batch, Req from lightllm.server.router.req_queue.base_queue import BaseQueue @@ -15,7 +14,7 @@ def _can_add_new_req(self, req: Req, estimated_peak_token_num: int, batch_req_nu # 新请求尚未进入 decode 阶段,缺少实际输出长度等运行信息,只能按输入长度加最大输出长度 # 保守估算该请求最多可能占用的 KV 资源。 req_token_num = req.input_len + req.sample_params.max_new_tokens - req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + req_token_num += self.args.page_size estimated_peak_token_num += req_token_num ok_token_num = estimated_peak_token_num < self.max_total_tokens batch_req_num += 1 @@ -48,7 +47,7 @@ def _caclu_batch_estimated_peak_token_num(self, batch: Batch): # 尚未进入 decode 阶段的请求没有足够的动态信息,仍按输入长度加最大输出长度 # 预留其最大 KV 资源。 req_token_num = req.input_len + req.sample_params.max_new_tokens - req_token_num = triton.cdiv(req_token_num, self.args.page_size) * self.args.page_size + req_token_num += self.args.page_size estimated_peak_token_num += req_token_num if decoding_req_list: diff --git a/unit_tests/server/router/req_queue/test_pd_queue_selection.py b/unit_tests/server/router/req_queue/test_pd_queue_selection.py index 16e57bac7d..f27a3ec570 100644 --- a/unit_tests/server/router/req_queue/test_pd_queue_selection.py +++ b/unit_tests/server/router/req_queue/test_pd_queue_selection.py @@ -54,7 +54,7 @@ def test_pd_prefill_peak_tokens_do_not_use_decode_estimation(): assert queue._caclu_batch_estimated_peak_token_num(batch) == 64 -def test_pd_decode_aligns_non_decode_requests_to_page_size(): +def test_pd_decode_adds_page_reserve_for_non_decode_requests(): queue = PDDecodeQueue.__new__(PDDecodeQueue) queue.args = SimpleNamespace(page_size=16) queue.dp_index = 0 @@ -75,7 +75,7 @@ def test_pd_decode_aligns_non_decode_requests_to_page_size(): ] batch = Batch(batch_id=1, reqs=reqs, dp_size_in_node=2) - assert queue._caclu_batch_estimated_peak_token_num(batch) == 32 + assert queue._caclu_batch_estimated_peak_token_num(batch) == 37 def test_pd_decode_uses_tuple_estimation():