KernelBench cuda · RTX PRO 6000

GLM-5.2 Fused MoE Grok 4.7

9.48%geomean peak fraction across shapes

Grok 4.7 wrote real fused-MoE plumbing in CUDA - histogram, expert-grouped gather with per-expert padding, SiLU-and-multiply, weighted scatter-reduce, and a hand-written decode GEMV - but every routed and shared GEMM is `cublasGemmStridedBatchedEx` (solution.py:297-301), not authored MMA. Same class as the published grok-4.6 (0.0939) and deepseek-v4-pro (0.0968) cells on this problem, both annotated `interesting` for the same reason: cuBLAS is not on the forbidden list, but the peak measures a library GEMM. Routing is taken as given, the shared expert always fires, every GEMM runs bf16-in / fp32-accum at full K and N on the real weights, nothing is precomputed outside the timed window, and the correctness margin is 13x inside the gate.

harnessgrokagent session1h 23mtotal wall1h 41mcheck4mbenchmark4moutput tokens—gpu-lock wait10mgpu-lock held9mregimecompute

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

shape 011.140 ms33.3%1.17 TB/s · 65% of 1.8 TB/s HBM · also 167 TFLOPS (33% of compute)
shape 111.425 ms32.7%1.14 TB/s · 63% of 1.8 TB/s HBM · also 164 TFLOPS (33% of compute)
shape 20.321 ms0.3%40.31 TB/s · 100% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
shape 318.023 ms41.2%206 TFLOPS · 41% of 500 TF bf16 peak · also 0.73 TB/s (40% of HBM)
shape 48.464 ms5.5%1.53 TB/s · 85% of 1.8 TB/s HBM · also 27 TFLOPS (5% of compute)
shape 58.667 ms10.4%1.49 TB/s · 83% of 1.8 TB/s HBM · also 52 TFLOPS (10% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(33.3% · 32.7% · 0.3% · 41.2% · 5.5% · 10.4%) = 9.5%

Kernel source (redacted)
"""GLM-5.2 fused MoE: CUDA gather + cuBLAS tensor-core GEMM + CUDA epilogue.

Per expert: h = silu(x @ gate.T) * (x @ up.T); y = h @ down.T
out = sum_shared(y_s) + sum_k weight[t, k] * y_routed[e_k]

Large T: tokens grouped by expert, padded to a cuBLAS-friendly M, one
strided-batched GEMM per projection (weights read once). T=1: fused GEMV
over the 8 routed experts + the shared expert only.
"""
from __future__ import annotations

import os

os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
os.environ.setdefault("CUDA_HOME", "/usr/local/cuda")
os.environ.setdefault("MAX_JOBS", "4")

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline

_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>

#include <algorithm>
#include <cstdint>
#include <vector>

#define CHECK_CUBLAS(expr)                                                     \
  do {                                                                         \
    cublasStatus_t _st = (expr);                                               \
    TORCH_CHECK(_st == CUBLAS_STATUS_SUCCESS, "cublas ", int(_st), " @ ",     \
                __LINE__);                                                     \
  } while (0)

struct Ctx {
  cublasHandle_t handle = nullptr;
  int device = -1;
  torch::Tensor workspace;
  torch::Tensor counts_h;  // pinned int32
  torch::Tensor Xg, gate_up, hidden, Y;
  torch::Tensor shared_gate, shared_h, shared_y;
  torch::Tensor counts_d, cursor_d, row_of_pair;
  int cap_E = 0;
  int cap_pairs = 0;
  int64_t cap_rows = 0;
  int cap_T = 0;
  int cap_I = 0;
  int cap_H = 0;
};

static Ctx& ctx() {
  static Ctx c;
  return c;
}

static void ensure_handle() {
  Ctx& c = ctx();
  int dev = 0;
  cudaGetDevice(&dev);
  if (c.handle && c.device == dev) {
    CHECK_CUBLAS(cublasSetStream(c.handle, at::cuda::getCurrentCUDAStream()));
    return;
  }
  if (c.handle) {
    cublasDestroy(c.handle);
    c.handle = nullptr;
  }
  c.device = dev;
  CHECK_CUBLAS(cublasCreate(&c.handle));
  CHECK_CUBLAS(cublasSetStream(c.handle, at::cuda::getCurrentCUDAStream()));
  cublasSetMathMode(c.handle, CUBLAS_DEFAULT_MATH);
  auto opt = torch::TensorOptions().device(torch::kCUDA).dtype(torch::kUInt8);
  c.workspace = torch::empty({256 << 20}, opt);
  CHECK_CUBLAS(cublasSetWorkspace_v2(c.handle, c.workspace.data_ptr(),
                                     c.workspace.numel()));
}

// Measured on RTX PRO 6000 sm_120: strided-batched bf16 GEMM falls off a cliff
// just above M=256 (7.5 ms -> 10.5 ms) and again above 384. Stay on the fast tiles.
static int choose_mpad(int maxc) {
  if (maxc <= 0) return 0;
  const int good[] = {32, 48, 64, 96, 112, 128, 160, 176, 192,
                      208, 224, 240, 256, 384};
  for (int g : good) {
    if (g >= maxc) return g;
  }
  return (maxc + 63) & ~63;
}

__device__ __forceinline__ float silu_f(float x) {
  float y = expf(-fabsf(x));
  float sig = (x >= 0.f) ? (1.f / (1.f + y)) : (y / (1.f + y));
  return x * sig;
}

__device__ __forceinline__ float dot8(uint4 a, uint4 b) {
  const __nv_bfloat16* ap = reinterpret_cast<const __nv_bfloat16*>(&a);
  const __nv_bfloat16* bp = reinterpret_cast<const __nv_bfloat16*>(&b);
  float acc = 0.f;
#pragma unroll
  for (int j = 0; j < 8; ++j) {
    acc = fmaf(__bfloat162float(ap[j]), __bfloat162float(bp[j]), acc);
  }
  return acc;
}

__device__ __forceinline__ uint4 ldcs_u4(const uint4* p) {
  uint4 v;
  asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
               : "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
               : "l"(p));
  return v;
}

__global__ void histogram_i64(const int64_t* __restrict__ ids, int n, int E,
                              int* __restrict__ counts) {
  int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i < n) {
    int64_t e = ids[i];
    if ((unsigned)e < (unsigned)E) atomicAdd(counts + e, 1);
  }
}

// One block per (token, slot). Assigns a row inside the expert's M_pad tile and copies x.
__global__ void assign_and_gather(const int64_t* __restrict__ ids,
                                  const __nv_bfloat16* __restrict__ x,
                                  __nv_bfloat16* __restrict__ Xg,
                                  int* __restrict__ cursor,
                                  int* __restrict__ row_of_pair,
                                  int n_pairs, int top_k, int H, int M_pad) {
  int pair = blockIdx.x;
  if (pair >= n_pairs) return;
  int t = pair / top_k;
  int64_t e64 = ids[pair];
  int e = (int)e64;
  int local = 0;
  if (threadIdx.x == 0) {
    local = atomicAdd(cursor + e, 1);
    row_of_pair[pair] = e * M_pad + local;
  }
  local = row_of_pair[pair] - e * M_pad;
  // row_of_pair is written by thread 0; publish before the copy.
  __syncthreads();
  int row = row_of_pair[pair];
  const uint4* src =
      reinterpret_cast<const uint4*>(x + (int64_t)t * H);
  uint4* dst = reinterpret_cast<uint4*>(Xg + (int64_t)row * H);
  int nvec = H >> 3;
  for (int i = threadIdx.x; i < nvec; i += blockDim.x) {
    dst[i] = src[i];
  }
  (void)local;
}

__global__ void zero_pad_rows(const int* __restrict__ counts,
                              __nv_bfloat16* __restrict__ Xg, int E, int M_pad,
                              int H) {
  int e = blockIdx.y;
  if (e >= E) return;
  int local = counts[e] + blockIdx.x;
  if (local >= M_pad) return;
  uint4* dst =
      reinterpret_cast<uint4*>(Xg + ((int64_t)e * M_pad + local) * H);
  int nvec = H >> 3;
  for (int i = threadIdx.x; i < nvec; i += blockDim.x) {
    uint4 z = make_uint4(0, 0, 0, 0);
    dst[i] = z;
  }
}

__global__ void add_bf16_kernel(const __nv_bfloat16* __restrict__ src,
                              float* __restrict__ dst, int n, int accumulate) {
  int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i >= n) return;
  float v = __bfloat162float(src[i]);
  dst[i] = accumulate ? (dst[i] + v) : v;
}

__global__ void cast_fp32_bf16_kernel(const float* __restrict__ src,
                                      __nv_bfloat16* __restrict__ dst, int n) {
  int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i < n) dst[i] = __float2bfloat16(src[i]);
}

__global__ void silu_mul_kernel(const __nv_bfloat16* __restrict__ gate_up,
                                __nv_bfloat16* __restrict__ hidden, int rows,
                                int I) {
  int row = blockIdx.x;
  if (row >= rows) return;
  const __nv_bfloat16* g = gate_up + (int64_t)row * (I * 2);
  const __nv_bfloat16* u = g + I;
  __nv_bfloat16* h = hidden + (int64_t)row * I;
  for (int i = threadIdx.x; i < I; i += blockDim.x) {
    float hv = silu_f(__bfloat162float(g[i])) * __bfloat162float(u[i]);
    h[i] = __float2bfloat16(hv);
  }
}

// out[t] = shared[t] + sum_k weight[t,k] * Y[row_of_pair[t,k]]
template <int TOPK>
__global__ void reduce_kernel(const __nv_bfloat16* __restrict__ shared,
                              const __nv_bfloat16* __restrict__ Y,
                              const int* __restrict__ row_of_pair,
                              const __nv_bfloat16* __restrict__ weights,
                              __nv_bfloat16* __restrict__ out, int T, int H,
                              int top_k) {
  int t = blockIdx.x;
  if (t >= T) return;
  const int kmax = TOPK > 0 ? TOPK : top_k;
  for (int h = threadIdx.x; h < H; h += blockDim.x) {
    float acc = __bfloat162float(shared[(int64_t)t * H + h]);
    const int* rows = row_of_pair + (int64_t)t * kmax;
    const __nv_bfloat16* w = weights + (int64_t)t * kmax;
#pragma unroll 1
    for (int k = 0; k < kmax; ++k) {
      int row = rows[k];
      float yv = __bfloat162float(Y[(int64_t)row * H + h]);
      acc = fmaf(__bfloat162float(w[k]), yv, acc);
    }
    out[(int64_t)t * H + h] = __float2bfloat16(acc);
  }
}

// T=1 GEMV. One warp = one output. Streaming weight loads; x stays in L2.
// grid: (ceil(N / warps_per_block), n_exp)
constexpr int GEMV_THREADS = 128;

__global__ void gemv_nt_kernel(const __nv_bfloat16* __restrict__ x_base,
                               int64_t x_stride,
                               const __nv_bfloat16* __restrict__ W_base,
                               const int* __restrict__ expert_index,
                               int64_t expert_stride,
                               __nv_bfloat16* __restrict__ y, int N, int K) {
  const int lane = threadIdx.x & 31;
  const int warp = threadIdx.x >> 5;
  const int warps = blockDim.x >> 5;
  const int n_out = blockIdx.x * warps + warp;
  const int exp = blockIdx.y;
  if (n_out >= N) return;

  const __nv_bfloat16* x = x_base + (int64_t)exp * x_stride;
  int e = expert_index[exp];
  const __nv_bfloat16* row =
      W_base + ((int64_t)e * expert_stride) + (int64_t)n_out * K;
  const uint4* row4 = reinterpret_cast<const uint4*>(row);
  const uint4* x4 = reinterpret_cast<const uint4*>(x);
  const int nvec = K >> 3;
  float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
  int i = lane;
  for (; i + 96 < nvec; i += 128) {
    uint4 w0 = ldcs_u4(row4 + i);
    uint4 w1 = ldcs_u4(row4 + i + 32);
    uint4 w2 = ldcs_u4(row4 + i + 64);
    uint4 w3 = ldcs_u4(row4 + i + 96);
    a0 += dot8(w0, x4[i]);
    a1 += dot8(w1, x4[i + 32]);
    a2 += dot8(w2, x4[i + 64]);
    a3 += dot8(w3, x4[i + 96]);
  }
  float acc = (a0 + a1) + (a2 + a3);
  for (; i < nvec; i += 32) acc += dot8(ldcs_u4(row4 + i), x4[i]);
#pragma unroll
  for (int off = 16; off > 0; off >>= 1)
    acc += __shfl_down_sync(0xffffffff, acc, off);
  if (lane == 0) y[(int64_t)exp * N + n_out] = __float2bfloat16(acc);
}

__global__ void reduce_t1_kernel(const __nv_bfloat16* __restrict__ y_routed,
                                 const __nv_bfloat16* __restrict__ y_shared,
                                 const __nv_bfloat16* __restrict__ weights,
                                 __nv_bfloat16* __restrict__ out, int top_k,
                                 int n_shared, int H) {
  for (int h = blockIdx.x * blockDim.x + threadIdx.x; h < H;
       h += blockDim.x * gridDim.x) {
    float acc = 0.f;
    for (int s = 0; s < n_shared; ++s) {
      acc += __bfloat162float(y_shared[(int64_t)s * H + h]);
    }
    for (int k = 0; k < top_k; ++k) {
      float yv = __bfloat162float(y_routed[(int64_t)k * H + h]);
      acc = fmaf(__bfloat162float(weights[k]), yv, acc);
    }
    out[h] = __float2bfloat16(acc);
  }
}

// C = X @ W^T
// W: (batch, N, K) bf16, X: (batch, M, K) bf16, C: (batch, M, N) bf16
static void gemm_x_wt(int M, int N, int K, int batch, const void* W,
                      const void* X, void* C) {
  float alpha = 1.f, beta = 0.f;
  CHECK_CUBLAS(cublasGemmStridedBatchedEx(
      ctx().handle, CUBLAS_OP_T, CUBLAS_OP_N, N, M, K, &alpha, W, CUDA_R_16BF,
      K, (long long)N * K, X, CUDA_R_16BF, K, (long long)M * K, &beta, C,
      CUDA_R_16BF, N, (long long)M * N, batch, CUBLAS_COMPUTE_32F,
      CUBLAS_GEMM_DEFAULT));
}

static void launch_gemv(const __nv_bfloat16* x_base, int64_t x_stride,
                        const __nv_bfloat16* W_base, const int* expert_index,
                        int64_t expert_stride, __nv_bfloat16* y, int n_exp,
                        int N, int K, cudaStream_t stream) {
  constexpr int warps = GEMV_THREADS / 32;
  dim3 grid((N + warps - 1) / warps, n_exp);
  gemv_nt_kernel<<<grid, GEMV_THREADS, 0, stream>>>(
      x_base, x_stride, W_base, expert_index, expert_stride, y, N, K);
}

static torch::Tensor forward_t1(torch::Tensor x, torch::Tensor expert_ids,
                                torch::Tensor expert_weights,
                                torch::Tensor w1_routed, torch::Tensor w2_routed,
                                torch::Tensor w1_shared, torch::Tensor w2_shared) {
  const int H = (int)x.size(1);
  const int top_k = (int)expert_ids.size(1);
  const int E = (int)w1_routed.size(0);
  const int twoI = (int)w1_routed.size(1);
  const int I = twoI / 2;
  const int n_shared = (int)w1_shared.size(0);
  TORCH_CHECK(H % 8 == 0 && I % 8 == 0, "H and I must be multiples of 8");
  TORCH_CHECK(top_k > 0 && n_shared >= 0, "bad top_k/n_shared");

  auto opt_i = torch::TensorOptions().dtype(torch::kInt32).device(x.device());
  auto opt_b = x.options();
  auto exp_idx = torch::empty({top_k + std::max(n_shared, 1)}, opt_i);

  // ids are 8 int64s; pull them so the GEMV knows which rows of W to read.
  auto ids_cpu = expert_ids.contiguous().view({-1}).to(torch::kCPU);
  const int64_t* ids = ids_cpu.data_ptr<int64_t>();
  std::vector<int> hidx(top_k + n_shared);
  for (int k = 0; k < top_k; ++k) {
    int e = (int)ids[k];
    TORCH_CHECK(e >= 0 && e < E, "expert id out of range");
    hidx[k] = e;
  }
  for (int s = 0; s < n_shared; ++s) hidx[top_k + s] = s;
  cudaMemcpyAsync(exp_idx.data_ptr<int>(), hidx.data(),
                  (size_t)(top_k + n_shared) * sizeof(int),
                  cudaMemcpyHostToDevice, at::cuda::getCurrentCUDAStream());

  auto y_routed = torch::empty({top_k, H}, opt_b);
  auto y_shared = torch::empty({std::max(n_shared, 1), H}, opt_b);
  auto gate_r = torch::empty({top_k, twoI}, opt_b);
  auto hid_r = torch::empty({top_k, I}, opt_b);
  auto gate_s = torch::empty({std::max(n_shared, 1), twoI}, opt_b);
  auto hid_s = torch::empty({std::max(n_shared, 1), I}, opt_b);

  cudaStream_t stream = at::cuda::getCurrentCUDAStream();
  const __nv_bfloat16* xptr =
      reinterpret_cast<const __nv_bfloat16*>(x.data_ptr());

  if (top_k > 0) {
    launch_gemv(xptr, /*x_stride=*/0,
                reinterpret_cast<const __nv_bfloat16*>(w1_routed.data_ptr()),
                exp_idx.data_ptr<int>(), (int64_t)twoI * H,
                reinterpret_cast<__nv_bfloat16*>(gate_r.data_ptr()), top_k,
                twoI, H, stream);
    silu_mul_kernel<<<top_k, 256, 0, stream>>>(
        reinterpret_cast<const __nv_bfloat16*>(gate_r.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(hid_r.data_ptr()), top_k, I);
    launch_gemv(reinterpret_cast<const __nv_bfloat16*>(hid_r.data_ptr()),
                /*x_stride=*/I,
                reinterpret_cast<const __nv_bfloat16*>(w2_routed.data_ptr()),
                exp_idx.data_ptr<int>(), (int64_t)H * I,
                reinterpret_cast<__nv_bfloat16*>(y_routed.data_ptr()), top_k,
                H, I, stream);
  }
  if (n_shared > 0) {
    launch_gemv(xptr, 0,
                reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr()),
                exp_idx.data_ptr<int>() + top_k, (int64_t)twoI * H,
                reinterpret_cast<__nv_bfloat16*>(gate_s.data_ptr()), n_shared,
                twoI, H, stream);
    silu_mul_kernel<<<n_shared, 256, 0, stream>>>(
        reinterpret_cast<const __nv_bfloat16*>(gate_s.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(hid_s.data_ptr()), n_shared, I);
    // shared experts are indexed 0..n_shared-1 inside w1_shared, which we
    // stored at exp_idx[top_k + s].
    launch_gemv(reinterpret_cast<const __nv_bfloat16*>(hid_s.data_ptr()), I,
                reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr()),
                exp_idx.data_ptr<int>() + top_k, (int64_t)H * I,
                reinterpret_cast<__nv_bfloat16*>(y_shared.data_ptr()), n_shared,
                H, I, stream);
  }

  auto out = torch::empty({1, H}, opt_b);
  reduce_t1_kernel<<<(H + 255) / 256, 256, 0, stream>>>(
      reinterpret_cast<const __nv_bfloat16*>(y_routed.data_ptr()),
      reinterpret_cast<const __nv_bfloat16*>(y_shared.data_ptr()),
      reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
      reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), top_k, n_shared, H);
  return out;
}

static torch::Tensor forward_grouped(torch::Tensor x, torch::Tensor expert_ids,
                                     torch::Tensor expert_weights,
                                     torch::Tensor w1_routed,
                                     torch::Tensor w2_routed,
                                     torch::Tensor w1_shared,
                                     torch::Tensor w2_shared) {
  ensure_handle();
  const int T = (int)x.size(0);
  const int H = (int)x.size(1);
  const int top_k = (int)expert_ids.size(1);
  const int E = (int)w1_routed.size(0);
  const int twoI = (int)w1_routed.size(1);
  const int I = twoI / 2;
  const int n_shared = (int)w1_shared.size(0);
  const int n_pairs = T * top_k;
  TORCH_CHECK(H % 8 == 0 && I % 8 == 0, "H and I must be multiples of 8");
  TORCH_CHECK(w1_routed.size(2) == H && w2_routed.size(1) == H, "H mismatch");
  TORCH_CHECK(w2_routed.size(2) == I && w2_routed.size(0) == E, "I/E mismatch");

  Ctx& c = ctx();
  auto dev = x.device();
  if (c.cap_E < E) {
    c.counts_h = torch::empty(
        {E}, torch::TensorOptions().dtype(torch::kInt32).pinned_memory(true));
    c.counts_d = torch::empty(
        {E}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
    c.cursor_d = torch::empty(
        {E}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
    c.cap_E = E;
  }
  if (c.cap_pairs < n_pairs) {
    c.row_of_pair = torch::empty(
        {n_pairs}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
    c.cap_pairs = n_pairs;
  }
  if (c.cap_T < T || c.cap_H < H || c.cap_I < I) {
    auto bf = x.options();
    c.shared_gate = torch::empty({(int64_t)T * twoI}, bf);
    c.shared_h = torch::empty({(int64_t)T * I}, bf);
    c.shared_y = torch::empty({(int64_t)T * H}, bf);
    c.cap_T = T;
    c.cap_H = H;
    c.cap_I = I;
  }

  cudaStream_t stream = at::cuda::getCurrentCUDAStream();
  cudaMemsetAsync(c.counts_d.data_ptr(), 0, (size_t)E * sizeof(int), stream);
  histogram_i64<<<(n_pairs + 255) / 256, 256, 0, stream>>>(
      expert_ids.data_ptr<int64_t>(), n_pairs, E, c.counts_d.data_ptr<int>());
  cudaMemcpyAsync(c.counts_h.data_ptr<int>(), c.counts_d.data_ptr<int>(),
                  (size_t)E * sizeof(int), cudaMemcpyDeviceToHost, stream);
  cudaStreamSynchronize(stream);

  const int* hc = c.counts_h.data_ptr<int>();
  int maxc = 0;
  for (int e = 0; e < E; ++e) maxc = std::max(maxc, hc[e]);
  int M_pad = choose_mpad(std::max(maxc, 1));
  // Always batch all E experts. Empty experts have zero rows of x and are
  // never scattered into the output. At serving T this set is full anyway;
  // skipping the rare empty expert would force a weight compact copy.
  int64_t rows = (int64_t)E * M_pad;
  if (c.cap_rows < rows || c.cap_H < H || c.cap_I < I) {
    auto bf = x.options();
    c.Xg = torch::empty({rows, H}, bf);
    c.gate_up = torch::empty({rows, twoI}, bf);
    c.hidden = torch::empty({rows, I}, bf);
    c.Y = torch::empty({rows, H}, bf);
    c.cap_rows = rows;
  }

  // Zero only the padded tail of each expert (not the rows gather overwrites).
  int pad_span = 0;
  for (int e = 0; e < E; ++e) pad_span = std::max(pad_span, M_pad - hc[e]);
  if (pad_span > 0) {
    // counts_d still holds the real counts; cursor is separate.
    zero_pad_rows<<<dim3(pad_span, E), 128, 0, stream>>>(
        c.counts_d.data_ptr<int>(),
        reinterpret_cast<__nv_bfloat16*>(c.Xg.data_ptr()), E, M_pad, H);
  }
  cudaMemsetAsync(c.cursor_d.data_ptr(), 0, (size_t)E * sizeof(int), stream);
  assign_and_gather<<<n_pairs, 128, 0, stream>>>(
      expert_ids.data_ptr<int64_t>(),
      reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
      reinterpret_cast<__nv_bfloat16*>(c.Xg.data_ptr()),
      c.cursor_d.data_ptr<int>(), c.row_of_pair.data_ptr<int>(), n_pairs, top_k,
      H, M_pad);

  // Routed w1: (E, M_pad, H) @ (E, 2I, H)^T -> (E, M_pad, 2I)
  gemm_x_wt(M_pad, twoI, H, E, w1_routed.data_ptr(), c.Xg.data_ptr(),
            c.gate_up.data_ptr());
  silu_mul_kernel<<<(int)rows, 256, 0, stream>>>(
      reinterpret_cast<const __nv_bfloat16*>(c.gate_up.data_ptr()),
      reinterpret_cast<__nv_bfloat16*>(c.hidden.data_ptr()), (int)rows, I);
  gemm_x_wt(M_pad, H, I, E, w2_routed.data_ptr(), c.hidden.data_ptr(),
            c.Y.data_ptr());

  // Shared experts. n_shared==1 writes bf16 directly; more than one sums in fp32.
  torch::Tensor shared_bf;
  if (n_shared <= 0) {
    shared_bf = torch::zeros({T, H}, x.options());
  } else if (n_shared == 1) {
    const __nv_bfloat16* w1s =
        reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr());
    const __nv_bfloat16* w2s =
        reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr());
    gemm_x_wt(T, twoI, H, 1, w1s, x.data_ptr(), c.shared_gate.data_ptr());
    silu_mul_kernel<<<T, 256, 0, stream>>>(
        reinterpret_cast<const __nv_bfloat16*>(c.shared_gate.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(c.shared_h.data_ptr()), T, I);
    gemm_x_wt(T, H, I, 1, w2s, c.shared_h.data_ptr(), c.shared_y.data_ptr());
    // shared_y is cached at the high-water T; only the first T rows are live.
    shared_bf = c.shared_y.narrow(0, 0, (int64_t)T * H).view({T, H});
  } else {
    auto shared_fp = torch::empty({T, H}, x.options().dtype(torch::kFloat32));
    const int n_out = T * H;
    for (int s = 0; s < n_shared; ++s) {
      const __nv_bfloat16* w1s =
          reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr()) +
          (int64_t)s * twoI * H;
      const __nv_bfloat16* w2s =
          reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr()) +
          (int64_t)s * H * I;
      gemm_x_wt(T, twoI, H, 1, w1s, x.data_ptr(), c.shared_gate.data_ptr());
      silu_mul_kernel<<<T, 256, 0, stream>>>(
          reinterpret_cast<const __nv_bfloat16*>(c.shared_gate.data_ptr()),
          reinterpret_cast<__nv_bfloat16*>(c.shared_h.data_ptr()), T, I);
      gemm_x_wt(T, H, I, 1, w2s, c.shared_h.data_ptr(), c.shared_y.data_ptr());
      add_bf16_kernel<<<(n_out + 255) / 256, 256, 0, stream>>>(
          reinterpret_cast<const __nv_bfloat16*>(c.shared_y.data_ptr()),
          shared_fp.data_ptr<float>(), n_out, s > 0);
    }
    shared_bf = torch::empty({T, H}, x.options());
    cast_fp32_bf16_kernel<<<(n_out + 255) / 256, 256, 0, stream>>>(
        shared_fp.data_ptr<float>(),
        reinterpret_cast<__nv_bfloat16*>(shared_bf.data_ptr()), n_out);
  }

  auto out = torch::empty({T, H}, x.options());
  auto shared_c = shared_bf.contiguous();
  if (top_k == 8) {
    reduce_kernel<8><<<T, 256, 0, stream>>>(
        reinterpret_cast<const __nv_bfloat16*>(shared_c.data_ptr()),
        reinterpret_cast<const __nv_bfloat16*>(c.Y.data_ptr()),
        c.row_of_pair.data_ptr<int>(),
        reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), T, H, top_k);
  } else {
    reduce_kernel<0><<<T, 256, 0, stream>>>(
        reinterpret_cast<const __nv_bfloat16*>(shared_c.data_ptr()),
        reinterpret_cast<const __nv_bfloat16*>(c.Y.data_ptr()),
        c.row_of_pair.data_ptr<int>(),
        reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), T, H, top_k);
  }
  return out;
}

torch::Tensor fused_moe(torch::Tensor x, torch::Tensor expert_ids,
                        torch::Tensor expert_weights, torch::Tensor w1_routed,
                        torch::Tensor w2_routed, torch::Tensor w1_shared,
                        torch::Tensor w2_shared) {
  TORCH_CHECK(x.is_cuda(), "x must be CUDA");
  x = x.contiguous();
  expert_ids = expert_ids.contiguous();
  expert_weights = expert_weights.contiguous();
  w1_routed = w1_routed.contiguous();
  w2_routed = w2_routed.contiguous();
  w1_shared = w1_shared.contiguous();
  w2_shared = w2_shared.contiguous();
  if (expert_ids.scalar_type() != torch::kLong) {
    expert_ids = expert_ids.to(torch::kLong);
  }
  if (x.scalar_type() != torch::kBFloat16) {
    x = x.to(torch::kBFloat16);
  }
  TORCH_CHECK(x.dim() == 2, "x is (T, H)");
  const int T = (int)x.size(0);
  if (T == 0) {
    return torch::empty({0, x.size(1)}, x.options());
  }
  ensure_handle();
  if (T == 1) {
    return forward_t1(x, expert_ids, expert_weights, w1_routed, w2_routed,
                      w1_shared, w2_shared);
  }
  return forward_grouped(x, expert_ids, expert_weights, w1_routed, w2_routed,
                         w1_shared, w2_shared);
}
"""

_CPP = r"""
#include <torch/extension.h>
torch::Tensor fused_moe(torch::Tensor x, torch::Tensor expert_ids,
                        torch::Tensor expert_weights, torch::Tensor w1_routed,
                        torch::Tensor w2_routed, torch::Tensor w1_shared,
                        torch::Tensor w2_shared);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("fused_moe", &fused_moe, "GLM-5.2 fused MoE");
}
"""

_EXT = None


def _ext():
    global _EXT
    if _EXT is None:
        _EXT = load_inline(
            name="glm52_fused_moe_v1",
            cpp_sources=[_CPP],
            cuda_sources=[_CUDA],
            extra_ldflags=["-lcublas"],
            extra_cuda_cflags=["-O3", "-std=c++20"],
            extra_cflags=["-O3"],
            verbose=False,
        )
    return _EXT


class Model(nn.Module):
    def __init__(self, T: int, E: int, top_k: int, n_shared: int, H: int, I: int):
        super().__init__()
        self.T, self.E, self.top_k = T, E, top_k
        self.n_shared, self.H, self.I = n_shared, H, I
        self.w1_routed = nn.Parameter(torch.empty(E, 2 * I, H, dtype=torch.bfloat16))
        self.w2_routed = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
        self.w1_shared = nn.Parameter(torch.empty(n_shared, 2 * I, H, dtype=torch.bfloat16))
        self.w2_shared = nn.Parameter(torch.empty(n_shared, H, I, dtype=torch.bfloat16))
        for p in self.parameters():
            nn.init.normal_(p, std=0.02)

    def forward(self, x, expert_ids, expert_weights):
        return _ext().fused_moe(
            x, expert_ids, expert_weights,
            self.w1_routed, self.w2_routed, self.w1_shared, self.w2_shared,
        )

20260917_005805_grok_grok-4.7_01_glm52_fused_moe