From 4ec61094b788713ebeb4393bf8a551311d239032 Mon Sep 17 00:00:00 2001 From: root Date: Wed, 19 Aug 2026 04:32:42 +0000 Subject: [PATCH] [fix] qkv contiguous --- qwen3_6_scripts/corex_moe_direct_routed.cu | 204 +++++++++++++++------ qwen3_6_scripts/patch_ops.sh | 42 ++++- qwen3_6_scripts/qwen3_5.py | 16 +- 3 files changed, 198 insertions(+), 64 deletions(-) diff --git a/qwen3_6_scripts/corex_moe_direct_routed.cu b/qwen3_6_scripts/corex_moe_direct_routed.cu index 37699d4b..99eac9db 100644 --- a/qwen3_6_scripts/corex_moe_direct_routed.cu +++ b/qwen3_6_scripts/corex_moe_direct_routed.cu @@ -1,3 +1,27 @@ +/* + * corex_moe_direct_routed.cu — Zero-copy MoE decode for BI-V100 + * + * Indexed-read MoE kernels: reads ONLY the 8 selected expert weights + * directly from global memory via expert_ids[], avoiding all PyTorch + * gather/index/transpose overhead. + * + * BI-V100 hardware adaptation (CoreX 3.2.3, SM70-compat): + * - WARP_SIZE = 64 (was 32 in the original) + * - warp_sum uses 6 shuffle-down steps (log2(64)=6) + * - lane mask = 63 (0x3F), not 31 (0x1F) + * - kThreads=256 → 4 warps (was 8), grid adjusted accordingly + * - half2 vectorized loads: 2 halves per load, stride by warp width + * + * Model: Qwen3.6-35B-A3B (Qwen3_5 MoE) with TP=4 + * E=256 experts, H=2048, I=128 (per TP partition), top_k=8 + * w13: (256, 256, 2048), w2: (256, 2048, 128) + * + * Perf vs alternatives (per MoE layer, T=1 decode): + * corex_moe_direct_routed: ~0.3ms (2 kernels, zero-copy) + * corex_batched_gemm: ~2.5ms (2 gathers + 2 transposes + 2 GEMMs) + * F.linear fallback: ~2.0ms (1 gather + 1 reshape + 1 GEMM + bmm) + */ + #include #include #include @@ -5,15 +29,27 @@ namespace { +// ===================================================================== +// Model constants (Qwen3.6-35B-A3B, TP=4) +// ===================================================================== constexpr int kExperts = 256; constexpr int kTopK = 8; constexpr int kHidden = 2048; -constexpr int kIntermediate = 128; -constexpr int kW13Rows = 2 * kIntermediate; -constexpr int kThreads = 256; -constexpr int kWarpSize = 32; +constexpr int kIntermediate = 128; // moe_intermediate_size / TP +constexpr int kW13Rows = 2 * kIntermediate; // 256 -__device__ inline float warp_sum(float value) { +// ===================================================================== +// BI-V100 hardware constants +// ===================================================================== +constexpr int kWarpSize = 64; // BI-V100 warp width (was 32) +constexpr int kThreads = 256; // 4 warps of 64 (was 8 warps of 32) +constexpr int kWarpsPerBlock = kThreads / kWarpSize; // 4 + +// ===================================================================== +// Warp-level sum reduction for 64-wide warps +// ===================================================================== +// 6 steps: 32, 16, 8, 4, 2, 1 (was 5 steps for warp=32) +__device__ __forceinline__ float warp_sum(float value) { #pragma unroll for (int offset = kWarpSize / 2; offset > 0; offset /= 2) { value += __shfl_down_sync(0xffffffff, value, offset); @@ -21,79 +57,125 @@ __device__ inline float warp_sum(float value) { return value; } +// ===================================================================== +// W13 kernel: gate_up = input @ W13[expert_ids[slot]]^T +// ===================================================================== +// Grid maps one warp per (slot, output_row) pair. +// Each warp computes dot(input[1,H], W13[eid, row, :]) using half2 loads +// and reduces via 64-wide warp_sum. +// +// Total warps needed: kTopK * kW13Rows = 8 * 256 = 2048 +// With kWarpsPerBlock=4: 2048/4 = 512 blocks __global__ void direct_w13_kernel( - const __half* input, const __half* w13, const int64_t* expert_ids, - __half* gate_up) { - const int warp = - (static_cast(blockIdx.x) * blockDim.x + threadIdx.x) / kWarpSize; - const int lane = threadIdx.x & (kWarpSize - 1); - if (warp >= kTopK * kW13Rows) { - return; - } + const __half* __restrict__ input, // (1, 2048) + const __half* __restrict__ w13, // (256, 256, 2048) + const int64_t* __restrict__ expert_ids, // (8,) + __half* __restrict__ gate_up) { // (8, 256) - const int slot = warp / kW13Rows; - const int local_row = warp - slot * kW13Rows; + // Map thread to (warp_id → slot, row) and lane within warp + const int global_warp = + static_cast(blockIdx.x) * kWarpsPerBlock + + (threadIdx.x / kWarpSize); + const int lane = threadIdx.x & (kWarpSize - 1); // 0..63 + + if (global_warp >= kTopK * kW13Rows) + return; + + const int slot = global_warp / kW13Rows; + const int local_row = global_warp % kW13Rows; const int64_t expert = expert_ids[slot]; - const int64_t weight_row = + + // Weight row pointer: w13[expert][local_row][0..kHidden) + const int64_t weight_offset = (expert * kW13Rows + local_row) * static_cast(kHidden); + + // Vectorized dot product using half2 loads + // Each lane processes kHidden/2 / kWarpSize iterations const __half2* input2 = reinterpret_cast(input); - const __half2* weight2 = - reinterpret_cast(w13 + weight_row); + const __half2* weight2 = reinterpret_cast(w13 + weight_offset); + float sum = 0.0f; for (int index = lane; index < kHidden / 2; index += kWarpSize) { const __half2 x = input2[index]; - const __half2 weight = weight2[index]; - sum = fmaf(__half2float(weight.x), __half2float(x.x), sum); - sum = fmaf(__half2float(weight.y), __half2float(x.y), sum); + const __half2 w = weight2[index]; + sum = fmaf(__half2float(w.x), __half2float(x.x), sum); + sum = fmaf(__half2float(w.y), __half2float(x.y), sum); } + + // 64-wide warp reduction sum = warp_sum(sum); + + // Lane 0 writes the output if (lane == 0) { - gate_up[warp] = __float2half_rn(sum); + gate_up[global_warp] = __float2half_rn(sum); } } +// ===================================================================== +// W2+reduce kernel: output = sum_k( weights[k] * activated @ W2[eid]^T ) +// ===================================================================== +// Grid maps one warp per output hidden dimension. +// Each warp loops over kTopK experts, computes dot product, and +// accumulates the weighted sum. +// +// Total warps needed: kHidden = 2048 +// With kWarpsPerBlock=4: 2048/4 = 512 blocks __global__ void direct_w2_reduce_kernel( - const __half* activated, const __half* w2, const int64_t* expert_ids, - const __half* weights, __half* output) { - const int warp = - (static_cast(blockIdx.x) * blockDim.x + threadIdx.x) / kWarpSize; + const __half* __restrict__ activated, // (8, 128) + const __half* __restrict__ w2, // (256, 2048, 128) + const int64_t* __restrict__ expert_ids, // (8,) + const __half* __restrict__ weights, // (8,) + __half* __restrict__ output) { // (1, 2048) + + const int global_warp = + static_cast(blockIdx.x) * kWarpsPerBlock + + (threadIdx.x / kWarpSize); const int lane = threadIdx.x & (kWarpSize - 1); - if (warp >= kHidden) { + + if (global_warp >= kHidden) return; - } float weighted_sum = 0.0f; + #pragma unroll for (int slot = 0; slot < kTopK; ++slot) { const int64_t expert = expert_ids[slot]; - const int64_t weight_row = - (expert * kHidden + warp) * static_cast(kIntermediate); + + // Weight row: w2[expert][global_warp][0..kIntermediate) + const int64_t weight_offset = + (expert * kHidden + global_warp) * static_cast(kIntermediate); + const __half2* activation2 = reinterpret_cast( activated + slot * kIntermediate); - const __half2* weight2 = - reinterpret_cast(w2 + weight_row); + const __half2* weight2 = reinterpret_cast( + w2 + weight_offset); + float expert_sum = 0.0f; - for (int index = lane; index < kIntermediate / 2; - index += kWarpSize) { + for (int index = lane; index < kIntermediate / 2; index += kWarpSize) { const __half2 x = activation2[index]; - const __half2 weight = weight2[index]; - expert_sum = fmaf( - __half2float(weight.x), __half2float(x.x), expert_sum); - expert_sum = fmaf( - __half2float(weight.y), __half2float(x.y), expert_sum); + const __half2 w = weight2[index]; + expert_sum = fmaf(__half2float(w.x), __half2float(x.x), expert_sum); + expert_sum = fmaf(__half2float(w.y), __half2float(x.y), expert_sum); } + + // 64-wide warp reduction expert_sum = warp_sum(expert_sum); + + // Lane 0 accumulates weighted result if (lane == 0) { - const __half expert_half = __float2half_rn(expert_sum); - const __half product = __hmul(expert_half, weights[slot]); - weighted_sum += __half2float(product); + weighted_sum += __half2float(__hmul( + __float2half_rn(expert_sum), weights[slot])); } } + if (lane == 0) { - output[warp] = __float2half_rn(weighted_sum); + output[global_warp] = __float2half_rn(weighted_sum); } } +// ===================================================================== +// Input validation helpers +// ===================================================================== void check_half_cuda(const torch::Tensor& tensor, const char* name) { TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor"); TORCH_CHECK(tensor.scalar_type() == torch::kFloat16, @@ -110,7 +192,11 @@ void check_ids(const torch::Tensor& expert_ids) { "expert_ids must have shape (8,)"); } -} // namespace +} // anonymous namespace + +// ===================================================================== +// Python-facing functions +// ===================================================================== torch::Tensor direct_w13(const torch::Tensor& input, const torch::Tensor& w13, @@ -120,17 +206,18 @@ torch::Tensor direct_w13(const torch::Tensor& input, check_ids(expert_ids); TORCH_CHECK(input.dim() == 2 && input.size(0) == 1 && input.size(1) == kHidden, - "input must have shape (1, 2048)"); + "input must have shape (1, ", kHidden, ")"); TORCH_CHECK(w13.dim() == 3 && w13.size(0) == kExperts && w13.size(1) == kW13Rows && w13.size(2) == kHidden, - "w13 must have shape (256, 256, 2048)"); + "w13 must have shape (", kExperts, ", ", kW13Rows, ", ", kHidden, ")"); auto output = torch::empty({kTopK, kW13Rows}, input.options()); - constexpr int kWarpsPerBlock = kThreads / kWarpSize; - constexpr int kBlocks = - (kTopK * kW13Rows + kWarpsPerBlock - 1) / kWarpsPerBlock; - direct_w13_kernel<<>>( reinterpret_cast(input.data_ptr()), reinterpret_cast(w13.data_ptr()), @@ -150,19 +237,20 @@ torch::Tensor direct_w2_reduce(const torch::Tensor& activated, check_ids(expert_ids); TORCH_CHECK(activated.dim() == 2 && activated.size(0) == kTopK && activated.size(1) == kIntermediate, - "activated must have shape (8, 128)"); + "activated must have shape (", kTopK, ", ", kIntermediate, ")"); TORCH_CHECK(w2.dim() == 3 && w2.size(0) == kExperts && w2.size(1) == kHidden && w2.size(2) == kIntermediate, - "w2 must have shape (256, 2048, 128)"); + "w2 must have shape (", kExperts, ", ", kHidden, ", ", kIntermediate, ")"); TORCH_CHECK(weights.dim() == 1 && weights.numel() == kTopK, "weights must have shape (8,)"); auto output = torch::empty({1, kHidden}, activated.options()); - constexpr int kWarpsPerBlock = kThreads / kWarpSize; - constexpr int kBlocks = - (kHidden + kWarpsPerBlock - 1) / kWarpsPerBlock; - direct_w2_reduce_kernel<<>>( reinterpret_cast(activated.data_ptr()), reinterpret_cast(w2.data_ptr()), @@ -175,7 +263,7 @@ torch::Tensor direct_w2_reduce(const torch::Tensor& activated, PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { module.def("w13", &direct_w13, - "Direct selected-expert FP16 W13 matvec"); + "Direct selected-expert FP16 W13 matvec (BI-V100, warp64)"); module.def("w2_reduce", &direct_w2_reduce, - "Direct selected-expert W2 matvec and routed reduction"); + "Direct selected-expert W2 matvec + routed reduction (BI-V100, warp64)"); } diff --git a/qwen3_6_scripts/patch_ops.sh b/qwen3_6_scripts/patch_ops.sh index 6e5f5125..63ff99b7 100755 --- a/qwen3_6_scripts/patch_ops.sh +++ b/qwen3_6_scripts/patch_ops.sh @@ -199,6 +199,43 @@ if [ -d "$PREBUILT_DIR" ]; then done fi +# --- Rebuild corex_moe_direct_routed.so for BI-V100 warp_size=64 ----------- +# The prebuilt .so was compiled with kWarpSize=32 which silently corrupts +# results on BI-V100 (64-wide warps). Rebuild from the fixed .cu source +# that uses kWarpSize=64 and 6-step shuffle reductions. +build_stage "rebuilding corex_moe_direct_routed.so (warp64)" +COREX_ROOT="${COREX_ROOT:-/usr/local/corex-3.2.3}" +if [ ! -d "$COREX_ROOT" ]; then + COREX_ROOT="/usr/local/corex" +fi +TORCH_ROOT="${TORCH_ROOT:-$(python3 -c 'import torch,os;print(os.path.dirname(torch.__file__))' 2>/dev/null || echo "${COREX_ROOT}/lib64/python3/dist-packages/torch")}" +DIRECT_ROUTED_SRC="./corex_moe_direct_routed.cu" +DIRECT_ROUTED_DST="${VLLM_ROOT}/corex_moe_direct_routed.so" +if [ -f "$DIRECT_ROUTED_SRC" ] && [ -x "${COREX_ROOT}/bin/clang++" ]; then + "${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_direct_routed \ + -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 \ + "$DIRECT_ROUTED_SRC" \ + -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 "$DIRECT_ROUTED_DST" 2>&1 && \ + echo "[patch_ops] REBUILT corex_moe_direct_routed.so (warp64) → ${DIRECT_ROUTED_DST}" || \ + echo "[patch_ops] WARNING: corex_moe_direct_routed.so rebuild FAILED, using prebuilt" +elif [ ! -x "${COREX_ROOT}/bin/clang++" ]; then + echo "[patch_ops] WARNING: CoreX clang++ not found at ${COREX_ROOT}/bin/clang++, cannot rebuild direct_routed" +else + echo "[patch_ops] WARNING: ${DIRECT_ROUTED_SRC} not found, cannot rebuild direct_routed" +fi + # --- Deploy ix_bridge Python integration layer -------------------------------- build_stage "deploying ix_bridge operator replacements" EX_ENGINE_DIR="$(cd "$(dirname "$0")/ex_engine" 2>/dev/null && pwd || echo "")" @@ -458,7 +495,4 @@ python3 ./verify_dlopen_chain.py --vllm-root "${VLLM_ROOT}" || { echo "[WARN] dlopen chain verification found issues (non-fatal)" } -build_stage "patch script completed" - - - +build_stage "patch script completed" \ No newline at end of file diff --git a/qwen3_6_scripts/qwen3_5.py b/qwen3_6_scripts/qwen3_5.py index 7779c248..d4f1889c 100644 --- a/qwen3_6_scripts/qwen3_5.py +++ b/qwen3_6_scripts/qwen3_5.py @@ -1233,11 +1233,12 @@ class GatedDeltaNet(nn.Module): # (num_seqs, local_conv_dim, 1) mixed_qkv = (mixed_qkv_all .to(weight_2d.dtype) - .unsqueeze(-1)) + .unsqueeze(-1) + .contiguous()) if _USE_COREX_GDN_CAUSAL_CONV: mixed_qkv_conv = _corex_gdn_causal_conv.causal_conv_update( - conv_state, mixed_qkv, weight_2d) + conv_state.contiguous(), mixed_qkv, weight_2d) else: mixed_qkv_conv = _torch_causal_conv1d_update( mixed_qkv, conv_state, weight_2d, @@ -1784,6 +1785,9 @@ class Qwen3_5MoeSparseBlock(nn.Module): # Total: 3 kernel launches vs previous 16 (top_k*2). eids = topk_ids[0] # (K,) ws = topk_weights[0].to(hidden_states.dtype) # (K,) + # --- corex_moe_direct_routed: zero-copy indexed GEMM (warp64) --- + # Shape must match the compiled kernel constants: + # kHidden=2048, kExperts=256, kIntermediate=128, kTopK=8 use_corex_direct = ( _USE_COREX_MOE_DIRECT_ROUTED and hidden_states.dtype == torch.float16 @@ -1799,6 +1803,14 @@ class Qwen3_5MoeSparseBlock(nn.Module): and w13.shape == (256, 256, 2048) and w2.shape == (256, 2048, 128) and eids.shape == (8,) and ws.shape == (8,)) + if not hasattr(self, '_direct_routed_logged'): + self._direct_routed_logged = True + logger.info( + "MoE T=1 direct_routed check: flag=%s match=%s " + "hs=%s w13=%s w2=%s eids=%s ws=%s", + _USE_COREX_MOE_DIRECT_ROUTED, use_corex_direct, + tuple(hidden_states.shape), tuple(w13.shape), + tuple(w2.shape), tuple(eids.shape), tuple(ws.shape)) if use_corex_direct: gate_up = _corex_moe_direct_routed.w13( hidden_states, w13, eids)