diff --git a/ex_engine/xllm_kernels/cuda/moe_batched_gemm.cu b/ex_engine/xllm_kernels/cuda/moe_batched_gemm.cu deleted file mode 100644 index e8d0a845..00000000 --- a/ex_engine/xllm_kernels/cuda/moe_batched_gemm.cu +++ /dev/null @@ -1,114 +0,0 @@ -// moe_batched_gemm.cu — Fused MoE expert GEMM for BI-V100 -// -// Replaces the Python for-loop over 256 experts with: -// 1. Gather tokens by expert (sorted by moe_compute_index) -// 2. Per-expert GEMM via cublas (torch::mm) -// 3. Fused silu activation -// 4. Per-expert down GEMM -// 5. Weighted scatter-add back to output -// -// This eliminates Python loop overhead (~256 iterations) and reduces -// kernel launch overhead by batching small GEMMs. -// -// Dimensions (Qwen3.5-35B-A3B, TP=4): -// w13: (256, 256, 2048) -> E=256, 2*I=256, H=2048 -// w2: (256, 2048, 128) -> E=256, H=2048, I=128 - -#include -#include -#include - -namespace xllm::kernel::cuda { - -// Fused SiLU-and-mul kernel (gate_up → act) -// gate_up: (N, 2*I), output: (N, I) -__global__ void silu_and_mul_inplace_kernel( - const __half* __restrict__ gate_up, - __half* __restrict__ output, - int64_t N, int64_t I) { - int64_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= N * I) return; - - int64_t row = idx / I; - int64_t col = idx % I; - - float gate = __half2float(gate_up[row * 2 * I + col]); - float up = __half2float(gate_up[row * 2 * I + I + col]); - float silu_gate = gate / (1.0f + expf(-gate)); - output[idx] = __float2half(silu_gate * up); -} - -// Main function: batched expert forward -// Called from Python with pre-sorted token indices -torch::Tensor moe_experts_forward( - torch::Tensor hidden_states, // (T, H) — all tokens - torch::Tensor w13, // (E, 2*I, H) — gate+up weights - torch::Tensor w2, // (E, H, I) — down weights - torch::Tensor sorted_tok_ids, // (T*topk,) — which token for each slot - torch::Tensor sorted_weights, // (T*topk,) — routing weight for each slot - torch::Tensor expert_offsets, // (E+1,) — cumsum of expert_sizes, expert_offsets[0]=0 - int64_t topk) { - - auto stream = at::cuda::getCurrentCUDAStream(); - int64_t T = hidden_states.size(0); - int64_t H = hidden_states.size(1); - int64_t E = w13.size(0); - int64_t two_I = w13.size(1); // 2*I - int64_t I = two_I / 2; - - auto out = torch::zeros({T, H}, hidden_states.options()); - - // Get expert offsets on CPU for loop control - auto offsets_cpu = expert_offsets.to(torch::kCPU, torch::kInt64); - auto offsets_ptr = offsets_cpu.data_ptr(); - - for (int64_t eid = 0; eid < E; ++eid) { - int64_t start = offsets_ptr[eid]; - int64_t end = offsets_ptr[eid + 1]; - int64_t count = end - start; - if (count == 0) continue; - - // Gather tokens for this expert - auto tok_ids = sorted_tok_ids.slice(0, start, end); // (count,) - auto tokens = hidden_states.index_select(0, tok_ids); // (count, H) - - // GEMM 1: gate+up projection - // tokens (count, H) × w13[eid].T (H, 2*I) → (count, 2*I) - auto gate_up = torch::mm(tokens, w13[eid].t()); // (count, 2*I) - - // Fused SiLU activation - auto act = torch::empty({count, I}, hidden_states.options()); - if (hidden_states.dtype() == torch::kFloat16) { - int64_t total = count * I; - int block = 256; - int grid = (total + block - 1) / block; - silu_and_mul_inplace_kernel<<>>( - gate_up.data_ptr(), - act.data_ptr(), - count, I); - } else { - auto chunks = gate_up.chunk(2, /*dim=*/1); - act = torch::silu(chunks[0]) * chunks[1]; - } - - // GEMM 2: down projection - // act (count, I) × w2[eid].T (I, H) → (count, H) - auto expert_out = torch::mm(act, w2[eid].t()); // (count, H) - - // Weighted scatter-add - auto weights = sorted_weights.slice(0, start, end).unsqueeze(1); // (count, 1) - out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype())); - } - - return out; -} - -} // namespace xllm::kernel::cuda - -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.def("moe_experts_forward", &xllm::kernel::cuda::moe_experts_forward, - "Batched MoE expert forward (gather → GEMM → silu → GEMM → scatter)", - py::arg("hidden_states"), py::arg("w13"), py::arg("w2"), - py::arg("sorted_tok_ids"), py::arg("sorted_weights"), - py::arg("expert_offsets"), py::arg("topk")); -} diff --git a/qwen3_6_scripts/build_xllm_kernels.sh b/qwen3_6_scripts/build_xllm_kernels.sh index 8f61d41f..ce8f7653 100644 --- a/qwen3_6_scripts/build_xllm_kernels.sh +++ b/qwen3_6_scripts/build_xllm_kernels.sh @@ -67,8 +67,5 @@ build_kernel "xllm_cache" \ build_kernel "xllm_moe" \ "${CUDA_DIR}/moe/moe_fused_topk.cu" "${CUDA_DIR}/moe/moe_compute_index.cu" "${CUDA_DIR}/moe/moe_combine.cu" "${BIND_DIR}/xllm_moe_bind.cpp" -build_kernel "xllm_moe_gemm" \ - "${CUDA_DIR}/moe_batched_gemm.cu" - echo "=== All kernels built ===" ls -lh "${PREBUILT_DIR}"/xllm_*.so 2>/dev/null || echo "No .so files found" diff --git a/qwen3_6_scripts/qwen3_5.py b/qwen3_6_scripts/qwen3_5.py index 2ce9e06d..c86632b5 100644 --- a/qwen3_6_scripts/qwen3_5.py +++ b/qwen3_6_scripts/qwen3_5.py @@ -156,14 +156,6 @@ except ImportError: except ImportError: _xllm_moe = None -try: - from vllm import xllm_moe_gemm as _xllm_moe_gemm -except ImportError: - try: - import xllm_moe_gemm as _xllm_moe_gemm - except ImportError: - _xllm_moe_gemm = None - try: from vllm import corex_gdn_chunk_recurrent as _corex_gdn_chunk_recurrent except ImportError: @@ -221,13 +213,8 @@ _USE_COREX_MOE_INDEX_COMBINE = ( _USE_XLLM_MOE = ( _xllm_moe is not None and env_bool("BI100_MOE_XLLM", True)) -_USE_XLLM_MOE_GEMM = ( - _xllm_moe_gemm is not None - and env_bool("BI100_MOE_XLLM_GEMM", True)) if _USE_XLLM_MOE: logger.info("xllm_moe ENABLED — fused_topk + compute_index + combine_result") -if _USE_XLLM_MOE_GEMM: - logger.info("xllm_moe_gemm ENABLED — batched expert GEMM (replaces Python loop)") _USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True) # ix_fused_moe: full 7-step fused MoE pipeline via ixformer C++ API @@ -1821,37 +1808,21 @@ class Qwen3_5MoeSparseBlock(nn.Module): expert_counts = torch.bincount( flat_eids, minlength=w13.shape[0]).tolist() - if _USE_XLLM_MOE_GEMM: - # Batched expert GEMM via CUDA — eliminates Python for-loop - # Build expert_offsets from expert_counts (cumsum with leading 0) - expert_counts_t = torch.tensor( - expert_counts if isinstance(expert_counts, list) - else expert_counts.tolist(), - dtype=torch.int64, device=hidden_states.device) - expert_offsets = torch.zeros( - len(expert_counts) + 1, dtype=torch.int64, - device=hidden_states.device) - torch.cumsum(expert_counts_t, dim=0, out=expert_offsets[1:]) - out = _xllm_moe_gemm.moe_experts_forward( - hidden_states, w13, w2, - sorted_tok_ids, sorted_weights.float(), - expert_offsets, self.top_k) - else: - start = 0 - for eid, count in enumerate(expert_counts): - end = start + count - if count == 0: - start = end - continue - tok_ids = sorted_tok_ids[start:end] - tokens = hidden_states[tok_ids] # (n, H) - gate_up = F.linear(tokens, w13[eid]) # (n, 2*I) - gate, up = gate_up.chunk(2, dim=-1) - act = F.silu(gate) * up # (n, I) - expert_out = F.linear(act, w2[eid]) # (n, H) - weights = sorted_weights[start:end].unsqueeze(-1) - out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype)) + start = 0 + for eid, count in enumerate(expert_counts): + end = start + count + if count == 0: start = end + continue + tok_ids = sorted_tok_ids[start:end] + tokens = hidden_states[tok_ids] # (n, H) + gate_up = F.linear(tokens, w13[eid]) # (n, 2*I) + gate, up = gate_up.chunk(2, dim=-1) + act = F.silu(gate) * up # (n, I) + expert_out = F.linear(act, w2[eid]) # (n, H) + weights = sorted_weights[start:end].unsqueeze(-1) + out.index_add_(0, tok_ids, (expert_out * weights).to(out.dtype)) + start = end return out # partial, all-reduce done in forward()