submission 683270
Aniket Sadashiva · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 241 lines, June 9 Researcher Reciprocity License v1.0.
submission_probe_v362_morefp8_p416_shapeamax13.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-683270?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:d7f6d8dcaef6dc4d3b6e30be0cd5e1c6881abb03e8aeec38db9872aadf1bbaa1
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
tl.store(out_ptr + offsets, x.to(tl.float8e4nv), mask=mask)Kernel source
submission_probe_v362_morefp8_p416_shapeamax13.py241 lines
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min
DEFAULT_Q_AMAX = 0.12
Q_AMAX_TABLE = {
(32, 1024): 0.13,
}
BF16_NP_SHAPES = {(4, 1024),}
FP8_SPLITS_TABLE = {
(4, 8192): 16,
(32, 1024): 4,
(32, 8192): 32,
(64, 1024): 16,
(64, 8192): 32,
(256, 1024): 16,
(256, 8192): 32,
}
_cache = {}
_stage1 = None
_reduce = None
_q_scale_cache = {}
_q_fp8_bufs = {}
def _compute_np_splits(bs, kv_seq_len):
cu_num = 304
overhead = 84.1
best_score = -1
best_i = 1
for i in range(1, 17):
waves = ((bs * i + cu_num - 1) // cu_num) * cu_num
score = (bs * i / waves) * kv_seq_len / (kv_seq_len + overhead * i)
if score > best_score:
best_score = score
best_i = i
return best_i
@triton.jit
def _reduce_kernel(sd_ptr, sl_ptr, o_ptr, num_valid, NUM_SPLITS: tl.constexpr):
bid = tl.program_id(0)
hid = tl.program_id(1)
offs_d = tl.arange(0, 512)
sd_base = (bid * NUM_SPLITS * 16 * 512 + hid * 512).to(tl.int64)
sl_base = (bid * NUM_SPLITS * 16 + hid).to(tl.int64)
e_max = -float("inf")
e_sum = 0.0
acc = tl.zeros((512,), dtype=tl.float32)
for s in range(NUM_SPLITS):
if s < num_valid:
v = tl.load(sd_ptr + sd_base + s * 16 * 512 + offs_d)
lse = tl.load(sl_ptr + sl_base + s * 16)
new_max = tl.maximum(e_max, lse)
old_scale = tl.exp(e_max - new_max)
new_scale = tl.exp(lse - new_max)
acc = acc * old_scale + new_scale * v
e_sum = e_sum * old_scale + new_scale
e_max = new_max
tl.store(o_ptr + (bid * 16 * 512 + hid * 512).to(tl.int64) + offs_d, (acc / e_sum).to(tl.bfloat16))
@triton.jit
def _quant_q_kernel(
q_ptr, out_ptr, n_elements, inv_scale,
fp8_min: tl.constexpr, fp8_max: tl.constexpr, block_size: tl.constexpr,
):
pid = tl.program_id(0)
offsets = pid * block_size + tl.arange(0, block_size)
mask = offsets < n_elements
x = tl.load(q_ptr + offsets, mask=mask).to(tl.float32)
x = x * inv_scale
x = tl.minimum(tl.maximum(x, fp8_min), fp8_max)
tl.store(out_ptr + offsets, x.to(tl.float8e4nv), mask=mask)
def _get_bf16_np_cached(batch_size, kv_seq_len):
key = ("bf16np", batch_size, kv_seq_len)
if key in _cache:
return _cache[key]
total_kv = batch_size * kv_seq_len
ns = _compute_np_splits(batch_size, kv_seq_len)
c = {
"output": torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"qo_indptr": torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"),
"kv_indptr": torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * kv_seq_len,
"kv_indices": torch.arange(total_kv, dtype=torch.int32, device="cuda"),
"kv_last_page_len": torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda"),
"ns": ns,
"num_kv_splits_indptr": torch.arange(0, (batch_size + 1) * ns, ns, dtype=torch.int32, device="cuda"),
"logits": torch.empty((batch_size, ns, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((batch_size, ns, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
}
_cache[key] = c
return c
def _get_fp8_ps_cached(batch_size, kv_seq_len, num_splits):
key = ("fp8ps", batch_size, kv_seq_len, num_splits)
if key in _cache:
return _cache[key]
total_kv = batch_size * kv_seq_len
qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * kv_seq_len
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
batch_size, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=num_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=False,
max_split_per_batch=num_splits, intra_batch_mode=True,
dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
)
logits = torch.empty(
(reduce_partial_map.size(0), 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda",
)
attn_lse = torch.empty(
(reduce_partial_map.size(0), 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda",
)
c = {
"output": torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
"logits": logits,
"attn_lse": attn_lse,
}
_cache[key] = c
return c
def custom_kernel(data: input_t) -> output_t:
global _stage1, _reduce, _q_scale_cache
q, kv_data, _, _, config = data
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
if _stage1 is None:
_stage1 = aiter.mla_decode_stage1_asm_fwd
_reduce = aiter.mla_reduce_v1
if (batch_size, kv_seq_len) in BF16_NP_SHAPES:
c = _get_bf16_np_cached(batch_size, kv_seq_len)
kv_4d = kv_data["bf16"].view(-1, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
_stage1(
q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,
c["qo_indptr"], c["kv_indptr"], c["kv_indices"], c["kv_last_page_len"],
c["num_kv_splits_indptr"], None, None, None,
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["output"],
q_scale=None, kv_scale=None,
)
ns = c["ns"]
_reduce_kernel[(batch_size, NUM_HEADS)](
c["logits"].view(batch_size * ns, NUM_HEADS, V_HEAD_DIM),
c["attn_lse"].view(batch_size * ns, NUM_HEADS),
c["output"], ns, NUM_SPLITS=ns,
)
return c["output"]
if batch_size not in _q_fp8_bufs:
_q_fp8_bufs[batch_size] = torch.empty(
(batch_size, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda"
)
q_fp8 = _q_fp8_bufs[batch_size]
n_elements = batch_size * NUM_HEADS * QK_HEAD_DIM
q_amax = Q_AMAX_TABLE.get((batch_size, kv_seq_len), DEFAULT_Q_AMAX)
q_scale = _q_scale_cache.get(q_amax)
if q_scale is None:
q_scale = torch.tensor([q_amax / FP8_MAX], dtype=torch.float32, device="cuda")
_q_scale_cache[q_amax] = q_scale
_quant_q_kernel[((n_elements + 1023) // 1024,)](
q, q_fp8, n_elements, FP8_MAX / q_amax,
fp8_min=FP8_MIN, fp8_max=FP8_MAX, block_size=1024,
)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
num_splits = FP8_SPLITS_TABLE[(batch_size, kv_seq_len)]
c = _get_fp8_ps_cached(batch_size, kv_seq_len, num_splits)
_stage1(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_buffer_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM),
c["qo_indptr"], c["kv_indptr"],
c["kv_indices"], c["kv_last_page_len"],
None,
c["work_meta_data"], c["work_indptr"], c["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["output"],
q_scale=q_scale, kv_scale=kv_scale,
)
_reduce(
c["logits"], c["attn_lse"],
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
1, c["output"], None,
)
return c["output"]
scrolls · 241 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON