submission 695064
Jingze · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2927 lines, June 9 Researcher Reciprocity License v1.0.
submission_sparse.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-695064?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:befb391973e6b0148aac88caa3cb8af954336de751de1388bd77e084d0549cf4
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-15
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.py2927 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)
# tile_n, batch_size, seqlen_k, is_local, window_size_left, window_size_right
_BENCHMARK_SPLIT_PRESETS = {
(16, 4, 1024, True, 1024, 0): (8, [0, 9, 18, 27, 35, 43, 50, 57, 64]),
(16, 32, 1024, True, 1024, 0): (8, [0, 9, 17, 25, 33, 41, 49, 57, 64]),
(16, 64, 1024, True, 1024, 0): (8, [0, 9, 17, 25, 33, 41, 49, 57, 64]),
(16, 256, 1024, True, 1024, 0): (4, [0, 17, 33, 49, 64]),
(16, 4, 8192, True, 4096, 0): (16, [255, 273, 291, 309, 327, 344, 361, 378, 394, 410, 426, 441, 456, 470, 484, 498, 512]),
(16, 32, 8192, True, 4096, 0): (16, [255, 273, 290, 307, 324, 341, 358, 374, 390, 406, 422, 437, 452, 467, 482, 497, 512]),
(16, 64, 8192, True, 4096, 0): (16, [255, 273, 290, 307, 324, 341, 358, 374, 390, 406, 422, 437, 452, 467, 482, 497, 512]),
(16, 256, 8192, True, 4096, 0): (16, [255, 272, 289, 306, 322, 338, 354, 370, 386, 402, 418, 434, 450, 466, 481, 497, 512]),
(32, 4, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
(32, 32, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
(32, 64, 1024, True, 1024, 0): (4, [0, 9, 17, 25, 32]),
(32, 256, 1024, True, 1024, 0): (4, [0, 8, 16, 24, 32]),
(32, 4, 8192, True, 4096, 0): (8, [127, 146, 164, 181, 197, 213, 228, 242, 256]),
(32, 32, 8192, True, 4096, 0): (8, [127, 145, 162, 179, 195, 211, 226, 241, 256]),
(32, 64, 8192, True, 4096, 0): (8, [127, 145, 162, 179, 195, 211, 226, 241, 256]),
(32, 256, 8192, True, 4096, 0): (8, [127, 144, 161, 177, 193, 209, 225, 241, 256]),
}
_SPLIT_TENSOR_CACHE: dict[tuple[Any, ...], tuple[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_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
# # Compute n_block range for this m_block
# 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=IS_SPLIT_KV,
# 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)
if IS_SPLIT_KV:
n_block_min = tl.load(SplitBoundaries + split_idx)
n_block_max = tl.load(SplitBoundaries + split_idx + 1)
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
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,
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)
# Compute max across splits
for s in tl.range(0, num_splits):
lse_s = tl.load(
lse_part_row_base + s * stride_lps,
cache_modifier=".cg",
)
o_s = tl.load(
out_part_row_base + s * 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,
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,
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),
)
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, preset_boundaries = _BENCHMARK_SPLIT_PRESETS[preset_key]
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,
)
_SPLIT_TENSOR_CACHE[split_cache_key] = (
split_counts,
split_boundaries,
num_splits,
)
else:
split_counts, split_boundaries, num_splits = cached_split_tensors
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,
)
out_partial = _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_partial = _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,
)
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_partial,
lse_partial,
split_counts,
split_boundaries,
softmax_scale_log2,
0,
query.stride(-2),
query.stride(0),
0,
kv.stride(-2),
kv.stride(0),
0,
out_partial.stride(-2),
out_partial.stride(-3),
out_partial.stride(0),
0,
lse_partial.stride(-2),
lse_partial.stride(0),
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=True,
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,
)
_flash_attn_fwd_combine(
out_partial,
lse_partial,
out,
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
is_split_kv = seqlen_q == 1 and seqlen_q != seqlen_k
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),
)
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, preset_boundaries = _BENCHMARK_SPLIT_PRESETS[preset_key]
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,
)
_SPLIT_TENSOR_CACHE[split_cache_key] = (
split_counts,
split_boundaries,
num_splits,
)
else:
split_counts, split_boundaries, num_splits = cached_split_tensors
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,
)
out_partial = _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_partial = _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,
)
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_partial,
lse_partial,
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),
out_partial.stride(1),
out_partial.stride(-2),
out_partial.stride(-3),
out_partial.stride(0),
lse_partial.stride(1),
lse_partial.stride(-2),
lse_partial.stride(0),
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,
)
_flash_attn_fwd_combine(
out_partial,
lse_partial,
out,
)
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 · 2927 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 692590.
⋯ 526 unchanged linesrow_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].⋯ 9 unchanged lines# (row_max * scale_log2 + tl.log2(row_sum)) * ln2,# )- row_scale = 1.0 / row_sum+ row_scale = 1.0 / row_sum * final_scaleln2 = 0.6931471805599453lse = (row_max * scale_log2 + tl.log2(row_sum)) * ln2return row_scale, lse⋯ 109 unchanged linesdef _fwd_inner_sparse_fp8_kernel(q_tile,q_tile_tail,- q_scale,- kv_scale,k_tile,k_tile_tail,k_ptrs,⋯ 1 unchanged linesacc_o,row_max,row_sum,- softmax_scale_log2: tl.constexpr,+ score_scale_log2,n_block,n_block_min,TILE_M: tl.constexpr,⋯ 2 unchanged lines):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_scaled(- q_tile_tail,- None,- "e4m3",- k_tile_tail,- None,- "e4m3",- acc=acc_s,- fast_math=True,- )- acc_s *= q_scale * kv_scale+ # 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 pointerk_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 pointerk_tail_ptrs = tl.advance(k_tail_ptrs, (0, -TILE_N))if n_block > n_block_min:# Load next key tile- k_tile_next = tl.load(k_ptrs, boundary_check=(0, 1))- k_tile_tail = tl.load(k_tail_ptrs, boundary_check=(0, 1))+ k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")# Apply online softmaxp, row_max, row_sum, row_scale = (⋯ 1 unchanged linesacc_s=acc_s,row_max=row_max,row_sum=row_sum,- scale_log2=softmax_scale_log2,+ scale_log2=score_scale_log2,CHECK_INF=CHECK_INF,))+ p = p.to(k_tile.dtype)+v_tile = tl.trans(k_tile)# Rescale output accumulatoracc_o = rescale_o(acc_o, row_scale)# Update output accumulator- acc_o += tl.dot_scaled(- p.to(v_tile.dtype),- None,- "e4m3",- v_tile,- None,- "e4m3",- fast_math=True,- ) * kv_scale+ # acc_o += tl.dot_scaled(+ # p,+ # None,+ # "e4m3",+ # v_tile,+ # None,+ # "e4m3",+ # fast_math=True,+ # )+ acc_o += tl.dot(p, v_tile)- k_tile = k_tile_next-return (- k_tile,+ k_tile_next,k_tile_tail,k_ptrs,k_tail_ptrs,⋯ 13 unchanged linesLse,SplitCounts,SplitBoundaries,- softmax_scale_log2: tl.constexpr,+ softmax_scale_log2,stride_qb,stride_qh,stride_qm,⋯ 105 unchanged linesout_base += split_idx * stride_oslse_base += split_idx * stride_ls- # Compute n_block range for this m_block- 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=IS_SPLIT_KV,- WINDOW_SIZE_LEFT=WINDOW_SIZE_LEFT,- WINDOW_SIZE_RIGHT=WINDOW_SIZE_RIGHT,- QHEAD_PER_KVHEAD_PACKGQA=QHEADS_PER_KVHEAD_PACKGQA,- )+ # # Compute n_block range for this m_block+ # 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=IS_SPLIT_KV,+ # 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)+if IS_SPLIT_KV:- split_boundaries = SplitBoundaries- split_n_block_min = tl.load(split_boundaries + split_idx)- split_n_block_max = tl.load(split_boundaries + 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 = tl.load(SplitBoundaries + split_idx)+ n_block_max = tl.load(SplitBoundaries + split_idx + 1)n_block_min_no_mask = get_n_block_min_before_local_mask(seqlen_q=actual_seqlen_q,⋯ 84 unchanged linesq_scale = tl.load(Q_S)kv_scale = tl.load(KV_S)+ score_scale_log2 = softmax_scale_log2 * q_scale * kv_scale# Load query tileq_tile = tl.load(⋯ 11 unchanged linesacc_o = tl.zeros((TILE_M, TILE_K), dtype=tl.float32)# Load key tile- k_tile = tl.load(k_ptrs)- k_tile_tail = tl.load(k_tail_ptrs)+ 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_maxn_block_min_no_mask = tl.minimum(n_block_min_no_mask, n_block_max_no_mask)⋯ 12 unchanged lines) = _fwd_inner_sparse_fp8_kernel(q_tile=q_tile,q_tile_tail=q_tile_tail,- q_scale=q_scale,- kv_scale=kv_scale,k_tile=k_tile,k_tile_tail=k_tile_tail,k_ptrs=k_ptrs,⋯ 1 unchanged linesacc_o=acc_o,row_max=row_max,row_sum=row_sum,- softmax_scale_log2=softmax_scale_log2,+ score_scale_log2=score_scale_log2,n_block=n_block,n_block_min=n_block_min_no_mask,TILE_M=TILE_M,⋯ 20 unchanged linesorder=(0, 1),)- k_tile = tl.load(k_ptrs)- k_tile_tail = tl.load(k_tail_ptrs)+ 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,⋯ 6 unchanged lines) = _fwd_inner_sparse_fp8_kernel(q_tile=q_tile,q_tile_tail=q_tile_tail,- q_scale=q_scale,- kv_scale=kv_scale,k_tile=k_tile,k_tile_tail=k_tile_tail,k_ptrs=k_ptrs,⋯ 1 unchanged linesacc_o=acc_o,row_max=row_max,row_sum=row_sum,- softmax_scale_log2=softmax_scale_log2,+ score_scale_log2=score_scale_log2,n_block=n_block,n_block_min=n_block_min,TILE_M=TILE_M,⋯ 5 unchanged linesrow_scale, lse_tile = finalize(row_max=row_max,row_sum=row_sum,- scale_log2=softmax_scale_log2,+ scale_log2=score_scale_log2,+ final_scale=kv_scale,)acc_o = rescale_o(acc_o, row_scale)⋯ 214 unchanged linesTILE_K = 512TILE_M = 16- TILE_N = 16- num_warps = 4+ TILE_N = 32+ num_warps = 8num_stages = 2waves_per_eu = 0matrix_instr_nonkdim = 16⋯ 124 unchanged linesreturn 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+ is_split_kv = seqlen_q == 1 and seqlen_q != seqlen_k+ 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),+ )+ 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, preset_boundaries = _BENCHMARK_SPLIT_PRESETS[preset_key]+ 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,+ )+ _SPLIT_TENSOR_CACHE[split_cache_key] = (+ split_counts,+ split_boundaries,+ num_splits,+ )+ else:+ split_counts, split_boundaries, num_splits = cached_split_tensors++ 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,+ )++ out_partial = _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_partial = _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,+ )++ 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_partial,+ lse_partial,+ 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),+ out_partial.stride(1),+ out_partial.stride(-2),+ out_partial.stride(-3),+ out_partial.stride(0),+ lse_partial.stride(1),+ lse_partial.stride(-2),+ lse_partial.stride(0),+ 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,+ )++ _flash_attn_fwd_combine(+ out_partial,+ lse_partial,+ out,+ )++ 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,⋯ 581 unchanged lineselse:kv_input, kv_scale = kv_data["bf16"], None- return flash_sparse_attn_varlen_forward_func(- q_input, kv_input, qo_indptr, kv_indptr, config,+ 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,+ # )# ---------------------------------------------------------------------------⋯ 317 unchanged linesreturn 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 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 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+ 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)+ 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.+ 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)+ 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+ 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,- # )+ # 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,- # )+ # 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,- # }+ 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)+ 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+ 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,- # }+ 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)+ return (q, kv_data, qo_indptr, kv_indptr, config)- # _BENCHMARK_CASES = [- # {"batchsize": 4, "qseqlen": 1, "kvseqlen": 1024, "seed": 4217},+ _BENCHMARK_CASES = [+ {"batchsize": 4, "qseqlen": 1, "kvseqlen": 1024, "seed": 4217},- # {"batchsize": 32, "qseqlen": 1, "kvseqlen": 1024, "seed": 5412},+ {"batchsize": 32, "qseqlen": 1, "kvseqlen": 1024, "seed": 5412},- # {"batchsize": 64, "qseqlen": 1, "kvseqlen": 1024, "seed": 1357},+ {"batchsize": 64, "qseqlen": 1, "kvseqlen": 1024, "seed": 1357},- # {"batchsize": 256, "qseqlen": 1, "kvseqlen": 1024, "seed": 9823},+ {"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},- # ]+ {"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_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 _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"])+ 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,- # )+ 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())+ 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()+ 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)+ # 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+ 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_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+ 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+ 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+ 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)+ 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)+ 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,- # }+ 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)+ 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)+ 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),- # }+ 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()+ 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)}")+ 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}")+ 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_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:+ 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 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()+ if __name__ == "__main__":+ run_tests()
scrolls · 1159 diff lines total
Best evidence level for this revision: reported
JSON