[AMD] refactor sparse MLA decode kernel for Deepseek V4 triton backend (#28265)

Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
Co-authored-by: yichiche@amd.com <jacky.cheng>
This commit is contained in:
amd-danli103
2026-06-15 03:26:59 -07:00
committed by GitHub
co-authored by Raiden-Makoto yichiche@amd.com
parent da12f36629
commit c4ec39a785
5 changed files with 212 additions and 2918 deletions
@@ -1,585 +0,0 @@
"""
Common utilities and attention kernels for Triton MLA Decode.
This module contains shared code for the DeepSeek V4 Triton decode implementation:
- Attention kernels (unified sparse decode)
- Helper functions for chunked attention
- Token range computation for memory-based chunking
"""
from typing import List, Tuple
import torch
import triton
import triton.language as tl
LOG2E = tl.constexpr(1.4426950408889634)
# ============================================================================
# Bucketing for autotune keys to avoid recompilation per unique batch size
# ============================================================================
def _bucket_total_tokens(total_tokens: int) -> int:
"""Round total_tokens up to the nearest power of 2 for autotune key stability.
In serving, total_tokens (= batch_size * seq_len) varies with every batch.
Using the exact value as an autotune key causes recompilation for each unique
value. Bucketing to powers of 2 limits the number of unique keys to ~15,
dramatically reducing autotuning overhead.
Returns:
Power-of-2 bucket: 1, 2, 4, 8, ..., up to the next power of 2.
"""
if total_tokens <= 0:
return 1
# Round up to next power of 2
n = 1
while n < total_tokens:
n <<= 1
return n
# ============================================================================
# Helper function to compute workload size category for autotune
# ============================================================================
def _get_workload_size_category(total_tokens: int, topk: int) -> int:
"""
Compute workload size category for autotune key.
Returns:
0: small (< 10K elements)
1: medium (10K - 100K elements)
2: large (100K - 1M elements)
3: very large (> 1M elements)
"""
total_elements = total_tokens * topk
if total_elements < 10000:
return 0
elif total_elements < 100000:
return 1
elif total_elements < 1000000:
return 2
else:
return 3
# ============================================================================
# Unified Attention Kernels
# ============================================================================
# ============================================================================
# CDNA4 (gfx950) Optimized: Added high-performance configs for MI355X
# Best config for h_q=128, large topk: BLOCK_H=64, BLOCK_N=256, num_warps=8
# ============================================================================
@triton.autotune(
configs=[
# Selected based on CDNA4 architecture analysis:
# - BLOCK_D=128 is fixed (matches KV tile structure for d_qk=512).
# - BLOCK_N=256: best for amortizing memory access over topk dimension.
# (decode attention is memory-bound; larger BLOCK_N = fewer iterations)
# - num_warps=8: memory-bound decode benefits from more warps for latency hiding.
# - BLOCK_H varies to cover different batch sizes:
# * BLOCK_H=16: cdiv(128,16)=8 H-blocks, best for small batches (bs=1-8)
# * BLOCK_H=32: cdiv(128,32)=4 H-blocks, good for medium batches (bs=8-32)
# * BLOCK_H=64: cdiv(128,64)=2 H-blocks, best for large batches (bs=32+)
# (original comment: "Best for h_q=128, large topk")
# * BLOCK_H=128: cdiv(128,128)=1 H-block, for very large batches (bs=128+)
triton.Config(
{"BLOCK_H": 16, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=1
),
triton.Config(
{"BLOCK_H": 32, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=1
),
triton.Config(
{"BLOCK_H": 64, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=1
),
triton.Config(
{"BLOCK_H": 128, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=1
),
],
key=["total_tokens_bucket", "h_q", "total_topk", "d_qk"],
)
@triton.jit
def _unified_sparse_decode_kernel(
Q,
KV,
Mask,
AttnSink,
Output,
LSE,
sm_scale,
total_tokens,
total_tokens_bucket,
h_q,
total_topk,
d_qk,
d_v,
stride_q_t,
stride_q_h,
stride_q_d,
stride_kv_t,
stride_kv_k,
stride_kv_d,
stride_mask_t,
stride_mask_k,
stride_o_t,
stride_o_h,
stride_o_d,
stride_lse_t,
stride_lse_h,
HAS_ATTN_SINK: tl.constexpr,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_D: tl.constexpr,
):
"""Unified attention kernel with single KV buffer (int64 safe)."""
pid_t = tl.program_id(0)
pid_h = tl.program_id(1)
pid_t_64 = pid_t.to(tl.int64)
NEG_INF = float("-inf")
POS_INF = float("+inf")
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
mask_h = offs_h < h_q
m_i = tl.full([BLOCK_H], NEG_INF, dtype=tl.float32)
l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
acc_0 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_1 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_2 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_3 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
stride_q_t_64 = tl.cast(stride_q_t, tl.int64)
stride_kv_t_64 = tl.cast(stride_kv_t, tl.int64)
stride_mask_t_64 = tl.cast(stride_mask_t, tl.int64)
q_base = Q + pid_t_64 * stride_q_t_64
kv_base = KV + pid_t_64 * stride_kv_t_64
mask_base = Mask + pid_t_64 * stride_mask_t_64
for n_start in range(0, total_topk, BLOCK_N):
offs_n = n_start + tl.arange(0, BLOCK_N)
mask_n = offs_n < total_topk
mask_ptrs = mask_base + offs_n * stride_mask_k
invalid = tl.load(mask_ptrs, mask=mask_n, other=True)
valid = mask_n & ~invalid
qk = tl.zeros([BLOCK_H, BLOCK_N], dtype=tl.float32)
for d_start in range(0, d_qk, BLOCK_D):
offs_d = d_start + tl.arange(0, BLOCK_D)
mask_d = offs_d < d_qk
q_ptrs = (
q_base + offs_h[:, None] * stride_q_h + offs_d[None, :] * stride_q_d
)
q_chunk = tl.load(
q_ptrs, mask=mask_h[:, None] & mask_d[None, :], other=0.0
).to(tl.bfloat16)
k_ptrs = (
kv_base + offs_n[:, None] * stride_kv_k + offs_d[None, :] * stride_kv_d
)
k_chunk = tl.load(
k_ptrs, mask=valid[:, None] & mask_d[None, :], other=0.0
).to(tl.bfloat16)
qk += tl.dot(q_chunk, tl.trans(k_chunk))
qk = qk * sm_scale
qk = tl.where(valid[None, :], qk, NEG_INF)
m_ij = tl.max(qk, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.where(m_i == NEG_INF, 0.0, tl.math.exp2((m_i - m_new) * LOG2E))
p = tl.where(qk == NEG_INF, 0.0, tl.math.exp2((qk - m_new[:, None]) * LOG2E))
l_new = alpha * l_i + tl.sum(p, axis=1)
p_bf16 = p.to(tl.bfloat16)
offs_v = tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(v_ptrs, mask=valid[:, None], other=0.0).to(tl.bfloat16)
acc_0 = acc_0 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_1 = acc_1 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = 2 * BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_2 = acc_2 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = 3 * BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_3 = acc_3 * alpha[:, None] + tl.dot(p_bf16, v)
m_i = m_new
l_i = l_new
lse = m_i + tl.math.log2(tl.where(l_i == 0.0, 1.0, l_i)) / LOG2E
is_lonely_q = l_i == 0.0
if HAS_ATTN_SINK:
attn_sink_vals = tl.load(AttnSink + offs_h, mask=mask_h, other=0.0)
exp_attn_sink_minus_m = tl.math.exp2((attn_sink_vals - m_i) * LOG2E)
denominator = l_i + exp_attn_sink_minus_m
denominator = tl.where(denominator == 0.0, 1.0, denominator)
output_scale = 1.0 / denominator
else:
output_scale = tl.where(l_i == 0.0, 0.0, 1.0 / l_i)
# Pre-compute 2D versions for efficiency
is_lonely_q_2d = is_lonely_q[:, None]
output_scale_2d = output_scale[:, None]
acc_0 = tl.where(is_lonely_q_2d, 0.0, acc_0 * output_scale_2d)
acc_1 = tl.where(is_lonely_q_2d, 0.0, acc_1 * output_scale_2d)
acc_2 = tl.where(is_lonely_q_2d, 0.0, acc_2 * output_scale_2d)
acc_3 = tl.where(is_lonely_q_2d, 0.0, acc_3 * output_scale_2d)
lse = tl.where(is_lonely_q, POS_INF, lse)
stride_lse_t_64 = tl.cast(stride_lse_t, tl.int64)
tl.store(LSE + pid_t_64 * stride_lse_t_64 + offs_h * stride_lse_h, lse, mask=mask_h)
stride_o_t_64 = tl.cast(stride_o_t, tl.int64)
o_base = Output + pid_t_64 * stride_o_t_64
# Pre-compute 2D versions
offs_h_2d = offs_h[:, None]
mask_h_2d = mask_h[:, None]
offs_v_0 = tl.arange(0, BLOCK_D)
offs_v_1 = BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_2 = 2 * BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_3 = 3 * BLOCK_D + tl.arange(0, BLOCK_D)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_0[None, :] * stride_o_d,
acc_0.to(tl.bfloat16),
mask=mask_h_2d,
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_1[None, :] * stride_o_d,
acc_1.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_1[None, :] < d_v),
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_2[None, :] * stride_o_d,
acc_2.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_2[None, :] < d_v),
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_3[None, :] * stride_o_d,
acc_3.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_3[None, :] < d_v),
)
# ============================================================================
# Attention Runner Functions
# ============================================================================
def run_unified_attention(
q_reshaped,
gathered_kv,
invalid_mask,
d_v,
sm_scale,
total_tokens,
h_q,
total_topk,
d_qk,
attn_sink=None,
):
"""Run unified attention with single KV buffer.
Run unified sparse decode attention kernel.
"""
output = torch.empty(
(total_tokens, h_q, d_v), dtype=torch.bfloat16, device=q_reshaped.device
)
lse = torch.empty(
(total_tokens, h_q), dtype=torch.float32, device=q_reshaped.device
)
HAS_ATTN_SINK = attn_sink is not None
attn_sink_tensor = attn_sink if HAS_ATTN_SINK else lse[:1]
grid = lambda meta: (total_tokens, triton.cdiv(h_q, meta["BLOCK_H"]))
_unified_sparse_decode_kernel[grid](
q_reshaped,
gathered_kv,
invalid_mask,
attn_sink_tensor,
output,
lse,
sm_scale,
total_tokens,
_bucket_total_tokens(total_tokens),
h_q,
total_topk,
d_qk,
d_v,
q_reshaped.stride(0),
q_reshaped.stride(1),
q_reshaped.stride(2),
gathered_kv.stride(0),
gathered_kv.stride(1),
gathered_kv.stride(2),
invalid_mask.stride(0),
invalid_mask.stride(1),
output.stride(0),
output.stride(1),
output.stride(2),
lse.stride(0),
lse.stride(1),
HAS_ATTN_SINK=HAS_ATTN_SINK,
)
return output, lse
def run_chunked_attention_triton(
q_reshaped,
gathered_kv,
invalid_mask,
d_v,
sm_scale,
total_tokens,
h_q,
total_topk,
d_qk,
attn_sink=None,
chunk_size=8192,
):
"""Chunked attention using Triton kernels with cross-chunk softmax merging."""
device = q_reshaped.device
num_chunks = (total_topk + chunk_size - 1) // chunk_size
kv_chunks = []
mask_chunks = []
chunk_sizes = []
for chunk_idx in range(num_chunks):
start_k = chunk_idx * chunk_size
end_k = min(start_k + chunk_size, total_topk)
chunk_topk = end_k - start_k
chunk_sizes.append(chunk_topk)
kv_chunks.append(gathered_kv[:, start_k:end_k, :].contiguous())
mask_chunks.append(invalid_mask[:, start_k:end_k].contiguous())
lse_acc = torch.full(
(total_tokens, h_q), float("-inf"), dtype=torch.float32, device=device
)
acc = torch.zeros((total_tokens, h_q, d_v), dtype=torch.float32, device=device)
for chunk_idx in range(num_chunks):
kv_chunk = kv_chunks[chunk_idx]
mask_chunk = mask_chunks[chunk_idx]
chunk_topk = chunk_sizes[chunk_idx]
chunk_output, chunk_lse = run_unified_attention(
q_reshaped,
kv_chunk,
mask_chunk,
d_v,
sm_scale,
total_tokens,
h_q,
chunk_topk,
d_qk,
attn_sink=None,
)
is_chunk_lonely = torch.isinf(chunk_lse) & (chunk_lse > 0)
chunk_lse_for_merge = torch.where(
is_chunk_lonely, torch.full_like(chunk_lse, float("-inf")), chunk_lse
)
lse_max = torch.maximum(lse_acc, chunk_lse_for_merge)
exp_acc = torch.exp(lse_acc - lse_max)
exp_acc = torch.where(torch.isnan(exp_acc), torch.zeros_like(exp_acc), exp_acc)
exp_chunk = torch.exp(chunk_lse_for_merge - lse_max)
exp_chunk = torch.where(
torch.isnan(exp_chunk) | is_chunk_lonely,
torch.zeros_like(exp_chunk),
exp_chunk,
)
sum_exp = exp_acc + exp_chunk
lse_new = lse_max + torch.log(
torch.where(sum_exp == 0, torch.ones_like(sum_exp), sum_exp)
)
both_empty = (lse_acc == float("-inf")) & (chunk_lse_for_merge == float("-inf"))
lse_new = torch.where(
both_empty, torch.full_like(lse_new, float("-inf")), lse_new
)
weight_acc = torch.exp(lse_acc - lse_new)
weight_acc = torch.where(
torch.isnan(weight_acc) | torch.isinf(weight_acc),
torch.zeros_like(weight_acc),
weight_acc,
)
weight_chunk = torch.exp(chunk_lse_for_merge - lse_new)
weight_chunk = torch.where(
torch.isnan(weight_chunk) | torch.isinf(weight_chunk) | is_chunk_lonely,
torch.zeros_like(weight_chunk),
weight_chunk,
)
acc = (
weight_acc.unsqueeze(-1) * acc
+ weight_chunk.unsqueeze(-1) * chunk_output.float()
)
lse_acc = lse_new
output = acc
lse = lse_acc
is_lonely_final = lse == float("-inf")
lse = torch.where(is_lonely_final, torch.full_like(lse, float("+inf")), lse)
if attn_sink is not None:
attn_sink_expanded = attn_sink.view(1, h_q)
exp_diff = torch.exp(attn_sink_expanded - lse)
exp_diff = torch.where(
is_lonely_final, torch.full_like(exp_diff, float("inf")), exp_diff
)
scale = 1.0 / (1.0 + exp_diff)
output = output * scale.unsqueeze(-1)
output = torch.where(
is_lonely_final.unsqueeze(-1), torch.zeros_like(output), output
)
return output.to(torch.bfloat16), lse
# ============================================================================
# Helper class and functions for token-range based chunking
# ============================================================================
class SlicedKVScope:
"""A sliced view of KV scope for a specific token range."""
__slots__ = [
"blocked_k",
"blocked_k_quantized",
"indices_in_kvcache",
"topk_length",
]
def __init__(self, blocked_k, blocked_k_quantized, indices_in_kvcache, topk_length):
self.blocked_k = blocked_k
self.blocked_k_quantized = blocked_k_quantized
self.indices_in_kvcache = indices_in_kvcache
self.topk_length = topk_length
def slice_kv_scope_for_tokens(orig_scope, start_t: int, end_t: int, s_q: int):
"""Slice a KV scope to only include tokens in range [start_t, end_t)."""
if orig_scope is None:
return None
orig_indices = orig_scope.indices_in_kvcache.reshape(
-1, orig_scope.indices_in_kvcache.size(-1)
)
sliced_indices = orig_indices[start_t:end_t]
sliced_topk_length = None
if orig_scope.topk_length is not None:
batch_start = start_t // s_q
batch_end = (end_t + s_q - 1) // s_q
batch_topk_length = orig_scope.topk_length[batch_start:batch_end]
if s_q > 1:
chunk_tokens = end_t - start_t
expanded = batch_topk_length.unsqueeze(1).expand(-1, s_q).reshape(-1)
offset_in_first_batch = start_t % s_q
sliced_topk_length = expanded[
offset_in_first_batch : offset_in_first_batch + chunk_tokens
]
else:
sliced_topk_length = batch_topk_length
return SlicedKVScope(
blocked_k=orig_scope.blocked_k,
blocked_k_quantized=orig_scope.blocked_k_quantized,
indices_in_kvcache=sliced_indices,
topk_length=sliced_topk_length,
)
def compute_token_ranges(
total_tokens: int,
total_topk: int,
d_qk: int,
max_buffer_bytes: int = 2 * 1024 * 1024 * 1024,
) -> List[Tuple[int, int]]:
"""Compute token ranges for processing, chunking if buffer would exceed limit."""
buffer_size_bytes = total_tokens * total_topk * d_qk * 2
if buffer_size_bytes <= max_buffer_bytes:
return [(0, total_tokens)]
max_tokens_per_chunk = max_buffer_bytes // (total_topk * d_qk * 2)
chunk_size = max(1, max_tokens_per_chunk)
token_ranges = []
start_t = 0
while start_t < total_tokens:
end_t = min(start_t + chunk_size, total_tokens)
token_ranges.append((start_t, end_t))
start_t = end_t
return token_ranges
# ============================================================================
# Split-K Attention for Large TopK
# ============================================================================
def run_splitk_unified_attention(
q_reshaped,
gathered_kv,
invalid_mask,
d_v,
sm_scale,
total_tokens,
h_q,
total_topk,
d_qk,
attn_sink=None,
split_k=4,
):
"""Run split-K attention for large topk cases."""
from .triton_mla_kernels_decode_splitk import run_splitk_attention
return run_splitk_attention(
q_reshaped,
gathered_kv,
invalid_mask,
d_v,
sm_scale,
total_tokens,
h_q,
total_topk,
d_qk,
attn_sink=attn_sink,
split_k=split_k,
)
@@ -6,17 +6,9 @@ This module implements a fused kernel that combines:
2. Dequant: FP8 to BF16 dequantization 2. Dequant: FP8 to BF16 dequantization
3. Attention: Compute attention scores and output 3. Attention: Compute attention scores and output
Benefits for workloads without extra scope:
- Eliminates intermediate buffer (gathered_kv) write/read
- Reduces kernel launch overhead (1 kernel instead of 2)
- Better cache utilization
Supports: Supports:
- DSV4 (d_qk=512): 7 tiles of 64, uint8 scales - DSV4 (d_qk=512): 7 tiles of 64, uint8 scales
- All configs: with/without topk_length, with/without attn_sink - All configs: with/without topk_length, with/without attn_sink
OPTIMIZED VERSION: Reduced code duplication in dual-scope kernel by using
a helper function for KV block processing.
""" """
from typing import Optional, Tuple from typing import Optional, Tuple
@@ -25,7 +17,16 @@ import torch
import triton import triton
import triton.language as tl import triton.language as tl
from .triton_mla_kernels_decode_common import _bucket_total_tokens
def _bucket_total_tokens(total_tokens: int) -> int:
"""Round total_tokens up to the nearest power of 2 for autotune key stability."""
if total_tokens <= 0:
return 1
n = 1
while n < total_tokens:
n <<= 1
return n
# ============================================================================ # ============================================================================
# Constants for DSV4 layout # Constants for DSV4 layout
@@ -35,15 +36,93 @@ DSV4_D_NOPE = 448
DSV4_D_ROPE = 64 DSV4_D_ROPE = 64
DSV4_D_V = 512 DSV4_D_V = 512
DSV4_TILE_SIZE = 64 DSV4_TILE_SIZE = 64
# ============================================================================
# Dispatch thresholds for split-K decision
# ============================================================================
# Dual-scope topk threshold.
# Split-K is more beneficial for larger topk due to more work per token.
DUAL_SCOPE_SPLITK_TOPK_THRESHOLD = 2048
# Token thresholds for split-K vs no-splitk decision:
# - Below these thresholds, split-K provides better GPU utilization.
# - Above these thresholds, the combine kernel overhead dominates.
NOSPLITK_TOKEN_THRESHOLD_LOW_TOPK = (
64 # For total_topk < DUAL_SCOPE_SPLITK_TOPK_THRESHOLD
)
# Small batch threshold: below this, split-K=4/8 for parallelism
SMALL_BATCH_TOKEN_THRESHOLD = 8
# Topk threshold for split-K value selection within the split-K path
SPLITK_HIGH_TOPK_THRESHOLD = 512
# ============================================================================
# Shared split-K decision logic for dual-scope kernels
# ============================================================================
def _decide_splitk_dual_scope(total_tokens: int, h_q: int, total_topk: int) -> int:
"""Decide the split_k value for dual-scope attention.
Returns:
split_k value (0 means no split-K, use non-splitk kernel).
"""
# Conditions under which split-K is beneficial:
use_splitk_for_small_bs = total_tokens <= SMALL_BATCH_TOKEN_THRESHOLD and (
h_q >= 128 or total_topk >= 1024
)
use_splitk_for_h64_large_topk = (
h_q <= 64
and total_topk >= 1024
and total_tokens > SMALL_BATCH_TOKEN_THRESHOLD
and total_tokens <= 128
)
use_splitk_for_large_topk = (
total_tokens > NOSPLITK_TOKEN_THRESHOLD_LOW_TOPK
and total_topk >= DUAL_SCOPE_SPLITK_TOPK_THRESHOLD
)
# For large h_q, the non-splitk grid has very few blocks
# in the H dimension, leading to low GPU utilization.
use_splitk_for_large_hq = (
h_q > 64 and total_tokens > SMALL_BATCH_TOKEN_THRESHOLD and total_topk >= 256
)
if not (
use_splitk_for_small_bs
or use_splitk_for_h64_large_topk
or use_splitk_for_large_topk
or use_splitk_for_large_hq
):
return 0 # No split-K
# Select split_k value based on workload characteristics.
# Higher topk benefits from more splits; lower topk needs fewer to
# avoid combine overhead.
if total_tokens <= SMALL_BATCH_TOKEN_THRESHOLD:
if total_topk >= SPLITK_HIGH_TOPK_THRESHOLD and total_tokens <= 4:
return 8
return 4
elif use_splitk_for_large_hq:
if total_topk >= SPLITK_HIGH_TOPK_THRESHOLD:
return 4
return 2
elif use_splitk_for_h64_large_topk:
return 2
else:
return _select_split_k(total_topk, h_q, total_tokens)
DSV4_NUM_TILES = 7 DSV4_NUM_TILES = 7
DSV4_BYTES_PER_TOKEN_DATA = 576 # 448 nope + 128 rope DSV4_BYTES_PER_TOKEN_DATA = 576 # 448 nope + 128 rope
DSV4_BYTES_PER_TOKEN_SCALE = 8 # 7 scales + 1 padding DSV4_BYTES_PER_TOKEN_SCALE = 8 # 7 scales + 1 padding
# ============================================================================ # ============================================================================
# Helper: Process KV block and compute QK scores + accumulator update # Helper: Process KV block and compute QK scores + accumulator update
# This is the core computation shared by both single and dual scope kernels # This is the core computation shared by both single and dual scope kernels
# ============================================================================ # ============================================================================
@triton.jit @triton.jit
def _process_kv_block_aggressive( def _process_kv_block_aggressive(
# KV cache parameters # KV cache parameters
@@ -198,20 +277,20 @@ def _process_kv_block_aggressive(
# ============================================================================ # ============================================================================
# DSV4 Fused Gather+Dequant+Attention Kernel (Single Scope) # DSV4 Fused Gather+Dequant+Attention Kernel (Single Scope)
# ============================================================================ # ============================================================================
@triton.autotune( @triton.autotune(
configs=[ configs=[
# Fused gather+dequant+attention kernel. triton.Config({"BLOCK_H": 16, "BLOCK_N": 32}, num_warps=4, num_stages=1),
# Two axes: BLOCK_H × BLOCK_N, with BLOCK_N being the key perf knob
# for h_q=64 where fewer BLOCK_H values affect the grid.
# BLOCK_N=64: better for large topk (less register pressure per iter).
# BLOCK_N=128: better for small topk (fewer iterations).
# num_warps=4: fused kernel is compute-bound.
triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 32}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 64, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 64, "BLOCK_N": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 128, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 128, "BLOCK_N": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 64}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 64}, num_warps=8, num_stages=1),
], ],
key=["total_tokens_bucket", "h_q", "topk"], key=["total_tokens_bucket", "h_q", "topk"],
) )
@@ -494,9 +573,8 @@ def _fused_gather_attn_dsv4_kernel(
tl.store(lse_ptrs, lse, mask=mask_h) tl.store(lse_ptrs, lse, mask=mask_h)
# Threshold for disabling AMD buffer_ops optimization # Threshold for disabling buffer_ops optimization
# When KV cache size exceeds INT32_MAX, buffer_ops can cause int32 overflow # When KV cache size exceeds this threshold, buffer_ops may overflow
# INT32_MAX = 2^31 - 1 = 2,147,483,647 bytes (~2GB)
BUFFER_OPS_DISABLE_THRESHOLD = 2 * 1024 * 1024 * 1024 # 2GB BUFFER_OPS_DISABLE_THRESHOLD = 2 * 1024 * 1024 * 1024 # 2GB
@@ -739,34 +817,26 @@ def fused_gather_attn_decode_dsv4(
return output, lse return output, lse
# Uses helper function to eliminate code duplication
# ============================================================================ # ============================================================================
def _prune_dual_scope_configs(configs, named_args, **kwargs): def _prune_dual_scope_configs(configs, named_args, **kwargs):
"""Prune configs where BLOCK_H > h_q for the dual-scope kernel. """Prune autotune configs for the dual-scope kernel.
When BLOCK_H > h_q, cdiv(h_q, BLOCK_H) = 1 regardless of BLOCK_H value, For h_q <= 64: restrict to BLOCK_H=16 only (BLOCK_H >= 32 causes
so larger BLOCK_H gives the same grid but may have worse register allocation. precision issues in online softmax due to different MFMA reduction orders).
Keep only the smallest BLOCK_H that gives cdiv(h_q, BLOCK_H) = 1, plus For h_q > 64: prune BLOCK_H > h_q (same grid size, worse register usage).
any BLOCK_H <= h_q configs.
For h_q=64: keep BLOCK_H <= 64 (removes BLOCK_H=128 which gives same grid)
For h_q=128: keep all (all give different grid sizes)
""" """
h_q = named_args.get("h_q", 128) h_q = named_args.get("h_q", 128)
pruned = [c for c in configs if c.kwargs.get("BLOCK_H", 16) <= h_q] if h_q <= 64:
pruned = [c for c in configs if c.kwargs.get("BLOCK_H", 16) <= 16]
else:
pruned = [c for c in configs if c.kwargs.get("BLOCK_H", 16) <= h_q]
return pruned if pruned else configs return pruned if pruned else configs
@triton.autotune( @triton.autotune(
configs=[ configs=[
# Dual-scope fused gather+dequant+attention.
# Three axes: BLOCK_H × BLOCK_N × (warps, stages).
# - BLOCK_H: {16, 32, 64, 128} covers h_q=64 and h_q=128.
# - BLOCK_N: {64, 128}. BLOCK_N=64 better for large topk, 128 for small topk.
# - _prune_dual_scope_configs removes BLOCK_H > h_q configs (e.g. BLOCK_H=128
# is pruned when h_q=64 since it gives the same grid as BLOCK_H=64).
# warps=4: baseline configs # warps=4: baseline configs
triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1),
@@ -1173,47 +1243,22 @@ def _fused_gather_attn_dsv4_dual_scope_kernel(
tl.store(lse_ptrs, lse, mask=mask_h) tl.store(lse_ptrs, lse, mask=mask_h)
def _prune_splitk_configs(configs, named_args, **kwargs):
"""Prune BLOCK_H=16 configs for large batch sizes to avoid CU oversubscription.
With h_q=128 and BLOCK_H=16, the grid has cdiv(128,16)=8 H-blocks.
At bs=32 with split_k=2, this creates 8*32*2=512 blocks (200% CU),
causing performance regression from oversubscription.
For small batch sizes (bucket <= 8), BLOCK_H=16 provides better
parallelism and is ~10% faster in CUDA graph replay.
"""
total_tokens_bucket = named_args.get("total_tokens_bucket", 32)
if total_tokens_bucket > 8:
# Remove BLOCK_H=16 configs for large batch sizes
pruned = [c for c in configs if c.kwargs.get("BLOCK_H", 32) > 16]
if pruned:
return pruned
return configs
# ============================================================================ # ============================================================================
# Split-K Kernel for Dual Scope # Split-K Kernel for Dual Scope
# ============================================================================ # ============================================================================
@triton.autotune( @triton.autotune(
configs=[ configs=[
# Split-K dual-scope fused kernel. # BLOCK_H=16 only (BLOCK_H >= 32 causes precision issues).
# - Split-K adds parallelism in K dim (2-8 splits). triton.Config({"BLOCK_H": 16, "BLOCK_N": 32}, num_warps=4, num_stages=1),
# - BLOCK_N={64,128}: BLOCK_N=64 better for large topk_per_split.
# - num_warps=4: compute-bound fused kernel.
# - BLOCK_H={16,64}: covers h_q=64 and h_q=128.
triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 128}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 128}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 64}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 32}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 128}, num_warps=4, num_stages=1),
# BLOCK_H=32: critical for cc=32 with h_q=128 (gives 256 blocks with split_k=2)
triton.Config({"BLOCK_H": 32, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 32, "BLOCK_N": 128}, num_warps=4, num_stages=1),
], ],
key=["total_tokens_bucket", "h_q", "topk_per_split"], key=["total_tokens_bucket", "h_q", "topk_per_split"],
prune_configs_by={"early_config_prune": _prune_splitk_configs},
) )
@triton.jit @triton.jit
def _fused_gather_attn_dsv4_dual_scope_splitk_kernel( def _fused_gather_attn_dsv4_dual_scope_splitk_kernel(
@@ -1654,58 +1699,12 @@ def fused_gather_attn_decode_dsv4_dual_scope(
or kv_cache_size_extra > BUFFER_OPS_DISABLE_THRESHOLD or kv_cache_size_extra > BUFFER_OPS_DISABLE_THRESHOLD
) )
# When force_no_splitk is set, skip the split-K decision and fall split_k = (
# through to the non-splitk kernel path below. 0
use_splitk = not force_no_splitk if force_no_splitk
else _decide_splitk_dual_scope(total_tokens, h_q, total_topk)
# Use Split-K for dual scope in these cases:
# 1. Small batch sizes with h_q=128 or large topk to increase GPU parallelism
# 2. Large topk (>= 2048) with medium/large batch sizes
# 3. NEW: h_q=64 + large topk (>=1024) + medium batch sizes (~21% improvement)
SPLITK_DUAL_SCOPE_TOPK_THRESHOLD = 2048
# For small bs, only use splitk when h_q=128 or total_topk >= 1024
use_splitk_for_small_bs = total_tokens <= 8 and (h_q >= 128 or total_topk >= 1024)
# NEW: For h_q=64 with large topk, splitk is beneficial for medium batch sizes
# Only for tokens <= 128 based on benchmarking (bs=64 shows 13% improvement)
use_splitk_for_h64_large_topk = (
h_q <= 64 and total_topk >= 1024 and total_tokens > 8 and total_tokens <= 128
) )
use_splitk_for_large_topk = ( if split_k > 0:
total_tokens > 64 and total_topk >= SPLITK_DUAL_SCOPE_TOPK_THRESHOLD
)
# For h_q > 64 (e.g. h_q=128), the non-splitk grid has very few blocks
# in the H dimension, leading to low GPU utilization at medium batch sizes.
use_splitk_for_large_hq = h_q > 64 and total_tokens > 8 and total_topk >= 256
if use_splitk and (
use_splitk_for_small_bs
or use_splitk_for_h64_large_topk
or use_splitk_for_large_topk
or use_splitk_for_large_hq
):
# Select split_k based on workload and total_topk.
# CUDA graph replay benchmarks show optimal split_k depends on both:
# - High topk (>=512, c4 layers): more splits needed to parallelize
# - Low topk (<512, c128 layers): fewer splits, less combine overhead
if total_tokens <= 8:
if total_topk >= 512 and total_tokens <= 4:
# High topk + very small bs: split_k=8 is 8-33% faster than sk=4
split_k = 8
else:
# split_k=4 gives 2x more blocks than split_k=2
split_k = 4
elif use_splitk_for_large_hq:
# For h_q > 64 with bs > 8:
if total_topk >= 512:
# High topk: split_k=4 for all medium/large bs
split_k = 4
else:
# Low topk: split_k=2 is sufficient
split_k = 2
elif use_splitk_for_h64_large_topk:
# For h_q=64 + large topk + medium bs, split_k=2 is optimal
split_k = 2
else:
split_k = _select_split_k(total_topk, h_q, total_tokens)
topk_per_split = (total_topk + split_k - 1) // split_k topk_per_split = (total_topk + split_k - 1) // split_k
partial_output = torch.empty( partial_output = torch.empty(
@@ -1926,21 +1925,20 @@ def fused_gather_attn_decode_dsv4_dual_scope(
# Split-K Optimization for Large TopK (>= 8192) # Split-K Optimization for Large TopK (>= 8192)
# ============================================================================ # ============================================================================
SPLITK_TOPK_THRESHOLD = 8192 SPLITK_TOPK_THRESHOLD = 8192
SPLITK_DEFAULT = 4
@triton.autotune( @triton.autotune(
configs=[ configs=[
# Split-K fused kernel for large topk (≥8192).
# - BLOCK_N={16,32}: small blocks for scattered FP8 KV access pattern.
# - num_warps=4: balanced for fused dequant+attention compute.
# - BLOCK_H={16,64}: covers h_q=64 and h_q=128.
triton.Config({"BLOCK_H": 16, "BLOCK_N": 16}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 16}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 16, "BLOCK_N": 32}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_N": 32}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 16, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 16}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 64, "BLOCK_N": 16}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 32}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 64, "BLOCK_N": 32}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 64}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 16}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 128, "BLOCK_N": 16}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 32}, num_warps=4, num_stages=1), triton.Config({"BLOCK_H": 128, "BLOCK_N": 32}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_N": 32}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 128, "BLOCK_N": 32}, num_warps=8, num_stages=1),
], ],
key=["total_tokens_bucket", "h_q", "topk_per_split"], key=["total_tokens_bucket", "h_q", "topk_per_split"],
) )
@@ -2390,11 +2388,7 @@ def _combine_splitk_kernel(
@triton.autotune( @triton.autotune(
configs=[ configs=[
# Simple reduce kernel (weighted sum of 8 splits). triton.Config({"BLOCK_H": 16, "BLOCK_D": 512}, num_warps=4, num_stages=1),
# - BLOCK_D=512: covers d_v=512 in one pass (no D-dimension loop).
# - num_warps=8: memory-bound reduce benefits from more warps.
# - split_k=8 is only used at very small batch sizes (≤4 tokens),
# so BLOCK_H=16/32/64 covers the relevant parallelism range.
triton.Config({"BLOCK_H": 16, "BLOCK_D": 512}, num_warps=8, num_stages=1), triton.Config({"BLOCK_H": 16, "BLOCK_D": 512}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 32, "BLOCK_D": 512}, num_warps=8, num_stages=1), triton.Config({"BLOCK_H": 32, "BLOCK_D": 512}, num_warps=8, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_D": 512}, num_warps=8, num_stages=1), triton.Config({"BLOCK_H": 64, "BLOCK_D": 512}, num_warps=8, num_stages=1),
@@ -2731,12 +2725,11 @@ def _select_split_k(topk: int, h_q: int, total_tokens: int = 64) -> int:
the topk dimension. Larger split_k increases parallelism but also increases the topk dimension. Larger split_k increases parallelism but also increases
the overhead of the combine kernel. the overhead of the combine kernel.
Updated heuristics based on benchmarking with optimized BLOCK_N configs: Heuristics:
- For large topk (>= 16384): split_k=4 provides good balance with existing combine kernel - For large topk (>= SPLITK_TOPK_THRESHOLD): split_k=4
- For medium topk (8192-16383): split_k=4 - For small topk (< SPLITK_TOPK_THRESHOLD): split_k=2
- For small topk (< 8192): split_k=2
""" """
if topk >= 8192: if topk >= SPLITK_TOPK_THRESHOLD:
return 4 return 4
else: else:
return 2 return 2
@@ -2745,12 +2738,13 @@ def _select_split_k(topk: int, h_q: int, total_tokens: int = 64) -> int:
# ============================================================================ # ============================================================================
# Low-overhead buffer pool for splitk operations # Low-overhead buffer pool for splitk operations
# ============================================================================ # ============================================================================
class SplitKBufferPool: class SplitKBufferPool:
""" """
Pre-allocated buffer pool for split-K intermediate tensors. Pre-allocated buffer pool for split-K intermediate tensors.
Caches partial_output and partial_lse buffers to avoid repeated allocations. Caches intermediate buffers to avoid repeated allocations.
Output buffers are always freshly allocated to ensure correctness.
""" """
_buffers = {} _buffers = {}
@@ -2809,7 +2803,6 @@ def fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
to minimize Python overhead, which is significant for small batch sizes. to minimize Python overhead, which is significant for small batch sizes.
The kernel computation is identical to the original version. The kernel computation is identical to the original version.
Output buffers are always freshly allocated to ensure correctness.
""" """
total_tokens, h_q, d_qk = q.shape total_tokens, h_q, d_qk = q.shape
topk_main = indices_main.shape[1] topk_main = indices_main.shape[1]
@@ -2839,25 +2832,9 @@ def fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
indices_extra = indices_extra.contiguous() indices_extra = indices_extra.contiguous()
# Determine split_k # Determine split_k
SPLITK_DUAL_SCOPE_TOPK_THRESHOLD = 2048 split_k = _decide_splitk_dual_scope(total_tokens, h_q, total_topk)
use_splitk_for_small_bs = total_tokens <= 8 and (h_q >= 128 or total_topk >= 1024)
use_splitk_for_h64_large_topk = (
h_q <= 64 and total_topk >= 1024 and total_tokens > 8 and total_tokens <= 128
)
use_splitk_for_large_topk = (
total_tokens > 64 and total_topk >= SPLITK_DUAL_SCOPE_TOPK_THRESHOLD
)
# For h_q > 64 (e.g. h_q=128), the non-splitk grid has very few blocks
# in the H dimension (cdiv(128,64)=2), leading to low GPU utilization
# at medium batch sizes. Split-K doubles the parallelism.
use_splitk_for_large_hq = h_q > 64 and total_tokens > 8 and total_topk >= 256
if not ( if split_k == 0:
use_splitk_for_small_bs
or use_splitk_for_h64_large_topk
or use_splitk_for_large_topk
or use_splitk_for_large_hq
):
# Fall back to non-splitk version # Fall back to non-splitk version
return fused_gather_attn_decode_dsv4_dual_scope( return fused_gather_attn_decode_dsv4_dual_scope(
q, q,
@@ -2872,32 +2849,9 @@ def fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
topk_length_extra, topk_length_extra,
attn_sink, attn_sink,
s_q, s_q,
force_no_splitk=True,
) )
# Select split_k based on workload and total_topk.
# CUDA graph replay benchmarks show optimal split_k depends on both:
# - High topk (>=512, c4 layers): more splits needed to parallelize
# - Low topk (<512, c128 layers): fewer splits, less combine overhead
if total_tokens <= 8:
if total_topk >= 512 and total_tokens <= 4:
# High topk + very small bs: split_k=8 is 8-33% faster than sk=4
split_k = 8
else:
# split_k=4 gives 2x more blocks than split_k=2
split_k = 4
elif use_splitk_for_large_hq:
# For h_q > 64 with bs > 8:
if total_topk >= 512:
# High topk: split_k=4 for all medium/large bs
split_k = 4
else:
# Low topk: split_k=2 is sufficient
split_k = 2
elif use_splitk_for_h64_large_topk:
split_k = 2
else:
split_k = _select_split_k(total_topk, h_q, total_tokens)
topk_per_split = (total_topk + split_k - 1) // split_k topk_per_split = (total_topk + split_k - 1) // split_k
# Get pre-allocated intermediate buffers # Get pre-allocated intermediate buffers
@@ -2907,8 +2861,7 @@ def fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
stride_po = buffers["stride_po"] stride_po = buffers["stride_po"]
stride_plse = buffers["stride_plse"] stride_plse = buffers["stride_plse"]
# Reuse pre-allocated output buffers to avoid torch.empty() calls # Allocate output buffers
# that would be captured in CUDA graphs (each adds ~7-8us replay overhead).
output = torch.empty(total_tokens, h_q, d_v, dtype=torch.bfloat16, device=device) output = torch.empty(total_tokens, h_q, d_v, dtype=torch.bfloat16, device=device)
lse = torch.empty(total_tokens, h_q, dtype=torch.float32, device=device) lse = torch.empty(total_tokens, h_q, dtype=torch.float32, device=device)
@@ -1,13 +1,13 @@
""" """
Optimized Triton MLA Decode Kernels for DeepSeek V4. Optimized Triton MLA Decode Kernels for DeepSeek V4.
This module provides optimized sparse attention decode with reduced Python overhead. This module provides optimized sparse attention decode.
Key optimizations: Key optimizations:
1. Fused gather+dequant+attention kernels (eliminates intermediate buffers) 1. Fused gather+dequant+attention kernels (eliminates intermediate buffers)
2. Split-K for better GPU parallelism on small batches 2. Split-K for better GPU parallelism on small batches
3. Pre-allocated buffer pool for splitk intermediate results 3. Proper dispatch: no-splitk for large batches, split-K for small batches
4. Pre-computed strides to reduce tensor metadata operations 4. All paths use fused kernels (no 2-phase fallback)
Note: This implementation assumes KV cache is always FP8 quantized. Note: This implementation assumes KV cache is always FP8 quantized.
""" """
@@ -15,24 +15,38 @@ Note: This implementation assumes KV cache is always FP8 quantized.
from typing import Optional, Tuple from typing import Optional, Tuple
import torch import torch
import triton
from .triton_mla_kernels_decode_common import (
_bucket_total_tokens,
_unified_sparse_decode_kernel,
compute_token_ranges,
)
from .triton_mla_kernels_decode_dsv4 import (
DSV4_D_QK,
fused_gather_dequant_fp8_dsv4,
)
from .triton_mla_kernels_decode_fused import ( from .triton_mla_kernels_decode_fused import (
DSV4_D_QK,
fused_gather_attn_decode_dsv4, fused_gather_attn_decode_dsv4,
fused_gather_attn_decode_dsv4_dual_scope, fused_gather_attn_decode_dsv4_dual_scope,
fused_gather_attn_decode_dsv4_dual_scope_low_overhead, fused_gather_attn_decode_dsv4_dual_scope_low_overhead,
) )
def _should_use_fused_splitk(total_tokens: int, h_q: int, total_topk: int) -> bool:
"""Determine whether to use fused split-K kernel (low overhead).
The fused split-K kernel is preferred for small batch sizes because
split-K provides better GPU utilization when the grid is small.
This matches the original _should_use_fused_dual_scope() thresholds.
"""
if total_tokens <= 4:
return True
if h_q <= 64 and total_topk <= 800:
return total_tokens <= 256
if h_q <= 64 and total_topk >= 1024:
return total_tokens <= 128
# h_q > 64 (e.g. h_q=128 when q is padded to full n_heads).
if h_q > 64:
if total_topk >= 400:
return total_tokens <= 32
else:
return total_tokens <= 128
return True
def triton_sparse_attn_decode( def triton_sparse_attn_decode(
q: torch.Tensor, q: torch.Tensor,
kv_scope, kv_scope,
@@ -54,67 +68,6 @@ def triton_sparse_attn_decode(
) )
def _should_use_fused_dual_scope(total_tokens: int, h_q: int, total_topk: int) -> bool:
"""Determine whether to use fused kernel for dual-scope cases.
Returns True if the fused kernel (with splitk for small bs) should be used.
For large batch sizes (>= 256), use _should_use_fused_nosplitk instead.
The thresholds below were determined empirically on MI355X (256 CUs).
"""
if total_tokens <= 4:
return True
if h_q <= 64 and total_topk <= 800:
return total_tokens <= 256
if h_q <= 64 and total_topk >= 1024:
return total_tokens <= 128
# h_q > 64 (e.g. h_q=128 when q is padded to full n_heads).
if h_q > 64:
if total_topk >= 400:
return total_tokens <= 32
else:
return total_tokens <= 128
return True
def _should_use_fused_nosplitk(total_tokens: int, h_q: int, total_topk: int) -> bool:
"""Determine whether to use the fused no-splitk kernel for large batches.
Kernel-level benchmarking on MI355X shows that for large batch sizes
(total_tokens >= 256), the fused dual-scope kernel WITHOUT split-K
is ~10% faster than the separate gather+attention path:
total_tokens=256: fused-noSK=169us vs separate=194us (14% faster)
total_tokens=512: fused-noSK=350us vs separate=408us (14% faster)
total_tokens=1024: fused-noSK=700us vs separate=777us (10% faster)
total_tokens=4096: fused-noSK=2761us vs separate=3063us (10% faster)
The fused no-splitk kernel avoids:
1. Materializing the large intermediate gathered_kv buffer
2. The separate gather kernel launch
3. The split-K combine overhead
For total_tokens < 256, the separate path is faster because the
fused kernel has insufficient parallelism.
For extend (total_tokens >= 1024), the fused kernel always wins
regardless of h_q or total_topk because:
- The grid already has thousands of blocks (good GPU utilization)
- It eliminates 1.5-5 GB gathered_kv buffer allocation
- It eliminates 2x gather_dequant kernel launches (~414 us)
- It avoids chunking that TP>1 configs require with the separate path
"""
if total_tokens >= 1024:
return True
if h_q <= 64:
return False # Not benchmarked for h_q <= 64
if total_topk < 200:
return False # Small topk doesn't benefit
# For h_q > 64 and total_topk >= 200:
# Fused no-splitk wins for total_tokens >= 256
return total_tokens >= 256
def _triton_sparse_attn_decode_dsv4( def _triton_sparse_attn_decode_dsv4(
q: torch.Tensor, q: torch.Tensor,
kv_scope, kv_scope,
@@ -123,10 +76,17 @@ def _triton_sparse_attn_decode_dsv4(
d_v: int, d_v: int,
attn_sink: Optional[torch.Tensor], attn_sink: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
"""Optimized sparse attention decode for DeepSeek V4 (d_qk=512).""" """Sparse attention decode for DeepSeek V4 (d_qk=512).
All paths use fused kernels (no 2-phase fallback).
Dispatch logic:
- Single scope: always use fused kernel
- Dual scope, small total_tokens: fused split-K kernel (low overhead)
- Dual scope, otherwise: fused no-splitk kernel
"""
b, s_q, h_q, d_qk = q.shape b, s_q, h_q, d_qk = q.shape
total_tokens = b * s_q total_tokens = b * s_q
device = q.device
topk_main = kv_scope.indices_in_kvcache.shape[-1] topk_main = kv_scope.indices_in_kvcache.shape[-1]
kv_quantized_main = kv_scope.blocked_k_quantized kv_quantized_main = kv_scope.blocked_k_quantized
@@ -134,52 +94,55 @@ def _triton_sparse_attn_decode_dsv4(
# Single scope case # Single scope case
if extra_kv_scope is None: if extra_kv_scope is None:
if topk_main < 8192:
q_reshaped = q.reshape(total_tokens, h_q, d_qk)
if not q_reshaped.is_contiguous():
q_reshaped = q_reshaped.contiguous()
indices_main = kv_scope.indices_in_kvcache.reshape(total_tokens, topk_main)
if not indices_main.is_contiguous():
indices_main = indices_main.contiguous()
output, lse = fused_gather_attn_decode_dsv4(
q_reshaped,
kv_quantized_main,
indices_main,
block_size_main,
sm_scale,
topk_length=kv_scope.topk_length,
attn_sink=attn_sink,
s_q=s_q,
)
return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2)
else:
from .triton_mla_kernels_decode_dsv4 import triton_sparse_attn_decode_dsv4
return triton_sparse_attn_decode_dsv4(
q, kv_scope, extra_kv_scope, sm_scale, d_v, attn_sink
)
# Dual scope case
topk_extra = extra_kv_scope.indices_in_kvcache.shape[-1]
total_topk = topk_main + topk_extra
# For large batch sizes, use fused no-splitk kernel (10% faster than separate).
# This check is BEFORE the chunking check because the fused kernel does NOT
# allocate the intermediate gathered_kv buffer, so buffer size limits don't apply.
if _should_use_fused_nosplitk(total_tokens, h_q, total_topk):
q_reshaped = q.reshape(total_tokens, h_q, d_qk).contiguous() q_reshaped = q.reshape(total_tokens, h_q, d_qk).contiguous()
indices_main = kv_scope.indices_in_kvcache.reshape( indices_main = kv_scope.indices_in_kvcache.reshape(
total_tokens, topk_main total_tokens, topk_main
).contiguous() ).contiguous()
block_size_extra = extra_kv_scope.blocked_k.shape[1] output, lse = fused_gather_attn_decode_dsv4(
indices_extra = extra_kv_scope.indices_in_kvcache.reshape( q_reshaped,
total_tokens, topk_extra kv_quantized_main,
).contiguous() indices_main,
block_size_main,
sm_scale,
topk_length=kv_scope.topk_length,
attn_sink=attn_sink,
s_q=s_q,
)
return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2)
# Dual scope case
topk_extra = extra_kv_scope.indices_in_kvcache.shape[-1]
total_topk = topk_main + topk_extra
block_size_extra = extra_kv_scope.blocked_k.shape[1]
q_reshaped = q.reshape(total_tokens, h_q, d_qk).contiguous()
indices_main = kv_scope.indices_in_kvcache.reshape(
total_tokens, topk_main
).contiguous()
indices_extra = extra_kv_scope.indices_in_kvcache.reshape(
total_tokens, topk_extra
).contiguous()
# Dispatch: use split-K for small batches, no-splitk for everything else.
if _should_use_fused_splitk(total_tokens, h_q, total_topk):
# Small batch: fused split-K kernel (better GPU utilization)
output, lse = fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
q_reshaped,
kv_quantized_main,
indices_main,
block_size_main,
extra_kv_scope.blocked_k_quantized,
indices_extra,
block_size_extra,
sm_scale,
topk_length_main=kv_scope.topk_length,
topk_length_extra=extra_kv_scope.topk_length,
attn_sink=attn_sink,
s_q=s_q,
)
else:
# Large batch / extend / prefill: fused no-splitk kernel
output, lse = fused_gather_attn_decode_dsv4_dual_scope( output, lse = fused_gather_attn_decode_dsv4_dual_scope(
q_reshaped, q_reshaped,
kv_quantized_main, kv_quantized_main,
@@ -195,153 +158,5 @@ def _triton_sparse_attn_decode_dsv4(
s_q=s_q, s_q=s_q,
force_no_splitk=True, force_no_splitk=True,
) )
return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2)
# Check if chunking needed for separate path (fall back to original implementation)
token_ranges = compute_token_ranges(total_tokens, total_topk, d_qk)
if len(token_ranges) > 1:
from .triton_mla_kernels_decode_dsv4 import triton_sparse_attn_decode_dsv4
return triton_sparse_attn_decode_dsv4(
q, kv_scope, extra_kv_scope, sm_scale, d_v, attn_sink
)
# Use fused dual-scope kernel with low-overhead buffer pool
if _should_use_fused_dual_scope(total_tokens, h_q, total_topk):
q_reshaped = q.reshape(total_tokens, h_q, d_qk).contiguous()
indices_main = kv_scope.indices_in_kvcache.reshape(
total_tokens, topk_main
).contiguous()
block_size_extra = extra_kv_scope.blocked_k.shape[1]
indices_extra = extra_kv_scope.indices_in_kvcache.reshape(
total_tokens, topk_extra
).contiguous()
output, lse = fused_gather_attn_decode_dsv4_dual_scope_low_overhead(
q_reshaped,
kv_quantized_main,
indices_main,
block_size_main,
extra_kv_scope.blocked_k_quantized,
indices_extra,
block_size_extra,
sm_scale,
topk_length_main=kv_scope.topk_length,
topk_length_extra=extra_kv_scope.topk_length,
attn_sink=attn_sink,
s_q=s_q,
)
return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2)
# Fallback: Separate gather + attention path
return _fallback_gather_attention(
q,
kv_scope,
extra_kv_scope,
sm_scale,
d_v,
attn_sink,
total_tokens,
h_q,
d_qk,
topk_main,
topk_extra,
block_size_main,
kv_quantized_main,
fused_gather_dequant_fp8_dsv4,
)
def _fallback_gather_attention(
q: torch.Tensor,
kv_scope,
extra_kv_scope,
sm_scale: float,
d_v: int,
attn_sink: Optional[torch.Tensor],
total_tokens: int,
h_q: int,
d_qk: int,
topk_main: int,
topk_extra: int,
block_size_main: int,
kv_quantized_main,
fused_gather_fn,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Fallback path: separate gather + attention kernels."""
b = q.shape[0]
s_q = q.shape[1]
device = q.device
total_topk = topk_main + topk_extra
gathered_kv = torch.empty(
total_tokens, total_topk, d_qk, dtype=torch.bfloat16, device=device
)
invalid_mask = torch.empty(
total_tokens, total_topk, dtype=torch.bool, device=device
)
output = torch.empty(total_tokens, h_q, d_v, dtype=torch.bfloat16, device=device)
lse = torch.empty(total_tokens, h_q, dtype=torch.float32, device=device)
indices_main = kv_scope.indices_in_kvcache.reshape(total_tokens, topk_main)
block_size_extra = extra_kv_scope.blocked_k.shape[1]
indices_extra = extra_kv_scope.indices_in_kvcache.reshape(total_tokens, topk_extra)
fused_gather_fn(
kv_quantized_main,
indices_main,
block_size_main,
kv_scope.topk_length,
extra_kv_scope.blocked_k_quantized,
indices_extra,
block_size_extra,
extra_kv_scope.topk_length,
gathered_kv,
invalid_mask,
s_q,
)
if q.dtype == torch.bfloat16 and q.is_contiguous():
q_reshaped = q.view(total_tokens, h_q, d_qk)
else:
q_reshaped = q.to(torch.bfloat16).reshape(total_tokens, h_q, d_qk)
if not q_reshaped.is_contiguous():
q_reshaped = q_reshaped.contiguous()
HAS_ATTN_SINK = attn_sink is not None
attn_sink_tensor = attn_sink if HAS_ATTN_SINK else lse[:1]
grid = lambda meta: (total_tokens, triton.cdiv(h_q, meta["BLOCK_H"]))
_unified_sparse_decode_kernel[grid](
q_reshaped,
gathered_kv,
invalid_mask,
attn_sink_tensor,
output,
lse,
sm_scale,
total_tokens,
_bucket_total_tokens(total_tokens),
h_q,
total_topk,
d_qk,
d_v,
q_reshaped.stride(0),
q_reshaped.stride(1),
q_reshaped.stride(2),
gathered_kv.stride(0),
gathered_kv.stride(1),
gathered_kv.stride(2),
invalid_mask.stride(0),
invalid_mask.stride(1),
output.stride(0),
output.stride(1),
output.stride(2),
lse.stride(0),
lse.stride(1),
HAS_ATTN_SINK=HAS_ATTN_SINK,
)
return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2) return output.view(b, s_q, h_q, d_v), lse.view(b, s_q, h_q).transpose(1, 2)
@@ -1,534 +0,0 @@
"""
Split-K Attention Kernel for Large TopK Cases
This module implements a split-K version of the attention kernel that:
1. Splits the K (topk) dimension across multiple kernel instances
2. Each instance computes partial results with its own m_i, l_i, and accumulators
3. A combine kernel merges the partial results using online softmax
This reduces register pressure by processing fewer K tokens per kernel instance,
improving occupancy and overall performance for large topk cases.
"""
from typing import Optional, Tuple
import torch
import triton
import triton.language as tl
from .triton_mla_kernels_decode_common import _bucket_total_tokens
# ============================================================================
# Split-K Attention Kernel
# ============================================================================
@triton.autotune(
configs=[
# Split-K attention on already-gathered BF16 KV.
# - BLOCK_N=256: amortizes memory access over KV tokens (memory-bound kernel).
# - BLOCK_D=128: matches KV tile structure.
# - num_warps=8, num_stages=2: memory-bound kernel benefits from more warps
# and software pipelining (overlaps memory loads with compute).
# - BLOCK_H varies for different batch sizes:
triton.Config(
{"BLOCK_H": 16, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=2
),
triton.Config(
{"BLOCK_H": 32, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=2
),
triton.Config(
{"BLOCK_H": 64, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=2
),
triton.Config(
{"BLOCK_H": 128, "BLOCK_N": 256, "BLOCK_D": 128}, num_warps=8, num_stages=2
),
],
key=["total_tokens_bucket", "h_q", "topk_per_split", "d_qk"],
)
@triton.jit
def _splitk_attention_kernel(
Q,
KV,
Mask,
PartialOutput,
PartialLSE,
PartialM,
sm_scale,
total_tokens,
total_tokens_bucket,
h_q,
total_topk,
d_qk,
d_v,
topk_per_split,
stride_q_t,
stride_q_h,
stride_q_d,
stride_kv_t,
stride_kv_k,
stride_kv_d,
stride_mask_t,
stride_mask_k,
stride_po_s,
stride_po_t,
stride_po_h,
stride_po_d,
stride_plse_s,
stride_plse_t,
stride_plse_h,
stride_pm_s,
stride_pm_t,
stride_pm_h,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_D: tl.constexpr,
):
"""Split-K attention kernel that processes a subset of K tokens."""
LOG2E: tl.constexpr = 1.4426950408889634
pid_t = tl.program_id(0)
pid_h = tl.program_id(1)
pid_k = tl.program_id(2)
pid_t_64 = pid_t.to(tl.int64)
NEG_INF = float("-inf")
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
mask_h = offs_h < h_q
# Compute K range for this split
k_start = pid_k * topk_per_split
k_end = tl.minimum(k_start + topk_per_split, total_topk)
m_i = tl.full([BLOCK_H], NEG_INF, dtype=tl.float32)
l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
acc_0 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_1 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_2 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_3 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
stride_q_t_64 = tl.cast(stride_q_t, tl.int64)
stride_kv_t_64 = tl.cast(stride_kv_t, tl.int64)
stride_mask_t_64 = tl.cast(stride_mask_t, tl.int64)
q_base = Q + pid_t_64 * stride_q_t_64
kv_base = KV + pid_t_64 * stride_kv_t_64
mask_base = Mask + pid_t_64 * stride_mask_t_64
for n_start in range(k_start, k_end, BLOCK_N):
offs_n = n_start + tl.arange(0, BLOCK_N)
mask_n = offs_n < k_end
mask_ptrs = mask_base + offs_n * stride_mask_k
invalid = tl.load(mask_ptrs, mask=mask_n, other=True)
valid = mask_n & ~invalid
qk = tl.zeros([BLOCK_H, BLOCK_N], dtype=tl.float32)
for d_start in range(0, d_qk, BLOCK_D):
offs_d = d_start + tl.arange(0, BLOCK_D)
mask_d = offs_d < d_qk
q_ptrs = (
q_base + offs_h[:, None] * stride_q_h + offs_d[None, :] * stride_q_d
)
q_chunk = tl.load(
q_ptrs, mask=mask_h[:, None] & mask_d[None, :], other=0.0
).to(tl.bfloat16)
k_ptrs = (
kv_base + offs_n[:, None] * stride_kv_k + offs_d[None, :] * stride_kv_d
)
k_chunk = tl.load(
k_ptrs, mask=valid[:, None] & mask_d[None, :], other=0.0
).to(tl.bfloat16)
qk += tl.dot(q_chunk, tl.trans(k_chunk))
qk = qk * sm_scale
qk = tl.where(valid[None, :], qk, NEG_INF)
m_ij = tl.max(qk, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.where(m_i == NEG_INF, 0.0, tl.math.exp2((m_i - m_new) * LOG2E))
p = tl.where(qk == NEG_INF, 0.0, tl.math.exp2((qk - m_new[:, None]) * LOG2E))
l_new = alpha * l_i + tl.sum(p, axis=1)
p_bf16 = p.to(tl.bfloat16)
offs_v = tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(v_ptrs, mask=valid[:, None], other=0.0).to(tl.bfloat16)
acc_0 = acc_0 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_1 = acc_1 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = 2 * BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_2 = acc_2 * alpha[:, None] + tl.dot(p_bf16, v)
offs_v = 3 * BLOCK_D + tl.arange(0, BLOCK_D)
v_ptrs = kv_base + offs_n[:, None] * stride_kv_k + offs_v[None, :] * stride_kv_d
v = tl.load(
v_ptrs, mask=valid[:, None] & (offs_v[None, :] < d_v), other=0.0
).to(tl.bfloat16)
acc_3 = acc_3 * alpha[:, None] + tl.dot(p_bf16, v)
m_i = m_new
l_i = l_new
# Store partial results
stride_po_s_64 = tl.cast(stride_po_s, tl.int64)
stride_po_t_64 = tl.cast(stride_po_t, tl.int64)
po_base = PartialOutput + pid_k * stride_po_s_64 + pid_t_64 * stride_po_t_64
offs_h_2d = offs_h[:, None]
mask_h_2d = mask_h[:, None]
offs_v_0 = tl.arange(0, BLOCK_D)
offs_v_1 = BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_2 = 2 * BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_3 = 3 * BLOCK_D + tl.arange(0, BLOCK_D)
tl.store(
po_base + offs_h_2d * stride_po_h + offs_v_0[None, :] * stride_po_d,
acc_0,
mask=mask_h_2d,
)
tl.store(
po_base + offs_h_2d * stride_po_h + offs_v_1[None, :] * stride_po_d,
acc_1,
mask=mask_h_2d & (offs_v_1[None, :] < d_v),
)
tl.store(
po_base + offs_h_2d * stride_po_h + offs_v_2[None, :] * stride_po_d,
acc_2,
mask=mask_h_2d & (offs_v_2[None, :] < d_v),
)
tl.store(
po_base + offs_h_2d * stride_po_h + offs_v_3[None, :] * stride_po_d,
acc_3,
mask=mask_h_2d & (offs_v_3[None, :] < d_v),
)
stride_plse_s_64 = tl.cast(stride_plse_s, tl.int64)
stride_plse_t_64 = tl.cast(stride_plse_t, tl.int64)
plse_ptrs = (
PartialLSE
+ pid_k * stride_plse_s_64
+ pid_t_64 * stride_plse_t_64
+ offs_h * stride_plse_h
)
tl.store(plse_ptrs, l_i, mask=mask_h)
stride_pm_s_64 = tl.cast(stride_pm_s, tl.int64)
stride_pm_t_64 = tl.cast(stride_pm_t, tl.int64)
pm_ptrs = (
PartialM
+ pid_k * stride_pm_s_64
+ pid_t_64 * stride_pm_t_64
+ offs_h * stride_pm_h
)
tl.store(pm_ptrs, m_i, mask=mask_h)
# ============================================================================
# Combine Kernel for Split-K
# ============================================================================
@triton.autotune(
configs=[
# Simple reduce kernel merging split-K results.
# - BLOCK_D=128: 4 iterations to cover d_v=512.
# - num_warps=4: sufficient for this simple reduce operation.
# - BLOCK_H varies for different batch sizes:
triton.Config({"BLOCK_H": 16, "BLOCK_D": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 32, "BLOCK_D": 128}, num_warps=4, num_stages=1),
triton.Config({"BLOCK_H": 64, "BLOCK_D": 128}, num_warps=4, num_stages=1),
],
key=["total_tokens_bucket", "h_q", "split_k"],
)
@triton.jit
def _combine_splitk_attention_kernel(
PartialOutput,
PartialLSE,
PartialM,
AttnSink,
Output,
LSE,
total_tokens,
total_tokens_bucket,
h_q,
d_v,
split_k,
stride_po_s,
stride_po_t,
stride_po_h,
stride_po_d,
stride_plse_s,
stride_plse_t,
stride_plse_h,
stride_pm_s,
stride_pm_t,
stride_pm_h,
stride_o_t,
stride_o_h,
stride_o_d,
stride_lse_t,
stride_lse_h,
HAS_ATTN_SINK: tl.constexpr,
BLOCK_H: tl.constexpr,
BLOCK_D: tl.constexpr,
):
"""Combine partial results from split-K attention kernel."""
LOG2E: tl.constexpr = 1.4426950408889634
NEG_INF = float("-inf")
POS_INF = float("+inf")
pid_t = tl.program_id(0)
pid_h = tl.program_id(1)
pid_t_64 = pid_t.to(tl.int64)
offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
mask_h = offs_h < h_q
m_acc = tl.full([BLOCK_H], NEG_INF, dtype=tl.float32)
l_acc = tl.zeros([BLOCK_H], dtype=tl.float32)
acc_0 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_1 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_2 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
acc_3 = tl.zeros([BLOCK_H, BLOCK_D], dtype=tl.float32)
stride_po_s_64 = tl.cast(stride_po_s, tl.int64)
stride_po_t_64 = tl.cast(stride_po_t, tl.int64)
stride_plse_s_64 = tl.cast(stride_plse_s, tl.int64)
stride_plse_t_64 = tl.cast(stride_plse_t, tl.int64)
stride_pm_s_64 = tl.cast(stride_pm_s, tl.int64)
stride_pm_t_64 = tl.cast(stride_pm_t, tl.int64)
offs_h_2d = offs_h[:, None]
mask_h_2d = mask_h[:, None]
offs_v_0 = tl.arange(0, BLOCK_D)
offs_v_1 = BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_2 = 2 * BLOCK_D + tl.arange(0, BLOCK_D)
offs_v_3 = 3 * BLOCK_D + tl.arange(0, BLOCK_D)
for k in range(split_k):
k_64 = tl.cast(k, tl.int64)
po_base = PartialOutput + k_64 * stride_po_s_64 + pid_t_64 * stride_po_t_64
p_acc_0 = tl.load(
po_base + offs_h_2d * stride_po_h + offs_v_0[None, :] * stride_po_d,
mask=mask_h_2d,
other=0.0,
)
p_acc_1 = tl.load(
po_base + offs_h_2d * stride_po_h + offs_v_1[None, :] * stride_po_d,
mask=mask_h_2d & (offs_v_1[None, :] < d_v),
other=0.0,
)
p_acc_2 = tl.load(
po_base + offs_h_2d * stride_po_h + offs_v_2[None, :] * stride_po_d,
mask=mask_h_2d & (offs_v_2[None, :] < d_v),
other=0.0,
)
p_acc_3 = tl.load(
po_base + offs_h_2d * stride_po_h + offs_v_3[None, :] * stride_po_d,
mask=mask_h_2d & (offs_v_3[None, :] < d_v),
other=0.0,
)
plse_ptrs = (
PartialLSE
+ k_64 * stride_plse_s_64
+ pid_t_64 * stride_plse_t_64
+ offs_h * stride_plse_h
)
p_l = tl.load(plse_ptrs, mask=mask_h, other=0.0)
pm_ptrs = (
PartialM
+ k_64 * stride_pm_s_64
+ pid_t_64 * stride_pm_t_64
+ offs_h * stride_pm_h
)
p_m = tl.load(pm_ptrs, mask=mask_h, other=NEG_INF)
m_new = tl.maximum(m_acc, p_m)
alpha_acc = tl.where(
m_acc == NEG_INF, 0.0, tl.math.exp2((m_acc - m_new) * LOG2E)
)
alpha_p = tl.where(p_m == NEG_INF, 0.0, tl.math.exp2((p_m - m_new) * LOG2E))
l_new = alpha_acc * l_acc + alpha_p * p_l
acc_0 = acc_0 * alpha_acc[:, None] + p_acc_0 * alpha_p[:, None]
acc_1 = acc_1 * alpha_acc[:, None] + p_acc_1 * alpha_p[:, None]
acc_2 = acc_2 * alpha_acc[:, None] + p_acc_2 * alpha_p[:, None]
acc_3 = acc_3 * alpha_acc[:, None] + p_acc_3 * alpha_p[:, None]
m_acc = m_new
l_acc = l_new
lse = m_acc + tl.math.log2(tl.where(l_acc == 0.0, 1.0, l_acc)) / LOG2E
is_lonely_q = l_acc == 0.0
if HAS_ATTN_SINK:
attn_sink_vals = tl.load(AttnSink + offs_h, mask=mask_h, other=0.0)
exp_attn_sink_minus_m = tl.math.exp2((attn_sink_vals - m_acc) * LOG2E)
denominator = l_acc + exp_attn_sink_minus_m
denominator = tl.where(denominator == 0.0, 1.0, denominator)
output_scale = 1.0 / denominator
else:
output_scale = tl.where(l_acc == 0.0, 0.0, 1.0 / l_acc)
is_lonely_q_2d = is_lonely_q[:, None]
output_scale_2d = output_scale[:, None]
acc_0 = tl.where(is_lonely_q_2d, 0.0, acc_0 * output_scale_2d)
acc_1 = tl.where(is_lonely_q_2d, 0.0, acc_1 * output_scale_2d)
acc_2 = tl.where(is_lonely_q_2d, 0.0, acc_2 * output_scale_2d)
acc_3 = tl.where(is_lonely_q_2d, 0.0, acc_3 * output_scale_2d)
lse = tl.where(is_lonely_q, POS_INF, lse)
stride_o_t_64 = tl.cast(stride_o_t, tl.int64)
o_base = Output + pid_t_64 * stride_o_t_64
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_0[None, :] * stride_o_d,
acc_0.to(tl.bfloat16),
mask=mask_h_2d,
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_1[None, :] * stride_o_d,
acc_1.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_1[None, :] < d_v),
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_2[None, :] * stride_o_d,
acc_2.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_2[None, :] < d_v),
)
tl.store(
o_base + offs_h_2d * stride_o_h + offs_v_3[None, :] * stride_o_d,
acc_3.to(tl.bfloat16),
mask=mask_h_2d & (offs_v_3[None, :] < d_v),
)
stride_lse_t_64 = tl.cast(stride_lse_t, tl.int64)
tl.store(LSE + pid_t_64 * stride_lse_t_64 + offs_h * stride_lse_h, lse, mask=mask_h)
# ============================================================================
# Runner Function
# ============================================================================
def run_splitk_attention(
q_reshaped: torch.Tensor,
gathered_kv: torch.Tensor,
invalid_mask: torch.Tensor,
d_v: int,
sm_scale: float,
total_tokens: int,
h_q: int,
total_topk: int,
d_qk: int,
attn_sink: Optional[torch.Tensor] = None,
split_k: int = 4,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run split-K attention kernel."""
device = q_reshaped.device
topk_per_split = (total_topk + split_k - 1) // split_k
partial_output = torch.empty(
split_k, total_tokens, h_q, d_v, dtype=torch.float32, device=device
)
partial_lse = torch.empty(
split_k, total_tokens, h_q, dtype=torch.float32, device=device
)
partial_m = torch.empty(
split_k, total_tokens, h_q, dtype=torch.float32, device=device
)
output = torch.empty(total_tokens, h_q, d_v, dtype=torch.bfloat16, device=device)
lse = torch.empty(total_tokens, h_q, dtype=torch.float32, device=device)
grid_splitk = lambda meta: (
total_tokens,
triton.cdiv(h_q, meta["BLOCK_H"]),
split_k,
)
_splitk_attention_kernel[grid_splitk](
q_reshaped,
gathered_kv,
invalid_mask,
partial_output,
partial_lse,
partial_m,
sm_scale,
total_tokens,
_bucket_total_tokens(total_tokens),
h_q,
total_topk,
d_qk,
d_v,
topk_per_split,
q_reshaped.stride(0),
q_reshaped.stride(1),
q_reshaped.stride(2),
gathered_kv.stride(0),
gathered_kv.stride(1),
gathered_kv.stride(2),
invalid_mask.stride(0),
invalid_mask.stride(1),
partial_output.stride(0),
partial_output.stride(1),
partial_output.stride(2),
partial_output.stride(3),
partial_lse.stride(0),
partial_lse.stride(1),
partial_lse.stride(2),
partial_m.stride(0),
partial_m.stride(1),
partial_m.stride(2),
)
HAS_ATTN_SINK = attn_sink is not None
attn_sink_tensor = attn_sink if HAS_ATTN_SINK else lse[:1]
grid_combine = lambda meta: (total_tokens, triton.cdiv(h_q, meta["BLOCK_H"]))
_combine_splitk_attention_kernel[grid_combine](
partial_output,
partial_lse,
partial_m,
attn_sink_tensor,
output,
lse,
total_tokens,
_bucket_total_tokens(total_tokens),
h_q,
d_v,
split_k,
partial_output.stride(0),
partial_output.stride(1),
partial_output.stride(2),
partial_output.stride(3),
partial_lse.stride(0),
partial_lse.stride(1),
partial_lse.stride(2),
partial_m.stride(0),
partial_m.stride(1),
partial_m.stride(2),
output.stride(0),
output.stride(1),
output.stride(2),
lse.stride(0),
lse.stride(1),
HAS_ATTN_SINK=HAS_ATTN_SINK,
)
return output, lse