Skip to content
KernelIndex
Search⌘K

submission 629960

ưhat's up · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-629960?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
271.0µs
#703 of 766
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dd9a3db6a2b6f47fc7e665498f5ab627f2250fa150d02e5b48e398e0f567f8e7
license declaredunknown
license concludedunknown
authorsưhat's up
imported2026-08-26

Techniques

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

fp4kv_data["mxfp4"] — (Tensor, Tensor): kv_buffer fp4x2 (total_kv,1,288) + fp8_e8m0 scale
persistent-kernelThe reference uses aiter's a8w8 persistent MLA kernel (fp8 Q + fp8 KV),
shared-memory__shared__ float shared_qk_max[N_WARPS][16];

Kernel source

v0.py634 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA (Multi-head Latent Attention) decode kernel — submission template.

Implement custom_kernel() to beat the aiter a8w8 reference (fp8 Q + fp8 KV).

DeepSeek R1 forward_absorb MLA config:
  num_heads        = 16     (query heads, after TP split)
  num_kv_heads     = 1      (shared latent KV head)
  kv_lora_rank     = 512    (latent dim)
  qk_rope_head_dim = 64     (RoPE dim)
  qk_head_dim      = 576    (kv_lora_rank + qk_rope_head_dim, absorbed q/k dim)
  v_head_dim       = 512    (= kv_lora_rank, output dim)
  sm_scale         = 1/sqrt(576)

KV buffer format (forward_absorb):
  - Full 576 dims used as keys (for Q@K^T score computation)
  - First 512 dims (kv_lora_rank) used as values (for output computation)

Input tuple:
  q:          (total_q, 16, 576)       bfloat16 — absorbed query
  kv_data:    dict with three KV cache formats:
    kv_data["bf16"]  — Tensor (total_kv, 1, 576) bfloat16
    kv_data["fp8"]   — (Tensor, Tensor): kv_buffer fp8 (total_kv,1,576) + scalar scale
    kv_data["mxfp4"] — (Tensor, Tensor): kv_buffer fp4x2 (total_kv,1,288) + fp8_e8m0 scale
  qo_indptr:  (batch_size + 1,)        int32    — query segment pointers
  kv_indptr:  (batch_size + 1,)        int32    — KV segment pointers
  config:     dict with MLA parameters

Output:
  attention output: (total_q, 16, 512) bfloat16

The reference uses aiter's a8w8 persistent MLA kernel (fp8 Q + fp8 KV),
which is ~2-3x faster than bf16. To beat it, consider:
  1. Use mxfp4 KV cache for even lower memory bandwidth
     - Fuse dequantization with attention to avoid bf16 materialization
  2. Custom kernel with tighter memory access patterns
  3. MQA: 1 KV head shared across 16 query heads — minimize redundant memory loads
  4. Variable-length batching: indptr-based segmented attention
  5. Split K/V from buffer: full 576 dims for keys, first 512 dims for values
"""

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

if "PYTORCH_ROCM_ARCH" not in os.environ:
    os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"

# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
# Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
# ---------------------------------------------------------------------------


def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.

    Block size = 32. Each block gets an E8M0 scale factor.
    Two FP4 E2M1 values are packed per byte.

    Args:
        tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)

    Returns:
        (fp4_data, scale_e8m0)
        - fp4_data:   shape [B, M, N//2] in aiter_dtypes.fp4x2
        - scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
    """
    from aiter.ops.triton.quant import dynamic_mxfp4_quant

    orig_shape = tensor.shape  # (B, M, N)
    B, M, N = orig_shape

    # dynamic_mxfp4_quant expects 2D: (B*M, N)
    tensor_2d = tensor.reshape(B * M, N)
    fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)

    # Reshape fp4_data back to 3D: (B, M, N//2)
    fp4_data = fp4_data_2d.view(B, M, N // 2)

    return fp4_data, scale_e8m0


kernel_cpp = r"""
#undef __HIP_NO_HALF_OPERATORS__
#undef __HIP_NO_HALF_CONVERSIONS__

#include <hip/hip_bf16.h>
#include <hip/hip_ext_ocp.h>
#include <hip/hip_fp16.h>
#include <hip/hip_fp4.h>
#include <hip/hip_fp8.h>
#include <hip/hip_runtime.h>
#include <pybind11/pybind11.h>

#include <math.h>

using fp4x2_t = __amd_fp4x2_storage_t;
using fp4x16_t = fp4x2_t __attribute__((ext_vector_type(8)));
using fp4x32_t = fp4x2_t __attribute__((ext_vector_type(16)));
using fp4x64_t = fp4x2_t __attribute__((ext_vector_type(32)));
using uint8x2_t = uint8_t __attribute__((ext_vector_type(2)));
using uint8x16_t = uint8_t __attribute__((ext_vector_type(16)));
using uint8x32_t = uint8_t __attribute__((ext_vector_type(32)));
using bfp16 = __hip_bfloat16;
using bfp16x2_t = ushort __attribute__((ext_vector_type(2)));
using bfp16x4_t = ushort __attribute__((ext_vector_type(4)));
using bfp16x8_t = ushort __attribute__((ext_vector_type(8)));
using floatx4_t = float __attribute__((ext_vector_type(4)));
using floatx8_t = float __attribute__((ext_vector_type(8)));
using floatx16_t = float __attribute__((ext_vector_type(16)));
using floatx32_t = float __attribute__((ext_vector_type(32)));

#define WARP_SIZE 64
#define CEIL_DIV(x, y) (((x) + ((y) - 1)) / (y))

__device__ inline ushort float_2_bfloatraw(float f) {
#if HIP_BF16_AVX512_OP
  union {
    __bf16 bf16;
    unsigned short us;
  } u = {_mm_cvtness_sbh(f)};
  return u.us;
#else
  union {
    float fp32;
    unsigned int u32;
  } u = {f};
  if (~u.u32 & 0x7f800000) {
    u.u32 += 0x7fff + ((u.u32 >> 16) & 1); // Round to nearest, round to even
  } else if (u.u32 & 0xffff) {
    u.u32 |= 0x10000; // Preserve signaling NaN
  }
  return static_cast<unsigned short>(u.u32 >> 16);
#endif
}

__device__ inline fp4x64_t fp4x32_to_fp4x64(const fp4x32_t &p) {
  fp4x64_t v;
  reinterpret_cast<fp4x32_t *>(&v)[0] = p;
  reinterpret_cast<fp4x32_t *>(&v)[1] = 0;
  return v;
}

__device__ inline uint8x32_t fp4x32_to_uint8x32(const fp4x32_t &p) {
  uint8x32_t v;
#pragma unroll
  for (int i = 0; i < 16; ++i) {
    uint8x2_t tmp;
    tmp[0] = p[i] & 0x0f;
    tmp[1] = (p[i] & 0xf0) >> 4;
    reinterpret_cast<uint8x2_t *>(&v)[i] = tmp;
  }
  return v;
}

__device__ inline fp4x32_t uint8x32_to_fp4x32(const uint8x32_t &p) {
  fp4x32_t v;
#pragma unroll
  for (int i = 0; i < 16; ++i) {
    v[i] = (p[i << 1 | 1] << 4) | (p[i << 1]);
  }
  return v;
}

__device__ __forceinline__ bfp16x2_t fp4x2_to_bfp16x2(uint8_t packed) {

  ushort em0 = packed & 0x07; // Low nibble E+M
  ushort em1 = packed & 0x70; // High nibble E+M

  // BFloat16 Constants
  // Bias difference (127 - 1) = 126. Shifted to BF16 exp position: 126 << 7
  constexpr ushort bfp16_bias = 0x3F00;
  constexpr ushort bfp16_0p5 = 0x3F00; // 0.5 in Float32 hex

  // 2. Direct alignment to BFloat16 positions
  // Shift E+M left by 6. This perfectly drops the Mantissa into bit 6,
  // and the Exponent into bits 7 and 8.
  // Sign bits (Bit 3 and Bit 7) are shifted to the 15st bit.
  ushort x0 = (em0 << 6) | ((packed & 0x08) << 12);
  ushort x1 = (em1 << 2) | ((packed & 0x80) << 8);

  // 3. Apply Bias via Integer Addition (Triton's trick)
  // CMOV (Conditional Move): Only add bias if exponent is not zero
  x0 = ((em0 & 0x06) != 0) ? (x0 + bfp16_bias) : x0;
  x1 = ((em1 & 0x60) != 0) ? (x1 + bfp16_bias) : x1;

  // 4. Handle Subnormals
  // CMOV: Force to +/- 0.5 if it's a subnormal (em == 0x01)
  x0 = (em0 == 0x01) ? (bfp16_0p5 | (x0 & 0x8000)) : x0;
  x1 = (em1 == 0x10) ? (bfp16_0p5 | (x1 & 0x8000)) : x1;

  return {x0, x1};
}

__device__ inline floatx4_t mfma_16x16x32(bfp16x8_t a, bfp16x8_t b,
                                          floatx4_t c) {
  return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, c, 0, 0, 0);
}

__device__ inline floatx4_t mfma_16x16x128(fp4x64_t a, fp4x64_t b, floatx4_t c,
                                           uint8_t scale_a, uint8_t scale_b) {
  return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a, b, c, 4, 4, 0,
                                                          scale_a, 0, scale_b);
}

__device__ inline floatx16_t mfma_32x32x64(fp4x64_t a, fp4x64_t b, floatx16_t c,
                                           uint8_t scale_a, uint8_t scale_b) {
  return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a, b, c, 4, 4, 0,
                                                         scale_a, 0, scale_b);
}

__global__ void mla_decode(
    const fp4x2_t *q,        // [total_q, num_heads, (kv_lora + rope_dim) / 2]
    const uint8_t *q_scale,  // [total_q, num_heads, (kv_lora + rope_dim) / 32]
    const fp4x2_t *kv_cache, // [total_kv, 1, (kv_lora + rope_dim) / 2]
    const uint8_t *kv_scale, // [total_kv, 1, (kv_lora + rope_dim) / 32]
    const int *qo_indptr,    // [num_seqs + 1]
    const int *kv_indptr,    // [num_seqs + 1]
    bfp16 *output,           // [total_q, num_heads, num_partition, kv_lora]
    float *lse,              // [total_q, num_heads, num_partition]
    const int num_seqs,
    const int max_q_seq_len, // = 1
    const int max_kv_seq_len, const float scale, const int max_num_partitions) {
  // template arguments
  constexpr int N_WARPS = 4;
  constexpr int PARTITION_SIZE = 256;
  constexpr int QK_HEAD_DIM = 576;
  constexpr int V_HEAD_DIM = 512;
  constexpr int N_HEADS = 16;
  //
  constexpr int QKHE_PER_LOOP = 128;
  constexpr int QKHE_LOOP = CEIL_DIV(QK_HEAD_DIM, QKHE_PER_LOOP);

  constexpr int VHE_PER_LOOP = 16;
  constexpr int VHE_LOOP = CEIL_DIV(V_HEAD_DIM, 16);

  constexpr int TOKENS_PER_WARP = PARTITION_SIZE / N_WARPS;
  constexpr int TOKENS_PER_LOOP = 16;
  constexpr int TOKEN_LOOP = CEIL_DIV(TOKENS_PER_WARP, 16);

  constexpr int VT_PER_LOOP = 32;
  constexpr int VT_LOOP = TOKENS_PER_WARP / 32;

  const int seq_id = blockIdx.y;
  const int query_idx = qo_indptr[seq_id];

  const int part_id = blockIdx.x;
  const int kv_idx = kv_indptr[seq_id];
  const int ctx_len = kv_indptr[seq_id + 1] - kv_idx;
  const int start_token_idx = part_id * PARTITION_SIZE;
  const int num_partitions = CEIL_DIV(ctx_len, PARTITION_SIZE);
  if (part_id >= num_partitions)
    return;
  const int end_token_idx = min(start_token_idx + PARTITION_SIZE, ctx_len);

  const int tid = threadIdx.x;
  const int wid = tid / WARP_SIZE;
  const int lane_id = tid % WARP_SIZE;
  const int lane_id16 = lane_id % 16;
  const int lane4_id = lane_id / 16;

  // load q: vram -> reg
  const fp4x2_t *q_ptr =
      q + (query_idx * N_HEADS + lane_id16) * (QK_HEAD_DIM / 2);
  fp4x64_t q_local[QKHE_LOOP];
  for (int i = 0; i < QKHE_LOOP; ++i) {
    const int head_elem = i * QKHE_PER_LOOP + lane4_id * 32;
    fp4x32_t tmp;
    if (head_elem < QK_HEAD_DIM) {
      tmp = *reinterpret_cast<const fp4x32_t *>(q_ptr + head_elem / 2);
    } else {
      tmp = 0;
    }
    q_local[i] = fp4x32_to_fp4x64(tmp);
  }
  // load q_scale : vram->reg
  uint8_t q_scale_local[QKHE_LOOP];
  const uint8_t *q_scale_ptr =
      q_scale + (query_idx * N_HEADS + lane_id16) * (QK_HEAD_DIM / 32);
  for (int i = 0; i < QKHE_LOOP; ++i) {
    const int head_elem = i * QKHE_PER_LOOP + lane4_id * 32;
    if (head_elem < QK_HEAD_DIM) {
      q_scale_local[i] = q_scale_ptr[head_elem / 32];
    } else {
      q_scale_local[i] = 0;
    }
  }

  // load k: vram -> reg
  fp4x64_t k_local[TOKEN_LOOP][QKHE_LOOP];
  for (int i = 0; i < TOKEN_LOOP; ++i) {
    const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
                          i * TOKENS_PER_LOOP + lane_id16;
    const fp4x2_t *k_ptr = kv_cache + (kv_idx + token_idx) * (QK_HEAD_DIM / 2);
    for (int j = 0; j < QKHE_LOOP; ++j) {
      const int head_elem = j * QKHE_PER_LOOP + lane4_id * 32;
      fp4x32_t tmp;
      if (token_idx < end_token_idx && head_elem < QK_HEAD_DIM) {
        tmp = *reinterpret_cast<const fp4x32_t *>(k_ptr + head_elem / 2);
      } else {
        tmp = 0;
      }
      k_local[i][j] = fp4x32_to_fp4x64(tmp);
    }
  }
  // load k_scale: vram -> reg
  uint8_t k_scale_local[TOKEN_LOOP][QKHE_LOOP];
  for (int i = 0; i < TOKEN_LOOP; ++i) {
    const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
                          i * TOKENS_PER_LOOP + lane_id16;
    const uint8_t *k_scale_ptr =
        kv_scale + (kv_idx + token_idx) * (QK_HEAD_DIM / 32);
    for (int j = 0; j < QKHE_LOOP; ++j) {
      const int head_elem = j * QKHE_PER_LOOP + lane4_id * 32;
      if (token_idx < end_token_idx && head_elem < QK_HEAD_DIM) {
        k_scale_local[i][j] = k_scale_ptr[head_elem / 32];
      } else {
        k_scale_local[i][j] = 0;
      }
    }
  }

  // q @ k
  floatx4_t acc[TOKEN_LOOP];
  for (int i = 0; i < TOKEN_LOOP; ++i) {
    acc[i] = 0;
    for (int j = 0; j < QKHE_LOOP; ++j) {
      acc[i] = mfma_16x16x128(k_local[i][j], q_local[j], acc[i],
                              k_scale_local[i][j], q_scale_local[j]);
    }
    acc[i] *= scale;
  }
  float qk_max = -INFINITY;
  for (int i = 0; i < TOKEN_LOOP; ++i) {
    for (int j = 0; j < 4; ++j) {
      qk_max = fmaxf(qk_max, acc[i][j]);
    }
  }
  qk_max = fmaxf(__shfl_xor(qk_max, 16), qk_max);
  qk_max = fmaxf(__shfl_xor(qk_max, 32), qk_max);

  __shared__ float shared_qk_max[N_WARPS][16];
  if (lane_id < 16) {
    shared_qk_max[wid][lane_id16] = qk_max;
  }
  __syncthreads();
  for (int w = 0; w < N_WARPS; ++w) {
    qk_max = fmaxf(qk_max, shared_qk_max[w][lane_id16]);
  }

  float exp_sum = 0.f;
  for (int i = 0; i < TOKEN_LOOP; ++i) {
    const int token_idx = start_token_idx + wid * TOKENS_PER_WARP +
                          i * TOKENS_PER_LOOP + lane4_id * 4;
    for (int j = 0; j < 4; ++j) {
      float tmp =
          (token_idx + j < end_token_idx) ? __expf(acc[i][j] - qk_max) : 0.f;
      acc[i][j] = tmp;
      exp_sum += tmp;
    }
  }

  exp_sum += __shfl_xor(exp_sum, 16);
  exp_sum += __shfl_xor(exp_sum, 32);

  __shared__ float shared_exp_sum[N_WARPS][16];
  if (lane_id < 16) {
    shared_exp_sum[wid][lane_id16] = exp_sum;
  }
  __syncthreads();
  exp_sum = 0.f;
  for (int w = 0; w < N_WARPS; ++w) {
    exp_sum += shared_exp_sum[w][lane_id16];
  }

  if (lse != nullptr && wid == 0 && lane_id < 16) {
    lse[(query_idx * N_HEADS + lane_id) * max_num_partitions + part_id] =
        logf(exp_sum) + qk_max;
  }

  bfp16x8_t p[VT_LOOP];
  for (int i = 0; i < VT_LOOP; ++i) {
    for (int j = 0; j < 4; ++j) {
      p[i][j] = float_2_bfloatraw(acc[i * 2][j]);
      p[i][j + 4] = float_2_bfloatraw(acc[i * 2 + 1][j]);
    }
  }

  bfp16x8_t v_local[VHE_LOOP][VT_LOOP];

  __shared__ bfp16 transposed_v[N_WARPS][QKHE_PER_LOOP][VT_PER_LOOP];

  for (int qkhe_d = 0; qkhe_d < V_HEAD_DIM / QKHE_PER_LOOP; ++qkhe_d) {
    const int qkh_elem = lane4_id * 32;
    for (int vt_d = 0; vt_d < VT_LOOP; ++vt_d) {
      for (int i = 0; i < (VT_PER_LOOP / TOKENS_PER_LOOP); ++i) {
        int t_d = vt_d * (VT_PER_LOOP / TOKENS_PER_LOOP) + i;
        const int token_id = i * 4 + lane_id16 / 4 * 8 + lane_id16 % 4;
        fp4x32_t tmp = *reinterpret_cast<fp4x32_t *>(&k_local[t_d][qkhe_d]);
        ushort raw_v_scale = (ushort)k_scale_local[t_d][qkhe_d] << 7;
        bfp16 v_scale = *reinterpret_cast<bfp16 *>(&raw_v_scale);
        for (int j = 0; j < 16; ++j) {
          const int head_elem = lane4_id * 32 + j * 2;
          bfp16x2_t quan_v = fp4x2_to_bfp16x2(tmp[j]);
          bfp16 v0 = reinterpret_cast<bfp16 *>(&quan_v)[0] * v_scale;
          bfp16 v1 = reinterpret_cast<bfp16 *>(&quan_v)[1] * v_scale;
          transposed_v[wid][head_elem][token_id] = v0;
          transposed_v[wid][head_elem + 1][token_id] = v1;
        }
      }
      __syncthreads();
      for (int i = 0; i < (QKHE_PER_LOOP / VHE_PER_LOOP); ++i) {
        int vhe_d = qkhe_d * (QKHE_PER_LOOP / VHE_PER_LOOP) + i;
        const int head_elem = i * VHE_PER_LOOP + lane_id16;
        const int token_id = lane4_id * 8;
        v_local[vhe_d][vt_d] = *reinterpret_cast<bfp16x8_t *>(
            &transposed_v[wid][head_elem][token_id]);
      }
      __syncthreads();
    }
  }

  floatx4_t o_acc[VHE_LOOP];
  for (int i = 0; i < VHE_LOOP; ++i) {
    o_acc[i] = 0;
    for (int j = 0; j < VT_LOOP; ++j) {
      o_acc[i] = mfma_16x16x32(v_local[i][j], p[j], o_acc[i]);
    }
  }

  floatx4_t *shared_out = reinterpret_cast<floatx4_t *>(transposed_v);
  constexpr int chunk_vhe = (sizeof(bfp16) * QKHE_PER_LOOP * VT_PER_LOOP) /
                            (sizeof(floatx4_t) * WARP_SIZE);
  for (int i = 0; i < VHE_LOOP; i += chunk_vhe) {
    for (int j = 0; j < chunk_vhe; ++j) {
      shared_out[(j * N_WARPS + wid) * WARP_SIZE + lane_id] = o_acc[i + j];
    }
    __syncthreads();
    if (wid == 0) {
      for (int j = 0; j < chunk_vhe; ++j) {
        for (int w = 1; w < N_WARPS; ++w) {
          o_acc[i + j] += shared_out[(j * N_WARPS + w) * WARP_SIZE + lane_id];
        }
      }
    }
    __syncthreads();
  }

  if (wid == 0) {
    float inv_exp_sum = exp_sum == 0 ? 1.f : 1.f / exp_sum;
    for (int i = 0; i < VHE_LOOP; ++i) {
      o_acc[i] *= inv_exp_sum;
    }
    bfp16 *out_ptr =
        output +
        (query_idx * N_HEADS + lane_id16) * max_num_partitions * V_HEAD_DIM +
        part_id * V_HEAD_DIM;

    for (int i = 0; i < VHE_LOOP; ++i) {
      const int head_elem = i * VHE_PER_LOOP + lane4_id * 4;
      bfp16x4_t tmp = {
          float_2_bfloatraw(o_acc[i][0]), float_2_bfloatraw(o_acc[i][1]),
          float_2_bfloatraw(o_acc[i][2]), float_2_bfloatraw(o_acc[i][3])};
      *reinterpret_cast<bfp16x4_t *>(out_ptr + head_elem) = tmp;
    }
  }
}

__global__ void mla_decode_reduce(
    const bfp16 *input,   // [total_q, n_heads, max_num_partitions, kv_lora]
    const float *lse,     // [num_seqs, n_heads, max_num_partitions]
    bfp16 *output,        // [total_q, num_heads, kv_lora]
    const int *qo_indptr, // [num_seqs + 1]
    const int *kv_indptr, // [num_seqs + 1]
    const int num_seqs, const int max_num_partitions) {
  // HEAD_DIM = kv_lora = 512
  constexpr int HEAD_DIM = 512;
  constexpr int PACKED_SIZE = 8;
  constexpr int NUM_HEADS = 16;
  constexpr int PARTITION_SIZE = 256;

  using packed_io_t = ushort __attribute__((ext_vector_type(PACKED_SIZE)));
  using packed_process_t = float __attribute__((ext_vector_type(PACKED_SIZE)));

  const int tid = blockIdx.x * blockDim.x + threadIdx.x;
  const int global_idx = tid * PACKED_SIZE;
  const int seq_id = global_idx / (NUM_HEADS * HEAD_DIM);
  const int head_id = (global_idx / HEAD_DIM) % NUM_HEADS;
  const int head_element = global_idx % HEAD_DIM;

  const int query_id = qo_indptr[seq_id];
  const int context_len = kv_indptr[seq_id + 1] - kv_indptr[seq_id];
  const int num_partition = CEIL_DIV(context_len, PARTITION_SIZE);

  float accum_lse = -INFINITY;
  packed_process_t accum_out = 0.0f;
  for (int part_id = 0; part_id < num_partition; ++part_id) {
    const int input_idx =
        ((query_id * NUM_HEADS + head_id) * max_num_partitions + part_id) *
            HEAD_DIM +
        head_element;
    packed_io_t local_out_bf16 =
        reinterpret_cast<const packed_io_t *>(&input[input_idx])[0];
    packed_process_t local_out;
#pragma unroll
    for (int i = 0; i < PACKED_SIZE; ++i) {
      local_out[i] = __bfloat162float(__ushort_as_bfloat16(local_out_bf16[i]));
    }
    float local_lse =
        lse[(query_id * NUM_HEADS + head_id) * max_num_partitions + part_id];

    const float new_lse = logf(expf(accum_lse) + expf(local_lse));
    accum_out = accum_out * expf(accum_lse - new_lse) +
                local_out * expf(local_lse - new_lse);
    accum_lse = new_lse;
  }

  packed_io_t output_bf16;
#pragma unroll
  for (int i = 0; i < PACKED_SIZE; ++i) {
    output_bf16[i] = float_2_bfloatraw(accum_out[i]);
  }
  const int output_idx =
      (query_id * NUM_HEADS + head_id) * HEAD_DIM + head_element;
  reinterpret_cast<packed_io_t *>(&output[output_idx])[0] = output_bf16;
}

void launch_mla_decode(uintptr_t q,         //
                       uintptr_t q_scale,   //
                       uintptr_t kv,        //
                       uintptr_t kv_scale,  //
                       uintptr_t qo_indptr, //
                       uintptr_t kv_indptr, //
                       uintptr_t out,       //
                       uintptr_t lse,       //
                       uintptr_t tmp_out,   //
                       int num_seqs, int max_q_seq_len, int max_kv_seq_len,
                       float scale, int max_num_partitions) {
  const auto *d_q = reinterpret_cast<fp4x2_t *>(q);
  const auto *d_q_scale = reinterpret_cast<uint8_t *>(q_scale);
  const auto *d_kv = reinterpret_cast<fp4x2_t *>(kv);
  const auto *d_kv_scale = reinterpret_cast<uint8_t *>(kv_scale);
  const auto *d_qo_indptr = reinterpret_cast<int *>(qo_indptr);
  const auto *d_kv_indptr = reinterpret_cast<int *>(kv_indptr);
  auto *d_out = reinterpret_cast<bfp16 *>(out);
  auto *d_lse = lse == 0 ? nullptr : reinterpret_cast<float *>(lse);
  auto *d_tmp_out = tmp_out == 0 ? d_out : reinterpret_cast<bfp16 *>(tmp_out);

  dim3 block_size(256);
  dim3 grid_size(max_num_partitions, num_seqs);
  hipLaunchKernelGGL(mla_decode, grid_size, block_size, 0, 0, d_q, d_q_scale,
                     d_kv, d_kv_scale, d_qo_indptr, d_kv_indptr, d_tmp_out,
                     d_lse, num_seqs, max_q_seq_len, max_kv_seq_len, scale,
                     max_num_partitions);

  if (max_num_partitions == 1)
    return;
  block_size = dim3(256);
  grid_size = dim3(num_seqs * 4);
  hipLaunchKernelGGL(mla_decode_reduce, grid_size, block_size, 0, 0, d_tmp_out,
                     d_lse, d_out, d_qo_indptr, d_kv_indptr, num_seqs,
                     max_num_partitions);
}

PYBIND11_MODULE(mxfp4, m) {
  m.def("launch_mla_decode", &launch_mla_decode, "HIP MLA Decode kernel");
}

"""

hip_module = load_inline(
    name="mxfp4",
    cpp_sources="",
    cuda_sources=kernel_cpp,
    with_cuda=True,
    verbose=False,
    extra_cuda_cflags=["-std=c++20", "-O3"],
    no_implicit_headers=True,
)


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

    q, q_scale = quantize_mxfp4(q_bf16)

    kv_cache, kv_scale = kv_data["mxfp4"]

    num_seqs = config["batch_size"]
    max_q_seq_len = config["q_seq_len"]
    max_kv_seq_len = config["kv_seq_len"]
    scale = config["sm_scale"]

    out = torch.empty((num_seqs, 16, 512), dtype=torch.bfloat16, device=q.device)

    max_num_partitions = (max_kv_seq_len + PARITTION_SIZE - 1) // PARITTION_SIZE

    lse = None
    tmp_out = None
    if max_num_partitions > 1:
        lse = torch.empty(
            (num_seqs, 16, max_num_partitions), dtype=torch.float32, device=q.device
        )
        tmp_out = torch.empty(
            (num_seqs, 16, max_num_partitions, 512),
            dtype=torch.bfloat16,
            device=q.device,
        )

    hip_module.launch_mla_decode(
        q.data_ptr(),
        q_scale.data_ptr(),
        kv_cache.data_ptr(),
        kv_scale.data_ptr(),
        qo_indptr.data_ptr(),
        kv_indptr.data_ptr(),
        out.data_ptr(),
        lse.data_ptr() if lse is not None else 0,
        tmp_out.data_ptr() if lse is not None else 0,
        num_seqs,
        max_q_seq_len,
        max_kv_seq_len,
        scale,
        max_num_partitions,
    )
    return out
scrolls · 634 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