submission 585963
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 48 lines, June 9 Researcher Reciprocity License v1.0.
sub74_bmm_bf16.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-585963?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:355b53538dd91445bb63b845d890e3b324e70aad1d5831436c9d8dd9e2036dbb
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Kernel source
sub74_bmm_bf16.py48 lines
"""
sub74: Pure PyTorch bmm attention with bf16 KV.
No custom kernel, no aiter. Tests batched GEMM approach.
Advantages:
- Single torch.bmm call for QK^T (hipBLAS uses MFMA internally)
- Single torch.bmm call for OV
- No per-batch loops, no kernel launch overhead
- MQA: K/V naturally broadcast via batched matmul
Disadvantages:
- bf16 KV = 2x bandwidth of fp8
- Materializes full [bs, qseq*nh, kv_seq] scores tensor
"""
import torch
import torch.nn.functional as F
from task import input_t, output_t
SM_SCALE = 1.0 / (576 ** 0.5)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nh = config["num_heads"]
qseq = config["q_seq_len"]
kv_lora_rank = config["kv_lora_rank"]
kv_seq = config["kv_seq_len"]
kv_bf16 = kv_data["bf16"] # [total_kv, 1, 576] bf16
# Reshape for batched matmul (assumes uniform kv_len across batch)
Q = q.view(bs, qseq, nh, 576).reshape(bs, qseq * nh, 576) # [bs, M, 576]
K = kv_bf16.view(bs, kv_seq, 576) # [bs, N, 576]
V = K[:, :, :kv_lora_rank] # [bs, N, 512]
# QK^T: [bs, M, N] via bf16 GEMM (hipBLAS → MFMA)
scores = torch.bmm(Q, K.transpose(1, 2))
# Softmax in fp32 for numerical stability
scores = F.softmax(scores.float() * SM_SCALE, dim=-1)
# OV: [bs, M, 512] via bf16 GEMM
output = torch.bmm(scores.to(torch.bfloat16), V)
return output.view(bs, qseq, nh, 512).reshape(bs * qseq, nh, 512)
scrolls · 48 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