kernelbench.com

KernelBench cuda · H100

GLM-5.2 Fused MoE Claude Opus 5

7.88%geomean peak fraction across shapes

manually audited: clean

Hand-written Hopper CUDA: wgmma.mma_async.m64n256k16 grouped GEMMs with hand-built 128B-swizzled GMMA smem descriptors, cp.async multistage pipeline, fused SwiGLU epilogue in registers, routing kernels (histogram -> single-CTA scan -> scatter), dedicated bandwidth-bound path for T<=8 decode. Direct nvcc -arch=sm_90a + ctypes. _ws[T] caches only workspace buffers; every forward launches the full chain on live pointers, no output cache. Grader files Read-only (tool_use parse), template_mutated false. The 8 sibling-run references are ps/argv listings of the concurrently running NSA agent on the same box -- benign, no archive reads. Passed full check.py + numeric stress on the sequential isolated re-grade (idle H100, 2026-07-27); clean 0.0788 vs contended 0.0801.

harnessor-opusagent session3h 2mtotal wall3h 11mcheck5mbenchmark5moutput tokensgpu-lock wait16sgpu-lock held1h 33mregimecompute

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

shape 08.103 ms30.3%1.60 TB/s · 79% of 2.0 TB/s HBM · also 229 TFLOPS (30% of compute)
shape 18.192 ms30.2%1.59 TB/s · 78% of 2.0 TB/s HBM · also 228 TFLOPS (30% of compute)
shape 20.304 ms0.2%42.60 TB/s · 2089% of 2.0 TB/s HBM · also 1 TFLOPS (0% of compute)
shape 313.597 ms36.1%0.96 TB/s · 47% of 2.0 TB/s HBM · also 273 TFLOPS (36% of compute)
shape 46.999 ms4.4%1.85 TB/s · 91% of 2.0 TB/s HBM · also 33 TFLOPS (4% of compute)
shape 57.130 ms8.4%1.82 TB/s · 89% of 2.0 TB/s HBM · also 64 TFLOPS (8% of compute)

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

geomean(30.3% · 30.2% · 0.2% · 36.1% · 4.4% · 8.4%) = 7.9%

Kernel source (redacted)
"""GLM-5.2 fused MoE (E=256 routed top-8 + 1 always-on shared expert) for H100.

Real CUDA: hand-written Hopper SM90a WGMMA (wgmma.mma_async.sync.aligned.m64n256k16)
grouped GEMMs with a cp.async multi-stage smem pipeline, 128B-swizzled GMMA smem
descriptors built by hand, and a fused SwiGLU epilogue.

Pipeline per forward:
  1. route      : histogram expert loads, exclusive scan, scatter (token,slot) pairs
                  into expert-contiguous order, build the (expert, m-tile) tile list.
                  The shared expert is treated as expert index E with weight 1.0, so
                  all work becomes one grouped GEMM over T*(top_k+n_shared) pairs.
  2. gemm1      : grouped [192 x 256 x 64] WGMMA GEMM over the permuted rows, one
                  warpgroup per 64 rows of the tile.  The 256-wide N tile packs 128
                  gate columns and the matching 128 up columns so silu(gate)*up is
                  fused entirely in registers (accumulator reg i pairs with reg
                  i+64), writing bf16 h.
  3. gemm2      : grouped [192 x 256 x 64] WGMMA GEMM h @ w2^T, routing weight folded
                  into the epilogue, writing bf16 y in permuted order.
  4. reduce     : fp32 sum of the 9 permuted rows belonging to each token -> bf16 out.

T <= 8 (decode microbatch) uses a dedicated bandwidth-bound CUDA-core path instead:
at M<=8 the WGMMA M=64 granularity would quantize the grid to a fraction of a wave,
and the layer is pure weight streaming there.

The extension is compiled by invoking nvcc directly and loaded with ctypes, so the
build needs neither ninja nor pybind11 headers.
"""
from __future__ import annotations

import ctypes
import hashlib
import os
import subprocess
import sys

import torch
import torch.nn as nn

# ---------------------------------------------------------------------------
# tunables (env-overridable so the same source can be swept)
# ---------------------------------------------------------------------------
BN, BK = 256, 64
# BM = 64 * warpgroups; each warpgroup owns 64 rows of the accumulator tile.  The
# BM*BN accumulator floats have to leave room in the 64K-register file for addressing,
# which caps BM at 192 (3 warpgroups, 168 regs/thread, no spills).  Bytes streamed
# L2->smem for a grouped GEMM are 2*N*K*(P/BN + sum_e ceil(m_e/BM)), minimised at
# BM ~ sqrt(BM*BN) ~ 222, so 192x256 is the best tile the register file allows.
BM = int(os.environ.get("MOE_BM", "192"))
_NWG = BM // 64
_STAGES = int(os.environ.get("MOE_STAGES", "4" if BM <= 128 else "3"))
_PRODIST = int(os.environ.get("MOE_PRODIST", "2"))          # producer run-ahead
_PREFETCH = os.environ.get("MOE_PREFETCH", ".L2::256B")      # "", ".L2::128B", ".L2::256B"
_CTAS_PER_SM = int(os.environ.get("MOE_CTAS_PER_SM", "1"))
_SMALL_T = int(os.environ.get("MOE_SMALL_T", "8"))
# CTAs per SM for the decode path: pure weight streaming, so it only wants enough
# resident warps to keep DRAM busy.  Flat from 4 to 12 CTAs/SM (0.275-0.291ms at T=1,
# ~89% of the 453MB/1.85TB/s floor); 5 measured marginally best.  0 = absolute grid.
_SMALL_CTAS = int(os.environ.get("MOE_SMALL_CTAS", "5"))
_SMALL_GRID = int(os.environ.get("MOE_SMALL_GRID", "0"))


def _gen_wgmma(n: int) -> str:
    nr = n // 2
    regs = ",".join(f"%{i}" for i in range(nr))
    outs = ", ".join(f'"+f"(d[{i}])' for i in range(nr))
    return f"""
__device__ __forceinline__ void wgmma_n{n}(uint64_t da, uint64_t db, float* d) {{
  asm volatile(
    "wgmma.mma_async.sync.aligned.m64n{n}k16.f32.bf16.bf16 "
    "{{{regs}}}, %{nr}, %{nr + 1}, 1, 1, 1, 0, 0;\\n"
    : {outs}
    : "l"(da), "l"(db));
}}
"""


_CU = r"""
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cstdint>

typedef __nv_bfloat16 bf16;

#define BM @BM@
#define BN 256
#define BK 64
#define NSTG @STAGES@
#define PDIST @PDIST@
#define DOUT (NSTG - PDIST)
#define NWG @NWG@                   /* warpgroups: each owns 64 of the BM rows */
#define NTHR (128 * NWG)

#define ASTG (BM * BK * 2)
#define BSTG (BN * BK * 2)          /* 32768 */
#define STGSZ (ASTG + BSTG)
#define SMEM_MAIN (NSTG * STGSZ + 1024)

/* smem fill: 8 threads cover one 128B row, so a pass moves NTHR/8 rows.  B always
   uses the first 256 threads (BN=256 rows = 8 passes of 32); A spreads over all. */
#define AROW (NTHR / 8)
#define ASTEP (AROW * 128)
#define APASS (BM / AROW)
#define SROW1 (NTHR / 16)           /* epilogue store: 16 threads per 128-elem row */
#define SROW2 (NTHR / 32)           /* epilogue store: 32 threads per 256-elem row */
#if NTHR > 256
#define BTHREAD (tid < 256)
#else
#define BTHREAD 1
#endif

/* ---------------------------------------------------------------- helpers */
__device__ __forceinline__ uint32_t smem_addr(const void* p) {
  return static_cast<uint32_t>(__cvta_generic_to_shared(p));
}

/* GMMA shared-memory descriptor for a K-major tile, 64 bf16 (128B) per row,
   128B-swizzled: leading-byte-offset 1, stride-byte-offset 64 (16B units). */
__device__ __forceinline__ uint64_t gmma_desc(uint32_t addr) {
  uint64_t d = (uint64_t)((addr >> 4) & 0x3FFFu);
  d |= ((uint64_t)1u) << 16;
  d |= ((uint64_t)64u) << 32;
  d |= ((uint64_t)1u) << 62;   /* layout_type = B128 */
  return d;
}

/* byte offset of 16B chunk (r,c) inside a 128B-swizzled K-major tile */
__device__ __forceinline__ uint32_t swz(int r, int c) {
  return (uint32_t)(r * 128 + ((c ^ (r & 7)) << 4));
}

__device__ __forceinline__ void cp16(uint32_t dst, const void* src, int nbytes) {
  asm volatile("cp.async.cg.shared.global@PREFETCH@ [%0], [%1], 16, %2;\n"
               :: "r"(dst), "l"(src), "r"(nbytes) : "memory");
}
__device__ __forceinline__ void cp_commit() {
  asm volatile("cp.async.commit_group;\n" ::: "memory");
}
template <int N> __device__ __forceinline__ void cp_wait() {
  asm volatile("cp.async.wait_group %0;\n" :: "n"(N) : "memory");
}
__device__ __forceinline__ void wg_fence() {
  asm volatile("wgmma.fence.sync.aligned;\n" ::: "memory");
}
__device__ __forceinline__ void wg_commit() {
  asm volatile("wgmma.commit_group.sync.aligned;\n" ::: "memory");
}
template <int N> __device__ __forceinline__ void wg_wait() {
  asm volatile("wgmma.wait_group.sync.aligned %0;\n" :: "n"(N) : "memory");
}
__device__ __forceinline__ void async_proxy_fence() {
  asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__ float silu(float v) { return v / (1.0f + __expf(-v)); }

@WGMMA@

/* ------------------------------------------------------------- routing */
__global__ void k_hist(const int64_t* __restrict__ ids, int* __restrict__ counts, int n) {
  int i = blockIdx.x * 256 + threadIdx.x;
  if (i < n) atomicAdd(&counts[(int)ids[i]], 1);
}

/* single CTA: exclusive scan of expert loads + (expert, m-tile) work list */
__global__ __launch_bounds__(1024) void k_scan(const int* __restrict__ counts,
                                              int* __restrict__ offsets,
                                              int* __restrict__ cursor,
                                              int* __restrict__ tile_e,
                                              int* __restrict__ tile_p0,
                                              int* __restrict__ tile_mv,
                                              int* __restrict__ nmt,
                                              int T, int E, int NS, int bm) {
  __shared__ int sc[1024], sm[1024];
  int tid = threadIdx.x;
  int NE = E + NS;
  int c = (tid < NE) ? ((tid < E) ? counts[tid] : T) : 0;
  int mt = (c + bm - 1) / bm;
  sc[tid] = c;
  sm[tid] = mt;
  __syncthreads();
  for (int off = 1; off < 1024; off <<= 1) {
    int a = (tid >= off) ? sc[tid - off] : 0;
    int b = (tid >= off) ? sm[tid - off] : 0;
    __syncthreads();
    sc[tid] += a;
    sm[tid] += b;
    __syncthreads();
  }
  int ec = sc[tid] - c;
  int em = sm[tid] - mt;
  if (tid == 1023) nmt[0] = sm[1023];
  if (tid < NE) {
    offsets[tid] = ec;
    cursor[tid] = ec;
    for (int k = 0; k < mt; ++k) {
      int idx = em + k;
      tile_e[idx] = tid;
      tile_p0[idx] = ec + k * bm;
      tile_mv[idx] = min(bm, c - k * bm);
    }
  }
}

__global__ void k_scatter(const int64_t* __restrict__ ids, const bf16* __restrict__ wts,
                          const int* __restrict__ offsets, int* __restrict__ cursor,
                          int* __restrict__ perm_tok, float* __restrict__ perm_w,
                          int* __restrict__ pos_of, int T, int topk, int E, int NS) {
  int i = blockIdx.x * 256 + threadIdx.x;
  int nr = T * topk;
  int NJ = topk + NS;
  if (i < nr) {
    int t = i / topk, j = i - t * topk;
    int e = (int)ids[i];
    int pos = atomicAdd(&cursor[e], 1);
    perm_tok[pos] = t;
    perm_w[pos] = __bfloat162float(wts[i]);
    pos_of[t * NJ + j] = pos;
  } else if (i < nr + T * NS) {
    int r = i - nr;
    int s = r / T, t = r - s * T;
    int pos = offsets[E + s] + t;
    perm_tok[pos] = t;
    perm_w[pos] = 1.0f;
    pos_of[t * NJ + topk + s] = pos;
  }
}

/* ------------------------------------------------------------- gemm1
   C[64 x 256] = A[64 x K] @ B[256 x K]^T,  K = H
   B rows   0..127 : w1[e][i0 + r]        (gate)
   B rows 128..255 : w1[e][I + i0 + r]    (up)
   epilogue: h[p, i0 + col] = silu(acc[i]) * acc[i + 64]                       */
__global__ __launch_bounds__(NTHR) void k_gemm1(
    const bf16* __restrict__ x, const bf16* __restrict__ w1r, const bf16* __restrict__ w1s,
    bf16* __restrict__ hb, const int* __restrict__ ptok, const int* __restrict__ tile_e,
    const int* __restrict__ tile_p0, const int* __restrict__ tile_mv,
    const int* __restrict__ nmt, int H, int I, int E) {
  extern __shared__ __align__(16) char smem_raw[];
  const uint32_t sal = (smem_addr(smem_raw) + 1023u) & ~1023u;
  char* sgen = smem_raw + (sal - smem_addr(smem_raw));

  const int tid = threadIdx.x;
  const int wg = tid >> 7, warp = (tid >> 5) & 3, lane = tid & 31;
  const int rr = tid >> 3, cc = tid & 7;      /* rr in [0,32): 32 rows per pass */
  const uint32_t aoff = swz(rr, cc);

  const int NT = I >> 7;                 /* n-tiles: 128 gate cols each */
  const int NK = H / BK;
  const int total = (*nmt) * NT;

  for (int tile = blockIdx.x; tile < total; tile += gridDim.x) {
    const int mt = tile / NT, nt = tile - mt * NT;
    const int e = tile_e[mt], p0 = tile_p0[mt], mv = tile_mv[mt];
    const int i0 = nt << 7;
    const bf16* wb = (e < E) ? (w1r + (size_t)e * 2 * (size_t)I * H)
                             : (w1s + (size_t)(e - E) * 2 * (size_t)I * H);

    const bf16* xp[APASS];
    int av[APASS];
#pragma unroll
    for (int i = 0; i < APASS; ++i) {
      int r = rr + AROW * i;
      av[i] = (r < mv) ? 16 : 0;
      xp[i] = x + (size_t)ptok[p0 + r] * H + cc * 8;
    }
    const bf16* pg = wb + (size_t)(i0 + rr) * H + cc * 8;
    const bf16* pu = wb + (size_t)(I + i0 + rr) * H + cc * 8;

    float d[BN / 2];
#pragma unroll
    for (int i = 0; i < BN / 2; ++i) d[i] = 0.f;

    int st = 0;
#pragma unroll 1
    for (int s = 0; s < PDIST; ++s) {
      int k0 = s * BK;
      uint32_t da = sal + s * STGSZ, db = da + ASTG;
#pragma unroll
      for (int i = 0; i < APASS; ++i) cp16(da + aoff + i * ASTEP, xp[i] + k0, av[i]);
      if (BTHREAD) {
#pragma unroll
        for (int i = 0; i < 4; ++i) cp16(db + aoff + i * 4096, pg + (size_t)(32 * i) * H + k0, 16);
#pragma unroll
        for (int i = 0; i < 4; ++i)
          cp16(db + aoff + 16384 + i * 4096, pu + (size_t)(32 * i) * H + k0, 16);
      }
      cp_commit();
    }

#pragma unroll 1
    for (int kk = 0; kk < NK; ++kk) {
      cp_wait<PDIST - 1>();
      __syncthreads();
      async_proxy_fence();
      const uint32_t da = sal + st * STGSZ, db = da + ASTG;
      uint64_t descA = gmma_desc(da + wg * 8192), descB = gmma_desc(db);
      wg_fence();
#pragma unroll
      for (int j = 0; j < BK / 16; ++j) wgmma_n256(descA + 2 * j, descB + 2 * j, d);
      wg_commit();
      wg_wait<DOUT>();
      __syncthreads();
      int nk = kk + PDIST;
      if (nk < NK) {
        int ns = st + PDIST;
        if (ns >= NSTG) ns -= NSTG;
        int k0 = nk * BK;
        uint32_t na = sal + ns * STGSZ, nb = na + ASTG;
#pragma unroll
        for (int i = 0; i < APASS; ++i) cp16(na + aoff + i * ASTEP, xp[i] + k0, av[i]);
        if (BTHREAD) {
#pragma unroll
          for (int i = 0; i < 4; ++i)
            cp16(nb + aoff + i * 4096, pg + (size_t)(32 * i) * H + k0, 16);
#pragma unroll
          for (int i = 0; i < 4; ++i)
            cp16(nb + aoff + 16384 + i * 4096, pu + (size_t)(32 * i) * H + k0, 16);
        }
      }
      cp_commit();
      if (++st >= NSTG) st = 0;
    }
    wg_wait<0>();

    /* ---- fused SwiGLU epilogue, staged through stage-0 smem ---- */
    bf16* sh = (bf16*)sgen;            /* BM rows x 136 elems (272B) = 34816B */
    __syncthreads();
#pragma unroll
    for (int i = 0; i < 64; ++i) {
      int g = i >> 2, jj = i & 3;
      int row = (wg << 6) + (warp << 4) + ((jj >> 1) << 3) + (lane >> 2);
      int col = (g << 3) + ((lane & 3) << 1) + (jj & 1);
      sh[row * 136 + col] = __float2bfloat16(silu(d[i]) * d[i + 64]);
    }
    __syncthreads();
#pragma unroll 1
    for (int p = 0; p < BM / SROW1; ++p) {
      int row = p * SROW1 + (tid >> 4);
      if (row < mv) {
        uint4 v = *(const uint4*)(sh + row * 136 + ((tid & 15) << 3));
        *(uint4*)(hb + (size_t)(p0 + row) * I + i0 + ((tid & 15) << 3)) = v;
      }
    }
    __syncthreads();
  }
}

/* ------------------------------------------------------------- gemm2
   C[64 x 256] = A[64 x K] @ B[256 x K]^T,  K = I
   A rows: h[p0 + r]      B rows: w2[e][n0 + r]
   epilogue: y[p, n0 + col] = acc * routing_weight[p]                          */
__global__ __launch_bounds__(NTHR) void k_gemm2(
    const bf16* __restrict__ hb, const bf16* __restrict__ w2r, const bf16* __restrict__ w2s,
    bf16* __restrict__ yb, const float* __restrict__ pw, const int* __restrict__ tile_e,
    const int* __restrict__ tile_p0, const int* __restrict__ tile_mv,
    const int* __restrict__ nmt, int H, int I, int E) {
  extern __shared__ __align__(16) char smem_raw[];
  const uint32_t sal = (smem_addr(smem_raw) + 1023u) & ~1023u;
  char* sgen = smem_raw + (sal - smem_addr(smem_raw));

  const int tid = threadIdx.x;
  const int wg = tid >> 7, warp = (tid >> 5) & 3, lane = tid & 31;
  const int rr = tid >> 3, cc = tid & 7;
  const uint32_t aoff = swz(rr, cc);

  const int NT = H / BN;
  const int NK = I / BK;
  const int total = (*nmt) * NT;

  for (int tile = blockIdx.x; tile < total; tile += gridDim.x) {
    const int mt = tile / NT, nt = tile - mt * NT;
    const int e = tile_e[mt], p0 = tile_p0[mt], mv = tile_mv[mt];
    const int n0 = nt * BN;
    const bf16* wb = (e < E) ? (w2r + (size_t)e * (size_t)H * I)
                             : (w2s + (size_t)(e - E) * (size_t)H * I);

    int av[APASS];
#pragma unroll
    for (int i = 0; i < APASS; ++i) av[i] = ((rr + AROW * i) < mv) ? 16 : 0;
    const bf16* ph = hb + (size_t)(p0 + rr) * I + cc * 8;
    const bf16* pn = wb + (size_t)(n0 + rr) * I + cc * 8;

    float d[BN / 2];
#pragma unroll
    for (int i = 0; i < BN / 2; ++i) d[i] = 0.f;

    int st = 0;
#pragma unroll 1
    for (int s = 0; s < PDIST; ++s) {
      int k0 = s * BK;
      uint32_t da = sal + s * STGSZ, db = da + ASTG;
#pragma unroll
      for (int i = 0; i < APASS; ++i)
        cp16(da + aoff + i * ASTEP, ph + (size_t)(AROW * i) * I + k0, av[i]);
      if (BTHREAD) {
#pragma unroll
        for (int i = 0; i < 8; ++i)
          cp16(db + aoff + i * 4096, pn + (size_t)(32 * i) * I + k0, 16);
      }
      cp_commit();
    }

#pragma unroll 1
    for (int kk = 0; kk < NK; ++kk) {
      cp_wait<PDIST - 1>();
      __syncthreads();
      async_proxy_fence();
      const uint32_t da = sal + st * STGSZ, db = da + ASTG;
      uint64_t descA = gmma_desc(da + wg * 8192), descB = gmma_desc(db);
      wg_fence();
#pragma unroll
      for (int j = 0; j < BK / 16; ++j) wgmma_n256(descA + 2 * j, descB + 2 * j, d);
      wg_commit();
      wg_wait<DOUT>();
      __syncthreads();
      int nk = kk + PDIST;
      if (nk < NK) {
        int ns = st + PDIST;
        if (ns >= NSTG) ns -= NSTG;
        int k0 = nk * BK;
        uint32_t na = sal + ns * STGSZ, nb = na + ASTG;
#pragma unroll
        for (int i = 0; i < APASS; ++i)
          cp16(na + aoff + i * ASTEP, ph + (size_t)(AROW * i) * I + k0, av[i]);
        if (BTHREAD) {
#pragma unroll
          for (int i = 0; i < 8; ++i)
            cp16(nb + aoff + i * 4096, pn + (size_t)(32 * i) * I + k0, 16);
        }
      }
      cp_commit();
      if (++st >= NSTG) st = 0;
    }
    wg_wait<0>();

    bf16* sy = (bf16*)sgen;            /* BM rows x 264 elems (528B) = 67584B */
    __shared__ float swt[BM];
    if (tid < BM) swt[tid] = (tid < mv) ? pw[p0 + tid] : 0.f;
    __syncthreads();
#pragma unroll
    for (int i = 0; i < BN / 2; ++i) {
      int g = i >> 2, jj = i & 3;
      int row = (wg << 6) + (warp << 4) + ((jj >> 1) << 3) + (lane >> 2);
      int col = (g << 3) + ((lane & 3) << 1) + (jj & 1);
      sy[row * 264 + col] = __float2bfloat16(d[i] * swt[row]);
    }
    __syncthreads();
#pragma unroll 1
    for (int p = 0; p < BM / SROW2; ++p) {
      int row = p * SROW2 + (tid >> 5);
      if (row < mv) {
        uint4 v = *(const uint4*)(sy + row * 264 + (lane << 3));
        *(uint4*)(yb + (size_t)(p0 + row) * H + n0 + (lane << 3)) = v;
      }
    }
    __syncthreads();
  }
}

/* ------------------------------------------------------------- reduce */
__global__ __launch_bounds__(256) void k_reduce(const bf16* __restrict__ yb,
                                               const int* __restrict__ pos_of,
                                               bf16* __restrict__ out, int H, int NJ) {
  __shared__ int sp[64];
  const int t = blockIdx.y;
  if (threadIdx.x < NJ) sp[threadIdx.x] = pos_of[t * NJ + threadIdx.x];
  __syncthreads();
  int c = (blockIdx.x * 256 + threadIdx.x) * 8;
  if (c >= H) return;
  float a[8];
#pragma unroll
  for (int i = 0; i < 8; ++i) a[i] = 0.f;
  for (int j = 0; j < NJ; ++j) {
    const bf16* p = yb + (size_t)sp[j] * H + c;
    uint4 v = *(const uint4*)p;
    const bf16* q = (const bf16*)&v;
#pragma unroll
    for (int i = 0; i < 8; ++i) a[i] += __bfloat162float(q[i]);
  }
  bf16 o[8];
#pragma unroll
  for (int i = 0; i < 8; ++i) o[i] = __float2bfloat16(a[i]);
  *(uint4*)(out + (size_t)t * H + c) = *(const uint4*)o;
}

/* ------------------------------------------------------- small-T path */
__device__ __forceinline__ void dot8(float& acc, const uint4& a, const uint4& b) {
  const __nv_bfloat162* pa = (const __nv_bfloat162*)&a;
  const __nv_bfloat162* pb = (const __nv_bfloat162*)&b;
#pragma unroll
  for (int i = 0; i < 4; ++i) {
    float2 fa = __bfloat1622float2(pa[i]);
    float2 fb = __bfloat1622float2(pb[i]);
    acc = fmaf(fa.x, fb.x, acc);
    acc = fmaf(fa.y, fb.y, acc);
  }
}
__device__ __forceinline__ float wsum(float v) {
#pragma unroll
  for (int o = 16; o; o >>= 1) v += __shfl_down_sync(0xffffffffu, v, o);
  return v;
}

template <int MMAX>
__global__ __launch_bounds__(256) void k_small1(
    const bf16* __restrict__ x, const bf16* __restrict__ w1r, const bf16* __restrict__ w1s,
    bf16* __restrict__ hb, const int* __restrict__ ptok, const int* __restrict__ tile_e,
    const int* __restrict__ tile_p0, const int* __restrict__ tile_mv,
    const int* __restrict__ nmt, int H, int I, int E) {
  const int nib = I >> 3;
  const int total = (*nmt) * nib;
  const int lane = threadIdx.x & 31, w = threadIdx.x >> 5;
  for (int item = blockIdx.x; item < total; item += gridDim.x) {
    int mt = item / nib, ib = item - mt * nib;
    int e = tile_e[mt], p0 = tile_p0[mt], mv = tile_mv[mt];
    const bf16* wb = (e < E) ? (w1r + (size_t)e * 2 * (size_t)I * H)
                             : (w1s + (size_t)(e - E) * 2 * (size_t)I * H);
    int i = ib * 8 + w;
    const bf16* pg = wb + (size_t)i * H;
    const bf16* pu = wb + (size_t)(I + i) * H;
    const bf16* xr[MMAX];
    float ag[MMAX], au[MMAX];
#pragma unroll
    for (int m = 0; m < MMAX; ++m) {
      ag[m] = 0.f;
      au[m] = 0.f;
      xr[m] = x + (size_t)ptok[p0 + (m < mv ? m : 0)] * H;
    }
    for (int k = lane * 8; k < H; k += 256) {
      uint4 vg = *(const uint4*)(pg + k);
      uint4 vu = *(const uint4*)(pu + k);
#pragma unroll
      for (int m = 0; m < MMAX; ++m)
        if (m < mv) {
          uint4 vx = *(const uint4*)(xr[m] + k);
          dot8(ag[m], vg, vx);
          dot8(au[m], vu, vx);
        }
    }
#pragma unroll
    for (int m = 0; m < MMAX; ++m)
      if (m < mv) {
        float g = wsum(ag[m]), u = wsum(au[m]);
        if (lane == 0) hb[(size_t)(p0 + m) * I + i] = __float2bfloat16(silu(g) * u);
      }
  }
}

template <int MMAX>
__global__ __launch_bounds__(256) void k_small2(
    const bf16* __restrict__ hb, const bf16* __restrict__ w2r, const bf16* __restrict__ w2s,
    bf16* __restrict__ yb, const float* __restrict__ pw, const int* __restrict__ tile_e,
    const int* __restrict__ tile_p0, const int* __restrict__ tile_mv,
    const int* __restrict__ nmt, int H, int I, int E) {
  const int nib = H >> 3;
  const int total = (*nmt) * nib;
  const int lane = threadIdx.x & 31, w = threadIdx.x >> 5;
  for (int item = blockIdx.x; item < total; item += gridDim.x) {
    int mt = item / nib, ib = item - mt * nib;
    int e = tile_e[mt], p0 = tile_p0[mt], mv = tile_mv[mt];
    const bf16* wb = (e < E) ? (w2r + (size_t)e * (size_t)H * I)
                             : (w2s + (size_t)(e - E) * (size_t)H * I);
    int i = ib * 8 + w;
    const bf16* pn = wb + (size_t)i * I;
    const bf16* hr[MMAX];
    float ac[MMAX];
#pragma unroll
    for (int m = 0; m < MMAX; ++m) {
      ac[m] = 0.f;
      hr[m] = hb + (size_t)(p0 + (m < mv ? m : 0)) * I;
    }
    for (int k = lane * 8; k < I; k += 256) {
      uint4 vn = *(const uint4*)(pn + k);
#pragma unroll
      for (int m = 0; m < MMAX; ++m)
        if (m < mv) {
          uint4 vh = *(const uint4*)(hr[m] + k);
          dot8(ac[m], vn, vh);
        }
    }
#pragma unroll
    for (int m = 0; m < MMAX; ++m)
      if (m < mv) {
        float v = wsum(ac[m]);
        if (lane == 0) yb[(size_t)(p0 + m) * H + i] = __float2bfloat16(v * pw[p0 + m]);
      }
  }
}

/* ------------------------------------------------------------- driver */
static int g_init = 0;
static void init_once() {
  if (g_init) return;
  cudaFuncSetAttribute(k_gemm1, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_MAIN);
  cudaFuncSetAttribute(k_gemm2, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_MAIN);
  g_init = 1;
}

extern "C" int moe_smem_main() { return SMEM_MAIN; }

extern "C" void moe_forward(
    const void* x, const void* ids, const void* wts, const void* w1r, const void* w2r,
    const void* w1s, const void* w2s, void* out, void* hb, void* yb, int* counts,
    int* offsets, int* cursor, int* perm_tok, float* perm_w, int* pos_of, int* tile_e,
    int* tile_p0, int* tile_mv, int* nmt, int T, int E, int topk, int NS, int H, int I,
    int grid_main, int grid_small, int small_path, int mask, void* stream) {
  init_once();
  cudaStream_t s = (cudaStream_t)stream;
  const int NE = E + NS;
  const int NJ = topk + NS;

  if (mask & 1) {
  cudaMemsetAsync(counts, 0, NE * sizeof(int), s);
  k_hist<<<(T * topk + 255) / 256, 256, 0, s>>>((const int64_t*)ids, counts, T * topk);
  k_scan<<<1, 1024, 0, s>>>(counts, offsets, cursor, tile_e, tile_p0, tile_mv, nmt, T, E, NS,
                            small_path ? T : BM);
  k_scatter<<<(T * NJ + 255) / 256, 256, 0, s>>>((const int64_t*)ids, (const bf16*)wts, offsets,
                                                 cursor, perm_tok, perm_w, pos_of, T, topk, E, NS);
  }

  if (small_path) {
    if (T == 1) {
      if (mask & 2)
      k_small1<1><<<grid_small, 256, 0, s>>>((const bf16*)x, (const bf16*)w1r, (const bf16*)w1s,
                                             (bf16*)hb, perm_tok, tile_e, tile_p0, tile_mv, nmt,
                                             H, I, E);
      if (mask & 4)
      k_small2<1><<<grid_small, 256, 0, s>>>((const bf16*)hb, (const bf16*)w2r, (const bf16*)w2s,
                                             (bf16*)yb, perm_w, tile_e, tile_p0, tile_mv, nmt,
                                             H, I, E);
    } else {
      if (mask & 2)
      k_small1<8><<<grid_small, 256, 0, s>>>((const bf16*)x, (const bf16*)w1r, (const bf16*)w1s,
                                             (bf16*)hb, perm_tok, tile_e, tile_p0, tile_mv, nmt,
                                             H, I, E);
      if (mask & 4)
      k_small2<8><<<grid_small, 256, 0, s>>>((const bf16*)hb, (const bf16*)w2r, (const bf16*)w2s,
                                             (bf16*)yb, perm_w, tile_e, tile_p0, tile_mv, nmt,
                                             H, I, E);
    }
  } else {
    if (mask & 2)
    k_gemm1<<<grid_main, NTHR, SMEM_MAIN, s>>>((const bf16*)x, (const bf16*)w1r, (const bf16*)w1s,
                                               (bf16*)hb, perm_tok, tile_e, tile_p0, tile_mv, nmt,
                                               H, I, E);
    if (mask & 4)
    k_gemm2<<<grid_main, NTHR, SMEM_MAIN, s>>>((const bf16*)hb, (const bf16*)w2r, (const bf16*)w2s,
                                               (bf16*)yb, perm_w, tile_e, tile_p0, tile_mv, nmt,
                                               H, I, E);
  }
  if (mask & 8) {
  dim3 rg((H + 2047) / 2048, T);
  k_reduce<<<rg, 256, 0, s>>>((const bf16*)yb, pos_of, (bf16*)out, H, NJ);
  }
}

extern "C" const char* moe_last_error() { return cudaGetErrorString(cudaGetLastError()); }
"""


def _cuda_source() -> str:
    return (
        _CU.replace("@WGMMA@", _gen_wgmma(BN))
        .replace("@BM@", str(BM))
        .replace("@NWG@", str(_NWG))
        .replace("@STAGES@", str(_STAGES))
        .replace("@PDIST@", str(_PRODIST))
        .replace("@PREFETCH@", _PREFETCH)
    )


_LIB = None


def _cache_dir() -> str:
    for env in ("TORCH_EXTENSIONS_DIR", "TMPDIR", "TMP"):
        v = os.environ.get(env)
        if v:
            return os.path.join(v, "moe_glm52")
    return os.path.join(os.path.expanduser("~"), ".cache", "moe_glm52")


def _build():
    """Compile the kernels with nvcc and load them through ctypes.

    Avoids torch.utils.cpp_extension.load_inline (ninja + pybind11 headers) by
    exporting plain `extern "C"` launchers instead of a python module.
    """
    global _LIB
    if _LIB is not None:
        return _LIB
    src = _cuda_source()
    key = hashlib.sha1(src.encode()).hexdigest()[:16]
    d = os.path.join(_cache_dir(), key)
    so = os.path.join(d, "libmoe.so")
    if not os.path.exists(so):
        os.makedirs(d, exist_ok=True)
        cu = os.path.join(d, "moe.cu")
        with open(cu, "w") as f:
            f.write(src)
        nvcc = os.environ.get("REAL_NVCC") or "nvcc"
        tmp = so + f".{os.getpid()}.tmp"
        cmd = [
            nvcc, "-O3", "-std=c++17", "-gencode=arch=compute_90a,code=sm_90a",
            "--shared", "-Xcompiler", "-fPIC", "--use_fast_math", "-lineinfo",
            "--cudart", "shared", cu, "-o", tmp,
        ]
        r = subprocess.run(cmd, capture_output=True, text=True)
        if r.returncode != 0:
            sys.stderr.write(r.stdout[-4000:] + "\n" + r.stderr[-8000:] + "\n")
            raise RuntimeError("nvcc build failed for GLM-5.2 fused MoE kernels")
        os.replace(tmp, so)
    lib = ctypes.CDLL(so)
    lib.moe_forward.restype = None
    lib.moe_forward.argtypes = [ctypes.c_void_p] * 20 + [ctypes.c_int] * 10 + [ctypes.c_void_p]
    lib.moe_smem_main.restype = ctypes.c_int
    lib.moe_last_error.restype = ctypes.c_char_p
    _LIB = lib
    return lib


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))
        # Same N(0, 0.02) init as the reference, but drawn on the GPU when one is
        # available: 6.5G bf16 samples through the single-threaded CPU generator
        # would otherwise dominate construction.
        with torch.no_grad():
            gpu = torch.cuda.is_available()
            for p in self.parameters():
                if gpu and p.device.type == "cpu":
                    p.copy_(torch.empty(p.shape, dtype=p.dtype, device="cuda").normal_(0.0, 0.02))
                else:
                    nn.init.normal_(p, std=0.02)
        self._ws: dict = {}
        self._mask = -1          # stage bitmask; only touched by the profiling harness
        self._supported = (
            H % BN == 0 and I % (BN // 2) == 0 and H % BK == 0 and I % BK == 0
            and E + n_shared <= 1024 and top_k + n_shared <= 64
        )

    # ------------------------------------------------------------------ torch
    def _ref(self, x, expert_ids, expert_weights):
        """Fallback for geometries the kernel does not cover."""
        T, H = x.shape
        out = torch.zeros(T, H, device=x.device, dtype=torch.float32)
        xf = x.float()
        for s in range(self.n_shared):
            w1 = self.w1_shared[s].float()
            w2 = self.w2_shared[s].float()
            g = xf @ w1[: self.I].T
            u = xf @ w1[self.I:].T
            out += (torch.nn.functional.silu(g) * u) @ w2.T
        wts = expert_weights.float()
        for e in range(self.E):
            mask = expert_ids == e
            if not mask.any():
                continue
            ti, ki = mask.nonzero(as_tuple=True)
            w1 = self.w1_routed[e].float()
            w2 = self.w2_routed[e].float()
            xe = xf[ti]
            g = xe @ w1[: self.I].T
            u = xe @ w1[self.I:].T
            y = (torch.nn.functional.silu(g) * u) @ w2.T
            out.index_add_(0, ti, y * wts[ti, ki].unsqueeze(1))
        return out.to(torch.bfloat16)

    # ------------------------------------------------------------------ cuda
    def _workspace(self, T: int, dev: torch.device):
        ws = self._ws.get(T)
        if ws is not None:
            return ws
        lib = _build()
        H, I, E, NS = self.H, self.I, self.E, self.n_shared
        NJ = self.top_k + NS
        NE = E + NS
        P = T * NJ
        small = T <= _SMALL_T
        bm = T if small else BM
        MT = NE + (P + bm - 1) // bm
        i32 = dict(dtype=torch.int32, device=dev)
        nsm = torch.cuda.get_device_properties(dev).multi_processor_count
        maxtiles = MT * max(H // BN, I // (BN // 2))
        ws = dict(
            counts=torch.zeros(NE + 1, **i32),
            offsets=torch.zeros(NE + 1, **i32),
            cursor=torch.zeros(NE + 1, **i32),
            perm_tok=torch.zeros(P + BM, **i32),
            perm_w=torch.zeros(P + BM, dtype=torch.float32, device=dev),
            pos_of=torch.zeros(T * NJ, **i32),
            tile_e=torch.zeros(MT + 1, **i32),
            tile_p0=torch.zeros(MT + 1, **i32),
            tile_mv=torch.zeros(MT + 1, **i32),
            nmt=torch.zeros(1, **i32),
            hb=torch.empty((P + BM) * I, dtype=torch.bfloat16, device=dev),
            yb=torch.empty((P + BM) * H, dtype=torch.bfloat16, device=dev),
            out=torch.empty(T, H, dtype=torch.bfloat16, device=dev),
            grid_main=min(nsm * _CTAS_PER_SM, maxtiles),
            grid_small=_SMALL_GRID or min(nsm * _SMALL_CTAS, MT * (min(H, I) // 8)),
            small=1 if small else 0,
            lib=lib,
        )
        self._ws[T] = ws
        return ws

    def forward(self, x: torch.Tensor, expert_ids: torch.Tensor, expert_weights: torch.Tensor):
        T = x.shape[0]
        if not self._supported or not x.is_cuda:
            return self._ref(x, expert_ids, expert_weights)
        ws = self._workspace(T, x.device)
        lib = ws["lib"]
        p = ctypes.c_void_p
        lib.moe_forward(
            p(x.data_ptr()), p(expert_ids.data_ptr()), p(expert_weights.data_ptr()),
            p(self.w1_routed.data_ptr()), p(self.w2_routed.data_ptr()),
            p(self.w1_shared.data_ptr()), p(self.w2_shared.data_ptr()),
            p(ws["out"].data_ptr()), p(ws["hb"].data_ptr()), p(ws["yb"].data_ptr()),
            p(ws["counts"].data_ptr()), p(ws["offsets"].data_ptr()), p(ws["cursor"].data_ptr()),
            p(ws["perm_tok"].data_ptr()), p(ws["perm_w"].data_ptr()), p(ws["pos_of"].data_ptr()),
            p(ws["tile_e"].data_ptr()), p(ws["tile_p0"].data_ptr()), p(ws["tile_mv"].data_ptr()),
            p(ws["nmt"].data_ptr()),
            T, self.E, self.top_k, self.n_shared, self.H, self.I,
            ws["grid_main"], ws["grid_small"], ws["small"], self._mask,
            p(torch.cuda.current_stream(x.device).cuda_stream),
        )
        return ws["out"]

20260725_064140_or-opus_anthropic_claude-opus-5_01_glm52_fused_moe