From d136a483ac91b4abccdddef799d5c8f6afe0a3bb Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Wed, 9 Sep 2026 15:25:56 +0000 Subject: [PATCH 1/4] refactor: generalize hybrid attention checkpoint management --- .../linear_att_cpu_cache_copy.py | 2 +- .../operator/linear_att.py | 26 +- .../qwen3next_mem_manager.py | 21 +- .../linear_att_cache_manager/__init__.py | 3 - .../linear_att_buffer_manager.py | 83 ----- lightllm/common/req_manager/__init__.py | 3 +- lightllm/common/req_manager/hybrid_base.py | 70 +++++ lightllm/common/req_manager/linear_att.py | 72 +++-- .../common/state_cache_manager/__init__.py | 13 + lightllm/common/state_cache_manager/base.py | 57 ++++ .../layer_cache.py | 0 .../linear_att.py} | 45 +++ lightllm/models/qwen3next/model.py | 2 +- lightllm/server/api_start.py | 12 +- lightllm/server/core/objs/req.py | 32 +- .../core/objs/token_chunck_hash_list.py | 2 +- lightllm/server/pd_io_struct.py | 5 +- ...dix_cache.py => hybrid_att_radix_cache.py} | 138 ++++----- .../server/router/model_infer/infer_batch.py | 285 ++++++++---------- .../model_infer/mode_backend/base_backend.py | 20 +- .../mode_backend/chunked_prefill/impl.py | 4 +- .../mode_backend/dp_backend/impl.py | 12 +- .../mode_backend/multi_level_kv_cache.py | 17 +- .../pd/decode_node_impl/decode_impl.py | 15 +- .../pd/prefill_node_impl/prefill_impl.py | 8 +- .../model_infer/mtp_speculative/utils.py | 6 +- lightllm/utils/backend_validator.py | 2 +- lightllm/utils/config_utils.py | 5 + lightllm/utils/kv_cache_utils.py | 25 +- test/cpu_cache_kernel/test_speed.py | 2 +- .../basemodel/attention/linear/test_gdn.py | 2 +- .../mode_backend/test_multi_level_kv_cache.py | 20 +- 32 files changed, 550 insertions(+), 459 deletions(-) delete mode 100644 lightllm/common/linear_att_cache_manager/__init__.py delete mode 100644 lightllm/common/linear_att_cache_manager/linear_att_buffer_manager.py create mode 100644 lightllm/common/req_manager/hybrid_base.py create mode 100644 lightllm/common/state_cache_manager/__init__.py create mode 100644 lightllm/common/state_cache_manager/base.py rename lightllm/common/{linear_att_cache_manager => state_cache_manager}/layer_cache.py (100%) rename lightllm/common/{linear_att_cache_manager/config_objs.py => state_cache_manager/linear_att.py} (80%) rename lightllm/server/router/dynamic_prompt/{linear_att_radix_cache.py => hybrid_att_radix_cache.py} (82%) diff --git a/lightllm/common/basemodel/triton_kernel/linear_att_cpu_cache_copy.py b/lightllm/common/basemodel/triton_kernel/linear_att_cpu_cache_copy.py index ed0e742d73..b0c9a60649 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att_cpu_cache_copy.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att_cpu_cache_copy.py @@ -1,7 +1,7 @@ import torch import triton import triton.language as tl -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import LinearAttCacheConfig @triton.jit diff --git a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py index 49e2549265..71158ac97a 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -6,7 +6,7 @@ from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.dist_utils import get_current_rank_in_dp, get_dp_world_size from lightllm.utils.log_utils import init_logger -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import LinearAttCacheConfig if TYPE_CHECKING: from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuKvCacheClient @@ -46,9 +46,9 @@ def load_cpu_cache_to_gpu( big_page_buffer_ids_cpu = [] for i in range(big_page_num): - page_id = mem_manager.linear_att_big_page_buffers.alloc_one_state_cache() + page_id = mem_manager.big_page_buffers.alloc_one_state_cache() assert page_id is not None - req.linear_att_len_to_big_page_id[max_kv_len] = page_id + req.hybrid_len_to_big_page_id[max_kv_len] = page_id big_page_buffer_ids_cpu.append(page_id) max_kv_len -= args.cpu_cache_token_page_size assert max_kv_len % args.cpu_cache_token_page_size == 0 @@ -81,8 +81,8 @@ def load_cpu_cache_to_gpu( big_page_buffer_ids=big_page_buffer_ids_gpu, page_indexes=page_indexes, gpu_full_att_kv_state=mem_manager.kv_buffer, - cpu_kv_conv_state=mem_manager.linear_att_big_page_buffers.conv_state_cache.buffer, - cpu_kv_ssm_state=mem_manager.linear_att_big_page_buffers.ssm_state_cache.buffer, + cpu_kv_conv_state=mem_manager.big_page_buffers.conv_state_cache.buffer, + cpu_kv_ssm_state=mem_manager.big_page_buffers.ssm_state_cache.buffer, cpu_cache_tensor=cpu_cache_client.cpu_kv_cache_tensor, tp_rank=get_current_rank_in_dp(), tp_world_size=get_dp_world_size(), @@ -92,7 +92,7 @@ def load_cpu_cache_to_gpu( from lightllm.server.router.model_infer.infer_batch import g_infer_context - g_infer_context.req_manager.copy_big_page_buffer_to_linear_att_state( + g_infer_context.req_manager.restore_big_page_state( big_page_buffer_idx=big_page_buffer_ids_cpu[-1], req=req, ) @@ -131,7 +131,7 @@ def offload_gpu_kv_to_cpu_cache( max_kv_len = (len(mem_indexes) // args.cpu_cache_token_page_size) * args.cpu_cache_token_page_size start_kv_len = (len(big_page_buffer_ids_cpu) + 1) * args.cpu_cache_token_page_size for seq_len in range(start_kv_len, max_kv_len + 1, args.cpu_cache_token_page_size): - page_id = req.linear_att_len_to_big_page_id[seq_len] + page_id = req.hybrid_len_to_big_page_id[seq_len] big_page_buffer_ids_cpu.append(page_id) if len(mem_indexes) % args.cpu_cache_token_page_size != 0: @@ -141,15 +141,15 @@ def offload_gpu_kv_to_cpu_cache( dst_mem_indexes = self.mem_indexes_buffer[0:dst_len].fill_(-1) dst_mem_indexes[0 : len(mem_indexes)].copy_(mem_indexes, non_blocking=True) mem_indexes = dst_mem_indexes - assert req.tail_linear_att_small_page_buffer_id is not None + assert req.tail_small_page_buffer_id is not None from lightllm.common.basemodel.triton_kernel.linear_att_cpu_cache_copy import ( copy_linear_att_state_to_linear_att_state, ) - src_conv_state, src_ssm_state = g_infer_context.radix_cache.linear_att_small_page_buffers.get_state_cache( - buffer_idx=req.tail_linear_att_small_page_buffer_id + src_conv_state, src_ssm_state = g_infer_context.radix_cache.small_page_buffers.get_state_cache( + buffer_idx=req.tail_small_page_buffer_id ) - dst_conv_state, dst_ssm_state = mem_manager.linear_att_big_page_buffers.get_state_cache( + dst_conv_state, dst_ssm_state = mem_manager.big_page_buffers.get_state_cache( buffer_idx=mem_manager.CPU_CACHE_BIG_PAGE_OFFLOAD_TEMP_BUFFER_ID, ) copy_linear_att_state_to_linear_att_state( @@ -179,8 +179,8 @@ def offload_gpu_kv_to_cpu_cache( page_readies=page_readies, big_page_buffer_ids=big_page_buffer_ids_gpu, gpu_kv_full_att_state=mem_manager.kv_buffer, - cpu_kv_conv_state=mem_manager.linear_att_big_page_buffers.conv_state_cache.buffer, - cpu_kv_ssm_state=mem_manager.linear_att_big_page_buffers.ssm_state_cache.buffer, + cpu_kv_conv_state=mem_manager.big_page_buffers.conv_state_cache.buffer, + cpu_kv_ssm_state=mem_manager.big_page_buffers.ssm_state_cache.buffer, cpu_cache_tensor=cpu_cache_client.cpu_kv_cache_tensor, tp_rank=get_current_rank_in_dp(), tp_world_size=get_dp_world_size(), diff --git a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py index 907cc494a6..384269cb17 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -1,9 +1,10 @@ import torch +from lightllm.server.pd_io_struct import HYBRID_ATT_STATE_PAGE_KIND import triton from lightllm.utils.log_utils import init_logger from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager from lightllm.utils.envs_utils import get_env_start_args -from lightllm.common.linear_att_cache_manager import LinearAttCacheConfig, LinearAttCacheManager +from lightllm.common.state_cache_manager import LinearAttCacheConfig, LinearAttCacheManager from .operator import LinearAttMemOperator from typing import Tuple, Any, List @@ -45,14 +46,14 @@ def _init_linear_att_buffers(self): # 申请大页可能需要对应的资源, 多申请了两个linear att的状态,理论上这个状态 # 永远不会被 alloc 申请到,只会在 cpu cache中,用于过渡和存储碎页情况下的 # cpu cache 的页面拷贝。 - self.linear_att_big_page_buffers = LinearAttCacheManager( + self.big_page_buffers = LinearAttCacheManager( size=triton.cdiv(self.size, big_page_token_num) + 2, linear_config=self.linear_config, keep_num=2, ) - self.CPU_CACHE_BIG_PAGE_LOAD_TEMP_BUFFER_ID = self.linear_att_big_page_buffers.size - 2 - self.CPU_CACHE_BIG_PAGE_OFFLOAD_TEMP_BUFFER_ID = self.linear_att_big_page_buffers.size - 1 + self.CPU_CACHE_BIG_PAGE_LOAD_TEMP_BUFFER_ID = self.big_page_buffers.size - 2 + self.CPU_CACHE_BIG_PAGE_OFFLOAD_TEMP_BUFFER_ID = self.big_page_buffers.size - 1 return def _free_buffers(self): @@ -61,7 +62,7 @@ def _free_buffers(self): return def _free_linear_att_buffers(self): - self.linear_att_big_page_buffers = None + self.big_page_buffers = None return def write_to_shm(self, req_manager): @@ -72,12 +73,12 @@ def write_to_shm(self, req_manager): # pinned(cudaHostAlloc) 的内存退化为普通 shm mmap,之后 Triton kernel 携带该指针 # 启动会报 "Pointer argument cannot be accessed from Triton (cpu tensor?)"。 # 跨进程消费方并不使用 cpu 侧大页 state cache,序列化期间临时剔除以保住 pinned。 - big_page_buffers = self.linear_att_big_page_buffers - self.linear_att_big_page_buffers = None + big_page_buffers = self.big_page_buffers + self.big_page_buffers = None try: return super().write_to_shm(req_manager) finally: - self.linear_att_big_page_buffers = big_page_buffers + self.big_page_buffers = big_page_buffers def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: kv_move_buffer = super().alloc_paged_kv_move_buffer(page_num, page_size) @@ -104,7 +105,7 @@ def write_mem_to_page_kv_move_buffer( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" + assert page_kind == HYBRID_ATT_STATE_PAGE_KIND, f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) @@ -131,7 +132,7 @@ def read_page_kv_move_buffer_to_mem( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" + assert page_kind == HYBRID_ATT_STATE_PAGE_KIND, f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) diff --git a/lightllm/common/linear_att_cache_manager/__init__.py b/lightllm/common/linear_att_cache_manager/__init__.py deleted file mode 100644 index ab3c8e2cd9..0000000000 --- a/lightllm/common/linear_att_cache_manager/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .linear_att_buffer_manager import LinearAttCacheManager -from .config_objs import LinearAttCacheConfig -from .layer_cache import LayerCache diff --git a/lightllm/common/linear_att_cache_manager/linear_att_buffer_manager.py b/lightllm/common/linear_att_cache_manager/linear_att_buffer_manager.py deleted file mode 100644 index 30dc4d937c..0000000000 --- a/lightllm/common/linear_att_cache_manager/linear_att_buffer_manager.py +++ /dev/null @@ -1,83 +0,0 @@ -import torch -import collections -from lightllm.utils.log_utils import init_logger -from .layer_cache import LayerCache -from typing import List, Optional, Tuple, Union -from .config_objs import LinearAttCacheConfig - -logger = init_logger(__name__) - - -class LinearAttCacheManager: - def __init__( - self, - size: int, - linear_config: LinearAttCacheConfig, - keep_num: int = 0, # 用于记录需要保留的缓存数量,用于支持含有 linear_att 的如qwen3.5 模型的cpu cache的碎页处理。 - ): - # init the mem state - self.size = size - self.linear_config = linear_config - self.keep_num = keep_num - assert 0 <= self.keep_num <= self.size, f"invalid keep_num {self.keep_num} for size {self.size}" - # init the layer cache - self.conv_state_cache = LayerCache( - size=self.size, - dtype=self.linear_config.conv_state_dtype, - shape=self.linear_config.get_conv_state_shape(), - layer_num=self.linear_config.linear_layer_num, - device="cpu", - size_first=True, - ) - self.ssm_state_cache = LayerCache( - size=self.size, - dtype=self.linear_config.ssm_state_dtype, - shape=self.linear_config.get_ssm_state_shape(), - layer_num=self.linear_config.linear_layer_num, - device="cpu", - size_first=True, - ) - self.clear_to_init_state() - return - - def get_state_cache(self, buffer_idx: int): - return self.conv_state_cache.buffer[buffer_idx, ...], self.ssm_state_cache.buffer[buffer_idx, ...] - - def alloc_one_state_cache(self) -> Optional[int]: - if len(self.free_list) == 0: - return None - - alloc_index = self.free_list.popleft() - return alloc_index - - def alloc_state_cache(self, need_size: int) -> Optional[List[int]]: - if need_size > len(self.free_list): - logger.error(f"warn no enough cache need_size {need_size} free_size {len(self.free_list)}") - return None - - alloc_indexes = [self.free_list.popleft() for _ in range(need_size)] - return alloc_indexes - - def free_state_cache(self, free_indexes: List[int]): - alloc_upper_bound = self.size - self.keep_num - for idx in free_indexes: - assert 0 <= idx < alloc_upper_bound, ( - f"free index {idx} out of alloc range [0, {alloc_upper_bound}), " f"reserved tail num {self.keep_num}" - ) - self.free_list.extend(free_indexes) - assert ( - len(self.free_list) <= alloc_upper_bound - ), f"free cache num {len(self.free_list)} should not be larger than alloc size {alloc_upper_bound}" - return - - def get_free_cache_num(self): - return len(self.free_list) - - def get_used_cache_num(self): - return self.size - len(self.free_list) - - def clear_to_init_state(self): - self.conv_state_cache.buffer.zero_() - self.ssm_state_cache.buffer.zero_() - self.free_list = collections.deque(range(self.size - self.keep_num)) - return diff --git a/lightllm/common/req_manager/__init__.py b/lightllm/common/req_manager/__init__.py index 078aa8c155..abcd4f491e 100644 --- a/lightllm/common/req_manager/__init__.py +++ b/lightllm/common/req_manager/__init__.py @@ -1,5 +1,6 @@ from .base import ReqManager from .linear_att import ReqManagerForMamba +from .hybrid_base import HybridAttentionReqManager from .req_sampling_params import ReqSamplingParamsManager -__all__ = ["ReqManager", "ReqManagerForMamba", "ReqSamplingParamsManager"] +__all__ = ["ReqManager", "HybridAttentionReqManager", "ReqManagerForMamba", "ReqSamplingParamsManager"] diff --git a/lightllm/common/req_manager/hybrid_base.py b/lightllm/common/req_manager/hybrid_base.py new file mode 100644 index 0000000000..948313116d --- /dev/null +++ b/lightllm/common/req_manager/hybrid_base.py @@ -0,0 +1,70 @@ +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, List + +import torch + +from .base import ReqManager + + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.infer_batch import InferReq + + +class HybridAttentionReqManager(ReqManager, ABC): + """混合 attention 的请求运行态与大小页 checkpoint 管理接口。 + + 大小页沿同一虚拟 token 索引空间匹配前缀,full attention KV 保持 token 粒度存储。 + linear/sliding-window 状态在大页边界及请求可缓存尾部的小页边界保存 checkpoint, + 缓存命中后,再将相应 checkpoint 恢复到请求运行态。 + + 请求的 GPU 计算状态由具体实现管理;small_page_buffers 持有 CPU 小页快照; + big_page_buffers 引用 mem_manager 的 CPU 大页快照,与 full KV 容量一起创建、调整。 + 大小页不包含额外的 GPU 运行态,命中后恢复到请求状态,不重新分配 checkpoint 池。 + 公共缓存流程负责 checkpoint 槽位分配、边界、匹配与淘汰;本接口负责运行态与 checkpoint 保存/恢复。 + """ + + def __init__(self, max_request_num, max_sequence_length, mem_manager): + super().__init__(max_request_num, max_sequence_length, mem_manager) + self.small_page_buffers = None + + @property + def big_page_buffers(self): + return self.mem_manager.big_page_buffers + + @abstractmethod + def create_small_page_cache_manager(self, size: int): + """创建并持有 CPU 小页池,返回同一池供 radix 使用;size 是槽位数,不是 token 数。""" + + @abstractmethod + def init_hybrid_attention_state(self, req: "InferReq"): + """无前缀缓存命中时,初始化已分配请求槽位的 GPU 运行态。""" + + def restore_big_page_state(self, big_page_buffer_idx: int, req: "InferReq"): + """将指定大页槽位的 CPU checkpoint 恢复到请求 GPU 运行态。""" + self.restore_state(req, self.big_page_buffers, big_page_buffer_idx) + + def restore_small_page_state(self, req: "InferReq"): + """将 req.shared_kv_node 对应的小页 checkpoint 恢复到请求 GPU 运行态。""" + self.restore_state(req, self.small_page_buffers, req.shared_kv_node.small_page_buffer_idx) + + @abstractmethod + def restore_state(self, req: "InferReq", state_cache_manager, buffer_idx: int): + """CPU checkpoint → 请求 GPU 运行态;大小页共用,不负责前缀匹配或 full KV 索引恢复。""" + + def save_big_page_states(self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int]): + """批量保存请求 GPU 运行态到已分配的大页槽位,buffer_indexes 中的 -1 表示跳过。 + + b_req_idx 与 req_indexes 分别为同一批请求的 GPU 索引张量和 CPU 索引列表。 + 默认逐请求保存,模型可覆盖为批量拷贝算子。 + """ + for req_idx, buffer_idx in zip(req_indexes, buffer_indexes): + if buffer_idx != -1: + self.save_state(req_idx, buffer_idx, self.big_page_buffers) + + @abstractmethod + def save_state(self, req_idx: int, buffer_idx: int, state_cache_manager): + """请求 GPU 运行态 → 指定 CPU checkpoint 槽位;大小页共用,调用方负责分配槽位。""" + + def update_mtp_state(self, b_req_mtp_start_loc, b_req_idx, b_mtp_index, accepted_index, verify_width): + """接受推测 token 后更新运行态位置;需要调整状态索引的模型覆写。""" + return diff --git a/lightllm/common/req_manager/linear_att.py b/lightllm/common/req_manager/linear_att.py index 967bc9bf7a..ccab8a06f6 100644 --- a/lightllm/common/req_manager/linear_att.py +++ b/lightllm/common/req_manager/linear_att.py @@ -1,20 +1,18 @@ -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, List import torch -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig -from lightllm.common.linear_att_cache_manager.layer_cache import LayerCache -from lightllm.common.linear_att_cache_manager.linear_att_buffer_manager import LinearAttCacheManager +from lightllm.common.state_cache_manager import LayerCache, LinearAttCacheConfig, LinearAttCacheManager from lightllm.utils.envs_utils import get_env_start_args -from .base import ReqManager +from .hybrid_base import HybridAttentionReqManager if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq -class ReqManagerForMamba(ReqManager): +class ReqManagerForMamba(HybridAttentionReqManager): def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_config: LinearAttCacheConfig): super().__init__(max_request_num, max_sequence_length, mem_manager) self.mtp_step = get_env_start_args().mtp_step @@ -50,7 +48,7 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_con ) return - def init_linear_att_state(self, req: "InferReq"): + def init_hybrid_attention_state(self, req: "InferReq"): conv_index = req.req_idx ssm_start = req.req_idx * (self.mtp_step + 1) self.req_to_conv_state.buffer[:, conv_index, ...].fill_(0) @@ -61,6 +59,35 @@ def init_linear_att_state(self, req: "InferReq"): self.req_to_mtp_state_index[req.req_idx] = 0 return + def create_small_page_cache_manager(self, size: int): + self.small_page_buffers = LinearAttCacheManager(size=size, linear_config=self.linear_config) + return self.small_page_buffers + + def save_big_page_states(self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int]): + from lightllm.common.basemodel.triton_kernel.linear_att_copy import copy_linear_att_state_to_kv_buffer + + buffer_indexes = torch.tensor(buffer_indexes, dtype=torch.int32, device="cpu").cuda(non_blocking=True) + state_cache_manager = self.big_page_buffers + copy_linear_att_state_to_kv_buffer( + b_req_idx=b_req_idx, + big_page_buffer_ids=buffer_indexes, + gpu_conv_state=self.req_to_conv_state.buffer, + gpu_ssm_state=self.req_to_ssm_state.buffer, + cpu_kv_conv_state=state_cache_manager.conv_state_cache.buffer, + cpu_kv_ssm_state=state_cache_manager.ssm_state_cache.buffer, + mtp_step=self.mtp_step, + ) + return + + def save_state(self, req_idx: int, buffer_idx: int, state_cache_manager: LinearAttCacheManager): + # checkpoint 只保存标准 conv 窗口和请求的基准 SSM 状态,不包含 MTP 扩展运行态。 + conv_cache_width = self.linear_config.get_conv_state_shape()[-1] + gpu_conv_state = self.req_to_conv_state.buffer[:, req_idx, ..., :conv_cache_width] + gpu_ssm_state = self.req_to_ssm_state.buffer[:, req_idx * (self.mtp_step + 1), ...] + dst_conv_state, dst_ssm_state = state_cache_manager.get_state_cache(buffer_idx=buffer_idx) + dst_conv_state.copy_(gpu_conv_state, non_blocking=True) + dst_ssm_state.copy_(gpu_ssm_state, non_blocking=True) + def get_mamba_cache(self, layer_idx_in_all: int): assert ( 0 <= layer_idx_in_all < self.linear_config.all_layer_num @@ -70,30 +97,23 @@ def get_mamba_cache(self, layer_idx_in_all: int): ssm_states = self.req_to_ssm_state.buffer[layer_idx_in_linear] return conv_states, ssm_states - def copy_big_page_buffer_to_linear_att_state(self, big_page_buffer_idx: int, req: "InferReq"): - big_page_buffers: LinearAttCacheManager = self.mem_manager.linear_att_big_page_buffers - - conv_state, ssm_state = big_page_buffers.get_state_cache(buffer_idx=big_page_buffer_idx) - conv_dest = req.req_idx - ssm_dest = req.req_idx * (self.mtp_step + 1) - conv_cache_width = conv_state.shape[-1] - self.req_to_conv_state.buffer[:, conv_dest, ..., :conv_cache_width] = conv_state - self.req_to_ssm_state.buffer[:, ssm_dest, ...] = ssm_state - if self.req_to_mtp_state_index is not None: - self.req_to_mtp_state_index[req.req_idx] = 0 - return + def update_mtp_state(self, b_req_mtp_start_loc, b_req_idx, b_mtp_index, accepted_index, verify_width): + from lightllm.common.basemodel.triton_kernel.mtp_utils import linear_att_mtp_state_index_update - def copy_small_page_buffer_to_linear_att_state( - self, req: "InferReq", linear_att_small_page_buffers: LinearAttCacheManager - ): - conv_state, ssm_state = linear_att_small_page_buffers.get_state_cache( - buffer_idx=req.shared_kv_node.small_page_buffer_idx + linear_att_mtp_state_index_update( + req_to_mtp_state_index=self.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + accepted_index=accepted_index, + verify_width=verify_width, ) + + def restore_state(self, req: "InferReq", state_cache_manager: LinearAttCacheManager, buffer_idx: int): + conv_state, ssm_state = state_cache_manager.get_state_cache(buffer_idx=buffer_idx) conv_dest = req.req_idx ssm_dest = req.req_idx * (self.mtp_step + 1) conv_cache_width = conv_state.shape[-1] - # TODO 下面这个从 cpu cache 拷贝数据的 gpu的操作,是否是阻塞的操作。 - # 同时,非连续对象的拷贝,可能存在效率问题。 self.req_to_conv_state.buffer[:, conv_dest, ..., :conv_cache_width] = conv_state self.req_to_ssm_state.buffer[:, ssm_dest, ...] = ssm_state if self.req_to_mtp_state_index is not None: diff --git a/lightllm/common/state_cache_manager/__init__.py b/lightllm/common/state_cache_manager/__init__.py new file mode 100644 index 0000000000..af0479a158 --- /dev/null +++ b/lightllm/common/state_cache_manager/__init__.py @@ -0,0 +1,13 @@ +from .base import StateCacheManager +from .layer_cache import LayerCache +from .linear_att import LinearAttCacheConfig, LinearAttCacheManager + + +def get_hybrid_cache_config(): + """Return the model-specific layout used by hybrid CPU/disk cache pages.""" + from lightllm.utils.config_utils import is_linear_att_mixed_model + from lightllm.utils.envs_utils import get_env_start_args + + if is_linear_att_mixed_model(get_env_start_args().model_dir): + return LinearAttCacheConfig.load_from_args() + raise ValueError("No hybrid state-cache layout registered for this model") diff --git a/lightllm/common/state_cache_manager/base.py b/lightllm/common/state_cache_manager/base.py new file mode 100644 index 0000000000..aa131fe784 --- /dev/null +++ b/lightllm/common/state_cache_manager/base.py @@ -0,0 +1,57 @@ +import collections +from abc import ABC, abstractmethod +from typing import List, Optional + +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + + +class StateCacheManager(ABC): + """CPU checkpoint 槽位池;大小页分别实例化,具体状态布局由子类负责。 + + 尾部 keep_num 个槽位保留给 CPU cache 碎页传输,不参与普通分配。 + 本类不管理 GPU 运行态,也不判断页面边界、前缀匹配或淘汰策略。 + """ + + def __init__(self, size: int, keep_num: int = 0): + self.size = size + self.keep_num = keep_num + assert 0 <= keep_num <= size, f"invalid keep_num {keep_num} for size {size}" + self.free_list = collections.deque(range(size - keep_num)) + + @abstractmethod + def get_state_cache(self, buffer_idx: int): + """返回指定槽位的状态视图;可以是单个 Tensor 或多个 Tensor。""" + + def alloc_one_state_cache(self) -> Optional[int]: + return None if not self.free_list else self.free_list.popleft() + + def alloc_state_cache(self, need_size: int) -> Optional[List[int]]: + if need_size > len(self.free_list): + logger.error(f"warn no enough cache need_size {need_size} free_size {len(self.free_list)}") + return None + return [self.free_list.popleft() for _ in range(need_size)] + + def free_state_cache(self, free_indexes: List[int]): + alloc_upper_bound = self.size - self.keep_num + for idx in free_indexes: + assert 0 <= idx < alloc_upper_bound, ( + f"free index {idx} out of alloc range [0, {alloc_upper_bound}), " f"reserved tail num {self.keep_num}" + ) + self.free_list.extend(free_indexes) + assert ( + len(self.free_list) <= alloc_upper_bound + ), f"free cache num {len(self.free_list)} should not be larger than alloc size {alloc_upper_bound}" + + def get_free_cache_num(self): + return len(self.free_list) + + def get_used_cache_num(self): + # Preserve the existing accounting: reserved slots count as used. + return self.size - len(self.free_list) + + def clear_to_init_state(self): + """重置空闲槽位;子类同时清零自身的 checkpoint buffer。""" + self.free_list = collections.deque(range(self.size - self.keep_num)) diff --git a/lightllm/common/linear_att_cache_manager/layer_cache.py b/lightllm/common/state_cache_manager/layer_cache.py similarity index 100% rename from lightllm/common/linear_att_cache_manager/layer_cache.py rename to lightllm/common/state_cache_manager/layer_cache.py diff --git a/lightllm/common/linear_att_cache_manager/config_objs.py b/lightllm/common/state_cache_manager/linear_att.py similarity index 80% rename from lightllm/common/linear_att_cache_manager/config_objs.py rename to lightllm/common/state_cache_manager/linear_att.py index f588ec7d5c..a9ee5cf1b4 100644 --- a/lightllm/common/linear_att_cache_manager/config_objs.py +++ b/lightllm/common/state_cache_manager/linear_att.py @@ -1,10 +1,15 @@ import torch import dataclasses import triton + from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num, get_env_start_args from lightllm.utils.log_utils import init_logger from lightllm.utils.torch_dtype_utils import get_torch_dtype +from .base import StateCacheManager +from .layer_cache import LayerCache + + logger = init_logger(__name__) @@ -143,3 +148,43 @@ def load_from_args() -> "LinearAttCacheConfig": all_layer_num=n_layer, draft_full_att_kv_layer_num=get_added_mtp_kv_layer_num(), ) + + +class LinearAttCacheManager(StateCacheManager): + """CPU pinned conv/SSM checkpoint,两个 buffer 均保持 size-first 布局。""" + + def __init__( + self, + size: int, + linear_config: LinearAttCacheConfig, + keep_num: int = 0, # 用于记录需要保留的缓存数量,用于支持含有 linear_att 的如qwen3.5 模型的cpu cache的碎页处理。 + ): + super().__init__(size, keep_num) + self.linear_config = linear_config + # init the layer cache + self.conv_state_cache = LayerCache( + size=self.size, + dtype=self.linear_config.conv_state_dtype, + shape=self.linear_config.get_conv_state_shape(), + layer_num=self.linear_config.linear_layer_num, + device="cpu", + size_first=True, + ) + self.ssm_state_cache = LayerCache( + size=self.size, + dtype=self.linear_config.ssm_state_dtype, + shape=self.linear_config.get_ssm_state_shape(), + layer_num=self.linear_config.linear_layer_num, + device="cpu", + size_first=True, + ) + return + + def get_state_cache(self, buffer_idx: int): + return self.conv_state_cache.buffer[buffer_idx, ...], self.ssm_state_cache.buffer[buffer_idx, ...] + + def clear_to_init_state(self): + self.conv_state_cache.buffer.zero_() + self.ssm_state_cache.buffer.zero_() + super().clear_to_init_state() + return diff --git a/lightllm/models/qwen3next/model.py b/lightllm/models/qwen3next/model.py index 5e443abc95..dd1a9c883a 100644 --- a/lightllm/models/qwen3next/model.py +++ b/lightllm/models/qwen3next/model.py @@ -16,7 +16,7 @@ from lightllm.common.kv_cache_mem_manager.qwen3next_mem_manager import Qwen3NextMemManager from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.common.req_manager import ReqManagerForMamba -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import LinearAttCacheConfig logger = init_logger(__name__) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 7c7ac9fe48..fe41c05aa0 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -19,7 +19,7 @@ from lightllm.utils.config_utils import ( has_audio_module, has_vision_module, - is_linear_att_mixed_model, + is_hybrid_att_model, auto_set_max_req_total_len, auto_set_fused_shared_experts, auto_set_response_parsers, @@ -256,23 +256,23 @@ def _launch_subprocesses(args: StartArgs): f"but got {args.batch_max_tokens}, {args.chunked_prefill_size}" ) - # linear att cache 参数自动设置 + # hybrid checkpoint 参数自动设置;保留现有 linear_att_* 启动参数名。 if args.linear_att_cache_size is None: - # linear_att_cache_size 只会在 qwen3.5 等混合线性层模型中生效。 + # 小页池大小只对 hybrid 模型生效。 default_cache_size = args.running_max_req_size * 2 dp_size_in_node = max(1, args.dp // args.nnodes) per_dp_cache_size = max(1, math.ceil(args.running_max_req_size / dp_size_in_node) * 2) args.linear_att_cache_size = min(default_cache_size, per_dp_cache_size) if args.run_mode == "decode": - # PD Decode 节点只接收 prompt 末尾位置的 linear attention state,不具备 + # PD Decode 节点只接收 prompt 末尾位置的 hybrid checkpoint,不具备 # 中间大页边界对应的 state。因此 Decode 节点必须使用默认值关闭大页功能, # 避免请求释放时将不完整的大页 state 写入 radix cache 并触发断言。 args.linear_att_page_block_num = 10000000 - if args.enable_cpu_cache and is_linear_att_mixed_model(args.model_dir): + if args.enable_cpu_cache and is_hybrid_att_model(args.model_dir): args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num - logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for linear hybrid att model") + logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for hybrid att model") # help to manage data stored on Ceph if "s3://" in args.model_dir: diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 9729a8205c..d9072bfba2 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -11,7 +11,7 @@ from lightllm.server.req_id_generator import convert_sub_id_to_group_id from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.envs_utils import get_env_start_args -from lightllm.utils.config_utils import is_linear_att_mixed_model +from lightllm.utils.config_utils import is_hybrid_att_model from lightllm.utils.kv_cache_utils import compute_token_list_hash from typing import Any, Dict, List, Union from lightllm.utils.log_utils import init_logger @@ -145,11 +145,11 @@ class Req(ctypes.Structure): # 当 stop_str_matched 条件满足的时候,对应的最后一个生成 token 所在的index位置。 # 该变量为 detokenization 进程写入,http_server 读取 ("stop_str_matched_token_index", ctypes.c_int), - # 用于在 包含linear att 混合模型中,进行输入的提前hash,方便在对应的page radix tree中进行快速操作。 - ("linear_att_token_hash_list", TokenHashList), + # hybrid 模型按 checkpoint 粒度提前计算输入 hash,供大小页 radix 匹配。 + ("hybrid_token_hash_list", TokenHashList), # 用于在开启cpu cache 或者 硬盘 cache时,预先计算,分块输入token的hash值。 ("token_hash_list", TokenHashList), - # 用于存储每个cpu cache 页面对应的真实token数量,用于linear att的qwen3.5等模型的碎片化处理最后一个页面的问题 + # 每个 CPU cache 页的真实 token 数,包含 hybrid 模型不足大页的尾页。 ("token_hash_page_len_list", TokenPageLenList), # 用于保存查找匹配到的可以被复用的cpu cache 页面信息。 ("cpu_cache_match_page_indexes", CpuCachePageList), @@ -216,10 +216,10 @@ def init( self.post_init() args = get_env_start_args() - if is_linear_att_mixed_model(args.model_dir): - self._fill_linear_att_token_hash() + if is_hybrid_att_model(args.model_dir): + self._fill_hybrid_token_hash() if args.enable_cpu_cache: - cpu_cache_hash_list, cpu_cache_page_len_list = self._calcu_linear_att_cpu_cache_page_len_list() + cpu_cache_hash_list, cpu_cache_page_len_list = self._calcu_hybrid_cpu_cache_page_len_list() self.token_hash_list = TokenHashList() self.token_hash_list.clear() self.token_hash_list.fill(cpu_cache_hash_list) @@ -243,12 +243,12 @@ def post_init(self): # 子类继承进行一些额外的初始化操作 pass - def _calcu_linear_att_cpu_cache_page_len_list(self): - token_hash_list = self.linear_att_token_hash_list.get_all() - linear_att_hash_page_size = get_env_start_args().linear_att_hash_page_size + def _calcu_hybrid_cpu_cache_page_len_list(self): + token_hash_list = self.hybrid_token_hash_list.get_all() + hash_page_size = get_env_start_args().linear_att_hash_page_size block_num = get_env_start_args().linear_att_page_block_num cpu_cache_page_size = get_env_start_args().cpu_cache_token_page_size - assert cpu_cache_page_size == linear_att_hash_page_size * block_num + assert cpu_cache_page_size == hash_page_size * block_num cpu_cache_hash_list = [] cpu_cache_page_len_list = [] cum_sum_len = 0 @@ -260,7 +260,7 @@ def _calcu_linear_att_cpu_cache_page_len_list(self): elif i == len(token_hash_list) - 1: cpu_cache_hash_list.append(token_hash_list[len(token_hash_list) - 1]) page_num = (i % block_num) + 1 - cum_sum_len += page_num * linear_att_hash_page_size + cum_sum_len += page_num * hash_page_size cpu_cache_page_len_list.append(cum_sum_len) return cpu_cache_hash_list, cpu_cache_page_len_list @@ -272,11 +272,11 @@ def _fill_input_token_hash(self): self.token_hash_list.fill(hash_values) return - def _fill_linear_att_token_hash(self): - self.linear_att_token_hash_list = TokenHashList() - self.linear_att_token_hash_list.clear() + def _fill_hybrid_token_hash(self): + self.hybrid_token_hash_list = TokenHashList() + self.hybrid_token_hash_list.clear() hash_values = compute_token_list_hash(self.get_prompt_ids(), get_env_start_args().linear_att_hash_page_size) - self.linear_att_token_hash_list.fill(hash_values) + self.hybrid_token_hash_list.fill(hash_values) return def create_prompt_ids_shm_array(self): diff --git a/lightllm/server/core/objs/token_chunck_hash_list.py b/lightllm/server/core/objs/token_chunck_hash_list.py index 23de10353b..eaff4bc1f5 100644 --- a/lightllm/server/core/objs/token_chunck_hash_list.py +++ b/lightllm/server/core/objs/token_chunck_hash_list.py @@ -92,7 +92,7 @@ def get_all(self): class TokenPageLenList(CpuCachePageList): """ - 用于记录cpu cache 每个 page 对应的真实prefix token数量, 用于支持含有 linear_att 的如qwen3.5 模型的cpu cache的 + 用于记录 CPU cache 每个 page 对应的真实 prefix token 数量,支持 hybrid 模型 CPU cache 的 的最后一个页面的非满页面的碎片化处理。 """ diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 78f5fedc93..db477df8c6 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -10,6 +10,9 @@ logger = init_logger(__name__) +# Keep the existing wire tag so P/D nodes can be upgraded independently. +HYBRID_ATT_STATE_PAGE_KIND = "linear_att_state" + # 节点的行为 class NodeRole(enum.Enum): @@ -190,7 +193,7 @@ def __post_init__(self): raise ValueError(error_info) if self.page_kind == "kv": assert len(self.mem_indexes) == (self.end_kv_index - self.start_kv_index) - elif self.page_kind == "linear_att_state": + elif self.page_kind == HYBRID_ATT_STATE_PAGE_KIND: assert self.start_kv_index == self.end_kv_index assert len(self.mem_indexes) == 0 else: diff --git a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py b/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py similarity index 82% rename from lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py rename to lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py index 6a8e0a3917..d911e37320 100644 --- a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py @@ -2,19 +2,19 @@ import numpy as np from typing import Tuple, Dict, Set, List, Optional from sortedcontainers import SortedSet, SortedDict -from lightllm.common.linear_att_cache_manager import LinearAttCacheManager +from lightllm.common.state_cache_manager import StateCacheManager from .shared_arr import SharedArray from .radix_cache import time_gen -class LinearAttPagedTreeNode: +class HybridAttPagedTreeNode: def __init__(self, hash_page_size: int, big_page_num: int): self.hash_page_size = hash_page_size self.big_page_num = big_page_num # children are keyed by the last ``block_hash`` of each child - self.children: Dict[int, "LinearAttPagedTreeNode"] = {} - self.parent: "LinearAttPagedTreeNode" = None + self.children: Dict[int, "HybridAttPagedTreeNode"] = {} + self.parent: "HybridAttPagedTreeNode" = None # Hash of the last page in this node (None for the empty root). self.page_num = None # 页面数量,只能是 1 或者 big_page_num @@ -62,9 +62,9 @@ def add_and_return_new_child( token_mem_index_value: torch.Tensor, block_hash: int, small_page_buffer_idx: Optional[int], - ) -> "LinearAttPagedTreeNode": + ) -> "HybridAttPagedTreeNode": assert len(token_id_key) == self.hash_page_size == len(token_mem_index_value) - child = LinearAttPagedTreeNode(hash_page_size=self.hash_page_size, big_page_num=self.big_page_num) + child = HybridAttPagedTreeNode(hash_page_size=self.hash_page_size, big_page_num=self.big_page_num) child.page_hash = block_hash child.small_page_buffer_idx = small_page_buffer_idx child.token_id_key = token_id_key @@ -81,9 +81,9 @@ def add_and_return_new_child( def add_and_return_new_big_page_child( self, token_id_key: torch.Tensor, token_mem_index_value: torch.Tensor, block_hash: int, big_page_buffer_idx: int - ) -> "LinearAttPagedTreeNode": + ) -> "HybridAttPagedTreeNode": assert len(token_id_key) == self.hash_page_size * self.big_page_num == len(token_mem_index_value) - child = LinearAttPagedTreeNode(hash_page_size=self.hash_page_size, big_page_num=self.big_page_num) + child = HybridAttPagedTreeNode(hash_page_size=self.hash_page_size, big_page_num=self.big_page_num) child.page_hash = block_hash child.token_id_key = token_id_key child.token_mem_index_value = token_mem_index_value @@ -98,7 +98,7 @@ def add_and_return_new_big_page_child( child.node_prefix_total_len = child.parent.node_prefix_total_len + new_len return child - def remove_child(self, child_node: "LinearAttPagedTreeNode"): + def remove_child(self, child_node: "HybridAttPagedTreeNode"): del self.children[child_node.page_hash] child_node.parent = None @@ -109,7 +109,7 @@ def is_leaf(self): return len(self.children) == 0 -class LinearAttPagedRadixCache: +class HybridAttPagedRadixCache: def __init__( self, unique_name: str, @@ -118,7 +118,7 @@ def __init__( hash_page_size: int, big_page_num: int, kv_cache_mem_manager=None, - linear_att_small_page_buffers=None, + small_page_buffers=None, ): from lightllm.common.kv_cache_mem_manager import MemoryManager @@ -132,18 +132,18 @@ def __init__( self.mem_manager: MemoryManager = kv_cache_mem_manager - self.linear_att_big_page_buffers: LinearAttCacheManager = self.mem_manager.linear_att_big_page_buffers + self.big_page_buffers: StateCacheManager = self.mem_manager.big_page_buffers self._key_dtype = torch.int64 self._value_dtype = torch.int64 - self.root_node = LinearAttPagedTreeNode(hash_page_size=hash_page_size, big_page_num=big_page_num) + self.root_node = HybridAttPagedTreeNode(hash_page_size=hash_page_size, big_page_num=big_page_num) 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 self.root_node.page_num = self.big_page_num - self._evict_tree_set: Set[LinearAttPagedTreeNode] = SortedSet(key=lambda x: x.get_compare_key()) - self._evict_tree_set_for_linear_att: Set[LinearAttPagedTreeNode] = SortedSet( + self._evict_tree_set: Set[HybridAttPagedTreeNode] = SortedSet(key=lambda x: x.get_compare_key()) + self._evict_tree_set_for_state_cache: Set[HybridAttPagedTreeNode] = SortedSet( key=lambda x: x.get_compare_key_for_buffer_idx() ) @@ -153,23 +153,23 @@ def __init__( f"{unique_name}_tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64 ) self.tree_total_tokens_num.arr[0] = 0 - self.linear_att_small_page_buffers: LinearAttCacheManager = linear_att_small_page_buffers + self.small_page_buffers: StateCacheManager = small_page_buffers - def _discard_node(self, node: LinearAttPagedTreeNode): + def _discard_node(self, node: HybridAttPagedTreeNode): if node.is_leaf(): self._evict_tree_set.discard(node) if node.small_page_buffer_idx is not None: - self._evict_tree_set_for_linear_att.discard(node) + self._evict_tree_set_for_state_cache.discard(node) return - def _add_node(self, node: LinearAttPagedTreeNode): + def _add_node(self, node: HybridAttPagedTreeNode): # root 永远不参与回收:当树为空时 root 自身也满足 is_leaf(),若加入 _evict_tree_set, # 会与 _evict 中 "node is not self.root_node" 的断言相矛盾(当前仅靠 root 的 ref_counter>=1 # 和回收水位 guard 掩盖)。这里显式排除,使数据结构与回收逻辑的意图一致。 if node.is_leaf() and node is not self.root_node: self._evict_tree_set.add(node) if node.small_page_buffer_idx is not None: - self._evict_tree_set_for_linear_att.add(node) + self._evict_tree_set_for_state_cache.add(node) return def insert( @@ -177,17 +177,17 @@ def insert( key: torch.Tensor, value: Optional[torch.Tensor] = None, block_hashs: Optional[List[int]] = None, - block_linear_idxs: Optional[List[int]] = None, + block_state_idxs: Optional[List[int]] = None, len_to_big_page_id: Optional[SortedDict] = None, - ) -> Tuple[int, Optional[LinearAttPagedTreeNode]]: + ) -> Tuple[int, Optional[HybridAttPagedTreeNode]]: assert key is not None if value is None: value = key assert len(key) == len(value) if block_hashs is None: block_hashs = [] - if block_linear_idxs is None: - block_linear_idxs = [] + if block_state_idxs is None: + block_state_idxs = [] if len_to_big_page_id is None: len_to_big_page_id = SortedDict() @@ -197,38 +197,38 @@ def insert( len(key) == len(block_hashs) * self.hash_page_size ), f"key length {len(key)} does not match block_hashs length {len(block_hashs)} * {self.hash_page_size}" assert len(block_hashs) == len( - block_linear_idxs - ), f"block_hashs length {len(block_hashs)} does not match block_linear_idxs length {len(block_linear_idxs)}" + block_state_idxs + ), f"block_hashs length {len(block_hashs)} does not match block_state_idxs length {len(block_state_idxs)}" if len(block_hashs) == 0: return 0, None if len(block_hashs) % self.big_page_num == 0: assert all( - e is None for e in block_linear_idxs - ), "all block_linear_idxs must be None when block_hashs length is a multiple of big_page_num" + e is None for e in block_state_idxs + ), "all block_state_idxs must be None when block_hashs length is a multiple of big_page_num" else: # TODO, test stable then to delete this assertion assert all( - e is None for e in block_linear_idxs[:-1] - ), "only the last block_linear_idx can be non-None, for compatibility with non-paged radix cache" + e is None for e in block_state_idxs[:-1] + ), "only the last block_state_idx can be non-None, for compatibility with non-paged radix cache" assert ( - block_linear_idxs[-1] is not None - ), "the last block_linear_idx must not be None, for compatibility with non-paged radix cache" + block_state_idxs[-1] is not None + ), "the last block_state_idx must not be None, for compatibility with non-paged radix cache" - ans = self._insert_helper(self.root_node, key, value, block_hashs, block_linear_idxs, len_to_big_page_id) + ans = self._insert_helper(self.root_node, key, value, block_hashs, block_state_idxs, len_to_big_page_id) assert len(len_to_big_page_id) == 0 return ans def _insert_helper( self, - node: LinearAttPagedTreeNode, + node: HybridAttPagedTreeNode, key: torch.Tensor, value: torch.Tensor, block_hashs: List[int], - block_linear_idxs: List[int], + block_state_idxs: List[int], len_to_big_page_id: SortedDict, - ) -> Tuple[int, Optional[LinearAttPagedTreeNode]]: + ) -> Tuple[int, Optional[HybridAttPagedTreeNode]]: self._discard_node(node) node.update_time() @@ -250,7 +250,7 @@ def _insert_helper( new_big_page_buffer_id = len_to_big_page_id.pop(child.node_prefix_total_len, None) if new_big_page_buffer_id is not None: # 因为节点已经存在,所以无法插入,但是要释放对应的buffer_id 节点 - self.linear_att_big_page_buffers.free_state_cache([new_big_page_buffer_id]) + self.big_page_buffers.free_state_cache([new_big_page_buffer_id]) # 已经存在了 sub_prefix_len, ans_node = self._insert_helper( @@ -258,7 +258,7 @@ def _insert_helper( key[self.big_page_tokens :], value[self.big_page_tokens :], block_hashs[self.big_page_num :], - block_linear_idxs[self.big_page_num :], + block_state_idxs[self.big_page_num :], len_to_big_page_id, ) return self.big_page_tokens + sub_prefix_len, ans_node @@ -284,7 +284,7 @@ def _insert_helper( key[self.big_page_tokens :], value[self.big_page_tokens :], block_hashs[self.big_page_num :], - block_linear_idxs[self.big_page_num :], + block_state_idxs[self.big_page_num :], len_to_big_page_id, ) return 0, ans_node @@ -296,23 +296,23 @@ def _insert_helper( if block_hashs[0] in node.children: child = node.children[block_hashs[0]] - if block_linear_idxs[0] is not None: - assert len(block_hashs) == 1 == len(block_linear_idxs) + if block_state_idxs[0] is not None: + assert len(block_hashs) == 1 == len(block_state_idxs) if child.small_page_buffer_idx is None: # 将这个buffer id 移交给这个存在的节点。 self._discard_node(child) - child.small_page_buffer_idx = block_linear_idxs[0] + child.small_page_buffer_idx = block_state_idxs[0] self._add_node(child) else: - # 说明节点已经存在了,直接提前移除掉这个节点占用的线性缓存,外部不用处理这个细节了 - self.linear_att_small_page_buffers.free_state_cache(free_indexes=[block_linear_idxs[0]]) + # 节点已有 checkpoint,释放本次重复申请的槽位。 + self.small_page_buffers.free_state_cache(free_indexes=[block_state_idxs[0]]) sub_prefix_len, ans_node = self._insert_helper( child, key[self.hash_page_size :], value[self.hash_page_size :], block_hashs[1:], - block_linear_idxs[1:], + block_state_idxs[1:], len_to_big_page_id, ) return self.hash_page_size + sub_prefix_len, ans_node @@ -321,7 +321,7 @@ def _insert_helper( key[: self.hash_page_size], value[: self.hash_page_size], block_hashs[0], - block_linear_idxs[0], + block_state_idxs[0], ) assert not new_node.is_big_page_node() assert new_node.page_num == 1 @@ -331,7 +331,7 @@ def _insert_helper( key[self.hash_page_size :], value[self.hash_page_size :], block_hashs[1:], - block_linear_idxs[1:], + block_state_idxs[1:], len_to_big_page_id, ) return 0, ans_node @@ -357,7 +357,7 @@ def match_prefix( if len(block_hashs) == 0 or len(key) == 0: return None, 0, None - ans_node_list: List[LinearAttPagedTreeNode] = [] + ans_node_list: List[HybridAttPagedTreeNode] = [] self._match_prefix_helper( self.root_node, key=key, @@ -381,7 +381,7 @@ def match_prefix( def _match_prefix_helper( self, - node: LinearAttPagedTreeNode, + node: HybridAttPagedTreeNode, key: torch.Tensor, block_hashs: Optional[List[int]], ans_node_list: list, @@ -434,7 +434,7 @@ def _match_prefix_helper( finally: self._add_node(node) - def _trim_unusable_match_tail(self, nodes: List[LinearAttPagedTreeNode]) -> List[LinearAttPagedTreeNode]: + def _trim_unusable_match_tail(self, nodes: List[HybridAttPagedTreeNode]) -> List[HybridAttPagedTreeNode]: removed_list = [] for node in reversed(nodes): if node.is_big_page_node(): @@ -459,7 +459,7 @@ def _trim_unusable_match_tail(self, nodes: List[LinearAttPagedTreeNode]) -> List else: return nodes[: -len(removed_list)] - def _try_merge(self, child_node: LinearAttPagedTreeNode) -> Optional[LinearAttPagedTreeNode]: + def _try_merge(self, child_node: HybridAttPagedTreeNode) -> Optional[HybridAttPagedTreeNode]: raise NotImplementedError() def merge_unreferenced_nodes(self): @@ -474,7 +474,7 @@ def flush_cache(self): self.free_radix_cache_to_get_enough_token(need_token_num=self.total_token_num) return - def deref_to_first_big_page_node(self, node: LinearAttPagedTreeNode) -> Optional[LinearAttPagedTreeNode]: + def deref_to_first_big_page_node(self, node: HybridAttPagedTreeNode) -> Optional[HybridAttPagedTreeNode]: assert not node.is_big_page_node() iter_node = node while not iter_node.is_big_page_node(): @@ -493,7 +493,7 @@ def deref_to_first_big_page_node(self, node: LinearAttPagedTreeNode) -> Optional else: return iter_node - def dec_node_ref_counter(self, node: LinearAttPagedTreeNode): + def dec_node_ref_counter(self, node: HybridAttPagedTreeNode): if node is None: return old_node = node @@ -508,7 +508,7 @@ def dec_node_ref_counter(self, node: LinearAttPagedTreeNode): self._add_node(old_node) return - def add_node_ref_counter(self, node: LinearAttPagedTreeNode): + def add_node_ref_counter(self, node: HybridAttPagedTreeNode): if node is None: return old_node = node @@ -523,7 +523,7 @@ def add_node_ref_counter(self, node: LinearAttPagedTreeNode): self._add_node(old_node) return - def get_mem_index_value_by_node(self, node: LinearAttPagedTreeNode) -> Optional[torch.Tensor]: + def get_mem_index_value_by_node(self, node: HybridAttPagedTreeNode) -> Optional[torch.Tensor]: if node is None: return None @@ -535,7 +535,7 @@ def get_mem_index_value_by_node(self, node: LinearAttPagedTreeNode) -> Optional[ ans_list.reverse() return torch.concat(ans_list, dim=0) - def get_big_page_ids_by_node(self, node: LinearAttPagedTreeNode) -> List[int]: + def get_big_page_ids_by_node(self, node: HybridAttPagedTreeNode) -> List[int]: if node is None: return [] if node is self.root_node: @@ -559,7 +559,7 @@ def get_tree_total_tokens_num(self): def print_self(self, indent=0): self._print_helper(self.root_node, indent) - def _print_helper(self, node: LinearAttPagedTreeNode, indent): + def _print_helper(self, node: HybridAttPagedTreeNode, indent): print( " " * indent, f"hash_info: {node.page_hash} " @@ -580,9 +580,9 @@ def free_radix_cache_to_get_enough_token(self, need_token_num): release_mems = [] small_page_buffer_ids = [] - def release_mem(mem_index, linear_att_small_page_id): + def release_mem(mem_index, small_page_buffer_id): release_mems.append(mem_index) - small_page_buffer_ids.append(linear_att_small_page_id) + small_page_buffer_ids.append(small_page_buffer_id) return self._evict(need_evict_token_num, release_mem) @@ -590,22 +590,22 @@ def release_mem(mem_index, linear_att_small_page_id): self.mem_manager.free(mem_index) small_page_buffer_ids = [idx for idx in small_page_buffer_ids if idx is not None] if len(small_page_buffer_ids) > 0: - self.linear_att_small_page_buffers.free_state_cache(small_page_buffer_ids) + self.small_page_buffers.free_state_cache(small_page_buffer_ids) return - def free_one_small_page_linear_att_buffer(self): - if self.linear_att_small_page_buffers is None: + def free_one_small_page_buffer(self): + if self.small_page_buffers is None: return - if self.linear_att_small_page_buffers.get_free_cache_num() > 0: + if self.small_page_buffers.get_free_cache_num() > 0: return - if len(self._evict_tree_set_for_linear_att) == 0: + if len(self._evict_tree_set_for_state_cache) == 0: return - node: LinearAttPagedTreeNode = self._evict_tree_set_for_linear_att.pop(0) + node: HybridAttPagedTreeNode = self._evict_tree_set_for_state_cache.pop(0) self._discard_node(node) assert node.small_page_buffer_idx is not None - self.linear_att_small_page_buffers.free_state_cache(free_indexes=[node.small_page_buffer_idx]) + self.small_page_buffers.free_state_cache(free_indexes=[node.small_page_buffer_idx]) node.small_page_buffer_idx = None self._add_node(node) @@ -618,7 +618,7 @@ def _evict(self, need_remove_tokens, evict_callback): refed_tokens_num {self.refed_tokens_num.arr[0]}""" num_evicted = 0 while num_evicted < need_remove_tokens: - node: LinearAttPagedTreeNode = self._evict_tree_set.pop(0) + node: HybridAttPagedTreeNode = self._evict_tree_set.pop(0) self._discard_node(node) assert ( @@ -628,11 +628,11 @@ def _evict(self, need_remove_tokens, evict_callback): if node.is_big_page_node(): assert node.big_page_buffer_idx is not None - self.linear_att_big_page_buffers.free_state_cache([node.big_page_buffer_idx]) + self.big_page_buffers.free_state_cache([node.big_page_buffer_idx]) evict_callback(node.token_mem_index_value, node.small_page_buffer_idx) self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value) - parent_node: LinearAttPagedTreeNode = node.parent + parent_node: HybridAttPagedTreeNode = node.parent parent_node.remove_child(node) self._add_node(parent_node) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 3585c223e2..dd37dace9e 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -8,13 +8,13 @@ from sortedcontainers import SortedDict from dataclasses import dataclass, field from typing import TYPE_CHECKING, List, Dict, Tuple, Optional, Callable, Any, Union -from lightllm.common.req_manager import ReqManager, ReqManagerForMamba +from lightllm.common.req_manager import ReqManager, HybridAttentionReqManager from lightllm.utils.infer_utils import mark_start, mark_end from lightllm.server.core.objs import Req, SamplingParams, FinishStatus, ShmReqManager from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache, TreeNode -from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import ( - LinearAttPagedRadixCache, - LinearAttPagedTreeNode, +from lightllm.server.router.dynamic_prompt.hybrid_att_radix_cache import ( + HybridAttPagedRadixCache, + HybridAttPagedTreeNode, ) from lightllm.utils.log_utils import init_logger from lightllm.server.req_id_generator import convert_sub_id_to_group_id @@ -34,8 +34,8 @@ @dataclass class InferenceContext: - req_manager: Union[ReqManager, ReqManagerForMamba] = None # gpu 请求管理 - radix_cache: Union[LinearAttPagedRadixCache, RadixCache] = None + req_manager: ReqManager = None # gpu 请求管理 + radix_cache: Union[HybridAttPagedRadixCache, RadixCache] = None shm_req_manager: ShmReqManager = None # 共享内存请求对象管理 requests_mapping: Dict[int, "InferReq"] = None infer_req_ids = None @@ -45,13 +45,13 @@ class InferenceContext: overlap_stream: torch.cuda.Stream = None # 一些情况下推理进程进行异步折叠操作的异步流对象。 cpu_kv_cache_stream: torch.cuda.Stream = None # 用 cpu kv cache 操作的 stream - is_linear_att_mixed_model: bool = False # 标记模型是否是full att 混合 linear att 的混合模型。 + is_hybrid_att_model: bool = False # 使用大小页 checkpoint 的混合 attention 模型。 def register( self, backend: "ModeBackend", - req_manager: Union[ReqManager, ReqManagerForMamba], - radix_cache: Union[LinearAttPagedRadixCache, RadixCache], + req_manager: ReqManager, + radix_cache: Union[HybridAttPagedRadixCache, RadixCache], shm_req_manager: ShmReqManager, vocab_size: int, cache_placement_controller: Optional[CachePlacementController] = None, @@ -69,7 +69,7 @@ def register( self.vocab_size = vocab_size - self.is_linear_att_mixed_model = isinstance(self.req_manager, ReqManagerForMamba) + self.is_hybrid_att_model = isinstance(self.req_manager, HybridAttentionReqManager) return @@ -133,11 +133,11 @@ def free_a_req_mem(self, free_token_index: List, req: "InferReq"): elif CacheTier.GPU not in req.cache_tiers: self._free_req_mem_without_radix_insert(free_token_index=free_token_index, req=req) else: - if not self.is_linear_att_mixed_model: + if not self.is_hybrid_att_model: self._full_att_free_req(free_token_index=free_token_index, req=req) else: - self._linear_att_free_req(free_token_index=free_token_index, req=req) - assert len(req.linear_att_len_to_big_page_id) == 0 + self._hybrid_att_free_req(free_token_index=free_token_index, req=req) + assert len(req.hybrid_len_to_big_page_id) == 0 req.cur_kv_len = 0 req.shm_req.shm_cur_kv_len = req.cur_kv_len return @@ -146,19 +146,15 @@ def _free_req_mem_without_radix_insert(self, free_token_index: List, req: "Infer shared_kv_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][shared_kv_len : req.cur_kv_len]) - if self.is_linear_att_mixed_model: - # 释放请求尾部尚未移交给 radix cache 的 linear attention 小页状态。 - if req.tail_linear_att_small_page_buffer_id is not None: - self.radix_cache.linear_att_small_page_buffers.free_state_cache( - [req.tail_linear_att_small_page_buffer_id] - ) - req.tail_linear_att_small_page_buffer_id = None + if self.is_hybrid_att_model: + # 释放请求尾部尚未移交给 radix cache 的 hybrid attention 小页状态。 + if req.tail_small_page_buffer_id is not None: + self.radix_cache.small_page_buffers.free_state_cache([req.tail_small_page_buffer_id]) + req.tail_small_page_buffer_id = None # 释放请求执行期间申请、但不再插入 radix cache 的大页状态。 - if req.linear_att_len_to_big_page_id: - self.radix_cache.linear_att_big_page_buffers.free_state_cache( - list(req.linear_att_len_to_big_page_id.values()) - ) - req.linear_att_len_to_big_page_id.clear() + if req.hybrid_len_to_big_page_id: + self.radix_cache.big_page_buffers.free_state_cache(list(req.hybrid_len_to_big_page_id.values())) + req.hybrid_len_to_big_page_id.clear() # 解除请求对已命中 GPU radix cache 前缀节点的引用。 if req.shared_kv_node is not None: @@ -181,43 +177,43 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): req.shared_kv_node = None return - def _linear_att_free_req(self, free_token_index: List, req: "InferReq"): - assert g_infer_context.is_linear_att_mixed_model is True + def _hybrid_att_free_req(self, free_token_index: List, req: "InferReq"): + assert g_infer_context.is_hybrid_att_model is True args = get_env_start_args() hash_page_size = args.linear_att_hash_page_size big_page_num = args.linear_att_page_block_num shared_kv_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len tail_big_page_token_num = ( - req.linear_att_cache_len // (hash_page_size * big_page_num) * (hash_page_size * big_page_num) + req.hybrid_cache_len // (hash_page_size * big_page_num) * (hash_page_size * big_page_num) ) - page_num = req.linear_att_cache_len // hash_page_size - assert req.linear_att_cache_len >= shared_kv_len - if req.tail_linear_att_small_page_buffer_id is not None: - assert req.linear_att_cache_len <= req.cur_kv_len + page_num = req.hybrid_cache_len // hash_page_size + assert req.hybrid_cache_len >= shared_kv_len + if req.tail_small_page_buffer_id is not None: + assert req.hybrid_cache_len <= req.cur_kv_len if req.cur_kv_len == 0: return - if req.linear_att_cache_len <= req.cur_kv_len and req.tail_linear_att_small_page_buffer_id is not None: - # 只有小页可以有 tail_linear_att_small_page_buffer_id,然后进行小页插入。 + if req.hybrid_cache_len <= req.cur_kv_len and req.tail_small_page_buffer_id is not None: + # 只有小页可以有 tail_small_page_buffer_id,然后进行小页插入。 assert page_num % big_page_num != 0 free_token_index.append( - self.req_manager.req_to_token_indexs[req.req_idx][req.linear_att_cache_len : req.cur_kv_len] + self.req_manager.req_to_token_indexs[req.req_idx][req.hybrid_cache_len : req.cur_kv_len] ) - req.cur_kv_len = req.linear_att_cache_len + req.cur_kv_len = req.hybrid_cache_len input_token_ids = req.get_input_token_ids() key = torch.tensor(input_token_ids[0 : req.cur_kv_len], dtype=torch.int64, device="cpu") value = self.req_manager.req_to_token_indexs[req.req_idx][: req.cur_kv_len].detach().cpu() - block_hashs = req.shm_req.linear_att_token_hash_list.get_all()[:page_num] - linear_idxs = [None for _ in range(page_num)] - linear_idxs[-1] = req.tail_linear_att_small_page_buffer_id - req.tail_linear_att_small_page_buffer_id = None + block_hashs = req.shm_req.hybrid_token_hash_list.get_all()[:page_num] + state_idxs = [None for _ in range(page_num)] + state_idxs[-1] = req.tail_small_page_buffer_id + req.tail_small_page_buffer_id = None prefix_len, _ = self.radix_cache.insert( key, value, block_hashs=block_hashs, - block_linear_idxs=linear_idxs, - len_to_big_page_id=req.linear_att_len_to_big_page_id, + block_state_idxs=state_idxs, + len_to_big_page_id=req.hybrid_len_to_big_page_id, ) 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]) @@ -233,20 +229,20 @@ def _linear_att_free_req(self, free_token_index: List, req: "InferReq"): ) req.cur_kv_len = tail_big_page_token_num - assert req.tail_linear_att_small_page_buffer_id is None + assert req.tail_small_page_buffer_id is None input_token_ids = req.get_input_token_ids() key = torch.tensor(input_token_ids[0 : req.cur_kv_len], dtype=torch.int64, device="cpu") value = self.req_manager.req_to_token_indexs[req.req_idx][: req.cur_kv_len].detach().cpu() cur_page_num = tail_big_page_token_num // hash_page_size assert tail_big_page_token_num % hash_page_size == 0 - block_hashs = req.shm_req.linear_att_token_hash_list.get_all()[:cur_page_num] - linear_idxs = [None for _ in range(cur_page_num)] + block_hashs = req.shm_req.hybrid_token_hash_list.get_all()[:cur_page_num] + state_idxs = [None for _ in range(cur_page_num)] prefix_len, _ = self.radix_cache.insert( key, value, block_hashs=block_hashs, - block_linear_idxs=linear_idxs, - len_to_big_page_id=req.linear_att_len_to_big_page_id, + block_state_idxs=state_idxs, + len_to_big_page_id=req.hybrid_len_to_big_page_id, ) 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]) @@ -265,14 +261,12 @@ def _linear_att_free_req(self, free_token_index: List, req: "InferReq"): # state buffer。仅当请求未走 insert 分支(小页/大页插入)就被释放时才会有残留,典型场景: # big page 模式下请求在 prefill 跨过 big page 边界后、到达末尾前被 pause / abort。 # 若不释放,会泄漏 big page state slot,并触发 free_a_req_mem 中 dict 为空的断言。 - if req.linear_att_len_to_big_page_id: - self.radix_cache.linear_att_big_page_buffers.free_state_cache( - list(req.linear_att_len_to_big_page_id.values()) - ) - req.linear_att_len_to_big_page_id.clear() + if req.hybrid_len_to_big_page_id: + self.radix_cache.big_page_buffers.free_state_cache(list(req.hybrid_len_to_big_page_id.values())) + req.hybrid_len_to_big_page_id.clear() req.cur_kv_len = shared_kv_len - assert req.tail_linear_att_small_page_buffer_id is None + assert req.tail_small_page_buffer_id is None if req.shared_kv_node is not None: assert req.shared_kv_node.node_prefix_total_len == req.cur_kv_len self.radix_cache.dec_node_ref_counter(req.shared_kv_node) @@ -375,8 +369,8 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo if prefill_need_token_num > can_alloc_token_num: break - if g_infer_context.is_linear_att_mixed_model: - req._linear_match_radix_cache() + if g_infer_context.is_hybrid_att_model: + req._hybrid_match_radix_cache() else: req._match_radix_cache() @@ -396,74 +390,50 @@ def get_can_alloc_token_num(self): ) return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num - def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): - """ - 该函数用于在线性混合模型prefill后,如果存在大页匹配的情况下,将线性层状态复制到 - """ - if not self.is_linear_att_mixed_model: + def save_hybrid_state_to_cache(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): + """Snapshot request-level attention state at big/small-page boundaries.""" + if not self.is_hybrid_att_model or self.radix_cache is None: return - # 大页对应的 linear att 的拷贝 + # Request-state snapshot at a big-page boundary. big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num big_page_buffer_ids = [] for req in reqs: cur_input_len = req.get_chuncked_input_token_len() - if cur_input_len % big_page_token_num == 0 and cur_input_len <= req.linear_att_cache_len: - big_page_id = self.radix_cache.linear_att_big_page_buffers.alloc_one_state_cache() + if cur_input_len % big_page_token_num == 0 and cur_input_len <= req.hybrid_cache_len: + big_page_id = self.radix_cache.big_page_buffers.alloc_one_state_cache() assert big_page_id is not None big_page_buffer_ids.append(big_page_id) - assert cur_input_len not in req.linear_att_len_to_big_page_id - req.linear_att_len_to_big_page_id[cur_input_len] = big_page_id + assert cur_input_len not in req.hybrid_len_to_big_page_id + req.hybrid_len_to_big_page_id[cur_input_len] = big_page_id else: big_page_buffer_ids.append(-1) assert len(b_req_idx) == len(big_page_buffer_ids) if any(buffer_id != -1 for buffer_id in big_page_buffer_ids): - big_page_buffer_ids = torch.tensor( - big_page_buffer_ids, dtype=torch.int32, requires_grad=False, device="cpu" - ) - big_page_buffer_ids = big_page_buffer_ids.cuda(non_blocking=True) - - from lightllm.common.basemodel.triton_kernel.linear_att_copy import copy_linear_att_state_to_kv_buffer - - copy_linear_att_state_to_kv_buffer( + self.req_manager.save_big_page_states( b_req_idx=b_req_idx, - big_page_buffer_ids=big_page_buffer_ids, - gpu_conv_state=self.req_manager.req_to_conv_state.buffer, - gpu_ssm_state=self.req_manager.req_to_ssm_state.buffer, - cpu_kv_conv_state=self.radix_cache.linear_att_big_page_buffers.conv_state_cache.buffer, - cpu_kv_ssm_state=self.radix_cache.linear_att_big_page_buffers.ssm_state_cache.buffer, - mtp_step=self.args.mtp_step, + req_indexes=[req.req_idx for req in reqs], + buffer_indexes=big_page_buffer_ids, ) - assert not self.args.disable_chunked_prefill, "chunked prefill mode must be enabled for linear att mixed model" + assert not self.args.disable_chunked_prefill, "chunked prefill must be enabled for hybrid attention models" - # tail small page 的linear att 状态的存储 + # Request-state snapshot at the final small-page boundary. for req in reqs: - # 判断本次prefill 完以后 kv 的长度是否到达linear att 块存储的临界点。 - if req.get_chuncked_input_token_len() == req.linear_att_cache_len: - assert req.tail_linear_att_small_page_buffer_id is None - if req.linear_att_cache_len % big_page_token_num != 0: - self.radix_cache.free_one_small_page_linear_att_buffer() - req.tail_linear_att_small_page_buffer_id = ( - self.radix_cache.linear_att_small_page_buffers.alloc_one_state_cache() - ) - if req.tail_linear_att_small_page_buffer_id is not None: - conv_src_idx = req.req_idx - ssm_src_idx = req.req_idx * (self.args.mtp_step + 1) - conv_cache_width = self.req_manager.linear_config.get_conv_state_shape()[-1] - gpu_conv_state = self.req_manager.req_to_conv_state.buffer[ - :, conv_src_idx, ..., :conv_cache_width - ] - gpu_ssm_state = self.req_manager.req_to_ssm_state.buffer[:, ssm_src_idx, ...] - dst_buffer_idx = req.tail_linear_att_small_page_buffer_id - - dst_conv_state, dst_ssm_state = self.radix_cache.linear_att_small_page_buffers.get_state_cache( - buffer_idx=dst_buffer_idx + # 判断本次prefill 完以后 kv 的长度是否到达 hybrid checkpoint 的存储边界。 + if req.get_chuncked_input_token_len() == req.hybrid_cache_len: + assert req.tail_small_page_buffer_id is None + if req.hybrid_cache_len % big_page_token_num != 0: + self.radix_cache.free_one_small_page_buffer() + req.tail_small_page_buffer_id = self.radix_cache.small_page_buffers.alloc_one_state_cache() + if req.tail_small_page_buffer_id is not None: + dst_buffer_idx = req.tail_small_page_buffer_id + self.req_manager.save_state( + req_idx=req.req_idx, + buffer_idx=dst_buffer_idx, + state_cache_manager=self.radix_cache.small_page_buffers, ) - # TODO 对于非连续对象调用 copy_ 效率并不高 - dst_conv_state.copy_(gpu_conv_state, non_blocking=True) - dst_ssm_state.copy_(gpu_ssm_state, non_blocking=True) return @@ -582,16 +552,15 @@ def __init__( self.pd_task_failed_num: int = 0 self.pd_trans_device_id: int = -1 - # 类似 qwen3.5 这种混合linear att 模型使用的状态,记录申请来用于保存对应的线性att缓存的 buffer id - # 当 prefill 阶段结束后, 对应长度的 linear att state 会写入到申请 buffer id 对应的块中, 方便插入到 radix cache中 + # hybrid checkpoint 槽位:prefill 到达边界后保存运行态,供请求释放时插入 radix cache。 # 方便被后续的请求使用,因为这种资源是有限的,也可能不存在的情况,申请不到时, 为None,则这种小块对应长度的 kv 无法 # 在后续被插入到radix cache中. 这个id 是对应radix cache中的small page的buffer. # 对应请求最尾巴上那一个块,对应的 small page buffer id - self.tail_linear_att_small_page_buffer_id: Optional[int] = None - # linear cache 对应的长度位置。 - self.linear_att_cache_len: Optional[int] = None + self.tail_small_page_buffer_id: Optional[int] = None + # 本请求预计可缓存的 checkpoint 尾部位置。 + self.hybrid_cache_len: Optional[int] = None # 存储对应长度位置的大页buffer_id - self.linear_att_len_to_big_page_id: Optional[SortedDict] = None + self.hybrid_len_to_big_page_id: Optional[SortedDict] = None # 在开启 enable_cpu_cache 的情况下,当请求结束后,会将请求的 kv cache # 卸载到 cpu cache 中,该标志变量用于标记请求的卸载任务的状态 @@ -611,9 +580,9 @@ def __init__( 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 - self.get_chuncked_input_token_ids = self.get_chuncked_input_token_ids_for_linear_att + if g_infer_context.is_hybrid_att_model: + self.get_chuncked_input_token_len = self.get_chuncked_input_token_len_for_hybrid_att + self.get_chuncked_input_token_ids = self.get_chuncked_input_token_ids_for_hybrid_att self._init_all_state() @@ -623,8 +592,8 @@ def __init__( self.generator.manual_seed(self.sampling_param.shm_param.seed) if init_prefix_cache: - if g_infer_context.is_linear_att_mixed_model: - self._linear_match_radix_cache() + if g_infer_context.is_hybrid_att_model: + self._hybrid_match_radix_cache() else: self._match_radix_cache() return @@ -648,22 +617,22 @@ def _init_all_state(self): self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() self.multimodal_params = self.multimodal_params.to_dict() - self.shared_kv_node: Union[TreeNode, LinearAttPagedTreeNode] = None + self.shared_kv_node: Union[TreeNode, HybridAttPagedTreeNode] = None self.finish_status = FinishStatus() - # 申请线性att混合模型使用的缓存资源 - if g_infer_context.is_linear_att_mixed_model: - linear_block_num = self.shm_req.linear_att_token_hash_list.size - self.linear_att_cache_len = linear_block_num * self.args.linear_att_hash_page_size - self.linear_att_len_to_big_page_id = SortedDict() + # 申请 hybrid attention 模型使用的缓存资源 + if g_infer_context.is_hybrid_att_model: + block_num = self.shm_req.hybrid_token_hash_list.size + self.hybrid_cache_len = block_num * self.args.linear_att_hash_page_size + self.hybrid_len_to_big_page_id = SortedDict() return def _match_radix_cache(self): assert ( - g_infer_context.is_linear_att_mixed_model is False - ), "current _match_radix_cache does not support linear att hybrid model, to do..." + g_infer_context.is_hybrid_att_model is False + ), "current _match_radix_cache does not support hybrid attention models, to do..." enable_prompt_cache = (not self.sampling_param.disable_prompt_cache) and g_infer_context.radix_cache is not None if enable_prompt_cache and self.get_cur_total_len() > 1 and self.cur_kv_len == 0: input_token_ids = self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] @@ -681,30 +650,30 @@ def _match_radix_cache(self): self.shm_req.shm_cur_kv_len = self.cur_kv_len return - def _linear_match_radix_cache(self): + def _hybrid_match_radix_cache(self): assert ( - g_infer_context.is_linear_att_mixed_model is True - ), "current _linear_match_radix_cache only support linear att hybrid model, to do..." + g_infer_context.is_hybrid_att_model is True + ), "current _hybrid_match_radix_cache only supports hybrid attention models, to do..." enable_prompt_cache = (not self.sampling_param.disable_prompt_cache) and g_infer_context.radix_cache is not None - linear_hash_list = self.shm_req.linear_att_token_hash_list.get_all() - linear_att_hash_page_size = self.args.linear_att_hash_page_size - match_tokens = min(len(linear_hash_list) * linear_att_hash_page_size, self.get_cur_total_len() - 1) + block_hashs = self.shm_req.hybrid_token_hash_list.get_all() + hash_page_size = self.args.linear_att_hash_page_size + match_tokens = min(len(block_hashs) * hash_page_size, self.get_cur_total_len() - 1) match_tokens = max(0, match_tokens) - match_tokens = (match_tokens // linear_att_hash_page_size) * linear_att_hash_page_size - match_block_num = match_tokens // linear_att_hash_page_size - linear_hash_list = linear_hash_list[:match_block_num] - assert len(linear_hash_list) == self.shm_req.linear_att_token_hash_list.size - big_page_token_num = linear_att_hash_page_size * self.args.linear_att_page_block_num + match_tokens = (match_tokens // hash_page_size) * hash_page_size + match_block_num = match_tokens // hash_page_size + block_hashs = block_hashs[:match_block_num] + assert len(block_hashs) == self.shm_req.hybrid_token_hash_list.size + big_page_token_num = hash_page_size * self.args.linear_att_page_block_num big_page_is_disable = big_page_token_num > self.args.max_req_total_len - if enable_prompt_cache and match_tokens > 1 and len(linear_hash_list) > 0 and self.cur_kv_len == 0: + if enable_prompt_cache and match_tokens > 1 and len(block_hashs) > 0 and self.cur_kv_len == 0: input_token_ids = self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] key = torch.tensor(input_token_ids[0:match_tokens], dtype=torch.int64, device="cpu") - assert len(key) == len(linear_hash_list) * linear_att_hash_page_size + assert len(key) == len(block_hashs) * hash_page_size share_node, kv_len, value_tensor = g_infer_context.radix_cache.match_prefix( - key, block_hashs=linear_hash_list, update_refs=True + key, block_hashs=block_hashs, update_refs=True ) if share_node is not None: - assert self.tail_linear_att_small_page_buffer_id is None + assert self.tail_small_page_buffer_id is None if share_node.is_big_page_node(): # 大页匹配 self.shared_kv_node = share_node @@ -713,9 +682,9 @@ def _linear_match_radix_cache(self): g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 - assert self.tail_linear_att_small_page_buffer_id is None - # 恢复linear att 状态 - g_infer_context.req_manager.copy_big_page_buffer_to_linear_att_state( + assert self.tail_small_page_buffer_id is None + # 恢复 hybrid checkpoint + g_infer_context.req_manager.restore_big_page_state( big_page_buffer_idx=share_node.big_page_buffer_idx, req=self ) else: @@ -728,11 +697,10 @@ def _linear_match_radix_cache(self): g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 - assert self.tail_linear_att_small_page_buffer_id is None - # 恢复linear att 状态 - g_infer_context.req_manager.copy_small_page_buffer_to_linear_att_state( + assert self.tail_small_page_buffer_id is None + # 恢复 hybrid checkpoint + g_infer_context.req_manager.restore_small_page_state( req=self, - linear_att_small_page_buffers=g_infer_context.radix_cache.linear_att_small_page_buffers, ) else: # 如果 大页本质是被启用的,则需要使用小页的匹配结果, 将小页的kv 复制到的新申请的kv位置,同时释放 @@ -762,10 +730,9 @@ def _linear_match_radix_cache(self): destination_indexes=tail_mems, ) - self.shared_kv_node = share_node # 只是为了保证 copy_small_page_buffer_to_linear_att_state 正确调用 - g_infer_context.req_manager.copy_small_page_buffer_to_linear_att_state( + self.shared_kv_node = share_node # 只是为了保证 restore_small_page_state 正确调用 + g_infer_context.req_manager.restore_small_page_state( req=self, - linear_att_small_page_buffers=g_infer_context.radix_cache.linear_att_small_page_buffers, ) self.shared_kv_node = None @@ -787,9 +754,9 @@ def _linear_match_radix_cache(self): ] = value_tensor[0:ready_cache_len] self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 - assert self.tail_linear_att_small_page_buffer_id is None - # 恢复linear att 状态 - g_infer_context.req_manager.copy_big_page_buffer_to_linear_att_state( + assert self.tail_small_page_buffer_id is None + # 恢复 hybrid checkpoint + g_infer_context.req_manager.restore_big_page_state( big_page_buffer_idx=share_node.big_page_buffer_idx, req=self ) @@ -797,7 +764,7 @@ def _linear_match_radix_cache(self): if self.cur_kv_len == 0: # 说明没有任何命中 - g_infer_context.req_manager.init_linear_att_state(req=self) + g_infer_context.req_manager.init_hybrid_attention_state(req=self) return def is_master_req(self): @@ -848,7 +815,7 @@ def get_chuncked_input_token_ids(self): chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) return self.shm_req.shm_prompt_ids.arr[0:chunked_end] - def get_chuncked_input_token_ids_for_linear_att(self): + def get_chuncked_input_token_ids_for_hybrid_att(self): big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num chunked_start = self.cur_kv_len @@ -857,9 +824,9 @@ def get_chuncked_input_token_ids_for_linear_att(self): total_end = self.get_cur_total_len() end = min(total_end, chunked_end, big_page_end) - if chunked_start < self.linear_att_cache_len < end: - # linear att cache 对应需要存储的部分。 - end = self.linear_att_cache_len + if chunked_start < self.hybrid_cache_len < end: + # hybrid checkpoint 对应需要存储的部分。 + end = self.hybrid_cache_len return self.shm_req.shm_prompt_ids.arr[0:end] @@ -868,15 +835,15 @@ def get_chuncked_input_token_len(self): chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) return chunked_end - def get_chuncked_input_token_len_for_linear_att(self): + def get_chuncked_input_token_len_for_hybrid_att(self): big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num chunked_start = self.cur_kv_len chunked_end = chunked_start + self.args.chunked_prefill_size big_page_end = ((chunked_start // big_page_token_num) + 1) * big_page_token_num total_end = self.get_cur_total_len() end = min(total_end, chunked_end, big_page_end) - if chunked_start < self.linear_att_cache_len < end: - end = self.linear_att_cache_len + if chunked_start < self.hybrid_cache_len < end: + end = self.hybrid_cache_len return end def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: int, rank: int = -1): 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 560c3b6de4..8ce61685b0 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -15,9 +15,8 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager -from lightllm.common.req_manager import ReqManagerForMamba -from lightllm.common.linear_att_cache_manager import LinearAttCacheManager -from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import LinearAttPagedRadixCache +from lightllm.common.req_manager import HybridAttentionReqManager +from lightllm.server.router.dynamic_prompt.hybrid_att_radix_cache import HybridAttPagedRadixCache from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput from lightllm.utils.dist_utils import init_distributed_env @@ -151,28 +150,27 @@ def init_model(self, kvargs): self.model, self.is_multimodal = get_model(model_cfg, model_kvargs) self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) - self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) + self.is_hybrid_att_model = isinstance(self.model.req_manager, HybridAttentionReqManager) - if self.is_linear_att_mixed_model: - self.linear_att_cache_manager = LinearAttCacheManager( + if self.is_hybrid_att_model: + self.small_page_buffers = self.model.req_manager.create_small_page_cache_manager( size=self.args.linear_att_cache_size, - linear_config=self.model.req_manager.linear_config, ) else: - self.linear_att_cache_manager = None + self.small_page_buffers = None if not self.use_dynamic_prompt_cache: self.radix_cache = None else: - if self.is_linear_att_mixed_model: - self.radix_cache = LinearAttPagedRadixCache( + if self.is_hybrid_att_model: + self.radix_cache = HybridAttPagedRadixCache( unique_name=get_unique_server_name(), total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, hash_page_size=self.args.linear_att_hash_page_size, big_page_num=self.args.linear_att_page_block_num, kv_cache_mem_manager=self.model.mem_manager, - linear_att_small_page_buffers=self.linear_att_cache_manager, + small_page_buffers=self.small_page_buffers, ) else: self.radix_cache = RadixCache( 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..2326c8c515 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 @@ -119,7 +119,7 @@ def prefill_normal( b_prefill_has_output_cpu=model_input.b_prefill_has_output_cpu, mask_func=self.prefill_mask_func, ) - g_infer_context.copy_linear_att_state_to_cache_buffer( + g_infer_context.save_hybrid_state_to_cache( b_req_idx=model_input.b_req_idx, reqs=run_reqs, ) @@ -215,7 +215,7 @@ def prefill_mtp( target_model_output=model_output, target_next_token_ids=next_token_ids, ) - g_infer_context.copy_linear_att_state_to_cache_buffer( + g_infer_context.save_hybrid_state_to_cache( b_req_idx=model_input.b_req_idx, reqs=run_reqs, ) 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..ce21d01987 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 @@ -195,7 +195,7 @@ def prefill_normal( b_prefill_has_output_cpu=model_input.b_prefill_has_output_cpu, mask_func=None, ) - g_infer_context.copy_linear_att_state_to_cache_buffer( + g_infer_context.save_hybrid_state_to_cache( b_req_idx=model_input.b_req_idx, reqs=run_reqs, ) @@ -316,8 +316,8 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer mask_func=None, ) - if g_infer_context.is_linear_att_mixed_model: - g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) + if g_infer_context.is_hybrid_att_model: + g_infer_context.save_hybrid_state_to_cache(b_req_idx=b_req_idx, reqs=run_reqs) sync_event = torch.cuda.Event() sync_event.record() @@ -440,7 +440,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] target_next_token_ids=next_token_ids, ) if req_num > 0: - g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) + g_infer_context.save_hybrid_state_to_cache(b_req_idx=b_req_idx, reqs=run_reqs) sync_event = torch.cuda.Event() sync_event.record() @@ -695,8 +695,8 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I target_next_token_ids1=target_next_token_ids_gpu1, ) - if req_num > 0 and g_infer_context.is_linear_att_mixed_model: - g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) + if req_num > 0 and g_infer_context.is_hybrid_att_model: + g_infer_context.save_hybrid_state_to_cache(b_req_idx=b_req_idx, reqs=run_reqs) sync_event = torch.cuda.Event() sync_event.record() diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py index 00489b9c27..2a366f0256 100644 --- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py @@ -8,7 +8,7 @@ from collections import deque from lightllm.server.multi_level_kv_cache import CacheTier from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuKvCacheClient -from lightllm.utils.config_utils import is_linear_att_mixed_model +from lightllm.utils.config_utils import is_hybrid_att_model from lightllm.utils.envs_utils import get_env_start_args from ..infer_batch import InferReq from lightllm.utils.dist_utils import create_new_group_for_current_dp @@ -115,8 +115,7 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): page_indexes_cuda = torch.tensor(need_pages, dtype=torch.int32, device="cpu").cuda( non_blocking=True ) - # 因为在支持 linear att 以后,所有的页面加载必须要按照 page页面的整数倍来做, - # 不然可能导致页面数据不完整,导致无法从kv中恢复完整的 linear att状态,所以 + # hybrid 页面加载必须按完整 page 处理,否则可能缺失恢复运行态所需的 checkpoint,所以 # 这里需要进行pad操作,使操作的页面是完整的。 _start = page_len_start_list[ready_page_num] @@ -176,7 +175,7 @@ def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> L continue # 过滤不适合进行 kv 卸载到 cpu cache 的请求。 - if g_infer_context.is_linear_att_mixed_model: + if g_infer_context.is_hybrid_att_model: offload_limit_size = self.args.linear_att_hash_page_size else: offload_limit_size = self.args.cpu_cache_token_page_size @@ -231,8 +230,8 @@ def _start_kv_cache_offload_task( find_index = bisect.bisect_right(page_len_list, req.cur_kv_len) move_block_size = find_index - # 对于 linear att 模型, 如果最后一个页面是碎页,需要做特殊处理,判断该碎页是否满足卸载条件。 - move_block_size = self._handle_linear_att_last_page( + # hybrid 模型的最后一个页面可能是碎页,需判断该碎页是否满足卸载条件。 + move_block_size = self._handle_hybrid_att_last_page( req=req, move_block_size=move_block_size, page_len_list=page_len_list ) @@ -308,8 +307,8 @@ def _start_kv_cache_offload_task( return trans_task - def _handle_linear_att_last_page(self, req: InferReq, move_block_size: int, page_len_list: List[int]) -> int: - if not g_infer_context.is_linear_att_mixed_model: + def _handle_hybrid_att_last_page(self, req: InferReq, move_block_size: int, page_len_list: List[int]) -> int: + if not g_infer_context.is_hybrid_att_model: return move_block_size if move_block_size == 0: @@ -322,7 +321,7 @@ def _handle_linear_att_last_page(self, req: InferReq, move_block_size: int, page if self.args.disable_linear_att_small_page_cpu_cache: return move_block_size - 1 # 说明是碎页,碎页需要判定是否满足cpu cache 的offload条件。 - if req.tail_linear_att_small_page_buffer_id is None: + if req.tail_small_page_buffer_id is None: return move_block_size - 1 return move_block_size diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 472049442f..17b4e92e48 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -1,6 +1,11 @@ import random import torch.multiprocessing as mp -from lightllm.server.pd_io_struct import PDChunckedTransTask, PDChunckedTransTaskGroup, PDAbortReq +from lightllm.server.pd_io_struct import ( + HYBRID_ATT_STATE_PAGE_KIND, + PDChunckedTransTask, + PDChunckedTransTaskGroup, + PDAbortReq, +) from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq @@ -158,15 +163,15 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): req_obj.cur_kv_len += len(mem_indexes) - # 如果当前是linear att 混合模型,则需要创建一个linear att 状态的传输任务 - if g_infer_context.is_linear_att_mixed_model: + # hybrid 模型额外传输 prompt 末尾的请求状态页。 + if g_infer_context.is_hybrid_att_model: self._create_pd_trans_task( req_obj=req_obj, mem_indexes=[], kv_start_index=input_len, kv_end_index=input_len, group=group, - page_kind="linear_att_state", + page_kind=HYBRID_ATT_STATE_PAGE_KIND, ) else: assert req_obj.cur_kv_len == input_len - 1 @@ -205,7 +210,7 @@ def _create_pd_trans_task( if page_kind == "kv": req_idx = None - elif page_kind == "linear_att_state": + elif page_kind == HYBRID_ATT_STATE_PAGE_KIND: req_idx = req_obj.req_idx else: raise ValueError(f"unknown PD trans page kind {page_kind}") diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 084e51e1e5..90a26ac601 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -2,7 +2,7 @@ import random from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import InferReq -from lightllm.server.pd_io_struct import PDAbortReq, PDChunckedTransTask +from lightllm.server.pd_io_struct import HYBRID_ATT_STATE_PAGE_KIND, PDAbortReq, PDChunckedTransTask from lightllm.utils.log_utils import init_logger from lightllm.utils.device_utils import kv_trans_use_p2p from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -87,13 +87,13 @@ def _prefill_chuncked_handle_func( break if prefill_finished and len(trans_task_list) != 0 and output_len == 1: - if g_infer_context.is_linear_att_mixed_model: + if g_infer_context.is_hybrid_att_model: trans_task_list.append( self._create_pd_trans_task( req_obj=req_obj, kv_start_index=input_len, kv_end_index=input_len, - page_kind="linear_att_state", + page_kind=HYBRID_ATT_STATE_PAGE_KIND, ) ) trans_task_list[-1].first_gen_token_id = next_token_id @@ -127,7 +127,7 @@ def _create_pd_trans_task( .tolist() ) req_idx = None - elif page_kind == "linear_att_state": + elif page_kind == HYBRID_ATT_STATE_PAGE_KIND: mem_indexes = [] req_idx = req_obj.req_idx else: diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 935fe92d33..cf84f8617e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -6,7 +6,6 @@ import torch from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_mtp_state_index_update, mtp_scatter_next_token_ids, mtp_verify, ) @@ -49,9 +48,8 @@ def verify_mtp_tokens( new_next_token_ids=next_token_ids, b_req_idx=b_req_idx, ) - if backend.is_linear_att_mixed_model: - linear_att_mtp_state_index_update( - req_to_mtp_state_index=backend.model.req_manager.req_to_mtp_state_index, + if backend.is_hybrid_att_model: + backend.model.req_manager.update_mtp_state( b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, diff --git a/lightllm/utils/backend_validator.py b/lightllm/utils/backend_validator.py index cd535871f0..0903b9b33c 100644 --- a/lightllm/utils/backend_validator.py +++ b/lightllm/utils/backend_validator.py @@ -99,7 +99,7 @@ def _validate_flashqla(): from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import ( chunk_gated_delta_rule as fla_chunk_gated_delta_rule, ) - from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig + from lightllm.common.state_cache_manager import LinearAttCacheConfig linear_config = LinearAttCacheConfig.load_from_args() num_k_heads = linear_config.num_linear_k_heads diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 0903955aed..4bd78d887f 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -468,6 +468,11 @@ def is_linear_att_mixed_model(model_path: str) -> bool: return False +def is_hybrid_att_model(model_path: str) -> bool: + """Models whose non-full attention state follows hybrid checkpoint pages.""" + return is_linear_att_mixed_model(model_path) + + def get_model_type(model_path: str) -> Optional[str]: """Get model type from config.json""" try: diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index e81caafe7a..40383b2eba 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -16,21 +16,20 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger -from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model +from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_hybrid_att_model from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( MemoryManager, PPLINT8KVMemoryManager, PPLINT4KVMemoryManager, Deepseek2MemoryManager, - Qwen3NextMemManager, ) from typing import List, Tuple, Optional from tqdm import tqdm from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup from lightllm.utils.dist_utils import get_current_device_id -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import get_hybrid_cache_config logger = init_logger(__name__) @@ -64,20 +63,16 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": args = get_env_start_args() assert args.enable_cpu_cache - if is_linear_att_mixed_model(args.model_dir): - # 对于 qwen3.5 等 linear att 混合模型的特殊处理。 - mem_manager_class = Qwen3NextMemManager - else: - mem_manager_class = select_mem_manager_class() - - if mem_manager_class is Qwen3NextMemManager: - linear_config = LinearAttCacheConfig.load_from_args() + is_hybrid_model = is_hybrid_att_model(args.model_dir) + mem_manager_class = None if is_hybrid_model else select_mem_manager_class() + if is_hybrid_model: + hybrid_config = get_hybrid_cache_config() cpu_cache_meta = CpuKVCacheMeta( page_num=0, token_page_size=1, layer_num=1, num_heads=1, - head_dim=linear_config.get_cpu_cache_big_page_bytes(), + head_dim=hybrid_config.get_cpu_cache_big_page_bytes(), data_type=torch.uint8, scale_head_dim=0, scale_data_type=get_llm_data_type(), @@ -121,9 +116,9 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": if args.mtp_mode is not None: # TODO 可能会存在不同mtp模式的精度问题 - if not is_linear_att_mixed_model(args.model_dir): - # 对于非 linear att 混合模型,需要额外增加 mtp 的 kv 层数, - # 对于 linear att 混合模型,如qwen 3.5 mtp,已经将 kv 数据 + if not is_hybrid_model: + # 对于非 hybrid 模型,需要额外增加 mtp 的 kv 层数, + # 对于 hybrid 模型,如 qwen 3.5 mtp,已经将 kv 数据 # 打包成一个块了,所以不需要额外增加,其 layer_num 一直都保持为 1 cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() diff --git a/test/cpu_cache_kernel/test_speed.py b/test/cpu_cache_kernel/test_speed.py index 254142050c..f187088883 100644 --- a/test/cpu_cache_kernel/test_speed.py +++ b/test/cpu_cache_kernel/test_speed.py @@ -39,7 +39,7 @@ # --------------------------------------------------------------------------- # Step 1 – build LinearAttCacheConfig directly (avoids needing a real model dir) # --------------------------------------------------------------------------- -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import LinearAttCacheConfig linear_config = LinearAttCacheConfig( tp_world_size=8, diff --git a/unit_tests/common/basemodel/attention/linear/test_gdn.py b/unit_tests/common/basemodel/attention/linear/test_gdn.py index 12b6996a83..d6574de151 100644 --- a/unit_tests/common/basemodel/attention/linear/test_gdn.py +++ b/unit_tests/common/basemodel/attention/linear/test_gdn.py @@ -10,7 +10,7 @@ import lightllm.common.basemodel.triton_kernel.linear_att.fla.ops as fla_ops from lightllm.common.basemodel.attention.linear.flashqla import FlashQlaLinearAttBackend from lightllm.common.basemodel.attention.linear.triton import TritonLinearAttBackend -from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig +from lightllm.common.state_cache_manager import LinearAttCacheConfig from lightllm.server.api_cli import make_argument_parser import lightllm.utils.backend_validator as backend_validator diff --git a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py index 96f295efd2..96925b6cb2 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py @@ -41,7 +41,7 @@ def test_cache_tiers_reassignment_is_rejected(): def test_non_gpu_cache_tiers_release_owned_tokens_without_radix_insert(): released_refs = [] context = InferenceContext() - context.is_linear_att_mixed_model = False + context.is_hybrid_att_model = False context.req_manager = SimpleNamespace(req_to_token_indexs=torch.tensor([[10, 11, 12, 13, 14]])) context.radix_cache = SimpleNamespace(dec_node_ref_counter=released_refs.append) shared_node = SimpleNamespace(node_prefix_total_len=2) @@ -58,7 +58,7 @@ def test_non_gpu_cache_tiers_release_owned_tokens_without_radix_insert(): def test_legacy_cache_tiers_still_insert_gpu_radix_cache(): context = InferenceContext() context.radix_cache = object() - context.is_linear_att_mixed_model = False + context.is_hybrid_att_model = False inserted_reqs = [] context._full_att_free_req = lambda free_token_index, req: inserted_reqs.append(req) req = SimpleNamespace( @@ -98,7 +98,7 @@ def start_offload(req, cpu_kv_cache_stream): return SimpleNamespace(req=req) module._start_kv_cache_offload_task = start_offload - monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "is_linear_att_mixed_model", False) + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "is_hybrid_att_model", False) monkeypatch.setattr( multi_level_kv_cache_impl.g_infer_context, "get_cpu_kv_cache_stream", @@ -135,18 +135,18 @@ def test_non_gpu_linear_cache_tiers_release_pending_state_pages(): freed_small_pages = [] freed_big_pages = [] context = InferenceContext() - context.is_linear_att_mixed_model = True + context.is_hybrid_att_model = True context.req_manager = SimpleNamespace(req_to_token_indexs=torch.tensor([[10, 11, 12]])) context.radix_cache = SimpleNamespace( - linear_att_small_page_buffers=SimpleNamespace(free_state_cache=freed_small_pages.extend), - linear_att_big_page_buffers=SimpleNamespace(free_state_cache=freed_big_pages.extend), + small_page_buffers=SimpleNamespace(free_state_cache=freed_small_pages.extend), + big_page_buffers=SimpleNamespace(free_state_cache=freed_big_pages.extend), ) req = SimpleNamespace( req_idx=0, cur_kv_len=3, shared_kv_node=None, - tail_linear_att_small_page_buffer_id=7, - linear_att_len_to_big_page_id={128: 8, 256: 9}, + tail_small_page_buffer_id=7, + hybrid_len_to_big_page_id={128: 8, 256: 9}, ) free_token_indexes = [] @@ -155,5 +155,5 @@ def test_non_gpu_linear_cache_tiers_release_pending_state_pages(): assert free_token_indexes[0].tolist() == [10, 11, 12] assert freed_small_pages == [7] assert freed_big_pages == [8, 9] - assert req.tail_linear_att_small_page_buffer_id is None - assert req.linear_att_len_to_big_page_id == {} + assert req.tail_small_page_buffer_id is None + assert req.hybrid_len_to_big_page_id == {} From 5f1c96a34cc0eed71ce747aa5ef3a3f319719797 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 10 Sep 2026 03:09:51 +0000 Subject: [PATCH 2/4] refactor: preserve main behavior in hybrid cache extraction --- .../kv_cache_mem_manager/qwen3next_mem_manager.py | 5 ++--- lightllm/common/state_cache_manager/linear_att.py | 1 + lightllm/server/pd_io_struct.py | 5 +---- lightllm/server/router/model_infer/infer_batch.py | 2 +- .../mode_backend/pd/decode_node_impl/decode_impl.py | 11 +++-------- .../mode_backend/pd/prefill_node_impl/prefill_impl.py | 6 +++--- 6 files changed, 11 insertions(+), 19 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py index 384269cb17..78781f3300 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -1,5 +1,4 @@ import torch -from lightllm.server.pd_io_struct import HYBRID_ATT_STATE_PAGE_KIND import triton from lightllm.utils.log_utils import init_logger from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager @@ -105,7 +104,7 @@ def write_mem_to_page_kv_move_buffer( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == HYBRID_ATT_STATE_PAGE_KIND, f"unknown page_kind={page_kind}" + assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) @@ -132,7 +131,7 @@ def read_page_kv_move_buffer_to_mem( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == HYBRID_ATT_STATE_PAGE_KIND, f"unknown page_kind={page_kind}" + assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) diff --git a/lightllm/common/state_cache_manager/linear_att.py b/lightllm/common/state_cache_manager/linear_att.py index a9ee5cf1b4..aa7f5e439a 100644 --- a/lightllm/common/state_cache_manager/linear_att.py +++ b/lightllm/common/state_cache_manager/linear_att.py @@ -178,6 +178,7 @@ def __init__( device="cpu", size_first=True, ) + self.clear_to_init_state() return def get_state_cache(self, buffer_idx: int): diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index db477df8c6..78f5fedc93 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -10,9 +10,6 @@ logger = init_logger(__name__) -# Keep the existing wire tag so P/D nodes can be upgraded independently. -HYBRID_ATT_STATE_PAGE_KIND = "linear_att_state" - # 节点的行为 class NodeRole(enum.Enum): @@ -193,7 +190,7 @@ def __post_init__(self): raise ValueError(error_info) if self.page_kind == "kv": assert len(self.mem_indexes) == (self.end_kv_index - self.start_kv_index) - elif self.page_kind == HYBRID_ATT_STATE_PAGE_KIND: + elif self.page_kind == "linear_att_state": assert self.start_kv_index == self.end_kv_index assert len(self.mem_indexes) == 0 else: diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index dd37dace9e..18a8f6c04f 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -392,7 +392,7 @@ def get_can_alloc_token_num(self): def save_hybrid_state_to_cache(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """Snapshot request-level attention state at big/small-page boundaries.""" - if not self.is_hybrid_att_model or self.radix_cache is None: + if not self.is_hybrid_att_model: return # Request-state snapshot at a big-page boundary. diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 17b4e92e48..7170d6c5e5 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -1,11 +1,6 @@ import random import torch.multiprocessing as mp -from lightllm.server.pd_io_struct import ( - HYBRID_ATT_STATE_PAGE_KIND, - PDChunckedTransTask, - PDChunckedTransTaskGroup, - PDAbortReq, -) +from lightllm.server.pd_io_struct import PDChunckedTransTask, PDChunckedTransTaskGroup, PDAbortReq from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq @@ -171,7 +166,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): kv_start_index=input_len, kv_end_index=input_len, group=group, - page_kind=HYBRID_ATT_STATE_PAGE_KIND, + page_kind="linear_att_state", ) else: assert req_obj.cur_kv_len == input_len - 1 @@ -210,7 +205,7 @@ def _create_pd_trans_task( if page_kind == "kv": req_idx = None - elif page_kind == HYBRID_ATT_STATE_PAGE_KIND: + elif page_kind == "linear_att_state": req_idx = req_obj.req_idx else: raise ValueError(f"unknown PD trans page kind {page_kind}") diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 90a26ac601..ac287e7097 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -2,7 +2,7 @@ import random from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import InferReq -from lightllm.server.pd_io_struct import HYBRID_ATT_STATE_PAGE_KIND, PDAbortReq, PDChunckedTransTask +from lightllm.server.pd_io_struct import PDAbortReq, PDChunckedTransTask from lightllm.utils.log_utils import init_logger from lightllm.utils.device_utils import kv_trans_use_p2p from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -93,7 +93,7 @@ def _prefill_chuncked_handle_func( req_obj=req_obj, kv_start_index=input_len, kv_end_index=input_len, - page_kind=HYBRID_ATT_STATE_PAGE_KIND, + page_kind="linear_att_state", ) ) trans_task_list[-1].first_gen_token_id = next_token_id @@ -127,7 +127,7 @@ def _create_pd_trans_task( .tolist() ) req_idx = None - elif page_kind == HYBRID_ATT_STATE_PAGE_KIND: + elif page_kind == "linear_att_state": mem_indexes = [] req_idx = req_obj.req_idx else: From 13d32f5111e88d6e28184cafe6231d048e9950e8 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 10 Sep 2026 08:45:05 +0000 Subject: [PATCH 3/4] refactor: generalize PD attention state page naming --- .../common/kv_cache_mem_manager/qwen3next_mem_manager.py | 4 ++-- lightllm/server/pd_io_struct.py | 7 ++++++- .../mode_backend/pd/decode_node_impl/decode_impl.py | 6 +++--- .../mode_backend/pd/prefill_node_impl/prefill_impl.py | 5 +++-- skills/test_model/qwen3.5-0.8b-pd-nixl/SKILL.md | 6 +++--- 5 files changed, 17 insertions(+), 11 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py index 78781f3300..e03e9f08da 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -104,7 +104,7 @@ def write_mem_to_page_kv_move_buffer( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" + assert page_kind == "att_state", f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) @@ -131,7 +131,7 @@ def read_page_kv_move_buffer_to_mem( page_kind=page_kind, req_idx=req_idx, ) - assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" + assert page_kind == "att_state", f"unknown page_kind={page_kind}" assert req_idx is not None helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 78f5fedc93..39f878fb9d 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -179,6 +179,9 @@ class PDChunckedTransTask: error_info: Optional[str] = None transfer_time_out_secs: int = 66 + # kv: 通过 mem_indexes 传输 [start_kv_index, end_kv_index) 的 token KV。 + # att_state: start_kv_index == end_kv_index 处的 attention 续算状态(如 conv/SSM)。 + # 状态通过本地 req_idx 寻址,mem_indexes 为空;具体打包和恢复由模型 mem_manager 负责。 page_kind: str = "kv" # Only valid for the local task owner; remote notify copies may carry the sender-local req_idx. req_idx: Optional[int] = None @@ -190,7 +193,7 @@ def __post_init__(self): raise ValueError(error_info) if self.page_kind == "kv": assert len(self.mem_indexes) == (self.end_kv_index - self.start_kv_index) - elif self.page_kind == "linear_att_state": + elif self.page_kind == "att_state": assert self.start_kv_index == self.end_kv_index assert len(self.mem_indexes) == 0 else: @@ -216,6 +219,7 @@ def transfer_time(self): return time.time() - self.start_trans_time def get_key(self) -> str: + # page_kind 参与 P/D 任务匹配,发送端和接收端必须使用一致的协议取值。 return f"{self.request_id}_{self.page_kind}_{self.start_kv_index}_{self.end_kv_index}" def to_str(self): @@ -237,6 +241,7 @@ def transfer_kv_num(self): return self.end_kv_index - self.start_kv_index def need_transfer_page(self): + # att_state 虽然没有 token 区间,仍需传一页;空 kv 任务仅用于完成通知。 return self.page_kind != "kv" or self.transfer_kv_num() != 0 def createRetObj(self) -> "PDChunckedTransTaskRet": diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 7170d6c5e5..80294cc9f5 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -158,7 +158,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): req_obj.cur_kv_len += len(mem_indexes) - # hybrid 模型额外传输 prompt 末尾的请求状态页。 + # 额外接收 prompt 末尾的 attention 续算状态,由本地 req_idx 定位恢复位置。 if g_infer_context.is_hybrid_att_model: self._create_pd_trans_task( req_obj=req_obj, @@ -166,7 +166,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): kv_start_index=input_len, kv_end_index=input_len, group=group, - page_kind="linear_att_state", + page_kind="att_state", ) else: assert req_obj.cur_kv_len == input_len - 1 @@ -205,7 +205,7 @@ def _create_pd_trans_task( if page_kind == "kv": req_idx = None - elif page_kind == "linear_att_state": + elif page_kind == "att_state": req_idx = req_obj.req_idx else: raise ValueError(f"unknown PD trans page kind {page_kind}") diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index ac287e7097..67186f8ddb 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -88,12 +88,13 @@ def _prefill_chuncked_handle_func( if prefill_finished and len(trans_task_list) != 0 and output_len == 1: if g_infer_context.is_hybrid_att_model: + # KV 分段之外,单独发送 prompt 末尾的 attention 续算状态。 trans_task_list.append( self._create_pd_trans_task( req_obj=req_obj, kv_start_index=input_len, kv_end_index=input_len, - page_kind="linear_att_state", + page_kind="att_state", ) ) trans_task_list[-1].first_gen_token_id = next_token_id @@ -127,7 +128,7 @@ def _create_pd_trans_task( .tolist() ) req_idx = None - elif page_kind == "linear_att_state": + elif page_kind == "att_state": mem_indexes = [] req_idx = req_obj.req_idx else: diff --git a/skills/test_model/qwen3.5-0.8b-pd-nixl/SKILL.md b/skills/test_model/qwen3.5-0.8b-pd-nixl/SKILL.md index b5f775a581..ee562f68a7 100644 --- a/skills/test_model/qwen3.5-0.8b-pd-nixl/SKILL.md +++ b/skills/test_model/qwen3.5-0.8b-pd-nixl/SKILL.md @@ -24,7 +24,7 @@ Qwen3.5 与 Qwen3-8B 的关键差异: | 项 | Qwen3.5-0.8B NIXL PD 要点 | |---|---| -| linear-att 状态 | PD 传输除了 KV page,还会传 `linear_att_state` 特殊页 | +| attention 状态 | PD 传输除了 KV page,还会传 `att_state` 续算状态页(本模型为 conv/SSM) | | NIXL page size | 建议固定 **`--pd_kv_page_size 2048`**;`1024` 可能不足以容纳 linear-att 状态 | | page num | 建议 **`--pd_kv_page_num 16`** 起步,避免 page 池过大导致显存压力 | | cache 判断 | repeated prompt 可能只在 prefill 侧命中,decode 侧不一定 decode-only 命中 | @@ -242,7 +242,7 @@ rg -n 'flexible-extract|strict-match|exact_match|Traceback|ERROR|can not find wa - prefill 侧会按 512 token 粒度逐步命中,例如 513 的第二次可命中 512。 - decode 侧可能仍为 `gpu cache hit: False`、`gpu_prompt_cache_len:0`。 -- 只要 decode 未全命中,仍会出现 `recv WRITE request from prefill` 和 `linear_att_state` 传输。 +- 只要 decode 未全命中,仍会出现 `recv WRITE request from prefill` 和 `att_state` 传输。 ### 简单重复 prompt @@ -275,7 +275,7 @@ done ### 判定信号 ```bash -rg -n 'gpu cache hit:|recv WRITE request from prefill|start WRITE to decode node|linear_att_state|trans task ret success' \ +rg -n 'gpu cache hit:|recv WRITE request from prefill|start WRITE to decode node|att_state|trans task ret success' \ "${LOG_DIR}/prefill.log" "${LOG_DIR}/decode.log" \ | tee -a "${LOG_DIR}/summary.txt" ``` From b9ee6fceb4980428e962167c0334da04db3e4476 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 10 Sep 2026 10:49:46 +0000 Subject: [PATCH 4/4] docs: clarify hybrid PD request-state buffer transfer --- lightllm/server/pd_io_struct.py | 5 +++-- .../mode_backend/pd/decode_node_impl/decode_impl.py | 3 ++- .../mode_backend/pd/prefill_node_impl/prefill_impl.py | 2 +- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 39f878fb9d..f479279a26 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -180,8 +180,9 @@ class PDChunckedTransTask: error_info: Optional[str] = None transfer_time_out_secs: int = 66 # kv: 通过 mem_indexes 传输 [start_kv_index, end_kv_index) 的 token KV。 - # att_state: start_kv_index == end_kv_index 处的 attention 续算状态(如 conv/SSM)。 - # 状态通过本地 req_idx 寻址,mem_indexes 为空;具体打包和恢复由模型 mem_manager 负责。 + # att_state: 混合注意力模型的请求运行态 buffer(如 linear attention 的 conv/SSM 状态)。 + # start_kv_index == end_kv_index 标记状态对应的 token 位置,mem_indexes 为空。 + # 通过本地 req_idx 定位运行态 buffer,具体打包和恢复由模型 mem_manager 负责。 page_kind: str = "kv" # Only valid for the local task owner; remote notify copies may carry the sender-local req_idx. req_idx: Optional[int] = None diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 80294cc9f5..b95258ad92 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -158,7 +158,8 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): req_obj.cur_kv_len += len(mem_indexes) - # 额外接收 prompt 末尾的 attention 续算状态,由本地 req_idx 定位恢复位置。 + # 混合注意力模型还需接收请求运行态 buffer(如 linear attention 的 conv/SSM 状态)。 + # 通过本地 req_idx 定位运行态 buffer 的恢复位置。 if g_infer_context.is_hybrid_att_model: self._create_pd_trans_task( req_obj=req_obj, diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 67186f8ddb..0199c6e4e0 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -88,7 +88,7 @@ def _prefill_chuncked_handle_func( if prefill_finished and len(trans_task_list) != 0 and output_len == 1: if g_infer_context.is_hybrid_att_model: - # KV 分段之外,单独发送 prompt 末尾的 attention 续算状态。 + # 混合注意力模型除 KV 外,还需传输 prefill 完成时的请求运行态 buffer(如 linear attention 的 conv/SSM 状态)。 trans_task_list.append( self._create_pd_trans_task( req_obj=req_obj,