diff --git a/qwen3_6_scripts/build_corex_moe_index_combine.sh b/qwen3_6_scripts/build_corex_moe_index_combine.sh new file mode 100644 index 00000000..3912f7e1 --- /dev/null +++ b/qwen3_6_scripts/build_corex_moe_index_combine.sh @@ -0,0 +1,29 @@ +#!/usr/bin/env bash +set -euo pipefail + +VLLM_ROOT=${1:?usage: build_corex_moe_index_combine.sh VLLM_ROOT} +COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3} +TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch} +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +OUTPUT=${VLLM_ROOT}/corex_moe_index_combine.so + +"${COREX_ROOT}/bin/clang++" \ + -std=c++17 -O3 -shared -fPIC \ + --cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \ + --no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \ + -DTORCH_EXTENSION_NAME=corex_moe_index_combine \ + -DTORCH_API_INCLUDE_EXTENSION_H \ + -I"${TORCH_ROOT}/include" \ + -I"${TORCH_ROOT}/include/torch/csrc/api/include" \ + -I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \ + -I/usr/local/include/python3.10 \ + -I"${COREX_ROOT}/include" \ + -I"${SCRIPT_DIR}" \ + "${SCRIPT_DIR}/corex_moe_index_combine.cu" \ + -L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \ + -Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \ + -ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \ + -lc10_cuda -lc10 -lcudart -o "${OUTPUT}" + +test -s "${OUTPUT}" +printf '[ok] CoreX MoE index+combine extension %s\n' "${OUTPUT}" diff --git a/qwen3_6_scripts/corex_moe_index_combine.cu b/qwen3_6_scripts/corex_moe_index_combine.cu new file mode 100644 index 00000000..ac7984b7 --- /dev/null +++ b/qwen3_6_scripts/corex_moe_index_combine.cu @@ -0,0 +1,176 @@ +// corex_moe_index_combine.cu — Fused MoE index computation + combine +// +// Two kernels from xllm/core/kernels/cuda/moe/: +// 1. moe_compute_index: histogram + prefix_sum + place → {src_dst, dst_src, expert_sizes} +// 2. moe_combine_result: weighted sum of expert outputs → final output +// +// These replace Python argsort+bincount+loop in qwen3_5.py _pure_pytorch_experts prefill path. + +#include +#include +#include + +// ========== moe_compute_index ========== + +constexpr int32_t kMoeIndexBlock = 256; + +__global__ void moe_histogram_kernel( + const int32_t* __restrict__ expert_id, + int32_t* __restrict__ expert_sizes, + int64_t num_elements, + int32_t num_experts) { + int64_t tid = int64_t(blockIdx.x) * kMoeIndexBlock + threadIdx.x; + if (tid < num_elements) { + int32_t eid = expert_id[tid]; + if (eid >= 0 && eid < num_experts) { + atomicAdd(&expert_sizes[eid], 1); + } + } +} + +__global__ void moe_prefix_sum_kernel( + const int32_t* __restrict__ expert_sizes, + int32_t* __restrict__ expert_offsets, + int32_t num_experts, + int64_t* __restrict__ total_out) { + using BlockScan = cub::BlockScan; + __shared__ typename BlockScan::TempStorage s_scan; + + int32_t val = (threadIdx.x < num_experts) ? expert_sizes[threadIdx.x] : 0; + int32_t offset; + BlockScan(s_scan).ExclusiveSum(val, offset); + __syncthreads(); + + int32_t total = offset + val; + + if (threadIdx.x < num_experts) { + expert_offsets[threadIdx.x] = offset; + } + if (threadIdx.x == 0 && total_out != nullptr) { + *total_out = total; + } +} + +__global__ void moe_place_indices_kernel( + const int32_t* __restrict__ expert_id, + int32_t* __restrict__ expert_offsets, + int32_t* __restrict__ dst_src, + int32_t* __restrict__ src_dst, + int64_t num_elements, + int32_t num_experts) { + int64_t flat_idx = int64_t(blockIdx.x) * kMoeIndexBlock + threadIdx.x; + if (flat_idx >= num_elements) return; + + int32_t eid = expert_id[flat_idx]; + if (eid < 0 || eid >= num_experts) return; + + int32_t pos = atomicAdd(&expert_offsets[eid], 1); + dst_src[pos] = static_cast(flat_idx); + src_dst[flat_idx] = pos; +} + +std::tuple moe_compute_index( + const torch::Tensor& expert_id, + int64_t num_experts) { + auto device = expert_id.device(); + auto stream = at::cuda::getCurrentCUDAStream(); + int64_t N = expert_id.numel(); + int32_t E = static_cast(num_experts); + auto expert_id_i32 = expert_id.to(torch::kInt32).contiguous(); + auto opt_i32 = expert_id_i32.options(); + + auto expert_sizes = torch::zeros({num_experts}, opt_i32); + auto expert_offsets = torch::empty({num_experts}, opt_i32); + auto dst_src = torch::empty({N}, opt_i32); + auto src_dst = torch::empty({N}, opt_i32); + + int64_t grid = (N + kMoeIndexBlock - 1) / kMoeIndexBlock; + + moe_histogram_kernel<<>>( + expert_id_i32.data_ptr(), + expert_sizes.data_ptr(), + N, E); + + moe_prefix_sum_kernel<<<1, kMoeIndexBlock, 0, stream>>>( + expert_sizes.data_ptr(), + expert_offsets.data_ptr(), + E, nullptr); + + moe_place_indices_kernel<<>>( + expert_id_i32.data_ptr(), + expert_offsets.data_ptr(), + dst_src.data_ptr(), + src_dst.data_ptr(), + N, E); + + return std::make_tuple(src_dst, dst_src, expert_sizes); +} + +// ========== moe_combine_result ========== + +constexpr int32_t kCombineBlockSize = 256; + +template +__global__ void moe_combine_kernel( + const scalar_t* __restrict__ gemm2, + const float* __restrict__ reduce_weight, + scalar_t* __restrict__ output, + int64_t N, + int32_t topk, + int64_t H) { + int64_t token_id = blockIdx.x; + if (token_id >= N) return; + + int32_t tid = threadIdx.x; + int32_t stride = kCombineBlockSize; + + for (int64_t h = tid; h < H; h += stride) { + float acc = 0.0f; + for (int32_t k = 0; k < topk; ++k) { + int64_t flat_idx = token_id * topk + k; + float w = reduce_weight[flat_idx]; + acc += w * static_cast(gemm2[flat_idx * H + h]); + } + output[token_id * H + h] = static_cast(acc); + } +} + +torch::Tensor moe_combine_result( + const torch::Tensor& gemm2, + const torch::Tensor& reduce_weight, + int64_t N, + int64_t topk) { + auto stream = at::cuda::getCurrentCUDAStream(); + int64_t H = gemm2.size(1); + auto dtype = gemm2.scalar_type(); + + auto output = torch::empty({N, H}, gemm2.options()); + auto rw = reduce_weight.to(gemm2.device(), torch::kFloat32).contiguous(); + + if (dtype == torch::kFloat16) { + moe_combine_kernel + <<>>( + gemm2.data_ptr(), + rw.data_ptr(), + output.data_ptr(), + N, static_cast(topk), H); + } else { + moe_combine_kernel + <<>>( + gemm2.data_ptr(), + rw.data_ptr(), + output.data_ptr(), + N, static_cast(topk), H); + } + + return output; +} + +// ========== pybind ========== + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("moe_compute_index", &moe_compute_index, + "Fused MoE token-expert index computation (histogram+prefix_sum+place)"); + m.def("moe_combine_result", &moe_combine_result, + "Fused MoE expert output weighted combination"); +} diff --git a/verify_moe_index_combine.py b/verify_moe_index_combine.py new file mode 100644 index 00000000..ba975e29 --- /dev/null +++ b/verify_moe_index_combine.py @@ -0,0 +1,186 @@ +#!/usr/bin/env python3 +"""Verify moe_compute_index + moe_combine_result on real BI-V100. + +Step 1: Compile corex_moe_index_combine.cu → .so +Step 2: Test moe_compute_index vs PyTorch argsort+bincount +Step 3: Test moe_combine_result vs PyTorch weighted sum +Step 4: End-to-end MoE prefill path benchmark + +Run: python3 verify_moe_index_combine.py +""" + +import sys +import os +import time +import torch +import torch.nn.functional as F + +def compile_kernel(): + """Compile the .so using corex clang++.""" + script_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), + "qwen3_6_scripts") + build_sh = os.path.join(script_dir, "build_corex_moe_index_combine.sh") + # Use a temp vllm root for testing + tmp_root = "/tmp/moe_test" + os.makedirs(tmp_root, exist_ok=True) + ret = os.system(f"bash {build_sh} {tmp_root} 2>&1") + so_path = os.path.join(tmp_root, "corex_moe_index_combine.so") + if ret != 0 or not os.path.exists(so_path): + print(f"[FAIL] Compilation failed (exit={ret})") + return None + print(f"[OK] Compiled: {so_path}") + + import importlib.util + spec = importlib.util.spec_from_file_location( + "corex_moe_index_combine", so_path) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + return mod + + +def pytorch_compute_index(expert_ids_flat, num_experts): + """Reference: what qwen3_5.py does in prefill path.""" + order = torch.argsort(expert_ids_flat, stable=True) + expert_counts = torch.bincount( + expert_ids_flat, minlength=num_experts) + # dst_src[i] = which flat_idx goes to position i (sorted order) + dst_src = torch.arange(len(expert_ids_flat), + device=expert_ids_flat.device)[order] + # src_dst[flat_idx] = position in sorted order + src_dst = torch.empty_like(order) + src_dst[order] = torch.arange(len(order), device=order.device) + return src_dst, dst_src, expert_counts + + +def pytorch_combine(expert_outputs, weights, topk, num_tokens, H): + """Reference: weighted sum of expert outputs.""" + # expert_outputs: (N*topk, H), weights: (N, topk) + out = expert_outputs.view(num_tokens, topk, H) + w = weights.unsqueeze(-1) # (N, topk, 1) + return (out * w).sum(dim=1) # (N, H) + + +def main(): + print("=" * 60) + print("BI-V100 moe_compute_index + moe_combine verification") + print("=" * 60) + + if not torch.cuda.is_available(): + print("FATAL: No CUDA device") + return 1 + + mod = compile_kernel() + if mod is None: + return 1 + + # ---- Test 1: moe_compute_index ---- + print("\n--- Test 1: moe_compute_index (256 experts, 32 tokens, top_k=8) ---") + num_tokens = 32 + num_experts = 256 + topk = 8 + torch.manual_seed(42) + # Simulate topk routing: each token picks 8 experts + topk_ids = torch.randint(0, num_experts, (num_tokens, topk), + device="cuda", dtype=torch.int64) + flat_ids = topk_ids.reshape(-1) # (256,) + + # Kernel + kern_src_dst, kern_dst_src, kern_sizes = mod.moe_compute_index( + flat_ids, num_experts) + + # PyTorch reference + ref_src_dst, ref_dst_src, ref_sizes = pytorch_compute_index( + flat_ids, num_experts) + + # Compare sizes (must match exactly) + sizes_match = torch.equal(kern_sizes.cpu(), ref_sizes.cpu().to(torch.int32)) + print(f" Expert sizes match: {sizes_match}") + + # Compare mappings: verify kern_dst_src is a valid permutation + # that groups tokens by expert + kern_sorted_eids = flat_ids[kern_dst_src.long()] + ref_sorted_eids = flat_ids[ref_dst_src.long()] + # Both should be sorted by expert + kern_sorted = torch.all(kern_sorted_eids[:-1] <= kern_sorted_eids[1:]).item() + ref_sorted = torch.all(ref_sorted_eids[:-1] <= ref_sorted_eids[1:]).item() + print(f" Kernel produces sorted expert order: {kern_sorted}") + print(f" Ref produces sorted expert order: {ref_sorted}") + + # ---- Test 2: moe_combine_result ---- + print("\n--- Test 2: moe_combine_result (32 tokens, top_k=8, H=2048) ---") + H = 2048 + expert_outputs = torch.randn(num_tokens * topk, H, + device="cuda", dtype=torch.float16) + weights = torch.rand(num_tokens, topk, + device="cuda", dtype=torch.float32) + weights = weights / weights.sum(dim=-1, keepdim=True) # normalize + + kern_out = mod.moe_combine_result(expert_outputs, weights, num_tokens, topk) + ref_out = pytorch_combine(expert_outputs, weights, topk, num_tokens, H) + + max_diff = (kern_out.float() - ref_out.float()).abs().max().item() + print(f" Max diff: {max_diff:.8f}") + print(f" Match (tol=1e-3): {max_diff < 1e-3}") + + # ---- Test 3: Performance ---- + print("\n--- Performance: moe_compute_index ---") + flat_ids = torch.randint(0, 256, (256,), device="cuda", dtype=torch.int64) + + # Warmup + for _ in range(10): + mod.moe_compute_index(flat_ids, 256) + pytorch_compute_index(flat_ids, 256) + torch.cuda.synchronize() + + N = 200 + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(N): + mod.moe_compute_index(flat_ids, 256) + torch.cuda.synchronize() + kern_ms = (time.perf_counter() - t0) / N * 1000 + + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(N): + pytorch_compute_index(flat_ids, 256) + torch.cuda.synchronize() + pt_ms = (time.perf_counter() - t0) / N * 1000 + + print(f" Kernel: {kern_ms:.3f} ms") + print(f" PyTorch: {pt_ms:.3f} ms") + print(f" Speedup: {pt_ms/kern_ms:.2f}x") + + print("\n--- Performance: moe_combine_result ---") + expert_outputs = torch.randn(32 * 8, 2048, device="cuda", dtype=torch.float16) + weights = torch.rand(32, 8, device="cuda", dtype=torch.float32) + + for _ in range(10): + mod.moe_combine_result(expert_outputs, weights, 32, 8) + pytorch_combine(expert_outputs, weights, 8, 32, 2048) + torch.cuda.synchronize() + + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(N): + mod.moe_combine_result(expert_outputs, weights, 32, 8) + torch.cuda.synchronize() + kern_ms = (time.perf_counter() - t0) / N * 1000 + + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(N): + pytorch_combine(expert_outputs, weights, 8, 32, 2048) + torch.cuda.synchronize() + pt_ms = (time.perf_counter() - t0) / N * 1000 + + print(f" Kernel: {kern_ms:.3f} ms") + print(f" PyTorch: {pt_ms:.3f} ms") + print(f" Speedup: {pt_ms/kern_ms:.2f}x") + + print("\n" + "=" * 60) + return 0 + + +if __name__ == "__main__": + sys.exit(main())