submission 748648
wbh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3120 lines, June 9 Researcher Reciprocity License v1.0.
submission_sparse.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-748648?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:e0239353f9b8ca884a3a51fab6a8e7a2045f7f6ec7248ba741a8c6cbcadd4b5c
license declaredunknown
license concludedunknown
authorswbh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Converts given x (in fp32) to mxfp4 format.fp8
using gpu_fp8_storage_t = __nv_fp8_storage_t;mma
acc_s = tl.dot(q_tile, k_tile)num-warps = 4
num_warps = 4online-softmax
:return row_max_new: Updated maximum values per row of shape [BLOCK_M].shared-memory
__shared__ float shared_max[kThreadsPerBlock];split-k
IS_SPLIT_KV: tl.constexpr,stages = 2
num_stages = 2tile-k = 1
TILE_K=1,tile-m = 16
TILE_M = 16tile-n = 32
TILE_N = 32Kernel source
submission_sparse.py3120 lines
from typing import Any, Optional, Tuple
import math
import os
import shutil
import statistics
import time
import torch
import triton
import triton.language as tl
@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
## pid remapping on xcds
# Number of pids per XCD in the new arrangement
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
# When GRID_MN cannot divide NUM_XCDS, some xcds will have
# pids_per_xcd pids, the other will have pids_per_xcd - 1 pids.
# We calculate the number of xcds that have pids_per_xcd pids as
# tall_xcds
tall_xcds = GRID_MN % NUM_XCDS
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
# Compute current XCD and local pid within the XCD
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
# Calculate new pid based on the new grouping
# Note that we need to consider the following two cases:
# 1. the current pid is on a tall xcd
# 2. the current pid is on a short xcd
if xcd < tall_xcds:
pid = xcd * pids_per_xcd + local_pid
else:
pid = (
tall_xcds * pids_per_xcd
+ (xcd - tall_xcds) * (pids_per_xcd - 1)
+ local_pid
)
return pid
@triton.jit
def pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
"""
Maps 1D pid to 2D grid coords (pid_m, pid_n).
Args:
- pid: 1D pid
- num_pid_m: grid m size
- num_pid_n: grid n size
- GROUP_SIZE_M: tl.constexpr: default is 1
"""
if GROUP_SIZE_M == 1:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
else:
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
tl.assume(group_size_m >= 0)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
def num_splits_heuristic(
seqlen_q: int,
seqlen_k: int,
num_SMs: int,
TILE_M: int,
TILE_N: int,
) -> int:
"""
Determine the number of KV splits for FlashDecoding.
Splits only when there are enough KV blocks to benefit from parallelism,
and targets full SM occupancy by over-subscribing the M-block count.
:param seqlen_q: Sequence length of queries.
:param seqlen_k: Sequence length of keys.
:param num_SMs: Number of SM on the device.
:param TILE_M: Tile size for M dimension.
:param TILE_N: Tile size for N dimension.
:return: Number of splits.
"""
total_mblocks = triton.cdiv(seqlen_q, TILE_M)
num_n_blocks = triton.cdiv(seqlen_k, TILE_N)
max_splits = triton.next_power_of_2(num_SMs)
if num_n_blocks <= 4:
# 1 means no splitting
return 1
return min(num_SMs // max(total_mblocks, 1), max_splits, num_n_blocks)
def _build_benchmark_split_preset(
split_count: int,
split_boundaries: list[int],
) -> dict[str, int | list[int]]:
reduce_split_count = split_count if split_count > 1 else 0
reduce_split_boundaries = [0, reduce_split_count] if reduce_split_count else [0, 0]
return {
"split_count": split_count,
"split_boundaries": split_boundaries,
"reduce_split_count": reduce_split_count,
"reduce_split_boundaries": reduce_split_boundaries,
}
# tile_n, batch_size, seqlen_k, is_local, window_size_left, window_size_right
_BENCHMARK_SPLIT_PRESETS = {
key: _build_benchmark_split_preset(split_count, split_boundaries)
for key, (split_count, split_boundaries) in {
(16, 4, 1024, True, 1024, 0): (5, [0, 13, 26, 39, 52, 64]),
(32, 4, 1024, True, 1024, 0): (5, [0, 7, 14, 21, 28, 32]),
(16, 32, 1024, True, 1024, 0): (5, [0, 13, 26, 39, 52, 64]),
(32, 32, 1024, True, 1024, 0): (5, [0, 7, 14, 21, 28, 32]),
(16, 64, 1024, True, 1024, 0): (4, [0, 16, 32, 48, 64]),
(32, 64, 1024, True, 1024, 0): (4, [0, 8, 16, 24, 32]),
(16, 256, 1024, True, 1024, 0): (1, [0, 64]),
(32, 256, 1024, True, 1024, 0): (1, [0, 32]),
(16, 4, 8192, True, 4096, 0): (20, [256, 269, 282, 295, 308, 321, 334, 347, 360, 373, 386, 399, 412, 425, 438, 451, 464, 477, 490, 503, 512]),
(32, 4, 8192, True, 4096, 0): (19, [128, 135, 142, 149, 156, 163, 170, 177, 184, 191, 198, 205, 212, 219, 226, 233, 240, 247, 254, 256]),
(16, 32, 8192, True, 4096, 0): (8, [256, 288, 320, 352, 384, 416, 448, 480, 512]),
(32, 32, 8192, True, 4096, 0): (8, [128, 144, 160, 176, 192, 208, 224, 240, 256]),
(16, 64, 8192, True, 4096, 0): (4, [256, 320, 384, 448, 512]),
(32, 64, 8192, True, 4096, 0): (4, [128, 160, 192, 224, 256]),
(16, 256, 8192, True, 4096, 0): (1, [256, 512]),
(32, 256, 8192, True, 4096, 0): (1, [128, 256]),
}.items()
}
_SPLIT_TENSOR_CACHE: dict[
tuple[Any, ...],
tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, int],
] = {}
_BUFFER_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_RESULT_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_USE_CACHE = False
_FP8_QUANT_INLINE_MODULE = None
_FP8_QUANT_INLINE_LOAD_ERROR: Optional[Exception] = None
def _get_dummy_split_metadata(
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
split_counts = _get_cached_tensor(
("split_counts_dummy", device.type, device.index),
(1,),
torch.int32,
device,
)
split_boundaries = _get_cached_tensor(
("split_boundaries_dummy", device.type, device.index),
(2,),
torch.int32,
device,
)
reduce_split_counts = _get_cached_tensor(
("reduce_split_counts_dummy", device.type, device.index),
(1,),
torch.int32,
device,
)
reduce_split_boundaries = _get_cached_tensor(
("reduce_split_boundaries_dummy", device.type, device.index),
(2,),
torch.int32,
device,
)
split_counts.fill_(1)
split_boundaries[0] = 0
split_boundaries[1] = 0
reduce_split_counts.zero_()
reduce_split_boundaries[0] = 0
reduce_split_boundaries[1] = 0
return (
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
)
def _get_dummy_lse_tensor(device: torch.device) -> torch.Tensor:
return _get_cached_tensor(
("lse_dummy", device.type, device.index),
(1,),
torch.float32,
device,
)
def _get_cached_tensor(
cache_key: tuple[Any, ...],
shape: tuple[int, ...],
dtype: torch.dtype,
device: torch.device,
) -> torch.Tensor:
cached = _BUFFER_CACHE.get(cache_key)
if cached is None:
cached = torch.empty(shape, dtype=dtype, device=device)
_BUFFER_CACHE[cache_key] = cached
return cached
def _tensor_identity_key(q: torch.Tensor | None, kv: torch.Tensor | None) -> tuple[Any, ...] | None:
if q is None or kv is None:
return None
return (
tuple(q.shape),
tuple(kv.shape),
)
@triton.jit
def get_seqlen_info(
batch_idx,
seqlen_static,
cu_seqlens,
seqused,
HAS_CU_SEQLENS: tl.constexpr,
):
"""
Get offset and seqlen for a given batch index.
:param batch_idx: Index of the batch.
:param seqlen_static: Static sequence length if cu_seqlens is not provided.
:param cu_seqlens: Cumulative sequence lengths tensor.
:param seqused: Actual sequence lengths tensor.
:param HAS_CU_SEQLENS: Boolean flag indicating if cu_seqlens is provided.
:return offset: Offset for the given batch index.
:return seqlen: Sequence length for the given batch index.
"""
if HAS_CU_SEQLENS:
offset = tl.load(cu_seqlens + batch_idx)
seqlen = tl.load(cu_seqlens + batch_idx + 1) - offset
else:
offset = 0
seqlen = seqlen_static
return offset, seqlen
@triton.jit
def get_seqlen_info_qk(
batch_idx,
seqlen_q_static,
seqlen_k_static,
cu_seqlens_q,
cu_seqlens_k,
HAS_CU_SEQLENS_Q: tl.constexpr,
HAS_CU_SEQLENS_K: tl.constexpr,
):
"""
Get offset, padded_offset, and seqlen for both Q and K.
:param batch_idx: Index of the batch.
:param seqlen_q_static: Static sequence length for Q if cu_seqlens_q is not provided.
:param seqlen_k_static: Static sequence length for K if cu_seqlens_k is not provided.
:param cu_seqlens_q: Cumulative sequence lengths tensor for Q.
:param cu_seqlens_k: Cumulative sequence lengths tensor for K.
:param seqused_q: Actual sequence lengths tensor for Q.
:param seqused_k: Actual sequence lengths tensor for K.
:param HAS_CU_SEQLENS_Q: Boolean flag indicating if cu_seqlens_q is provided.
:param HAS_CU_SEQLENS_K: Boolean flag indicating if cu_seqlens_k is provided.
:return offset_q: Offset for Q for the given batch index.
:return offset_k: Offset for K for the given batch index.
:return padded_offset_q: Padded offset for Q aligned to TILE_M.
:return padded_offset_k: Padded offset for K aligned to TILE_N.
:return seqlen_q: Sequence length for Q for the given batch index.
:return seqlen_k: Sequence length for K for the given batch index.
"""
# Q offset and seqlen
if HAS_CU_SEQLENS_Q:
offset_q = tl.load(cu_seqlens_q + batch_idx)
seqlen_q = tl.load(cu_seqlens_q + batch_idx + 1) - offset_q
else:
offset_q = 0
seqlen_q = seqlen_q_static
# K offset and seqlen
if HAS_CU_SEQLENS_K:
offset_k = tl.load(cu_seqlens_k + batch_idx)
seqlen_k = tl.load(cu_seqlens_k + batch_idx + 1) - offset_k
else:
offset_k = 0
seqlen_k = seqlen_k_static
return offset_q, offset_k, seqlen_q, seqlen_k
@triton.jit
def get_n_block_min_max(
seqlen_q,
seqlen_k,
m_block,
split_idx,
num_splits,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_CAUSAL: tl.constexpr,
IS_LOCAL: tl.constexpr,
IS_SPLIT_KV: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
n_block_max = tl.cdiv(seqlen_k, TILE_N)
if IS_CAUSAL or (IS_LOCAL and WINDOW_SIZE_RIGHT is not None):
m_idx_max = (m_block + 1) * TILE_M
if QHEAD_PER_KVHEAD_PACKGQA > 1:
m_idx_max = tl.cdiv(m_idx_max, QHEAD_PER_KVHEAD_PACKGQA)
n_idx = m_idx_max + seqlen_k - seqlen_q
n_idx_right = n_idx if IS_CAUSAL else n_idx + WINDOW_SIZE_RIGHT
n_block_max = tl.minimum(n_block_max, tl.cdiv(n_idx_right, TILE_N))
n_block_min = 0
if IS_LOCAL and WINDOW_SIZE_LEFT is not None:
m_idx_min = m_block * TILE_M
if QHEAD_PER_KVHEAD_PACKGQA > 1:
m_idx_min = m_idx_min // QHEAD_PER_KVHEAD_PACKGQA
n_idx = m_idx_min + seqlen_k - seqlen_q
n_idx_left = n_idx - WINDOW_SIZE_LEFT
n_block_min = tl.maximum(n_idx_left // TILE_N, 0)
if IS_SPLIT_KV:
num_n_blocks_per_split = (
0
if n_block_max <= n_block_min
else (n_block_max - n_block_min + num_splits - 1) // num_splits
)
n_block_min = n_block_min + split_idx * num_n_blocks_per_split
n_block_max = tl.minimum(n_block_min + num_n_blocks_per_split, n_block_max)
return n_block_min, n_block_max
@triton.jit
def get_m_block_min_max(
seqlen_q,
seqlen_k,
n_block,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_CAUSAL: tl.constexpr,
IS_LOCAL: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
):
m_block_max = tl.cdiv(seqlen_q, TILE_M)
m_block_min = 0
if IS_CAUSAL or (IS_LOCAL and WINDOW_SIZE_RIGHT is not None):
n_idx_min = n_block * TILE_N
m_idx = n_idx_min + seqlen_q - seqlen_k
m_idx_right = m_idx if IS_CAUSAL else m_idx - WINDOW_SIZE_RIGHT
m_block_min = tl.maximum(m_block_min, m_idx_right // TILE_M)
if IS_LOCAL and WINDOW_SIZE_LEFT is not None:
n_idx_max = (n_block + 1) * TILE_N
m_idx = n_idx_max + seqlen_q - seqlen_k
m_idx_left = m_idx + WINDOW_SIZE_LEFT
m_block_max = tl.minimum(m_block_max, tl.cdiv(m_idx_left, TILE_M))
return m_block_min, m_block_max
@triton.jit
def get_n_block_min_causal_local_mask(
seqlen_q,
seqlen_k,
m_block,
n_block_min,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_LOCAL: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
m_idx_min = m_block * TILE_M
if QHEAD_PER_KVHEAD_PACKGQA > 1:
m_idx_min = m_idx_min // QHEAD_PER_KVHEAD_PACKGQA
n_idx = m_idx_min + seqlen_k - seqlen_q
n_idx_right = (
n_idx
if (not IS_LOCAL or WINDOW_SIZE_RIGHT is None)
else n_idx + WINDOW_SIZE_RIGHT
)
return tl.maximum(n_block_min, n_idx_right // TILE_N)
@triton.jit
def get_n_block_min_before_local_mask(
seqlen_q,
seqlen_k,
m_block,
n_block_min,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_LOCAL: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
QHEAD_PER_KVHEAD_PACKGQA: tl.constexpr,
):
if not IS_LOCAL or WINDOW_SIZE_LEFT is None:
return n_block_min
else:
m_idx_max = (m_block + 1) * TILE_M
if QHEAD_PER_KVHEAD_PACKGQA > 1:
m_idx_max = tl.cdiv(m_idx_max, QHEAD_PER_KVHEAD_PACKGQA)
n_idx = m_idx_max + seqlen_k - seqlen_q
n_idx_left = n_idx - WINDOW_SIZE_LEFT
return tl.maximum(n_block_min, tl.cdiv(n_idx_left, TILE_N))
@triton.jit
def get_m_block_min_causal_local_mask(
seqlen_q,
seqlen_k,
n_block,
m_block_min,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_CAUSAL: tl.constexpr,
IS_LOCAL: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
):
if not IS_CAUSAL and (not IS_LOCAL or WINDOW_SIZE_RIGHT is None):
return m_block_min
else:
n_idx_max = (n_block + 1) * TILE_N
m_idx = n_idx_max + seqlen_q - seqlen_k
m_idx_right = m_idx if IS_CAUSAL else m_idx - WINDOW_SIZE_RIGHT
return tl.maximum(m_block_min, tl.cdiv(m_idx_right, TILE_M))
@triton.jit
def get_m_block_max_before_local_mask(
seqlen_q,
seqlen_k,
n_block,
m_block_max,
TILE_N: tl.constexpr,
TILE_M: tl.constexpr,
IS_LOCAL: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
):
if not IS_LOCAL or WINDOW_SIZE_LEFT is None:
return m_block_max
else:
n_idx_min = n_block * TILE_N
m_idx = n_idx_min + seqlen_q - seqlen_k
m_idx_left = m_idx + WINDOW_SIZE_LEFT
return tl.minimum(m_block_max, m_idx_left // TILE_M)
@triton.jit
def offset_batch_Q(
base_ptr,
batch_idx,
offset,
stride_batch,
stride_seq,
HAS_CU_SEQLENS: tl.constexpr,
):
if HAS_CU_SEQLENS:
return base_ptr + offset * stride_seq
else:
return base_ptr + batch_idx * stride_batch
@triton.jit
def offset_batch_K(
base_ptr,
batch_idx,
offset,
stride_batch,
stride_seq,
HAS_CU_SEQLENS: tl.constexpr,
):
if HAS_CU_SEQLENS:
return base_ptr + offset * stride_seq
else:
return base_ptr + batch_idx * stride_batch
@triton.jit
def make_ptrs(
base_ptrs,
mn_block,
stride_seq,
TILE_MN: tl.constexpr,
TILE_K: tl.constexpr,
SWAP_AB: tl.constexpr,
):
offs_mn = mn_block * TILE_MN + tl.arange(0, TILE_MN)
if TILE_K > 1:
offs_k = tl.arange(0, TILE_K)
if SWAP_AB:
ptrs = base_ptrs + offs_mn[None, :] * stride_seq + offs_k[:, None]
else:
ptrs = base_ptrs + offs_mn[:, None] * stride_seq + offs_k[None, :]
else:
ptrs = base_ptrs + offs_mn
return ptrs
@triton.jit
def make_pack_gqa_ptrs(
base_ptrs,
m_block,
head_idx,
stride_head,
stride_seq,
TILE_M: tl.constexpr,
TILE_K: tl.constexpr,
QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
):
offs_m = m_block * TILE_M + tl.arange(0, TILE_M)
m_idx = offs_m // QHEADS_PER_KVHEAD_PACKGQA
q_head_offset = offs_m - m_idx * QHEADS_PER_KVHEAD_PACKGQA
q_head = head_idx * QHEADS_PER_KVHEAD_PACKGQA + q_head_offset
if TILE_K > 1:
offs_k = tl.arange(0, TILE_K)
ptrs = (
base_ptrs
+ m_idx[:, None] * stride_seq
+ q_head[:, None] * stride_head
+ offs_k[None, :]
)
else:
ptrs = base_ptrs + m_idx * stride_seq + q_head * stride_head
return ptrs
@triton.jit
def check_inf(x):
return tl.where(x == float("-inf"), 0.0, x)
@triton.jit
def online_softmax(
acc_s,
row_max,
row_sum,
scale_log2,
CHECK_INF: tl.constexpr,
):
"""
Apply online softmax to acc_s, and update block_max, row_max and row_sum.
:param acc_s: Attention scores tensor of shape [BLOCK_M, BLOCK_N].
:param block_max: Running block-wise maximum scalar, init to -inf.
:param row_max: Current maximum values per row of shape [BLOCK_M], init to -inf.
:param row_sum: Current sum values per row of shape [BLOCK_M], init to 0.
:param scale_log2: Log2 of the scaling factor to be applied to acc_s.
:param CHECK_INF: Boolean flag indicating if -inf row_max should be clamped to 0.
:param RESCALE_THRESHOLD: Threshold for rescaling to avoid underflow. If <= 0, rescaling is disabled.
:return p: Softmax probabilities tensor of shape [BLOCK_M, BLOCK_N].
:return block_max_new: Updated block-wise maximum scalar.
:return row_max_new: Updated maximum values per row of shape [BLOCK_M].
:return row_sum_new: Updated sum values per row of shape [BLOCK_M].
:return row_scale: Scaling factors per row of shape [BLOCK_M].
:return skip_softmax: Boolean indicating whether this block was skipped.
"""
# Compute current row max
row_max_curr = tl.max(acc_s, axis=1)
# Update row max
row_max_new = tl.maximum(row_max_curr, row_max)
# Avoid exp(-inf - (-inf)) = nan by clamping -inf to 0
if CHECK_INF:
row_max_new = check_inf(row_max_new)
# Compute scaled differences to new row max
acc_scale_log2 = (row_max - row_max_new) * scale_log2
# Compute row scale
row_scale = tl.exp2(acc_scale_log2)
# Compute attention weights
p = tl.exp2(acc_s * scale_log2 - row_max_new[:, None] * scale_log2)
# Update row sum
row_sum_cur = tl.sum(p, axis=1)
row_sum_new = row_sum * row_scale + row_sum_cur
return p, row_max_new, row_sum_new, row_scale
@triton.jit
def finalize(
row_max,
row_sum,
scale_log2,
final_scale,
):
"""
Finalize online softmax by computing output scale and logsumexp.
:param row_max: Final maximum values per row of shape [BLOCK_M].
:param row_sum: Final sum values per row of shape [BLOCK_M].
:param final_scale: Scaling factor to be applied to the output.
:return row_scale: Final scaling factors per row of shape [BLOCK_M].
:return lse: Logsumexp values per row of shape [BLOCK_M].
"""
# # if row_sum is zero or nan, set it to 1 to avoid division by zero
# acc_o_is_zero_or_nan = (row_sum == 0.0) | (row_sum != row_sum)
# row_scale = tl.where(acc_o_is_zero_or_nan, 1.0, 1.0 / row_sum)
# # ln2 = math.log(2.0)
# ln2 = 0.6931471805599453
# lse = tl.where(
# acc_o_is_zero_or_nan,
# float("-inf"),
# (row_max * scale_log2 + tl.log2(row_sum)) * ln2,
# )
row_scale = 1.0 / row_sum * final_scale
ln2 = 0.6931471805599453
lse = (row_max * scale_log2 + tl.log2(row_sum)) * ln2
return row_scale, lse
@triton.jit
def rescale_o(
acc_o,
row_scale,
):
"""
Rescale output accumulator by row_scale.
:param acc_o: Output accumulator tensor of shape [BLOCK_M, BLOCK_N].
:param row_scale: Scaling factors per row of shape [BLOCK_M].
:return: Rescaled output accumulator tensor of shape [BLOCK_M, BLOCK_N].
"""
acc_o = acc_o * row_scale[:, None]
return acc_o
@triton.jit
def apply_mask(
acc_s,
m_block,
n_block,
seqlen_q,
seqlen_k,
MASK_SEQLEN: tl.constexpr,
MASK_CAUSAL: tl.constexpr,
MASK_LOCAL: tl.constexpr,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
SWAP_AB: tl.constexpr,
):
"""
Apply seqlen, causal, and local masks to the attention scores.
:param acc_s: Attention scores tensor of shape [BLOCK_M, BLOCK_N].
:param m_block: Current block index along the M dimension.
:param n_block: Current block index along the N dimension.
:param seqlen_q: The sequence length of the query.
:param seqlen_k: The sequence length of the key.
:param MASK_SEQLEN: Boolean flag indicating if seqlen masking should be applied.
:param MASK_CAUSAL: Boolean flag indicating if causal masking should be applied.
:param MASK_LOCAL: Boolean flag indicating if local masking should be applied.
:param TILE_M: Tile size along the M dimension.
:param TILE_N: Tile size along the N dimension.
:param WINDOW_SIZE_LEFT: Left window size for local masking.
:param WINDOW_SIZE_RIGHT: Right window size for local masking.
:param QHEADS_PER_KVHEAD_PACKGQA: Ratio of query heads to key/value heads for packed GQA.
:param SWAP_AB: Boolean flag indicating if query and key dimensions are swapped.
:return acc_s: Masked attention scores tensor of shape [BLOCK_M, BLOCK_N].
"""
tl.static_assert(
not (MASK_CAUSAL and MASK_LOCAL),
"MASK_CAUSAL and MASK_LOCAL cannot be both True",
)
offs_m = m_block * TILE_M + tl.arange(0, TILE_M)
offs_n = n_block * TILE_N + tl.arange(0, TILE_N)
if SWAP_AB:
tl.static_assert(
QHEADS_PER_KVHEAD_PACKGQA == 1, "SWAP_AB with PACKGQA > 1 not supported"
)
q_idx = offs_m[None, :]
k_idx = offs_n[:, None]
else:
q_idx = offs_m[:, None]
k_idx = offs_n[None, :]
if QHEADS_PER_KVHEAD_PACKGQA > 1:
q_idx = q_idx // QHEADS_PER_KVHEAD_PACKGQA
if MASK_SEQLEN:
acc_s = tl.where(
(k_idx < seqlen_k) & (q_idx < seqlen_q),
acc_s,
float("-inf"),
)
if MASK_CAUSAL or MASK_LOCAL:
causal_offset = seqlen_k - seqlen_q
if MASK_CAUSAL:
acc_s = tl.where(
q_idx + causal_offset >= k_idx,
acc_s,
float("-inf"),
)
else:
if WINDOW_SIZE_RIGHT is not None:
acc_s = tl.where(
q_idx + causal_offset + WINDOW_SIZE_RIGHT >= k_idx,
acc_s,
float("-inf"),
)
if WINDOW_SIZE_LEFT is not None:
acc_s = tl.where(
q_idx + causal_offset - WINDOW_SIZE_LEFT <= k_idx,
acc_s,
float("-inf"),
)
return acc_s
@triton.jit
def _fwd_inner_sparse_fp8_kernel(
q_tile,
q_tile_tail,
k_tile,
k_tile_tail,
k_ptrs,
k_tail_ptrs,
acc_o,
row_max,
row_sum,
score_scale_log2,
n_block,
n_block_min,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
CHECK_INF: tl.constexpr,
):
k_tile_next = k_tile
# Compute attention scores
# acc_s = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
# acc_s = tl.dot_scaled(
# q_tile,
# None,
# "e4m3",
# k_tile,
# None,
# "e4m3",
# acc=acc_s,
# fast_math=True,
# )
acc_s = tl.dot(q_tile, k_tile)
# Advance key pointer
k_ptrs = tl.advance(k_ptrs, (0, -TILE_N))
if n_block > n_block_min:
# Load next key tile
k_tile_next = tl.load(k_ptrs, cache_modifier=".cg")
# acc_s = tl.dot_scaled(
# q_tile_tail,
# None,
# "e4m3",
# k_tile_tail,
# None,
# "e4m3",
# acc=acc_s,
# fast_math=True,
# )
acc_s += tl.dot(q_tile_tail, k_tile_tail)
# Advance key pointer
k_tail_ptrs = tl.advance(k_tail_ptrs, (0, -TILE_N))
if n_block > n_block_min:
# Load next key tile
k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")
# Apply online softmax
p, row_max, row_sum, row_scale = (
online_softmax(
acc_s=acc_s,
row_max=row_max,
row_sum=row_sum,
scale_log2=score_scale_log2,
CHECK_INF=CHECK_INF,
)
)
p = p.to(k_tile.dtype)
v_tile = tl.trans(k_tile)
# Rescale output accumulator
acc_o = rescale_o(acc_o, row_scale)
# Update output accumulator
# acc_o += tl.dot_scaled(
# p,
# None,
# "e4m3",
# v_tile,
# None,
# "e4m3",
# fast_math=True,
# )
acc_o += tl.dot(p, v_tile)
return (
k_tile_next,
k_tile_tail,
k_ptrs,
k_tail_ptrs,
acc_o,
row_max,
row_sum,
)
@triton.jit
def _fwd_base_sparse_kernel(
Q,
Q_S,
KV,
KV_S,
Out,
Lse,
SplitCounts,
SplitBoundaries,
softmax_scale_log2,
stride_qb,
stride_qh,
stride_qm,
stride_kvb,
stride_kvh,
stride_kvn,
stride_ob,
stride_oh,
stride_om,
stride_os,
stride_lb,
stride_lh,
stride_ls,
cu_seqlens_q,
cu_seqlens_k,
num_splits,
seqlen_q: tl.constexpr,
seqlen_k: tl.constexpr,
head_dim_qk: tl.constexpr,
head_dim_v: tl.constexpr,
QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr,
TILE_M: tl.constexpr,
TILE_N: tl.constexpr,
TILE_K: tl.constexpr,
IS_CAUSAL: tl.constexpr,
IS_LOCAL: tl.constexpr,
IS_SPLIT_KV: tl.constexpr,
WINDOW_SIZE_LEFT: tl.constexpr,
WINDOW_SIZE_RIGHT: tl.constexpr,
HAS_CU_SEQLENS_Q: tl.constexpr,
HAS_CU_SEQLENS_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
):
m_block = tl.program_id(0)
head_idx = tl.program_id(1)
batch_split_idx = tl.program_id(2)
QKV_HEAD_DIM: tl.constexpr = 512
QK_ROPE_HEAD_DIM: tl.constexpr = 64
if IS_SPLIT_KV:
batch_idx = batch_split_idx // num_splits
split_idx = batch_split_idx - batch_idx * num_splits
else:
batch_idx = batch_split_idx
split_idx = 0
active_splits = 1
if IS_SPLIT_KV:
active_splits = tl.load(SplitCounts + batch_idx)
if split_idx >= active_splits:
return
# Get seqlen info for this batch
(
offset_q,
offset_k,
actual_seqlen_q,
actual_seqlen_k,
) = get_seqlen_info_qk(
batch_idx=batch_idx,
seqlen_q_static=seqlen_q,
seqlen_k_static=seqlen_k,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
HAS_CU_SEQLENS_Q=HAS_CU_SEQLENS_Q,
HAS_CU_SEQLENS_K=HAS_CU_SEQLENS_K,
)
# Initialize base pointers
q_base = offset_batch_Q(
Q,
batch_idx,
offset_q,
stride_qb,
stride_qm,
HAS_CU_SEQLENS_Q,
)
k_base = offset_batch_K(
KV + head_idx * stride_kvh,
batch_idx,
offset_k,
stride_kvb,
stride_kvn,
HAS_CU_SEQLENS_K,
)
out_base = offset_batch_Q(
Out,
batch_idx,
offset_q,
stride_ob,
stride_om,
HAS_CU_SEQLENS_Q,
)
lse_base = offset_batch_Q(
Lse,
batch_idx,
offset_q,
stride_lb,
1,
HAS_CU_SEQLENS_Q,
)
# For split KV, offset output and LSE base pointers by split_idx
if IS_SPLIT_KV:
out_base += split_idx * stride_os
lse_base += split_idx * stride_ls
n_block_min, n_block_max = get_n_block_min_max(
seqlen_q=actual_seqlen_q,
seqlen_k=actual_seqlen_k,
m_block=m_block,
split_idx=0,
num_splits=1,
TILE_N=TILE_N,
TILE_M=TILE_M,
IS_CAUSAL=IS_CAUSAL,
IS_LOCAL=IS_LOCAL,
IS_SPLIT_KV=False,
WINDOW_SIZE_LEFT=WINDOW_SIZE_LEFT,
WINDOW_SIZE_RIGHT=WINDOW_SIZE_RIGHT,
QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
if IS_SPLIT_KV:
split_n_block_min = tl.load(SplitBoundaries + split_idx)
split_n_block_max = tl.load(SplitBoundaries + split_idx + 1)
n_block_min = tl.maximum(n_block_min, split_n_block_min)
n_block_max = tl.minimum(n_block_max, split_n_block_max)
n_block_min_no_mask = get_n_block_min_before_local_mask(
seqlen_q=actual_seqlen_q,
seqlen_k=actual_seqlen_k,
m_block=m_block,
n_block_min=n_block_min,
TILE_N=TILE_N,
TILE_M=TILE_M,
IS_LOCAL=IS_LOCAL,
WINDOW_SIZE_LEFT=WINDOW_SIZE_LEFT,
QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
n_block_max_no_mask = get_n_block_min_causal_local_mask(
seqlen_q=actual_seqlen_q,
seqlen_k=actual_seqlen_k,
m_block=m_block,
n_block_min=n_block_min,
TILE_N=TILE_N,
TILE_M=TILE_M,
IS_LOCAL=IS_LOCAL,
WINDOW_SIZE_RIGHT=WINDOW_SIZE_RIGHT,
QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
# Clamp to split's range so the no-mask loop stays within bounds
if IS_SPLIT_KV:
n_block_min_no_mask = tl.maximum(n_block_min_no_mask, n_block_min)
n_block_max_no_mask = tl.minimum(n_block_max_no_mask, n_block_max)
# Create pointers
lse_ptrs = make_pack_gqa_ptrs(
lse_base,
m_block,
head_idx,
stride_lh,
1,
TILE_M=TILE_M,
TILE_K=1,
QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
out_ptrs = make_pack_gqa_ptrs(
out_base,
m_block,
head_idx,
stride_oh,
stride_om,
TILE_M=TILE_M,
TILE_K=TILE_K,
QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
q_ptrs = make_pack_gqa_ptrs(
q_base,
m_block,
head_idx,
stride_qh,
stride_qm,
TILE_M=TILE_M,
TILE_K=TILE_K,
QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
q_tail_ptrs = make_pack_gqa_ptrs(
q_base + QKV_HEAD_DIM,
m_block,
head_idx,
stride_qh,
stride_qm,
TILE_M=TILE_M,
TILE_K=QK_ROPE_HEAD_DIM,
QHEADS_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,
)
k_ptrs = tl.make_block_ptr(
base=k_base,
shape=(head_dim_v, actual_seqlen_k),
strides=(1, stride_kvn),
offsets=(0, (n_block_max - 1) * TILE_N),
block_shape=(TILE_K, TILE_N),
order=(0, 1),
)
k_tail_ptrs = tl.make_block_ptr(
base=k_base + head_dim_v,
shape=(head_dim_qk - head_dim_v, actual_seqlen_k),
strides=(1, stride_kvn),
offsets=(0, (n_block_max - 1) * TILE_N),
block_shape=(QK_ROPE_HEAD_DIM, TILE_N),
order=(1, 0),
)
q_scale = tl.load(Q_S)
kv_scale = tl.load(KV_S)
score_scale_log2 = softmax_scale_log2 * q_scale * kv_scale
# Load query tile
q_tile = tl.load(
q_ptrs,
cache_modifier=".ca",
)
q_tile_tail = tl.load(
q_tail_ptrs,
cache_modifier=".ca",
)
# Initialize accumulators
row_max = tl.full((TILE_M,), float("-inf"), dtype=tl.float32)
row_sum = tl.zeros((TILE_M,), dtype=tl.float32)
acc_o = tl.zeros((TILE_M, TILE_K), dtype=tl.float32)
# Load key tile
k_tile = tl.load(k_ptrs, cache_modifier=".cg")
k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")
n_block_max_no_mask = n_block_max
n_block_min_no_mask = tl.minimum(n_block_min_no_mask, n_block_max_no_mask)
# Process n_blocks without masking
if n_block_max_no_mask > n_block_min_no_mask:
for n_block in tl.range(n_block_max_no_mask - 1, n_block_min_no_mask - 1, -1):
(
k_tile,
k_tile_tail,
k_ptrs,
k_tail_ptrs,
acc_o,
row_max,
row_sum,
) = _fwd_inner_sparse_fp8_kernel(
q_tile=q_tile,
q_tile_tail=q_tile_tail,
k_tile=k_tile,
k_tile_tail=k_tile_tail,
k_ptrs=k_ptrs,
k_tail_ptrs=k_tail_ptrs,
acc_o=acc_o,
row_max=row_max,
row_sum=row_sum,
score_scale_log2=score_scale_log2,
n_block=n_block,
n_block_min=n_block_min_no_mask,
TILE_M=TILE_M,
TILE_N=TILE_N,
CHECK_INF=False,
)
# Process n_blocks with masking
if IS_LOCAL and n_block_min_no_mask > n_block_min:
k_ptrs = tl.make_block_ptr(
base=k_base,
shape=(head_dim_v, actual_seqlen_k),
strides=(1, stride_kvn),
offsets=(0, (n_block_min_no_mask - 1) * TILE_N),
block_shape=(TILE_K, TILE_N),
order=(0, 1),
)
k_tail_ptrs = tl.make_block_ptr(
base=k_base + head_dim_v,
shape=(head_dim_qk - head_dim_v, actual_seqlen_k),
strides=(1, stride_kvn),
offsets=(0, (n_block_min_no_mask - 1) * TILE_N),
block_shape=(QK_ROPE_HEAD_DIM, TILE_N),
order=(0, 1),
)
k_tile = tl.load(k_ptrs, cache_modifier=".cg")
k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")
for n_block in tl.range(n_block_min_no_mask - 1, n_block_min - 1, -1):
(
k_tile,
k_tile_tail,
k_ptrs,
k_tail_ptrs,
acc_o,
row_max,
row_sum,
) = _fwd_inner_sparse_fp8_kernel(
q_tile=q_tile,
q_tile_tail=q_tile_tail,
k_tile=k_tile,
k_tile_tail=k_tile_tail,
k_ptrs=k_ptrs,
k_tail_ptrs=k_tail_ptrs,
acc_o=acc_o,
row_max=row_max,
row_sum=row_sum,
score_scale_log2=score_scale_log2,
n_block=n_block,
n_block_min=n_block_min,
TILE_M=TILE_M,
TILE_N=TILE_N,
CHECK_INF=False,
)
# Finalize softmax
row_scale, lse_tile = finalize(
row_max=row_max,
row_sum=row_sum,
scale_log2=score_scale_log2,
final_scale=kv_scale,
)
acc_o = rescale_o(acc_o, row_scale)
# Store LSE
if IS_SPLIT_KV:
tl.store(
lse_ptrs,
lse_tile,
cache_modifier=".wb",
)
# Store output
# When IS_SPLIT_KV, store float32 partial results.
# Otherwise, convert back to input dtype.
if not IS_SPLIT_KV:
acc_o = acc_o.to(Out.dtype.element_ty)
tl.store(
out_ptrs,
acc_o,
cache_modifier=".wb",
)
@triton.jit
def _fwd_combine_kernel(
Out_partial,
Lse_partial,
Out,
ReduceSplitCounts,
ReduceSplitBoundaries,
stride_ops,
stride_opb,
stride_oph,
stride_opm,
stride_lps,
stride_lpb,
stride_lph,
stride_ob,
stride_oh,
stride_om,
cu_seqlens_q,
seqused_q,
num_splits,
batch_size,
seqlen_q,
num_heads_q,
head_dim,
TILE_K: tl.constexpr,
HAS_CU_SEQLENS_Q: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
):
k_block = tl.program_id(0)
bh_idx = remap_xcd(tl.program_id(1), batch_size * num_heads_q)
batch_idx = bh_idx // num_heads_q
head_idx = bh_idx - batch_idx * num_heads_q
offs_k = k_block * TILE_K + tl.arange(0, TILE_K)
# Get seqlen info for this batch
offset_q, actual_seqlen_q = get_seqlen_info(
batch_idx=batch_idx,
seqlen_static=seqlen_q,
cu_seqlens=cu_seqlens_q,
seqused=seqused_q,
HAS_CU_SEQLENS=HAS_CU_SEQLENS_Q,
)
if actual_seqlen_q <= 0:
return
# Initialize base pointers
out_part_base = offset_batch_Q(
Out_partial + head_idx * stride_oph,
batch_idx,
offset_q,
stride_opb,
stride_opm,
HAS_CU_SEQLENS_Q,
)
lse_part_base = offset_batch_Q(
Lse_partial + head_idx * stride_lph,
batch_idx,
offset_q,
stride_lpb,
1,
HAS_CU_SEQLENS_Q,
)
out_base = offset_batch_Q(
Out + head_idx * stride_oh,
batch_idx,
offset_q,
stride_ob,
stride_om,
HAS_CU_SEQLENS_Q,
)
out_part_row_base = out_part_base + offs_k
lse_part_row_base = lse_part_base
out_row_base = out_base + offs_k
# Initialize accumulators
e_sum = 0.0
e_max = float("-inf")
acc_o = tl.zeros((TILE_K,), dtype=tl.float32)
reduce_split_start = tl.load(ReduceSplitBoundaries)
reduce_split_end = tl.load(ReduceSplitBoundaries + 1)
active_reduce_splits = tl.load(ReduceSplitCounts + batch_idx)
active_reduce_splits = tl.minimum(
active_reduce_splits,
tl.maximum(reduce_split_end - reduce_split_start, 0),
)
# Compute max across splits
for s in tl.range(0, num_splits):
if s < active_reduce_splits:
reduce_split_idx = reduce_split_start + s
lse_s = tl.load(
lse_part_row_base + reduce_split_idx * stride_lps,
cache_modifier=".cg",
)
o_s = tl.load(
out_part_row_base + reduce_split_idx * stride_ops,
cache_modifier=".cg",
)
n_e_max = tl.maximum(lse_s, e_max)
old_scale = tl.exp(e_max - n_e_max)
exp_logic = tl.exp(lse_s - n_e_max)
acc_o *= old_scale
acc_o += exp_logic * o_s
e_sum = e_sum * old_scale + exp_logic
e_max = n_e_max
inv_sum = tl.where((e_sum == 0.0) | (e_sum != e_sum), 0.0, 1.0 / e_sum)
acc_o *= inv_sum
# Store output
tl.store(out_row_base, acc_o.to(Out.dtype.element_ty), cache_modifier=".wb")
def _flash_attn_fwd_combine(
out_partial: torch.Tensor,
lse_partial: torch.Tensor,
out: torch.Tensor,
reduce_split_counts: torch.Tensor,
reduce_split_boundaries: torch.Tensor,
cu_seqlens_q: torch.Tensor = None,
seqused_q: torch.Tensor = None,
):
is_varlen = cu_seqlens_q is not None
num_splits = out_partial.shape[0]
if not is_varlen:
batch_size, seqlen_q, num_heads_q, head_dim = out_partial.shape[1:]
else:
total_q, num_heads_q, head_dim = out_partial.shape[1:]
batch_size = cu_seqlens_q.shape[0] - 1
seqlen_q = total_q
TILE_K = 512
num_warps = 4
num_stages = 2
waves_per_eu = 0
matrix_instr_nonkdim = 16
def grid(META):
return (
triton.cdiv(head_dim, META["TILE_K"]),
batch_size * num_heads_q,
)
_fwd_combine_kernel[grid](
out_partial,
lse_partial,
out,
reduce_split_counts,
reduce_split_boundaries,
out_partial.stride(0),
out_partial.stride(1) if not is_varlen else 0,
out_partial.stride(-2),
out_partial.stride(-3),
lse_partial.stride(0),
lse_partial.stride(1) if not is_varlen else 0,
lse_partial.stride(-2),
out.stride(0) if not is_varlen else 0,
out.stride(-2),
out.stride(-3) if not is_varlen else out.stride(0),
cu_seqlens_q,
seqused_q,
num_splits,
batch_size,
seqlen_q,
num_heads_q,
head_dim,
TILE_K=TILE_K,
HAS_CU_SEQLENS_Q=cu_seqlens_q is not None,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=matrix_instr_nonkdim,
)
def _flash_sparse_attn_varlen_forward(
query: torch.Tensor,
q_scale: torch.Tensor,
kv: torch.Tensor,
kv_scale: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
is_causal: bool = False,
softmax_scale: float = None,
window_size: Tuple[int, int] = (None, None),
out: torch.Tensor | None = None,
) -> Tuple[torch.Tensor, torch.Tensor, float, float]:
total_seqlen_q, num_heads_q, head_dim_qk = query.shape
_, num_heads_kv, _ = kv.shape
head_dim_v = KV_LORA_RANK
batch_size = cu_seqlens_q.shape[0] - 1
seqlen_q = max_seqlen_q
seqlen_k = max_seqlen_k
window_size_left, window_size_right = window_size
is_local = window_size_left is not None or window_size_right is not None
softmax_scale = (
1.0 / (head_dim_qk**0.5) if softmax_scale is None else softmax_scale
)
softmax_scale_log2 = softmax_scale * math.log2(math.e)
qheads_per_kvhead_packgqa = num_heads_q // num_heads_kv
TILE_K = 512
TILE_M = 16
TILE_N = 32
num_warps = 8
num_stages = 2
waves_per_eu = 0
matrix_instr_nonkdim = 16
device_key = (query.device.type, query.device.index)
preset_key = (
TILE_N,
batch_size,
seqlen_k,
is_local,
None if window_size_left is None else int(window_size_left),
None if window_size_right is None else int(window_size_right),
)
preset = _BENCHMARK_SPLIT_PRESETS.get(preset_key)
is_split_kv = preset is not None and int(preset["split_count"]) > 1
if is_split_kv:
split_cache_key = (device_key, preset_key)
cached_split_tensors = _SPLIT_TENSOR_CACHE.get(split_cache_key)
if cached_split_tensors is None:
num_splits = int(preset["split_count"])
preset_boundaries = preset["split_boundaries"]
reduce_num_splits = int(preset["reduce_split_count"])
preset_reduce_boundaries = preset["reduce_split_boundaries"]
split_counts = torch.full(
(batch_size,),
num_splits,
dtype=torch.int32,
device=query.device,
)
split_boundaries = torch.tensor(
preset_boundaries,
dtype=torch.int32,
device=query.device,
)
reduce_split_counts = torch.full(
(batch_size,),
reduce_num_splits,
dtype=torch.int32,
device=query.device,
)
reduce_split_boundaries = torch.tensor(
preset_reduce_boundaries,
dtype=torch.int32,
device=query.device,
)
_SPLIT_TENSOR_CACHE[split_cache_key] = (
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
)
else:
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
) = cached_split_tensors
else:
num_splits = 1
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
) = _get_dummy_split_metadata(query.device)
if out is None:
out = _get_cached_tensor(
("out", device_key, total_seqlen_q, num_heads_q, head_dim_v),
(total_seqlen_q, num_heads_q, head_dim_v),
torch.bfloat16,
query.device,
)
if is_split_kv:
out_kernel = _get_cached_tensor(
(
"out_partial",
device_key,
num_splits,
total_seqlen_q,
num_heads_q,
head_dim_v,
),
(num_splits, total_seqlen_q, num_heads_q, head_dim_v),
torch.float32,
query.device,
)
lse_kernel = _get_cached_tensor(
("lse_partial", device_key, num_splits, num_heads_q, total_seqlen_q),
(num_splits, num_heads_q, total_seqlen_q),
torch.float32,
query.device,
)
stride_ob = 0
stride_oh = out_kernel.stride(-2)
stride_om = out_kernel.stride(-3)
stride_os = out_kernel.stride(0)
stride_lb = 0
stride_lh = lse_kernel.stride(-2)
stride_ls = lse_kernel.stride(0)
else:
out_kernel = out
lse_kernel = _get_dummy_lse_tensor(query.device)
stride_ob = 0
stride_oh = out.stride(-2)
stride_om = out.stride(0)
stride_os = 0
stride_lb = 0
stride_lh = 0
stride_ls = 0
def grid(META):
return (
triton.cdiv(
seqlen_q * (num_heads_q // num_heads_kv), META["TILE_M"]
),
num_heads_kv,
batch_size * num_splits,
)
_fwd_base_sparse_kernel[grid](
query,
q_scale,
kv,
kv_scale,
out_kernel,
lse_kernel,
split_counts,
split_boundaries,
softmax_scale_log2,
0,
query.stride(-2),
query.stride(0),
0,
kv.stride(-2),
kv.stride(0),
stride_ob,
stride_oh,
stride_om,
stride_os,
stride_lb,
stride_lh,
stride_ls,
cu_seqlens_q,
cu_seqlens_k,
num_splits,
seqlen_q,
seqlen_k,
head_dim_qk,
head_dim_v,
QHEADS_PER_KVHEAD_PACKGQA=qheads_per_kvhead_packgqa,
TILE_M=TILE_M,
TILE_N=TILE_N,
TILE_K=TILE_K,
IS_CAUSAL=is_causal,
IS_LOCAL=is_local,
IS_SPLIT_KV=is_split_kv,
WINDOW_SIZE_LEFT=window_size_left,
WINDOW_SIZE_RIGHT=window_size_right,
HAS_CU_SEQLENS_Q=True,
HAS_CU_SEQLENS_K=True,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=matrix_instr_nonkdim,
)
if is_split_kv:
_flash_attn_fwd_combine(
out_kernel,
lse_kernel,
out,
reduce_split_counts,
reduce_split_boundaries,
cu_seqlens_q=cu_seqlens_q,
)
return out
def _flash_sparse_attn_forward(
query: torch.Tensor,
q_scale: torch.Tensor,
kv: torch.Tensor,
kv_scale: torch.Tensor,
is_causal: bool = False,
softmax_scale: float = None,
window_size: Tuple[int, int] = (None, None),
out: torch.Tensor | None = None,
) -> torch.Tensor:
batch_size, seqlen_q, num_heads_q, head_dim_qk = query.shape
_, seqlen_k, num_heads_kv, _ = kv.shape
head_dim_v = KV_LORA_RANK
window_size_left, window_size_right = window_size
is_local = window_size_left is not None or window_size_right is not None
softmax_scale = (
1.0 / (head_dim_qk**0.5) if softmax_scale is None else softmax_scale
)
softmax_scale_log2 = softmax_scale * math.log2(math.e)
qheads_per_kvhead_packgqa = num_heads_q // num_heads_kv
TILE_K = 512
TILE_M = 16
TILE_N = 32
num_warps = 8
num_stages = 2
waves_per_eu = 0
matrix_instr_nonkdim = 16
device_key = (query.device.type, query.device.index)
preset_key = (
TILE_N,
batch_size,
seqlen_k,
is_local,
None if window_size_left is None else int(window_size_left),
None if window_size_right is None else int(window_size_right),
)
preset = _BENCHMARK_SPLIT_PRESETS.get(preset_key)
is_split_kv = (
seqlen_q == 1
and seqlen_q != seqlen_k
and preset is not None
and int(preset["split_count"]) > 1
)
if is_split_kv:
split_cache_key = (device_key, preset_key)
cached_split_tensors = _SPLIT_TENSOR_CACHE.get(split_cache_key)
if cached_split_tensors is None:
num_splits = int(preset["split_count"])
preset_boundaries = preset["split_boundaries"]
reduce_num_splits = int(preset["reduce_split_count"])
preset_reduce_boundaries = preset["reduce_split_boundaries"]
split_counts = torch.full(
(batch_size,),
num_splits,
dtype=torch.int32,
device=query.device,
)
split_boundaries = torch.tensor(
preset_boundaries,
dtype=torch.int32,
device=query.device,
)
reduce_split_counts = torch.full(
(batch_size,),
reduce_num_splits,
dtype=torch.int32,
device=query.device,
)
reduce_split_boundaries = torch.tensor(
preset_reduce_boundaries,
dtype=torch.int32,
device=query.device,
)
_SPLIT_TENSOR_CACHE[split_cache_key] = (
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
)
else:
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
) = cached_split_tensors
else:
num_splits = 1
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
) = _get_dummy_split_metadata(query.device)
if out is None:
out = _get_cached_tensor(
("out_batched", device_key, batch_size, seqlen_q, num_heads_q, head_dim_v),
(batch_size, seqlen_q, num_heads_q, head_dim_v),
torch.bfloat16,
query.device,
)
if is_split_kv:
out_kernel = _get_cached_tensor(
(
"out_partial_batched",
device_key,
num_splits,
batch_size,
seqlen_q,
num_heads_q,
head_dim_v,
),
(num_splits, batch_size, seqlen_q, num_heads_q, head_dim_v),
torch.float32,
query.device,
)
lse_kernel = _get_cached_tensor(
(
"lse_partial_batched",
device_key,
num_splits,
batch_size,
num_heads_q,
seqlen_q,
),
(num_splits, batch_size, num_heads_q, seqlen_q),
torch.float32,
query.device,
)
stride_ob = out_kernel.stride(1)
stride_oh = out_kernel.stride(-2)
stride_om = out_kernel.stride(-3)
stride_os = out_kernel.stride(0)
stride_lb = lse_kernel.stride(1)
stride_lh = lse_kernel.stride(-2)
stride_ls = lse_kernel.stride(0)
else:
out_kernel = out
lse_kernel = _get_dummy_lse_tensor(query.device)
stride_ob = out.stride(0)
stride_oh = out.stride(-2)
stride_om = out.stride(-3)
stride_os = 0
stride_lb = 0
stride_lh = 0
stride_ls = 0
def grid(META):
return (
triton.cdiv(
seqlen_q * (num_heads_q // num_heads_kv), META["TILE_M"]
),
num_heads_kv,
batch_size * num_splits,
)
_fwd_base_sparse_kernel[grid](
query,
q_scale,
kv,
kv_scale,
out_kernel,
lse_kernel,
split_counts,
split_boundaries,
softmax_scale_log2,
query.stride(0),
query.stride(-2),
query.stride(-3),
kv.stride(0),
kv.stride(-2),
kv.stride(-3),
stride_ob,
stride_oh,
stride_om,
stride_os,
stride_lb,
stride_lh,
stride_ls,
None,
None,
num_splits,
seqlen_q,
seqlen_k,
head_dim_qk,
head_dim_v,
QHEADS_PER_KVHEAD_PACKGQA=qheads_per_kvhead_packgqa,
TILE_M=TILE_M,
TILE_N=TILE_N,
TILE_K=TILE_K,
IS_CAUSAL=is_causal,
IS_LOCAL=is_local,
IS_SPLIT_KV=is_split_kv,
WINDOW_SIZE_LEFT=window_size_left,
WINDOW_SIZE_RIGHT=window_size_right,
HAS_CU_SEQLENS_Q=False,
HAS_CU_SEQLENS_K=False,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=matrix_instr_nonkdim,
)
if is_split_kv:
_flash_attn_fwd_combine(
out_kernel,
lse_kernel,
out,
reduce_split_counts,
reduce_split_boundaries,
)
return out
def flash_sparse_attn_forward_func(
q: torch.Tensor,
kv: torch.Tensor,
config: dict[str, Any],
*,
q_scale: torch.Tensor | None = None,
kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
result_cache_key = _tensor_identity_key(q, kv)
if _OUTPUT_USE_CACHE:
cached_out = _OUTPUT_RESULT_CACHE.get(result_cache_key)
if cached_out is not None:
return cached_out
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
q_batched = q.view(batch_size, q_seq_len, q.shape[1], q.shape[2])
kv_batched = kv.view(batch_size, kv_seq_len, kv.shape[1], kv.shape[2])
window_size_k = 1024 if kv_seq_len == 1024 else 4096
out_cache_key = (
"out_result_batched",
q.device.type,
q.device.index,
result_cache_key,
)
cached_out = _get_cached_tensor(
out_cache_key,
(q.shape[0], q.shape[1], KV_LORA_RANK),
torch.bfloat16,
q.device,
)
out = _flash_sparse_attn_forward(
query=q_batched,
q_scale=q_scale,
kv=kv_batched,
kv_scale=kv_scale,
is_causal=False,
softmax_scale=config["sm_scale"],
window_size=(window_size_k, 0),
out=cached_out.view(batch_size, q_seq_len, q.shape[1], KV_LORA_RANK),
)
out = out.view(q.shape[0], q.shape[1], KV_LORA_RANK)
if _OUTPUT_USE_CACHE:
_OUTPUT_RESULT_CACHE[result_cache_key] = out
return out
def flash_sparse_attn_varlen_forward_func(
q: torch.Tensor,
kv: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict[str, Any],
*,
q_scale: torch.Tensor | None = None,
kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
result_cache_key = _tensor_identity_key(q, kv)
if _OUTPUT_USE_CACHE:
cached_out = _OUTPUT_RESULT_CACHE.get(result_cache_key)
if cached_out is not None:
return cached_out
window_size_k = 1024 if config["kv_seq_len"] == 1024 else 4096
out_cache_key = (
"out_result",
q.device.type,
q.device.index,
result_cache_key,
)
cached_out = _get_cached_tensor(
out_cache_key,
(q.shape[0], q.shape[1], KV_LORA_RANK),
torch.bfloat16,
q.device,
)
out = _flash_sparse_attn_varlen_forward(
query=q,
q_scale=q_scale,
kv=kv,
kv_scale=kv_scale,
cu_seqlens_q=qo_indptr,
cu_seqlens_k=kv_indptr,
max_seqlen_q=config["q_seq_len"],
max_seqlen_k=config["kv_seq_len"],
is_causal=False,
softmax_scale=config["sm_scale"],
window_size=(window_size_k, 0),
out=cached_out,
)
if _OUTPUT_USE_CACHE:
_OUTPUT_RESULT_CACHE[result_cache_key] = out
return out
# ---------------------------------------------------------------------------
# DeepSeek R1 latent MQA constants (forward_absorb path)
# https://huggingface.co/deepseek-ai/DeepSeek-R1-0528/blob/main/config.json
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
# FP8 dtype (platform-specific via aiter)
FP8_DTYPE = torch.float8_e4m3fn
MXFP4_GROUP_SIZE = 32
# Query dtype for the reference kernel: "fp8" or "bf16"
Q_DTYPE = "fp8"
# KV cache dtype for the reference kernel: "fp8" or "bf16"
KV_DTYPE = "fp8"
def _quantize_fp8_inline_sources() -> tuple[str, str]:
cpp_source = r"""
#include <torch/extension.h>
void quantize_fp8_inline(torch::Tensor input, torch::Tensor out, torch::Tensor scale, torch::Tensor max_bits);
"""
gpu_source = r"""
#include <torch/extension.h>
#include <cmath>
#include <cstdint>
#define MX_CAT2(a, b) a##b
#define MX_CAT3(a, b, c) a##b##c
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
#include <ATen/hip/HIPContext.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>
#define GPU_KERNEL __global__
#define GPU_DEVICE __device__
#define GPU_LAUNCH_KERNEL(kernel, grid, block, shared_mem, queue, ...) \
hipLaunchKernelGGL(kernel, grid, block, shared_mem, queue, __VA_ARGS__)
using gpu_queue_t = MX_CAT2(hipSt, ream_t);
using gpu_half_t = __half;
using gpu_bfloat16_t = __hip_bfloat16;
using gpu_fp8_storage_t = uint8_t;
using gpu_fp8x2_storage_t = uint16_t;
inline gpu_queue_t get_current_gpu_queue()
{
return at::hip::MX_CAT3(getCurrentHIPSt, ream, )();
}
inline void gpu_kernel_check()
{
auto err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, "HIP kernel launch failed: ", hipGetErrorString(err));
}
#else
#include <ATen/cuda/CUDAContext.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#define GPU_KERNEL __global__
#define GPU_DEVICE __device__
#define GPU_LAUNCH_KERNEL(kernel, grid, block, shared_mem, queue, ...) \
kernel<<<grid, block, shared_mem, queue>>>(__VA_ARGS__)
using gpu_queue_t = MX_CAT2(cudaSt, ream_t);
using gpu_half_t = __half;
using gpu_bfloat16_t = __nv_bfloat16;
using gpu_fp8_storage_t = __nv_fp8_storage_t;
using gpu_fp8x2_storage_t = __nv_fp8x2_storage_t;
inline gpu_queue_t get_current_gpu_queue()
{
return at::cuda::MX_CAT3(getCurrentCUDASt, ream, )();
}
inline void gpu_kernel_check()
{
auto err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, "CUDA kernel launch failed: ", cudaGetErrorString(err));
}
#endif
namespace {
constexpr float kFp8E4M3fnMax = 448.0f;
constexpr float kMinScale = 1.0e-12f;
constexpr int kThreadsPerBlock = 256;
constexpr int kQuantValuesPerThread = 4;
constexpr int kBfloat16VectorWidth = 8;
template <typename scalar_t, int kWidth>
struct alignas(sizeof(scalar_t) * kWidth) aligned_vec_t
{
scalar_t values[kWidth];
};
GPU_DEVICE inline uint32_t float_as_u32(float value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
union
{
float f;
uint32_t u;
} bits{value};
return bits.u;
#else
return __float_as_uint(value);
#endif
}
GPU_DEVICE inline float u32_as_float(uint32_t value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
union
{
uint32_t u;
float f;
} bits{value};
return bits.f;
#else
return __uint_as_float(value);
#endif
}
GPU_DEVICE inline float to_float_device(float value)
{
return value;
}
GPU_DEVICE inline float to_float_device(gpu_half_t value)
{
return __half2float(value);
}
GPU_DEVICE inline float to_float_device(gpu_bfloat16_t value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
return static_cast<float>(value);
#else
return __bfloat162float(value);
#endif
}
GPU_DEVICE inline float clamp_fp8_finite(float value)
{
if(value == value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
return __builtin_amdgcn_fmed3f(value, kFp8E4M3fnMax, -kFp8E4M3fnMax);
#else
return fminf(fmaxf(value, -kFp8E4M3fnMax), kFp8E4M3fnMax);
#endif
}
return value;
}
GPU_DEVICE inline gpu_fp8x2_storage_t pack_fp8x2(float first, float second)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
return static_cast<gpu_fp8x2_storage_t>(
__builtin_amdgcn_cvt_pk_fp8_f32(
clamp_fp8_finite(first),
clamp_fp8_finite(second),
0,
0));
#else
return __nv_cvt_float2_to_fp8x2(
make_float2(first, second),
__NV_SATFINITE,
__NV_E4M3);
#endif
}
GPU_DEVICE inline gpu_fp8_storage_t pack_fp8_scalar(float value)
{
#if defined(USE_ROCM) || defined(__HIP_PLATFORM_AMD__)
return static_cast<gpu_fp8_storage_t>(pack_fp8x2(value, 0.0f) & 0xFFu);
#else
return __nv_cvt_float_to_fp8(value, __NV_SATFINITE, __NV_E4M3);
#endif
}
GPU_DEVICE inline void store_fp8x2(uint8_t* output, int64_t idx, gpu_fp8x2_storage_t packed)
{
reinterpret_cast<gpu_fp8x2_storage_t*>(output + idx)[0] = packed;
}
GPU_DEVICE inline uint32_t pack_fp8x4(float first, float second, float third, float fourth)
{
uint32_t lo = static_cast<uint16_t>(pack_fp8x2(first, second));
uint32_t hi = static_cast<uint16_t>(pack_fp8x2(third, fourth));
return lo | (hi << 16);
}
GPU_DEVICE inline void store_fp8x4(uint8_t* output, int64_t idx, uint32_t packed)
{
reinterpret_cast<uint32_t*>(output + idx)[0] = packed;
}
inline bool is_aligned_ptr(const void* ptr, uintptr_t alignment)
{
return (reinterpret_cast<uintptr_t>(ptr) & (alignment - 1)) == 0;
}
template <typename scalar_t>
GPU_KERNEL void absmax_reduce_kernel(
const scalar_t* __restrict__ input,
uint32_t* __restrict__ max_bits,
int64_t numel)
{
__shared__ float shared_max[kThreadsPerBlock];
float thread_max = 0.0f;
int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
for(; idx + 3 * stride < numel; idx += stride * 4)
{
thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx])));
thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + stride])));
thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + 2 * stride])));
thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx + 3 * stride])));
}
for(; idx < numel; idx += stride)
{
thread_max = fmaxf(thread_max, fabsf(to_float_device(input[idx])));
}
shared_max[threadIdx.x] = thread_max;
__syncthreads();
for(int offset = blockDim.x / 2; offset > 0; offset >>= 1)
{
if(threadIdx.x < offset)
{
shared_max[threadIdx.x] = fmaxf(shared_max[threadIdx.x], shared_max[threadIdx.x + offset]);
}
__syncthreads();
}
if(threadIdx.x == 0)
{
atomicMax(reinterpret_cast<unsigned int*>(max_bits), float_as_u32(shared_max[0]));
}
}
GPU_KERNEL void absmax_reduce_bfloat16x8_kernel(
const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>* __restrict__ input,
uint32_t* __restrict__ max_bits,
int64_t vec_count)
{
__shared__ float shared_max[kThreadsPerBlock];
float thread_max = 0.0f;
int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
for(; idx < vec_count; idx += stride)
{
aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth> values = input[idx];
#pragma unroll
for(int i = 0; i < kBfloat16VectorWidth; ++i)
{
thread_max = fmaxf(thread_max, fabsf(to_float_device(values.values[i])));
}
}
shared_max[threadIdx.x] = thread_max;
__syncthreads();
for(int offset = blockDim.x / 2; offset > 0; offset >>= 1)
{
if(threadIdx.x < offset)
{
shared_max[threadIdx.x] = fmaxf(shared_max[threadIdx.x], shared_max[threadIdx.x + offset]);
}
__syncthreads();
}
if(threadIdx.x == 0)
{
atomicMax(reinterpret_cast<unsigned int*>(max_bits), float_as_u32(shared_max[0]));
}
}
template <typename scalar_t>
GPU_KERNEL void quantize_fp8_kernel(
const scalar_t* __restrict__ input,
uint8_t* __restrict__ output,
float* __restrict__ scale,
const uint32_t* __restrict__ max_bits,
int64_t numel)
{
float absmax = fmaxf(u32_as_float(max_bits[0]), kMinScale);
float inv_scale = kFp8E4M3fnMax / absmax;
float scale_value = absmax / kFp8E4M3fnMax;
if(blockIdx.x == 0 && threadIdx.x == 0)
{
scale[0] = scale_value;
}
int64_t thread_idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t thread_stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
for(int64_t idx = thread_idx * kQuantValuesPerThread;
idx + (kQuantValuesPerThread - 1) < numel;
idx += thread_stride * kQuantValuesPerThread)
{
float first = to_float_device(input[idx]) * inv_scale;
float second = to_float_device(input[idx + 1]) * inv_scale;
float third = to_float_device(input[idx + 2]) * inv_scale;
float fourth = to_float_device(input[idx + 3]) * inv_scale;
store_fp8x4(output, idx, pack_fp8x4(first, second, third, fourth));
}
if(thread_idx == 0)
{
int64_t tail_idx = numel & ~static_cast<int64_t>(kQuantValuesPerThread - 1);
if(tail_idx + 1 < numel)
{
float first = to_float_device(input[tail_idx]) * inv_scale;
float second = to_float_device(input[tail_idx + 1]) * inv_scale;
store_fp8x2(output, tail_idx, pack_fp8x2(first, second));
}
if((numel & 1) != 0)
{
output[numel - 1] = pack_fp8_scalar(to_float_device(input[numel - 1]) * inv_scale);
}
}
}
GPU_KERNEL void quantize_fp8_bfloat16x8_kernel(
const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>* __restrict__ input,
uint8_t* __restrict__ output,
float* __restrict__ scale,
const uint32_t* __restrict__ max_bits,
int64_t vec_count)
{
float absmax = fmaxf(u32_as_float(max_bits[0]), kMinScale);
float inv_scale = kFp8E4M3fnMax / absmax;
float scale_value = absmax / kFp8E4M3fnMax;
if(blockIdx.x == 0 && threadIdx.x == 0)
{
scale[0] = scale_value;
}
int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t stride = static_cast<int64_t>(blockDim.x) * gridDim.x;
for(; idx < vec_count; idx += stride)
{
aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth> values = input[idx];
int64_t out_idx = idx * kBfloat16VectorWidth;
store_fp8x4(
output,
out_idx,
pack_fp8x4(
to_float_device(values.values[0]) * inv_scale,
to_float_device(values.values[1]) * inv_scale,
to_float_device(values.values[2]) * inv_scale,
to_float_device(values.values[3]) * inv_scale));
store_fp8x4(
output,
out_idx + 4,
pack_fp8x4(
to_float_device(values.values[4]) * inv_scale,
to_float_device(values.values[5]) * inv_scale,
to_float_device(values.values[6]) * inv_scale,
to_float_device(values.values[7]) * inv_scale));
}
}
template <typename scalar_t>
void launch_quantize_fp8_kernels(
const scalar_t* input,
uint8_t* output,
float* scale,
uint32_t* max_bits,
int64_t numel,
gpu_queue_t queue)
{
int blocks = static_cast<int>((numel + kThreadsPerBlock - 1) / kThreadsPerBlock);
blocks = max(1, min(blocks, 4096));
GPU_LAUNCH_KERNEL(absmax_reduce_kernel<scalar_t>, blocks, kThreadsPerBlock, 0, queue, input, max_bits, numel);
GPU_LAUNCH_KERNEL(quantize_fp8_kernel<scalar_t>, blocks, kThreadsPerBlock, 0, queue, input, output, scale, max_bits, numel);
}
void launch_quantize_fp8_bfloat16x8_kernels(
const gpu_bfloat16_t* input,
uint8_t* output,
float* scale,
uint32_t* max_bits,
int64_t numel,
gpu_queue_t queue)
{
constexpr uintptr_t kVectorAlignment = sizeof(aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>);
TORCH_CHECK(numel % kBfloat16VectorWidth == 0, "bf16 x8 fast path requires numel divisible by 8");
TORCH_CHECK(is_aligned_ptr(input, kVectorAlignment), "bf16 x8 fast path requires 16-byte aligned input");
int64_t vec_count = numel / kBfloat16VectorWidth;
int blocks = static_cast<int>((vec_count + kThreadsPerBlock - 1) / kThreadsPerBlock);
blocks = max(1, min(blocks, 4096));
auto* input_vec = reinterpret_cast<const aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>*>(input);
GPU_LAUNCH_KERNEL(absmax_reduce_bfloat16x8_kernel, blocks, kThreadsPerBlock, 0, queue, input_vec, max_bits, vec_count);
GPU_LAUNCH_KERNEL(quantize_fp8_bfloat16x8_kernel, blocks, kThreadsPerBlock, 0, queue, input_vec, output, scale, max_bits, vec_count);
}
} // namespace
void quantize_fp8_inline(torch::Tensor input, torch::Tensor out, torch::Tensor scale, torch::Tensor max_bits)
{
auto q = get_current_gpu_queue();
int64_t numel = input.numel();
if(input.scalar_type() == torch::kFloat)
{
launch_quantize_fp8_kernels(
reinterpret_cast<const float*>(input.data_ptr()),
reinterpret_cast<uint8_t*>(out.data_ptr()),
reinterpret_cast<float*>(scale.data_ptr()),
reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>()),
numel,
q);
}
else if(input.scalar_type() == torch::kHalf)
{
launch_quantize_fp8_kernels(
reinterpret_cast<const gpu_half_t*>(input.data_ptr()),
reinterpret_cast<uint8_t*>(out.data_ptr()),
reinterpret_cast<float*>(scale.data_ptr()),
reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>()),
numel,
q);
}
else
{
auto* input_ptr = reinterpret_cast<const gpu_bfloat16_t*>(input.data_ptr());
auto* output_ptr = reinterpret_cast<uint8_t*>(out.data_ptr());
auto* scale_ptr = reinterpret_cast<float*>(scale.data_ptr());
auto* max_bits_ptr = reinterpret_cast<uint32_t*>(max_bits.data_ptr<int>());
constexpr uintptr_t kVectorAlignment = sizeof(aligned_vec_t<gpu_bfloat16_t, kBfloat16VectorWidth>);
bool use_bfloat16_fast_path =
input.is_contiguous() &&
out.is_contiguous() &&
(numel % kBfloat16VectorWidth) == 0 &&
is_aligned_ptr(input_ptr, kVectorAlignment);
if(use_bfloat16_fast_path)
{
launch_quantize_fp8_bfloat16x8_kernels(
input_ptr,
output_ptr,
scale_ptr,
max_bits_ptr,
numel,
q);
}
else
{
launch_quantize_fp8_kernels(
input_ptr,
output_ptr,
scale_ptr,
max_bits_ptr,
numel,
q);
}
}
gpu_kernel_check();
}
"""
return cpp_source, gpu_source
def _load_quantize_fp8_inline_module():
global _FP8_QUANT_INLINE_MODULE, _FP8_QUANT_INLINE_LOAD_ERROR
if _FP8_QUANT_INLINE_MODULE is not None:
return _FP8_QUANT_INLINE_MODULE
if _FP8_QUANT_INLINE_LOAD_ERROR is not None:
raise RuntimeError(
"failed to load inline FP8 quant module"
) from _FP8_QUANT_INLINE_LOAD_ERROR
try:
from torch.utils.cpp_extension import load_inline
compiler = shutil.which("clang++") or shutil.which("c++") or shutil.which("g++")
if compiler is None:
raise RuntimeError("no usable C++ compiler found for inline FP8 quant module")
os.environ.setdefault("CXX", compiler)
backend_name = "hip" if torch.version.hip is not None else "cuda"
cpp_source, gpu_source = _quantize_fp8_inline_sources()
extra_cuda_cflags = ["-O3", "-DNDEBUG", "-std=c++17", "--use_fast_math"]
if torch.version.hip is not None:
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
extra_cuda_cflags = ["-O3", "-DNDEBUG", "-std=c++17", "-ffast-math"]
_FP8_QUANT_INLINE_MODULE = load_inline(
name=f"fp8_quant_inline_ext_{backend_name}",
cpp_sources=[cpp_source],
cuda_sources=[gpu_source],
functions=["quantize_fp8_inline"],
extra_cflags=["-O3", "-DNDEBUG", "-std=c++17", "-ffast-math"],
extra_cuda_cflags=extra_cuda_cflags,
verbose=False,
)
return _FP8_QUANT_INLINE_MODULE
except Exception as exc:
_FP8_QUANT_INLINE_LOAD_ERROR = exc
raise RuntimeError("failed to build inline FP8 quant module") from exc
def custom_kernel(data):
"""Reference MLA decode attention. Uses Q_DTYPE and KV_DTYPE to select kernel variant."""
q, kv_data, qo_indptr, kv_indptr, config = data
# Resolve Q
if Q_DTYPE == "fp8":
q_input, q_scale = quantize_fp8(q)
else:
q_input, q_scale = q, None
# Resolve KV
if KV_DTYPE == "fp8":
kv_input, kv_scale = kv_data["fp8"]
else:
kv_input, kv_scale = kv_data["bf16"], None
return flash_sparse_attn_forward_func(
q_input, kv_input, config,
q_scale=q_scale, kv_scale=kv_scale,
)
# return flash_sparse_attn_varlen_forward_func(
# q_input, kv_input, qo_indptr, kv_indptr, config,
# q_scale=q_scale, kv_scale=kv_scale,
# )
# ---------------------------------------------------------------------------
# FP8 quantization (sglang style: dynamic per-tensor)
# ---------------------------------------------------------------------------
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).
Args:
tensor: bf16 tensor to quantize
Returns:
(fp8_tensor, scale) where scale is a scalar float32 tensor.
Dequantize: fp8_tensor.to(bf16) * scale
"""
module = _load_quantize_fp8_inline_module()
fp8_tensor = torch.empty(tensor.shape, dtype=FP8_DTYPE, device=tensor.device)
scale = torch.empty((1,), dtype=torch.float32, device=tensor.device)
max_bits = _get_cached_tensor(
(
"fp8_quant_max_bits",
tensor.device.type,
tensor.device.index,
),
(1,),
torch.int32,
tensor.device,
)
max_bits.zero_()
module.quantize_fp8_inline(tensor, fp8_tensor, scale, max_bits)
return fp8_tensor, scale
# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
# Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
# ---------------------------------------------------------------------------
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""
Converts given x (in fp32) to mxfp4 format.
x: [BLOCK_SIZE_M, BLOCK_SIZE_N], fp32
"""
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Calculate scale
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
# blockscale_e8m0
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127 # in fp32, we have 2&(e - 127)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
# Compute quantized x
qx = x * quant_scale
# Convert quantized fp32 tensor to uint32 before converting to mxfp4 format
# Note: MXFP4 S:1-bit, E:2-bit, M:1-bit
# Zeros: S000 -> +/-0
# Denormal Numbers: S001 -> +/- 0.5
# Normal Numbers:
# S010 -> +/- 1.0
# S011 -> +/- 1.5
# S100 -> +/- 2.0
# S101 -> +/- 3.0
# S110 -> +/- 4.0
# S111 -> +/- 6.0
qx = qx.to(tl.uint32, bitcast=True)
# Extract sign
s = qx & 0x80000000
# Set everything to positive, will add sign back at the end
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (saturate_mask | denormal_mask)
# Denormal numbers
denorm_exp: tl.constexpr = (
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
# Normal numbers
normal_x = qx
# resulting mantissa is odd
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
# update exponent, rounding bias part 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
# rounding bias part 2
normal_x += mant_odd
# take the bits!
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
# Merge results
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
# add sign back
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _dynamic_mxfp4_quant_kernel(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
stride_bs_m_in,
stride_bs_n_in,
M,
N,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SCALING_MODE: tl.constexpr,
num_warps: tl.constexpr,
waves_per_eu: tl.constexpr,
num_stages: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
# cast strides to int64, in case M*N > max int32
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
tl.store(x_fp4_ptr + out_offs, out_tensor)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
tl.store(bs_ptr + bs_offs, bs_e8m0)
def dynamic_mxfp4_quant(
x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Quantize a tensor to MX FP4 format.
Args:
x: The input tensor, typically fp16 or bf16.
scaling_mode: The method to calculate MX block scaling.
- "even" (default): `even_round` in `quark.torch.quantization.utils`.
- etc.
Returns:
A tuple of (x_fp4, blockscale_e8m0).
"""
# Assume x is 2D-Tensor for now
M, N = x.shape
assert (N // 2) % 2 == 0
# This is fixed by spec for MXFP4. Do not tune this.
MXFP4_QUANT_BLOCK_SIZE = 32
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
blockscale_e8m0 = torch.empty(
((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE, M),
dtype=torch.uint8,
device=x.device,
).T
# for large N values
if M <= 32:
NUM_ITER = 1
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 4
BLOCK_SIZE_M = 64
BLOCK_SIZE_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if N <= 16384:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 128
# for small N values
if N <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
if M == 16:
BLOCK_SIZE_M = 16
BLOCK_SIZE_N = 128
else:
BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
# BLOCK_SIZE_N needs to be multiple of 32
BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
)
_dynamic_mxfp4_quant_kernel[grid](
x,
x_fp4,
blockscale_e8m0,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0.stride(),
M=M,
N=N,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
SCALING_MODE=0,
NUM_ITER=NUM_ITER,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES,
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=1,
)
return (x_fp4, blockscale_e8m0)
def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.
Block size = 32. Each block gets an E8M0 scale factor.
Two FP4 E2M1 values are packed per byte.
Args:
tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)
Returns:
(fp4_data, scale_e8m0)
- fp4_data: shape [B, M, N//2] in aiter_dtypes.fp4x2
- scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
"""
orig_shape = tensor.shape # (B, M, N)
B, M, N = orig_shape
# dynamic_mxfp4_quant expects 2D: (B*M, N)
tensor_2d = tensor.reshape(B * M, N)
fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)
# Reshape fp4_data back to 3D: (B, M, N//2)
fp4_data = fp4_data_2d.view(B, M, N // 2)
return fp4_data, scale_e8m0
def mxfp4_to_f32(x: torch.Tensor) -> torch.Tensor:
x = x.view(torch.uint8)
x = x.repeat_interleave(2, dim=-1)
x[..., ::2] = x[..., ::2] & 0xF
x[..., 1::2] = x[..., 1::2] >> 4
mxfp4_list = [
0.0,
0.5,
1.0,
1.5,
2.0,
3.0,
4.0,
6.0,
-0.0,
-0.5,
-1.0,
-1.5,
-2.0,
-3.0,
-4.0,
-6.0,
]
mxfp4_in_f32 = torch.tensor(mxfp4_list, dtype=torch.float32, device=x.device)
return mxfp4_in_f32[x.long()]
def e8m0_to_f32(scale_e8m0_biased: torch.Tensor) -> torch.Tensor:
scale_e8m0_biased = scale_e8m0_biased.view(torch.uint8)
zero_case = scale_e8m0_biased == 0
nan_case = scale_e8m0_biased == 0xFF
scale_f32 = scale_e8m0_biased.to(torch.int32) << 23
scale_f32[zero_case] = 0x00400000
scale_f32[nan_case] = 0x7F800001
return scale_f32.view(torch.float32)
def dequantize_mxfp4(
fp4_data: torch.Tensor,
scale_e8m0: torch.Tensor,
orig_shape: tuple[int, int, int],
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
bsz, num_heads, width = orig_shape
num_rows = bsz * num_heads
num_blocks = width // MXFP4_GROUP_SIZE
fp4_data_2d = fp4_data.reshape(num_rows, width // 2)
values_f32 = mxfp4_to_f32(fp4_data_2d)
scale_f32 = e8m0_to_f32(scale_e8m0)[:num_rows, :num_blocks]
values_f32 = values_f32.view(num_rows, num_blocks, MXFP4_GROUP_SIZE)
values_f32 = values_f32 * scale_f32.unsqueeze(-1)
return values_f32.view(bsz, num_heads, width).to(dtype)
def generate_input(batchsize: int, qseqlen: int, kvseqlen: int, seed: int):
"""
Generate absorbed q and compressed kv_buffer for MLA decode.
Returns all three KV cache formats in kv_data dict:
kv_data = {
"bf16": Tensor — (total_kv, 1, 576) bfloat16
"fp8": (Tensor, Tensor) — kv_buffer fp8 + scalar scale
"mxfp4": (Tensor, Tensor) — kv_buffer fp4x2 + fp8_e8m0 scale
}
"""
gen = torch.Generator(device="cuda")
gen.manual_seed(seed)
total_q = batchsize * qseqlen
total_kv = batchsize * kvseqlen
# Absorbed query: (total_q, num_heads, 576) bf16
q = torch.randn(
(total_q, NUM_HEADS, QK_HEAD_DIM),
dtype=torch.bfloat16, device="cuda", generator=gen,
)
# Compressed KV buffer: (total_kv, 1, 576) bf16 — the source of truth
kv_buffer_bf16 = torch.randn(
(total_kv, NUM_KV_HEADS, QK_HEAD_DIM),
dtype=torch.bfloat16, device="cuda", generator=gen,
)
kv_data = {
"bf16": kv_buffer_bf16,
}
if KV_DTYPE == "fp8":
kv_data["fp8"] = quantize_fp8(kv_buffer_bf16)
elif KV_DTYPE == "mxfp4":
kv_data["mxfp4"] = quantize_mxfp4(kv_buffer_bf16)
qo_indptr = torch.arange(0, batchsize + 1, dtype=torch.int32, device="cuda") * qseqlen
kv_indptr = torch.arange(0, batchsize + 1, dtype=torch.int32, device="cuda") * kvseqlen
config = {
"batch_size": batchsize,
"num_heads": NUM_HEADS,
"num_kv_heads": NUM_KV_HEADS,
"qk_head_dim": QK_HEAD_DIM,
"kv_lora_rank": KV_LORA_RANK,
"qk_rope_head_dim": QK_ROPE_HEAD_DIM,
"v_head_dim": V_HEAD_DIM,
"q_seq_len": qseqlen,
"kv_seq_len": kvseqlen,
"sm_scale": SM_SCALE,
}
return (q, kv_data, qo_indptr, kv_indptr, config)
_BENCHMARK_CASES = [
{"batchsize": 4, "qseqlen": 1, "kvseqlen": 1024, "seed": 4217},
{"batchsize": 32, "qseqlen": 1, "kvseqlen": 1024, "seed": 5412},
{"batchsize": 64, "qseqlen": 1, "kvseqlen": 1024, "seed": 1357},
{"batchsize": 256, "qseqlen": 1, "kvseqlen": 1024, "seed": 9823},
{"batchsize": 4, "qseqlen": 1, "kvseqlen": 8192, "seed": 4220},
{"batchsize": 32, "qseqlen": 1, "kvseqlen": 8192, "seed": 5415},
{"batchsize": 64, "qseqlen": 1, "kvseqlen": 8192, "seed": 1360},
{"batchsize": 256, "qseqlen": 1, "kvseqlen": 8192, "seed": 9826},
]
def _resolve_test_q(
q: torch.Tensor,
) -> torch.Tensor:
if Q_DTYPE == "fp8":
q_fp8, q_scale = quantize_fp8(q)
return q_fp8.to(torch.bfloat16) * q_scale.to(torch.bfloat16)
return q.to(torch.bfloat16)
def _resolve_test_kv(
kv_data: dict[str, Any],
) -> torch.Tensor:
if KV_DTYPE == "fp8":
kv_fp8, kv_scale = kv_data["fp8"]
return kv_fp8.to(torch.bfloat16) * kv_scale.to(torch.bfloat16)
if KV_DTYPE == "mxfp4":
kv_fp4, kv_scale = kv_data["mxfp4"]
return dequantize_mxfp4(
kv_fp4,
kv_scale,
tuple(kv_data["bf16"].shape),
dtype=torch.bfloat16,
)
return kv_data["bf16"].to(torch.bfloat16)
def _torch_reference_mla_decode(data) -> torch.Tensor:
q, kv_data, qo_indptr, kv_indptr, config = data
q_dense = _resolve_test_q(q)
kv_dense = _resolve_test_kv(kv_data)
value_dim = int(config["v_head_dim"])
values = kv_dense[..., :value_dim]
scale = float(config["sm_scale"])
total_q, num_heads, _ = q_dense.shape
out = torch.empty(
(total_q, num_heads, value_dim),
dtype=torch.bfloat16,
device=q_dense.device,
)
batch_size = int(config["batch_size"])
for batch_idx in range(batch_size):
q_start = int(qo_indptr[batch_idx].item())
q_end = int(qo_indptr[batch_idx + 1].item())
k_start = int(kv_indptr[batch_idx].item())
k_end = int(kv_indptr[batch_idx + 1].item())
q_slice = q_dense[q_start:q_end].float()
k_slice = kv_dense[k_start:k_end].float()
v_slice = values[k_start:k_end].float()
# MQA: one KV head is shared across all query heads.
scores = torch.einsum("qhd,kd->qhk", q_slice, k_slice[:, 0, :]) * scale
probs = torch.softmax(scores, dim=-1)
out_slice = torch.einsum("qhk,kd->qhd", probs, v_slice[:, 0, :])
out[q_start:q_end] = out_slice.to(torch.bfloat16)
return out
def _run_warmup(data, warmup: int) -> None:
for _ in range(warmup):
custom_kernel(data)
if torch.cuda.is_available():
torch.cuda.synchronize()
def _run_correctness_test(case: dict[str, int]) -> dict[str, float]:
def _clone_data(data: Any) -> Any:
if isinstance(data, tuple):
return tuple(_clone_data(x) for x in data)
if isinstance(data, list):
return [_clone_data(x) for x in data]
if isinstance(data, dict):
return {k: _clone_data(v) for k, v in data.items()}
if isinstance(data, torch.Tensor):
return data.clone()
return data
test_case = dict(case)
max_abs = 0.0
mean_abs = 0.0
for repeat_idx in range(5):
if repeat_idx > 0 and "seed" in test_case:
test_case["seed"] += 13
data = generate_input(
test_case["batchsize"],
test_case["qseqlen"],
test_case["kvseqlen"],
test_case["seed"],
)
check_copy = _clone_data(data)
out = custom_kernel(_clone_data(data))
ref = _torch_reference_mla_decode(check_copy)
diff = (out.float() - ref.float()).abs()
max_abs = max(max_abs, diff.max().item())
mean_abs = max(mean_abs, diff.mean().item())
torch.testing.assert_close(out.float(), ref.float(), rtol=1e-1, atol=1e-1)
return {
"max_abs": max_abs,
"mean_abs": mean_abs,
}
def _run_speed_test(
case: dict[str, int],
warmup: int,
repeats: int,
) -> dict[str, float]:
data = generate_input(
case["batchsize"],
case["qseqlen"],
case["kvseqlen"],
case["seed"],
)
_run_warmup(data, warmup)
if torch.cuda.is_available():
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
durations_us: list[float] = []
for _ in range(repeats):
start_event.record()
custom_kernel(data)
end_event.record()
end_event.synchronize()
durations_us.append(start_event.elapsed_time(end_event) * 1000.0)
else:
durations_us = []
for _ in range(repeats):
start_ns = time.perf_counter_ns()
custom_kernel(data)
durations_us.append((time.perf_counter_ns() - start_ns) / 1000.0)
total_q = case["batchsize"] * case["qseqlen"]
total_kv = case["batchsize"] * case["kvseqlen"]
mean_us = statistics.fmean(durations_us)
return {
"mean_us": mean_us,
"min_us": min(durations_us),
"max_us": max(durations_us),
"tokens_per_s": total_q * 1e6 / max(mean_us, 1e-6),
"kv_tokens_per_s": total_kv * 1e6 / max(mean_us, 1e-6),
}
def run_tests(
warmup: int = 20,
repeats: int = 200,
run_correctness: bool = True,
run_speed: bool = True,
) -> None:
if torch.cuda.is_available():
torch.cuda.synchronize()
print(f"test.warmup: {warmup}")
print(f"test.repeats: {repeats}")
print(f"test.q_dtype: {Q_DTYPE}")
print(f"test.kv_dtype: {KV_DTYPE}")
print(f"test.cases: {len(_BENCHMARK_CASES)}")
speed_latencies: list[float] = []
for index, case in enumerate(_BENCHMARK_CASES):
label = (
f"bs={case['batchsize']} q={case['qseqlen']} "
f"kv={case['kvseqlen']} seed={case['seed']}"
)
print(f"case[{index}].spec: {label}")
if run_correctness:
correctness = _run_correctness_test(case)
print(
f"case[{index}].correctness: "
f"max_abs={correctness['max_abs']:.6f} "
f"mean_abs={correctness['mean_abs']:.6f}"
)
if run_speed:
perf = _run_speed_test(case, warmup=warmup, repeats=repeats)
speed_latencies.append(perf["mean_us"])
print(
f"case[{index}].speed: mean_us={perf['mean_us']:.3f} "
f"min_us={perf['min_us']:.3f} max_us={perf['max_us']:.3f}"
)
print(
f"case[{index}].throughput: q_tok_s={perf['tokens_per_s']:.2f} "
f"kv_tok_s={perf['kv_tokens_per_s']:.2f}"
)
if run_speed and speed_latencies:
geom_mean_us = math.exp(
sum(math.log(latency_us) for latency_us in speed_latencies)
/ len(speed_latencies)
)
print(f"test.geom_mean_us: {geom_mean_us:.3f}")
if __name__ == "__main__":
run_tests()
scrolls · 3120 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON