From 202b09c7add4bb47acbac017fc3b5965f03ce1df Mon Sep 17 00:00:00 2001 From: bri-prism <288398250+bri-prism@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:52:24 -0700 Subject: [PATCH] cuda: bound mmvq rows by the per-slot stride for MUL_MAT_ID With ids, stride_col_dst is the token stride, not the row count. When rows_per_cuda_block > 1 (small_k) and the row count is not a multiple of it, the tail block wrote past the end of its expert slot into the next one, racing with that slot's block. test-backend-ops MUL_MAT_ID with ptq1_0, m=70, n=1, k=2048 failed in 8 of 10 runs on an RTX 4090. --- ggml/src/ggml-cuda/mmvq.cu | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index bf51b61e17b1..f8736382c54e 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -597,6 +597,9 @@ static __global__ void mul_mat_vec_q( const uint32_t sample_x = fastdiv(sample_dst, sample_ratio); const uint32_t sample_y = sample_dst; + // dst rows are contiguous; with ids the per-slot row stride is stride_channel_dst, not stride_col_dst + const uint32_t nrows_dst = ids ? stride_channel_dst : stride_col_dst; + constexpr bool use_gate = has_gate; bool use_bias = false; bool use_gate_bias = false; @@ -634,7 +637,7 @@ static __global__ void mul_mat_vec_q( // 2. load only on threads that won't die after partial sum calculation const uint32_t channel_bias = ids ? channel_x : channel_dst; if (threadIdx.x < rows_per_cuda_block && threadIdx.y == 0 && - (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < stride_col_dst)) { + (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < nrows_dst)) { if (use_bias) { x_bias = x_bias + sample_dst * stride_sample_dst + channel_bias * stride_channel_dst + row0; #pragma unroll @@ -805,7 +808,7 @@ static __global__ void mul_mat_vec_q( } } - if (threadIdx.x == i && (rows_per_cuda_block == 1 || uint32_t(row0 + i) < stride_col_dst)) { + if (threadIdx.x == i && (rows_per_cuda_block == 1 || uint32_t(row0 + i) < nrows_dst)) { float result = tmp[j][i]; if constexpr (has_fusion) { if constexpr (type == GGML_TYPE_NVFP4) {