[fix] baseline5 fix __shfl_down_sync

This commit is contained in:
root
2026-08-19 05:30:20 +00:00
parent 4ec61094b7
commit 4d63d394ca

View File

@@ -6,10 +6,11 @@
* gather/index/transpose overhead. * gather/index/transpose overhead.
* *
* BI-V100 hardware adaptation (CoreX 3.2.3, SM70-compat): * BI-V100 hardware adaptation (CoreX 3.2.3, SM70-compat):
* - WARP_SIZE = 64 (was 32 in the original) * - WARP_SIZE = 64 (BI-V100 native)
* - warp_sum uses 6 shuffle-down steps (log2(64)=6) * - Reduction uses shared memory (warp-agnostic, no __shfl_down_sync)
* - lane mask = 63 (0x3F), not 31 (0x1F) * __shfl_down_sync(0xffffffff, ...) only masks 32 threads on BI-V100,
* - kThreads=256 → 4 warps (was 8), grid adjusted accordingly * silently producing wrong results for 64-wide warps.
* - kThreads=256 → 4 warps of 64, grid adjusted accordingly
* - half2 vectorized loads: 2 halves per load, stride by warp width * - half2 vectorized loads: 2 halves per load, stride by warp width
* *
* Model: Qwen3.6-35B-A3B (Qwen3_5 MoE) with TP=4 * Model: Qwen3.6-35B-A3B (Qwen3_5 MoE) with TP=4
@@ -41,20 +42,29 @@ constexpr int kW13Rows = 2 * kIntermediate; // 256
// ===================================================================== // =====================================================================
// BI-V100 hardware constants // BI-V100 hardware constants
// ===================================================================== // =====================================================================
constexpr int kWarpSize = 64; // BI-V100 warp width (was 32) constexpr int kWarpSize = 64; // BI-V100 warp width
constexpr int kThreads = 256; // 4 warps of 64 (was 8 warps of 32) constexpr int kThreads = 256; // 4 warps of 64
constexpr int kWarpsPerBlock = kThreads / kWarpSize; // 4 constexpr int kWarpsPerBlock = kThreads / kWarpSize; // 4
// ===================================================================== // =====================================================================
// Warp-level sum reduction for 64-wide warps // Shared-memory reduction (warp-agnostic, safe for warp=32 or warp=64)
// ===================================================================== // =====================================================================
// 6 steps: 32, 16, 8, 4, 2, 1 (was 5 steps for warp=32) // Each warp gets its own 64-float slice in shared memory.
__device__ __forceinline__ float warp_sum(float value) { // Total smem per block = kWarpsPerBlock * kWarpSize * sizeof(float)
#pragma unroll // = 4 * 64 * 4 = 1024 bytes
for (int offset = kWarpSize / 2; offset > 0; offset /= 2) { __device__ __forceinline__ float smem_warp_sum(
value += __shfl_down_sync(0xffffffff, value, offset); float value, float* warp_smem, int lane) {
warp_smem[lane] = value;
__syncwarp();
// Tree reduction within the warp's shared memory slice
#pragma unroll
for (int s = kWarpSize / 2; s > 0; s >>= 1) {
if (lane < s) {
warp_smem[lane] += warp_smem[lane + s];
}
__syncwarp();
} }
return value; return warp_smem[0];
} }
// ===================================================================== // =====================================================================
@@ -62,7 +72,7 @@ __device__ __forceinline__ float warp_sum(float value) {
// ===================================================================== // =====================================================================
// Grid maps one warp per (slot, output_row) pair. // Grid maps one warp per (slot, output_row) pair.
// Each warp computes dot(input[1,H], W13[eid, row, :]) using half2 loads // Each warp computes dot(input[1,H], W13[eid, row, :]) using half2 loads
// and reduces via 64-wide warp_sum. // and reduces via shared-memory sum.
// //
// Total warps needed: kTopK * kW13Rows = 8 * 256 = 2048 // Total warps needed: kTopK * kW13Rows = 8 * 256 = 2048
// With kWarpsPerBlock=4: 2048/4 = 512 blocks // With kWarpsPerBlock=4: 2048/4 = 512 blocks
@@ -72,11 +82,15 @@ __global__ void direct_w13_kernel(
const int64_t* __restrict__ expert_ids, // (8,) const int64_t* __restrict__ expert_ids, // (8,)
__half* __restrict__ gate_up) { // (8, 256) __half* __restrict__ gate_up) { // (8, 256)
// Map thread to (warp_id → slot, row) and lane within warp // Shared memory: each warp gets kWarpSize floats
__shared__ float smem[kWarpsPerBlock * kWarpSize];
const int warp_in_block = threadIdx.x / kWarpSize; // 0..3
const int lane = threadIdx.x & (kWarpSize - 1); // 0..63
float* warp_smem = smem + warp_in_block * kWarpSize;
const int global_warp = const int global_warp =
static_cast<int>(blockIdx.x) * kWarpsPerBlock + static_cast<int>(blockIdx.x) * kWarpsPerBlock + warp_in_block;
(threadIdx.x / kWarpSize);
const int lane = threadIdx.x & (kWarpSize - 1); // 0..63
if (global_warp >= kTopK * kW13Rows) if (global_warp >= kTopK * kW13Rows)
return; return;
@@ -90,7 +104,6 @@ __global__ void direct_w13_kernel(
(expert * kW13Rows + local_row) * static_cast<int64_t>(kHidden); (expert * kW13Rows + local_row) * static_cast<int64_t>(kHidden);
// Vectorized dot product using half2 loads // Vectorized dot product using half2 loads
// Each lane processes kHidden/2 / kWarpSize iterations
const __half2* input2 = reinterpret_cast<const __half2*>(input); const __half2* input2 = reinterpret_cast<const __half2*>(input);
const __half2* weight2 = reinterpret_cast<const __half2*>(w13 + weight_offset); const __half2* weight2 = reinterpret_cast<const __half2*>(w13 + weight_offset);
@@ -102,8 +115,8 @@ __global__ void direct_w13_kernel(
sum = fmaf(__half2float(w.y), __half2float(x.y), sum); sum = fmaf(__half2float(w.y), __half2float(x.y), sum);
} }
// 64-wide warp reduction // Warp-agnostic shared memory reduction
sum = warp_sum(sum); sum = smem_warp_sum(sum, warp_smem, lane);
// Lane 0 writes the output // Lane 0 writes the output
if (lane == 0) { if (lane == 0) {
@@ -127,10 +140,14 @@ __global__ void direct_w2_reduce_kernel(
const __half* __restrict__ weights, // (8,) const __half* __restrict__ weights, // (8,)
__half* __restrict__ output) { // (1, 2048) __half* __restrict__ output) { // (1, 2048)
const int global_warp = __shared__ float smem[kWarpsPerBlock * kWarpSize];
static_cast<int>(blockIdx.x) * kWarpsPerBlock +
(threadIdx.x / kWarpSize); const int warp_in_block = threadIdx.x / kWarpSize;
const int lane = threadIdx.x & (kWarpSize - 1); const int lane = threadIdx.x & (kWarpSize - 1);
float* warp_smem = smem + warp_in_block * kWarpSize;
const int global_warp =
static_cast<int>(blockIdx.x) * kWarpsPerBlock + warp_in_block;
if (global_warp >= kHidden) if (global_warp >= kHidden)
return; return;
@@ -158,8 +175,8 @@ __global__ void direct_w2_reduce_kernel(
expert_sum = fmaf(__half2float(w.y), __half2float(x.y), expert_sum); expert_sum = fmaf(__half2float(w.y), __half2float(x.y), expert_sum);
} }
// 64-wide warp reduction // Warp-agnostic shared memory reduction
expert_sum = warp_sum(expert_sum); expert_sum = smem_warp_sum(expert_sum, warp_smem, lane);
// Lane 0 accumulates weighted result // Lane 0 accumulates weighted result
if (lane == 0) { if (lane == 0) {
@@ -263,7 +280,7 @@ torch::Tensor direct_w2_reduce(const torch::Tensor& activated,
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("w13", &direct_w13, module.def("w13", &direct_w13,
"Direct selected-expert FP16 W13 matvec (BI-V100, warp64)"); "Direct selected-expert FP16 W13 matvec (BI-V100, smem reduction)");
module.def("w2_reduce", &direct_w2_reduce, module.def("w2_reduce", &direct_w2_reduce,
"Direct selected-expert W2 matvec + routed reduction (BI-V100, warp64)"); "Direct selected-expert W2 matvec + routed reduction (BI-V100, smem reduction)");
} }