[FlashAttn] Add fused triton kernel for normal_decode_set_metadata (#20778)
Co-authored-by: kinza99 <dh18324568312@163.com>
This commit is contained in:
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user