Skip to content
Open
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
45 changes: 45 additions & 0 deletions docs/kv_cache_page_size.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
# KV Cache 多 Token 页设计

## 目标与不变量

启动参数 `--page_size N` 控制注意力 KV Cache 的物理页大小,默认值为 `1`。`N > 1` 时保持 token 级的
KV 存储和请求表,但资源的申请、缓存和释放以完整物理页为单位。

实现依赖以下不变量:

1. `cur_kv_len` 是请求已经写入有效 KV 的逻辑 token 数;`hold_kv_len` 是请求拥有的物理容量,始终满足
`cur_kv_len <= hold_kv_len` 且 `hold_kv_len % page_size == 0`。
2. 一个物理页内的 KV 槽连续,页首索引可被 `page_size` 整除。请求表保存已持有页的全部 token 索引,
包括尚未使用的尾部槽位。
3. `req_to_token_indexs` 保存请求拥有的完整页;创建 `InferState` 时根据请求索引和序列位置直接聚合本轮
真实参与计算的 token,`ModelInput` 不携带物理 KV 索引,`InferState.mem_index` 不包含预留尾部。
4. Radix Cache 只插入、拆分、命中和淘汰完整页;不足一页的请求尾部在请求结束或暂停时整页回收。
5. `page_size=1` 与多 token 页使用相同的调度期预留和模型执行期索引选择路径。

页容量计算和整页申请由调度层在请求获准进入 Prefill/Decode batch 时完成;输入构造层只构造本轮输入,
真实推理索引由模型执行层从请求表选取。
`ReqManager` 只保留请求 ID、请求表及其原有生命周期职责,不提供 page allocator 接口。

## 生命周期

- Prefill:请求通过调度后,按目标 KV 长度将容量向上对齐,新页一次申请并完整写入请求表。模型执行时再选择
`[cur_kv_len, target_kv_len)` 对应的真实索引。
- Decode:请求通过调度后,若尾页仍有预留槽位则不再访问 allocator;容量不足时先补齐完整页再执行。
- Prefix Cache 命中:命中长度天然是整页,`cur_kv_len` 与 `hold_kv_len` 同时初始化为命中长度。
- 完成/暂停:完整逻辑页可进入 Radix Cache;重复页和未完成尾页按物理页展开后释放。
- Attention:FA3/FlashInfer 页表每项由物理 token 页首索引除以 `page_size` 得到,KV buffer 视为
`[page_count, page_size, ...]`。

## 当前兼容范围

首版支持普通与分块 Prefill、Decode、动态 Prompt Cache,以及非量化 FA3/FlashInfer(含 MLA Decode)。
MTP、PD 分离、CPU KV Cache、DP Prompt Cache 拉取、DP Prefill Balance、diverse mode 和混合线性注意力
拥有额外的 KV 申请或迁移语义;这些组合在启动阶段明确报错,后续应在各自模块接入统一的页所有权接口后再开放。

## 边界处理

- KV 总容量和请求表宽度分别向下、向上对齐,防止物理页越界。
- CUDA Graph 的 HOLD request 映射到额外保留的一整个物理页;请求表按
`[HOLD, HOLD+1, ..., HOLD+page_size-1]` 循环填充,而不是重复同一个 token 索引。
- Radix 子节点用“首个完整 token 页”作为键,避免不同序列仅首 token 相同造成页级分支冲突。
- 非法的 `page_size < 1` 以及尚未支持的功能组合在模型加载前失败。
37 changes: 23 additions & 14 deletions lightllm/common/basemodel/attention/fa3/fp.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import dataclasses
import torch
import triton
from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl
from typing import Optional, TYPE_CHECKING
from lightllm.utils.sgl_utils import flash_attn_with_kvcache, flash_attn_with_kvcache_autotune
Expand All @@ -18,6 +19,7 @@ class Fa3AttBackend(BaseAttBackend):

def __init__(self, model):
super().__init__(model=model)
self.page_size = model.args.page_size

# 延迟到首次获取 page table 时再初始化,避免 PD 分离模式下的 prefill 节点
# 分配仅供 decode 使用的 buffer,减少显存浪费。
Expand All @@ -37,7 +39,7 @@ def _init_page_table_buffers(self):
self.page_table_max_batch_size = max(running_max_batch_size, model.graph_max_batch_size)
# max_seq_length is max_req_total_len plus the MTP headroom reserved when
# the model is initialized.
self.page_table_max_seq_len = model.max_seq_length
self.page_table_max_seq_len = triton.cdiv(model.max_seq_length, self.page_size)
buffer_count = 2 if args.enable_decode_microbatch_overlap else 1
workspace_size = self.page_table_max_batch_size * self.page_table_max_seq_len
self.page_table_buffers = [
Expand All @@ -57,12 +59,13 @@ def get_page_table_view(self, att_batch_size, max_kv_len, microbatch_index):
f"FA3 attention batch size {att_batch_size} exceeds page-table capacity "
f"{self.page_table_max_batch_size}"
)
if max_kv_len > self.page_table_max_seq_len:
max_page_len = triton.cdiv(max_kv_len, self.page_size)
if max_page_len > self.page_table_max_seq_len:
raise RuntimeError(
f"FA3 max KV sequence length {max_kv_len} exceeds page-table capacity " f"{self.page_table_max_seq_len}"
)
return self.page_table_buffers[microbatch_index][: att_batch_size * max_kv_len].reshape(
att_batch_size, max_kv_len
return self.page_table_buffers[microbatch_index][: att_batch_size * max_page_len].reshape(
att_batch_size, max_page_len
)

def create_att_prefill_state(self, infer_state) -> "Fa3PrefillAttState":
Expand All @@ -84,14 +87,18 @@ def init_state(self):
self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int()
self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int()
self.page_table = torch.empty(
(self.infer_state.batch_size, self.infer_state.max_kv_seq_len),
(
self.infer_state.batch_size,
triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size),
),
dtype=torch.int32,
device=self.infer_state.input_ids.device,
)
self.page_table.copy_(
self.infer_state.req_manager.req_to_token_indexs[
self.infer_state.b_req_idx, : self.infer_state.max_kv_seq_len
]
page_table_copy(
page_table=self.page_table,
req_to_token_indexs=self.infer_state.req_manager.req_to_token_indexs,
b_req_idx=self.infer_state.b_req_idx,
page_size=self.backend.page_size,
)

def prefill_att(
Expand Down Expand Up @@ -131,8 +138,8 @@ def _nomarl_prefill_att(
sm_scale = 1.0 / (Lq ** 0.5)
o = flash_attn_with_kvcache(
q=q,
k_cache=k.view(k.shape[0], 1, k.shape[1], k.shape[2]),
v_cache=v.view(v.shape[0], 1, v.shape[1], v.shape[2]),
k_cache=k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]),
v_cache=v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]),
page_table=self.page_table,
cache_seqlens=self.infer_state.b_seq_len,
cu_seqlens_q=self.cu_seqlens_q,
Expand Down Expand Up @@ -225,6 +232,7 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor):
att_batch_size = b_att_req_idx.shape[0]
model = self.backend.model
actual_max_kv_len = self.infer_state.max_kv_seq_len
actual_max_page_len = triton.cdiv(actual_max_kv_len, self.backend.page_size)
page_table_width = actual_max_kv_len
if model.graph is not None and model.graph.can_run(
batch_size=self.infer_state.batch_size,
Expand All @@ -242,9 +250,10 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor):
)

page_table_copy(
page_table=self.page_table[:, :actual_max_kv_len],
page_table=self.page_table[:, :actual_max_page_len],
req_to_token_indexs=model.req_manager.req_to_token_indexs,
b_req_idx=b_att_req_idx,
page_size=self.backend.page_size,
)

def copy_for_decode_cuda_graph(self, new_state: "Fa3DecodeAttState"):
Expand Down Expand Up @@ -290,8 +299,8 @@ def _normal_decode_att(
sm_scale = 1.0 / (Lq ** 0.5)
o = flash_attn_with_kvcache_autotune(
q=q,
k_cache=k.view(k.shape[0], 1, k.shape[1], k.shape[2]),
v_cache=v.view(v.shape[0], 1, v.shape[1], v.shape[2]),
k_cache=k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]),
v_cache=v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]),
page_table=self.page_table,
cache_seqlens=self.b_att_seq_len,
cu_seqlens_q=self.cu_seqlens_q,
Expand Down
9 changes: 6 additions & 3 deletions lightllm/common/basemodel/attention/fa3/mla.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import dataclasses
import torch
import triton
from ..base_att import BasePrefillAttState, BaseDecodeAttState, AttControl
from typing import Optional, TYPE_CHECKING, Tuple
from lightllm.utils.sgl_utils import flash_attn_with_kvcache
Expand Down Expand Up @@ -153,6 +154,7 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor):
att_batch_size = b_att_req_idx.shape[0]
model = self.backend.model
actual_max_kv_len = self.infer_state.max_kv_seq_len
actual_max_page_len = triton.cdiv(actual_max_kv_len, self.backend.page_size)
page_table_width = actual_max_kv_len
if model.graph is not None and model.graph.can_run(
batch_size=self.infer_state.batch_size,
Expand All @@ -170,9 +172,10 @@ def _init_page_table(self, b_att_req_idx: torch.Tensor):
)

page_table_copy(
page_table=self.page_table[:, :actual_max_kv_len],
page_table=self.page_table[:, :actual_max_page_len],
req_to_token_indexs=model.req_manager.req_to_token_indexs,
b_req_idx=b_att_req_idx,
page_size=self.backend.page_size,
)

def copy_for_decode_cuda_graph(self, new_state: "MlaFa3DecodeAttState"):
Expand Down Expand Up @@ -213,8 +216,8 @@ def _mla_decode_att(
kv = k
qk_rope_head_dim = 64
kv_lora_rank = kv.shape[-1] - qk_rope_head_dim
k_rope = kv[:, :, -qk_rope_head_dim:].view(-1, 1, 1, qk_rope_head_dim)
kv_nope = kv[:, :, :-qk_rope_head_dim].view(-1, 1, 1, kv_lora_rank)
k_rope = kv[:, :, -qk_rope_head_dim:].view(-1, self.backend.page_size, 1, qk_rope_head_dim)
kv_nope = kv[:, :, :-qk_rope_head_dim].view(-1, self.backend.page_size, 1, kv_lora_rank)
k_descale, v_descale = None, None
assert att_control.mla_decode
softmax_scale = att_control.mla_decode_dict["softmax_scale"]
Expand Down
105 changes: 71 additions & 34 deletions lightllm/common/basemodel/attention/flashinfer/fp.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import dataclasses
import torch
import triton
from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl
from lightllm.utils.dist_utils import get_dp_world_size, get_current_device_id
from ...triton_kernel.repack_kv_index import repack_kv_index
from ...triton_kernel.repack_kv_index import repack_kv_index, repack_page_kv_index
from .env_utils import set_flashinfer_envs
from .utils import should_init_decode_wrapper

Expand All @@ -14,6 +15,7 @@ class FlashInferAttBackend(BaseAttBackend):
def __init__(self, model):
set_flashinfer_envs()
super().__init__(model=model)
self.page_size = model.args.page_size
tp_world_size = get_dp_world_size()
self.tp_q_head_num = model.config["num_attention_heads"] // tp_world_size
self.tp_kv_head_num = max(model.config["num_key_value_heads"] // tp_world_size, 1)
Expand All @@ -22,10 +24,14 @@ def __init__(self, model):
self.max_seq_length = model.max_seq_length
self.kv_indices_buffer = [
torch.empty(
model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id()
model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size),
dtype=torch.int32,
device=get_current_device_id(),
),
torch.empty(
model.graph_max_batch_size * self.max_seq_length, dtype=torch.int32, device=get_current_device_id()
model.graph_max_batch_size * triton.cdiv(self.max_seq_length, self.page_size),
dtype=torch.int32,
device=get_current_device_id(),
),
]
self.q_data_type = model.data_type
Expand All @@ -51,21 +57,34 @@ def init_state(self):
device = self.infer_state.input_ids.device

q_starts = self.infer_state.b1_cu_q_seq_len.int()
kv_starts = self.infer_state.b1_cu_kv_seq_len.int()
kv_last_page_len = torch.full((batch_size,), 1, dtype=torch.int32, device=device)
kv_starts = self.infer_state.b1_cu_kv_seq_len.int().clone()
b_page_len = triton.cdiv(self.infer_state.b_seq_len, self.backend.page_size)
kv_starts[1:] = b_page_len.cumsum(0)
kv_last_page_len = self.infer_state.b_seq_len - (b_page_len - 1) * self.backend.page_size
kv_indices = torch.empty(
batch_size * self.backend.max_seq_length,
batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size),
dtype=torch.int32,
device=device,
)
repack_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
self.infer_state.b_seq_len,
kv_starts[:-1],
self.infer_state.max_kv_seq_len,
kv_indices,
)
if self.backend.page_size == 1:
repack_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
self.infer_state.b_seq_len,
kv_starts[:-1],
self.infer_state.max_kv_seq_len,
kv_indices,
)
else:
repack_page_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
b_page_len,
kv_starts[:-1],
triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size),
kv_indices,
self.backend.page_size,
)
self.prefill_wrapper = flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper(
self.backend.get_gpu_workspace_buffer(
key_name=self.backend.workspace_buffer_key,
Expand All @@ -84,7 +103,7 @@ def init_state(self):
self.backend.tp_q_head_num,
self.backend.tp_kv_head_num,
self.backend.head_dim,
1,
self.backend.page_size,
causal=True,
pos_encoding_mode="NONE",
logits_soft_cap=0.0,
Expand Down Expand Up @@ -119,7 +138,10 @@ def _nomarl_prefill_att(
o_tensor = alloc_func(q.shape, q.dtype, device="cuda")
self.prefill_wrapper.run(
q,
(k.unsqueeze(1), v.unsqueeze(1)),
(
k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]),
v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]),
),
out=o_tensor,
)
return o_tensor
Expand All @@ -141,30 +163,42 @@ def init_state(self):
self.backend: FlashInferAttBackend = self.backend
device = self.infer_state.input_ids.device
model = self.backend.model
self.kv_last_page_len_buffer = torch.full((self.infer_state.batch_size,), 1, dtype=torch.int32, device=device)
b_page_len = triton.cdiv(self.infer_state.b_seq_len, self.backend.page_size)
self.kv_last_page_len_buffer = self.infer_state.b_seq_len - (b_page_len - 1) * self.backend.page_size
buffer_len = self.infer_state.batch_size * triton.cdiv(self.backend.max_seq_length, self.backend.page_size)
if (
self.infer_state.batch_size <= model.graph_max_batch_size
and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch
):
self.kv_indices = self.backend.kv_indices_buffer[self.infer_state.microbatch_index][
: self.infer_state.batch_size * self.backend.max_seq_length
]
self.kv_indices = self.backend.kv_indices_buffer[self.infer_state.microbatch_index][:buffer_len]
else:
self.kv_indices = torch.empty(
self.infer_state.batch_size * self.backend.max_seq_length,
buffer_len,
dtype=torch.int32,
device=device,
)

repack_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
self.infer_state.b_seq_len,
self.infer_state.b_kv_start_loc,
self.infer_state.max_kv_seq_len,
self.kv_indices,
)
self.kv_starts = self.infer_state.b1_cu_kv_seq_len.int()
self.kv_starts = self.infer_state.b1_cu_kv_seq_len.int().clone()
self.kv_starts[1:] = b_page_len.cumsum(0)
if self.backend.page_size == 1:
repack_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
self.infer_state.b_seq_len,
self.kv_starts[:-1],
self.infer_state.max_kv_seq_len,
self.kv_indices,
)
else:
repack_page_kv_index(
self.infer_state.req_manager.req_to_token_indexs,
self.infer_state.b_req_idx,
b_page_len,
self.kv_starts[:-1],
triton.cdiv(self.infer_state.max_kv_seq_len, self.backend.page_size),
self.kv_indices,
self.backend.page_size,
)
if not self._should_init_decode_wrapper():
# 处于 graph replay 回放阶段,不需要特殊初始化 decode wrapper。
return
Expand All @@ -189,7 +223,7 @@ def init_state(self):
self.backend.tp_q_head_num,
self.backend.tp_kv_head_num,
self.backend.head_dim,
1,
self.backend.page_size,
q_data_type=self.backend.q_data_type,
kv_data_type=self.backend.kv_data_type,
non_blocking=True,
Expand All @@ -210,7 +244,7 @@ def _refresh_cuda_graph_decode_plan(self, max_kv_len: int):
dtype=torch.int32,
device="cpu",
)
* max_kv_len
* triton.cdiv(max_kv_len, self.backend.page_size)
)

fast_decode_plan(
Expand All @@ -221,7 +255,7 @@ def _refresh_cuda_graph_decode_plan(self, max_kv_len: int):
num_qo_heads=self.backend.tp_q_head_num,
num_kv_heads=self.backend.tp_kv_head_num,
head_dim=self.backend.head_dim,
page_size=1,
page_size=self.backend.page_size,
q_data_type=self.backend.q_data_type,
kv_data_type=self.backend.kv_data_type,
non_blocking=True,
Expand Down Expand Up @@ -259,7 +293,10 @@ def _normal_decode_att(
o_tensor = alloc_func(q.shape, q.dtype)
self.decode_wrapper.run(
q,
(k.unsqueeze(1), v.unsqueeze(1)),
(
k.view(-1, self.backend.page_size, k.shape[1], k.shape[2]),
v.view(-1, self.backend.page_size, v.shape[1], v.shape[2]),
),
out=o_tensor,
)
return o_tensor
Loading
Loading