[perf] parallelize create_flashmla_kv_indices over page-blocks (#27320)
This commit is contained in:
@@ -20,6 +20,7 @@ from sglang.srt.layers.attention.triton_ops.aiter_unified_attention import (
|
|||||||
from sglang.srt.layers.attention.utils import (
|
from sglang.srt.layers.attention.utils import (
|
||||||
create_flashinfer_kv_indices_triton,
|
create_flashinfer_kv_indices_triton,
|
||||||
create_flashmla_kv_indices_triton,
|
create_flashmla_kv_indices_triton,
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_tp_size,
|
get_attention_tp_size,
|
||||||
@@ -887,7 +888,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
bs, max_kv_len, dtype=torch.int32, device=self.device
|
bs, max_kv_len, dtype=torch.int32, device=self.device
|
||||||
)
|
)
|
||||||
|
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(bs, get_num_kv_index_blocks_flashmla(max_kv_len, 1))
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
|
|||||||
@@ -12,7 +12,10 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
|
|
||||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashmla_kv_indices_triton
|
from sglang.srt.layers.attention.utils import (
|
||||||
|
create_flashmla_kv_indices_triton,
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
@@ -88,7 +91,14 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
|||||||
spec_info = forward_batch.spec_info
|
spec_info = forward_batch.spec_info
|
||||||
|
|
||||||
if forward_mode.is_decode_or_idle() and spec_info is None:
|
if forward_mode.is_decode_or_idle() and spec_info is None:
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(
|
||||||
|
bs,
|
||||||
|
get_num_kv_index_blocks_flashmla(
|
||||||
|
self.cuda_graph_kv_indices.stride(0), PAGE_SIZE
|
||||||
|
),
|
||||||
|
)
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices[:bs],
|
forward_batch.req_pool_indices[:bs],
|
||||||
forward_batch.seq_lens[:bs],
|
forward_batch.seq_lens[:bs],
|
||||||
@@ -124,7 +134,9 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=forward_batch.seq_lens.device,
|
device=forward_batch.seq_lens.device,
|
||||||
)
|
)
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
|
|||||||
@@ -13,7 +13,10 @@ import triton
|
|||||||
from sgl_kernel.flash_mla import flash_mla_with_kvcache, get_mla_metadata
|
from sgl_kernel.flash_mla import flash_mla_with_kvcache, get_mla_metadata
|
||||||
|
|
||||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashmla_kv_indices_triton
|
from sglang.srt.layers.attention.utils import (
|
||||||
|
create_flashmla_kv_indices_triton,
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
|
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
@@ -116,7 +119,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=forward_batch.seq_lens.device,
|
device=forward_batch.seq_lens.device,
|
||||||
)
|
)
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
@@ -147,7 +152,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=seq_lens.device,
|
device=seq_lens.device,
|
||||||
)
|
)
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
forward_batch.req_pool_indices,
|
forward_batch.req_pool_indices,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
@@ -233,7 +240,14 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
)
|
)
|
||||||
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
|
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
|
||||||
|
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(
|
||||||
|
bs,
|
||||||
|
get_num_kv_index_blocks_flashmla(
|
||||||
|
self.cuda_graph_kv_indices.stride(0), PAGE_SIZE
|
||||||
|
),
|
||||||
|
)
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices[:bs],
|
req_pool_indices[:bs],
|
||||||
seq_lens,
|
seq_lens,
|
||||||
|
|||||||
@@ -49,6 +49,15 @@ def get_num_page_per_block_flashmla(page_size: int = 64) -> int:
|
|||||||
return num_page_per_block
|
return num_page_per_block
|
||||||
|
|
||||||
|
|
||||||
|
def get_num_kv_index_blocks_flashmla(kv_indices_width: int, page_size: int) -> int:
|
||||||
|
"""Grid axis-1 size for create_flashmla_kv_indices_triton: the number of
|
||||||
|
page-blocks spanning the widest sequence (one CTA per block). kv_indices_width
|
||||||
|
is the per-row width of the kv_indices buffer (the kernel's kv_indices_ptr_stride).
|
||||||
|
"""
|
||||||
|
npb = get_num_page_per_block_flashmla(page_size)
|
||||||
|
return (kv_indices_width + npb - 1) // npb
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def create_flashmla_kv_indices_triton(
|
def create_flashmla_kv_indices_triton(
|
||||||
req_to_token_ptr, # [max_batch, max_context_len]
|
req_to_token_ptr, # [max_batch, max_context_len]
|
||||||
@@ -79,7 +88,10 @@ def create_flashmla_kv_indices_triton(
|
|||||||
num_paged = tl.cdiv(kv_end - kv_start, PAGED_SIZE)
|
num_paged = tl.cdiv(kv_end - kv_start, PAGED_SIZE)
|
||||||
num_pages_loop = tl.cdiv(kv_end - kv_start, FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON)
|
num_pages_loop = tl.cdiv(kv_end - kv_start, FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON)
|
||||||
|
|
||||||
for i in range(num_pages_loop):
|
# One CTA per page-block (grid axis 1) rather than one CTA looping all blocks;
|
||||||
|
# CTAs beyond this sequence's block count are guarded out.
|
||||||
|
i = tl.program_id(axis=1)
|
||||||
|
if i < num_pages_loop:
|
||||||
# index into req_to_token_ptr needs to be int64
|
# index into req_to_token_ptr needs to be int64
|
||||||
paged_offset = (
|
paged_offset = (
|
||||||
tl.arange(0, NUM_PAGE_PER_BLOCK).to(tl.int64) + i * NUM_PAGE_PER_BLOCK
|
tl.arange(0, NUM_PAGE_PER_BLOCK).to(tl.int64) + i * NUM_PAGE_PER_BLOCK
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
|
|||||||
from sglang.srt.layers.attention.utils import (
|
from sglang.srt.layers.attention.utils import (
|
||||||
concat_mla_absorb_q_general,
|
concat_mla_absorb_q_general,
|
||||||
create_flashmla_kv_indices_triton,
|
create_flashmla_kv_indices_triton,
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
get_num_page_per_block_flashmla,
|
get_num_page_per_block_flashmla,
|
||||||
mla_quantize_and_rope_for_fp8,
|
mla_quantize_and_rope_for_fp8,
|
||||||
)
|
)
|
||||||
@@ -388,7 +389,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
|
|
||||||
create_flashmla_kv_indices_triton[(batch_size,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(
|
||||||
|
batch_size,
|
||||||
|
get_num_kv_index_blocks_flashmla(max_blocks, self.page_size),
|
||||||
|
)
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
@@ -534,7 +540,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
metadata.seq_lens_k.copy_(seq_lens.to(torch.int32))
|
metadata.seq_lens_k.copy_(seq_lens.to(torch.int32))
|
||||||
|
|
||||||
# Update block indices for new sequences.
|
# Update block indices for new sequences.
|
||||||
create_flashmla_kv_indices_triton[(bs,)](
|
create_flashmla_kv_indices_triton[
|
||||||
|
(
|
||||||
|
bs,
|
||||||
|
get_num_kv_index_blocks_flashmla(
|
||||||
|
metadata.block_kv_indices.shape[1], self.page_size
|
||||||
|
),
|
||||||
|
)
|
||||||
|
](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices[:bs],
|
req_pool_indices[:bs],
|
||||||
seq_lens,
|
seq_lens,
|
||||||
|
|||||||
@@ -19,6 +19,9 @@ from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
|||||||
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
||||||
create_flashmla_kv_indices_triton as create_flashmla_kv_indices_triton,
|
create_flashmla_kv_indices_triton as create_flashmla_kv_indices_triton,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
||||||
|
get_num_kv_index_blocks_flashmla as get_num_kv_index_blocks_flashmla,
|
||||||
|
)
|
||||||
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
||||||
get_num_page_per_block_flashmla as get_num_page_per_block_flashmla,
|
get_num_page_per_block_flashmla as get_num_page_per_block_flashmla,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user