Skip to content
KernelIndex
Search⌘K

submission 644275

nawal-yantrion · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 208 lines, June 9 Researcher Reciprocity License v1.0.

sub_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-644275?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
4.51ms
#748 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6757f69a224868c99b2f30ae931577256f9309486c59b51b1b82b83b668b02b2
license declaredunknown
license concludedunknown
authorsnawal-yantrion
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4elif QKV_DTYPE == "mxfp4":

Kernel source

sub_v2.py208 lines
import torch
import torch.nn.functional as F
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes

FP8_DTYPE = aiter_dtypes.fp8

# QKV dtype for custom_kernel dispatch: "bf16", "fp8", or "mxfp4"
QKV_DTYPE = "fp8"


# ---------------------------------------------------------------------------
# Dispatcher: select kernel based on QKV_DTYPE
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    """Dispatch to the appropriate kernel based on QKV_DTYPE."""
    if QKV_DTYPE == "fp8":
        return custom_kernel_fp8(data)
    elif QKV_DTYPE == "bf16":
        return custom_kernel_bf16(data)
    elif QKV_DTYPE == "mxfp4":
        return custom_kernel_mxfp4(data)
    else:
        raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")


# ---------------------------------------------------------------------------
# FP8 quantization helper (per-tensor, sglang style)
# ---------------------------------------------------------------------------
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Dynamic per-tensor FP8 quantization. Returns (fp8_tensor, scale)."""
    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)


# ---------------------------------------------------------------------------
# FP8 Q + FP8 KV — optimized FP8 attention kernel using torch._scaled_mm
# ---------------------------------------------------------------------------
def custom_kernel_fp8(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    num_heads = config["num_heads"]
    kv_lora_rank = config["kv_lora_rank"]
    qk_head_dim = config["qk_head_dim"]
    sm_scale = config["sm_scale"]

    # FP8 KV buffer and scale
    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)    # (total_kv, 576) fp8

    # Quantize Q to fp8
    q_fp8, q_scale = quantize_fp8(q)                     # q_fp8: (total_q, 16, 576) fp8

    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    scale_one = torch.ones(1, dtype=torch.float32, device="cuda")

    # Use torch.jit.script to optimize the loop, reducing Python overhead
    for i in range(batch_size):
        q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
        kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
        seq_q = q_e - q_s
        seq_kv = kv_e - kv_s

        # Q: (seq_q * nhead, 576) fp8,  K: (seq_kv, 576) fp8
        qi_fp8 = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim)   # (seq_q*16, 576)
        ki_fp8 = kv_fp8_2d[kv_s:kv_e]                                      # (seq_kv, 576)

        # QK^T via _scaled_mm: (seq_q*16, 576) @ (seq_kv, 576).T -> (seq_q*16, seq_kv)
        # _scaled_mm expects (M,K) @ (N,K).T  where b is row-major contiguous
        raw_scores = torch._scaled_mm(
            qi_fp8, ki_fp8.t(),
            scale_a=q_scale, scale_b=kv_scale_fp8,
            out_dtype=torch.float32,
        )
        # raw_scores: (seq_q*16, seq_kv)
        scores = raw_scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)  # (nhead, seq_q, seq_kv)
        scores = scores * sm_scale
        scores = F.softmax(scores, dim=-1)

        # V: first 512 dims of KV buffer (bf16 for softmax@V since scores are float)
        kv_bf16 = kv_data["bf16"]
        vi = kv_bf16[kv_s:kv_e, 0, :kv_lora_rank].float()  # (seq_kv, 512)

        # softmax @ V: (nhead, seq_q, seq_kv) @ (seq_kv, 512) -> (nhead, seq_q, 512)
        oi = torch.matmul(scores, vi)
        oi = oi.permute(1, 0, 2)                             # (seq_q, nhead, 512)
        out_list.append(oi.to(torch.bfloat16))

    return torch.cat(out_list, dim=0)


# ---------------------------------------------------------------------------
# Baseline: bf16 Q + bf16 KV — naive torch attention
# ---------------------------------------------------------------------------
def custom_kernel_bf16(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    num_heads = config["num_heads"]
    kv_lora_rank = config["kv_lora_rank"]
    sm_scale = config["sm_scale"]

    # This naive baseline uses bf16 KV directly.
    kv_buffer_bf16 = kv_data["bf16"]

    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
        kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())

        qi = q[q_s:q_e]                       # (seq_q, nhead, 576)
        kvc = kv_buffer_bf16[kv_s:kv_e, 0]    # (seq_kv, 576)  squeeze kv_heads dim

        # Key: full 576 dims; Value: first 512 dims (kv_lora_rank)
        ki = kvc                       # (seq_kv, 576) — broadcast over heads
        vi = kvc[:, :kv_lora_rank]     # (seq_kv, 512)

        # Attention: (nhead, seq_q, 576) @ (576, seq_kv) → (nhead, seq_q, seq_kv)
        qi_t = qi.float().permute(1, 0, 2)  # (nhead, seq_q, 576)
        scores = torch.matmul(qi_t * sm_scale, ki.float().T)  # (nhead, seq_q, seq_kv)

        scores = F.softmax(scores, dim=-1)

        # Output: (nhead, seq_q, seq_kv) @ (seq_kv, 512) → (nhead, seq_q, 512)
        oi = torch.matmul(scores, vi.float())  # (nhead, seq_q, 512)
        oi = oi.permute(1, 0, 2)               # (seq_q, nhead, 512)
        out_list.append(oi.to(torch.bfloat16))

    return torch.cat(out_list, dim=0)


def unpack_mxfp4(packed: torch.Tensor) -> torch.Tensor:
    """
    Unpack fp4x2 → int8 values in range [-8, 7]
    packed: (..., 288) uint8
    returns: (..., 576) int8
    """
    high = (packed >> 4) & 0xF
    low = packed & 0xF

    # convert to signed [-8, 7]
    high = high.to(torch.int8) - 8
    low = low.to(torch.int8) - 8

    return torch.stack((high, low), dim=-1).reshape(*packed.shape[:-1], -1)


def custom_kernel_mxfp4(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    num_heads = config["num_heads"]
    kv_lora_rank = config["kv_lora_rank"]   # 512
    qk_head_dim = config["qk_head_dim"]     # 576
    sm_scale = config["sm_scale"]

    # MXFP4 KV
    kv_packed, kv_scale = kv_data["mxfp4"]   # (kv,1,288), (scale)
    kv_packed = kv_packed.view(-1, 288)

    # Quantize Q → fp8
    q_fp8, q_scale = quantize_fp8(q)

    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s, q_e = int(qo_indptr[i]), int(qo_indptr[i + 1])
        kv_s, kv_e = int(kv_indptr[i]), int(kv_indptr[i + 1])

        seq_q = q_e - q_s
        seq_kv = kv_e - kv_s

        qi = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim)

        # ---- UNPACK + DEQUANT (FUSED STYLE) ----
        kv_chunk = kv_packed[kv_s:kv_e]                 # (seq_kv, 288)
        kv_int8 = unpack_mxfp4(kv_chunk)                # (seq_kv, 576)

        # apply scale (fp8 → fp32)
        kv_fp32 = kv_int8.float() * kv_scale

        # split K / V without extra copies
        ki = kv_fp32                                    # (seq_kv, 576)
        vi = kv_fp32[:, :kv_lora_rank]                  # (seq_kv, 512)

        # ---- QK ----
        scores = torch.matmul(
            qi.float(), 
            ki.T
        )  # (seq_q * nhead, seq_kv)

        scores = scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)
        scores *= sm_scale
        scores = F.softmax(scores, dim=-1)

        # ---- ATTENTION @ V ----
        oi = torch.matmul(scores, vi)  # (nhead, seq_q, 512)
        oi = oi.permute(1, 0, 2)

        out_list.append(oi.to(torch.bfloat16))

    return torch.cat(out_list, dim=0)
scrolls · 208 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