Skip to content
KernelIndex
Search⌘K

submission 589134

Purple rain · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_amd_moe_mxfp4_hip_fused.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589134?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
185.6µs
#650 of 782
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1427acc0ad3bd553a730bdd60ed139b3001a0ec63307792f27ffcf2aa7f3aa15
license declaredunknown
license concludedunknown
authorsPurple rain
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float smem[];

Kernel source

submission_amd_moe_mxfp4_hip_fused.py505 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!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
_MODE_ENV = "MOE_MXFP4_HIP_FUSED_MODE"  # base | hip
_DEFAULT_MODE = "base"


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

torch::Tensor hip_moe_forward(
    torch::Tensor hidden_states,
    torch::Tensor gate_up_weight,
    torch::Tensor down_weight,
    torch::Tensor gate_up_weight_scale,
    torch::Tensor down_weight_scale,
    torch::Tensor topk_weights,
    torch::Tensor topk_ids,
    int64_t d_expert);
"""


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

#include <cstdint>
#include <vector>

namespace {

constexpr int kThreads = 256;
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 fp4_nibble_to_float(uint8_t nibble) {
  return kFp4Lut[nibble & 0x0F];
}

__device__ __forceinline__ float e8m0_to_float(uint8_t e8m0_val) {
  // Match aiter.utility.fp4_utils.e8m0_to_f32:
  //   normal: exponent bits in fp32
  //   0x00  : smallest positive value used by the kernel path
  //   0xFF  : NaN payload
  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 fp4x2_at(const uint8_t* packed_row, int col) {
  uint8_t packed = packed_row[col >> 1];
  uint8_t nib = (col & 1) ? static_cast<uint8_t>(packed >> 4)
                          : static_cast<uint8_t>(packed & 0x0F);
  return fp4_nibble_to_float(nib);
}

__device__ __forceinline__ float silu(float x) {
  return x / (1.0f + expf(-x));
}

__device__ __forceinline__ float warp_sum(float v) {
  for (int offset = warpSize / 2; offset > 0; offset >>= 1) {
    v += __shfl_down(v, offset);
  }
  return v;
}

__global__ void fused_moe_kernel(
    const hip_bfloat16* hidden_states,     // [M, d_hidden]
    const uint8_t* gate_up_weight,         // [E, 2*d_expert_pad, d_hidden_pad/2]
    const uint8_t* down_weight,            // [E, d_hidden_pad, d_expert_pad/2]
    const uint8_t* gate_up_weight_scale,   // [E, 2*d_expert_pad, d_hidden_pad/32]
    const uint8_t* down_weight_scale,      // [E, d_hidden_pad, d_expert_pad/32]
    const float* topk_weights,             // [M, topk]
    const int32_t* topk_ids,               // [M, topk]
    float* output_accum,                   // [M, d_hidden]
    int32_t m,
    int32_t topk,
    int32_t num_experts,
    int32_t d_hidden,
    int32_t d_hidden_pad,
    int32_t d_expert,
    int32_t d_expert_pad,
    int32_t gateup_k_pack,
    int32_t down_k_pack,
    int32_t gateup_scale_k,
    int32_t down_scale_k) {
  int32_t token = static_cast<int32_t>(blockIdx.x);
  int32_t k_idx = static_cast<int32_t>(blockIdx.y);
  if (token >= m || k_idx >= topk) {
    return;
  }

  int32_t tid = static_cast<int32_t>(threadIdx.x);
  int32_t lane = tid % warpSize;
  int32_t warp_id = tid / warpSize;
  int32_t warps_per_block = static_cast<int32_t>(blockDim.x) / warpSize;

  int64_t route_offset = static_cast<int64_t>(token) * topk + k_idx;
  int32_t expert = topk_ids[route_offset];
  if (expert < 0 || expert >= num_experts) {
    return;
  }
  float route_weight = topk_weights[route_offset];

  extern __shared__ float smem[];
  float* x_sh = smem;                  // [d_hidden]
  float* inter_sh = smem + d_hidden;   // [d_expert]

  const hip_bfloat16* x_row =
      hidden_states + static_cast<int64_t>(token) * d_hidden;
  for (int32_t col = tid; col < d_hidden; col += static_cast<int32_t>(blockDim.x)) {
    x_sh[col] = static_cast<float>(x_row[col]);
  }
  __syncthreads();

  const int64_t gateup_expert_stride =
      static_cast<int64_t>(2 * d_expert_pad) * gateup_k_pack;
  const int64_t gateup_scale_expert_stride =
      static_cast<int64_t>(2 * d_expert_pad) * gateup_scale_k;
  const int64_t down_expert_stride =
      static_cast<int64_t>(d_hidden_pad) * down_k_pack;
  const int64_t down_scale_expert_stride =
      static_cast<int64_t>(d_hidden_pad) * down_scale_k;

  const int64_t gateup_base = static_cast<int64_t>(expert) * gateup_expert_stride;
  const int64_t gateup_scale_base =
      static_cast<int64_t>(expert) * gateup_scale_expert_stride;

  for (int32_t j = warp_id; j < d_expert; j += warps_per_block) {
    const uint8_t* gate_row = gate_up_weight + gateup_base + static_cast<int64_t>(j) * gateup_k_pack;
    const uint8_t* up_row =
        gate_up_weight + gateup_base + static_cast<int64_t>(d_expert_pad + j) * gateup_k_pack;
    const uint8_t* gate_scale_row =
        gate_up_weight_scale + gateup_scale_base + static_cast<int64_t>(j) * gateup_scale_k;
    const uint8_t* up_scale_row =
        gate_up_weight_scale + gateup_scale_base + static_cast<int64_t>(d_expert_pad + j) * gateup_scale_k;

    float gate_sum = 0.0f;
    float up_sum = 0.0f;
    for (int32_t col = lane; col < d_hidden; col += warpSize) {
      float x = x_sh[col];

      float gate_w =
          fp4x2_at(gate_row, col) * e8m0_to_float(gate_scale_row[col / kScaleBlock]);
      float up_w =
          fp4x2_at(up_row, col) * e8m0_to_float(up_scale_row[col / kScaleBlock]);

      gate_sum += x * gate_w;
      up_sum += x * up_w;
    }

    gate_sum = warp_sum(gate_sum);
    up_sum = warp_sum(up_sum);

    if (lane == 0) {
      inter_sh[j] = silu(gate_sum) * up_sum;
    }
  }
  __syncthreads();

  const int64_t down_base = static_cast<int64_t>(expert) * down_expert_stride;
  const int64_t down_scale_base =
      static_cast<int64_t>(expert) * down_scale_expert_stride;

  for (int32_t out_col = warp_id; out_col < d_hidden; out_col += warps_per_block) {
    float out_sum = 0.0f;
    if (out_col < d_hidden_pad) {
      const uint8_t* down_row =
          down_weight + down_base + static_cast<int64_t>(out_col) * down_k_pack;
      const uint8_t* down_scale_row =
          down_weight_scale + down_scale_base + static_cast<int64_t>(out_col) * down_scale_k;

      for (int32_t j = lane; j < d_expert; j += warpSize) {
        float w =
            fp4x2_at(down_row, j) * e8m0_to_float(down_scale_row[j / kScaleBlock]);
        out_sum += inter_sh[j] * w;
      }
    }

    out_sum = warp_sum(out_sum);
    if (lane == 0) {
      atomicAdd(
          output_accum + static_cast<int64_t>(token) * d_hidden + out_col,
          route_weight * out_sum);
    }
  }
}

__global__ void cast_float_to_bf16_kernel(
    const float* in,
    hip_bfloat16* out,
    int64_t total) {
  int64_t idx = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
  if (idx < total) {
    out[idx] = static_cast<hip_bfloat16>(in[idx]);
  }
}

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

}  // namespace


torch::Tensor hip_moe_forward(
    torch::Tensor hidden_states,
    torch::Tensor gate_up_weight,
    torch::Tensor down_weight,
    torch::Tensor gate_up_weight_scale,
    torch::Tensor down_weight_scale,
    torch::Tensor topk_weights,
    torch::Tensor topk_ids,
    int64_t d_expert) {
  TORCH_CHECK(hidden_states.is_cuda(), "hidden_states must be CUDA/HIP tensor");
  TORCH_CHECK(gate_up_weight.is_cuda(), "gate_up_weight must be CUDA/HIP tensor");
  TORCH_CHECK(down_weight.is_cuda(), "down_weight must be CUDA/HIP tensor");
  TORCH_CHECK(gate_up_weight_scale.is_cuda(), "gate_up_weight_scale must be CUDA/HIP tensor");
  TORCH_CHECK(down_weight_scale.is_cuda(), "down_weight_scale must be CUDA/HIP tensor");
  TORCH_CHECK(topk_weights.is_cuda(), "topk_weights must be CUDA/HIP tensor");
  TORCH_CHECK(topk_ids.is_cuda(), "topk_ids must be CUDA/HIP tensor");

  TORCH_CHECK(hidden_states.scalar_type() == torch::kBFloat16, "hidden_states must be bfloat16");
  TORCH_CHECK(gate_up_weight.element_size() == 1, "gate_up_weight must be 1-byte packed fp4x2");
  TORCH_CHECK(down_weight.element_size() == 1, "down_weight must be 1-byte packed fp4x2");
  TORCH_CHECK(gate_up_weight_scale.element_size() == 1, "gate_up_weight_scale must be 1-byte e8m0");
  TORCH_CHECK(down_weight_scale.element_size() == 1, "down_weight_scale must be 1-byte e8m0");
  TORCH_CHECK(topk_weights.scalar_type() == torch::kFloat32, "topk_weights must be float32");
  TORCH_CHECK(topk_ids.scalar_type() == torch::kInt32, "topk_ids must be int32");

  TORCH_CHECK(hidden_states.dim() == 2, "hidden_states must be [M, d_hidden]");
  TORCH_CHECK(gate_up_weight.dim() == 3, "gate_up_weight must be [E, 2*d_expert_pad, d_hidden_pad/2]");
  TORCH_CHECK(down_weight.dim() == 3, "down_weight must be [E, d_hidden_pad, d_expert_pad/2]");
  TORCH_CHECK(gate_up_weight_scale.dim() == 3, "gate_up_weight_scale must be [E, 2*d_expert_pad, d_hidden_pad/32]");
  TORCH_CHECK(down_weight_scale.dim() == 3, "down_weight_scale must be [E, d_hidden_pad, d_expert_pad/32]");
  TORCH_CHECK(topk_weights.dim() == 2, "topk_weights must be [M, topk]");
  TORCH_CHECK(topk_ids.dim() == 2, "topk_ids must be [M, topk]");

  hidden_states = hidden_states.contiguous();
  gate_up_weight = gate_up_weight.contiguous();
  down_weight = down_weight.contiguous();
  gate_up_weight_scale = gate_up_weight_scale.contiguous();
  down_weight_scale = down_weight_scale.contiguous();
  topk_weights = topk_weights.contiguous();
  topk_ids = topk_ids.contiguous();

  const int64_t m = hidden_states.size(0);
  const int64_t d_hidden = hidden_states.size(1);
  const int64_t topk = topk_ids.size(1);
  const int64_t num_experts = gate_up_weight.size(0);
  const int64_t d_hidden_pad = down_weight.size(1);
  const int64_t d_expert_pad = down_weight.size(2) * 2;

  TORCH_CHECK(topk_weights.size(0) == m, "topk_weights M mismatch");
  TORCH_CHECK(topk_ids.size(0) == m, "topk_ids M mismatch");
  TORCH_CHECK(topk_weights.size(1) == topk, "topk shape mismatch");

  TORCH_CHECK(gate_up_weight.size(0) == num_experts, "gate_up_weight E mismatch");
  TORCH_CHECK(down_weight.size(0) == num_experts, "down_weight E mismatch");
  TORCH_CHECK(gate_up_weight_scale.size(0) == num_experts, "gate_up_weight_scale E mismatch");
  TORCH_CHECK(down_weight_scale.size(0) == num_experts, "down_weight_scale E mismatch");

  TORCH_CHECK(gate_up_weight.size(1) == 2 * d_expert_pad, "gate_up_weight row mismatch vs d_expert_pad");
  TORCH_CHECK(gate_up_weight_scale.size(1) == 2 * d_expert_pad, "gate_up_weight_scale row mismatch");
  TORCH_CHECK(gate_up_weight.size(2) * 2 == d_hidden_pad, "gate_up_weight K mismatch vs d_hidden_pad");
  TORCH_CHECK(gate_up_weight_scale.size(2) == d_hidden_pad / kScaleBlock, "gate_up_weight_scale K mismatch");

  TORCH_CHECK(down_weight_scale.size(1) == d_hidden_pad, "down_weight_scale row mismatch");
  TORCH_CHECK(down_weight_scale.size(2) == d_expert_pad / kScaleBlock, "down_weight_scale K mismatch");

  TORCH_CHECK(d_hidden <= d_hidden_pad, "d_hidden must be <= d_hidden_pad");
  TORCH_CHECK(d_expert > 0 && d_expert <= d_expert_pad, "d_expert out of valid range");
  TORCH_CHECK(d_hidden_pad % kScaleBlock == 0, "d_hidden_pad must be divisible by 32");
  TORCH_CHECK(d_expert_pad % kScaleBlock == 0, "d_expert_pad must be divisible by 32");

  auto out_accum = torch::zeros({m, d_hidden}, hidden_states.options().dtype(torch::kFloat32));
  auto output = torch::empty({m, d_hidden}, hidden_states.options().dtype(torch::kBFloat16));

  if (m == 0 || d_hidden == 0 || topk == 0) {
    return output.zero_();
  }

  size_t smem_bytes = static_cast<size_t>(d_hidden + d_expert) * sizeof(float);

  hipLaunchKernelGGL(
      fused_moe_kernel,
      dim3(static_cast<unsigned int>(m), static_cast<unsigned int>(topk)),
      dim3(kThreads),
      smem_bytes,
      0,
      reinterpret_cast<const hip_bfloat16*>(hidden_states.data_ptr<at::BFloat16>()),
      reinterpret_cast<const uint8_t*>(gate_up_weight.data_ptr()),
      reinterpret_cast<const uint8_t*>(down_weight.data_ptr()),
      reinterpret_cast<const uint8_t*>(gate_up_weight_scale.data_ptr()),
      reinterpret_cast<const uint8_t*>(down_weight_scale.data_ptr()),
      topk_weights.data_ptr<float>(),
      topk_ids.data_ptr<int32_t>(),
      out_accum.data_ptr<float>(),
      static_cast<int32_t>(m),
      static_cast<int32_t>(topk),
      static_cast<int32_t>(num_experts),
      static_cast<int32_t>(d_hidden),
      static_cast<int32_t>(d_hidden_pad),
      static_cast<int32_t>(d_expert),
      static_cast<int32_t>(d_expert_pad),
      static_cast<int32_t>(gate_up_weight.size(2)),
      static_cast<int32_t>(down_weight.size(2)),
      static_cast<int32_t>(gate_up_weight_scale.size(2)),
      static_cast<int32_t>(down_weight_scale.size(2)));
  check_hip_error("fused_moe_kernel");

  int64_t total = m * d_hidden;
  int32_t blocks = static_cast<int32_t>((total + kThreads - 1) / kThreads);
  hipLaunchKernelGGL(
      cast_float_to_bf16_kernel,
      dim3(static_cast<unsigned int>(blocks)),
      dim3(kThreads),
      0,
      0,
      out_accum.data_ptr<float>(),
      reinterpret_cast<hip_bfloat16*>(output.data_ptr<at::BFloat16>()),
      total);
  check_hip_error("cast_float_to_bf16_kernel");

  return output;
}
"""


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="moe_mxfp4_hip_fused_v1",
                cpp_sources=[CPP_WRAPPER],
                cuda_sources=[HIP_SRC],
                functions=["hip_moe_forward"],
                verbose=False,
                extra_cuda_cflags=["-O3", "-std=c++20"],
            )
        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 _get_mode() -> str:
    mode = (os.getenv(_MODE_ENV, _DEFAULT_MODE) or _DEFAULT_MODE).strip().lower()
    if mode in {"base", "hip"}:
        return mode
    return _DEFAULT_MODE


def _normalize_scale_layout(
    scale: torch.Tensor,
    expected_shape: tuple[int, int, int],
) -> torch.Tensor:
    """
    Normalize e8m0 scale layout to [E, rows, scale_k].
    Runtime may provide either native 3D or flattened 2D buffers.
    """
    e, rows, k_blocks = expected_shape
    need = e * rows * k_blocks
    s = scale.contiguous()

    if s.dim() == 3:
        if tuple(s.shape) == expected_shape:
            return s
        flat = s.reshape(-1)
    else:
        flat = s.reshape(-1)

    if flat.numel() < need:
        raise RuntimeError(
            f"scale buffer too small: got {flat.numel()} elems, need {need} for shape {expected_shape}"
        )
    return flat[:need].view(expected_shape).contiguous()


def _run_base_fused_moe(data: input_t) -> output_t:
    from aiter import ActivationType, QuantType
    from aiter.fused_moe import fused_moe

    (
        hidden_states,
        _gate_up_weight,
        _down_weight,
        _gate_up_weight_scale,
        _down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
    intermediate_pad = int(config["d_expert_pad"] - config["d_expert"])

    return fused_moe(
        hidden_states.contiguous(),
        gate_up_weight_shuffled.contiguous(),
        down_weight_shuffled.contiguous(),
        topk_weights.to(torch.float32).contiguous(),
        topk_ids.to(torch.int32).contiguous(),
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled.contiguous(),
        w2_scale=down_weight_scale_shuffled.contiguous(),
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )


def _run_hip_fused(data: input_t) -> output_t:
    (
        hidden_states,
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        _gate_up_weight_shuffled,
        _down_weight_shuffled,
        _gate_up_weight_scale_shuffled,
        _down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    module = _get_hip_module()
    n_experts = int(config["n_routed_experts"] + config["n_shared_experts"])
    d_hidden_pad = int(config["d_hidden_pad"])
    d_expert_pad = int(config["d_expert_pad"])

    gate_up_weight_scale = _normalize_scale_layout(
        gate_up_weight_scale,
        (n_experts, 2 * d_expert_pad, d_hidden_pad // 32),
    )
    down_weight_scale = _normalize_scale_layout(
        down_weight_scale,
        (n_experts, d_hidden_pad, d_expert_pad // 32),
    )

    return module.hip_moe_forward(
        hidden_states.contiguous(),
        gate_up_weight.contiguous(),
        down_weight.contiguous(),
        gate_up_weight_scale.contiguous(),
        down_weight_scale.contiguous(),
        topk_weights.to(torch.float32).contiguous(),
        topk_ids.to(torch.int32).contiguous(),
        int(config["d_expert"]),
    )


def custom_kernel(data: input_t) -> output_t:
    if _get_mode() == "hip":
        return _run_hip_fused(data)
    return _run_base_fused_moe(data)
scrolls · 505 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