"""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 #include #include #include #include #include #include #include #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(&a); const __nv_bfloat16* bp = reinterpret_cast(&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(x + (int64_t)t * H); uint4* dst = reinterpret_cast(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(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 __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(row); const uint4* x4 = reinterpret_cast(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<<>>( 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(); std::vector 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(), 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(x.data_ptr()); if (top_k > 0) { launch_gemv(xptr, /*x_stride=*/0, reinterpret_cast(w1_routed.data_ptr()), exp_idx.data_ptr(), (int64_t)twoI * H, reinterpret_cast<__nv_bfloat16*>(gate_r.data_ptr()), top_k, twoI, H, stream); silu_mul_kernel<<>>( reinterpret_cast(gate_r.data_ptr()), reinterpret_cast<__nv_bfloat16*>(hid_r.data_ptr()), top_k, I); launch_gemv(reinterpret_cast(hid_r.data_ptr()), /*x_stride=*/I, reinterpret_cast(w2_routed.data_ptr()), exp_idx.data_ptr(), (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(w1_shared.data_ptr()), exp_idx.data_ptr() + top_k, (int64_t)twoI * H, reinterpret_cast<__nv_bfloat16*>(gate_s.data_ptr()), n_shared, twoI, H, stream); silu_mul_kernel<<>>( reinterpret_cast(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(hid_s.data_ptr()), I, reinterpret_cast(w2_shared.data_ptr()), exp_idx.data_ptr() + 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(y_routed.data_ptr()), reinterpret_cast(y_shared.data_ptr()), reinterpret_cast(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(), n_pairs, E, c.counts_d.data_ptr()); cudaMemcpyAsync(c.counts_h.data_ptr(), c.counts_d.data_ptr(), (size_t)E * sizeof(int), cudaMemcpyDeviceToHost, stream); cudaStreamSynchronize(stream); const int* hc = c.counts_h.data_ptr(); 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<<>>( c.counts_d.data_ptr(), 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<<>>( expert_ids.data_ptr(), reinterpret_cast(x.data_ptr()), reinterpret_cast<__nv_bfloat16*>(c.Xg.data_ptr()), c.cursor_d.data_ptr(), c.row_of_pair.data_ptr(), 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(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(w1_shared.data_ptr()); const __nv_bfloat16* w2s = reinterpret_cast(w2_shared.data_ptr()); gemm_x_wt(T, twoI, H, 1, w1s, x.data_ptr(), c.shared_gate.data_ptr()); silu_mul_kernel<<>>( reinterpret_cast(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(w1_shared.data_ptr()) + (int64_t)s * twoI * H; const __nv_bfloat16* w2s = reinterpret_cast(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<<>>( reinterpret_cast(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(c.shared_y.data_ptr()), shared_fp.data_ptr(), 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(), 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><<>>( reinterpret_cast(shared_c.data_ptr()), reinterpret_cast(c.Y.data_ptr()), c.row_of_pair.data_ptr(), reinterpret_cast(expert_weights.data_ptr()), reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), T, H, top_k); } else { reduce_kernel<0><<>>( reinterpret_cast(shared_c.data_ptr()), reinterpret_cast(c.Y.data_ptr()), c.row_of_pair.data_ptr(), reinterpret_cast(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::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, )