From 27b873b67716e66b6faaf98b926d5822de5bd934 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sun, 27 Sep 2026 13:29:47 +0100
Subject: [PATCH 15/18] vulkan : concat with a transposed src1 through
 shared-memory tiles

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan-types.h      |  1 +
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          |  9 ++-
 .../vulkan-shaders/concat_transpose.comp      | 61 +++++++++++++++++++
 .../vulkan-shaders/vulkan-shaders-gen.cpp     |  1 +
 tests/test-backend-ops.cpp                    | 17 +++++-
 5 files changed, 87 insertions(+), 2 deletions(-)
 create mode 100644 ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
index de6b06390..cc472b891 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
@@ -836,6 +836,7 @@ struct vk_device_struct {
     vk_pipeline pipeline_add_id_f32;
 
     vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64;
+    vk_pipeline pipeline_concat_transpose_i32;
     vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32;
     vk_pipeline pipeline_scale_f32;
     vk_pipeline pipeline_log[2];
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index f1cd6a636..636848e05 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -3386,6 +3386,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
     ggml_vk_create_pipeline(device, device->pipeline_concat_i8, "concat_i8", concat_i8_len, concat_i8_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_concat_i16, "concat_i16", concat_i16_len, concat_i16_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_concat_i32, "concat_i32", concat_i32_len, concat_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1);
+    ggml_vk_create_pipeline(device, device->pipeline_concat_transpose_i32, "concat_transpose_i32", concat_transpose_i32_len, concat_transpose_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_concat_i64, "concat_i64", concat_i64_len, concat_i64_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1);
 
     ggml_vk_create_pipeline(device, device->pipeline_upscale_nearest_f32, "upscale_f32", upscale_f32_len, upscale_f32_data, "main", 2, sizeof(vk_op_upscale_push_constants), {512, 1, 1}, {GGML_SCALE_MODE_NEAREST}, 1);
@@ -8680,6 +8681,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const
         if (!ggml_vk_concat_supported(src0, src1, dst)) {
             return nullptr;
         }
+        // src1 is a transposed view (e.g. the ssm conv state concat): stage tiles through shared memory
+        if (ggml_get_op_params_i32(dst, 0) == 0 && ggml_type_size(src0->type) == 4 && !ggml_is_quantized(src0->type) &&
+            src1->nb[1] == 4 && src1->nb[0] > src1->nb[1] && dst->nb[0] == 4) {
+            return ctx->device->pipeline_concat_transpose_i32;
+        }
         switch (ggml_vk_concat_unit_size(src0->type)) {
         case 1:
             return ctx->device->pipeline_concat_i8;
@@ -9642,7 +9648,8 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
                 elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
                 elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
             } else if (pipeline == ctx->device->pipeline_cpy_transpose_32 ||
-                pipeline == ctx->device->pipeline_cpy_transpose_16) {
+                pipeline == ctx->device->pipeline_cpy_transpose_16 ||
+                pipeline == ctx->device->pipeline_concat_transpose_i32) {
                 // 32x32 tiles
                 elements[0] = (uint32_t)CEIL_DIV(dst->ne[0], 32);
                 elements[1] = (uint32_t)CEIL_DIV(dst->ne[1], 32);
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp
new file mode 100644
index 000000000..e365d2441
--- /dev/null
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp
@@ -0,0 +1,61 @@
+#version 450
+
+#include "types.glsl"
+#include "generic_binary_head.glsl"
+
+// Concat along dim 0 with a transposed src1 (dim 1 innermost): stage 32x32 tiles in shared memory so reads and writes are contiguous.
+#define TILE_DIM 32
+layout(local_size_x = 32, local_size_y = 8, local_size_z = 1) in;
+
+shared uint sh[TILE_DIM][TILE_DIM + 1];
+
+void iter(uvec3 wg_id) {
+    const uint tid_col = gl_LocalInvocationID.x;
+    const uint tid_row = gl_LocalInvocationID.y;
+
+    const uint i2 = wg_id.z % p.ne22;
+    const uint i3 = wg_id.z / p.ne22;
+
+    // src reads: tid.x walks i1
+    [[unroll]] for (uint y = 0; y < 4; ++y) {
+        const uint i0 = wg_id.x * TILE_DIM + tid_row + 8 * y;
+        const uint i1 = wg_id.y * TILE_DIM + tid_col;
+        if (i0 < p.ne20 && i1 < p.ne21) {
+            uint v;
+            if (i0 < p.ne00) {
+                v = uint(data_a[get_aoffset() + i0 * p.nb00 + i1 * p.nb01 + i2 * p.nb02 + i3 * p.nb03]);
+            } else {
+                v = uint(data_b[get_boffset() + (i0 - p.ne00) * p.nb10 + i1 * p.nb11 + i2 * p.nb12 + i3 * p.nb13]);
+            }
+            sh[tid_row + 8 * y][tid_col] = v;
+        }
+    }
+
+    barrier();
+
+    // dst writes: tid.x walks i0
+    [[unroll]] for (uint y = 0; y < 4; ++y) {
+        const uint i0 = wg_id.x * TILE_DIM + tid_col;
+        const uint i1 = wg_id.y * TILE_DIM + tid_row + 8 * y;
+        if (i0 < p.ne20 && i1 < p.ne21) {
+            data_d[get_doffset() + i0 * p.nb20 + i1 * p.nb21 + i2 * p.nb22 + i3 * p.nb23] = D_TYPE(sh[tid_col][tid_row + 8 * y]);
+        }
+    }
+}
+
+#define CEIL_DIV(a, b) (((a) + (b) - 1) / (b))
+
+void main() {
+    bool need_barrier = false;
+    for (uint z = gl_WorkGroupID.z; z < p.ne22 * p.ne23; z += gl_NumWorkGroups.z) {
+        for (uint y = gl_WorkGroupID.y; y < CEIL_DIV(p.ne21, TILE_DIM); y += gl_NumWorkGroups.y) {
+            for (uint x = gl_WorkGroupID.x; x < CEIL_DIV(p.ne20, TILE_DIM); x += gl_NumWorkGroups.x) {
+                if (need_barrier) {
+                    barrier();
+                }
+                need_barrier = true;
+                iter(uvec3(x, y, z));
+            }
+        }
+    }
+}
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 eb14e065c..058f4e921 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -964,6 +964,7 @@ void process_shaders() {
     string_to_spv("concat_i8", "concat.comp", {{"A_TYPE", "uint8_t"}, {"B_TYPE", "uint8_t"}, {"D_TYPE", "uint8_t"}});
     string_to_spv("concat_i16", "concat.comp", {{"A_TYPE", "uint16_t"}, {"B_TYPE", "uint16_t"}, {"D_TYPE", "uint16_t"}});
     string_to_spv("concat_i32", "concat.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}});
+    string_to_spv("concat_transpose_i32", "concat_transpose.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}});
     string_to_spv("concat_i64", "concat.comp", {{"A_TYPE", "uvec2"}, {"B_TYPE", "uvec2"}, {"D_TYPE", "uvec2"}});
 
     string_to_spv("upscale_f32", "upscale.comp", {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}});
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 7ab411442..78f6d8260 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -6607,7 +6607,15 @@ struct test_concat : public test_case {
             ggml_set_name(a, "a");
         }
         ggml_tensor * b;
-        if (v & 2) {
+        if (v & 16) {
+            // transposed b, like the ssm conv state concat
+            auto ne = ne_b; std::swap(ne[0], ne[1]);
+            b = ggml_new_tensor(ctx, type, 4, ne.data());
+            ggml_set_name(b, "b");
+
+            b = ggml_transpose(ctx, b);
+            ggml_set_name(b, "transpose_of_b");
+        } else if (v & 2) {
             auto ne = ne_b; ne[0] *= 3; ne[1] *= 2; ne[2] *= 4;
             b = ggml_new_tensor(ctx, type, 4, ne.data());
             ggml_set_name(b, "b");
@@ -10675,6 +10683,12 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
         }
     }
 
+    for (ggml_type type : { GGML_TYPE_F32, GGML_TYPE_I32, GGML_TYPE_F16 }) {
+        test_cases.emplace_back(new test_concat(type, {11, 12, 13, 14}, 7, 0, 16));
+        test_cases.emplace_back(new test_concat(type, {3, 1000, 2, 1}, 100, 0, 16));
+        test_cases.emplace_back(new test_concat(type, {3, 1000, 1, 1}, 100, 0, 17));
+    }
+
     for (ggml_type type_a : { GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0 }) {
         for (int v : { 0, 4, 8, 12 }) {
             for (int dim : { 0, 1, 2, 3, }) {
@@ -11283,6 +11297,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
 static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
 
     std::vector<std::unique_ptr<test_case>> test_cases;
+    test_cases.emplace_back(new test_concat(GGML_TYPE_F32, {3, 10240, 1, 1}, 1024, 0, 16));
 
     // Qwen3.x 27B style attention: 24 Q / 4 KV heads, head size 256, q8_0 KV cache (prefill + decode)
     for (int kv : { 8192, 32768, 33792, 34816, 65536, 66560, }) {
-- 
2.55.0

