Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 6 additions & 4 deletions onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Comment thread
qjia7 marked this conversation as resolved.
Outdated
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)
{
Expand Down Expand Up @@ -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<uint32_t>(parameters.sequence_length_)},
Expand All @@ -454,7 +455,8 @@ Status ApplyFlashAttention(const Tensor* Q, const Tensor* K, const Tensor* V, co
{static_cast<uint32_t>(parameters.total_sequence_length_ - parameters.kv_sequence_length_)},
{static_cast<uint32_t>(parameters.is_gqa_ ? 1 : 0)},
{static_cast<uint32_t>(parameters.n_reps)},
{alpha}});
{alpha},
{num_seq_tile}});

return context.RunProgram(program);
}
Expand Down
3 changes: 2 additions & 1 deletion onnxruntime/contrib_ops/webgpu/bert/flash_attention.h
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,8 @@ class FlashAttentionProgram final : public Program<FlashAttentionProgram> {
{"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_;
Expand Down
13 changes: 7 additions & 6 deletions onnxruntime/contrib_ops/webgpu/quantization/dp4a_matmul_nbits.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Comment thread
qjia7 marked this conversation as resolved.
Outdated
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);
Expand Down Expand Up @@ -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<int>(kVec4Components)},
{&a_scale, ProgramTensorMetadataDependency::TypeAndRank, 1},
{b, ProgramTensorMetadataDependency::TypeAndRank, static_cast<int>(kVec2Components * kU32Components)},
Expand All @@ -298,7 +298,8 @@ Status ApplyDP4AMatrixMatMulNBits(const Tensor* a, const Tensor* b, const Tensor
{static_cast<uint32_t>(N)},
{static_cast<uint32_t>(K)},
{static_cast<uint32_t>(K / 8)},
{static_cast<uint32_t>(K / 16)}})
{static_cast<uint32_t>(K / 16)},
{num_N_tile}})
.AddOutput({y, ProgramTensorMetadataDependency::TypeAndRank, reshaped_y_shape, static_cast<int>(kVec4Components)})
.CacheHint("Block" + std::to_string(block_size));
return context.RunProgram(mul_program);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,8 @@ class DP4AMatMulNBitsProgram final : public Program<DP4AMatMulNBitsProgram> {
{"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_;
Expand Down