KernelBench cuda · RTX PRO 6000

GLM-5.2 Fused MoE DeepSeek V4.1 Flash

9.46%geomean peak fraction across shapes

manually audited: clean

DeepSeek V4.1 Flash hand-rolled the whole GLM-5.2 MoE layer in inline PTX: mma.sync bf16 tensor cores, ldmatrix, xor-swizzled cp.async, hist/scan/scatter token packing, one code path for every T. Every expert GEMM runs on real weights at full K and N; row padding makes it burn 25-50% more FLOPs than the roofline charges. Error margin 2.5x inside the gate.

harnessdeepseek-claudeagent session2h 54mtotal wall3h 6mcheck4mbenchmark5moutput tokens437,222cost$23.60gpu-lock wait52sgpu-lock held1h 55mregimecompute

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

shape 010.781 ms34.4%1.21 TB/s · 67% of 1.8 TB/s HBM · also 172 TFLOPS (34% of compute)
shape 110.751 ms34.8%1.21 TB/s · 67% of 1.8 TB/s HBM · also 174 TFLOPS (35% of compute)
shape 20.396 ms0.2%32.65 TB/s · 100% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
shape 314.236 ms52.1%261 TFLOPS · 52% of 500 TF bf16 peak · also 0.92 TB/s (51% of HBM)
shape 49.021 ms5.1%1.43 TB/s · 80% of 1.8 TB/s HBM · also 26 TFLOPS (5% of compute)
shape 59.293 ms9.8%1.39 TB/s · 77% of 1.8 TB/s HBM · also 49 TFLOPS (10% of compute)

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

geomean(34.4% · 34.8% · 0.2% · 52.1% · 5.1% · 9.8%) = 9.5%

Kernel source (redacted)
"""GLM-5.2 fused MoE layer in raw CUDA (SM120) — KernelBench-CUDA.

Structure (per problem statement):
  E=256 routed experts, top_k=8, n_shared=1 shared expert, H=4096, I=2048.
  w1_routed (E,2I,H) packed gate|up, w2_routed (E,H,I)
  w1_shared (S,2I,H), w2_shared (S,H,I)

  out = sum_s silu(x W1s^T) . (x W2s^T)  +  sum_k w_k * silu(x W1[e_k]^T) . W2[e_k]^T

Implementation (one CUDA extension; no torch compute ops):
  1. histogram of expert ids + exclusive scan -> BM-aligned padded row offsets
  2. scatter -> sorted_token / sorted_weight, grouped by expert (vLLM-style
     moe_align_block_size).  The shared expert is appended as group index E.
  3. grouped GEMM #1 (mma.sync m16n8k16 bf16, cp.async pipeline, xor-swizzled
     smem): one CTA owns (expert, BM rows, BN columns of I); computes the gate
     and up tiles together and writes h = silu(gate)*up as bf16.
  4. grouped GEMM #2: h @ W2^T -> weighted atomic-add into fp32 out
  5. cast fp32 -> bf16
"""
from __future__ import annotations

import os

import torch
import torch.nn as nn

os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")

# make the venv's ninja visible to torch's cpp_extension bootstrap
import sys as _sys  # noqa: E402

_bin = os.path.dirname(_sys.executable)
if os.path.isdir(_bin) and _bin not in os.environ.get("PATH", ""):
    os.environ["PATH"] = _bin + os.pathsep + os.environ.get("PATH", "")

from torch.utils.cpp_extension import load_inline  # noqa: E402

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

#define DEV __device__ __forceinline__
using bf16 = __nv_bfloat16;

// tile geometry (BK is fixed at 64 halves = 128 byte smem rows)
constexpr int BM1 = 128;   // gate|up rows per CTA
constexpr int BN1 = 128;   // gate|up columns-of-I per CTA
constexpr int ST1 = 2;
constexpr int BM2 = 128;   // down rows per CTA
constexpr int BN2 = 256;   // down columns of H per CTA
constexpr int ST2 = 2;
// The activation tile is re-read once per column tile, so its L2 cost scales
// as 1/BN; the weight tile is re-read once per row tile, so its cost scales as
// 1/BM.  Making gate|up twice as wide halves the dominant A traffic.  That
// costs BN*BK*2 = 16 KB more smem per stage, which the 99 KB budget only
// affords at two stages -- still a real pipeline, because the wait leaves one
// group in flight.  16 warps keep the accumulator count at 128 registers.
constexpr int WM1 = 4, WN1 = 4;   // gate|up warp grid (16 warps = 512 threads)
constexpr int WM2 = 4, WN2 = 4;   // down warp grid (16 warps = 512 threads)

// ---------------------------------------------------------------------------
// pipeline / mma helpers
// ---------------------------------------------------------------------------
DEV uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }

DEV void cp_async16(void* dst, const void* src) {
  asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" ::"r"(smem_u32(dst)), "l"(src));
}
DEV void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N>
DEV void cp_wait() { asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); }

DEV void ldsm4(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, uint32_t addr) {
  asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
               : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
               : "r"(addr));
}

DEV void mma_16816(float* c, const uint32_t* a, uint32_t b0, uint32_t b1) {
  asm volatile(
      "mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
      "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
      : "+f"(c[0]), "+f"(c[1]), "+f"(c[2]), "+f"(c[3])
      : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1));
}

DEV void red_add2(float* p, float a, float b) {
  asm volatile("red.global.add.v2.f32 [%0], {%1,%2};\n" ::"l"(p), "f"(a), "f"(b) : "memory");
}

DEV float silu_f(float x) { return x / (1.0f + __expf(-x)); }

DEV int count_nonempty(const int* counts, int e) { return __ldg(&counts[e]); }

// xor-swizzled byte offset inside a row-major smem tile with 128-byte rows
// (BK = 64 halves -> 8 chunks of 16 bytes).  chunk in [0,8).
template <int BK>
DEV uint32_t swz(int row, int chunk) {
  return (uint32_t)(row * (BK * 2) + ((chunk ^ (row & 7)) << 4));
}

// ---------------------------------------------------------------------------
// prep kernels
// ---------------------------------------------------------------------------
__global__ void k_count(const long long* __restrict__ ids, int n, int E, int* __restrict__ counts) {
  int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i < n) atomicAdd(&counts[(int)ids[i]], 1);
}

// one block: exclusive scan of per-expert counts on BM-padded boundaries.
// counts[E] is forced to T (the shared expert sees every token).
// NOTE: counts_in and counts_out intentionally alias the same buffer; the scan
// snapshot lives in shared memory, so no __restrict__ here.
__global__ void k_scan(const int* counts_in, int E, int T, int BM,
                       int* __restrict__ row_off, int* __restrict__ tile_expert,
                       int* counts_out, int* __restrict__ num_tiles) {
  // smem: two ping-pong rows of E+2 ints.
  extern __shared__ int sm[];
  const int n = E + 1;  // experts 0..E-1 plus the shared pseudo-expert at E
  const int i = threadIdx.x;
  if (n > (int)blockDim.x) {
    // serialize for expert counts that do not fit in one block
    if (i == 0) {
      int acc = 0, t = 0;
      for (int e = 0; e <= E; ++e) {
        int v = (e < E) ? __ldg(&counts_in[e]) : T;
        row_off[e] = acc;
        counts_out[e] = v;
        int nt = (v + BM - 1) / BM;
        for (int j = 0; j < nt; ++j) tile_expert[t++] = e;
        acc += nt * BM;
      }
      row_off[E + 1] = acc;
      *num_tiles = t;
    }
    return;
  }
  int* s = sm;
  int* d = sm + (E + 2);
  int tc = 0;
  if (i < n) {
    int c = (i < E) ? __ldg(&counts_in[i]) : T;
    counts_out[i] = c;
    tc = (c + BM - 1) / BM;   // tiles this group occupies
    s[i] = tc;
  }
  __syncthreads();
  // Hillis-Steele inclusive scan of the tile counts
  for (int off = 1; off < n; off <<= 1) {
    int v = (i < n) ? s[i] : 0;
    int p = (i < n && i >= off) ? s[i - off] : 0;
    if (i < n) d[i] = v + p;
    __syncthreads();
    int* tmp = s; s = d; d = tmp;
    __syncthreads();
  }
  if (i < n) {
    const int base = s[i] - tc;          // exclusive tile prefix
    row_off[i] = base * BM;
    if (i == E) {
      row_off[E + 1] = s[i] * BM;
      *num_tiles = s[i];
    }
#pragma unroll 4
    for (int j = 0; j < tc; ++j) tile_expert[base + j] = i;
  }
}

// group the (token,k) assignment list by expert
__global__ void k_scatter(const long long* __restrict__ ids, const bf16* __restrict__ wts,
                          int n_routed, int top_k, const int* __restrict__ row_off,
                          int* __restrict__ cursor, int* __restrict__ sorted_token,
                          float* __restrict__ sorted_weight) {
  int i = blockIdx.x * blockDim.x + threadIdx.x;
  if (i < n_routed) {
    int e = (int)ids[i];
    int pos = atomicAdd(&cursor[e], 1);
    int row = row_off[e] + pos;
    sorted_token[row] = i / top_k;
    sorted_weight[row] = __bfloat162float(wts[i]);
  }
}

// shared expert rows: every token, weight 1
__global__ void k_shared_rows(int T, const int* __restrict__ row_off, int E,
                              int* __restrict__ sorted_token, float* __restrict__ sorted_weight) {
  int t = blockIdx.x * blockDim.x + threadIdx.x;
  if (t < T) {
    int row = row_off[E] + t;
    sorted_token[row] = t;
    sorted_weight[row] = 1.0f;
  }
}

// ---------------------------------------------------------------------------
// grouped GEMM 1: gate|up
// ---------------------------------------------------------------------------
template <int BM, int BN, int STAGES, int WM, int WN>
__global__ void __launch_bounds__(WM* WN * 32, 1) k_gate_up(
    const bf16* __restrict__ x, int H, const int* __restrict__ tile_expert,
    const int* __restrict__ row_off, const int* __restrict__ counts,
    const int* __restrict__ ntiles, const int* __restrict__ sorted_token,
    const bf16* __restrict__ w1r, const bf16* __restrict__ w1s, int E, int I,
    bf16* __restrict__ hbuf) {
  constexpr int BK = 64;
  constexpr int NT = WM * WN * 32;
  constexpr int AU = BM * 8;
  constexpr int BU = BN * 8;
  extern __shared__ char smem[];
  bf16* sA = (bf16*)smem;                        // STAGES * BM * BK
  bf16* sBg = sA + (size_t)STAGES * BM * BK;     // STAGES * BN * BK
  bf16* sBu = sBg + (size_t)STAGES * BN * BK;    // STAGES * BN * BK

  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  const int wm = warp % WM;
  const int wn = warp / WM;
  constexpr int WTM = BM / WM;  // warp tile rows
  constexpr int WTN = BN / WN;  // warp tile cols

  const int pid_n = blockIdx.x;
  const int pid_m = blockIdx.y;
  if (pid_m >= *ntiles) return;

  const int e = __ldg(&tile_expert[pid_m]);
  const int row0 = pid_m * BM;
  const int cnt = count_nonempty(counts, e);
  const int mbase = row0 - __ldg(&row_off[e]);
  const bf16* W1 = (e < E) ? (w1r + (size_t)e * 2 * I * H) : w1s;
  const bf16* Wg = W1;
  const bf16* Wu = W1 + (size_t)I * H;
  const int n0 = pid_n * BN;
  const int NK = H / BK;

  // The gather index and the smem destination of every A chunk are invariant
  // across the K loop, so resolve them once instead of re-loading the token
  // inside the pipeline (that load otherwise sits in the issue path of every
  // cp.async).
  constexpr int AIT = (AU + NT - 1) / NT;
  const bf16* abase[AIT];
  int asox[AIT];
  bool arun[AIT], aok[AIT];
#pragma unroll
  for (int it = 0; it < AIT; ++it) {
    int u = tid + it * NT;
    bool run = u < AU;
    bool valid = run && ((mbase + (u >> 3)) < cnt);
    int tok = valid ? __ldg(&sorted_token[row0 + (u >> 3)]) : 0;
    abase[it] = x + (size_t)tok * H + (size_t)((u & 7) * 8);
    asox[it] = run ? swz<BK>(u >> 3, u & 7) / 2 : 0;
    arun[it] = run;
    aok[it] = valid;
  }
  constexpr int BIT = (2 * BU + NT - 1) / NT;
  const char* bbase[BIT];
  int bsox[BIT];
  bool bok[BIT];
#pragma unroll
  for (int it = 0; it < BIT; ++it) {
    int u = tid + it * NT;
    bool valid = u < 2 * BU;
    int half = valid && (u >= BU);
    int v = half ? (u - BU) : u;
    bbase[it] = (const char*)(half ? Wu : Wg) + (size_t)(n0 + (v >> 3)) * H * 2 +
                (size_t)((v & 7) * 8) * 2;
    bsox[it] = valid ? swz<BK>(v >> 3, v & 7) / 2 : 0;
    bok[it] = valid;
  }

  auto load_a = [&](int stage, int k0) {
    bf16* sb = sA + (size_t)stage * BM * BK;
#pragma unroll
    for (int it = 0; it < AIT; ++it) {
      if (!arun[it]) continue;
      bf16* sp = sb + asox[it];
      if (aok[it])
        cp_async16(sp, (const char*)(abase[it] + k0));
      else
        *reinterpret_cast<uint4*>(sp) = make_uint4(0, 0, 0, 0);
    }
  };
  auto load_b = [&](int stage, int k0) {
#pragma unroll
    for (int it = 0; it < BIT; ++it) {
      if (!bok[it]) continue;
      int u = tid + it * NT;
      int half = u >= BU;
      bf16* sp = (half ? sBu : sBg) + (size_t)stage * BN * BK + bsox[it];
      cp_async16(sp, bbase[it] + (size_t)k0 * 2);
    }
  };

  float accg[WTM / 16][WTN / 8][4];
  float accu[WTM / 16][WTN / 8][4];
#pragma unroll
  for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
    for (int ni = 0; ni < WTN / 8; ++ni)
#pragma unroll
      for (int q = 0; q < 4; ++q) {
        accg[mi][ni][q] = 0.f;
        accu[mi][ni][q] = 0.f;
      }

#pragma unroll 1
  for (int s = 0; s < STAGES - 1; ++s) {
    if (s < NK) {
      load_a(s, s * BK);
      load_b(s, s * BK);
    }
    cp_commit();
  }

#pragma unroll 1
  for (int ks = 0; ks < NK; ++ks) {
    int lk = ks + STAGES - 1;
    if (lk < NK) {
      load_a(lk % STAGES, lk * BK);
      load_b(lk % STAGES, lk * BK);
    }
    cp_commit();
    cp_wait<1>();
    __syncthreads();

    const int st = ks % STAGES;
    const bf16* sAb = sA + (size_t)st * BM * BK;
    const bf16* sBgb = sBg + (size_t)st * BN * BK;
    const bf16* sBub = sBu + (size_t)st * BN * BK;

#pragma unroll
    for (int kk = 0; kk < BK / 16; ++kk) {
      const int c0 = 2 * kk;
      uint32_t aA[WTM / 16][4];
#pragma unroll
      for (int mi = 0; mi < WTM / 16; ++mi) {
        int row = wm * WTM + mi * 16 + ((lane >> 3) & 1) * 8 + (lane & 7);
        int ch = c0 + (lane >> 4);
        ldsm4(aA[mi][0], aA[mi][1], aA[mi][2], aA[mi][3], smem_u32(sAb) + swz<BK>(row, ch));
      }
      uint32_t bg[WTN / 8][2], bu[WTN / 8][2];
#pragma unroll
      for (int h2 = 0; h2 < (WTN / 8) / 2; ++h2) {
        int row = wn * WTN + h2 * 16 + ((lane >> 4) & 1) * 8 + (lane & 7);
        int ch = c0 + ((lane >> 3) & 1);
        ldsm4(bg[h2 * 2][0], bg[h2 * 2][1], bg[h2 * 2 + 1][0], bg[h2 * 2 + 1][1],
              smem_u32(sBgb) + swz<BK>(row, ch));
        ldsm4(bu[h2 * 2][0], bu[h2 * 2][1], bu[h2 * 2 + 1][0], bu[h2 * 2 + 1][1],
              smem_u32(sBub) + swz<BK>(row, ch));
      }
#pragma unroll
      for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
        for (int ni = 0; ni < WTN / 8; ++ni) {
          mma_16816(accg[mi][ni], aA[mi], bg[ni][0], bg[ni][1]);
          mma_16816(accu[mi][ni], aA[mi], bu[ni][0], bu[ni][1]);
        }
    }
    __syncthreads();
  }

  // epilogue: h = silu(gate) * up
#pragma unroll
  for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
    for (int ni = 0; ni < WTN / 8; ++ni) {
      int r = row0 + wm * WTM + mi * 16 + (lane >> 2);
      int col = n0 + wn * WTN + ni * 8 + 2 * (lane & 3);
      float h0 = silu_f(accg[mi][ni][0]) * accu[mi][ni][0];
      float h1 = silu_f(accg[mi][ni][1]) * accu[mi][ni][1];
      float h2 = silu_f(accg[mi][ni][2]) * accu[mi][ni][2];
      float h3 = silu_f(accg[mi][ni][3]) * accu[mi][ni][3];
      __nv_bfloat162 p0 = __floats2bfloat162_rn(h0, h1);
      __nv_bfloat162 p1 = __floats2bfloat162_rn(h2, h3);
      *(uint32_t*)(hbuf + (size_t)r * I + col) = *(uint32_t*)&p0;
      *(uint32_t*)(hbuf + (size_t)(r + 8) * I + col) = *(uint32_t*)&p1;
    }
}

// ---------------------------------------------------------------------------
// grouped GEMM 2: h @ W2^T  -> weighted atomic add into fp32 out
// ---------------------------------------------------------------------------
template <int BM, int BN, int STAGES, int WM, int WN>
__global__ void __launch_bounds__(WM* WN * 32, 1) k_down(
    const bf16* __restrict__ hbuf, int I, const int* __restrict__ tile_expert,
    const int* __restrict__ row_off, const int* __restrict__ counts,
    const int* __restrict__ ntiles, const int* __restrict__ sorted_token,
    const float* __restrict__ sorted_weight, const bf16* __restrict__ w2r,
    const bf16* __restrict__ w2s, int E, int H, float* __restrict__ out) {
  constexpr int BK = 64;
  constexpr int NT = WM * WN * 32;
  constexpr int AU = BM * 8;
  constexpr int BU = BN * 8;
  extern __shared__ char smem[];
  bf16* sA = (bf16*)smem;
  bf16* sB = sA + (size_t)STAGES * BM * BK;
  constexpr int WTM = BM / WM;
  constexpr int WTN = BN / WN;

  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  const int wm = warp % WM;
  const int wn = warp / WM;

  const int pid_n = blockIdx.x;
  const int pid_m = blockIdx.y;
  if (pid_m >= *ntiles) return;

  const int e = __ldg(&tile_expert[pid_m]);
  const int row0 = pid_m * BM;
  const int cnt = count_nonempty(counts, e);
  const int mbase = row0 - __ldg(&row_off[e]);
  const int n0 = pid_n * BN;
  const bf16* W2 = (e < E) ? (w2r + (size_t)e * H * I) : w2s;
  const int NK = I / BK;

  auto load_a = [&](int stage, int k0) {
#pragma unroll
    for (int it = 0; it < (AU + NT - 1) / NT; ++it) {
      int u = tid + it * NT;
      if (u < AU) {
        int r = u >> 3;
        int ch = u & 7;
        const char* src =
            (const char*)hbuf + (size_t)(row0 + r) * I * 2 + (size_t)(k0 + ch * 8) * 2;
        bf16* sp = sA + (size_t)stage * BM * BK + swz<BK>(r, ch) / 2;
        cp_async16(sp, src);
      }
    }
  };
  auto load_b = [&](int stage, int k0) {
#pragma unroll
    for (int it = 0; it < (BU + NT - 1) / NT; ++it) {
      int u = tid + it * NT;
      if (u < BU) {
        int r = u >> 3;
        int ch = u & 7;
        const char* src =
            (const char*)W2 + (size_t)(n0 + r) * I * 2 + (size_t)(k0 + ch * 8) * 2;
        bf16* sp = sB + (size_t)stage * BN * BK + swz<BK>(r, ch) / 2;
        cp_async16(sp, src);
      }
    }
  };

  float acc[WTM / 16][WTN / 8][4];
#pragma unroll
  for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
    for (int ni = 0; ni < WTN / 8; ++ni)
#pragma unroll
      for (int q = 0; q < 4; ++q) acc[mi][ni][q] = 0.f;

#pragma unroll 1
  for (int s = 0; s < STAGES - 1; ++s) {
    if (s < NK) {
      load_a(s, s * BK);
      load_b(s, s * BK);
    }
    cp_commit();
  }

#pragma unroll 1
  for (int ks = 0; ks < NK; ++ks) {
    int lk = ks + STAGES - 1;
    if (lk < NK) {
      load_a(lk % STAGES, lk * BK);
      load_b(lk % STAGES, lk * BK);
    }
    cp_commit();
    cp_wait<1>();
    __syncthreads();

    const int st = ks % STAGES;
    const bf16* sAb = sA + (size_t)st * BM * BK;
    const bf16* sBb = sB + (size_t)st * BN * BK;
#pragma unroll
    for (int kk = 0; kk < BK / 16; ++kk) {
      const int c0 = 2 * kk;
      uint32_t aA[WTM / 16][4];
#pragma unroll
      for (int mi = 0; mi < WTM / 16; ++mi) {
        int row = wm * WTM + mi * 16 + ((lane >> 3) & 1) * 8 + (lane & 7);
        int ch = c0 + (lane >> 4);
        ldsm4(aA[mi][0], aA[mi][1], aA[mi][2], aA[mi][3], smem_u32(sAb) + swz<BK>(row, ch));
      }
      uint32_t bb[WTN / 8][2];
#pragma unroll
      for (int h2 = 0; h2 < (WTN / 8) / 2; ++h2) {
        int row = wn * WTN + h2 * 16 + ((lane >> 4) & 1) * 8 + (lane & 7);
        int ch = c0 + ((lane >> 3) & 1);
        ldsm4(bb[h2 * 2][0], bb[h2 * 2][1], bb[h2 * 2 + 1][0], bb[h2 * 2 + 1][1],
              smem_u32(sBb) + swz<BK>(row, ch));
      }
#pragma unroll
      for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
        for (int ni = 0; ni < WTN / 8; ++ni) mma_16816(acc[mi][ni], aA[mi], bb[ni][0], bb[ni][1]);
    }
    __syncthreads();
  }

  // epilogue: out[token] += w * y.  The two rows of each m16 fragment
  // (r and r+8) must be guarded independently: rows past the group's count
  // carry no sorted_token entry, so writing them would scatter atoms of the
  // output all over the device heap.
  const int nvalid = cnt - mbase;
#pragma unroll
  for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
    for (int ni = 0; ni < WTN / 8; ++ni) {
      int col = n0 + wn * WTN + ni * 8 + 2 * (lane & 3);
      int l0 = wm * WTM + mi * 16 + (lane >> 2);
      if (l0 < nvalid) {
        int tok = __ldg(&sorted_token[row0 + l0]);
        float w = __ldg(&sorted_weight[row0 + l0]);
        red_add2(out + (size_t)tok * H + col, w * acc[mi][ni][0], w * acc[mi][ni][1]);
      }
      int l1 = l0 + 8;
      if (l1 < nvalid) {
        int tok = __ldg(&sorted_token[row0 + l1]);
        float w = __ldg(&sorted_weight[row0 + l1]);
        red_add2(out + (size_t)tok * H + col, w * acc[mi][ni][2], w * acc[mi][ni][3]);
      }
    }
}

// ---------------------------------------------------------------------------
// cast
// ---------------------------------------------------------------------------
// n is the element count; each thread converts a pair of floats.
__global__ void k_cast(const float* __restrict__ in, bf16* __restrict__ outp, long long n) {
  long long i = ((long long)blockIdx.x * blockDim.x + threadIdx.x) * 2;
  if (i >= n) return;
  float2 v = *(const float2*)(in + i);
  __nv_bfloat162 b = __floats2bfloat162_rn(v.x, v.y);
  *(uint32_t*)(outp + i) = *(uint32_t*)&b;
}

// ---------------------------------------------------------------------------
// host side
// ---------------------------------------------------------------------------
struct Scratch {
  int* counts = nullptr;         // E+1
  int* cursor = nullptr;         // E
  int* row_off = nullptr;        // E+2
  int* tile_expert = nullptr;    // max_tiles
  int* num_tiles = nullptr;      // 1
  int* [REDACTED credential assignment];   // max_rows
  float* sorted_weight = nullptr;
  bf16* hbuf = nullptr;          // max_rows * I
  float* out32 = nullptr;        // T * H
  long long cap_rows = 0, cap_T = 0, cap_tiles = 0;
};
static Scratch g;

static void ensure_scratch(long long max_rows, long long T, int H, int I, long long max_tiles) {
  if (g.cap_rows >= max_rows && g.cap_T >= T && g.cap_tiles >= max_tiles) return;
  if (g.sorted_token) cudaFree(g.sorted_token);
  if (g.sorted_weight) cudaFree(g.sorted_weight);
  if (g.hbuf) cudaFree(g.hbuf);
  if (g.out32) cudaFree(g.out32);
  if (g.tile_expert) cudaFree(g.tile_expert);
  [REDACTED credential assignment];
  g.sorted_weight = nullptr;
  g.hbuf = nullptr;
  g.out32 = nullptr;
  g.tile_expert = nullptr;
  size_t r = (size_t)max_rows;
  cudaMalloc(&g.sorted_token, r * sizeof(int));
  cudaMalloc(&g.sorted_weight, r * sizeof(float));
  cudaMalloc(&g.hbuf, r * (size_t)I * sizeof(bf16));
  cudaMalloc(&g.out32, (size_t)T * H * sizeof(float));
  cudaMalloc(&g.tile_expert, (size_t)max_tiles * sizeof(int));
  g.cap_rows = max_rows;
  g.cap_T = T;
  g.cap_tiles = max_tiles;
}

static void ensure_static(int E) {
  static int cap_E = 0;
  if (cap_E >= E + 2 && g.counts) return;
  if (g.counts) cudaFree(g.counts);
  if (g.cursor) cudaFree(g.cursor);
  if (g.row_off) cudaFree(g.row_off);
  if (g.num_tiles) cudaFree(g.num_tiles);
  cudaMalloc(&g.counts, (E + 1) * sizeof(int));
  cudaMalloc(&g.cursor, E * sizeof(int));
  cudaMalloc(&g.row_off, (E + 2) * sizeof(int));
  cudaMalloc(&g.num_tiles, sizeof(int));
  cap_E = E + 2;
}

extern "C" void moe_launch(uintptr_t x, uintptr_t ids, uintptr_t wts, uintptr_t w1r,
                           uintptr_t w2r, uintptr_t w1s, uintptr_t w2s, uintptr_t out_bf,
                           long long T, long long E, long long top_k, long long H, long long I,
                           long long stream) {
  cudaStream_t st = (cudaStream_t)stream;

  long long max_rows = T * (top_k + 1) + (E + 1) * (BM1 - 1) + BM1;
  // Tiles per expert are padded to BM1 rows, so the count is bounded by what a
  // single fully packed expert would need, plus one partial tile per routed
  // expert, plus the shared expert's tiles.  A tight bound keeps the mostly
  // empty grid rows out of the small-T shapes.
  long long n_routed_tok = T * top_k;
  long long partial = n_routed_tok < E ? n_routed_tok : E;
  long long max_tiles = (n_routed_tok + BM1 - 1) / BM1 + partial + (T + BM1 - 1) / BM1;
  ensure_static((int)E);
  ensure_scratch(max_rows, T, (int)H, (int)I, max_tiles);

  int n_routed = (int)(T * top_k);
  cudaMemsetAsync(g.counts, 0, (E + 1) * sizeof(int), st);
  cudaMemsetAsync(g.cursor, 0, E * sizeof(int), st);
  cudaMemsetAsync(g.out32, 0, (size_t)T * H * sizeof(float), st);

  k_count<<<(n_routed + 255) / 256, 256, 0, st>>>((const long long*)ids, n_routed, (int)E, g.counts);
  k_scan<<<1, 512, (int)(2 * (E + 2) * sizeof(int)), st>>>(g.counts, (int)E, (int)T, BM1, g.row_off,
                                                       g.tile_expert, g.counts, g.num_tiles);
  k_scatter<<<(n_routed + 255) / 256, 256, 0, st>>>((const long long*)ids, (const bf16*)wts,
                                                    n_routed, (int)top_k, g.row_off, g.cursor,
                                                    g.sorted_token, g.sorted_weight);
  k_shared_rows<<<(int)((T + 255) / 256), 256, 0, st>>>((int)T, g.row_off, (int)E,
                                                        g.sorted_token, g.sorted_weight);


  int sm1 = (ST1 * BM1 * 64 + 2 * ST1 * BN1 * 64) * (int)sizeof(bf16);
  auto k1 = k_gate_up<BM1, BN1, ST1, WM1, WN1>;
  static bool s1 = false;
  if (!s1) {
    cudaFuncSetAttribute(k1, cudaFuncAttributeMaxDynamicSharedMemorySize, sm1);
    s1 = true;
  }
  dim3 gr1((unsigned)(I / BN1), (unsigned)max_tiles);
  k1<<<gr1, WM1 * WN1 * 32, sm1, st>>>((const bf16*)x, (int)H, g.tile_expert, g.row_off, g.counts,
                                       g.num_tiles, g.sorted_token, (const bf16*)w1r,
                                       (const bf16*)w1s, (int)E, (int)I, g.hbuf);

  int sm2 = (ST2 * BM2 * 64 + ST2 * BN2 * 64) * (int)sizeof(bf16);
  auto k2 = k_down<BM2, BN2, ST2, WM2, WN2>;
  static bool s2 = false;
  if (!s2) {
    cudaFuncSetAttribute(k2, cudaFuncAttributeMaxDynamicSharedMemorySize, sm2);
    s2 = true;
  }
  dim3 gr2((unsigned)(H / BN2), (unsigned)max_tiles);
  k2<<<gr2, WM2 * WN2 * 32, sm2, st>>>(g.hbuf, (int)I, g.tile_expert, g.row_off, g.counts,
                                       g.num_tiles, g.sorted_token, g.sorted_weight,
                                       (const bf16*)w2r, (const bf16*)w2s, (int)E, (int)H,
                                       g.out32);

  long long n2 = T * H;
  k_cast<<<(unsigned)((n2 / 2 + 255) / 256), 256, 0, st>>>(g.out32, (bf16*)out_bf, n2);
}

"""

_CPP_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>

extern "C" void moe_launch(uintptr_t x, uintptr_t ids, uintptr_t wts, uintptr_t w1r,
                           uintptr_t w2r, uintptr_t w1s, uintptr_t w2s, uintptr_t out_bf,
                           long long T, long long E, long long top_k, long long H, long long I,
                           long long stream);
torch::Tensor moe_forward(torch::Tensor x, torch::Tensor ids, torch::Tensor wts,
                          torch::Tensor w1r, torch::Tensor w2r, torch::Tensor w1s,
                          torch::Tensor w2s, int64_t top_k) {
  TORCH_CHECK(x.is_cuda() && x.dtype() == torch::kBFloat16);
  TORCH_CHECK(x.is_contiguous() && ids.is_contiguous() && wts.is_contiguous());
  TORCH_CHECK(w1r.is_contiguous() && w2r.is_contiguous());
  TORCH_CHECK(w1s.is_contiguous() && w2s.is_contiguous());
  long long T = x.size(0);
  long long H = x.size(1);
  long long E = w1r.size(0);
  long long I = w2r.size(2);
  auto out = torch::empty({T, H}, x.options());
  cudaStream_t st = at::cuda::getCurrentCUDAStream();
  moe_launch((uintptr_t)x.data_ptr(), (uintptr_t)ids.data_ptr(), (uintptr_t)wts.data_ptr(),
             (uintptr_t)w1r.data_ptr(), (uintptr_t)w2r.data_ptr(), (uintptr_t)w1s.data_ptr(),
             (uintptr_t)w2s.data_ptr(), (uintptr_t)out.data_ptr(), T, E, top_k, H, I,
             (long long)st);
  return out;
}
"""

ext = load_inline(
    name="glm52_moe_cuda",
    cpp_sources=_CPP_SRC,
    cuda_sources=_CUDA_SRC,
    functions=["moe_forward"],
    extra_cuda_cflags=["-O3", "-std=c++17", "-lineinfo"],
    verbose=False,
)


class Model(nn.Module):
    def __init__(self, T, E, top_k, n_shared, H, I):
        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.moe_forward(
            x.contiguous(),
            expert_ids.contiguous(),
            expert_weights.contiguous(),
            self.w1_routed,
            self.w2_routed,
            self.w1_shared,
            self.w2_shared,
            self.top_k,
        )

20260910_115701_deepseek-claude_deepseek-flash_01_glm52_fused_moe