Skip to content

Commit 8b9755d

Browse files
feich-msclaude
andcommitted
Reuse RotaryEmbeddingProgram for GQA kv_empty rotary path
Remove RotaryEmbeddingWithOffsetProgram and instead pass a scalar position_ids tensor ([1,1] with value=past_sequence_length) to the existing RotaryEmbeddingProgram. The shader's broadcast logic already computes position_id = raw_pos + bsnh[1] when position_ids is scalar, eliminating the need for a separate class. Co-Authored-By: Claude Opus 4 <noreply@anthropic.com>
1 parent 5ca4c9e commit 8b9755d

4 files changed

Lines changed: 23 additions & 55 deletions

File tree

onnxruntime/contrib_ops/webgpu/bert/group_query_attention.cc

Lines changed: 22 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
#include "contrib_ops/webgpu/webgpu_contrib_kernels.h"
88
#include "contrib_ops/webgpu/bert/rotary_embedding.h"
99
#include "contrib_ops/webgpu/bert/flash_attention.h"
10+
#include "core/providers/webgpu/generator/range.h"
1011

1112
#include "core/common/narrow.h"
1213
#include "core/providers/webgpu/webgpu_supported_types.h"
@@ -183,8 +184,8 @@ Status RunFusedQKRotaryEmbedding(onnxruntime::webgpu::ComputeContext& context,
183184
return context.RunProgram(program);
184185
}
185186

186-
// Apply rotary embedding to a single tensor using RotaryEmbeddingWithOffsetProgram.
187-
// Position for each token = past_sequence_length + sequence_index.
187+
// Apply rotary embedding to a single tensor using RotaryEmbeddingProgram with a scalar position_ids.
188+
// Position for each token = past_sequence_length + sequence_index (computed in shader via broadcast).
188189
Status RunRotaryEmbedding(onnxruntime::webgpu::ComputeContext& context,
189190
const Tensor* input,
190191
const Tensor* cos_cache,
@@ -218,19 +219,35 @@ Status RunRotaryEmbedding(onnxruntime::webgpu::ComputeContext& context,
218219
gsl::narrow_cast<uint32_t>(head_size),
219220
1u});
220221

221-
RotaryEmbeddingWithOffsetProgram program(rotary_interleaved);
222+
// Create scalar position_ids [1,1] with value = past_sequence_length.
223+
// The shader broadcasts this to all threads and adds bsnh[1] (sequence index).
224+
const TensorShape pos_ids_shape({1, 1});
225+
Tensor pos_ids_tensor = context.CreateGPUTensor(DataTypeImpl::GetType<int64_t>(), pos_ids_shape);
226+
{
227+
RangeProgram range_program{ONNX_NAMESPACE::TensorProto_DataType_INT64};
228+
int32_t start_i32 = static_cast<int32_t>(past_sequence_length);
229+
int32_t delta_i32 = 1;
230+
range_program
231+
.AddOutput({&pos_ids_tensor, ProgramTensorMetadataDependency::Type})
232+
.SetDispatchGroupSize(1)
233+
.AddUniformVariables({1u, std::bit_cast<uint32_t>(start_i32), std::bit_cast<uint32_t>(delta_i32)});
234+
ORT_RETURN_IF_ERROR(context.RunProgram(range_program));
235+
}
236+
237+
RotaryEmbeddingProgram program(rotary_interleaved);
222238
program
223239
.CacheHint(rotary_interleaved)
224240
.AddInputs({{input, ProgramTensorMetadataDependency::TypeAndRank},
241+
{&pos_ids_tensor, ProgramTensorMetadataDependency::Rank},
225242
{cos_cache, ProgramTensorMetadataDependency::Rank},
226243
{sin_cache, ProgramTensorMetadataDependency::Rank}})
227244
.AddOutput({output, ProgramTensorMetadataDependency::None})
228245
.SetDispatchGroupSize((output_size + WORKGROUP_SIZE - 1) / WORKGROUP_SIZE)
229246
.AddUniformVariables({{scale},
230247
{gsl::make_span(global_dims)},
231248
{gsl::make_span(global_strides)},
232-
{gsl::make_span(input_output_strides)},
233-
{static_cast<uint32_t>(past_sequence_length)}});
249+
{gsl::make_span(input_output_strides)}})
250+
.AddIndices(TensorShape{1, 1});
234251
return context.RunProgram(program);
235252
}
236253

onnxruntime/contrib_ops/webgpu/bert/rotary_embedding.cc

Lines changed: 0 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -134,38 +134,6 @@ Status FusedQKRotaryEmbeddingProgram::GenerateShaderCode(ShaderHelper& shader) c
134134
return Status::OK();
135135
}
136136

137-
Status RotaryEmbeddingWithOffsetProgram::GenerateShaderCode(ShaderHelper& shader) const {
138-
const auto& input = shader.AddInput("input", ShaderUsage::UseUniform);
139-
const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform);
140-
const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform);
141-
const auto& output = shader.AddOutput("output", ShaderUsage::UseUniform);
142-
const auto interleaved_str = interleaved_ ? "true" : "false";
143-
shader.MainFunctionBody() << " let half_rotary_emb_dim = uniforms.cos_cache_shape[1];\n"
144-
" let bsnh = global_idx / uniforms.global_stride % uniforms.global_shape;\n"
145-
" let size = uniforms.global_shape[0] * uniforms.global_stride[0];\n"
146-
" if (global_idx >= size) { return; }\n"
147-
" if (bsnh[3] < half_rotary_emb_dim) {\n"
148-
" let position_id = uniforms.position_offset + bsnh[1];\n"
149-
<< " let i = dot(bsnh, uniforms.input_output_stride) + select(0, bsnh[3], " << interleaved_str << ");\n"
150-
<< " let j = i + select(half_rotary_emb_dim, 1, " << interleaved_str << ");\n"
151-
" let max_position = uniforms.cos_cache_shape[0];\n"
152-
" if (position_id >= max_position) {\n"
153-
<< " " << output.SetByOffset("i", input.GetByOffset("i")) << "\n"
154-
<< " " << output.SetByOffset("j", input.GetByOffset("j")) << "\n"
155-
" } else {\n"
156-
<< " let re = " << input.GetByOffset("i") << " * " << cos_cache.GetByIndices("vec2<u32>(position_id, bsnh[3])") << " - " << input.GetByOffset("j") << " * " << sin_cache.GetByIndices("vec2<u32>(position_id, bsnh[3])") << ";\n"
157-
<< " " << output.SetByOffset("i", "re") << "\n"
158-
<< " let im = " << input.GetByOffset("i") << " * " << sin_cache.GetByIndices("vec2<u32>(position_id, bsnh[3])") << " + " << input.GetByOffset("j") << " * " << cos_cache.GetByIndices("vec2<u32>(position_id, bsnh[3])") << ";\n"
159-
<< " " << output.SetByOffset("j", "im") << "\n"
160-
" }\n"
161-
<< " } else { \n"
162-
" let k = dot(bsnh, uniforms.input_output_stride) + half_rotary_emb_dim;\n"
163-
<< " " << output.SetByOffset("k", input.GetByOffset("k")) << "\n"
164-
<< " }";
165-
166-
return Status::OK();
167-
}
168-
169137
RotaryEmbedding::RotaryEmbedding(const OpKernelInfo& info) : WebGpuKernel(info) {
170138
scale_ = info.GetAttrOrDefault<float>("scale", 1.0);
171139
rotary_embedding_dim_ = static_cast<int>(info.GetAttrOrDefault<int64_t>("rotary_embedding_dim", 0));

onnxruntime/contrib_ops/webgpu/bert/rotary_embedding.h

Lines changed: 0 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -29,23 +29,6 @@ class RotaryEmbeddingProgram final : public Program<RotaryEmbeddingProgram> {
2929
const bool interleaved_;
3030
};
3131

32-
class RotaryEmbeddingWithOffsetProgram final : public Program<RotaryEmbeddingWithOffsetProgram> {
33-
public:
34-
RotaryEmbeddingWithOffsetProgram(bool interleaved)
35-
: Program{"RotaryEmbeddingWithOffset"}, interleaved_{interleaved} {}
36-
37-
Status GenerateShaderCode(ShaderHelper& sh) const override;
38-
39-
WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"scale", ProgramUniformVariableDataType::Float32},
40-
{"global_shape", ProgramUniformVariableDataType::Uint32},
41-
{"global_stride", ProgramUniformVariableDataType::Uint32},
42-
{"input_output_stride", ProgramUniformVariableDataType::Uint32},
43-
{"position_offset", ProgramUniformVariableDataType::Uint32});
44-
45-
private:
46-
const bool interleaved_;
47-
};
48-
4932
class FusedQKRotaryEmbeddingProgram final : public Program<FusedQKRotaryEmbeddingProgram> {
5033
public:
5134
FusedQKRotaryEmbeddingProgram(bool interleaved) : Program{"FusedQKRotaryEmbedding"}, interleaved_{interleaved} {}

onnxruntime/test/contrib_ops/group_query_attention_op_test.cc

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1888,7 +1888,7 @@ TEST(GroupQueryAttentionTest, WebGPU_SharedKV_Rotary_Prefill) {
18881888
}
18891889

18901890
// WebGPU: kv_sequence_length=0 with do_rotary=1 and batch_size > 1.
1891-
// Validates batch stride calculations in RotaryEmbeddingWithOffsetProgram.
1891+
// Validates batch stride calculations in the rotary embedding path.
18921892
TEST(GroupQueryAttentionTest, WebGPU_SharedKV_Rotary_MultiBatch) {
18931893
auto webgpu_ep = DefaultWebGpuExecutionProvider();
18941894
if (!webgpu_ep) {

0 commit comments

Comments
 (0)