Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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
Expand Down
26 changes: 13 additions & 13 deletions lightllm/common/kv_cache_mem_manager/operator/linear_att.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(),
Expand All @@ -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,
)
Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand Down Expand Up @@ -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(),
Expand Down
20 changes: 10 additions & 10 deletions lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
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

Expand Down Expand Up @@ -45,14 +45,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):
Expand All @@ -61,7 +61,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):
Expand All @@ -72,12 +72,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)
Expand All @@ -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)
Expand All @@ -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)
Expand Down
3 changes: 0 additions & 3 deletions lightllm/common/linear_att_cache_manager/__init__.py

This file was deleted.

This file was deleted.

3 changes: 2 additions & 1 deletion lightllm/common/req_manager/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
70 changes: 70 additions & 0 deletions lightllm/common/req_manager/hybrid_base.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading