[FlashAttn] Add fused triton kernel for normal_decode_set_metadata (#20778)

Co-authored-by: kinza99 <dh18324568312@163.com>
This commit is contained in:
Bowen Li
2026-03-22 15:12:29 +08:00
committed by GitHub
parent f7fc2c8592
commit 3bc595acbc
2 changed files with 705 additions and 14 deletions

View File

@@ -2663,9 +2663,167 @@ class FlashAttentionMultiStepBackend:
)
# @torch.compile(dynamic=True, backend=get_compiler_backend())
# TODO: fuse these kernels
# NOTE: torch.compile makes it slower in speculative decoding
@triton.jit
def _fused_metadata_kernel_general(
# Input tensors
seq_lens,
seq_lens_stride_0,
req_to_token,
req_to_token_stride_0,
req_to_token_stride_1,
req_pool_indices,
req_pool_indices_stride_0,
# Output buffers
cache_seqlens_int32,
cache_seqlens_int32_stride_0,
cu_seqlens_k,
cu_seqlens_k_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
swa_page_table,
swa_page_table_stride_0,
swa_page_table_stride_1,
full_to_swa_mapping,
full_to_swa_mapping_stride_0,
# Scalar parameters
B,
max_seq_pages,
page_size: tl.constexpr,
seq_len_delta: tl.constexpr,
use_swa: tl.constexpr,
SHIFT: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
pid_b = tl.program_id(0) # batch index
pid_c = tl.program_id(1) # column chunk index
# 1. Prefix sum (only one block does it)
if pid_b == 0 and pid_c == 0:
acc = 0
for idx in range(B):
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
val = (seq + seq_len_delta).to(tl.int32)
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
acc += val
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
# 2. Gather for this batch and column chunk
if max_seq_pages == 0:
return
i = pid_b
# Load row index for this batch (all threads in block have same i)
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
row_offset = row_idx * req_to_token_stride_0
col_start = pid_c * BLOCK_COLS
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
mask = col_offsets < max_seq_pages
# Compute column indices in the source tensor (token offset)
if page_size == 1:
col_idx = col_offsets
else:
col_idx = col_offsets << SHIFT # faster than multiplication for power-of-two
# Load page indices from req_to_token
rt_offsets = row_offset + col_idx * req_to_token_stride_1
page_index = tl.load(
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
)
# Compute page_table
if page_size == 1:
page_table_val = page_index
else:
page_table_val = page_index >> SHIFT
# Store to page_table
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
tl.store(page_table + pt_offsets, page_table_val, mask=mask, cache_modifier=".cg")
if use_swa:
swa_slot = tl.load(
full_to_swa_mapping + page_index * full_to_swa_mapping_stride_0,
mask=mask,
other=0,
cache_modifier=".cg",
)
if page_size == 1:
swa_val = swa_slot
else:
swa_val = swa_slot >> SHIFT
swa_offsets = (
i * swa_page_table_stride_0 + col_offsets * swa_page_table_stride_1
)
tl.store(swa_page_table + swa_offsets, swa_val, mask=mask, cache_modifier=".cg")
@triton.jit
def _fused_metadata_kernel_ps1_no_swa(
# Input tensors
seq_lens,
seq_lens_stride_0,
req_to_token,
req_to_token_stride_0,
req_to_token_stride_1,
req_pool_indices,
req_pool_indices_stride_0,
# Output buffers
cache_seqlens_int32,
cache_seqlens_int32_stride_0,
cu_seqlens_k,
cu_seqlens_k_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
# Scalar parameters
B,
max_seq_pages,
seq_len_delta: tl.constexpr,
BLOCK_COLS: tl.constexpr,
):
pid_b = tl.program_id(0) # batch index
pid_c = tl.program_id(1) # column chunk index
# 1. Prefix sum (only one block does it)
if pid_b == 0 and pid_c == 0:
acc = 0
for idx in range(B):
seq = tl.load(seq_lens + idx * seq_lens_stride_0)
val = (seq + seq_len_delta).to(tl.int32)
tl.store(cache_seqlens_int32 + idx * cache_seqlens_int32_stride_0, val)
tl.store(cu_seqlens_k + idx * cu_seqlens_k_stride_0, acc)
acc += val
tl.store(cu_seqlens_k + B * cu_seqlens_k_stride_0, acc)
# 2. Gather for this batch and column chunk
if max_seq_pages == 0:
return
i = pid_b
# Load row index for this batch (all threads in block have same i)
row_idx = tl.load(req_pool_indices + i * req_pool_indices_stride_0)
row_offset = row_idx * req_to_token_stride_0
col_start = pid_c * BLOCK_COLS
col_offsets = col_start + tl.arange(0, BLOCK_COLS)
mask = col_offsets < max_seq_pages
# page_size = 1: col_idx = col_offsets
rt_offsets = row_offset + col_offsets * req_to_token_stride_1
page_index = tl.load(
req_to_token + rt_offsets, mask=mask, other=0, cache_modifier=".cg"
)
# page_table = page_index // 1 = page_index
pt_offsets = i * page_table_stride_0 + col_offsets * page_table_stride_1
tl.store(page_table + pt_offsets, page_index, mask=mask, cache_modifier=".cg")
# Fused Triton kernel implementation
def normal_decode_set_metadata(
cache_seqlens_int32: torch.Tensor,
cu_seqlens_k: torch.Tensor,
@@ -2680,18 +2838,133 @@ def normal_decode_set_metadata(
swa_page_table: Optional[torch.Tensor] = None,
token_to_kv_pool: Optional[SWAKVPool] = None,
):
cache_seqlens_int32.copy_(seq_lens + seq_len_delta)
cu_seqlens_k[1:].copy_(torch.cumsum(cache_seqlens_int32, dim=0, dtype=torch.int32))
page_indices = req_to_token[
req_pool_indices[:, None],
strided_indices[:max_seq_pages][None, :],
]
page_table[:, :max_seq_pages].copy_(page_indices // page_size)
"""
Fused Triton implementation that replaces 4-5 sequential CUDA kernels with 1-2 kernels:
1. cache_seqlens = seq_lens + seq_len_delta (int64→int32 cast)
2. cu_seqlens_k = cumsum(cache_seqlens) (prefix-sum)
3. page_indices = req_to_token[pool_idx, stride_idx] (2-D gather)
4. page_table = page_indices // page_size (floor-divide)
5. (optional) swa_page_table for sliding window attention
if swa_page_table is not None and token_to_kv_pool is not None:
assert isinstance(token_to_kv_pool, SWAKVPool)
swa_page_indices = token_to_kv_pool.translate_loc_from_full_to_swa(page_indices)
swa_page_table[:, :max_seq_pages].copy_(swa_page_indices // page_size)
Achieves ~5.2x speedup on H200 hardware for typical decode workloads.
"""
assert (
page_size > 0 and (page_size & (page_size - 1)) == 0
), f"page_size must be a power of two, got {page_size}"
batch_size = cache_seqlens_int32.shape[0]
device = seq_lens.device
# Ensure contiguous memory layout for efficient Triton access
seq_lens = seq_lens.contiguous()
req_to_token = req_to_token.contiguous()
req_pool_indices = req_pool_indices.contiguous()
# Prepare tensor strides
seq_lens_stride_0 = seq_lens.stride(0)
req_to_token_stride_0 = req_to_token.stride(0)
req_to_token_stride_1 = req_to_token.stride(1)
req_pool_indices_stride_0 = req_pool_indices.stride(0)
cache_seqlens_int32_stride_0 = cache_seqlens_int32.stride(0)
cu_seqlens_k_stride_0 = cu_seqlens_k.stride(0)
page_table_stride_0 = page_table.stride(0)
page_table_stride_1 = page_table.stride(1)
# Check if we should use the specialized fast path for page_size=1, no SWA
use_swa = swa_page_table is not None and token_to_kv_pool is not None
if page_size == 1 and not use_swa:
# Specialized kernel for the common case (page_size=1, no SWA)
BLOCK_COLS = 256
if max_seq_pages == 0:
grid = (1, 1)
else:
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
grid = (batch_size, num_blocks_j)
_fused_metadata_kernel_ps1_no_swa[grid](
seq_lens,
seq_lens_stride_0,
req_to_token,
req_to_token_stride_0,
req_to_token_stride_1,
req_pool_indices,
req_pool_indices_stride_0,
cache_seqlens_int32,
cache_seqlens_int32_stride_0,
cu_seqlens_k,
cu_seqlens_k_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
batch_size,
max_seq_pages,
seq_len_delta,
BLOCK_COLS=BLOCK_COLS,
num_warps=8,
num_stages=3,
)
else:
# General kernel for page_size > 1 or SWA cases
# SWA parameters
if use_swa:
assert isinstance(token_to_kv_pool, SWAKVPool)
swa_page_table = swa_page_table.contiguous()
swa_page_table_stride_0 = swa_page_table.stride(0)
swa_page_table_stride_1 = swa_page_table.stride(1)
# Extract the full_to_swa_index_mapping from token_to_kv_pool
full_to_swa_mapping = (
token_to_kv_pool.full_to_swa_index_mapping.contiguous()
)
full_to_swa_mapping_stride_0 = full_to_swa_mapping.stride(0)
else:
# Dummy tensors (not used)
swa_page_table = torch.empty(0, dtype=torch.int32, device=device)
swa_page_table_stride_0 = 0
swa_page_table_stride_1 = 0
full_to_swa_mapping = torch.empty(0, dtype=torch.int32, device=device)
full_to_swa_mapping_stride_0 = 0
# Kernel configuration
BLOCK_COLS = 128
shift = (page_size).bit_length() - 1 if page_size > 1 else 0
if max_seq_pages == 0:
grid = (1, 1)
else:
num_blocks_j = triton.cdiv(max_seq_pages, BLOCK_COLS)
grid = (batch_size, num_blocks_j)
_fused_metadata_kernel_general[grid](
seq_lens,
seq_lens_stride_0,
req_to_token,
req_to_token_stride_0,
req_to_token_stride_1,
req_pool_indices,
req_pool_indices_stride_0,
cache_seqlens_int32,
cache_seqlens_int32_stride_0,
cu_seqlens_k,
cu_seqlens_k_stride_0,
page_table,
page_table_stride_0,
page_table_stride_1,
swa_page_table,
swa_page_table_stride_0,
swa_page_table_stride_1,
full_to_swa_mapping,
full_to_swa_mapping_stride_0,
batch_size,
max_seq_pages,
page_size,
seq_len_delta,
use_swa,
shift,
BLOCK_COLS=BLOCK_COLS,
num_warps=4,
num_stages=3,
)
@torch.compile(dynamic=True, backend=get_compiler_backend())