This commit is contained in:
Yineng Zhang
2025-08-25 01:29:06 -07:00
committed by GitHub
parent 938e986e15
commit ebd9dbe71b
5 changed files with 103 additions and 290 deletions
@@ -24,7 +24,9 @@ if os.environ["SGLANG_ENABLE_TORCH_COMPILE"] == "1":
from sglang.global_config import global_config
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.layers.attention.flashinfer_backend import (
create_flashinfer_kv_indices_triton,
)
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.utils import is_sm100_supported
from sglang.srt.managers.schedule_batch import global_server_args_dict
@@ -179,6 +181,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
q_indptr_decode_buf: Optional[torch.Tensor] = None,
):
super().__init__()
# Parse constants
self.max_context_len = model_runner.model_config.context_len
self.device = model_runner.device
@@ -210,25 +213,15 @@ class FlashInferMLAAttnBackend(AttentionBackend):
else:
self.kv_indptr = kv_indptr_buf
self.kv_indices = torch.empty(
(max_bs * (self.max_context_len + self.page_size - 1) // self.page_size,),
dtype=torch.int32,
device=model_runner.device,
)
if not self.skip_prefill:
self.qo_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
)
if q_indptr_decode_buf is None:
# A hack to pre-initialize large batch size for dp attention
if model_runner.server_args.enable_dp_attention:
max_bs = model_runner.server_args.dp_size * max_bs
self.q_indptr_decode = torch.arange(
0, max_bs + 1, dtype=torch.int32, device=model_runner.device
)
else:
self.q_indptr_decode = q_indptr_decode_buf
@@ -273,7 +266,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.prefill_cuda_graph_metadata = {} # For verify
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update(
forward_batch.req_pool_indices,
@@ -331,9 +323,16 @@ class FlashInferMLAAttnBackend(AttentionBackend):
max_num_tokens: int,
kv_indices_buf: Optional[torch.Tensor] = None,
):
self.cuda_graph_kv_indices = (
self.kv_indices.clone() if kv_indices_buf is None else kv_indices_buf
)
if kv_indices_buf is None:
cuda_graph_kv_indices = torch.zeros(
(max_bs * self.max_context_len,),
dtype=torch.int32,
device="cuda",
)
else:
cuda_graph_kv_indices = kv_indices_buf
self.cuda_graph_kv_indices = cuda_graph_kv_indices
self.cuda_graph_qo_indptr = self.q_indptr_decode.clone()
self.cuda_graph_kv_indptr = self.kv_indptr.clone()
self.cuda_graph_kv_lens = torch.ones(
@@ -359,7 +358,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
forward_mode: ForwardMode,
spec_info: Optional[SpecInfo],
):
if forward_mode.is_decode_or_idle():
decode_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
@@ -370,6 +368,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
kv_len_arr=self.cuda_graph_kv_lens[:num_tokens],
backend="auto",
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_decode.update(
req_pool_indices,
@@ -440,13 +439,11 @@ class FlashInferMLAAttnBackend(AttentionBackend):
spec_info: Optional[SpecInfo],
seq_lens_cpu: Optional[torch.Tensor],
):
if forward_mode.is_decode_or_idle():
assert seq_lens_cpu is not None
kv_len_arr_cpu = seq_lens_cpu[:bs]
num_pages_per_req = (seq_lens_cpu + self.page_size - 1) // self.page_size
self.cuda_graph_kv_indptr_cpu[1 : bs + 1] = torch.cumsum(
num_pages_per_req, dim=0
kv_len_arr_cpu, dim=0
)
self.fast_decode_kwargs.update(
{
@@ -455,6 +452,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
"kv_len_arr_cpu": kv_len_arr_cpu,
}
)
self.indices_updater_decode.update(
req_pool_indices[:bs],
seq_lens[:bs],
@@ -534,6 +532,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
q_rope = q_rope.view(
-1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim
)
if self.forward_metadata.use_ragged:
# ragged prefill
if q_rope is not None:
@@ -554,8 +553,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
k_buf = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
q.dtype
)
k_buf = k_buf.view(-1, self.page_size, k_buf.shape[-1])
if q_rope is None:
qall = q.view(-1, layer.tp_q_head_num, layer.head_dim)
q, q_rope = (
@@ -617,17 +614,17 @@ class FlashInferMLAAttnBackend(AttentionBackend):
q_nope = reshaped_q[:, :, : layer.v_head_dim]
q_rope = reshaped_q[:, :, layer.v_head_dim :]
k_buf = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
q.dtype
)
k_buf = k_buf.view(-1, self.page_size, k_buf.shape[-1])
o = q_nope.new_empty(q_nope.shape)
# Direct call to run without the wrapper
o = decode_wrapper.run(
q_nope,
q_rope,
k_buf[:, :, : layer.v_head_dim],
k_buf[:, :, layer.v_head_dim :],
k_buffer[:, :, : layer.v_head_dim],
k_buffer[:, :, layer.v_head_dim :],
out=o,
)
@@ -646,10 +643,9 @@ class FlashInferMLAIndicesUpdaterDecode:
self.scaling = model_runner.model_config.scaling
self.data_type = model_runner.dtype
self.attn_backend = attn_backend
self.page_size = model_runner.page_size
# Buffers and wrappers
self.kv_indptr = attn_backend.kv_indptr
self.kv_indices = attn_backend.kv_indices
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.q_indptr = attn_backend.q_indptr_decode
@@ -693,17 +689,13 @@ class FlashInferMLAIndicesUpdaterDecode:
kv_lens = paged_kernel_lens.to(torch.int32)
sm_scale = self.scaling
if spec_info is None:
num_pages_per_req = (
paged_kernel_lens + self.page_size - 1
) // self.page_size
kv_indptr[1 : bs + 1] = torch.cumsum(num_pages_per_req, dim=0)
kv_indptr[1 : bs + 1] = torch.cumsum(paged_kernel_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = (
self.kv_indices[: kv_indptr[-1]]
torch.empty(paged_kernel_lens_sum, dtype=torch.int32, device="cuda")
if not init_metadata_replay
else fast_decode_kwargs["kv_indices"]
)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
@@ -712,40 +704,39 @@ class FlashInferMLAIndicesUpdaterDecode:
None,
kv_indices,
self.req_to_token.shape[1],
self.page_size,
)
else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
if not init_metadata_replay:
wrapper.plan(
qo_indptr=q_indptr,
kv_indptr=kv_indptr,
kv_indices=kv_indices,
kv_len_arr=kv_lens,
num_heads=self.num_local_heads,
head_dim_ckv=self.kv_lora_rank,
head_dim_kpe=self.qk_rope_head_dim,
page_size=self.page_size,
causal=False,
sm_scale=sm_scale,
q_data_type=self.data_type,
kv_data_type=self.data_type,
q_indptr,
kv_indptr,
kv_indices,
kv_lens,
self.num_local_heads,
self.kv_lora_rank,
self.qk_rope_head_dim,
1,
False,
sm_scale,
self.data_type,
self.data_type,
)
else:
wrapper.plan(
qo_indptr_cpu=fast_decode_kwargs["qo_indptr_cpu"],
kv_indptr_cpu=fast_decode_kwargs["kv_indptr_cpu"],
kv_indices=kv_indices,
kv_len_arr_cpu=fast_decode_kwargs["kv_len_arr_cpu"],
num_heads=self.num_local_heads,
head_dim_ckv=self.kv_lora_rank,
head_dim_kpe=self.qk_rope_head_dim,
page_size=self.page_size,
causal=False,
sm_scale=sm_scale,
q_data_type=self.data_type,
kv_data_type=self.data_type,
fast_decode_kwargs["qo_indptr_cpu"],
fast_decode_kwargs["kv_indptr_cpu"],
kv_indices,
fast_decode_kwargs["kv_len_arr_cpu"],
self.num_local_heads,
self.kv_lora_rank,
self.qk_rope_head_dim,
1,
False,
sm_scale,
self.data_type,
self.data_type,
)
@@ -767,14 +758,12 @@ class FlashInferMLAIndicesUpdaterPrefill:
# Buffers and wrappers
self.kv_indptr = attn_backend.kv_indptr
self.qo_indptr = attn_backend.qo_indptr
self.kv_indices = attn_backend.kv_indices
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.prefill_wrapper_ragged = attn_backend.prefill_wrapper_ragged
self.page_size = model_runner.page_size
def update(
self,
req_pool_indices: torch.Tensor,
req_pool_indices: torch.Tnesor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
@@ -788,6 +777,7 @@ class FlashInferMLAIndicesUpdaterPrefill:
else:
paged_kernel_lens = seq_lens
paged_kernel_lens_sum = seq_lens_sum
self.call_begin_forward(
self.prefill_wrapper_ragged,
prefill_wrapper_paged,
@@ -821,12 +811,13 @@ class FlashInferMLAIndicesUpdaterPrefill:
if spec_info is None:
assert len(seq_lens) == len(req_pool_indices)
num_pages_per_req = (
paged_kernel_lens + self.page_size - 1
) // self.page_size
kv_indptr[1 : bs + 1] = torch.cumsum(num_pages_per_req, dim=0)
kv_indptr[1 : bs + 1] = torch.cumsum(paged_kernel_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = self.kv_indices[: kv_indptr[-1]]
kv_indices = torch.empty(
paged_kernel_lens_sum,
dtype=torch.int32,
device=req_pool_indices.device,
)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
@@ -835,7 +826,6 @@ class FlashInferMLAIndicesUpdaterPrefill:
None,
kv_indices,
self.req_to_token.shape[1],
self.page_size,
)
qo_indptr[1 : bs + 1] = torch.cumsum(seq_lens - prefix_lens, dim=0)
qo_indptr = qo_indptr[: bs + 1]
@@ -853,6 +843,7 @@ class FlashInferMLAIndicesUpdaterPrefill:
self.req_to_token,
)
)
if use_ragged:
# ragged prefill
wrapper_ragged.begin_forward(
@@ -867,26 +858,20 @@ class FlashInferMLAIndicesUpdaterPrefill:
)
else:
# mla paged prefill
if spec_info is not None:
assert (
self.page_size == 1
), "Only page_size=1 is supported for flashinfer backend with speculative decoding"
kv_lens = kv_indptr[1:] - kv_indptr[:-1]
else:
kv_lens = paged_kernel_lens.to(torch.int32)
kv_len_arr = kv_indptr[1:] - kv_indptr[:-1]
wrapper_paged.plan(
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
kv_indices=kv_indices,
kv_len_arr=kv_lens,
num_heads=self.num_local_heads,
head_dim_ckv=self.kv_lora_rank,
head_dim_kpe=self.qk_rope_head_dim,
page_size=self.page_size,
causal=True,
sm_scale=sm_scale,
q_data_type=self.q_data_type,
kv_data_type=self.data_type,
qo_indptr,
kv_indptr,
kv_indices,
kv_len_arr,
self.num_local_heads,
self.kv_lora_rank,
self.qk_rope_head_dim,
1,
True,
sm_scale,
self.q_data_type,
self.data_type,
)
@@ -981,7 +966,6 @@ class FlashInferMLAMultiStepDraftBackend:
call_fn(i, forward_batch)
def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_indices = torch.zeros(
(
self.speculative_num_steps,
@@ -1017,7 +1001,6 @@ class FlashInferMLAMultiStepDraftBackend:
)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
@@ -1034,7 +1017,6 @@ class FlashInferMLAMultiStepDraftBackend:
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
def call_fn(i, forward_batch):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,