submission 594826
aosudh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 546 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-594826?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:2ea0cda0ecdac7e9964a7bd777c359662f30cecfddc2207302e99eafe39c061c
license declaredunknown
license concludedunknown
authorsaosudh
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"mxfp4": (kv_buffer_mxfp4, kv_scale_mxfp4),online-softmax
m_new = tl.maximum(m_i, score)Kernel source
submission.py546 lines
import torch
from task import input_t, output_t
from utils import make_match_reference
try:
import triton
import triton.language as tl
_TRITON_AVAILABLE = True
except Exception:
triton = None
tl = None
_TRITON_AVAILABLE = False
try:
from aiter import dtypes as aiter_dtypes
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.utility.fp4_utils import dynamic_mxfp4_quant, e8m0_to_f32, mxfp4_to_f32
_AITER_AVAILABLE = True
except Exception:
aiter_dtypes = None
mla_decode_fwd = None
get_mla_metadata_info_v1 = None
get_mla_metadata_v1 = None
dynamic_mxfp4_quant = None
e8m0_to_f32 = None
mxfp4_to_f32 = None
_AITER_AVAILABLE = False
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
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
_META_CACHE = {}
_KV_INDICES_CACHE = {}
_Q_SCALE_CACHE = {}
_KV_LAST_PAGE_LEN_CACHE = {}
def _select_profile(batch_size: int, kv_seq_len: int, q_seq_len: int) -> tuple[int, bool]:
profile_key = (batch_size, kv_seq_len, q_seq_len)
profile_table = {
(4, 1024, 1): (12, True),
(4, 8192, 1): (24, True),
(32, 1024, 1): (12, True),
(32, 8192, 1): (48, True),
(64, 1024, 1): (16, True),
(64, 8192, 1): (72, True),
(256, 1024, 1): (12, False),
(256, 8192, 1): (192, True),
}
if profile_key in profile_table:
return profile_table[profile_key]
if kv_seq_len >= 8192:
num_kv_splits = 96 if batch_size >= 128 else 24
elif batch_size >= 256:
num_kv_splits = 12
elif batch_size >= 64:
num_kv_splits = 16
else:
num_kv_splits = 12
return num_kv_splits, batch_size <= 128
FP8_DTYPE = aiter_dtypes.fp8 if _AITER_AVAILABLE else (
torch.float8_e4m3fn if hasattr(torch, "float8_e4m3fn") else torch.float16
)
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
def _quantize_fp8_with_cached_scale(
tensor: torch.Tensor,
cache_key: tuple,
) -> tuple[torch.Tensor, torch.Tensor]:
cached = _Q_SCALE_CACHE.get(cache_key)
if cached is None:
fp8_tensor, scale = quantize_fp8(tensor)
_Q_SCALE_CACHE[cache_key] = scale
return fp8_tensor, scale
finfo = torch.finfo(FP8_DTYPE)
fp8_tensor = (tensor / cached).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, cached
def _fallback_quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
b, m, n = tensor.shape
assert n % 32 == 0
x = tensor.reshape(b * m, n).to(torch.float32)
blocks = x.view(b * m, n // 32, 32)
amax = blocks.abs().amax(dim=-1).clamp(min=1e-8)
scale = amax / 6.0
q = torch.clamp(torch.round(blocks / scale.unsqueeze(-1) * 2.0), -12.0, 12.0) / 2.0
q = q.view(b * m, n)
q2 = q.view(b * m, n // 2, 2)
lo = torch.clamp((q2[..., 0] * 2).round().to(torch.int32) + 8, 0, 15).to(torch.uint8)
hi = torch.clamp((q2[..., 1] * 2).round().to(torch.int32) + 8, 0, 15).to(torch.uint8)
packed = (lo | (hi << 4)).contiguous()
return packed.view(b, m, n // 2), scale
def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if _AITER_AVAILABLE:
b, m, n = tensor.shape
fp4_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor.reshape(b * m, n))
return fp4_2d.view(b, m, n // 2), scale_e8m0
return _fallback_quantize_mxfp4(tensor)
def dequantize_mxfp4(
fp4_data: torch.Tensor,
scale_e8m0: torch.Tensor,
orig_shape: tuple[int, int, int],
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
b, m, n = orig_shape
rows = b * m
if _AITER_AVAILABLE:
float_vals = mxfp4_to_f32(fp4_data.reshape(rows, n // 2))
scale_f32 = e8m0_to_f32(scale_e8m0)[:, : n // 32]
return (float_vals.view(rows, n // 32, 32) * scale_f32[:rows].unsqueeze(-1)).view(b, m, n).to(dtype)
packed = fp4_data.reshape(rows, n // 2).to(torch.uint8)
lo = (packed & 0x0F).to(torch.int32) - 8
hi = ((packed >> 4) & 0x0F).to(torch.int32) - 8
x = torch.empty((rows, n), device=fp4_data.device, dtype=torch.float32)
x[:, 0::2] = lo.to(torch.float32) * 0.5
x[:, 1::2] = hi.to(torch.float32) * 0.5
s = scale_e8m0[:, : n // 32].to(torch.float32)
return (x.view(rows, n // 32, 32) * s[:rows].unsqueeze(-1)).view(b, m, n).to(dtype)
def _get_mxfp4_scale_f32(scale_e8m0: torch.Tensor, total_kv: int) -> torch.Tensor:
if _AITER_AVAILABLE:
return e8m0_to_f32(scale_e8m0)[:total_kv, : QK_HEAD_DIM // 32].contiguous()
return scale_e8m0[:total_kv, : QK_HEAD_DIM // 32].to(torch.float32).contiguous()
def _get_fp4x2_uint8(fp4_data: torch.Tensor) -> torch.Tensor:
if fp4_data.dtype == torch.uint8:
out = fp4_data
else:
out = fp4_data.view(torch.uint8)
if out.dim() == 3:
out = out.squeeze(1)
return out.contiguous()
def _make_mla_decode_metadata(
batch_size: int,
max_q_len: int,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_last_page_len: torch.Tensor,
num_kv_splits: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
):
info = get_mla_metadata_info_v1(
batch_size,
max_q_len,
NUM_HEADS,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=True,
num_kv_splits=num_kv_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=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=True,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return {
"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,
}
def _aiter_mla_decode_fp8(
q: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
num_kv_splits, q_use_fp8 = _select_profile(batch_size, kv_seq_len, q_seq_len)
q_cache_key = (q.device.index, batch_size, kv_seq_len, q_seq_len)
if q_use_fp8:
q_prebuilt = config.get("_q_fp8_prebuilt")
if q_prebuilt is not None and q_prebuilt[0].shape == q.shape:
q_input, q_scale = q_prebuilt
else:
q_input, q_scale = _quantize_fp8_with_cached_scale(q, q_cache_key)
q_dtype = FP8_DTYPE
else:
q_input, q_scale = q, None
q_dtype = q.dtype
total_kv_len = int(kv_indptr[-1].item())
is_uniform = (q.shape[0] == batch_size * q_seq_len) and (total_kv_len == batch_size * kv_seq_len)
kv_indices_key = (q.device.index, total_kv_len)
kv_indices = _KV_INDICES_CACHE.get(kv_indices_key)
if kv_indices is None:
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=q.device)
_KV_INDICES_CACHE[kv_indices_key] = kv_indices
kv_buffer_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
kv_last_page_key = (q.device.index, batch_size, kv_seq_len, is_uniform)
kv_last_page_len = _KV_LAST_PAGE_LEN_CACHE.get(kv_last_page_key)
if kv_last_page_len is None:
if is_uniform:
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=q.device)
else:
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
_KV_LAST_PAGE_LEN_CACHE[kv_last_page_key] = kv_last_page_len
prebuilt_meta_key = config.get("_prebuilt_meta_key")
if prebuilt_meta_key is not None and prebuilt_meta_key in _META_CACHE:
meta = _META_CACHE[prebuilt_meta_key]
else:
meta = None
meta_key = (q.device.index, batch_size, q_seq_len, kv_seq_len, num_kv_splits, is_uniform, q_dtype)
if meta is None:
meta = _META_CACHE.get(meta_key)
if meta is None:
if is_uniform:
qo_meta = torch.arange(0, batch_size + 1, dtype=torch.int32, device=q.device) * q_seq_len
kv_meta = torch.arange(0, batch_size + 1, dtype=torch.int32, device=q.device) * kv_seq_len
kv_last_meta = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=q.device)
else:
qo_meta = qo_indptr
kv_meta = kv_indptr
kv_last_meta = kv_last_page_len
meta = _make_mla_decode_metadata(
batch_size,
q_seq_len,
qo_meta,
kv_meta,
kv_last_meta,
num_kv_splits,
q_dtype,
kv_fp8.dtype,
)
_META_CACHE[meta_key] = meta
o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
mla_decode_fwd(
q_input.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return o
def _torch_mla_decode_mxfp4(
q: torch.Tensor,
kv_packed: torch.Tensor,
scale_f32: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
total_q = q.shape[0]
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
kv_u8 = kv_packed
kv_rows = kv_u8.shape[0]
kv_bytes = kv_u8.reshape(kv_rows, -1)
dim = torch.arange(QK_HEAD_DIM, device=q.device, dtype=torch.int32)
byte_idx = torch.div(dim, 2, rounding_mode="floor")
nibble_sel = (dim & 1) == 0
batch = int(config["batch_size"])
qo_host = qo_indptr.detach().cpu().tolist()
kv_host = kv_indptr.detach().cpu().tolist()
lut = torch.tensor(
[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],
device=q.device,
dtype=torch.float32,
)
for b in range(batch):
qs, qe = qo_host[b], qo_host[b + 1]
ks, ke = kv_host[b], kv_host[b + 1]
if qs == qe:
continue
if ke == ks:
out[qs:qe].zero_()
continue
packed = kv_bytes[ks:ke]
scales = scale_f32[ks:ke]
packed_dim = packed[:, byte_idx]
nibbles = torch.where(nibble_sel, packed_dim & 0x0F, packed_dim >> 4)
vals = lut[nibbles.to(torch.int64)]
vals = vals.view(ke - ks, QK_HEAD_DIM // 32, 32) * scales.unsqueeze(-1)
kv = vals.view(ke - ks, QK_HEAD_DIM)
k = kv
v = kv[:, :V_HEAD_DIM]
score = torch.einsum("qhd,kd->qhk", q[qs:qe].to(torch.float32), k.to(torch.float32)) * float(config["sm_scale"])
prob = torch.softmax(score, dim=-1)
out[qs:qe] = torch.einsum("qhk,kd->qhd", prob, v.to(torch.float32)).to(torch.bfloat16)
return out
if _TRITON_AVAILABLE:
@triton.jit
def _mla_decode_mxfp4_kernel(
q_ptr,
kv_ptr,
scale_ptr,
lut_ptr,
q_kv_start_ptr,
q_kv_len_ptr,
o_ptr,
sm_scale,
stride_q0,
stride_q1,
stride_q2,
stride_kv0,
stride_kv1,
stride_s0,
stride_s1,
stride_o0,
stride_o1,
stride_o2,
N_HEADS: tl.constexpr,
H_BLOCK: tl.constexpr,
V_DIM: tl.constexpr,
NUM_K_BLOCKS: tl.constexpr,
QK_DIM: tl.constexpr,
):
pid0 = tl.program_id(0)
n_hg = (N_HEADS + H_BLOCK - 1) // H_BLOCK
q_idx = pid0 // n_hg
hg = pid0 % n_hg
h = hg * H_BLOCK + tl.arange(0, H_BLOCK)
h_mask = h < N_HEADS
dv = tl.arange(0, V_DIM)
dv_mask = dv < V_DIM
kv_start = tl.load(q_kv_start_ptr + q_idx).to(tl.int32)
kv_len = tl.load(q_kv_len_ptr + q_idx).to(tl.int32)
m_i = -float("inf")
m_i = tl.full((H_BLOCK,), m_i, dtype=tl.float32)
l_i = tl.zeros((H_BLOCK,), dtype=tl.float32)
acc = tl.zeros((H_BLOCK, V_DIM), dtype=tl.float32)
kv_rel = 0
while kv_rel < kv_len:
kv_idx = kv_start + kv_rel
score = tl.zeros((H_BLOCK,), dtype=tl.float32)
for kb in range(NUM_K_BLOCKS):
d = kb * 32 + tl.arange(0, 32)
q_ptrs = q_ptr + q_idx * stride_q0 + h[:, None] * stride_q1 + d[None, :] * stride_q2
qv = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0).to(tl.float32)
byte_idx = d // 2
packed = tl.load(kv_ptr + kv_idx * stride_kv0 + byte_idx * stride_kv1).to(tl.uint8)
nibble = tl.where((d & 1) == 0, packed & 0x0F, packed >> 4).to(tl.int32)
kval = tl.load(lut_ptr + nibble).to(tl.float32)
s = tl.load(scale_ptr + kv_idx * stride_s0 + kb * stride_s1).to(tl.float32)
score += tl.sum(qv * (kval * s)[None, :], axis=1)
score = score * sm_scale
byte_v = dv // 2
packed_v = tl.load(kv_ptr + kv_idx * stride_kv0 + byte_v * stride_kv1, mask=dv_mask, other=0).to(tl.uint8)
nibble_v = tl.where((dv & 1) == 0, packed_v & 0x0F, packed_v >> 4).to(tl.int32)
v = tl.load(lut_ptr + nibble_v, mask=dv_mask, other=0.0).to(tl.float32)
sb = dv // 32
sv = tl.load(scale_ptr + kv_idx * stride_s0 + sb * stride_s1, mask=dv_mask, other=0.0).to(tl.float32)
v = v * sv
m_new = tl.maximum(m_i, score)
alpha = tl.exp(m_i - m_new)
beta = tl.exp(score - m_new)
acc = acc * alpha[:, None] + beta[:, None] * v[None, :]
l_i = l_i * alpha + beta
m_i = m_new
kv_rel += 1
out = acc / l_i[:, None]
o_ptrs = o_ptr + q_idx * stride_o0 + h[:, None] * stride_o1 + dv[None, :] * stride_o2
tl.store(o_ptrs, out.to(tl.bfloat16), mask=h_mask[:, None] & dv_mask[None, :])
def _triton_mla_decode_mxfp4(
q: torch.Tensor,
kv_packed: torch.Tensor,
scale_f32: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
if not _TRITON_AVAILABLE:
return _torch_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)
total_q = q.shape[0]
qo_sizes = (qo_indptr[1:] - qo_indptr[:-1]).to(torch.int32)
kv_sizes = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
kv_starts = kv_indptr[:-1].to(torch.int32)
q_kv_start = torch.repeat_interleave(kv_starts, qo_sizes).contiguous()
q_kv_len = torch.repeat_interleave(kv_sizes, qo_sizes).contiguous()
if q_kv_start.numel() != total_q:
return _torch_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)
o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)
lut = torch.tensor(
[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],
device=q.device,
dtype=torch.float32,
)
h_block = 4
n_hg = (NUM_HEADS + h_block - 1) // h_block
grid = (total_q * n_hg,)
_mla_decode_mxfp4_kernel[grid](
q,
kv_packed,
scale_f32,
lut,
q_kv_start,
q_kv_len,
o,
float(config["sm_scale"]),
q.stride(0),
q.stride(1),
q.stride(2),
kv_packed.stride(0),
kv_packed.stride(1),
scale_f32.stride(0),
scale_f32.stride(1),
o.stride(0),
o.stride(1),
o.stride(2),
N_HEADS=NUM_HEADS,
H_BLOCK=h_block,
V_DIM=V_HEAD_DIM,
NUM_K_BLOCKS=QK_HEAD_DIM // 32,
QK_DIM=QK_HEAD_DIM,
)
return o
def generate_input(batchsize: int, qseqlen: int, kvseqlen: int, seed: int) -> input_t:
gen = torch.Generator(device="cuda")
gen.manual_seed(seed)
total_q = batchsize * qseqlen
total_kv = batchsize * kvseqlen
q = torch.randn((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=torch.bfloat16, device="cuda", generator=gen)
kv_buffer_bf16 = torch.randn((total_kv, NUM_KV_HEADS, QK_HEAD_DIM), dtype=torch.bfloat16, device="cuda", generator=gen)
kv_buffer_fp8, kv_scale_fp8 = quantize_fp8(kv_buffer_bf16)
kv_buffer_mxfp4, kv_scale_mxfp4 = quantize_mxfp4(kv_buffer_bf16)
kv_data = {
"bf16": kv_buffer_bf16,
"fp8": (kv_buffer_fp8, kv_scale_fp8),
"mxfp4": (kv_buffer_mxfp4, kv_scale_mxfp4),
}
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,
}
if _AITER_AVAILABLE:
num_kv_splits, q_use_fp8 = _select_profile(batchsize, kvseqlen, qseqlen)
if q_use_fp8:
config["_q_fp8_prebuilt"] = quantize_fp8(q)
kv_last_page_len = torch.full((batchsize,), kvseqlen, dtype=torch.int32, device="cuda")
meta_key = ("prebuilt", q.device.index, batchsize, qseqlen, kvseqlen, num_kv_splits, q.dtype if not q_use_fp8 else FP8_DTYPE)
if meta_key not in _META_CACHE:
_META_CACHE[meta_key] = _make_mla_decode_metadata(
batchsize,
qseqlen,
qo_indptr,
kv_indptr,
kv_last_page_len,
num_kv_splits,
FP8_DTYPE if q_use_fp8 else q.dtype,
kv_buffer_fp8.dtype,
)
config["_prebuilt_meta_key"] = meta_key
return (q, kv_data, qo_indptr, kv_indptr, config)
def ref_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
if _AITER_AVAILABLE and "fp8" in kv_data:
kv_fp8, kv_scale = kv_data["fp8"]
return _aiter_mla_decode_fp8(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config)
kv_fp4, kv_scale_e8m0 = kv_data["mxfp4"]
kv_packed = _get_fp4x2_uint8(kv_fp4)
scale_f32 = _get_mxfp4_scale_f32(kv_scale_e8m0, kv_packed.shape[0])
return _triton_mla_decode_mxfp4(q, kv_packed, scale_f32, qo_indptr, kv_indptr, config)
def custom_kernel(data: input_t) -> output_t:
return ref_kernel(data)
check_implementation = make_match_reference(ref_kernel, rtol=5e-3, atol=5e-3)
scrolls · 546 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