diff --git a/docs/kv_cache_page_size.md b/docs/kv_cache_page_size.md new file mode 100644 index 0000000000..a89df89afc --- /dev/null +++ b/docs/kv_cache_page_size.md @@ -0,0 +1,45 @@ +# 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. `req_to_token_indexs` 保存请求拥有的完整页;创建 `InferState` 时根据请求索引和序列位置直接聚合本轮 + 真实参与计算的 token,`ModelInput` 不携带物理 KV 索引,`InferState.mem_index` 不包含预留尾部。 +4. Radix Cache 只插入、拆分、命中和淘汰完整页;不足一页的请求尾部在请求结束或暂停时整页回收。 +5. `page_size=1` 与多 token 页使用相同的调度期预留和模型执行期索引选择路径。 + +页容量计算和整页申请由调度层在请求获准进入 Prefill/Decode batch 时完成;输入构造层只构造本轮输入, +真实推理索引由模型执行层从请求表选取。 +`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..1081c94d7f 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -15,9 +15,11 @@ 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 +from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import ( + 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 @@ -123,6 +125,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() @@ -371,13 +381,29 @@ def _init_hidden_collector(self): @torch.no_grad() def forward(self, model_input: ModelInput): model_input.to_cuda() - 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_mem_indexes(self, model_input: ModelInput): + if model_input.is_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, + 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], + ) + 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() infer_state.hidden_collector = self.hidden_collector_prototype.new_instance() @@ -408,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) @@ -453,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, - ) new_model_input.multimodal_params = new_model_input.multimodal_params + [ {"images": [], "audios": []} for _ in range(padded_batch_size) ] @@ -496,12 +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) - 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, - ) 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 ) @@ -589,15 +603,6 @@ 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, - ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -657,12 +662,6 @@ 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, - ) infer_state.init_some_extra_state(self) infer_state.init_att_state() @@ -793,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] @@ -818,28 +814,10 @@ 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, - ) 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, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -888,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 @@ -904,23 +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 - 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, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -942,22 +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) - 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, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -1099,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") @@ -1113,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, @@ -1178,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") @@ -1192,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, @@ -1241,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 @@ -1258,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, @@ -1277,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 ae645d4b7b..bf73a10f91 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -32,7 +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 is_prefill: bool = False b_ready_cache_len: torch.Tensor = None # Request/row-aligned MRoPE position offset. It is decode-only; prefill @@ -41,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 不使用。 @@ -59,8 +56,6 @@ def to_cuda(self): self.check_input() # Prefill 和 decode 都必须提供的公共张量。 - if self.mem_indexes is 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) @@ -94,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 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 @@ -123,10 +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 - 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 5849cccf54..bbb661af6b 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -255,7 +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() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) @@ -270,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, @@ -284,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 @@ -316,7 +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() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) @@ -333,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 bf6039a48f..28d44fefbf 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -194,7 +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() 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) @@ -210,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, @@ -225,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 @@ -255,7 +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() 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) @@ -271,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/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..ad995b2b5d 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -324,7 +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() - 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/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/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/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/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/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/req.py b/lightllm/server/core/objs/req.py index 87f54fd9e7..55e00f40a7 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]) @@ -456,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 @@ -485,33 +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 = max(0, cur_max_new_token_len - has_out_len - 1) + 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: - # 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 - - 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/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/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 9e6f77e58d..fe82689f39 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -619,13 +619,23 @@ 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) + # 调度器会把单请求的 KV 资源向上扩展到 page_size 的整数倍。额外预留一个页面, + # 可以在输入阶段截断物理容量不足的请求,避免请求进入等待队列后始终无法被调度。 + return min( + self.shm_max_total_token_num.get_value() - 36 - mtp_overlap_token_reserve - self.args.page_size, + 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/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/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..3b128312ec 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -125,7 +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: - free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_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) @@ -133,18 +133,22 @@ 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]) + 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 +534,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 @@ -570,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 @@ -640,6 +643,8 @@ 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 + 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 @@ -905,22 +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() - - seq_len = len(input_token_ids) - input_token_len = seq_len - self.cur_kv_len - return input_token_len + 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 1 - - def _mtp_decode_need_token_num(self) -> int: - return (1 + self.mtp_step) * 2 + 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: + page_size = self.args.page_size + target_hold_len = (target_kv_len + page_size - 1) // page_size * page_size + 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 28f2abf74b..8a28f44652 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: @@ -404,10 +405,14 @@ 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: - 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: @@ -437,7 +442,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, ) @@ -629,6 +634,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, @@ -717,10 +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: + 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 @@ -732,13 +755,20 @@ 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) + 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 - if token_num <= can_alloc_token_num: + 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 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 +981,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/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 4d09476849..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,12 +376,6 @@ 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, - ), - ) 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..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,13 +603,6 @@ 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, - ), - ) - select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, @@ -889,19 +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, ) - proposal.extra_mem_indexes_cpu.extend( - ( - MtpMemIndexesToFree( - mem_indexes_cpu=model_input0.mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu0 == 0, - ), - 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 22731439c4..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 @@ -65,11 +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 - 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]) - model_input = ModelInput( batch_size=b_seq_len.shape[0], total_token_num=total_token_num, @@ -77,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_cpu=mem_indexes, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, @@ -140,18 +134,12 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In device="cpu", ) - # 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]) - 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_cpu=mem_indexes, 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 9d59afd8a7..83a73c7d9e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -81,20 +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 - - # 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) 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/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/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 9af7afd1b4..90d4b0cbc0 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 962c9c31cb..57296bf22f 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,27 +43,19 @@ 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() ok_token_num = need_max_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(need_max_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( need_max_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): @@ -85,16 +77,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] # 等待判断的组 @@ -103,9 +92,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) @@ -144,7 +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) - return ( - need_max_token_num, - need_max_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 f8fe510989..7723ee51b3 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 @@ -37,22 +33,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(need_max_token_num, self.dp_index) self.router.shared_token_load.set_dynamic_max_load( need_max_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): @@ -71,10 +60,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 @@ -82,9 +67,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) @@ -109,7 +92,4 @@ 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, - ) + return (need_max_token_num, need_max_token_num / self.max_total_tokens) 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..d31501f0e1 --- /dev/null +++ b/lightllm/server/router/req_queue/chunked_prefill/impl_for_pd_decode.py @@ -0,0 +1,107 @@ +import uuid +import numpy as np +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 += 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_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 += 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..8c80a8296b 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,26 @@ import uuid -import numpy as np 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 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]: - 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 += 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 +36,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 后直接线性相加即可完成估算。 + req_token_num = req.input_len + req.sample_params.max_new_tokens + req_token_num += self.args.page_size + estimated_peak_token_num += req_token_num return estimated_peak_token_num 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/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/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 056513bdde..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,13 +118,12 @@ 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": []}], ) 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), ) @@ -154,14 +151,13 @@ 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=[], 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), ) @@ -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,12 +189,11 @@ 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=[], ) 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), ) @@ -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 6f9477e294..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), @@ -116,9 +113,9 @@ 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) + 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 = [] @@ -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 b286998e04..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=[], ) @@ -92,7 +89,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 +135,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()), ) @@ -174,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() @@ -182,11 +179,11 @@ 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()) - 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): @@ -194,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, ) @@ -204,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_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/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..0dab062934 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_select_kv_index_from_req.py @@ -0,0 +1,75 @@ +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, + 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 = model._select_mem_indexes(model_input) + + 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 new file mode 100644 index 0000000000..ed508fbcdd --- /dev/null +++ b/unit_tests/common/test_req_manager_page.py @@ -0,0 +1,166 @@ +from types import SimpleNamespace + +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: + 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) + 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): + req = SimpleNamespace(req_idx=req_idx, cur_kv_len=0, hold_kv_len=0) + req._kv_cache_alloc_need = lambda target_len: InferReq._kv_cache_alloc_need(req, target_len) + return req + + +def test_request_reuses_reserved_page_tail_before_allocating_next_page(monkeypatch): + context, backend = _make_context(monkeypatch) + req = _make_req(0) + + 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, alloc_token_num=0) + assert context.req_manager.mem_manager.alloc_sizes == [4] + + req.cur_kv_len = 4 + 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)) + + +def test_reservation_fills_each_request_table_row(monkeypatch): + _, backend = _make_context(monkeypatch) + req0 = _make_req(0) + req1 = _make_req(1) + + 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) + 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.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 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 + 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.decode_need_token_num(req) + backend._alloc_req_kv_mem(req, alloc_token_num) + + model_input, _ = generic_pre_process.prepare_decode_inputs([req]) + + 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(): + 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 = [] + + infer_context.free_a_req_mem(free_token_indexes, req) + + assert free_token_indexes[0].tolist() == list(range(12)) + assert req.cur_kv_len == req.hold_kv_len == 0 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/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 51f8d1c82c..78908087c4 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 @@ -16,6 +17,7 @@ def setup_module_env(): "cpu_cache_token_page_size": 256, "enable_cpu_cache": False, "model_dir": "", + "page_size": 4, } ) ) @@ -71,11 +73,15 @@ 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_adds_page_reserve(): + req = SimpleNamespace( + input_len=10, + shm_cur_output_len=0, + shm_cur_kv_len=0, + sample_params=SimpleNamespace(ignore_eos=True, max_new_tokens=5), + ) + + assert ChunkedPrefillReq.get_tuple_tokens(req, False, 10) == (11, 8) def test_finish_status(req): diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..c0f1dd6061 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, page_size=4) + + assert manager.get_real_supported_max_req_total_len() == 951 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_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 62296634a9..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 @@ -10,7 +10,12 @@ def _patch_empty_input_context(monkeypatch): alloc=lambda size: torch.empty((size,), dtype=torch.int32), ) infer_context = SimpleNamespace( - req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager), + args=SimpleNamespace(page_size=1), + 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) @@ -26,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, @@ -37,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, @@ -53,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.shape == (0,) 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 == [] @@ -69,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.shape == (0,) 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,) @@ -169,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.shape == (0,) 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 2e24262e7e..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, @@ -112,17 +109,12 @@ 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), ) 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,34 +140,22 @@ 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, 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), ) 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), @@ -207,7 +187,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: ( @@ -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( @@ -314,17 +289,12 @@ 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), ) 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 7baf34061b..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,8 +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), ) plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=1, pre_draft_step=1) @@ -715,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"], 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..f27a3ec570 --- /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) == 64 + + +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 + 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) == 37 + + +def test_pd_decode_uses_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_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 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()