From 85481d4837cdc68d923fbb99e17a300d1357dba7 Mon Sep 17 00:00:00 2001 From: Jiajia Qin Date: Tue, 18 Mar 2025 21:03:26 +0800 Subject: [PATCH 1/2] [webgpu] Use 1d dispatch group size This PR uses 1d disptach group size and uses workgroup_idx instead of workgroup.x|workgroup.y in case they are normalized. --- .../contrib_ops/webgpu/bert/flash_attention.cc | 10 ++++++---- .../contrib_ops/webgpu/bert/flash_attention.h | 3 ++- .../webgpu/quantization/dp4a_matmul_nbits.cc | 13 +++++++------ .../webgpu/quantization/dp4a_matmul_nbits.h | 3 ++- 4 files changed, 17 insertions(+), 12 deletions(-) diff --git a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc index 58ddf60df79f0..9e6dae696c254 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc @@ -224,12 +224,12 @@ Status FlashAttentionProgram::GenerateShaderCode(ShaderHelper& shader) const { // Shader is designed to be dispatched as Dispatch(num_heads, new_sequence_length / workgroup_size_x, 1) // Each lane/thread is responsible for a single q. shader.MainFunctionBody() << R"MAIN_FN( - let head_idx = workgroup_id.x; + let head_idx = workgroup_idx / uniforms.num_seq_tile; let capped_sg_id = min(sg_id, max_k_step); let capped_sg_size = min(sg_size, max_k_step); // Load Q - let q_idx_global = workgroup_id.y * workgroup_size_x + local_idx; + let q_idx_global = (workgroup_idx % uniforms.num_seq_tile) * workgroup_size_x + local_idx; let valid_q = q_idx_global < uniforms.new_sequence_length; if (valid_q) { @@ -445,7 +445,8 @@ Status ApplyFlashAttention(const Tensor* Q, const Tensor* K, const Tensor* V, co std::string cache_hint = std::to_string(has_attention_bias) + std::to_string(parameters.head_size_) + std::to_string(parameters.num_heads_); - program.SetDispatchGroupSize(parameters.num_heads_, (parameters.sequence_length_ + tile_size - 1) / tile_size, 1) + const uint32_t num_seq_tile = (parameters.sequence_length_ + tile_size - 1) / tile_size; + program.SetDispatchGroupSize(parameters.num_heads_ * num_seq_tile) .SetWorkgroupSize(tile_size) .CacheHint(cache_hint) .AddUniformVariables({{static_cast(parameters.sequence_length_)}, @@ -454,7 +455,8 @@ Status ApplyFlashAttention(const Tensor* Q, const Tensor* K, const Tensor* V, co {static_cast(parameters.total_sequence_length_ - parameters.kv_sequence_length_)}, {static_cast(parameters.is_gqa_ ? 1 : 0)}, {static_cast(parameters.n_reps)}, - {alpha}}); + {alpha}, + {num_seq_tile}}); return context.RunProgram(program); } diff --git a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.h b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.h index 2c2b888538843..8931403641a81 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.h +++ b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.h @@ -52,7 +52,8 @@ class FlashAttentionProgram final : public Program { {"past_sequence_length", ProgramUniformVariableDataType::Uint32}, {"is_gqa", ProgramUniformVariableDataType::Uint32}, {"n_reps", ProgramUniformVariableDataType::Uint32}, - {"alpha", ProgramUniformVariableDataType::Float32}); + {"alpha", ProgramUniformVariableDataType::Float32}, + {"num_seq_tile", ProgramUniformVariableDataType::Uint32}); private: bool has_attention_bias_; diff --git a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc index 05cbfb1f99c48..4df59117a8927 100644 --- a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc +++ b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc @@ -147,8 +147,8 @@ Status DP4AMatMulNBitsProgram::GenerateShaderCode(ShaderHelper& shader) const { shader.MainFunctionBody() << R"MAIN_FN( // During the load phase we use all 256 threads to load 64 rows of A/B. // For each row we load tile_size_k_vec (2) vectorized elements, which are 32 elements of K. - let a_global_base = workgroup_id.x * tile_size; - let b_global_base = workgroup_id.y * tile_size; + let a_global_base = (workgroup_idx / uniforms.num_N_tile) * tile_size; + let b_global_base = (workgroup_idx % uniforms.num_N_tile) * tile_size; let load_AorB = u32(local_idx/128); let load_row = u32((local_idx%128)/2); let load_col = u32(local_idx%2); @@ -285,11 +285,11 @@ Status ApplyDP4AMatrixMatMulNBits(const Tensor* a, const Tensor* b, const Tensor constexpr uint32_t kTileSize = 64; TensorShape reshaped_y_shape{1, M, N / kVec4Components}; + uint32_t num_M_tile = (M + kTileSize - 1) / kTileSize; + uint32_t num_N_tile = (N + kTileSize - 1) / kTileSize; DP4AMatMulNBitsProgram mul_program{block_size}; mul_program.SetWorkgroupSize(256); - mul_program.SetDispatchGroupSize( - (M + kTileSize - 1) / kTileSize, - (N + kTileSize - 1) / kTileSize, 1); + mul_program.SetDispatchGroupSize(num_M_tile * num_N_tile); mul_program.AddInputs({{&a_quant, ProgramTensorMetadataDependency::TypeAndRank, static_cast(kVec4Components)}, {&a_scale, ProgramTensorMetadataDependency::TypeAndRank, 1}, {b, ProgramTensorMetadataDependency::TypeAndRank, static_cast(kVec2Components * kU32Components)}, @@ -298,7 +298,8 @@ Status ApplyDP4AMatrixMatMulNBits(const Tensor* a, const Tensor* b, const Tensor {static_cast(N)}, {static_cast(K)}, {static_cast(K / 8)}, - {static_cast(K / 16)}}) + {static_cast(K / 16)}, + {num_N_tile}}) .AddOutput({y, ProgramTensorMetadataDependency::TypeAndRank, reshaped_y_shape, static_cast(kVec4Components)}) .CacheHint("Block" + std::to_string(block_size)); return context.RunProgram(mul_program); diff --git a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.h b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.h index 15b86d78301ad..4250fc2dc4aeb 100644 --- a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.h +++ b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.h @@ -28,7 +28,8 @@ class DP4AMatMulNBitsProgram final : public Program { {"N", ProgramUniformVariableDataType::Uint32}, {"K", ProgramUniformVariableDataType::Uint32}, {"K8", ProgramUniformVariableDataType::Uint32}, - {"K16", ProgramUniformVariableDataType::Uint32}); + {"K16", ProgramUniformVariableDataType::Uint32}, + {"num_N_tile", ProgramUniformVariableDataType::Uint32}); private: uint32_t block_size_; From 99f41082da57e1b6f14db751773e0b113b085519 Mon Sep 17 00:00:00 2001 From: Jiajia Qin Date: Wed, 19 Mar 2025 10:02:49 +0800 Subject: [PATCH 2/2] address comments --- onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc | 2 +- .../contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc index 9e6dae696c254..52c705abb1003 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc @@ -224,7 +224,7 @@ Status FlashAttentionProgram::GenerateShaderCode(ShaderHelper& shader) const { // Shader is designed to be dispatched as Dispatch(num_heads, new_sequence_length / workgroup_size_x, 1) // Each lane/thread is responsible for a single q. shader.MainFunctionBody() << R"MAIN_FN( - let head_idx = workgroup_idx / uniforms.num_seq_tile; + let head_idx = u32(workgroup_idx / uniforms.num_seq_tile); let capped_sg_id = min(sg_id, max_k_step); let capped_sg_size = min(sg_size, max_k_step); diff --git a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc index 4df59117a8927..528042b065ab0 100644 --- a/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc +++ b/onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc @@ -147,7 +147,7 @@ Status DP4AMatMulNBitsProgram::GenerateShaderCode(ShaderHelper& shader) const { shader.MainFunctionBody() << R"MAIN_FN( // During the load phase we use all 256 threads to load 64 rows of A/B. // For each row we load tile_size_k_vec (2) vectorized elements, which are 32 elements of K. - let a_global_base = (workgroup_idx / uniforms.num_N_tile) * tile_size; + let a_global_base = u32(workgroup_idx / uniforms.num_N_tile) * tile_size; let b_global_base = (workgroup_idx % uniforms.num_N_tile) * tile_size; let load_AorB = u32(local_idx/128); let load_row = u32((local_idx%128)/2);