Skip to content
KernelIndex
Search⌘K

submission 589625

Purple rain · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_amd_mixed_mla_mxfp4_hip_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-589625?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
1.76ms
#734 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7d7a38996bfc85788aa29103e0515b40022d039877e56b5a312789fa114b0c3d
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26

Techniques

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

fp4kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
shared-memory__shared__ float smem_k[kQkHeadDim];
split-kname="mla_mxfp4_splitk_hip_v1",

Kernel source

submission_amd_mixed_mla_mxfp4_hip_splitk.py486 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

import os
import threading

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline


_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None

_CHUNK_ENV = "MLA_MXFP4_CHUNK_SIZE"
_DEFAULT_CHUNK_SIZE = 128


CPP_WRAPPER = r"""
#include <torch/extension.h>

torch::Tensor mla_decode_forward(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor qo_indptr,
    torch::Tensor kv_indptr,
    int64_t config_v_dim,
    int64_t chunk_size,
    double sm_scale);
"""


HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>

#include <cstdint>
#include <cmath>

namespace {

constexpr int kNumHeads = 16;
constexpr int kQkHeadDim = 576;
constexpr int kVHeadDim = 512;
constexpr int kScaleBlock = 32;

__constant__ float kFp4Lut[16] = {
    0.0f,  0.5f,  1.0f,  1.5f,
    2.0f,  3.0f,  4.0f,  6.0f,
   -0.0f, -0.5f, -1.0f, -1.5f,
   -2.0f, -3.0f, -4.0f, -6.0f
};

__device__ __forceinline__ float e8m0_to_float(uint8_t e8m0_val) {
  uint32_t bits = static_cast<uint32_t>(e8m0_val) << 23;
  if (e8m0_val == 0) {
    bits = 0x00400000u;
  } else if (e8m0_val == 0xFF) {
    bits = 0x7F800001u;
  }
  return __uint_as_float(bits);
}

__device__ __forceinline__ float unpack_mxfp4(uint8_t val, float scale, bool high_nibble) {
  uint8_t nibble = high_nibble ? static_cast<uint8_t>(val >> 4) : static_cast<uint8_t>(val & 0x0F);
  return kFp4Lut[nibble & 0x0F] * scale;
}

__device__ __forceinline__ float wavefront_sum_64(float v) {
  for (int offset = 32; offset > 0; offset >>= 1) {
    v += __shfl_down(v, offset);
  }
  return v;
}

__global__ void mla_decode_chunk_kernel(
    const hip_bfloat16* __restrict__ q,
    const uint8_t* __restrict__ kv_data,
    const uint8_t* __restrict__ kv_scale,
    const int32_t* __restrict__ qo_indptr,
    const int32_t* __restrict__ kv_indptr,
    float* __restrict__ out_partials,
    float* __restrict__ lse_partials,
    int kv_row_stride,
    int kv_scale_row_stride,
    int qk_head_dim,
    int v_head_dim,
    float sm_scale,
    int chunk_size,
    int num_chunks) {
  const int chunk_idx = static_cast<int>(blockIdx.x);
  const int batch_idx = static_cast<int>(blockIdx.y);
  const int lane = static_cast<int>(threadIdx.x);
  const int head = static_cast<int>(threadIdx.y);

  if (head >= kNumHeads || chunk_idx >= num_chunks) {
    return;
  }

  const int32_t q_start = qo_indptr[batch_idx];
  const int32_t q_end = qo_indptr[batch_idx + 1];
  if (q_end <= q_start) {
    return;
  }
  const int32_t q_idx = q_start;

  const int32_t kv_start = kv_indptr[batch_idx];
  const int32_t kv_end = kv_indptr[batch_idx + 1];
  const int32_t token_begin = kv_start + chunk_idx * chunk_size;
  const int32_t token_end = token_begin + chunk_size < kv_end ? token_begin + chunk_size : kv_end;

  const hip_bfloat16* q_row = q + (static_cast<int64_t>(q_idx) * kNumHeads + head) * qk_head_dim;

  float q_frag[9];
  #pragma unroll
  for (int i = 0; i < 9; ++i) {
    const int d = lane + i * 64;
    q_frag[i] = (d < qk_head_dim) ? static_cast<float>(q_row[d]) : 0.0f;
  }

  float o_frag[8];
  #pragma unroll
  for (int i = 0; i < 8; ++i) {
    o_frag[i] = 0.0f;
  }

  float m = -INFINITY;
  float l = 0.0f;

  __shared__ float smem_k[kQkHeadDim];
  __shared__ float smem_v[kVHeadDim];

  for (int32_t token = token_begin; token < token_end; ++token) {
    if (head == 0) {
      const uint8_t* kv_row = kv_data + static_cast<int64_t>(token) * kv_row_stride;
      const uint8_t* scale_row = kv_scale + static_cast<int64_t>(token) * kv_scale_row_stride;

      for (int d = lane; d < qk_head_dim; d += 64) {
        const uint8_t packed = kv_row[d >> 1];
        const float scale = e8m0_to_float(scale_row[d / kScaleBlock]);
        const float val = unpack_mxfp4(packed, scale, (d & 1) != 0);
        smem_k[d] = val;
        if (d < v_head_dim) {
          smem_v[d] = val;
        }
      }
    }
    __syncthreads();

    float qk_local = 0.0f;
    #pragma unroll
    for (int i = 0; i < 9; ++i) {
      const int d = lane + i * 64;
      if (d < qk_head_dim) {
        qk_local += q_frag[i] * smem_k[d];
      }
    }

    float score = wavefront_sum_64(qk_local);
    score = (lane == 0) ? score * sm_scale : 0.0f;
    score = __shfl(score, 0);

    const float new_m = fmaxf(m, score);
    const float alpha = expf(m - new_m);
    const float p = expf(score - new_m);
    const float new_l = l * alpha + p;

    #pragma unroll
    for (int i = 0; i < 8; ++i) {
      const int d = lane + i * 64;
      o_frag[i] = o_frag[i] * alpha + p * smem_v[d];
    }

    m = new_m;
    l = new_l;
    __syncthreads();
  }

  const int64_t partial_base =
      (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + chunk_idx) * v_head_dim);

  #pragma unroll
  for (int i = 0; i < 8; ++i) {
    const int d = lane + i * 64;
    out_partials[partial_base + d] = o_frag[i];
  }

  if (lane == 0) {
    const int64_t lse_idx =
        (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + chunk_idx) * 2);
    lse_partials[lse_idx] = m;
    lse_partials[lse_idx + 1] = l;
  }
}

__global__ void mla_decode_reduce_kernel(
    const float* __restrict__ out_partials,
    const float* __restrict__ lse_partials,
    hip_bfloat16* __restrict__ out,
    const int32_t* __restrict__ qo_indptr,
    int num_chunks,
    int v_head_dim) {
  const int batch_idx = static_cast<int>(blockIdx.x);
  const int head = static_cast<int>(blockIdx.y);
  const int tid = static_cast<int>(threadIdx.x);

  __shared__ float s_m;
  __shared__ float s_l;

  if (tid == 0) {
    float m = -INFINITY;
    for (int c = 0; c < num_chunks; ++c) {
      const int64_t lse_idx =
          (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
      const float mc = lse_partials[lse_idx];
      const float lc = lse_partials[lse_idx + 1];
      if (lc > 0.0f) {
        m = fmaxf(m, mc);
      }
    }

    if (!isfinite(m)) {
      s_m = -INFINITY;
      s_l = 0.0f;
    } else {
      float l = 0.0f;
      for (int c = 0; c < num_chunks; ++c) {
        const int64_t lse_idx =
            (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
        const float mc = lse_partials[lse_idx];
        const float lc = lse_partials[lse_idx + 1];
        if (lc > 0.0f) {
          l += lc * expf(mc - m);
        }
      }
      s_m = m;
      s_l = l;
    }
  }
  __syncthreads();

  const int32_t q_idx = qo_indptr[batch_idx];
  const int64_t out_base = (static_cast<int64_t>(q_idx) * kNumHeads + head) * v_head_dim;

  for (int d = tid; d < v_head_dim; d += static_cast<int>(blockDim.x)) {
    float acc = 0.0f;
    if (s_l > 0.0f) {
      for (int c = 0; c < num_chunks; ++c) {
        const int64_t lse_idx =
            (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * 2);
        const float mc = lse_partials[lse_idx];
        const float lc = lse_partials[lse_idx + 1];
        if (lc > 0.0f) {
          const float beta = expf(mc - s_m);
          const int64_t p_idx =
              (((static_cast<int64_t>(batch_idx) * kNumHeads + head) * num_chunks + c) * v_head_dim + d);
          acc += out_partials[p_idx] * beta;
        }
      }
      acc = acc / s_l;
    }
    out[out_base + d] = static_cast<hip_bfloat16>(acc);
  }
}

inline void check_hip_error(const char* where) {
  hipError_t err = hipGetLastError();
  TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}

}  // namespace

torch::Tensor mla_decode_forward(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor qo_indptr,
    torch::Tensor kv_indptr,
    int64_t config_v_dim,
    int64_t chunk_size,
    double sm_scale) {
  TORCH_CHECK(q.is_cuda(), "q must be CUDA/HIP tensor");
  TORCH_CHECK(kv_buffer.is_cuda(), "kv_buffer must be CUDA/HIP tensor");
  TORCH_CHECK(kv_scale.is_cuda(), "kv_scale must be CUDA/HIP tensor");
  TORCH_CHECK(qo_indptr.is_cuda(), "qo_indptr must be CUDA/HIP tensor");
  TORCH_CHECK(kv_indptr.is_cuda(), "kv_indptr must be CUDA/HIP tensor");

  TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "q must be bfloat16");
  TORCH_CHECK(q.dim() == 3, "q must be [total_q, 16, 576]");
  TORCH_CHECK(q.size(1) == kNumHeads, "q num_heads must be 16");
  TORCH_CHECK(q.size(2) == kQkHeadDim, "q head dim must be 576");

  TORCH_CHECK(kv_buffer.element_size() == 1, "kv_buffer must be 1-byte packed fp4x2");
  TORCH_CHECK(kv_buffer.dim() == 3, "kv_buffer must be [total_kv, 1, 288]");
  TORCH_CHECK(kv_buffer.size(1) == 1, "kv_buffer num_kv_heads must be 1");
  TORCH_CHECK(kv_buffer.size(2) == kQkHeadDim / 2, "kv_buffer packed width must be 288");

  TORCH_CHECK(kv_scale.element_size() == 1, "kv_scale must be 1-byte e8m0");
  TORCH_CHECK(kv_scale.dim() == 2, "kv_scale must be [total_kv, 18]");
  TORCH_CHECK(kv_scale.size(1) == kQkHeadDim / kScaleBlock, "kv_scale width must be 18");

  TORCH_CHECK(qo_indptr.scalar_type() == torch::kInt32, "qo_indptr must be int32");
  TORCH_CHECK(kv_indptr.scalar_type() == torch::kInt32, "kv_indptr must be int32");
  TORCH_CHECK(qo_indptr.dim() == 1, "qo_indptr must be 1D");
  TORCH_CHECK(kv_indptr.dim() == 1, "kv_indptr must be 1D");
  TORCH_CHECK(qo_indptr.size(0) == kv_indptr.size(0), "indptr size mismatch");

  TORCH_CHECK(config_v_dim == kVHeadDim, "config_v_dim must be 512");
  TORCH_CHECK(chunk_size > 0, "chunk_size must be > 0");

  q = q.contiguous();
  kv_buffer = kv_buffer.contiguous();
  kv_scale = kv_scale.contiguous();
  qo_indptr = qo_indptr.contiguous();
  kv_indptr = kv_indptr.contiguous();

  int current_device = 0;
  hipDeviceProp_t prop;
  TORCH_CHECK(hipGetDevice(&current_device) == hipSuccess, "hipGetDevice failed");
  TORCH_CHECK(hipGetDeviceProperties(&prop, current_device) == hipSuccess,
              "hipGetDeviceProperties failed");
  TORCH_CHECK(prop.warpSize == 64, "This kernel requires wavefront size 64, got ", prop.warpSize);

  const int64_t total_q = q.size(0);
  const int64_t batch_size = qo_indptr.size(0) - 1;
  auto out = torch::empty({total_q, kNumHeads, config_v_dim}, q.options().dtype(torch::kBFloat16));
  if (total_q == 0 || batch_size <= 0) {
    return out.zero_();
  }

  auto q_lens = qo_indptr.slice(0, 1, batch_size + 1) - qo_indptr.slice(0, 0, batch_size);
  TORCH_CHECK(q_lens.min().item<int64_t>() == 1 && q_lens.max().item<int64_t>() == 1,
              "decode-only kernel requires q_seq_len == 1 per batch segment");

  auto kv_lens = kv_indptr.slice(0, 1, batch_size + 1) - kv_indptr.slice(0, 0, batch_size);
  const int64_t max_kv_len = kv_lens.max().item<int64_t>();
  const int64_t num_chunks = (max_kv_len + chunk_size - 1) / chunk_size;
  if (num_chunks <= 0) {
    return out.zero_();
  }

  auto out_partials = torch::zeros(
      {batch_size, kNumHeads, num_chunks, config_v_dim},
      q.options().dtype(torch::kFloat32));
  auto lse_partials = torch::zeros(
      {batch_size, kNumHeads, num_chunks, 2},
      q.options().dtype(torch::kFloat32));

  const int kv_row_stride = static_cast<int>(kv_buffer.stride(0));
  const int kv_scale_row_stride = static_cast<int>(kv_scale.stride(0));

  hipLaunchKernelGGL(
      mla_decode_chunk_kernel,
      dim3(static_cast<unsigned int>(num_chunks), static_cast<unsigned int>(batch_size), 1),
      dim3(64, 16, 1),
      0,
      0,
      reinterpret_cast<const hip_bfloat16*>(q.data_ptr<at::BFloat16>()),
      reinterpret_cast<const uint8_t*>(kv_buffer.data_ptr()),
      reinterpret_cast<const uint8_t*>(kv_scale.data_ptr()),
      qo_indptr.data_ptr<int32_t>(),
      kv_indptr.data_ptr<int32_t>(),
      out_partials.data_ptr<float>(),
      lse_partials.data_ptr<float>(),
      kv_row_stride,
      kv_scale_row_stride,
      kQkHeadDim,
      static_cast<int>(config_v_dim),
      static_cast<float>(sm_scale),
      static_cast<int>(chunk_size),
      static_cast<int>(num_chunks));
  check_hip_error("mla_decode_chunk_kernel");

  hipLaunchKernelGGL(
      mla_decode_reduce_kernel,
      dim3(static_cast<unsigned int>(batch_size), static_cast<unsigned int>(kNumHeads), 1),
      dim3(256, 1, 1),
      0,
      0,
      out_partials.data_ptr<float>(),
      lse_partials.data_ptr<float>(),
      reinterpret_cast<hip_bfloat16*>(out.data_ptr<at::BFloat16>()),
      qo_indptr.data_ptr<int32_t>(),
      static_cast<int>(num_chunks),
      static_cast<int>(config_v_dim));
  check_hip_error("mla_decode_reduce_kernel");

  return out;
}
"""


def _get_chunk_size() -> int:
    raw = os.getenv(_CHUNK_ENV)
    if raw is None:
        return _DEFAULT_CHUNK_SIZE
    try:
        return max(1, int(raw))
    except ValueError:
        return _DEFAULT_CHUNK_SIZE


def _normalize_mxfp4_scale(
    kv_scale: torch.Tensor,
    total_kv: int,
    qk_head_dim: int,
) -> torch.Tensor:
    num_blocks = qk_head_dim // 32
    need = total_kv * num_blocks
    flat = kv_scale.contiguous().reshape(-1)
    if flat.numel() < need:
        raise RuntimeError(
            f"kv_scale has {flat.numel()} elements, need at least {need} for [{total_kv}, {num_blocks}]"
        )
    return flat[:need].view(total_kv, num_blocks).contiguous()


def _get_hip_module():
    global _HIP_MODULE, _HIP_BUILD_ERROR
    if _HIP_MODULE is not None:
        return _HIP_MODULE
    if _HIP_BUILD_ERROR is not None:
        raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

    with _HIP_LOCK:
        if _HIP_MODULE is not None:
            return _HIP_MODULE
        if _HIP_BUILD_ERROR is not None:
            raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")

        try:
            os.environ.setdefault("CXX", "clang++")
            _HIP_MODULE = load_inline(
                name="mla_mxfp4_splitk_hip_v1",
                cpp_sources=[CPP_WRAPPER],
                cuda_sources=[HIP_SRC],
                functions=["mla_decode_forward"],
                verbose=False,
                extra_cuda_cflags=["-O3", "-std=c++20", "-DUSE_ROCM"],
            )
        except Exception as e:  # pragma: no cover
            _HIP_BUILD_ERROR = e
            raise RuntimeError(f"HIP inline build failed: {e}") from e

    return _HIP_MODULE


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

    kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]

    q = q.contiguous()
    qo_indptr = qo_indptr.to(torch.int32).contiguous()
    kv_indptr = kv_indptr.to(torch.int32).contiguous()

    q_lens = qo_indptr[1:] - qo_indptr[:-1]
    if int(q_lens.min().item()) != 1 or int(q_lens.max().item()) != 1:
        raise RuntimeError("This decode kernel supports q_seq_len == 1 only")

    total_kv = int(kv_buffer_mxfp4.size(0))
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])

    kv_buffer_mxfp4 = kv_buffer_mxfp4.contiguous()
    kv_scale_mxfp4 = _normalize_mxfp4_scale(
        kv_scale_mxfp4,
        total_kv=total_kv,
        qk_head_dim=qk_head_dim,
    )

    module = _get_hip_module()
    return module.mla_decode_forward(
        q,
        kv_buffer_mxfp4,
        kv_scale_mxfp4,
        qo_indptr,
        kv_indptr,
        v_head_dim,
        _get_chunk_size(),
        float(config["sm_scale"]),
    )
scrolls · 486 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