-
Notifications
You must be signed in to change notification settings - Fork 24k
CUDA: extend MOE fusion to specdec, earlier MOE glu fusion and topk-router fusion were restricted to 1 token #27621
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||
|---|---|---|---|---|---|---|
|
|
@@ -773,10 +773,10 @@ static __global__ void mul_mat_vec_q( | |||||
| // Grid: (ceil(nrows_x / c_rows_per_block), nchannels_dst) | ||||||
| // Block: (warp_size, ncols_dst) - each warp handles one token independently. | ||||||
| // No shared memory reduction needed since each warp works alone. | ||||||
| template <ggml_type type, int c_rows_per_block> | ||||||
| template <ggml_type type, int c_rows_per_block, bool has_fusion = false> | ||||||
| __launch_bounds__(get_mmvq_mmid_max_batch_for_device<type>()*ggml_cuda_get_physical_warp_size(), 1) | ||||||
| static __global__ void mul_mat_vec_q_moe( | ||||||
| const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, | ||||||
| const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, const ggml_cuda_mm_fusion_args_device fusion, | ||||||
| float * dst_ptr, | ||||||
| const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, | ||||||
| const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, | ||||||
|
|
@@ -794,6 +794,29 @@ static __global__ void mul_mat_vec_q_moe( | |||||
|
|
||||||
| constexpr vec_dot_q_cuda_t vec_dot_q_cuda = get_vec_dot_q_cuda(type); | ||||||
|
|
||||||
| // fuse gate, bias, scales, and glu_op into the up projection | ||||||
| bool use_gate = false; | ||||||
| const void * vgate = nullptr; | ||||||
| const float * x_bias = nullptr; | ||||||
| const float * gate_bias = nullptr; | ||||||
| const float * x_scale = nullptr; | ||||||
| const float * gate_scale = nullptr; | ||||||
| ggml_glu_op active_glu = GGML_GLU_OP_SWIGLU; | ||||||
| float glu_limit = 0.0f; | ||||||
|
|
||||||
| if constexpr (has_fusion) { | ||||||
| use_gate = fusion.gate != nullptr; | ||||||
| vgate = fusion.gate; | ||||||
| x_bias = (const float *) fusion.x_bias; | ||||||
| gate_bias = (const float *) fusion.gate_bias; | ||||||
| active_glu = fusion.glu_op; | ||||||
| glu_limit = fusion.glu_limit; | ||||||
| if constexpr (type == GGML_TYPE_NVFP4) { | ||||||
| x_scale = (const float *) fusion.x_scale; | ||||||
| gate_scale = (const float *) fusion.gate_scale; | ||||||
| } | ||||||
| } | ||||||
|
|
||||||
| const uint32_t token_idx = threadIdx.y; | ||||||
| const int row0 = c_rows_per_block*blockIdx.x; | ||||||
| const int blocks_per_row_x = ncols_x / qk; | ||||||
|
|
@@ -814,6 +837,7 @@ static __global__ void mul_mat_vec_q_moe( | |||||
|
|
||||||
| // partial sum for each thread | ||||||
| float tmp[c_rows_per_block] = {0.0f}; | ||||||
| float tmp_gate[c_rows_per_block] = {0.0f}; | ||||||
|
|
||||||
| for (int kbx = threadIdx.x / (qi/vdr); kbx < blocks_per_row_x; kbx += blocks_per_iter) { | ||||||
| const int kby = kbx * (qk/QK8_1); | ||||||
|
|
@@ -822,6 +846,11 @@ static __global__ void mul_mat_vec_q_moe( | |||||
| #pragma unroll | ||||||
| for (int i = 0; i < c_rows_per_block; ++i) { | ||||||
| tmp[i] += vec_dot_q_cuda(vx, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); | ||||||
| if constexpr (has_fusion) { | ||||||
| if (use_gate) { | ||||||
| tmp_gate[i] += vec_dot_q_cuda(vgate, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); | ||||||
| } | ||||||
| } | ||||||
| } | ||||||
| } | ||||||
|
|
||||||
|
|
@@ -831,11 +860,63 @@ static __global__ void mul_mat_vec_q_moe( | |||||
| #pragma unroll | ||||||
| for (int i = 0; i < c_rows_per_block; ++i) { | ||||||
| tmp[i] = warp_reduce_sum<warp_size>(tmp[i]); | ||||||
| if constexpr (has_fusion) { | ||||||
| if (use_gate) { | ||||||
| tmp_gate[i] = warp_reduce_sum<warp_size>(tmp_gate[i]); | ||||||
| } | ||||||
| } | ||||||
| } | ||||||
|
|
||||||
| // Write results | ||||||
| if (threadIdx.x < c_rows_per_block && (c_rows_per_block == 1 || uint32_t(row0 + threadIdx.x) < nrows_x)) { | ||||||
| dst[channel_dst*stride_channel_dst + token_idx*stride_col_dst + row0 + threadIdx.x] = tmp[threadIdx.x]; | ||||||
| float result = tmp[threadIdx.x]; | ||||||
| if constexpr (has_fusion) { | ||||||
| const uint32_t bias_idx = channel_x*stride_channel_dst + row0 + threadIdx.x; | ||||||
|
|
||||||
| if constexpr (type == GGML_TYPE_NVFP4) { | ||||||
| if (x_scale) { | ||||||
| result *= x_scale[channel_x]; | ||||||
| } | ||||||
| } | ||||||
| if (x_bias) { | ||||||
| result += x_bias[bias_idx]; | ||||||
| } | ||||||
| if (use_gate) { | ||||||
| float gate_value = tmp_gate[threadIdx.x]; | ||||||
| if constexpr (type == GGML_TYPE_NVFP4) { | ||||||
| if (gate_scale) { | ||||||
| gate_value *= gate_scale[channel_x]; | ||||||
| } | ||||||
| } | ||||||
| if (gate_bias) { | ||||||
| gate_value += gate_bias[bias_idx]; | ||||||
| } | ||||||
| switch (active_glu) { | ||||||
| case GGML_GLU_OP_SWIGLU: | ||||||
| result *= ggml_cuda_op_silu_single(gate_value); | ||||||
| break; | ||||||
| case GGML_GLU_OP_GEGLU: | ||||||
| result *= ggml_cuda_op_gelu_single(gate_value); | ||||||
| break; | ||||||
| case GGML_GLU_OP_SWIGLU_OAI: | ||||||
| result = ggml_cuda_op_swiglu_oai_single(gate_value, result); | ||||||
| break; | ||||||
| case GGML_GLU_OP_SWIGLU_CLAMP: | ||||||
| result = ggml_cuda_op_swiglu_clamp_single(gate_value, result, glu_limit); | ||||||
| break; | ||||||
| default: | ||||||
| result = result * gate_value; | ||||||
|
Comment on lines
+894
to
+908
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. let's also add sqrtsoftplus used for deepseek4
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I might be missing something here, I don't see sqrtsoftplus in the glu op enum Line 623 in c060ca9
llama.cpp/ggml/src/ggml-cuda/ggml-cuda.cu Line 2825 in c060ca9
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Ah okay, then that can be done in a follow-up PR. I was thinking I already added it in #25896 |
||||||
| break; | ||||||
| } | ||||||
| } | ||||||
| } | ||||||
| dst[channel_dst*stride_channel_dst + token_idx*stride_col_dst + row0 + threadIdx.x] = result; | ||||||
| } | ||||||
|
|
||||||
| if constexpr (!has_fusion) { | ||||||
| GGML_UNUSED_VARS(use_gate, tmp_gate, vgate, x_bias, gate_bias, active_glu, glu_limit, x_scale, gate_scale); | ||||||
| } else if constexpr (type != GGML_TYPE_NVFP4) { | ||||||
| GGML_UNUSED_VARS(x_scale, gate_scale); | ||||||
| } | ||||||
| } | ||||||
|
|
||||||
|
|
@@ -885,7 +966,7 @@ static void mul_mat_vec_q_switch_fusion( | |||||
|
|
||||||
| template <ggml_type type> | ||||||
| static void mul_mat_vec_q_moe_launch( | ||||||
| const void * vx, const void * vy, const int32_t * ids, float * dst, | ||||||
| const void * vx, const void * vy, const int32_t * ids, const ggml_cuda_mm_fusion_args_device fusion, float * dst, | ||||||
| const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, | ||||||
| const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, | ||||||
| const uint32_t stride_channel_x, const uint32_t stride_channel_y, const uint32_t stride_channel_dst, | ||||||
|
|
@@ -898,11 +979,22 @@ static void mul_mat_vec_q_moe_launch( | |||||
| const dim3 block_dims(warp_size, ncols_dst); | ||||||
| const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, 0, stream); | ||||||
|
|
||||||
| ggml_cuda_kernel_launch(mul_mat_vec_q_moe<type, rows_per_block>, launch_params, | ||||||
| vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, | ||||||
| stride_row_x, stride_col_y, stride_col_dst, | ||||||
| stride_channel_x, stride_channel_y, stride_channel_dst, | ||||||
| ncols_dst, ids_stride); | ||||||
| const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr || | ||||||
| fusion.x_scale != nullptr || fusion.gate_scale != nullptr; | ||||||
|
|
||||||
| if (has_fusion) { | ||||||
| ggml_cuda_kernel_launch(mul_mat_vec_q_moe<type, rows_per_block, true>, launch_params, | ||||||
| vx, vy, ids, fusion, dst, ncols_x, nchannels_y, nrows_x, | ||||||
| stride_row_x, stride_col_y, stride_col_dst, | ||||||
| stride_channel_x, stride_channel_y, stride_channel_dst, | ||||||
| ncols_dst, ids_stride); | ||||||
| } else { | ||||||
| ggml_cuda_kernel_launch(mul_mat_vec_q_moe<type, rows_per_block, false>, launch_params, | ||||||
| vx, vy, ids, fusion, dst, ncols_x, nchannels_y, nrows_x, | ||||||
| stride_row_x, stride_col_y, stride_col_dst, | ||||||
| stride_channel_x, stride_channel_y, stride_channel_dst, | ||||||
| ncols_dst, ids_stride); | ||||||
| } | ||||||
| } | ||||||
|
|
||||||
| template <ggml_type type> | ||||||
|
|
@@ -998,7 +1090,7 @@ static void mul_mat_vec_q_switch_ncols_dst( | |||||
| if (has_ids && ncols_dst > 1) { | ||||||
| // Multi-token MUL_MAT_ID path - dedicated MoE kernel | ||||||
| mul_mat_vec_q_moe_launch<type>( | ||||||
| vx, vy, ids, dst, ncols_x, nchannels_y_fd, nrows_x, | ||||||
| vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, nrows_x, | ||||||
| stride_row_x, stride_col_y, stride_col_dst, | ||||||
| stride_channel_x, stride_channel_y, stride_channel_dst, | ||||||
| ncols_dst, ids_stride, warp_size, nchannels_dst, stream); | ||||||
|
|
@@ -1280,7 +1372,8 @@ void ggml_cuda_mul_mat_vec_q( | |||||
| ggml_cuda_mm_fusion_args_device fusion_local{}; | ||||||
|
|
||||||
| if (fusion) { | ||||||
| GGML_ASSERT( !ids || dst->ne[2] == 1); | ||||||
| const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; | ||||||
| GGML_ASSERT( !ids || dst->ne[2] <= get_mmvq_mmid_max_batch(src0->type, cc)); | ||||||
| GGML_ASSERT( ids || dst->ne[1] == 1); | ||||||
| // Scale fusion is only allowed for NVFP4 currently as the cost of checking this at run-time in the prologue is | ||||||
| // non-negligible for some models such as gpt-oss-20b | ||||||
|
|
||||||
Uh oh!
There was an error while loading. Please reload this page.