diff --git a/src/ops/sparse_moe/decode/sparse_moe_decode_kernels.cu b/src/ops/sparse_moe/decode/sparse_moe_decode_kernels.cu index 8aa0e1bc62..f5dd00f701 100644 --- a/src/ops/sparse_moe/decode/sparse_moe_decode_kernels.cu +++ b/src/ops/sparse_moe/decode/sparse_moe_decode_kernels.cu @@ -104,98 +104,128 @@ __global__ void sparse_moe_d2_warp_kernel(const float* __restrict__ scores, int* sparse_moe_select_top8_warp(scores, ids, alpha, shared_scale, selected_logits); } +// The stored planes of one matrix. They travel together because a dot product takes one matrix at +// a time and each codec reads a different subset of them: a call site handing them over one by one +// had to know that subset before it could name the codec. +struct SparseMoePlanes { + const std::uint8_t* __restrict__ codes = nullptr; + const std::uint8_t* __restrict__ high = nullptr; + const std::uint8_t* __restrict__ scales = nullptr; +}; + +// Codecs come in two lane ownerships. Under `kPackedWord8` a lane owns eight consecutive K values +// and a single decode feeds eight FP32 FMAs, so four adjacent groups of eight lanes each form one +// 128-byte warp transaction. The other ownership is scalar - one value per lane in D3, a pair in +// D4 - which is what a 32-wide group leaves room for. Either way the codec takes the matrix's +// column count as a template argument, because it walks its own row and only it knows how its +// planes are laid out. + struct Q4Codec { - static constexpr int kGroupK = 64; - static constexpr bool kD3PackedWord8 = true; + static constexpr int kGroupK = 64; + static constexpr bool kPackedWord8 = true; + static constexpr bool kSingleValuePerLane = false; + + template + __device__ static __forceinline__ void load_eight(const SparseMoePlanes& planes, int row, + int group, int lane_in_group, + float (&weights)[8]) { + const std::int64_t index = static_cast(row) * (K / kGroupK) + group; + const std::uint32_t packed = *reinterpret_cast( + planes.codes + index * Q4RowSplitStorage::kCodeBytesPerGroup + lane_in_group * 4); + const auto scale_bits = *reinterpret_cast( + planes.scales + index * Q4RowSplitStorage::kScaleBytesPerGroup); + Q4SimtDecodeAtom::decode_eight(packed, scale_bits, weights); + } }; struct Q5Codec { - static constexpr int kGroupK = 64; - static constexpr bool kPackedWord8 = true; - - __device__ static __forceinline__ void - load_eight(const std::uint8_t* codes, const std::uint8_t* high, const std::uint8_t* scales, - std::int64_t group_index, int lane_in_group, float (&weights)[8]) { + static constexpr int kGroupK = 64; + static constexpr bool kPackedWord8 = true; + static constexpr bool kSingleValuePerLane = false; + + template + __device__ static __forceinline__ void load_eight(const SparseMoePlanes& planes, int row, + int group, int lane_in_group, + float (&weights)[8]) { + const std::int64_t index = static_cast(row) * (K / kGroupK) + group; const std::uint32_t packed = *reinterpret_cast( - codes + group_index * Q5RowSplitStorage::kCodeBytesPerGroup + lane_in_group * 4); + planes.codes + index * Q5RowSplitStorage::kCodeBytesPerGroup + lane_in_group * 4); const std::uint8_t high_bits = - high[group_index * Q5RowSplitStorage::kHighBytesPerGroup + lane_in_group]; + planes.high[index * Q5RowSplitStorage::kHighBytesPerGroup + lane_in_group]; const auto scale_bits = *reinterpret_cast( - scales + group_index * Q5RowSplitStorage::kScaleBytesPerGroup); + planes.scales + index * Q5RowSplitStorage::kScaleBytesPerGroup); Q5SimtDecodeAtom::decode_eight(packed, high_bits, scale_bits, weights); } }; struct Q6Codec { - static constexpr int kGroupK = 64; - static constexpr bool kPackedWord8 = true; - - __device__ static __forceinline__ void - load_eight(const std::uint8_t* codes, const std::uint8_t* high, const std::uint8_t* scales, - std::int64_t group_index, int lane_in_group, float (&weights)[8]) { + static constexpr int kGroupK = 64; + static constexpr bool kPackedWord8 = true; + static constexpr bool kSingleValuePerLane = false; + + template + __device__ static __forceinline__ void load_eight(const SparseMoePlanes& planes, int row, + int group, int lane_in_group, + float (&weights)[8]) { + const std::int64_t index = static_cast(row) * (K / kGroupK) + group; const std::uint32_t packed = *reinterpret_cast( - codes + group_index * Q6RowSplitStorage::kCodeBytesPerGroup + lane_in_group * 4); + planes.codes + index * Q6RowSplitStorage::kCodeBytesPerGroup + lane_in_group * 4); const std::uint16_t high_bits = *reinterpret_cast( - high + group_index * Q6RowSplitStorage::kHighBytesPerGroup + lane_in_group * 2); + planes.high + index * Q6RowSplitStorage::kHighBytesPerGroup + lane_in_group * 2); const auto scale_bits = *reinterpret_cast( - scales + group_index * Q6RowSplitStorage::kScaleBytesPerGroup); + planes.scales + index * Q6RowSplitStorage::kScaleBytesPerGroup); Q6SimtDecodeAtom::decode_eight(packed, high_bits, scale_bits, weights); } }; struct Q8Codec { - static constexpr int kGroupK = 32; - static constexpr bool kD3SingleValuePerLane = true; - static constexpr bool kD3PackedWord8 = false; - static constexpr bool kPackedWord8 = false; - - __device__ static __forceinline__ float load_one(const std::uint8_t* codes, - const std::uint8_t* scales, - std::int64_t group_index, int lane) { - const float scale = __half2float( - __ushort_as_half(*reinterpret_cast(scales + group_index * 2))); - return static_cast(static_cast(codes[group_index * kGroupK + lane])) * + static constexpr int kGroupK = 32; + static constexpr bool kPackedWord8 = false; + static constexpr bool kSingleValuePerLane = true; + + template + __device__ static __forceinline__ float load_one(const SparseMoePlanes& planes, int row, + int group, int lane) { + const std::int64_t index = static_cast(row) * (K / kGroupK) + group; + const float scale = __half2float(__ushort_as_half(*reinterpret_cast( + planes.scales + index * Q8RowSplitStorage::kScaleBytesPerGroup))); + return static_cast(static_cast( + planes.codes[index * Q8RowSplitStorage::kCodeBytesPerGroup + lane])) * scale; } - __device__ static __forceinline__ void - load_pair(const std::uint8_t* codes, const std::uint8_t* high, const std::uint8_t* scales, - std::int64_t group_index, int lane, float& w0, float& w1) { - Q8ScalarDecodeAtom::load_pair(codes, high, scales, group_index, lane, w0, w1); + template + __device__ static __forceinline__ void load_pair(const SparseMoePlanes& planes, int row, + int group, int lane, float& w0, float& w1) { + Q8ScalarDecodeAtom::load_pair(planes.codes, planes.high, planes.scales, + static_cast(row) * (K / kGroupK) + group, lane, + w0, w1); } }; +// A codec declares exactly one lane ownership. The trap below catches a codec that declares +// neither, which would otherwise compile and reduce a zero accumulator. +template +inline constexpr bool kNoLaneOwnership = false; + template -__device__ __forceinline__ void dot_two_rows(const std::uint8_t* codes, const std::uint8_t* high, - const std::uint8_t* scales, int row0, int row1, +__device__ __forceinline__ void dot_two_rows(const SparseMoePlanes& planes, int row0, int row1, const __nv_bfloat16* x, int k_begin, int k_end, float& result0, float& result1) { - constexpr int kGroups = K / Codec::kGroupK; const int lane = static_cast(threadIdx.x) & 31; float acc0 = 0.0f; float acc1 = 0.0f; const int first_group = k_begin / Codec::kGroupK; const int last_group = k_end / Codec::kGroupK; - if constexpr (Codec::kD3PackedWord8) { - // Four adjacent Q4 groups form one 128-byte warp transaction. Each lane owns eight - // consecutive K values, so one mantissa decode feeds eight FP32 FMAs instead of issuing - // four scalar code-pair/decode iterations. + if constexpr (Codec::kPackedWord8) { const int lane_group = lane >> 3; const int lane_in_group = lane & 7; for (int group_base = first_group; group_base < last_group; group_base += 4) { - const int group = group_base + lane_group; - const std::int64_t index0 = static_cast(row0) * kGroups + group; - const std::int64_t index1 = static_cast(row1) * kGroups + group; - const std::uint32_t packed0 = - *reinterpret_cast(codes + index0 * 32 + lane_in_group * 4); - const std::uint32_t packed1 = - *reinterpret_cast(codes + index1 * 32 + lane_in_group * 4); - const auto scale0 = *reinterpret_cast(scales + index0 * 2); - const auto scale1 = *reinterpret_cast(scales + index1 * 2); + const int group = group_base + lane_group; float weights0[8]; float weights1[8]; - Q4SimtDecodeAtom::decode_eight(packed0, scale0, weights0); - Q4SimtDecodeAtom::decode_eight(packed1, scale1, weights1); + Codec::template load_eight(planes, row0, group, lane_in_group, weights0); + Codec::template load_eight(planes, row1, group, lane_in_group, weights1); const uint4 input = load_vec(x + group * Codec::kGroupK + lane_in_group * 8); const float2 x0 = bf16x2_bits_to_float2(input.x); const float2 x1 = bf16x2_bits_to_float2(input.y); @@ -208,41 +238,25 @@ __device__ __forceinline__ void dot_two_rows(const std::uint8_t* codes, const st acc1 = fmaf(weights1[item], values[item], acc1); } } - } else if constexpr (Codec::kD3SingleValuePerLane) { - for (int group = first_group; group < last_group; ++group) { - const std::int64_t index0 = static_cast(row0) * kGroups + group; - const std::int64_t index1 = static_cast(row1) * kGroups + group; - const float w0 = Codec::load_one(codes, scales, index0, lane); - const float w1 = Codec::load_one(codes, scales, index1, lane); - const float xv = __bfloat162float(x[group * Codec::kGroupK + lane]); - acc0 = fmaf(w0, xv, acc0); - acc1 = fmaf(w1, xv, acc1); - } - } else if (lane < Codec::kGroupK / 2) { + } else if constexpr (Codec::kSingleValuePerLane) { for (int group = first_group; group < last_group; ++group) { - float w00, w01, w10, w11; - Codec::load_pair(codes, high, scales, static_cast(row0) * kGroups + group, - lane, w00, w01); - Codec::load_pair(codes, high, scales, static_cast(row1) * kGroups + group, - lane, w10, w11); - const int k = group * Codec::kGroupK + lane * 2; - const float2 xv = __bfloat1622float2(load_vec<__nv_bfloat162>(x + k)); - acc0 = fmaf(w00, xv.x, acc0); - acc0 = fmaf(w01, xv.y, acc0); - acc1 = fmaf(w10, xv.x, acc1); - acc1 = fmaf(w11, xv.y, acc1); + const float w0 = Codec::template load_one(planes, row0, group, lane); + const float w1 = Codec::template load_one(planes, row1, group, lane); + const float xv = __bfloat162float(x[group * Codec::kGroupK + lane]); + acc0 = fmaf(w0, xv, acc0); + acc1 = fmaf(w1, xv, acc1); } + } else { + static_assert(kNoLaneOwnership, "dot_two_rows: codec declares no lane ownership"); } result0 = warp_reduce_sum(acc0); result1 = warp_reduce_sum(acc1); } -template -__global__ void sparse_moe_d3_nine_warp_kernel( - const __nv_bfloat16* __restrict__ x, const int* __restrict__ ids, - const std::uint8_t* __restrict__ routed_codes, const std::uint8_t* __restrict__ routed_high, - const std::uint8_t* __restrict__ routed_scales, const std::uint8_t* __restrict__ shared_codes, - const std::uint8_t* __restrict__ shared_scales, float* __restrict__ act) { +template +__global__ void sparse_moe_d3_nine_warp_kernel(const __nv_bfloat16* __restrict__ x, + const int* __restrict__ ids, SparseMoePlanes routed, + SparseMoePlanes shared, float* __restrict__ act) { __shared__ __align__(16) __nv_bfloat16 x_shared[kHidden]; const int tid = static_cast(threadIdx.x); const int warp = tid >> 5; @@ -258,23 +272,21 @@ __global__ void sparse_moe_d3_nine_warp_kernel( pdl::wait_for_dependencies(); const int expert = ids[warp]; const int row_base = expert * 1024; - dot_two_rows(routed_codes, routed_high, routed_scales, row_base + j, - row_base + kIntermediate + j, x_shared, 0, kHidden, gate, - up); + dot_two_rows(routed, row_base + j, row_base + kIntermediate + j, + x_shared, 0, kHidden, gate, up); } else { - dot_two_rows(shared_codes, nullptr, shared_scales, j, kIntermediate + j, - x_shared, 0, kHidden, gate, up); + dot_two_rows(shared, j, kIntermediate + j, x_shared, 0, kHidden, gate, + up); } if (lane == 0) { act[static_cast(warp) * kIntermediate + j] = silu(gate) * up; } } -template -__global__ void sparse_moe_d3_path_tiled_kernel( - const __nv_bfloat16* __restrict__ x, const int* __restrict__ token_ids, - const std::uint8_t* __restrict__ routed_codes, const std::uint8_t* __restrict__ routed_high, - const std::uint8_t* __restrict__ routed_scales, const std::uint8_t* __restrict__ shared_codes, - const std::uint8_t* __restrict__ shared_scales, float* __restrict__ token_activations, - int tokens, const int* __restrict__ adaptive_route_jobs) { +template +__global__ void sparse_moe_d3_path_tiled_kernel(const __nv_bfloat16* __restrict__ x, + const int* __restrict__ token_ids, + SparseMoePlanes routed, SparseMoePlanes shared, + float* __restrict__ token_activations, int tokens, + const int* __restrict__ adaptive_route_jobs) { // Three path CTAs per token/output row expose enough blocks for the 170-SM target and keep the // heavier shared Q8 path from holding eight completed routed warps resident. static_assert(PathsPerBlock > 0 && (kTopK + 1) % PathsPerBlock == 0); @@ -317,12 +329,11 @@ __global__ void sparse_moe_d3_path_tiled_kernel( if (path < kTopK) { const int expert = token_ids[token * kTopK + path]; const int row_base = expert * (2 * kIntermediate); - dot_two_rows(routed_codes, routed_high, routed_scales, - row_base + j, row_base + kIntermediate + j, x_shared, - 0, kHidden, gate, up); + dot_two_rows(routed, row_base + j, row_base + kIntermediate + j, + x_shared, 0, kHidden, gate, up); } else { - dot_two_rows(shared_codes, nullptr, shared_scales, j, - kIntermediate + j, x_shared, 0, kHidden, gate, up); + dot_two_rows(shared, j, kIntermediate + j, x_shared, 0, kHidden, + gate, up); } if (lane == 0) { token_activations[(static_cast(token) * (kTopK + 1) + path) * @@ -333,19 +344,18 @@ __global__ void sparse_moe_d3_path_tiled_kernel( } } -template -__device__ __forceinline__ void dot_fp32_rows(const std::uint8_t* codes, const std::uint8_t* high, - const std::uint8_t* scales, int row_base, +template +__device__ __forceinline__ void dot_fp32_rows(const SparseMoePlanes& planes, int row_base, const float* x, int first_group, int last_group, float (&result)[Rows]) { - constexpr int kGroups = kIntermediate / Codec::kGroupK; - const int lane = static_cast(threadIdx.x) & 31; + const int lane = static_cast(threadIdx.x) & 31; float acc[Rows]; #pragma unroll for (int row = 0; row < Rows; ++row) { acc[row] = 0.0f; } if constexpr (Codec::kPackedWord8) { - // Q5/Q6 use the same eight-value lane ownership as D3. The high plane and FP16 scale are - // decoded exactly from their registered row-split codec before FP32 accumulation. + // The same eight-value lane ownership as D3, over the FP32 SwiGLU result rather than the + // BF16 hidden state. Each row decodes from its own stored planes before the FP32 + // accumulation. const int lane_group = lane >> 3; const int lane_in_group = lane & 7; for (int group_base = first_group; group_base < last_group; group_base += 4) { @@ -355,42 +365,43 @@ __device__ __forceinline__ void dot_fp32_rows(const std::uint8_t* codes, const s const float values[8] = {x0.x, x0.y, x0.z, x0.w, x1.x, x1.y, x1.z, x1.w}; #pragma unroll for (int row = 0; row < Rows; ++row) { - const std::int64_t group_index = - static_cast(row_base + row) * kGroups + group; float weights[8]; - Codec::load_eight(codes, high, scales, group_index, lane_in_group, weights); + Codec::template load_eight(planes, row_base + row, group, lane_in_group, + weights); #pragma unroll for (int item = 0; item < 8; ++item) { acc[row] = fmaf(weights[item], values[item], acc[row]); } } } - } else if (lane < Codec::kGroupK / 2) { - for (int group = first_group; group < last_group; ++group) { - const int k = group * Codec::kGroupK + lane * 2; - const float2 xv = load_vec(x + k); + } else if constexpr (Codec::kSingleValuePerLane) { + // Two values a lane here rather than one: the FP32 input is half the width of the BF16 + // hidden state, so the same group needs half the lanes. + if (lane < Codec::kGroupK / 2) { + for (int group = first_group; group < last_group; ++group) { + const int k = group * Codec::kGroupK + lane * 2; + const float2 xv = load_vec(x + k); #pragma unroll - for (int row = 0; row < Rows; ++row) { - float w0, w1; - Codec::load_pair(codes, high, scales, - static_cast(row_base + row) * kGroups + group, lane, - w0, w1); - acc[row] = fmaf(w0, xv.x, acc[row]); - acc[row] = fmaf(w1, xv.y, acc[row]); + for (int row = 0; row < Rows; ++row) { + float w0, w1; + Codec::template load_pair(planes, row_base + row, group, lane, w0, w1); + acc[row] = fmaf(w0, xv.x, acc[row]); + acc[row] = fmaf(w1, xv.y, acc[row]); + } } } + } else { + static_assert(kNoLaneOwnership, "dot_fp32_rows: codec declares no lane ownership"); } #pragma unroll for (int row = 0; row < Rows; ++row) { result[row] = warp_reduce_sum(acc[row]); } } -template +template __global__ void sparse_moe_d4_nine_warp_kernel( const int* __restrict__ ids, const float* __restrict__ alpha, - const float* __restrict__ shared_scale, const float* __restrict__ act, - const std::uint8_t* __restrict__ routed_codes, const std::uint8_t* __restrict__ routed_high, - const std::uint8_t* __restrict__ routed_scales, const std::uint8_t* __restrict__ shared_codes, - const std::uint8_t* __restrict__ shared_scales, __nv_bfloat16* __restrict__ destination, + const float* __restrict__ shared_scale, const float* __restrict__ act, SparseMoePlanes routed, + SparseMoePlanes shared, __nv_bfloat16* __restrict__ destination, const char* __restrict__ prefetch_data, unsigned long long prefetch_bytes) { __shared__ float paths[kTopK + 1][Rows]; pdl::wait_for_dependencies(); @@ -400,19 +411,19 @@ __global__ void sparse_moe_d4_nine_warp_kernel( if (warp < kTopK) { const int expert = ids[warp]; float dot[Rows]; - dot_fp32_rows(routed_codes, routed_high, routed_scales, - expert * kHidden + row_base, - act + static_cast(warp) * kIntermediate, 0, - kIntermediate / RoutedCodec::kGroupK, dot); + dot_fp32_rows( + routed, expert * kHidden + row_base, + act + static_cast(warp) * kIntermediate, 0, + kIntermediate / RoutedCodec::kGroupK, dot); if (lane == 0) { #pragma unroll for (int row = 0; row < Rows; ++row) { paths[warp][row] = alpha[warp] * dot[row]; } } } else { float dot[Rows]; - dot_fp32_rows(shared_codes, nullptr, shared_scales, row_base, - act + static_cast(kTopK) * kIntermediate, 0, - kIntermediate / Q8Codec::kGroupK, dot); + dot_fp32_rows( + shared, row_base, act + static_cast(kTopK) * kIntermediate, 0, + kIntermediate / SharedCodec::kGroupK, dot); if (lane == 0) { #pragma unroll for (int row = 0; row < Rows; ++row) { paths[kTopK][row] = *shared_scale * dot[row]; } @@ -437,14 +448,13 @@ __global__ void sparse_moe_d4_nine_warp_kernel( } } -template -__global__ void sparse_moe_d4_token_kernel( - const int* __restrict__ token_ids, const float* __restrict__ token_alpha, - const float* __restrict__ shared_scale, const float* __restrict__ token_activations, - const std::uint8_t* __restrict__ routed_codes, const std::uint8_t* __restrict__ routed_high, - const std::uint8_t* __restrict__ routed_scales, const std::uint8_t* __restrict__ shared_codes, - const std::uint8_t* __restrict__ shared_scales, __nv_bfloat16* __restrict__ destination, - int tokens, const int* __restrict__ adaptive_route_jobs) { +template +__global__ void +sparse_moe_d4_token_kernel(const int* __restrict__ token_ids, const float* __restrict__ token_alpha, + const float* __restrict__ shared_scale, + const float* __restrict__ token_activations, SparseMoePlanes routed, + SparseMoePlanes shared, __nv_bfloat16* __restrict__ destination, + int tokens, const int* __restrict__ adaptive_route_jobs) { // Token is a grid dimension rather than an in-CTA serial loop. Rows lets one routed-weight // stream serve adjacent outputs while retaining the deterministic rank-order FP32 epilogue. __shared__ float paths[kTopK + 1][Rows]; @@ -468,10 +478,10 @@ __global__ void sparse_moe_d4_token_kernel( if (warp < kTopK) { const int expert = token_ids[token * kTopK + warp]; float dot[Rows]; - dot_fp32_rows(routed_codes, routed_high, routed_scales, - expert * kHidden + row_base, - act + static_cast(warp) * kIntermediate, - 0, kIntermediate / RoutedCodec::kGroupK, dot); + dot_fp32_rows( + routed, expert * kHidden + row_base, + act + static_cast(warp) * kIntermediate, 0, + kIntermediate / RoutedCodec::kGroupK, dot); if (lane == 0) { #pragma unroll for (int row = 0; row < Rows; ++row) { @@ -480,9 +490,9 @@ __global__ void sparse_moe_d4_token_kernel( } } else { float dot[Rows]; - dot_fp32_rows(shared_codes, nullptr, shared_scales, row_base, - act + static_cast(kTopK) * kIntermediate, 0, - kIntermediate / Q8Codec::kGroupK, dot); + dot_fp32_rows( + shared, row_base, act + static_cast(kTopK) * kIntermediate, 0, + kIntermediate / SharedCodec::kGroupK, dot); if (lane == 0) { #pragma unroll for (int row = 0; row < Rows; ++row) { @@ -514,20 +524,21 @@ void launch_d1(const Tensor& x, const SparseMoeWeights& weights, CUDA_CHECK(cudaGetLastError()); } -template +SparseMoePlanes matrix_planes(const Weight& weight) { + return {static_cast(weight.qdata), + static_cast(weight.qhigh), + static_cast(weight.scales)}; +} + +template void launch_d3_dependent_codec(const Tensor& x, const SparseMoeWeights& weights, const SparseMoeDecodeWorkspace& workspace, cudaStream_t stream) { - const auto* input = static_cast(x.data); - const auto* ids = static_cast(workspace.ids.data); - auto* act = static_cast(workspace.scratch.data); - const auto* routed_codes = static_cast(weights.routed_gate_up.qdata); - const auto* routed_high = static_cast(weights.routed_gate_up.qhigh); - const auto* routed_scales = static_cast(weights.routed_gate_up.scales); - const auto* shared_codes = static_cast(weights.shared_gate_up.qdata); - const auto* shared_scales = static_cast(weights.shared_gate_up.scales); CUDA_CHECK(pdl::launch_dependent( - {dim3(kIntermediate), dim3(9 * 32), 0, stream}, sparse_moe_d3_nine_warp_kernel, - input, ids, routed_codes, routed_high, routed_scales, shared_codes, shared_scales, act)); + {dim3(kIntermediate), dim3(9 * 32), 0, stream}, + sparse_moe_d3_nine_warp_kernel, + static_cast(x.data), static_cast(workspace.ids.data), + matrix_planes(weights.routed_gate_up), matrix_planes(weights.shared_gate_up), + static_cast(workspace.scratch.data))); } void launch_d2_d3(const Tensor& x, const SparseMoeWeights& weights, @@ -541,35 +552,30 @@ void launch_d2_d3(const Tensor& x, const SparseMoeWeights& weights, switch (weights.routed_gate_up.qtype) { case QType::Q4_G64_FP16: - launch_d3_dependent_codec(x, weights, workspace, stream); + launch_d3_dependent_codec(x, weights, workspace, stream); return; case QType::Q8_G32_FP16: - launch_d3_dependent_codec(x, weights, workspace, stream); + launch_d3_dependent_codec(x, weights, workspace, stream); return; default: throw std::invalid_argument("sparse_moe: unsupported D3 codec"); } } -template +template void launch_d4_dependent_codec(const SparseMoeWeights& weights, Tensor& destination, const SparseMoeDecodeWorkspace& workspace, cudaStream_t stream, const void* prefetch_data, std::size_t prefetch_bytes) { - const auto* ids = static_cast(workspace.ids.data); - const auto* alpha = static_cast(workspace.alpha.data); - const auto* shared_scale = static_cast(workspace.shared_scale.data); - const auto* act = static_cast(workspace.scratch.data); - const auto* routed_codes = static_cast(weights.routed_down.qdata); - const auto* routed_high = static_cast(weights.routed_down.qhigh); - const auto* routed_scales = static_cast(weights.routed_down.scales); - const auto* shared_codes = static_cast(weights.shared_down.qdata); - const auto* shared_scales = static_cast(weights.shared_down.scales); - auto* output = static_cast<__nv_bfloat16*>(destination.data); + constexpr int kRows = 1; CUDA_CHECK(pdl::launch_dependent( - {dim3(kHidden), dim3(9 * 32), 0, stream}, sparse_moe_d4_nine_warp_kernel, ids, - alpha, shared_scale, act, routed_codes, routed_high, routed_scales, shared_codes, - shared_scales, output, static_cast(prefetch_data), - static_cast(prefetch_bytes))); + {dim3(kHidden / kRows), dim3(9 * 32), 0, stream}, + sparse_moe_d4_nine_warp_kernel, + static_cast(workspace.ids.data), + static_cast(workspace.alpha.data), + static_cast(workspace.shared_scale.data), + static_cast(workspace.scratch.data), matrix_planes(weights.routed_down), + matrix_planes(weights.shared_down), static_cast<__nv_bfloat16*>(destination.data), + static_cast(prefetch_data), static_cast(prefetch_bytes))); } void launch_d4_dependent(const SparseMoeWeights& weights, Tensor& destination, @@ -577,89 +583,81 @@ void launch_d4_dependent(const SparseMoeWeights& weights, Tensor& destination, const void* prefetch_data, std::size_t prefetch_bytes) { switch (weights.routed_down.qtype) { case QType::Q5_G64_FP16: - launch_d4_dependent_codec(weights, destination, workspace, stream, prefetch_data, - prefetch_bytes); + launch_d4_dependent_codec(weights, destination, workspace, stream, + prefetch_data, prefetch_bytes); return; case QType::Q6_G64_FP16: - launch_d4_dependent_codec(weights, destination, workspace, stream, prefetch_data, - prefetch_bytes); + launch_d4_dependent_codec(weights, destination, workspace, stream, + prefetch_data, prefetch_bytes); return; case QType::Q8_G32_FP16: - launch_d4_dependent_codec(weights, destination, workspace, stream, prefetch_data, - prefetch_bytes); + launch_d4_dependent_codec(weights, destination, workspace, stream, + prefetch_data, prefetch_bytes); return; default: throw std::invalid_argument("sparse_moe: unsupported D4 codec"); } } -template +template void launch_d3_small_t_paths(const Tensor& x, const SparseMoeWeights& weights, const int* token_ids, float* token_activations, std::int32_t tokens, cudaStream_t stream, const int* adaptive_route_jobs) { - constexpr int kPathBlocks = (kTopK + 1) / PathsPerBlock; - const auto* input = static_cast(x.data); - const auto* routed_codes = static_cast(weights.routed_gate_up.qdata); - const auto* routed_high = static_cast(weights.routed_gate_up.qhigh); - const auto* routed_scales = static_cast(weights.routed_gate_up.scales); - const auto* shared_codes = static_cast(weights.shared_gate_up.qdata); - const auto* shared_scales = static_cast(weights.shared_gate_up.scales); + constexpr int kPathBlocks = (kTopK + 1) / PathsPerBlock; + const auto* input = static_cast(x.data); + const SparseMoePlanes routed = matrix_planes(weights.routed_gate_up); + const SparseMoePlanes shared = matrix_planes(weights.shared_gate_up); if constexpr (Adaptive) { - sparse_moe_d3_path_tiled_kernel + sparse_moe_d3_path_tiled_kernel <<>>( - input, token_ids, routed_codes, routed_high, routed_scales, shared_codes, - shared_scales, token_activations, tokens, adaptive_route_jobs); + input, token_ids, routed, shared, token_activations, tokens, adaptive_route_jobs); CUDA_CHECK(cudaGetLastError()); } else { CUDA_CHECK(pdl::launch_dependent( {dim3(kIntermediate, tokens * kPathBlocks), dim3(PathsPerBlock * 32), 0, stream}, - sparse_moe_d3_path_tiled_kernel, input, token_ids, - routed_codes, routed_high, routed_scales, shared_codes, shared_scales, - token_activations, tokens, nullptr)); + sparse_moe_d3_path_tiled_kernel, input, + token_ids, routed, shared, token_activations, tokens, nullptr)); } } -template +template void launch_d3_small_t_codec(const Tensor& x, const SparseMoeWeights& weights, const int* token_ids, float* token_activations, std::int32_t tokens, SparseMoeSmallTD3Schedule schedule, cudaStream_t stream, const int* adaptive_route_jobs) { switch (schedule) { case SparseMoeSmallTD3Schedule::Paths1: - launch_d3_small_t_paths(x, weights, token_ids, token_activations, - tokens, stream, adaptive_route_jobs); + launch_d3_small_t_paths( + x, weights, token_ids, token_activations, tokens, stream, adaptive_route_jobs); return; case SparseMoeSmallTD3Schedule::Paths3: - launch_d3_small_t_paths(x, weights, token_ids, token_activations, - tokens, stream, adaptive_route_jobs); + launch_d3_small_t_paths( + x, weights, token_ids, token_activations, tokens, stream, adaptive_route_jobs); return; case SparseMoeSmallTD3Schedule::Paths9: - launch_d3_small_t_paths(x, weights, token_ids, token_activations, - tokens, stream, adaptive_route_jobs); + launch_d3_small_t_paths( + x, weights, token_ids, token_activations, tokens, stream, adaptive_route_jobs); return; } throw std::logic_error("sparse_moe: unknown small-T D3 schedule"); } -template +template void launch_d4_small_t_rows(const SparseMoeWeights& weights, Tensor& destination, const int* token_ids, const float* token_alpha, const float* shared_scale, const float* token_activations, std::int32_t tokens, cudaStream_t stream, const int* adaptive_route_jobs) { const dim3 grid = Adaptive ? dim3(kAdaptiveD4Blocks) : dim3(kHidden / Rows, tokens); - sparse_moe_d4_token_kernel<<>>( - token_ids, token_alpha, shared_scale, token_activations, - static_cast(weights.routed_down.qdata), - static_cast(weights.routed_down.qhigh), - static_cast(weights.routed_down.scales), - static_cast(weights.shared_down.qdata), - static_cast(weights.shared_down.scales), - static_cast<__nv_bfloat16*>(destination.data), tokens, adaptive_route_jobs); + sparse_moe_d4_token_kernel + <<>>( + token_ids, token_alpha, shared_scale, token_activations, + matrix_planes(weights.routed_down), matrix_planes(weights.shared_down), + static_cast<__nv_bfloat16*>(destination.data), tokens, adaptive_route_jobs); CUDA_CHECK(cudaGetLastError()); } -template +template void launch_d4_small_t_codec(const SparseMoeWeights& weights, Tensor& destination, const int* token_ids, const float* token_alpha, const float* shared_scale, const float* token_activations, @@ -667,19 +665,19 @@ void launch_d4_small_t_codec(const SparseMoeWeights& weights, Tensor& destinatio cudaStream_t stream, const int* adaptive_route_jobs) { switch (schedule) { case SparseMoeSmallTD4Schedule::Rows1: - launch_d4_small_t_rows(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, stream, - adaptive_route_jobs); + launch_d4_small_t_rows( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, tokens, + stream, adaptive_route_jobs); return; case SparseMoeSmallTD4Schedule::Rows2: - launch_d4_small_t_rows(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, stream, - adaptive_route_jobs); + launch_d4_small_t_rows( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, tokens, + stream, adaptive_route_jobs); return; case SparseMoeSmallTD4Schedule::Rows4: - launch_d4_small_t_rows(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, stream, - adaptive_route_jobs); + launch_d4_small_t_rows( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, tokens, + stream, adaptive_route_jobs); return; } throw std::logic_error("sparse_moe: unknown small-T D4 schedule"); @@ -694,16 +692,17 @@ void sparse_moe_decode_launch_d3_small_t(const Tensor& x, const SparseMoeWeights switch (weights.routed_gate_up.qtype) { case QType::Q4_G64_FP16: if (adaptive_route_jobs == nullptr) { - launch_d3_small_t_codec(x, weights, token_ids, token_activations, - tokens, schedule, stream, nullptr); + launch_d3_small_t_codec( + x, weights, token_ids, token_activations, tokens, schedule, stream, nullptr); } else { - launch_d3_small_t_codec(x, weights, token_ids, token_activations, tokens, - schedule, stream, adaptive_route_jobs); + launch_d3_small_t_codec(x, weights, token_ids, + token_activations, tokens, schedule, + stream, adaptive_route_jobs); } return; case QType::Q8_G32_FP16: - launch_d3_small_t_codec(x, weights, token_ids, token_activations, tokens, - schedule, stream, nullptr); + launch_d3_small_t_codec(x, weights, token_ids, token_activations, + tokens, schedule, stream, nullptr); return; default: throw std::invalid_argument("sparse_moe: unsupported small-T D3 codec"); @@ -718,30 +717,30 @@ void sparse_moe_decode_launch_d4_small_t(const SparseMoeWeights& weights, Tensor switch (weights.routed_down.qtype) { case QType::Q5_G64_FP16: if (adaptive_route_jobs == nullptr) { - launch_d4_small_t_codec(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, - schedule, stream, nullptr); + launch_d4_small_t_codec( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, + tokens, schedule, stream, nullptr); } else { - launch_d4_small_t_codec(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, - schedule, stream, adaptive_route_jobs); + launch_d4_small_t_codec( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, + tokens, schedule, stream, adaptive_route_jobs); } return; case QType::Q6_G64_FP16: if (adaptive_route_jobs == nullptr) { - launch_d4_small_t_codec(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, - schedule, stream, nullptr); + launch_d4_small_t_codec( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, + tokens, schedule, stream, nullptr); } else { - launch_d4_small_t_codec(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, - schedule, stream, adaptive_route_jobs); + launch_d4_small_t_codec( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, + tokens, schedule, stream, adaptive_route_jobs); } return; case QType::Q8_G32_FP16: - launch_d4_small_t_codec(weights, destination, token_ids, token_alpha, - shared_scale, token_activations, tokens, schedule, - stream, nullptr); + launch_d4_small_t_codec( + weights, destination, token_ids, token_alpha, shared_scale, token_activations, tokens, + schedule, stream, nullptr); return; default: throw std::invalid_argument("sparse_moe: unsupported small-T D4 codec");