submission 727774
Jingze · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2067 lines, June 9 Researcher Reciprocity License v1.0.
submission_triton.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-727774?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:93fb15080cbff0d8fd4a90b25182ecc2186c96a15bcb670475bbb0bc38c00741
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
row_max_new = tl.maximum(row_max_curr, row_max)persistent-kernel
tl.num_programs(0),shared-memory
__shared__ float shared_max[kThreadsPerBlock];split-k
is_split_kv = (stages = 2
num_stages = 2tile-k = 512
TILE_K: tl.constexpr = 512tile-m = 16
TILE_M: tl.constexpr = 16tile-n = 32
TILE_N: tl.constexpr = 32Kernel source
submission_triton.py2067 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
def _build_benchmark_split_preset(
split_count: int,
split_boundaries: list[int],
) -> dict[str, int | list[int]]:
reduce_split_count = split_count if split_count > 1 else 0
reduce_split_boundaries = [0, reduce_split_count] if reduce_split_count else [0, 0]
return {
"split_count": split_count,
"split_boundaries": split_boundaries,
"reduce_split_count": reduce_split_count,
"reduce_split_boundaries": reduce_split_boundaries,
}
# tile_n, batch_size, seqlen_k, is_local, window_size_left, window_size_right
_BENCHMARK_SPLIT_PRESETS = {
key: _build_benchmark_split_preset(split_count, split_boundaries)
for key, (split_count, split_boundaries) in {
(32, 4, 1024, True, 1024, 0): (12, [0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 29, 31, 32]), # 28.2 ± 0.03 µs
(32, 32, 1024, True, 1024, 0): (8, [0, 4, 8, 12, 16, 20, 24, 28, 32]), # 32.3 ± 0.03 µs
(32, 64, 1024, True, 1024, 0): (8, [0, 4, 8, 12, 16, 20, 24, 28, 32]), # 38.6 ± 0.04 µs
(32, 256, 1024, True, 1024, 0): (4, [0, 8, 16, 24, 32]), # 85.1 ± 0.08 µs
(32, 4, 8192, True, 4096, 0): (32, [128, 132, 136, 140, 144, 148, 152, 156, 160, 164, 168, 172, 176, 180, 184, 188, 192, 196, 200, 204, 208, 212, 216, 220, 224, 228, 232, 236, 240, 244, 248, 252, 256]), # 37.3 ± 0.04 µs
(32, 32, 8192, True, 4096, 0): (16, [224, 226, 228, 230, 232, 234, 236, 238, 240, 242, 244, 246, 248, 250, 252, 254, 256]), # 48.8
# (32, 64, 8192, True, 4096, 0): (8, [128, 144, 160, 176, 192, 208, 224, 240, 256]), # 66.5 ± 0.07 µs
(32, 64, 8192, True, 4096, 0): (8, [224, 228, 232, 236, 240, 244, 248, 252, 256]), # 66.5 ± 0.07 µs
# (32, 256, 8192, True, 4096, 0): (4, [128, 160, 192, 224, 256]), # 199 ± 0.2 µs
# (32, 256, 8192, True, 4096, 0): (4, [192, 208, 224, 240, 256]),
(32, 256, 8192, True, 4096, 0): (4, [224, 232, 240, 248, 256]),
}.items()
}
_SPLIT_TENSOR_CACHE: dict[
tuple[Any, ...],
tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, int],
] = {}
_BUFFER_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_RESULT_CACHE: dict[tuple[Any, ...], torch.Tensor] = {}
_OUTPUT_USE_CACHE = False
_FP8_QUANT_INLINE_MODULE = None
_FP8_QUANT_INLINE_LOAD_ERROR: Optional[Exception] = None
def _get_dummy_split_metadata(
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
split_counts = _get_cached_tensor(
("split_counts_dummy", device.type, device.index),
(1,),
torch.int32,
device,
)
split_boundaries = _get_cached_tensor(
("split_boundaries_dummy", device.type, device.index),
(2,),
torch.int32,
device,
)
reduce_split_counts = _get_cached_tensor(
("reduce_split_counts_dummy", device.type, device.index),
(1,),
torch.int32,
device,
)
reduce_split_boundaries = _get_cached_tensor(
("reduce_split_boundaries_dummy", device.type, device.index),
(2,),
torch.int32,
device,
)
split_counts.fill_(1)
split_boundaries[0] = 0
split_boundaries[1] = 0
reduce_split_counts.zero_()
reduce_split_boundaries[0] = 0
reduce_split_boundaries[1] = 0
return (
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
)
def _get_dummy_lse_tensor(device: torch.device) -> torch.Tensor:
return _get_cached_tensor(
("lse_dummy", device.type, device.index),
(1,),
torch.float32,
device,
)
def _get_cached_tensor(
cache_key: tuple[Any, ...],
shape: tuple[int, ...],
dtype: torch.dtype,
device: torch.device,
) -> torch.Tensor:
cached = _BUFFER_CACHE.get(cache_key)
if cached is None:
cached = torch.empty(shape, dtype=dtype, device=device)
_BUFFER_CACHE[cache_key] = cached
return cached
def _tensor_identity_key(q: torch.Tensor | None, kv: torch.Tensor | None) -> tuple[Any, ...] | None:
if q is None or kv is None:
return None
return (
tuple(q.shape),
tuple(kv.shape),
)
@triton.jit
def _fwd_kernel(
Q,
Q_S,
KV,
KV_S,
Out,
Lse,
SplitCounts,
SplitBoundaries,
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,
num_heads_kv,
cu_seqlens_q,
cu_seqlens_k,
num_splits,
SEQLEN_K: tl.constexpr,
IS_FP8: 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,
):
QHEADS_PER_KVHEAD_PACKGQA: tl.constexpr = 16
TILE_M: tl.constexpr = 16
TILE_N: tl.constexpr = 32
TILE_K: tl.constexpr = 512
HEAD_DIM: tl.constexpr = 512
TAIL_HEAD_DIM: tl.constexpr = 64
SOFTMAX_SCALE_LOG2: tl.constexpr = 0.06011229337037347
head_batch_split_idx = remap_xcd(
tl.program_id(0),
tl.num_programs(0),
)
head_idx = head_batch_split_idx % num_heads_kv
batch_split_idx = head_batch_split_idx // num_heads_kv
batch_idx = batch_split_idx // num_splits
split_idx = batch_split_idx - batch_idx * num_splits
active_splits = tl.load(SplitCounts + batch_idx)
if split_idx >= active_splits:
return
offs_m = tl.arange(0, TILE_M)
offs_k = tl.arange(0, TILE_K)
offs_kt = tl.arange(0, TAIL_HEAD_DIM)
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
# Get seqlen info for this batch
if HAS_CU_SEQLENS_Q:
offset_q = tl.load(cu_seqlens_q + batch_idx)
else:
offset_q = 0
if HAS_CU_SEQLENS_K:
offset_k = tl.load(cu_seqlens_k + batch_idx)
else:
offset_k = 0
# Initialize base pointers
if HAS_CU_SEQLENS_Q:
q_base = Q + offset_q * stride_qm
else:
q_base = Q + batch_idx * stride_qb
if HAS_CU_SEQLENS_K:
k_base = KV + offset_k * stride_kvn
else:
k_base = KV + batch_idx * stride_kvb
if HAS_CU_SEQLENS_Q:
out_base = Out + offset_q * stride_om + split_idx * stride_os
else:
out_base = Out + batch_idx * stride_ob + split_idx * stride_os
if HAS_CU_SEQLENS_Q:
lse_base = Lse + offset_q + split_idx * stride_ls
else:
lse_base = Lse + batch_idx * stride_lb + split_idx * stride_ls
n_block_min = tl.load(SplitBoundaries + split_idx)
n_block_max = tl.load(SplitBoundaries + split_idx + 1)
# Create pointers
lse_ptrs = lse_base + q_head * stride_lh
out_ptrs = out_base + q_head[:, None] * stride_oh + offs_k[None, :]
q_ptrs = q_base + q_head[:, None] * stride_qh + offs_k[None, :]
q_tail_ptrs = q_base + HEAD_DIM + q_head[:, None] * stride_qh + offs_kt[None, :]
k_ptrs = tl.make_block_ptr(
base=k_base,
shape=(HEAD_DIM, 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,
shape=(TAIL_HEAD_DIM, SEQLEN_K),
strides=(1, stride_kvn),
offsets=(0, (n_block_max - 1) * TILE_N),
block_shape=(TAIL_HEAD_DIM, TILE_N),
order=(1, 0),
)
if IS_FP8:
q_scale = tl.load(Q_S)
kv_scale = tl.load(KV_S)
score_scale_log2 = SOFTMAX_SCALE_LOG2 * q_scale * kv_scale
final_scale = kv_scale
else:
score_scale_log2 = SOFTMAX_SCALE_LOG2
final_scale = 1.0
# 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 query tile
q_tile = tl.load(q_ptrs, cache_modifier=".ca")
# Load key tile
k_tile = tl.load(k_ptrs, cache_modifier=".cg")
q_tile_tail = tl.load(q_tail_ptrs, cache_modifier=".ca")
k_tile_tail = tl.load(k_tail_ptrs, cache_modifier=".cg")
# Process n_blocks without masking
for n_block in tl.range(n_block_max - 1, n_block_min - 1, -1):
# Compute attention scores
acc_s = tl.dot(q_tile, k_tile)
# Advance key pointer
k_ptrs = tl.advance(k_ptrs, (0, -TILE_N))
k_tile_next = k_tile
if n_block > n_block_min:
# Load next key tile
k_tile_next = tl.load(k_ptrs, cache_modifier=".cg")
acc_s += tl.dot(q_tile_tail, k_tile_tail)
# Advance key pointer
k_tail_ptrs = tl.advance(k_tail_ptrs, (0, -TILE_N))
k_tile_tail_next = k_tile_tail
if n_block > n_block_min:
# Load next key tile
k_tile_tail_next = tl.load(k_tail_ptrs, cache_modifier=".cg")
# 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)
# Compute scaled differences to new row max
acc_scale_log2 = (row_max - row_max_new) * score_scale_log2
# Update row max
row_max = row_max_new
# Compute row scale
row_scale = tl.exp2(acc_scale_log2)
# Compute attention weights
p = tl.exp2(acc_s * score_scale_log2 - row_max[:, None] * score_scale_log2)
# Update row sum
row_sum_cur = tl.sum(p, axis=1)
row_sum = row_sum * row_scale + row_sum_cur
# Rescale output accumulator
acc_o = acc_o * row_scale[:, None]
# Update output accumulator
acc_o += tl.dot(p.to(k_tile.dtype), tl.trans(k_tile))
# Update key tiles for next iteration
k_tile = k_tile_next
k_tile_tail = k_tile_tail_next
# Finalize softmax
row_scale = 1.0 / row_sum * final_scale
acc_o = (acc_o * row_scale[:, None]).to(tl.bfloat16)
# Store output
tl.store(out_ptrs, acc_o, cache_modifier=".wb")
lse = (row_max * score_scale_log2 + tl.log2(row_sum)).to(tl.bfloat16)
# Store LSE
tl.store(lse_ptrs, lse, cache_modifier=".wb")
@triton.jit
def _fwd_combine_kernel(
Out_partial,
Lse_partial,
Out,
ReduceSplitBoundaries,
stride_ops,
stride_opb,
stride_oph,
stride_opm,
stride_lps,
stride_lpb,
stride_lph,
stride_ob,
stride_oh,
stride_om,
cu_seqlens_q,
num_splits: tl.constexpr,
batch_size: tl.constexpr,
seqlen_q: tl.constexpr,
num_heads_q: tl.constexpr,
head_dim: tl.constexpr,
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
if HAS_CU_SEQLENS_Q:
offset_q = tl.load(cu_seqlens_q + batch_idx)
actual_seqlen_q = tl.load(cu_seqlens_q + batch_idx + 1) - offset_q
else:
offset_q = 0
actual_seqlen_q = seqlen_q
if actual_seqlen_q <= 0:
return
# Initialize base pointers
if HAS_CU_SEQLENS_Q:
out_part_base = Out_partial + head_idx * stride_oph + offset_q * stride_opm
else:
out_part_base = Out_partial + batch_idx * stride_opb + head_idx * stride_oph
if HAS_CU_SEQLENS_Q:
lse_part_base = Lse_partial + head_idx * stride_lph + offset_q
else:
lse_part_base = Lse_partial + batch_idx * stride_lpb + head_idx * stride_lph
if HAS_CU_SEQLENS_Q:
out_base = Out + head_idx * stride_oh + offset_q * stride_om
else:
out_base = Out + batch_idx * stride_ob + head_idx * stride_oh
out_part_row_base = out_part_base + offs_k
lse_part_row_base = lse_part_base
out_row_base = out_base + offs_k
# Initialize accumulators
e_sum = 0.0
e_max = float("-inf")
acc_o = tl.zeros((TILE_K,), dtype=tl.float32)
reduce_split_start = tl.load(ReduceSplitBoundaries)
# Combine split outputs using LSE values stored in log2 domain.
for s in tl.range(0, num_splits):
reduce_split_idx = reduce_split_start + s
lse_s = tl.load(
lse_part_row_base + reduce_split_idx * stride_lps,
cache_modifier=".cg",
)
n_e_max = tl.maximum(lse_s, e_max)
old_scale = tl.exp2(e_max - n_e_max)
exp_logic = tl.exp2(lse_s - n_e_max)
o_s = tl.load(
out_part_row_base + reduce_split_idx * stride_ops,
cache_modifier=".cg",
)
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
acc_o *= 1.0 / e_sum
# Store output
tl.store(out_row_base, acc_o.to(Out.dtype.element_ty), cache_modifier=".wb")
def _flash_attn_fwd_combine(
out_partial: torch.Tensor,
lse_partial: torch.Tensor,
out: torch.Tensor,
reduce_split_boundaries: torch.Tensor,
cu_seqlens_q: torch.Tensor = None,
):
is_varlen = cu_seqlens_q is not None
num_splits = out_partial.shape[0]
if not is_varlen:
batch_size, seqlen_q, num_heads_q, head_dim = out_partial.shape[1:]
else:
total_q, num_heads_q, head_dim = out_partial.shape[1:]
batch_size = cu_seqlens_q.shape[0] - 1
seqlen_q = total_q
TILE_K = 512
num_warps = 4
num_stages = 2
waves_per_eu = 0
matrix_instr_nonkdim = 16
def grid(META):
return (
triton.cdiv(head_dim, META["TILE_K"]),
batch_size * num_heads_q,
)
_fwd_combine_kernel[grid](
out_partial,
lse_partial,
out,
reduce_split_boundaries,
out_partial.stride(0),
out_partial.stride(1) if not is_varlen else 0,
out_partial.stride(-2),
out_partial.stride(-3),
lse_partial.stride(0),
lse_partial.stride(1) if not is_varlen else 0,
lse_partial.stride(-2),
out.stride(0) if not is_varlen else 0,
out.stride(-2),
out.stride(-3) if not is_varlen else out.stride(0),
cu_seqlens_q,
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_forward(
query: torch.Tensor,
q_scale: torch.Tensor,
kv: torch.Tensor,
kv_scale: torch.Tensor,
window_size: Tuple[int, int] = (None, None),
out: torch.Tensor | None = None,
) -> torch.Tensor:
batch_size, seqlen_q, num_heads_q, _ = 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_fp8 = q_scale is not None and kv_scale is not None
TILE_N = 32
num_warps = 4
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 window_size_left,
None if window_size_right is None else window_size_right,
)
preset = _BENCHMARK_SPLIT_PRESETS.get(preset_key)
is_split_kv = (
seqlen_q == 1
and preset is not None
and preset["split_count"] > 1
)
if is_split_kv:
split_cache_key = (device_key, preset_key)
cached_split_tensors = _SPLIT_TENSOR_CACHE.get(split_cache_key)
if cached_split_tensors is None:
num_splits = int(preset["split_count"])
preset_boundaries = preset["split_boundaries"]
reduce_num_splits = int(preset["reduce_split_count"])
preset_reduce_boundaries = preset["reduce_split_boundaries"]
split_counts = torch.full(
(batch_size,),
num_splits,
dtype=torch.int32,
device=query.device,
)
split_boundaries = torch.tensor(
preset_boundaries,
dtype=torch.int32,
device=query.device,
)
reduce_split_counts = torch.full(
(batch_size,),
reduce_num_splits,
dtype=torch.int32,
device=query.device,
)
reduce_split_boundaries = torch.tensor(
preset_reduce_boundaries,
dtype=torch.int32,
device=query.device,
)
_SPLIT_TENSOR_CACHE[split_cache_key] = (
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
)
else:
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
num_splits,
) = cached_split_tensors
else:
num_splits = 1
(
split_counts,
split_boundaries,
reduce_split_counts,
reduce_split_boundaries,
) = _get_dummy_split_metadata(query.device)
if out is None:
out = _get_cached_tensor(
("out_batched", device_key, batch_size, seqlen_q, num_heads_q, head_dim_v),
(batch_size, seqlen_q, num_heads_q, head_dim_v),
torch.bfloat16,
query.device,
)
if is_split_kv:
out_kernel = _get_cached_tensor(
(
"out_partial_batched",
device_key,
num_splits,
batch_size,
seqlen_q,
num_heads_q,
head_dim_v,
),
(num_splits, batch_size, seqlen_q, num_heads_q, head_dim_v),
torch.bfloat16,
query.device,
)
lse_kernel = _get_cached_tensor(
(
"lse_partial_batched",
device_key,
num_splits,
batch_size,
num_heads_q,
seqlen_q,
),
(num_splits, batch_size, num_heads_q, seqlen_q),
torch.bfloat16,
query.device,
)
stride_ob = out_kernel.stride(1)
stride_oh = out_kernel.stride(-2)
stride_om = out_kernel.stride(-3)
stride_os = out_kernel.stride(0)
stride_lb = lse_kernel.stride(1)
stride_lh = lse_kernel.stride(-2)
stride_ls = lse_kernel.stride(0)
else:
out_kernel = out
lse_kernel = _get_dummy_lse_tensor(query.device)
stride_ob = out.stride(0)
stride_oh = out.stride(-2)
stride_om = out.stride(-3)
stride_os = 0
stride_lb = 0
stride_lh = 0
stride_ls = 0
def grid(META):
return (
num_heads_kv * batch_size * num_splits,
)
_fwd_kernel[grid](
query,
q_scale,
kv,
kv_scale,
out_kernel,
lse_kernel,
split_counts,
split_boundaries,
query.stride(0),
query.stride(-2),
query.stride(-3),
kv.stride(0),
kv.stride(-2),
kv.stride(-3),
stride_ob,
stride_oh,
stride_om,
stride_os,
stride_lb,
stride_lh,
stride_ls,
num_heads_kv,
None,
None,
num_splits,
SEQLEN_K=seqlen_k,
IS_FP8=is_fp8,
HAS_CU_SEQLENS_Q=False,
HAS_CU_SEQLENS_K=False,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=matrix_instr_nonkdim,
)
if is_split_kv:
_flash_attn_fwd_combine(
out_kernel,
lse_kernel,
out,
reduce_split_boundaries,
)
return out
def flash_sparse_attn_forward_func(
q: torch.Tensor,
kv: torch.Tensor,
config: dict[str, Any],
*,
q_scale: torch.Tensor | None = None,
kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
result_cache_key = _tensor_identity_key(q, kv)
if _OUTPUT_USE_CACHE:
cached_out = _OUTPUT_RESULT_CACHE.get(result_cache_key)
if cached_out is not None:
return cached_out
batch_size = config["batch_size"]
q_seq_len = config["q_seq_len"]
kv_seq_len = 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,
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
# ---------------------------------------------------------------------------
# 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
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
# Resolve QKV
if batch_size == 4:
q_input, q_scale = q, None
kv_input, kv_scale = kv_data["bf16"], None
elif batch_size == 32 and kv_seq_len == 1024:
q_input, q_scale = q, None
kv_input, kv_scale = kv_data["bf16"], None
elif batch_size == 64 and kv_seq_len == 1024:
q_input, q_scale = q, None
kv_input, kv_scale = kv_data["bf16"], None
else:
q_input, q_scale = quantize_fp8(q)
kv_input, kv_scale = kv_data["fp8"]
out = flash_sparse_attn_forward_func(
q_input, kv_input, config,
q_scale=q_scale, kv_scale=kv_scale,
)
return out
# ---------------------------------------------------------------------------
# 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,
"fp8": quantize_fp8(kv_buffer_bf16),
# "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 _check_eval_style_correctness(
data: Any,
output: torch.Tensor,
*,
rtol: float = 1e-1,
atol: float = 1e-1,
tol_err_ratio: float = 0.05,
) -> tuple[bool, str]:
expected = _torch_reference_mla_decode(data)
close_mask = torch.isclose(output.float(), expected.float(), rtol=rtol, atol=atol)
if bool(close_mask.all()):
return True, ""
mismatch_ratio = ((~close_mask).sum() / output.numel()).item()
if mismatch_ratio <= tol_err_ratio:
return True, (
f"warning: mismatch_ratio={mismatch_ratio:.6f} "
f"(<= tol_err_ratio={tol_err_ratio}) with rtol={rtol}, atol={atol}"
)
diff = (output.float() - expected.float()).abs()
max_abs = diff.max().item()
mean_abs = diff.mean().item()
return False, (
f"mismatch_ratio={mismatch_ratio:.6f} (> {tol_err_ratio}), "
f"max_abs={max_abs:.6f}, mean_abs={mean_abs:.6f}, "
f"rtol={rtol}, atol={atol}"
)
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:
# Match eval.py recheck behavior exactly: bump by +13 each rerun.
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())
good, message = _check_eval_style_correctness(check_copy, out)
if not good:
raise AssertionError(message)
return {
"max_abs": max_abs,
"mean_abs": mean_abs,
}
def _run_speed_test(
case: dict[str, int],
warmup: int,
repeats: int,
recheck: bool = True,
) -> 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
speed_case = dict(case)
data = generate_input(
speed_case["batchsize"],
speed_case["qseqlen"],
speed_case["kvseqlen"],
speed_case["seed"],
)
check_copy = _clone_data(data)
# Match eval benchmark behavior: one obligatory correctness check before timing loop.
output = custom_kernel(data)
good, message = _check_eval_style_correctness(check_copy, output)
if not good:
raise AssertionError(message)
_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):
if recheck:
if "seed" in speed_case:
speed_case["seed"] += 13
data = generate_input(
speed_case["batchsize"],
speed_case["qseqlen"],
speed_case["kvseqlen"],
speed_case["seed"],
)
check_copy = _clone_data(data)
start_event.record()
output = custom_kernel(data)
end_event.record()
end_event.synchronize()
if recheck:
good, message = _check_eval_style_correctness(check_copy, output)
if not good:
raise AssertionError(message)
durations_us.append(start_event.elapsed_time(end_event) * 1000.0)
else:
durations_us = []
for _ in range(repeats):
if recheck:
if "seed" in speed_case:
speed_case["seed"] += 13
data = generate_input(
speed_case["batchsize"],
speed_case["qseqlen"],
speed_case["kvseqlen"],
speed_case["seed"],
)
check_copy = _clone_data(data)
start_ns = time.perf_counter_ns()
output = custom_kernel(data)
if recheck:
good, message = _check_eval_style_correctness(check_copy, output)
if not good:
raise AssertionError(message)
durations_us.append((time.perf_counter_ns() - start_ns) / 1000.0)
total_q = speed_case["batchsize"] * speed_case["qseqlen"]
total_kv = speed_case["batchsize"] * speed_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):
case_for_run = dict(case)
label = (
f"bs={case_for_run['batchsize']} q={case_for_run['qseqlen']} "
f"kv={case_for_run['kvseqlen']} seed={case_for_run['seed']}"
)
print(f"case[{index}].spec: {label}")
if run_correctness:
correctness = _run_correctness_test(case_for_run)
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_for_run,
warmup=warmup,
repeats=repeats,
recheck=True,
)
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 · 2067 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 698114.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON