]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
CUDA: fuse rms_norm + mul + rope (+ view + set_rows) (#26767)
authorRafail Giavrimis <redacted>
Sat, 8 Aug 2026 16:32:37 +0000 (17:32 +0100)
committerGitHub <redacted>
Sat, 8 Aug 2026 16:32:37 +0000 (00:32 +0800)
* CUDA: fuse rms_norm + mul + rope (+ view + set_rows)

* tests: add broadcast weight case to rms_norm_mul_rope

* CUDA: check memory ranges before rms_norm rope fusion

* CUDA: check memory ranges in rope set_rows fusion

ggml/src/ggml-cuda/ggml-cuda.cu
ggml/src/ggml-cuda/rope.cu
ggml/src/ggml-cuda/rope.cuh
tests/test-backend-ops.cpp

index 5446b31318957381fa2ff7fb399ecdfe23557b52..dec61932407097b7c82994b5e47bbab00cb6f60f 100644 (file)
@@ -2651,6 +2651,52 @@ static bool ggml_cuda_should_fuse_rope_set_rows(const ggml_tensor * rope,
     return true;
 }
 
+static bool ggml_cuda_should_fuse_rms_norm_mul_rope(const ggml_tensor * rms_norm,
+                                                    const ggml_tensor * mul,
+                                                    const ggml_tensor * rope) {
+    if (rms_norm->op != GGML_OP_RMS_NORM || mul->op != GGML_OP_MUL || rope->op != GGML_OP_ROPE) {
+        return false;
+    }
+
+    if (rms_norm->src[0]->type != GGML_TYPE_F32 || rms_norm->type != GGML_TYPE_F32 ||
+        mul->src[0]->type != GGML_TYPE_F32 || mul->src[1]->type != GGML_TYPE_F32 ||
+        mul->type != GGML_TYPE_F32 || rope->type != GGML_TYPE_F32) {
+        return false;
+    }
+
+    if (rope->src[0] != mul) {
+        return false;
+    }
+
+    //if rms norm is the B operand, then we don't handle broadcast
+    if (rms_norm == mul->src[1] && !ggml_are_same_shape(mul->src[0], rms_norm)) {
+        return false;
+    }
+
+    if (!ggml_are_same_shape(rms_norm, mul)) {
+        return false;
+    }
+
+    //rms_norm kernel assumes contiguous rows
+    if (!ggml_is_contiguous_rows(rms_norm->src[0]) ||
+        !ggml_is_contiguous_rows(mul->src[0]) || !ggml_is_contiguous_rows(mul->src[1])) {
+        return false;
+    }
+
+    // the fused kernel handles the norm/neox rope modes only
+    const int mode = ((const int32_t *) rope->op_params)[2];
+    if (mode != GGML_ROPE_TYPE_NORMAL && mode != GGML_ROPE_TYPE_NEOX) {
+        return false;
+    }
+
+    const int n_dims = ((const int32_t *) rope->op_params)[1];
+    if (n_dims % 2 != 0 || rope->src[0]->ne[0] % 2 != 0) {
+        return false;
+    }
+
+    return true;
+}
+
 // match gated_delta_net + the strided cpy that scatters its state snapshots into the cache
 // (slot i -> rollback group i, slot 0 newest), so the kernel can write them and skip the cpy.
 static int ggml_cuda_try_gdn_cache_fusion(
@@ -2980,6 +3026,36 @@ static bool ggml_cuda_can_fuse(const struct ggml_cgraph *                cgraph,
         }
     }
 
+    std::initializer_list<enum ggml_op> rms_norm_mul_rope_ops          = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE };
+    std::initializer_list<enum ggml_op> rms_norm_mul_rope_set_rows_ops = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS };
+
+    if (is_equal(rms_norm_mul_rope_set_rows_ops, ops) && ggml_can_fuse_subgraph(cgraph, node_idx, ops, { node_idx + 4 })) {
+        const ggml_tensor * rms_norm = cgraph->nodes[node_idx];
+        const ggml_tensor * mul      = cgraph->nodes[node_idx + 1];
+        const ggml_tensor * rope     = cgraph->nodes[node_idx + 2];
+        const ggml_tensor * view     = cgraph->nodes[node_idx + 3];
+        const ggml_tensor * set_rows = cgraph->nodes[node_idx + 4];
+
+        if (ggml_check_edges(cgraph, node_idx, {{1, 0, 0}, {2, 0, 1}, {3, 0, 2}, {4, 0, 3}}) &&
+            ggml_cuda_should_fuse_rms_norm_mul_rope(rms_norm, mul, rope) &&
+            ggml_cuda_should_fuse_rope_set_rows(rope, view, set_rows)) {
+            int out_nodes[] = { node_idx + 4 };
+            return ggml_cuda_check_fusion_memory_ranges(cgraph, node_idx, (int)ops.size(), out_nodes, 1);
+        }
+    }
+
+    if (is_equal(rms_norm_mul_rope_ops, ops) && ggml_can_fuse(cgraph, node_idx, ops)) {
+        const ggml_tensor * rms_norm = cgraph->nodes[node_idx];
+        const ggml_tensor * mul      = cgraph->nodes[node_idx + 1];
+        const ggml_tensor * rope     = cgraph->nodes[node_idx + 2];
+
+        if (ggml_cuda_should_fuse_rms_norm_mul_rope(rms_norm, mul, rope)) {
+            int out_nodes[] = { node_idx + 2 };
+            return ggml_cuda_check_fusion_memory_ranges(cgraph, node_idx, (int)ops.size(), out_nodes, 1);
+        }
+        return false;
+    }
+
     std::initializer_list<enum ggml_op> rope_set_rows_ops = { GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS };
 
     if (is_equal(rope_set_rows_ops, ops) && ggml_can_fuse_subgraph(cgraph, node_idx, ops, { node_idx + 2 })) {
@@ -2988,7 +3064,8 @@ static bool ggml_cuda_can_fuse(const struct ggml_cgraph *                cgraph,
         const ggml_tensor * set_rows = cgraph->nodes[node_idx + 2];
 
         if (ggml_cuda_should_fuse_rope_set_rows(rope, view, set_rows)) {
-            return true;
+            int out_nodes[] = { node_idx + 2 };
+            return ggml_cuda_check_fusion_memory_ranges(cgraph, node_idx, (int)ops.size(), out_nodes, 1);
         }
     }
 
@@ -3840,6 +3917,16 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
         return fused_node_count - 1;
     }
 
+    if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS }, {})) {
+        ggml_cuda_op_rms_norm_mul_rope_fused(*cuda_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2], cgraph->nodes[i + 4]);
+        return 4;
+    }
+
+    if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE }, {})) {
+        ggml_cuda_op_rms_norm_mul_rope_fused(*cuda_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2], nullptr);
+        return 2;
+    }
+
     if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD }, {})) {
         ggml_cuda_op_rms_norm_fused_add(*cuda_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]);
         return 2;
index e20a5cb6bed10e6fb42e2df0222038e5a4a8b6c1..504c6b818d4127548b0ce7dbb2d5458f7faebd9b 100644 (file)
@@ -670,3 +670,238 @@ void ggml_cuda_op_rope_back(ggml_backend_cuda_context & ctx, ggml_tensor * dst)
 void ggml_cuda_op_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * rope, ggml_tensor * set_rows) {
     ggml_cuda_op_rope_impl<true>(ctx, rope, set_rows);
 }
+
+// fused RMS_NORM + MUL + ROPE (+ VIEW + SET_ROWS)
+// 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 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,
+        const float * mul,
+        const int64_t mul_s01, const int64_t mul_s02, const int64_t mul_s03,
+        const uint3 mul_ncols_packed, const uint3 mul_nrows_packed,
+        const uint3 mul_nchannels_packed, const uint3 mul_nsamples_packed,
+        const int n_dims, const int32_t * pos,
+        const float freq_scale, const float ext_factor, const float attn_factor,
+        const rope_corr_dims corr_dims, const float theta_scale,
+        const float * freq_factors,
+        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;
+    }
+
+    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)];
+
+        if (i0 >= n_dims) {
+            dst[ix0] = ggml_cuda_cast<D>(x0);
+            dst[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);
+
+        dst[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
+        dst[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
+    }
+}
+
+template <typename D>
+static void rms_norm_mul_rope_cuda(
+        const float * x, D * dst,
+        const int ncols, const int nrows, 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,
+        const float * mul,
+        const int64_t mul_s01, const int64_t mul_s02, const int64_t mul_s03,
+        const uint32_t mul_ncols, const uint32_t mul_nrows,
+        const uint32_t mul_nchannels, const uint32_t mul_nsamples,
+        const int n_dims, const int32_t * pos,
+        const float freq_scale, const float freq_base, const float ext_factor, const float attn_factor,
+        const rope_corr_dims corr_dims,
+        const float * freq_factors,
+        const int64_t * row_indices, const int set_rows_stride,
+        const bool is_neox, cudaStream_t stream) {
+    GGML_ASSERT(ncols % 2 == 0);
+
+    const dim3 blocks_num(nrows, nchannels, nsamples);
+
+    const float theta_scale = powf(freq_base, -2.0f/n_dims);
+
+    const uint3 mul_ncols_packed     = init_fastdiv_values(mul_ncols);
+    const uint3 mul_nrows_packed     = init_fastdiv_values(mul_nrows);
+    const uint3 mul_nchannels_packed = init_fastdiv_values(mul_nchannels);
+    const uint3 mul_nsamples_packed  = init_fastdiv_values(mul_nsamples);
+
+    if (ncols < 1024) {
+        const dim3 block_dims(256, 1, 1);
+        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,
+                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,
+                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 {
+        const dim3 block_dims(1024, 1, 1);
+        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,
+                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,
+                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);
+        }
+    }
+}
+
+void ggml_cuda_op_rms_norm_mul_rope_fused(ggml_backend_cuda_context & ctx,
+        ggml_tensor * rms_norm, ggml_tensor * mul, ggml_tensor * rope, ggml_tensor * set_rows) {
+    const ggml_tensor * x = rms_norm->src[0];
+    const ggml_tensor * mul_src = mul->src[0] == rms_norm ? mul->src[1] : mul->src[0];
+
+    float eps = 0.0f;
+    memcpy(&eps, rms_norm->op_params, sizeof(float));
+    GGML_ASSERT(eps >= 0.0f);
+
+    GGML_ASSERT(x->type == GGML_TYPE_F32);
+    GGML_ASSERT(mul_src->type == GGML_TYPE_F32);
+    GGML_ASSERT(rope->type == GGML_TYPE_F32);
+
+    void *          dst_d           = rope->data;
+    ggml_type       dst_type        = rope->type;
+    const int64_t * row_indices     = nullptr;
+    int             set_rows_stride = 0;
+
+    if (set_rows != nullptr) {
+        dst_d           = set_rows->data;
+        dst_type        = set_rows->type;
+        row_indices     = (const int64_t *) set_rows->src[1]->data;
+        set_rows_stride = set_rows->nb[1] / ggml_type_size(set_rows->type);
+    }
+
+    const int n_dims     = ((const int32_t *) rope->op_params)[1];
+    const int mode       = ((const int32_t *) rope->op_params)[2];
+    const int n_ctx_orig = ((const int32_t *) rope->op_params)[4];
+
+    float freq_base;
+    float freq_scale;
+    float ext_factor;
+    float attn_factor;
+    float beta_fast;
+    float beta_slow;
+
+    memcpy(&freq_base,   (const int32_t *) rope->op_params +  5, sizeof(float));
+    memcpy(&freq_scale,  (const int32_t *) rope->op_params +  6, sizeof(float));
+    memcpy(&ext_factor,  (const int32_t *) rope->op_params +  7, sizeof(float));
+    memcpy(&attn_factor, (const int32_t *) rope->op_params +  8, sizeof(float));
+    memcpy(&beta_fast,   (const int32_t *) rope->op_params +  9, sizeof(float));
+    memcpy(&beta_slow,   (const int32_t *) rope->op_params + 10, sizeof(float));
+
+    const bool is_neox = mode & GGML_ROPE_TYPE_NEOX;
+
+    const int32_t * pos = (const int32_t *) rope->src[1]->data;
+
+    const float * freq_factors = rope->src[2] != nullptr ? (const float *) rope->src[2]->data : nullptr;
+
+    rope_corr_dims corr_dims;
+    ggml_rope_yarn_corr_dims(n_dims, n_ctx_orig, freq_base, beta_fast, beta_slow, corr_dims.v);
+
+    const size_t ts0 = ggml_type_size(x->type);
+    GGML_ASSERT(x->nb[0] == ts0);
+    const int64_t s01 = x->nb[1] / ts0;
+    const int64_t s02 = x->nb[2] / ts0;
+    const int64_t s03 = x->nb[3] / ts0;
+
+    const size_t ts_mul = ggml_type_size(mul_src->type);
+    GGML_ASSERT(mul_src->nb[0] == ts_mul);
+    const int64_t mul_s01 = mul_src->nb[1] / ts_mul;
+    const int64_t mul_s02 = mul_src->nb[2] / ts_mul;
+    const int64_t mul_s03 = mul_src->nb[3] / ts_mul;
+
+    const size_t ts_dst = ggml_type_size(rope->type);
+    const int64_t s1 = rope->nb[1] / ts_dst;
+    const int64_t s2 = rope->nb[2] / ts_dst;
+    const int64_t s3 = rope->nb[3] / ts_dst;
+
+    cudaStream_t stream = ctx.stream();
+
+    if (dst_type == GGML_TYPE_F32) {
+        rms_norm_mul_rope_cuda((const float *) x->data, (float *) dst_d,
+            x->ne[0], x->ne[1], x->ne[2], x->ne[3], s01, s02, s03, s1, s2, s3, eps,
+            (const float *) mul_src->data, mul_s01, mul_s02, mul_s03,
+            mul_src->ne[0], mul_src->ne[1], mul_src->ne[2], mul_src->ne[3],
+            n_dims, pos, freq_scale, freq_base, ext_factor, attn_factor, corr_dims,
+            freq_factors, row_indices, set_rows_stride, is_neox, stream);
+    } else if (dst_type == GGML_TYPE_F16) {
+        rms_norm_mul_rope_cuda((const float *) x->data, (half *) dst_d,
+            x->ne[0], x->ne[1], x->ne[2], x->ne[3], s01, s02, s03, s1, s2, s3, eps,
+            (const float *) mul_src->data, mul_s01, mul_s02, mul_s03,
+            mul_src->ne[0], mul_src->ne[1], mul_src->ne[2], mul_src->ne[3],
+            n_dims, pos, freq_scale, freq_base, ext_factor, attn_factor, corr_dims,
+            freq_factors, row_indices, set_rows_stride, is_neox, stream);
+    } else {
+        GGML_ABORT("fatal error");
+    }
+}
index 72af086cd1b42573975b095c45e7b272d362805d..7ce2d71c508ee880cf96806e93b2988c75c84978 100644 (file)
@@ -7,3 +7,5 @@ void ggml_cuda_op_rope(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
 void ggml_cuda_op_rope_back(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
 
 void ggml_cuda_op_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * set_rows);
+
+void ggml_cuda_op_rms_norm_mul_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * rms_norm, ggml_tensor * mul, ggml_tensor * rope, ggml_tensor * set_rows);
index 6688debd4adc282d98aa2d24faede1e30bb6bd9c..14a2340604afdf269598a0638071fd1f9e0fcf8c 100644 (file)
@@ -2584,6 +2584,7 @@ struct test_rms_norm_mul_rope : public test_case {
     const float eps;
     const bool multi_add; // test a sequence of adds feeding into rms_norm
     const bool set_rows;
+    const bool broadcast; // multiply by a 1D [ne0] weight, as model norm weights are
     int mode;
 
     std::string op_desc(ggml_tensor * t) override {
@@ -2594,12 +2595,12 @@ struct test_rms_norm_mul_rope : public test_case {
     bool run_whole_graph() override { return true; }
 
     std::string vars() override {
-        return VARS_TO_STR5(ne, eps, multi_add, set_rows, mode);
+        return VARS_TO_STR6(ne, eps, multi_add, set_rows, broadcast, mode);
     }
 
     test_rms_norm_mul_rope(std::array<int64_t, 4> ne, float eps = 1e-6f, bool multi_add = false,
-                           bool set_rows = false, int mode = GGML_ROPE_TYPE_NORMAL)
-        : ne(ne), eps(eps), multi_add(multi_add), set_rows(set_rows), mode(mode) {}
+                           bool set_rows = false, bool broadcast = false, int mode = GGML_ROPE_TYPE_NORMAL)
+        : ne(ne), eps(eps), multi_add(multi_add), set_rows(set_rows), broadcast(broadcast), mode(mode) {}
 
     ggml_tensor * build_graph(ggml_context * ctx) override {
         ggml_tensor * a = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, ne[0], ne[1], ne[2], 1);
@@ -2610,7 +2611,9 @@ struct test_rms_norm_mul_rope : public test_case {
             a = ggml_add(ctx, ggml_add(ctx, a, b), c);
         }
 
-        a = ggml_mul(ctx, ggml_rms_norm(ctx, a, eps), b);
+        ggml_tensor * w = broadcast ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, ne[0]) : b;
+
+        a = ggml_mul(ctx, ggml_rms_norm(ctx, a, eps), w);
 
         ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, ne[2]);
 
@@ -8756,16 +8759,18 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
 
     for (auto multi_add : {false, true}) {
         for (auto set_rows : {false, true}) {
-            for (auto rope : {GGML_ROPE_TYPE_NORMAL, GGML_ROPE_TYPE_NEOX}) {
-                test_cases.emplace_back(new test_rms_norm_mul_rope({768, 1, 1, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({768, 3, 1, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({768, 3, 5, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({128, 32, 2, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({128, 4, 2, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({128, 32, 50, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({128, 4, 50, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({8192, 2, 2, 1}, 1e-6f, multi_add, set_rows, rope));
-                test_cases.emplace_back(new test_rms_norm_mul_rope({8192, 2, 2, 1}, 1e-6f, multi_add, set_rows, rope));
+            for (auto broadcast : {false, true}) {
+                for (auto rope : {GGML_ROPE_TYPE_NORMAL, GGML_ROPE_TYPE_NEOX}) {
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({768, 1, 1, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({768, 3, 1, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({768, 3, 5, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({128, 32, 2, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({128, 4, 2, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({128, 32, 50, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({128, 4, 50, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({8192, 2, 2, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                    test_cases.emplace_back(new test_rms_norm_mul_rope({8192, 2, 2, 1}, 1e-6f, multi_add, set_rows, broadcast, rope));
+                }
             }
         }
     }