Skip to content

Commit 8b3b253

Browse files
authored
ggml-webgpu: support non-square subgroup matrix configs for Intel GPUs (ggml-org#21669)
1 parent 8c90049 commit 8b3b253

2 files changed

Lines changed: 27 additions & 20 deletions

File tree

ggml/src/ggml-webgpu/ggml-webgpu.cpp

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -3461,13 +3461,15 @@ static bool create_webgpu_device(ggml_backend_webgpu_reg_context * ctx) {
34613461
GGML_ASSERT(ctx->webgpu_global_ctx->adapter.HasFeature(wgpu::FeatureName::ShaderF16));
34623462

34633463
#ifndef __EMSCRIPTEN__
3464-
// Only support square f16 matrices of size 8 or 16 for now
3464+
// Accept f16 subgroup matrix configurations (square or non-square).
3465+
// NVIDIA GPUs typically report square configs (e.g. 16x16x16),
3466+
// while Intel Xe2 GPUs report non-square configs (e.g. 8x16x16).
3467+
// The shaders are already parameterized to handle any M/N/K dimensions.
34653468
bool valid_subgroup_matrix_config = false;
34663469
if (ctx->webgpu_global_ctx->adapter.HasFeature(wgpu::FeatureName::ChromiumExperimentalSubgroupMatrix)) {
34673470
for (size_t i = 0; i < subgroup_matrix_configs.configCount; i++) {
34683471
const wgpu::SubgroupMatrixConfig config = subgroup_matrix_configs.configs[i];
3469-
if (config.M == config.N && config.N == config.K && (config.K == 8 || config.K == 16) &&
3470-
config.componentType == wgpu::SubgroupMatrixComponentType::F16 &&
3472+
if (config.componentType == wgpu::SubgroupMatrixComponentType::F16 &&
34713473
config.resultComponentType == wgpu::SubgroupMatrixComponentType::F16) {
34723474
ctx->webgpu_global_ctx->capabilities.sg_mat_m = config.M;
34733475
ctx->webgpu_global_ctx->capabilities.sg_mat_n = config.N;
@@ -3805,6 +3807,11 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
38053807
if (!ctx->webgpu_global_ctx->capabilities.supports_subgroup_matrix) {
38063808
break;
38073809
}
3810+
// Head dimensions must be divisible by subgroup matrix dimensions
3811+
if (src0->ne[0] % ctx->webgpu_global_ctx->capabilities.sg_mat_k != 0 ||
3812+
src2->ne[0] % ctx->webgpu_global_ctx->capabilities.sg_mat_n != 0) {
3813+
break;
3814+
}
38083815
// Head dimensions must fit in workgroup memory with minimum tile sizes
38093816
size_t limit_bytes = ctx->webgpu_global_ctx->capabilities.limits.maxComputeWorkgroupStorageSize;
38103817
const bool has_mask = op->src[3] != nullptr;

ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl

Lines changed: 17 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -369,35 +369,35 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
369369
#endif
370370
for (var kv_block = subgroup_id; kv_block < KV_BLOCKS; kv_block += num_subgroups) {
371371
let inter_offset = kv_block * SG_MAT_N;
372-
var acc: subgroup_matrix_result<f16, SG_MAT_M, SG_MAT_N> = subgroupMatrixLoad<subgroup_matrix_result<f16, SG_MAT_M, SG_MAT_N>>(&inter_shmem, inter_offset, false, KV_TILE);
372+
var acc: subgroup_matrix_result<f16, SG_MAT_N, SG_MAT_M> = subgroupMatrixLoad<subgroup_matrix_result<f16, SG_MAT_N, SG_MAT_M>>(&inter_shmem, inter_offset, false, KV_TILE);
373373

374-
var q_cur = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K>>(&q_shmem, 0u, false, HEAD_DIM_QK);
374+
var q_cur = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M>>(&q_shmem, 0u, false, HEAD_DIM_QK);
375375

376376
#ifdef KV_DIRECT
377-
var k_cur = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&K, k_global_offset + 0u, true, params.stride_k1);
377+
var k_cur = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&K, k_global_offset + 0u, true, params.stride_k1);
378378
#else
379-
var k_cur = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&kv_shmem, k_block_offset + 0u, true, HEAD_DIM_QK);
379+
var k_cur = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&kv_shmem, k_block_offset + 0u, true, HEAD_DIM_QK);
380380
#endif
381381

382382
var t: u32 = 1u;
383383
for (; t + 1u < HEAD_DIM_QK / SG_MAT_K; t += 2u) {
384384
let h0 = t * SG_MAT_K;
385-
var q0 = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K>>(&q_shmem, h0, false, HEAD_DIM_QK);
385+
var q0 = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M>>(&q_shmem, h0, false, HEAD_DIM_QK);
386386
#ifdef KV_DIRECT
387-
var k0 = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&K, k_global_offset + h0, true, params.stride_k1);
387+
var k0 = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&K, k_global_offset + h0, true, params.stride_k1);
388388
#else
389-
var k0 = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&kv_shmem, k_block_offset + h0, true, HEAD_DIM_QK);
389+
var k0 = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&kv_shmem, k_block_offset + h0, true, HEAD_DIM_QK);
390390
#endif
391391
acc = subgroupMatrixMultiplyAccumulate(q_cur, k_cur, acc);
392392
q_cur = q0;
393393
k_cur = k0;
394394

395395
let h1 = (t + 1u) * SG_MAT_K;
396-
var q1g = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K>>(&q_shmem, h1, false, HEAD_DIM_QK);
396+
var q1g = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M>>(&q_shmem, h1, false, HEAD_DIM_QK);
397397
#ifdef KV_DIRECT
398-
var k1g = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&K, k_global_offset + h1, true, params.stride_k1);
398+
var k1g = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&K, k_global_offset + h1, true, params.stride_k1);
399399
#else
400-
var k1g = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&kv_shmem, k_block_offset + h1, true, HEAD_DIM_QK);
400+
var k1g = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&kv_shmem, k_block_offset + h1, true, HEAD_DIM_QK);
401401
#endif
402402
acc = subgroupMatrixMultiplyAccumulate(q_cur, k_cur, acc);
403403
q_cur = q1g;
@@ -407,11 +407,11 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
407407
// handle odd tail
408408
if (t < HEAD_DIM_QK / SG_MAT_K) {
409409
let h = t * SG_MAT_K;
410-
var qn = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K>>(&q_shmem, h, false, HEAD_DIM_QK);
410+
var qn = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M>>(&q_shmem, h, false, HEAD_DIM_QK);
411411
#ifdef KV_DIRECT
412-
var kn = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&K, k_global_offset + h, true, params.stride_k1);
412+
var kn = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&K, k_global_offset + h, true, params.stride_k1);
413413
#else
414-
var kn = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(&kv_shmem, k_block_offset + h, true, HEAD_DIM_QK);
414+
var kn = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(&kv_shmem, k_block_offset + h, true, HEAD_DIM_QK);
415415
#endif
416416
acc = subgroupMatrixMultiplyAccumulate(q_cur, k_cur, acc);
417417
q_cur = qn;
@@ -566,15 +566,15 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
566566
head_dim_block < HEAD_DIM_V;
567567
head_dim_block += num_subgroups * SG_MAT_N) {
568568
// load O submatrix from shared memory
569-
var o_sg_mat: subgroup_matrix_result<f16, SG_MAT_M, SG_MAT_N> = subgroupMatrixLoad<subgroup_matrix_result<f16, SG_MAT_M, SG_MAT_N>>(
569+
var o_sg_mat: subgroup_matrix_result<f16, SG_MAT_N, SG_MAT_M> = subgroupMatrixLoad<subgroup_matrix_result<f16, SG_MAT_N, SG_MAT_M>>(
570570
&o_shmem,
571571
head_dim_block,
572572
false,
573573
HEAD_DIM_V
574574
);
575575
for (var kv_block = 0u; kv_block < KV_BLOCKS; kv_block++) {
576576
let p_offset = kv_block * SG_MAT_N;
577-
var p_sg_mat: subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K> = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_M, SG_MAT_K>>(
577+
var p_sg_mat: subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M> = subgroupMatrixLoad<subgroup_matrix_left<f16, SG_MAT_K, SG_MAT_M>>(
578578
&inter_shmem,
579579
p_offset,
580580
false,
@@ -585,15 +585,15 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
585585
#ifdef KV_DIRECT
586586
let v_block_row = kv_tile + kv_block * SG_MAT_N;
587587
let v_global_offset = v_head_offset + v_block_row * params.stride_v1 + head_dim_block;
588-
var v_sg_mat: subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N> = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(
588+
var v_sg_mat: subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K> = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(
589589
&V,
590590
v_global_offset,
591591
false,
592592
params.stride_v1
593593
);
594594
#else
595595
let v_block_offset = kv_block * SG_MAT_N * HEAD_DIM_V;
596-
var v_sg_mat: subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N> = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_K, SG_MAT_N>>(
596+
var v_sg_mat: subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K> = subgroupMatrixLoad<subgroup_matrix_right<f16, SG_MAT_N, SG_MAT_K>>(
597597
&kv_shmem,
598598
v_block_offset + head_dim_block,
599599
false,

0 commit comments

Comments
 (0)