From 1c18b2cb06772bbc838d1198ec0bcc00c7894c6d Mon Sep 17 00:00:00 2001 From: Carl Philipp Klemm Date: Tue, 1 Sep 2026 13:06:10 +0200 Subject: [PATCH] feat: enhance mmq configuration for various architectures with moe_ncols_min_cc support --- ggml/src/ggml-cuda/mmq-config-ampere.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-blackwell.cuh | 1 + ggml/src/ggml-cuda/mmq-config-cdna.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-rdna2.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-rdna3.cuh | 3 +- ggml/src/ggml-cuda/mmq-config-rdna4.cuh | 3 +- ggml/src/ggml-cuda/mmq.cuh | 37 +++++++++---------- 9 files changed, 33 insertions(+), 26 deletions(-) diff --git a/ggml/src/ggml-cuda/mmq-config-ampere.cuh b/ggml/src/ggml-cuda/mmq-config-ampere.cuh index 9f9fd19738..7731d2242e 100644 --- a/ggml/src/ggml-cuda/mmq-config-ampere.cuh +++ b/ggml/src/ggml-cuda/mmq-config-ampere.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_ampere(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_VOLTA; CASE(GGML_TYPE_Q1_0, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_Q1_0, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_Q1_0, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); @@ -379,5 +380,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 1, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, false); CASE(GGML_TYPE_NVFP4, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-blackwell.cuh b/ggml/src/ggml-cuda/mmq-config-blackwell.cuh index 9fbe32b697..5e9dd340b3 100644 --- a/ggml/src/ggml-cuda/mmq-config-blackwell.cuh +++ b/ggml/src/ggml-cuda/mmq-config-blackwell.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_blackwell(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_BLACKWELL; CASE(GGML_TYPE_MXFP4, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_FP4, MMQ_ITER_K_FP4, true, true); CASE(GGML_TYPE_MXFP4, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_FP4, MMQ_ITER_K_FP4, true, true); CASE(GGML_TYPE_MXFP4, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_FP4, MMQ_ITER_K_FP4, true, true); diff --git a/ggml/src/ggml-cuda/mmq-config-cdna.cuh b/ggml/src/ggml-cuda/mmq-config-cdna.cuh index 4a8d89f720..4330f54bdf 100644 --- a/ggml/src/ggml-cuda/mmq-config-cdna.cuh +++ b/ggml/src/ggml-cuda/mmq-config-cdna.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_cdna(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_CDNA1; CASE(GGML_TYPE_Q1_0, 512, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_Q1_0, 512, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_Q1_0, 512, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); @@ -181,5 +182,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 512, 1, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, false); CASE(GGML_TYPE_NVFP4, 512, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 512, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 512, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh index 83eb7c146e..7df5b390e0 100644 --- a/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh +++ b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal_dp4a(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = 0; CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); @@ -269,5 +270,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-rdna2.cuh b/ggml/src/ggml-cuda/mmq-config-rdna2.cuh index 8324d9e1a8..59446fbf68 100644 --- a/ggml/src/ggml-cuda/mmq-config-rdna2.cuh +++ b/ggml/src/ggml-cuda/mmq-config-rdna2.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna2(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_RDNA2; CASE(GGML_TYPE_Q1_0, 256, 2, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); @@ -269,5 +270,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh b/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh index 180b2d9370..df60462b5c 100644 --- a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh +++ b/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna3_5(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_RDNA3_5; CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); @@ -286,5 +287,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-rdna3.cuh b/ggml/src/ggml-cuda/mmq-config-rdna3.cuh index 3a3ef7bd9c..b8b969fce2 100644 --- a/ggml/src/ggml-cuda/mmq-config-rdna3.cuh +++ b/ggml/src/ggml-cuda/mmq-config-rdna3.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna3(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_RDNA3; CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); @@ -270,5 +271,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq-config-rdna4.cuh b/ggml/src/ggml-cuda/mmq-config-rdna4.cuh index 9293d9d558..91d17beb21 100644 --- a/ggml/src/ggml-cuda/mmq-config-rdna4.cuh +++ b/ggml/src/ggml-cuda/mmq-config-rdna4.cuh @@ -1,4 +1,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna4(ggml_type type, int J, bool fallback) { + constexpr int moe_ncols_min_cc = GGML_CUDA_CC_RDNA4; CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); @@ -286,5 +287,5 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_NVFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, moe_ncols_min_cc, false, true); } diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index b235379dfb..60e8dfceb6 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -170,12 +170,17 @@ struct ggml_cuda_mmq_config { int J; // SRAM tile width in src1->ne[1]/dst->ne[1] direction. ggml_cuda_mmq_sram_layout sram_layout; // SRAM tile length in src0->ne[0]/src1->ne[0] direction (physical 32 bit elements). int K_vram; // VRAM tile length in src0->ne[0]/src1->ne[0] direction (logical elements). + int moe_ncols_min_cc; // Minimum architecture to use the typical routed expert width. bool stream_k; // Whether or not to use stream-k decomposition. bool fallback; // Whether a fallback for out-of-bounds check in src0->ne[1] direction is needed. constexpr __host__ __device__ ggml_cuda_mmq_config( - ggml_type type, int nthreads, int occupancy, int I, int J, ggml_cuda_mmq_sram_layout sram_layout, int K_vram, bool stream_k, bool fallback) : - type(type), nthreads(nthreads), occupancy(occupancy), I(I), J(J), sram_layout(sram_layout), K_vram(K_vram), stream_k(stream_k), fallback(fallback) {} + ggml_type type, int nthreads, int occupancy, int I, int J, ggml_cuda_mmq_sram_layout sram_layout, int K_vram, int moe_ncols_min_cc, bool stream_k, bool fallback) : + type(type), nthreads(nthreads), occupancy(occupancy), I(I), J(J), sram_layout(sram_layout), K_vram(K_vram), moe_ncols_min_cc(moe_ncols_min_cc), stream_k(stream_k), fallback(fallback) {} + + constexpr __host__ __device__ bool use_moe_ncols(const int cc) const { + return moe_ncols_min_cc != 0 && cc >= moe_ncols_min_cc; + } constexpr __device__ int rows_per_warp() const { #if defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) @@ -210,7 +215,7 @@ struct ggml_cuda_mmq_config { static_assert((I_) % 32 == 0, "bad I"); \ static_assert((J_) % 8 == 0, "bad J"); \ static_assert((K_vram_) % 256 == 0, "bad K_vram"); \ - return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), (stream_k_), (fallback_)); \ + return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), moe_ncols_min_cc, (stream_k_), (fallback_)); \ } \ #include "mmq-config-pascal-older.cuh" @@ -1472,14 +1477,6 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a ntx_fd); } -static bool mmq_use_routed_moe_ncols_picker(const int cc) { - return (GGML_CUDA_CC_IS_NVIDIA(cc) && cc >= GGML_CUDA_CC_VOLTA) || - GGML_CUDA_CC_IS_CDNA(cc) || - GGML_CUDA_CC_IS_RDNA2(cc) || - GGML_CUDA_CC_IS_RDNA3(cc) || - GGML_CUDA_CC_IS_RDNA4(cc); -} - template void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) { const int id = ggml_cuda_get_device(); @@ -1487,14 +1484,16 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, const size_t smpbo = ggml_cuda_info().devices[id].smpbo; int64_t ncols_picker = args.ncols_max; - if (args.expert_bounds != nullptr && mmq_use_routed_moe_ncols_picker(cc) && args.nchannels_x > 0) { - // In routed MoE, ncols_max is the worst-case per-expert width. Size the - // MMQ N-tile from the typical routed width while it is below the current - // architecture's max tile width. The launch grid still uses args.ncols_max. - const int J_max = ggml_cuda_mmq_get_J_max(type, fallback, cc, 128); - const int64_t ncols_typical = (args.ncols_dst + args.nchannels_x - 1) / args.nchannels_x; - if (ncols_typical >= 1 && ncols_typical < J_max && ncols_typical < ncols_picker) { - ncols_picker = ncols_typical; + if (args.expert_bounds != nullptr && args.nchannels_x > 0) { + const int J_max = ggml_cuda_mmq_get_J_max(type, fallback, cc, 128); + const ggml_cuda_mmq_config config_max = ggml_cuda_mmq_get_config(type, J_max, fallback, cc); + if (config_max.use_moe_ncols(cc)) { + // Use the typical expert width only for tile selection. + // The launch grid still uses args.ncols_max. + const int64_t ncols_typical = (args.ncols_dst + args.nchannels_x - 1) / args.nchannels_x; + if (ncols_typical >= 1 && ncols_typical < J_max && ncols_typical < ncols_picker) { + ncols_picker = ncols_typical; + } } }