Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
435c445
webgpu: Extend FlashAttention decode path to support any sequence length
qjia7 May 7, 2026
39ffe45
webgpu: Add m_tile optimization and remove Subgroups requirement for …
qjia7 May 11, 2026
ecfc56a
webgpu: Fuse QKT and SplitVx into single QKV kernel for FlashAttentio…
qjia7 May 12, 2026
934ed2a
webgpu: Route between prefill and split-reduce paths in FlashAttention
qjia7 May 12, 2026
35b35f4
webgpu: Fix clang-format in FlashAttentionDecodeQKVProgram constructor
qjia7 May 12, 2026
9b8dadf
Merge branch 'main' into webgpu-flash-attention-relax-subgroups
qjia7 May 19, 2026
b7e6592
Remove subgroups feature check from use_split_reduce routing
qjia7 May 19, 2026
46fd59a
Merge branch 'main' into webgpu-flash-attention-relax-subgroups
qjia7 May 21, 2026
07ef10d
Remove use_seqlen_k extension from decode kernels, keep only use_indi…
qjia7 May 21, 2026
b690d0f
Extend indirect dispatch to support sequence_length > 1 in split-redu…
qjia7 May 21, 2026
4241a91
Extract shared write_indirect_dispatch WGSL function for dispatch nor…
qjia7 May 21, 2026
42419a6
Merge branch 'main' into webgpu-flash-attention-relax-subgroups
qjia7 Jun 2, 2026
a16023e
Merge branch 'main' into webgpu-flash-attention-relax-subgroups
qjia7 Jun 8, 2026
10ef54d
Tune flash attention split-reduce threshold to seq_len < 32
qjia7 Jun 8, 2026
b0dcedf
Use getByOffset/setByOffset helpers in flash attention decode shaders
qjia7 Jun 9, 2026
ea60286
Normalize indirect-dispatch helper to accept (x, y, z) and consolidat…
qjia7 Jun 9, 2026
13ea0e4
Simplify flash attention split-reduce routing to seq_len < 32
qjia7 Jun 9, 2026
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
267 changes: 135 additions & 132 deletions onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc

Large diffs are not rendered by default.

59 changes: 26 additions & 33 deletions onnxruntime/contrib_ops/webgpu/bert/flash_attention.h
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,9 @@
{"half_rotary_dim", ProgramUniformVariableDataType::Uint32},
{"present_sequence_length", ProgramUniformVariableDataType::Uint32},
{"tile_size", ProgramUniformVariableDataType::Uint32},
{"dispatch_size", ProgramUniformVariableDataType::Uint32});
{"dispatch_size", ProgramUniformVariableDataType::Uint32},
{"batch_size", ProgramUniformVariableDataType::Uint32},
{"num_q_tiles", ProgramUniformVariableDataType::Uint32});

private:
const bool interleaved_;
Expand All @@ -56,7 +58,9 @@
{"total_sequence_length", ProgramUniformVariableDataType::Uint32},
{"kv_sequence_length", ProgramUniformVariableDataType::Uint32},
{"tile_size", ProgramUniformVariableDataType::Uint32},
{"num_heads", ProgramUniformVariableDataType::Uint32});
{"num_heads", ProgramUniformVariableDataType::Uint32},
{"batch_size", ProgramUniformVariableDataType::Uint32},
{"num_q_tiles", ProgramUniformVariableDataType::Uint32});

private:
bool has_past_;
Expand Down Expand Up @@ -138,11 +142,14 @@
int max_k_step_;
};

class FlashAttentionDecodeQKTProgram final : public Program<FlashAttentionDecodeQKTProgram> {
class FlashAttentionDecodeQKVProgram final : public Program<FlashAttentionDecodeQKVProgram> {
public:
FlashAttentionDecodeQKTProgram(const std::string& kernel_name,
bool has_attention_bias, uint32_t tile_size, bool use_indirect_dispatch)
: Program{kernel_name}, has_attention_bias_(has_attention_bias), tile_size_(tile_size), use_indirect_dispatch_(use_indirect_dispatch) {
FlashAttentionDecodeQKVProgram(const std::string& kernel_name,
bool has_attention_bias, uint32_t tile_size, int head_size_vec,
bool use_indirect_dispatch, bool q_BNSH = false,
bool is_unidirectional = false,
uint32_t m_tile = 1)
: Program{kernel_name}, has_attention_bias_(has_attention_bias), tile_size_(tile_size), head_size_vec_(head_size_vec), use_indirect_dispatch_(use_indirect_dispatch), q_BNSH_(q_BNSH), is_unidirectional_(is_unidirectional), m_tile_(m_tile) {
}

Status GenerateShaderCode(ShaderHelper& sh) const override;
Expand All @@ -156,41 +163,23 @@
{"num_heads", ProgramUniformVariableDataType::Uint32},
{"batch_size", ProgramUniformVariableDataType::Uint32},
{"attn_bias_dim0", ProgramUniformVariableDataType::Uint32},
{"attn_bias_dim1", ProgramUniformVariableDataType::Uint32});
{"attn_bias_dim1", ProgramUniformVariableDataType::Uint32},
{"new_sequence_length", ProgramUniformVariableDataType::Uint32});

private:
bool has_attention_bias_;
uint32_t tile_size_;
bool use_indirect_dispatch_;
};

class FlashAttentionDecodeSplitVxProgram final : public Program<FlashAttentionDecodeSplitVxProgram> {
public:
FlashAttentionDecodeSplitVxProgram(const std::string& kernel_name, uint32_t tile_size, int head_size_vec, bool use_indirect_dispatch, bool has_head_sink = false)
: Program{kernel_name}, tile_size_(tile_size), head_size_vec_(head_size_vec), use_indirect_dispatch_(use_indirect_dispatch), has_head_sink_(has_head_sink) {
}

Status GenerateShaderCode(ShaderHelper& sh) const override;

WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"total_sequence_length", ProgramUniformVariableDataType::Uint32},
{"head_size_vec", ProgramUniformVariableDataType::Uint32},
{"present_sequence_length", ProgramUniformVariableDataType::Uint32},
{"n_reps", ProgramUniformVariableDataType::Uint32},
{"num_present_sequence_length_tile", ProgramUniformVariableDataType::Uint32},
{"batch_heads", ProgramUniformVariableDataType::Uint32},
{"num_heads", ProgramUniformVariableDataType::Uint32});

private:
uint32_t tile_size_;
int head_size_vec_;
bool use_indirect_dispatch_;
bool has_head_sink_;
bool q_BNSH_;
bool is_unidirectional_;
uint32_t m_tile_;
};

class FlashAttentionDecodeVxReduceProgram final : public Program<FlashAttentionDecodeVxReduceProgram> {
public:
FlashAttentionDecodeVxReduceProgram(const std::string& kernel_name, uint32_t tile_size, uint32_t seq_tile_size, bool use_indirect_dispatch)
: Program{kernel_name}, tile_size_(tile_size), seq_tile_size_(seq_tile_size), use_indirect_dispatch_(use_indirect_dispatch) {
FlashAttentionDecodeVxReduceProgram(const std::string& kernel_name, uint32_t tile_size, uint32_t seq_tile_size, bool use_indirect_dispatch, bool has_head_sink = false, uint32_t m_tile = 1)

Check warning on line 181 in onnxruntime/contrib_ops/webgpu/bert/flash_attention.h

View workflow job for this annotation

GitHub Actions / Optional Lint C++

[cpplint] reported by reviewdog 🐶 Add #include <string> for string [build/include_what_you_use] [4] Raw Output: onnxruntime/contrib_ops/webgpu/bert/flash_attention.h:181: Add #include <string> for string [build/include_what_you_use] [4]
: Program{kernel_name}, tile_size_(tile_size), seq_tile_size_(seq_tile_size), use_indirect_dispatch_(use_indirect_dispatch), has_head_sink_(has_head_sink), m_tile_(m_tile) {
}

Status GenerateShaderCode(ShaderHelper& sh) const override;
Expand All @@ -199,12 +188,16 @@
{"num_total_seq_length_tile", ProgramUniformVariableDataType::Uint32},
{"num_present_sequence_length_tile", ProgramUniformVariableDataType::Uint32},
{"num_head_size_tile", ProgramUniformVariableDataType::Uint32},
{"batch_heads", ProgramUniformVariableDataType::Uint32});
{"batch_heads", ProgramUniformVariableDataType::Uint32},
{"new_sequence_length", ProgramUniformVariableDataType::Uint32},
{"num_heads", ProgramUniformVariableDataType::Uint32});

private:
uint32_t tile_size_;
uint32_t seq_tile_size_;
bool use_indirect_dispatch_;
bool has_head_sink_;
uint32_t m_tile_;
};

Status ApplyFlashAttention(const Tensor* Q, const Tensor* K, const Tensor* V, const Tensor* attention_bias,
Expand All @@ -225,7 +218,7 @@
Tensor* present_key,
Tensor* present_value,
Tensor* indirect_buffer,
uint32_t tile_size);
uint32_t tile_size, uint32_t num_q_tiles);
} // namespace webgpu
} // namespace contrib
} // namespace onnxruntime

This file was deleted.

Loading
Loading