Commit fda186661 for llama.cpp

commit fda186661304735bb9027abd472c1403a9902081
Author: Eden <aa9736195201@gmail.com>
Date:   Thu Oct 8 14:46:21 2026 +0800

    CUDA: fix norm family kernels when ne[2]/ne[3] exceed grid dim limits (#28175)

diff --git a/ggml/src/ggml-cuda/norm.cu b/ggml/src/ggml-cuda/norm.cu
index 5543307be..d804b933b 100644
--- a/ggml/src/ggml-cuda/norm.cu
+++ b/ggml/src/ggml-cuda/norm.cu
@@ -3,38 +3,46 @@

 template <int block_size>
 static __global__ void norm_f32(
-        const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
-        const int64_t stride_sample, const float eps) {
-    const int nrows     = gridDim.x;
-    const int nchannels = gridDim.y;
+        const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
+        const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
+    const int nrows = gridDim.x;
+    const int row   = blockIdx.x;
+    const int tid   = threadIdx.x;

-    const int row       = blockIdx.x;
-    const int channel   = blockIdx.y;
-    const int sample    = blockIdx.z;
-    const int tid       = threadIdx.x;
+    extern __shared__ float2 s_sum2[];

-    x   += sample*stride_sample + channel*stride_channel + row*stride_row;
-    dst += ((sample*nchannels + channel)*nrows + row)*ncols;
+    ggml_cuda_pdl_sync();

-    float2 mean_var = make_float2(0.0f, 0.0f);
+    // grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
+    for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
+        for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
+            const float * xc   = x   + sample*stride_sample + channel*stride_channel + row*stride_row;
+            float       * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;

-    ggml_cuda_pdl_sync();
-    for (int col = tid; col < ncols; col += block_size) {
-        const float xi = x[col];
-        mean_var.x += xi;
-        mean_var.y += xi * xi;
-    }
+            float2 mean_var = make_float2(0.0f, 0.0f);

-    // sum up partial sums
-    extern __shared__ float2 s_sum2[];
-    mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);
+            for (int col = tid; col < ncols; col += block_size) {
+                const float xi = xc[col];
+                mean_var.x += xi;
+                mean_var.y += xi * xi;
+            }

-    const float mean = mean_var.x / ncols;
-    const float var = mean_var.y / ncols - mean * mean;
-    const float inv_std = rsqrtf(var + eps);
+            // sum up partial sums
+            mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);

-    for (int col = tid; col < ncols; col += block_size) {
-        dst[col] = (x[col] - mean) * inv_std;
+            const float mean = mean_var.x / ncols;
+            const float var = mean_var.y / ncols - mean * mean;
+            const float inv_std = rsqrtf(var + eps);
+
+            for (int col = tid; col < ncols; col += block_size) {
+                dstc[col] = (xc[col] - mean) * inv_std;
+            }
+
+            if constexpr (block_size > WARP_SIZE) {
+                // sync is needed as we reuse s_sum2 across block_reduce invocations, see #26385
+                __syncthreads();
+            }
+        }
     }
 }

@@ -77,6 +85,8 @@ template <int block_size, bool do_multiply = false, bool do_add = false, bool do
 static __global__ void rms_norm_f32(const float * x,
                                     float *       dst,
                                     const int     ncols,
+                                    const int     nchannels,
+                                    const int     nsamples,
                                     const int64_t stride_row,
                                     const int64_t stride_channel,
                                     const int64_t stride_sample,
@@ -99,61 +109,71 @@ static __global__ void rms_norm_f32(const float * x,
                                     const uint3   add_nsamples_packed  = make_uint3(0, 0, 0),
                                     const float   scale_out            = 1.0f) {
     ggml_cuda_pdl_lc();
-    const int nrows     = gridDim.x;
-    const int nchannels = gridDim.y;
-
-    const int row       = blockIdx.x;
-    const int channel   = blockIdx.y;
-    const int sample    = blockIdx.z;
-    const int tid       = threadIdx.x;
+    const int nrows = gridDim.x;
+    const int row   = blockIdx.x;
+    const int tid   = threadIdx.x;

     static_assert(!do_add || do_multiply, "fusing add is not supported without multiplying");
     static_assert(!do_scale || !do_multiply, "fusing scale is not supported with multiplying");

-    x   += sample*stride_sample + channel*stride_channel + row*stride_row;
-    dst += ((sample*nchannels + channel)*nrows + row)*ncols;
-
-    if constexpr (do_multiply) {
-        const uint32_t mul_row     = fastmodulo(row, mul_nrows_packed);
-        const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
-        const uint32_t mul_sample  = fastmodulo(sample, mul_nsamples_packed);
-        mul += mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
-    }
-
-    if constexpr (do_add) {
-        const int add_row     = fastmodulo(row, add_nrows_packed);
-        const int add_channel = fastmodulo(channel, add_nchannels_packed);
-        const int add_sample  = fastmodulo(sample, add_nsamples_packed);
-        add += add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
-    }
-
-    float tmp = 0.0f; // partial sum for thread in warp
-
-    ggml_cuda_pdl_sync();
-    for (int col = tid; col < ncols; col += block_size) {
-        const float xi = x[col];
-        tmp += xi * xi;
-    }
-
-    // sum up partial sums
     extern __shared__ float s_sum[];
-    tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);

-    const float mean = tmp / ncols;
-    const float scale = rsqrtf(mean + eps);
+    ggml_cuda_pdl_sync();

-    for (int col = tid; col < ncols; col += block_size) {
-        if constexpr (do_multiply && do_add) {
-            const int mul_col = fastmodulo(col, mul_ncols_packed);
-            const int add_col = fastmodulo(col, add_ncols_packed);
-            dst[col]          = scale * x[col] * mul[mul_col] + add[add_col];
-        } else if constexpr (do_multiply) {
-            const int mul_col = fastmodulo(col, mul_ncols_packed);
-            dst[col]          = scale * x[col] * mul[mul_col];
-        } else if constexpr (do_scale) {
-            dst[col] = scale_out * (scale * x[col]);
-        } else {
-            dst[col] = scale * x[col];
+    // grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
+    for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
+        for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
+            const float * xc   = x   + sample*stride_sample + channel*stride_channel + row*stride_row;
+            float       * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
+
+            [[maybe_unused]] const float * mulc = nullptr;
+            if constexpr (do_multiply) {
+                const uint32_t mul_row     = fastmodulo(row, mul_nrows_packed);
+                const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
+                const uint32_t mul_sample  = fastmodulo(sample, mul_nsamples_packed);
+                mulc = mul + mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
+            }
+
+            [[maybe_unused]] const float * addc = nullptr;
+            if constexpr (do_add) {
+                const int add_row     = fastmodulo(row, add_nrows_packed);
+                const int add_channel = fastmodulo(channel, add_nchannels_packed);
+                const int add_sample  = fastmodulo(sample, add_nsamples_packed);
+                addc = add + add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
+            }
+
+            float tmp = 0.0f; // partial sum for thread in warp
+
+            for (int col = tid; col < ncols; col += block_size) {
+                const float xi = xc[col];
+                tmp += xi * xi;
+            }
+
+            // sum up partial sums
+            tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
+
+            const float mean = tmp / ncols;
+            const float scale = rsqrtf(mean + eps);
+
+            for (int col = tid; col < ncols; col += block_size) {
+                if constexpr (do_multiply && do_add) {
+                    const int mul_col = fastmodulo(col, mul_ncols_packed);
+                    const int add_col = fastmodulo(col, add_ncols_packed);
+                    dstc[col]         = scale * xc[col] * mulc[mul_col] + addc[add_col];
+                } else if constexpr (do_multiply) {
+                    const int mul_col = fastmodulo(col, mul_ncols_packed);
+                    dstc[col]         = scale * xc[col] * mulc[mul_col];
+                } else if constexpr (do_scale) {
+                    dstc[col] = scale_out * (scale * xc[col]);
+                } else {
+                    dstc[col] = scale * xc[col];
+                }
+            }
+
+            if constexpr (block_size > WARP_SIZE) {
+                // sync is needed as we reuse s_sum across block_reduce invocations, see #26385
+                __syncthreads();
+            }
         }
     }
 }
@@ -247,50 +267,57 @@ static __global__ void rms_norm_back_f32(

 template <int block_size>
 static __global__ void l2_norm_f32(
-        const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
-        const int64_t stride_sample, const float eps) {
-    const int nrows     = gridDim.x;
-    const int nchannels = gridDim.y;
+        const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
+        const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
+    const int nrows = gridDim.x;
+    const int row   = blockIdx.x;
+    const int tid   = threadIdx.x;

-    const int row       = blockIdx.x;
-    const int channel   = blockIdx.y;
-    const int sample    = blockIdx.z;
-    const int tid       = threadIdx.x;
+    extern __shared__ float s_sum[];

-    x   += sample*stride_sample + channel*stride_channel + row*stride_row;
-    dst += ((sample*nchannels + channel)*nrows + row)*ncols;
+    ggml_cuda_pdl_sync();

-    float tmp = 0.0f; // partial sum for thread in warp
+    // grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
+    for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
+        for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
+            const float * xc   = x   + sample*stride_sample + channel*stride_channel + row*stride_row;
+            float       * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;

-    ggml_cuda_pdl_sync();
-    for (int col = tid; col < ncols; col += block_size) {
-        const float xi = x[col];
-        tmp += xi * xi;
-    }
+            float tmp = 0.0f; // partial sum for thread in warp

-    // sum up partial sums
-    extern __shared__ float s_sum[];
-    tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
-    ggml_cuda_pdl_lc();
+            for (int col = tid; col < ncols; col += block_size) {
+                const float xi = xc[col];
+                tmp += xi * xi;
+            }

-    // from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
-    const float scale = rsqrtf(fmaxf(tmp, eps * eps));
+            // sum up partial sums
+            tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);

-    for (int col = tid; col < ncols; col += block_size) {
-        dst[col] = scale * x[col];
+            // from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
+            const float scale = rsqrtf(fmaxf(tmp, eps * eps));
+
+            for (int col = tid; col < ncols; col += block_size) {
+                dstc[col] = scale * xc[col];
+            }
+
+            if constexpr (block_size > WARP_SIZE) {
+                // sync is needed as we reuse s_sum across block_reduce invocations, see #26385
+                __syncthreads();
+            }
+        }
     }
 }

 static void norm_f32_cuda(
         const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
         const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
-    const dim3 blocks_num(nrows, nchannels, nsamples);
+    const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
     if (ncols < 1024) {
         const dim3 block_dims(WARP_SIZE, 1, 1);
-        norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
+        norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
     } else {
         const dim3 block_dims(1024, 1, 1);
-        norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
+        norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
     }
 }

@@ -310,19 +337,19 @@ static void rms_norm_f32_cuda(
         const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
         const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream,
         const float scale_out = 1.0f) {
-    const dim3 blocks_num(nrows, nchannels, nsamples);
+    const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
     if (ncols < 1024) {
         const dim3 block_dims(256, 1, 1);
         const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
         ggml_cuda_kernel_launch(rms_norm_f32<256, false, false, do_scale>, launch_params,
-            x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
+            x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
         // underlying cudaLaunchKernelEx does not support default params
         nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
         nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
     } else {
         const dim3 block_dims(1024, 1, 1);
         const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
-        ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
+        ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
         // underlying cudaLaunchKernelEx does not support default params
         nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
         nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
@@ -356,7 +383,7 @@ static void rms_norm_mul_f32_cuda(const float *  x,
                                   const uint32_t add_nsamples,
                                   const float    eps,
                                   cudaStream_t   stream) {
-    const dim3 blocks_num(nrows, nchannels, nsamples);
+    const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
     if (mul == nullptr) {
         rms_norm_f32_cuda(x, dst, ncols, nrows, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, stream);
         return;
@@ -370,7 +397,7 @@ static void rms_norm_mul_f32_cuda(const float *  x,
             const dim3 block_dims(256, 1, 1);
             const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
             ggml_cuda_kernel_launch(rms_norm_f32<256, true>, launch_params,
-                x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
+                x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
                 mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 // underlying cudaLaunchKernelEx does not support default params
             nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
@@ -378,7 +405,7 @@ static void rms_norm_mul_f32_cuda(const float *  x,
             const dim3 block_dims(1024, 1, 1);
             const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
             ggml_cuda_kernel_launch(rms_norm_f32<1024, true>, launch_params,
-                x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
+                x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
                 mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 // underlying cudaLaunchKernelEx does not support default params
             nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
@@ -397,7 +424,7 @@ static void rms_norm_mul_f32_cuda(const float *  x,
             const dim3 block_dims(256, 1, 1);
             const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims,block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
             ggml_cuda_kernel_launch(rms_norm_f32<256, true, true>, launch_params,
-                x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
+                x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
                 mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
                 add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
                 add_nchannels_packed, add_nsamples_packed, 1.0f);
@@ -405,7 +432,7 @@ static void rms_norm_mul_f32_cuda(const float *  x,
             const dim3 block_dims(1024, 1, 1);
             const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
             ggml_cuda_kernel_launch(rms_norm_f32<1024, true, true>, launch_params,
-                x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
+                x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
                 mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
                 add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
                 add_nchannels_packed, add_nsamples_packed, 1.0f);
@@ -426,15 +453,15 @@ static void rms_norm_back_f32_cuda(const float * grad, const float * xf, float *
 static void l2_norm_f32_cuda(
         const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
         const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
-    const dim3 blocks_num(nrows, nchannels, nsamples);
+    const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
     if (ncols < 1024) {
         const dim3 block_dims(WARP_SIZE, 1, 1);
         const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, 0, stream};
-        ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
+        ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
     } else {
         const dim3 block_dims(1024, 1, 1);
         const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
-        ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
+        ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
     }
 }

diff --git a/ggml/src/ggml-cuda/rope.cu b/ggml/src/ggml-cuda/rope.cu
index e546fb655..b4fca778f 100644
--- a/ggml/src/ggml-cuda/rope.cu
+++ b/ggml/src/ggml-cuda/rope.cu
@@ -709,7 +709,7 @@ void ggml_cuda_op_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * rope
 // one block per row: block_reduce gives the norm scale, then each thread applies mul and rope to the elements it owns
 template <int block_size, bool has_ff, typename D>
 static __global__ void rms_norm_mul_rope_f32(
-        const float * x, D * dst, const int ncols,
+        const float * x, D * dst, const int ncols, const int nchannels, const int nsamples,
         const int64_t s01, const int64_t s02, const int64_t s03,
         const int64_t s1, const int64_t s2, const int64_t s3,
         const float eps,
@@ -724,66 +724,76 @@ static __global__ void rms_norm_mul_rope_f32(
         const int64_t * row_indices, const int set_rows_stride,
         const bool is_neox) {
     ggml_cuda_pdl_lc();
-    const int row     = blockIdx.x;
-    const int channel = blockIdx.y;
-    const int sample  = blockIdx.z;
-    const int tid     = threadIdx.x;
-
-    x += sample*s03 + channel*s02 + row*s01;
-
-    const uint32_t mul_row     = fastmodulo(row,     mul_nrows_packed);
-    const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
-    const uint32_t mul_sample  = fastmodulo(sample,  mul_nsamples_packed);
-    mul += mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
-
-    float tmp = 0.0f;
-
-    ggml_cuda_pdl_sync();
-    for (int col = tid; col < ncols; col += block_size) {
-        const float xi = x[col];
-        tmp += xi * xi;
-    }
+    const int row = blockIdx.x;
+    const int tid = threadIdx.x;

     extern __shared__ float s_sum[];
-    tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
-
-    const float scale = rsqrtf(tmp/ncols + eps);
-
-    int64_t idst = sample*s3 + channel*s2 + row*s1;
-    if (set_rows_stride != 0) {
-        idst = row*s1 + row_indices[channel]*set_rows_stride;
-    }
-    dst += idst;
-
-    for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
-        int ix0;
-        int ix1;
-        if (is_neox && i0 < n_dims) {
-            ix0 = i0/2;
-            ix1 = i0/2 + n_dims/2;
-        } else {
-            ix0 = i0 + 0;
-            ix1 = i0 + 1;
-        }

-        const float x0 = scale * x[ix0] * mul[fastmodulo(ix0, mul_ncols_packed)];
-        const float x1 = scale * x[ix1] * mul[fastmodulo(ix1, mul_ncols_packed)];
+    ggml_cuda_pdl_sync();

-        if (i0 >= n_dims) {
-            dst[ix0] = ggml_cuda_cast<D>(x0);
-            dst[ix1] = ggml_cuda_cast<D>(x1);
-            continue;
+    // grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
+    for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
+        for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
+            const float * xc = x + sample*s03 + channel*s02 + row*s01;
+
+            const uint32_t mul_row     = fastmodulo(row,     mul_nrows_packed);
+            const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
+            const uint32_t mul_sample  = fastmodulo(sample,  mul_nsamples_packed);
+            const float * mulc = mul + mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
+
+            float tmp = 0.0f;
+
+            for (int col = tid; col < ncols; col += block_size) {
+                const float xi = xc[col];
+                tmp += xi * xi;
+            }
+
+            tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
+
+            const float scale = rsqrtf(tmp/ncols + eps);
+
+            int64_t idst = sample*s3 + channel*s2 + row*s1;
+            if (set_rows_stride != 0) {
+                idst = row*s1 + row_indices[channel]*set_rows_stride;
+            }
+            D * dstc = dst + idst;
+
+            for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
+                int ix0;
+                int ix1;
+                if (is_neox && i0 < n_dims) {
+                    ix0 = i0/2;
+                    ix1 = i0/2 + n_dims/2;
+                } else {
+                    ix0 = i0 + 0;
+                    ix1 = i0 + 1;
+                }
+
+                const float x0 = scale * xc[ix0] * mulc[fastmodulo(ix0, mul_ncols_packed)];
+                const float x1 = scale * xc[ix1] * mulc[fastmodulo(ix1, mul_ncols_packed)];
+
+                if (i0 >= n_dims) {
+                    dstc[ix0] = ggml_cuda_cast<D>(x0);
+                    dstc[ix1] = ggml_cuda_cast<D>(x1);
+                    continue;
+                }
+
+                const float theta_base  = pos[channel]*powf(theta_scale, i0/2.0f);
+                const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
+
+                float cos_theta;
+                float sin_theta;
+                rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
+
+                dstc[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
+                dstc[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
+            }
+
+            if constexpr (block_size > WARP_SIZE) {
+                // sync is needed as we reuse s_sum across block_reduce invocations, see #26385
+                __syncthreads();
+            }
         }
-
-        const float theta_base  = pos[channel]*powf(theta_scale, i0/2.0f);
-        const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
-
-        float cos_theta;
-        float sin_theta;
-        rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
-
-        dst[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
-        dst[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
     }
 }

@@ -806,7 +816,7 @@ static void rms_norm_mul_rope_cuda(
         const bool is_neox, cudaStream_t stream) {
     GGML_ASSERT(ncols % 2 == 0);

-    const dim3 blocks_num(nrows, nchannels, nsamples);
+    const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));

     const float theta_scale = powf(freq_base, -2.0f/n_dims);

@@ -820,13 +830,13 @@ static void rms_norm_mul_rope_cuda(
         const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
         if (freq_factors == nullptr) {
             ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, false, D>, launch_params,
-                x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
+                x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
                 mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
                 freq_factors, row_indices, set_rows_stride, is_neox);
         } else {
             ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, true, D>, launch_params,
-                x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
+                x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
                 mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
                 freq_factors, row_indices, set_rows_stride, is_neox);
@@ -836,13 +846,13 @@ static void rms_norm_mul_rope_cuda(
         const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
         if (freq_factors == nullptr) {
             ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, false, D>, launch_params,
-                x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
+                x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
                 mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
                 freq_factors, row_indices, set_rows_stride, is_neox);
         } else {
             ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, true, D>, launch_params,
-                x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
+                x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
                 mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
                 n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
                 freq_factors, row_indices, set_rows_stride, is_neox);
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index f31016bbb..a38766cba 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -2826,7 +2826,14 @@ struct test_rms_norm_mul_rope : public test_case {
     }

     double max_nmse_err() override {
-        return ne[0] == 8192 ? 5e-6 : test_case::max_nmse_err();
+        if (ne[0] == 8192) {
+            return 5e-6;
+        }
+        // large positions amplify the difference between the GPU and CPU trig functions
+        if (ne[2] > 8192) {
+            return 1e-5;
+        }
+        return test_case::max_nmse_err();
     }
 };

@@ -10273,6 +10280,12 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {

     // in-place tests
     test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, {64, 5, 4, 3}, false, 1e-6f, true));
+    // shapes below exceed the CUDA gridDim.y/gridDim.z limit of 65535 (#27901)
+    test_cases.emplace_back(new test_norm        (GGML_TYPE_F32, {4, 1, 65536, 1}, false, 1e-6f));
+    test_cases.emplace_back(new test_rms_norm    (GGML_TYPE_F32, {4, 1, 65536, 1}, false, 1e-6f, false));
+    test_cases.emplace_back(new test_rms_norm    (GGML_TYPE_F32, {4, 1, 1, 65536}, false, 1e-6f, false));
+    test_cases.emplace_back(new test_l2_norm     (GGML_TYPE_F32, {4, 1, 65536, 1}, 1e-12f, false, false));
+    test_cases.emplace_back(new test_rms_norm_mul_add(GGML_TYPE_F32, {4, 1, 65536, 1}, 1e-6f, false, false));

     for (ggml_type set_rows_type : { GGML_TYPE_F32, GGML_TYPE_F16 }) {
         test_cases.emplace_back(new test_rms_norm_mul_rope({ 256, 1, 1, 1 }, 1e-6f, false, true, false, GGML_ROPE_TYPE_NORMAL, false, false, set_rows_type));
@@ -10331,6 +10344,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
             }
         }
     }
+    // ne[2] > 65535 exceeds the CUDA gridDim.y limit (#27901)
+    test_cases.emplace_back(new test_rms_norm_mul_rope({4, 1, 65536, 1}, 1e-6f, false, false, false, GGML_ROPE_TYPE_NORMAL));
+    test_cases.emplace_back(new test_rms_norm_mul_rope({4, 1, 65536, 1}, 1e-6f, false, true,  false, GGML_ROPE_TYPE_NEOX));
     for (int64_t d_conv : {3, 4, 9}) {
         for (int64_t d_inner: {1024, 1536, 2048}) {
             test_cases.emplace_back(new test_ssm_conv(GGML_TYPE_F32, {d_conv, d_inner, 1, 1}, {d_conv, d_inner, 1, 1}));