From e0a73ac5bcd2e1bc37263df250cb8fc9b9b58957 Mon Sep 17 00:00:00 2001 From: bri-prism <288398250+bri-prism@users.noreply.github.com> Date: Wed, 30 Sep 2026 19:18:19 -0700 Subject: [PATCH] webgpu: route up to 8 columns through the mat-vec kernel The mat-vec shader is already generic in NUM_COLS, but ggml_webgpu_mul_mat only used it for ne11 <= 4 and sent 5..8 columns to the 32x32 register tile, which pads M and runs 2-4x slower at decode-batch sizes on Apple M5 Pro (Dawn, Metal): q4_0 17408x5120 n=8 1.91 ms -> 0.99 ms q8_0 17408x5120 n=8 2.01 ms -> 1.13 ms f16 17408x5120 n=8 1.83 ms -> 0.84 ms q1_0 17408x5120 n=8 1.95 ms -> 0.87 ms n=4 unchanged; n=12/16 stay on the tile path (mat-vec regresses there). 8 matches MMVQ_MAX_BATCH_SIZE on CUDA and mul_mat_vec_max_cols on Vulkan. Also build ggml-webgpu as C++20: the Dawn webgpu_cpp.h header needs std::span and std::type_identity. --- ggml/src/ggml-webgpu/CMakeLists.txt | 2 ++ ggml/src/ggml-webgpu/ggml-webgpu.cpp | 2 +- 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-webgpu/CMakeLists.txt b/ggml/src/ggml-webgpu/CMakeLists.txt index 1503a1ef8ba6..d63c1da3ae89 100644 --- a/ggml/src/ggml-webgpu/CMakeLists.txt +++ b/ggml/src/ggml-webgpu/CMakeLists.txt @@ -81,3 +81,5 @@ endif() target_include_directories(ggml-webgpu PRIVATE ${SHADER_OUTPUT_DIR}) target_link_libraries(ggml-webgpu PRIVATE ${DawnWebGPU_TARGET}) + +target_compile_features(ggml-webgpu PRIVATE cxx_std_20) diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 377f3d702c02..0f76b87852f2 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -1603,7 +1603,7 @@ static webgpu_encoded_op ggml_webgpu_mul_mat(webgpu_context & ctx, ggml_tensor * src1, ggml_tensor * dst) { // Determine if this is a mat-vec operation - bool use_mat_vec = (dst->ne[1] <= 4); + bool use_mat_vec = (dst->ne[1] <= 8); // use MMVQ path for mat-vec bool use_mmvq = ggml_webgpu_can_use_mmvq(src0, src1, ctx->global_ctx->capabilities.supports_dot_product,