Skip to content
KernelIndex
Search⌘K

submission 608791

XiaomingFun233 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 736 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-608791?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
625.5µs
#718 of 766
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c57282b48ff0b6f9e8cd55e75bfcd9e1567b2e1a8881898b1f61b290806b38c6
license declaredunknown
license concludedunknown
authorsXiaomingFun233
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__ __align__(16) float smem_kv[kSplitTileKv][kQkHeadDim];
split-kname="mla_mxfp4_splitk_hip_v5_metadata_graph",
vector-width = float4float4* smem_ptr = reinterpret_cast<float4*>(&smem_kv[token_local][dim_base]);

Kernel source

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

import math
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
_METADATA_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_BUFFER_CACHE: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}

_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 work_info,
    torch::Tensor reduce_indptr,
    torch::Tensor reduce_partial_map,
    torch::Tensor partial_output,
    torch::Tensor partial_lse,
    int64_t config_v_dim,
    double sm_scale);
"""


# The inline implementation below is reorganized to mirror the original MLA decode
# graph we documented in `e2e/cmd/mla_decode_fwd_csrc_graph.md`:
#   1. metadata/work-item construction on the Python side
#   2. asm_mla.cu-like stage1 partial accumulation
#   3. reduce.cu-like partial reduction into the final output
# The HK path from hk_decode_fwd.cu is intentionally not recreated here because
# this benchmark always uses 16 query heads with MXFP4 KV, while the HK kernel is
# an experimental fp8 h128 path.
HIP_SRC_COMMON = r"""
#include <torch/extension.h>
#include <hip/hip_bfloat16.h>
#include <hip/hip_runtime.h>

#include <algorithm>
#include <cmath>
#include <cstdint>

namespace {

constexpr int kNumHeads = 16;
constexpr int kQkHeadDim = 576;
constexpr int kVHeadDim = 512;
constexpr int kScaleBlock = 32;
constexpr int kSplitTileKv = 16;
constexpr int kPackedBytesPerToken = kQkHeadDim / 2;
constexpr int kVecBytes = 4;
constexpr int kVecPerToken = kPackedBytesPerToken / kVecBytes;
constexpr int kWavefrontSize = 64;
constexpr int kSplitWarpsPerBlock = kNumHeads;
constexpr int kSplitThreadsPerBlock = kWavefrontSize * kSplitWarpsPerBlock;

struct MlaShape {
  int64_t total_q;
  int64_t batch_size;
  int64_t num_work;
};

__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 packed, float scale, bool high_nibble) {
  const uint8_t nibble =
      high_nibble ? static_cast<uint8_t>(packed >> 4) : static_cast<uint8_t>(packed & 0x0F);
  return kFp4Lut[nibble & 0x0F] * scale;
}

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

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

inline MlaShape check_mla_inputs(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor qo_indptr,
    torch::Tensor work_info,
    torch::Tensor reduce_indptr,
    torch::Tensor reduce_partial_map,
    int64_t config_v_dim) {
  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(work_info.is_cuda(), "work_info must be CUDA/HIP tensor");
  TORCH_CHECK(reduce_indptr.is_cuda(), "reduce_indptr must be CUDA/HIP tensor");
  TORCH_CHECK(reduce_partial_map.is_cuda(), "reduce_partial_map 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 packed fp4x2 bytes");
  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 e8m0 bytes");
  TORCH_CHECK(kv_scale.dim() == 2, "kv_scale must be [total_kv, 18]");
  TORCH_CHECK(kv_scale.size(0) == kv_buffer.size(0), "kv_scale row count mismatch");
  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(work_info.scalar_type() == torch::kInt32, "work_info must be int32");
  TORCH_CHECK(reduce_indptr.scalar_type() == torch::kInt32, "reduce_indptr must be int32");
  TORCH_CHECK(reduce_partial_map.scalar_type() == torch::kInt32, "reduce_partial_map must be int32");

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

  TORCH_CHECK(q.is_contiguous(), "q must be contiguous");
  TORCH_CHECK(kv_buffer.is_contiguous(), "kv_buffer must be contiguous");
  TORCH_CHECK(kv_scale.is_contiguous(), "kv_scale must be contiguous");
  TORCH_CHECK(qo_indptr.is_contiguous(), "qo_indptr must be contiguous");
  TORCH_CHECK(work_info.is_contiguous(), "work_info must be contiguous");
  TORCH_CHECK(reduce_indptr.is_contiguous(), "reduce_indptr must be contiguous");
  TORCH_CHECK(reduce_partial_map.is_contiguous(), "reduce_partial_map must be contiguous");

  MlaShape shape{};
  shape.total_q = q.size(0);
  shape.batch_size = qo_indptr.size(0) - 1;

  TORCH_CHECK(shape.batch_size >= 0, "batch_size must be non-negative");

  shape.num_work = work_info.size(0);
  TORCH_CHECK(
      reduce_indptr.size(0) == shape.batch_size + 1,
      "reduce_indptr size mismatch");
  TORCH_CHECK(
      reduce_partial_map.size(0) == shape.num_work,
      "reduce_partial_map size mismatch");

  return shape;
}

}  // namespace
"""


HIP_SRC_STAGE1 = r"""
namespace {

__global__ void mla_stage1_split_mxfp4_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,
    const int32_t* __restrict__ work_info,
    float* __restrict__ partial_output,
    float* __restrict__ partial_lse,
    int kv_row_stride,
    int kv_scale_row_stride,
    int v_head_dim,
    float sm_scale,
    int split_size,
    int uniform_num_splits,
    int work_info_stride) {
  const int work_idx = static_cast<int>(blockIdx.x);
  const int lane = static_cast<int>(threadIdx.x);
  const int head = static_cast<int>(threadIdx.y);
  const int linear_tid = head * kWavefrontSize + lane;

  if (head >= kNumHeads) {
    return;
  }

  int32_t partial_idx = work_idx;
  int32_t q_idx = 0;
  int32_t token_begin = 0;
  int32_t token_end = 0;

  if (uniform_num_splits > 0) {
    const int32_t batch_idx = work_idx / uniform_num_splits;
    const int32_t split_idx = work_idx - batch_idx * uniform_num_splits;
    q_idx = qo_indptr[batch_idx];
    const int32_t kv_start = kv_indptr[batch_idx];
    const int32_t kv_end = kv_indptr[batch_idx + 1];
    token_begin = kv_start + split_idx * split_size;
    token_end = token_begin + split_size < kv_end ? token_begin + split_size : kv_end;
  } else {
    const int32_t* work_row = work_info + static_cast<int64_t>(work_idx) * work_info_stride;
    partial_idx = work_row[1];
    q_idx = work_row[2];
    token_begin = work_row[4];
    token_end = work_row[5];
  }

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

  float q_frag[(kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize];
  #pragma unroll
  for (int i = 0; i < (kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize; ++i) {
    const int d = lane + i * kWavefrontSize;
    q_frag[i] = (d < kQkHeadDim) ? static_cast<float>(q_row[d]) : 0.0f;
  }

  float out_frag[kVHeadDim / kWavefrontSize];
  #pragma unroll
  for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
    out_frag[i] = 0.0f;
  }

  float max_score = -INFINITY;
  float score_sum = 0.0f;

  __shared__ __align__(16) float smem_kv[kSplitTileKv][kQkHeadDim];

  for (int32_t tile_begin = token_begin; tile_begin < token_end; tile_begin += kSplitTileKv) {
    const int32_t tile_tokens =
        std::min(kSplitTileKv, static_cast<int>(token_end - tile_begin));
    const int total_vec_tasks = tile_tokens * kVecPerToken;

    for (int task = linear_tid; task < total_vec_tasks; task += kSplitThreadsPerBlock) {
      const int token_local = task / kVecPerToken;
      const int vec_idx = task % kVecPerToken;
      const int32_t token = tile_begin + token_local;

      const int scale_idx = vec_idx / 4;
      const uint8_t* scale_row = kv_scale + static_cast<int64_t>(token) * kv_scale_row_stride;
      const float scale = e8m0_to_float(scale_row[scale_idx]);

      const uint8_t* kv_row = kv_data + static_cast<int64_t>(token) * kv_row_stride;
      const uint32_t packed4 = reinterpret_cast<const uint32_t*>(kv_row)[vec_idx];
      const uint8_t* packed_bytes = reinterpret_cast<const uint8_t*>(&packed4);

      const int dim_base = vec_idx * 8;
      float4* smem_ptr = reinterpret_cast<float4*>(&smem_kv[token_local][dim_base]);

      float4 low4;
      float4 high4;
      low4.x = unpack_mxfp4(packed_bytes[0], scale, false);
      low4.y = unpack_mxfp4(packed_bytes[0], scale, true);
      low4.z = unpack_mxfp4(packed_bytes[1], scale, false);
      low4.w = unpack_mxfp4(packed_bytes[1], scale, true);

      high4.x = unpack_mxfp4(packed_bytes[2], scale, false);
      high4.y = unpack_mxfp4(packed_bytes[2], scale, true);
      high4.z = unpack_mxfp4(packed_bytes[3], scale, false);
      high4.w = unpack_mxfp4(packed_bytes[3], scale, true);

      smem_ptr[0] = low4;
      smem_ptr[1] = high4;
    }
    __syncthreads();

    #pragma unroll 1
    for (int token_local = 0; token_local < tile_tokens; ++token_local) {
      float qk_local = 0.0f;
      #pragma unroll
      for (int i = 0; i < (kQkHeadDim + kWavefrontSize - 1) / kWavefrontSize; ++i) {
        const int d = lane + i * kWavefrontSize;
        if (d < kQkHeadDim) {
          qk_local += q_frag[i] * smem_kv[token_local][d];
        }
      }

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

      const float new_max = fmaxf(max_score, score);
      const float alpha = __expf(max_score - new_max);
      const float prob = __expf(score - new_max);
      const float new_sum = score_sum * alpha + prob;

      #pragma unroll
      for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
        const int d = lane + i * kWavefrontSize;
        out_frag[i] = out_frag[i] * alpha + prob * smem_kv[token_local][d];
      }

      max_score = new_max;
      score_sum = new_sum;
    }
    __syncthreads();
  }

  const int64_t split_base =
      ((static_cast<int64_t>(partial_idx) * kNumHeads + head) * v_head_dim);
  #pragma unroll
  for (int i = 0; i < kVHeadDim / kWavefrontSize; ++i) {
    const int d = lane + i * kWavefrontSize;
    partial_output[split_base + d] = out_frag[i];
  }

  if (lane == 0) {
    const int64_t lse_base =
        ((static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2);
    partial_lse[lse_base] = max_score;
    partial_lse[lse_base + 1] = score_sum;
  }
}

void mla_decode_stage1_fp32_split(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor qo_indptr,
    torch::Tensor work_info,
    torch::Tensor partial_output,
    torch::Tensor partial_lse,
    double sm_scale,
    const MlaShape& shape) {
  if (shape.num_work <= 0) {
    return;
  }

  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));
  const int work_info_stride = static_cast<int>(work_info.stride(0));

  hipLaunchKernelGGL(
      mla_stage1_split_mxfp4_kernel,
      dim3(static_cast<unsigned int>(shape.num_work), 1, 1),
      dim3(kWavefrontSize, kNumHeads, 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>(),
      qo_indptr.data_ptr<int32_t>(),
      work_info.data_ptr<int32_t>(),
      partial_output.data_ptr<float>(),
      partial_lse.data_ptr<float>(),
      kv_row_stride,
      kv_scale_row_stride,
      kVHeadDim,
      static_cast<float>(sm_scale),
      0,
      0,
      work_info_stride);
  check_hip_error("mla_stage1_split_mxfp4_kernel");
}

}  // namespace
"""


HIP_SRC_REDUCE = r"""
namespace {

__global__ void mla_reduce_v1_like_kernel(
    const float* __restrict__ partial_output,
    const float* __restrict__ partial_lse,
    hip_bfloat16* __restrict__ out,
    const int32_t* __restrict__ qo_indptr,
    const int32_t* __restrict__ reduce_indptr,
    const int32_t* __restrict__ reduce_partial_map,
    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);
  const int num_threads = static_cast<int>(blockDim.x);

  const int32_t reduce_begin = reduce_indptr[batch_idx];
  const int32_t reduce_end = reduce_indptr[batch_idx + 1];
  const int num_splits = reduce_end - reduce_begin;

  // Shared memory: [0..num_threads-1] for warp-reduce temporaries
  extern __shared__ float smem_reduce[];
  float* smem_max = smem_reduce;               // [num_threads] for max reduction
  float* smem_sum = smem_reduce + num_threads; // [num_threads] for sum reduction

  // --- Phase 1: parallel max over splits ---
  float local_max = -INFINITY;
  for (int i = tid; i < num_splits; i += num_threads) {
    const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
    const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
    const float split_sum = partial_lse[lse_base + 1];
    if (split_sum > 0.0f) {
      local_max = fmaxf(local_max, partial_lse[lse_base]);
    }
  }
  smem_max[tid] = local_max;
  __syncthreads();
  for (int stride = num_threads >> 1; stride > 0; stride >>= 1) {
    if (tid < stride) {
      smem_max[tid] = fmaxf(smem_max[tid], smem_max[tid + stride]);
    }
    __syncthreads();
  }
  const float global_max = smem_max[0];

  if (!isfinite(global_max)) {
    // All splits empty — write zeros
    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 += num_threads) {
      out[out_base + d] = static_cast<hip_bfloat16>(0.0f);
    }
    return;
  }

  // --- Phase 2: parallel sum of rescaled weights ---
  float local_sum = 0.0f;
  for (int i = tid; i < num_splits; i += num_threads) {
    const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
    const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
    const float split_sum = partial_lse[lse_base + 1];
    if (split_sum > 0.0f) {
      local_sum += __expf(partial_lse[lse_base] - global_max) * split_sum;
    }
  }
  smem_sum[tid] = local_sum;
  __syncthreads();
  for (int stride = num_threads >> 1; stride > 0; stride >>= 1) {
    if (tid < stride) {
      smem_sum[tid] += smem_sum[tid + stride];
    }
    __syncthreads();
  }
  const float global_sum = smem_sum[0];

  // --- Phase 3: parallel V accumulation ---
  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 += num_threads) {
    float acc = 0.0f;
    for (int i = 0; i < num_splits; i++) {
      const int32_t partial_idx = reduce_partial_map[reduce_begin + i];
      const int64_t lse_base = (static_cast<int64_t>(partial_idx) * kNumHeads + head) * 2;
      const float split_sum_w = partial_lse[lse_base + 1];
      if (split_sum_w > 0.0f) {
        const float rescale = __expf(partial_lse[lse_base] - global_max);
        const int64_t split_base =
            (static_cast<int64_t>(partial_idx) * kNumHeads + head) * v_head_dim + d;
        acc += partial_output[split_base] * rescale;
      }
    }
    out[out_base + d] = static_cast<hip_bfloat16>(acc / global_sum);
  }
}

void mla_reduce_v1_like(
    torch::Tensor partial_output,
    torch::Tensor partial_lse,
    torch::Tensor qo_indptr,
    torch::Tensor reduce_indptr,
    torch::Tensor reduce_partial_map,
    torch::Tensor out,
    const MlaShape& shape) {
  if (shape.batch_size <= 0) {
    return;
  }

  constexpr int kReduceThreads = 256;
  const int smem_bytes = 2 * kReduceThreads * sizeof(float);
  hipLaunchKernelGGL(
      mla_reduce_v1_like_kernel,
      dim3(static_cast<unsigned int>(shape.batch_size), static_cast<unsigned int>(kNumHeads), 1),
      dim3(kReduceThreads, 1, 1),
      smem_bytes,
      0,
      partial_output.data_ptr<float>(),
      partial_lse.data_ptr<float>(),
      reinterpret_cast<hip_bfloat16*>(out.data_ptr<at::BFloat16>()),
      qo_indptr.data_ptr<int32_t>(),
      reduce_indptr.data_ptr<int32_t>(),
      reduce_partial_map.data_ptr<int32_t>(),
      kVHeadDim);
  check_hip_error("mla_reduce_v1_like_kernel");
}

}  // namespace
"""


HIP_SRC_ENTRY = r"""
torch::Tensor mla_decode_forward(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor qo_indptr,
    torch::Tensor work_info,
    torch::Tensor reduce_indptr,
    torch::Tensor reduce_partial_map,
    torch::Tensor partial_output,
    torch::Tensor partial_lse,
    int64_t config_v_dim,
    double sm_scale) {

  const MlaShape shape =
      check_mla_inputs(
          q,
          kv_buffer,
          kv_scale,
          qo_indptr,
          work_info,
          reduce_indptr,
          reduce_partial_map,
          config_v_dim);

  auto out = torch::empty({shape.total_q, kNumHeads, config_v_dim}, q.options().dtype(torch::kBFloat16));
  if (shape.total_q == 0 || shape.batch_size <= 0) {
    return out.zero_();
  }

  mla_decode_stage1_fp32_split(
      q,
      kv_buffer,
      kv_scale,
      qo_indptr,
      work_info,
      partial_output,
      partial_lse,
      sm_scale,
      shape);

  mla_reduce_v1_like(partial_output, partial_lse, qo_indptr, reduce_indptr, reduce_partial_map, out, shape);
  return out;
}
"""


HIP_SRC = "\n".join(
    [
        HIP_SRC_COMMON,
        HIP_SRC_STAGE1,
        HIP_SRC_REDUCE,
        HIP_SRC_ENTRY,
    ]
)


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 _build_decode_metadata(
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    split_size: int,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    qo_host = qo_indptr.detach().cpu().tolist()
    kv_host = kv_indptr.detach().cpu().tolist()
    cache_key = (
        str(device),
        split_size,
        tuple(qo_host),
        tuple(kv_host),
    )
    cached = _METADATA_CACHE.get(cache_key)
    if cached is not None:
        return cached

    batch_size = len(qo_host) - 1
    work_rows: list[list[int]] = []
    reduce_partial_map: list[int] = []
    reduce_indptr = [0]

    for batch_idx in range(batch_size):
        q_start = int(qo_host[batch_idx])
        q_end = int(qo_host[batch_idx + 1])
        kv_start = int(kv_host[batch_idx])
        kv_end = int(kv_host[batch_idx + 1])
        kv_len = kv_end - kv_start

        num_splits = math.ceil(kv_len / split_size) if kv_len > 0 else 0
        for split_idx in range(num_splits):
            split_begin = kv_start + split_idx * split_size
            split_end = min(split_begin + split_size, kv_end)
            partial_idx = len(work_rows)
            work_rows.append(
                [batch_idx, partial_idx, q_start, q_end, split_begin, split_end, 0, 0]
            )
            reduce_partial_map.append(partial_idx)

        reduce_indptr.append(len(reduce_partial_map))

    if work_rows:
        work_info = torch.tensor(work_rows, dtype=torch.int32, device=device)
    else:
        work_info = torch.empty((0, 8), dtype=torch.int32, device=device)

    metadata = (
        work_info.contiguous(),
        torch.tensor(reduce_indptr, dtype=torch.int32, device=device).contiguous(),
        torch.tensor(reduce_partial_map, dtype=torch.int32, device=device).contiguous(),
    )
    _METADATA_CACHE[cache_key] = metadata
    return metadata


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_v5_metadata_graph",
                cpp_sources=[CPP_WRAPPER],
                cuda_sources=[HIP_SRC],
                functions=["mla_decode_forward"],
                verbose=False,
                extra_cuda_cflags=["-O3", "-ffast-math", "-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"])
    batch_size = int(qo_indptr.size(0)) - 1

    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,
    )
    split_size = _get_chunk_size()
    work_info, reduce_indptr, reduce_partial_map = _build_decode_metadata(
        qo_indptr,
        kv_indptr,
        split_size=split_size,
        device=q.device,
    )

    num_work = int(work_info.size(0))
    num_heads = 16
    buf_key = (num_work, num_heads, v_head_dim)
    if buf_key not in _BUFFER_CACHE:
        _BUFFER_CACHE[buf_key] = (
            torch.empty((num_work, num_heads, v_head_dim), dtype=torch.float32, device=q.device),
            torch.empty((num_work, num_heads, 2), dtype=torch.float32, device=q.device),
        )
    partial_output, partial_lse = _BUFFER_CACHE[buf_key]

    module = _get_hip_module()
    return module.mla_decode_forward(
        q,
        kv_buffer_mxfp4,
        kv_scale_mxfp4,
        qo_indptr,
        work_info,
        reduce_indptr,
        reduce_partial_map,
        partial_output,
        partial_lse,
        v_head_dim,
        float(config["sm_scale"]),
    )
scrolls · 736 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