From 24da9240655ac2edc8cfc547d6447eeb0b0364ec Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sat, 26 Sep 2026 17:04:35 +0100
Subject: [PATCH 11/18] vulkan : use a subgroup reduction in rms_norm and keep
 the row in registers

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          | 10 +++++--
 .../ggml-vulkan/vulkan-shaders/rms_norm.comp  | 30 +++++++++++++++++--
 .../vulkan-shaders/vulkan-shaders-gen.cpp     |  1 +
 3 files changed, 36 insertions(+), 5 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index ae295b3de..13b7f3268 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -3256,8 +3256,14 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
     ggml_vk_create_pipeline(device, device->pipeline_norm_f32, "norm_f32", norm_f32_len, norm_f32_data, "main", 2, sizeof(vk_op_unary_push_constants), {1, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_group_norm_f32, "group_norm_f32", group_norm_f32_len, group_norm_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, {}, 1);
 
-    ggml_vk_create_pipeline(device, device->pipeline_rms_norm_f32, "rms_norm_f32", rms_norm_f32_len, rms_norm_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 0}, 1, true);
-    ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_f32, "rms_norm_mul_f32", rms_norm_f32_len, rms_norm_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1}, 1, true);
+    if (device->subgroup_arithmetic) {
+        // single-barrier subgroup reduction
+        ggml_vk_create_pipeline(device, device->pipeline_rms_norm_f32, "rms_norm_f32", rms_norm_subgroup_f32_len, rms_norm_subgroup_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 0}, 1, true);
+        ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_f32, "rms_norm_mul_f32", rms_norm_subgroup_f32_len, rms_norm_subgroup_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1}, 1, true);
+    } else {
+        ggml_vk_create_pipeline(device, device->pipeline_rms_norm_f32, "rms_norm_f32", rms_norm_f32_len, rms_norm_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 0}, 1, true);
+        ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_f32, "rms_norm_mul_f32", rms_norm_f32_len, rms_norm_f32_data, "main", 4, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1}, 1, true);
+    }
     ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_add_f32, "rms_norm_mul_add_f32", rms_norm_mul_add_f32_len, rms_norm_mul_add_f32_data, "main", 5, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1, 0}, 1, true);
     ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_add_mul_f32, "rms_norm_mul_add_mul_f32", rms_norm_mul_add_f32_len, rms_norm_mul_add_f32_data, "main", 5, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1, 1}, 1, true);
     ggml_vk_create_pipeline(device, device->pipeline_rms_norm_mul_add_partials_f32, "rms_norm_mul_add_partials_f32", rms_norm_mul_add_partials_f32_len, rms_norm_mul_add_partials_f32_data, "main", 6, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {0, 1, 0}, 1, true);
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
index ee813842c..367180fb7 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp
@@ -39,7 +39,12 @@ layout (binding = 3) readonly buffer I {uvec2 data_i[];};
 #endif
 
 #extension GL_EXT_control_flow_attributes : enable
+#if USE_SUBGROUP_ADD
+#extension GL_KHR_shader_subgroup_basic : require
+#extension GL_KHR_shader_subgroup_arithmetic : require
+#endif
 #define BLOCK_SIZE 512
+#define CACHE_ITERS 16
 
 layout (constant_id = 1) const bool do_multiply = false;
 #if RMS_NORM_ADD_FUSION
@@ -76,14 +81,32 @@ void rms_norm(uint num_iters) {
 #endif
     FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
 
+    // keep the row in registers for the second pass (num_iters is a small compile-time constant)
+    FLOAT_TYPE xs[CACHE_ITERS];
     [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
         FLOAT_TYPE xi = FLOAT_TYPE(0);
         if (col < ncols) {
             xi = FLOAT_TYPE(data_a[a_offset + col]);
         }
+        if (idx < CACHE_ITERS) {
+            xs[idx] = xi;
+        }
         sum += xi * xi;
     }
+#define LOAD_A(idx, col) ((idx) < CACHE_ITERS ? xs[(idx) < CACHE_ITERS ? (idx) : 0] : FLOAT_TYPE(data_a[a_offset + (col)]))
 
+#if USE_SUBGROUP_ADD
+    // reduce within subgroups, then across the (few) subgroups: a single barrier
+    sum = subgroupAdd(sum);
+    if (gl_SubgroupInvocationID == 0) {
+        sumsh[gl_SubgroupID] = sum;
+    }
+    barrier();
+    sum = FLOAT_TYPE(0.0f);
+    for (uint i = 0; i < gl_NumSubgroups; ++i) {
+        sum += sumsh[i];
+    }
+#else
     sumsh[tid] = sum;
     // sum up partial sums and write back result
     barrier();
@@ -95,6 +118,7 @@ void rms_norm(uint num_iters) {
         barrier();
     }
     sum = sumsh[0];
+#endif
 
     const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
     const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
@@ -105,7 +129,7 @@ void rms_norm(uint num_iters) {
                 if (col >= ncols) {
                     continue;
                 }
-                FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
+                FLOAT_TYPE value = scale * LOAD_A(idx, col) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
 #if RMS_NORM_ADD_FUSION
                 value += FLOAT_TYPE(data_c[d_offset + col]);
                 if (do_post_multiply) {
@@ -119,7 +143,7 @@ void rms_norm(uint num_iters) {
                 if (col >= ncols) {
                     continue;
                 }
-                FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
+                FLOAT_TYPE value = scale * LOAD_A(idx, col) * FLOAT_TYPE(data_b[b_offset + col]);
 #if RMS_NORM_ADD_FUSION
                 value += FLOAT_TYPE(data_c[d_offset + col]);
                 if (do_post_multiply) {
@@ -134,7 +158,7 @@ void rms_norm(uint num_iters) {
             if (col >= ncols) {
                 continue;
             }
-            data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
+            data_d[d_offset + col] = D_TYPE(scale * LOAD_A(idx, col));
         }
     }
 #if RMS_NORM_ROPE_FUSION
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
index e7e303e50..115534893 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -853,6 +853,7 @@ void process_shaders() {
     string_to_spv("norm_f32", "norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
     string_to_spv("group_norm_f32", "group_norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
     string_to_spv("rms_norm_f32", "rms_norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}}));
+    string_to_spv("rms_norm_subgroup_f32", "rms_norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}}));
     string_to_spv("rms_norm_mul_add_f32", "rms_norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}, {"RMS_NORM_ADD_FUSION", "1"}}));
     string_to_spv("rms_norm_mul_add_partials_f32", "rms_norm_partials.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}, {"RMS_NORM_ADD_FUSION", "1"}}));
     string_to_spv("rms_norm_set_rows_f32_f32", "rms_norm.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}, {"RMS_NORM_SET_ROWS_FUSION", "1"}}));
-- 
2.55.0

