Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
124.3µs
#497 of 766
2026-03-18

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