Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
172 changes: 118 additions & 54 deletions onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
// Licensed under the MIT License.

#include "core/providers/webgpu/shader_helper.h"
#include "core/providers/webgpu/string_macros.h"
#include "core/providers/webgpu/webgpu_utils.h"
#include "core/providers/webgpu/webgpu_supported_types.h"
#include "contrib_ops/webgpu/webgpu_contrib_kernels.h"
Expand All @@ -12,7 +13,7 @@ namespace contrib {
namespace webgpu {

Status SkipLayerNormProgram::GenerateShaderCode(ShaderHelper& shader) const {
const auto& x = shader.AddInput("x", ShaderUsage::UseUniform | ShaderUsage::UseValueTypeAlias);
const auto& x = shader.AddInput("x", ShaderUsage::UseUniform | ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias);
shader.AddInput("skip", ShaderUsage::UseUniform);
shader.AddInput("gamma", ShaderUsage::UseUniform);
if (hasBeta_) {
Expand All @@ -26,57 +27,112 @@ Status SkipLayerNormProgram::GenerateShaderCode(ShaderHelper& shader) const {
shader.AddOutput("input_skip_bias_sum", ShaderUsage::UseUniform);
}

int components = x.NumComponents();

std::string bias = (hasBias_) ? " + bias[offset1d + i] " : "";
std::string simpl1 = (simplified_) ? "" : "- mean * mean ";
std::string simpl2 = (simplified_) ? "" : "- element_t(mean) ";
std::string beta = (hasBeta_) ? " + beta[offset1d + i] " : "";
std::string input_skip_bias_sum = (has_input_skip_bias_sum_) ? "input_skip_bias_sum[offset + i] = value;\n" : "";

shader.AdditionalImplementation()
<< "alias element_t = " << (is_fp16_ ? "f16;\n" : "f32;\n")
<< "alias f32_val_t = " << (components == 4 ? "vec4<f32>" : (components == 2 ? "vec2<f32>" : "f32")) << ";\n"
<< "var<workgroup> sum_shared : array<f32_val_t, workgroup_size_x>;\n"
<< "var<workgroup> sum_squared_shared : array<f32_val_t, workgroup_size_x>;\n";

shader.MainFunctionBody()
<< "let ix = local_idx;\n"
<< "let iy = global_idx / workgroup_size_x;\n"
<< "let hidden_size_vectorized: u32 = uniforms.hidden_size / uniforms.components;\n"
<< "var stride = hidden_size_vectorized / workgroup_size_x;\n"
<< "let offset = ix * stride + iy * hidden_size_vectorized;\n"
<< "let offset1d = stride * ix;\n"
<< "if (ix == workgroup_size_x - 1) {\n"
<< " stride = hidden_size_vectorized - stride * ix;\n"
<< "}\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " let skip_value = skip[offset + i];\n"
<< " let input_value = x[offset + i];\n"
<< " let value = input_value + skip_value" << bias << ";\n"
<< " output[offset + i] = value;\n"
<< input_skip_bias_sum
<< " let f32_value = f32_val_t(value);\n"
<< " sum_shared[ix] += f32_value;\n"
<< " sum_squared_shared[ix] += f32_value * f32_value;\n"
<< "}\n"
<< "workgroupBarrier();\n"
<< "var reduce_size : u32 = workgroup_size_x;\n"
<< "for (var curr_size = reduce_size >> 1; curr_size > 0; curr_size = reduce_size >> 1) {\n"
<< " reduce_size = curr_size + (reduce_size & 1);\n"
<< " if (ix < curr_size) {\n"
<< " sum_shared[ix] += sum_shared[ix + reduce_size];\n"
<< " sum_squared_shared[ix] += sum_squared_shared[ix + reduce_size];\n"
<< " }\n"
<< " workgroupBarrier();\n"
<< "}\n"
<< "let sum = sum_shared[0];\n"
<< "let square_sum = sum_squared_shared[0];\n"
<< "let mean = " << SumVector("sum", components) << " / f32(uniforms.hidden_size);\n"
<< "let inv_std_dev = inverseSqrt(" << SumVector("square_sum", components) << " / f32(uniforms.hidden_size) " << simpl1 << "+ uniforms.epsilon);\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " output[offset + i] = (output[offset + i] " << simpl2 << ") * element_t(inv_std_dev) * gamma[offset1d + i]" << beta << ";\n"
<< "};\n";
std::string simpl2 = (simplified_) ? "" : "- x_element_t(mean) ";
if (split_hidden_dim_) {
shader.AdditionalImplementation()
<< "var<workgroup> sum_shared : array<f32, workgroup_size_x>;\n"
<< "var<workgroup> sum_squared_shared : array<f32, workgroup_size_x>;\n";

SS(input_skip_bias_sum_ss, 512);
if (has_input_skip_bias_sum_) {
input_skip_bias_sum_ss
<< " let workgroup_half_idx = uniforms.hidden_size / (workgroup_size_x * 4);\n"
<< " if (workgroup_idx >= workgroup_half_idx) {\n"
<< " offset = (workgroup_idx - workgroup_half_idx) * workgroup_size_x + local_idx;\n"
<< " let skip_value = skip[offset];\n"
<< " let input_value = x[offset];\n"
<< " let value = input_value + skip_value" << (hasBias_ ? " + bias[offset]" : "") << ";\n"
<< " input_skip_bias_sum[offset] = value;\n"
<< " return;\n"
<< " }\n";
}

shader.MainFunctionBody()
<< " var offset: u32 = 0;\n"
<< (has_input_skip_bias_sum_ ? SS_GET(input_skip_bias_sum_ss) : "")
<< " var sum_vec4 = vec4<f32>(0);\n"
<< " var sum_squared_vec4 = vec4<f32>(0);\n"
<< " var cur_input_skip_bias_sum = x_value_t(0);\n"
<< " for (var i: u32 = 0; i < uniforms.hidden_size / (workgroup_size_x * 4); i++) {\n"
<< " let input_offset = i * workgroup_size_x + local_idx;\n"
<< " let skip_value = skip[input_offset];\n"
<< " let input_value = x[input_offset];\n"
<< " let value = input_value + skip_value" << (hasBias_ ? " + bias[input_offset]" : "") << ";\n"
<< " if (i == workgroup_idx) {\n"
<< " cur_input_skip_bias_sum = value;\n"
<< " }\n"
<< " let f32_value = vec4<f32>(value);\n"
<< " sum_vec4 += f32_value;\n"
<< " sum_squared_vec4 += f32_value * f32_value;\n"
<< " }\n"
<< " var sum = " << SumVector("sum_vec4", 4) << ";\n"
<< " var sum_squared = " << SumVector("sum_squared_vec4", 4) << ";\n"
<< " sum_shared[local_idx] = sum;\n"
<< " sum_squared_shared[local_idx] = sum_squared;\n"
<< " workgroupBarrier();\n"
<< " var reduce_size : u32 = workgroup_size_x;\n"
<< " for (var curr_size = reduce_size >> 1; curr_size > 0; curr_size = reduce_size >> 1) {\n"
<< " reduce_size = curr_size + (reduce_size & 1);\n"
<< " if (local_idx < curr_size) {\n"
<< " sum_shared[local_idx] += sum_shared[local_idx + reduce_size];\n"
<< " sum_squared_shared[local_idx] += sum_squared_shared[local_idx + reduce_size];\n"
<< " }\n"
<< " workgroupBarrier();\n"
<< " }\n"
<< " let mean = sum_shared[0] / f32(uniforms.hidden_size);\n"
<< " let inv_std_dev = inverseSqrt(sum_squared_shared[0] / f32(uniforms.hidden_size) " << simpl1 << "+ uniforms.epsilon);\n"
<< " offset = workgroup_idx * workgroup_size_x + local_idx;\n"
<< " output[offset] = ((cur_input_skip_bias_sum " << simpl2 << ") * x_element_t(inv_std_dev) * gamma[offset]" << (hasBeta_ ? " + beta[offset] " : "") << ");\n";
} else {
int components = x.NumComponents();
std::string bias = (hasBias_) ? " + bias[offset1d + i] " : "";
std::string beta = (hasBeta_) ? " + beta[offset1d + i] " : "";
std::string input_skip_bias_sum = (has_input_skip_bias_sum_) ? "input_skip_bias_sum[offset + i] = value;\n" : "";

shader.AdditionalImplementation()
<< "alias f32_val_t = " << (components == 4 ? "vec4<f32>" : (components == 2 ? "vec2<f32>" : "f32")) << ";\n"
<< "var<workgroup> sum_shared : array<f32_val_t, workgroup_size_x>;\n"
<< "var<workgroup> sum_squared_shared : array<f32_val_t, workgroup_size_x>;\n";

shader.MainFunctionBody()
<< "let ix = local_idx;\n"
<< "let iy = global_idx / workgroup_size_x;\n"
<< "let hidden_size_vectorized: u32 = uniforms.hidden_size / uniforms.components;\n"
<< "var stride = hidden_size_vectorized / workgroup_size_x;\n"
<< "let offset = ix * stride + iy * hidden_size_vectorized;\n"
<< "let offset1d = stride * ix;\n"
<< "if (ix == workgroup_size_x - 1) {\n"
<< " stride = hidden_size_vectorized - stride * ix;\n"
<< "}\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " let skip_value = skip[offset + i];\n"
<< " let input_value = x[offset + i];\n"
<< " let value = input_value + skip_value" << bias << ";\n"
<< " output[offset + i] = value;\n"
<< input_skip_bias_sum
<< " let f32_value = f32_val_t(value);\n"
<< " sum_shared[ix] += f32_value;\n"
<< " sum_squared_shared[ix] += f32_value * f32_value;\n"
<< "}\n"
<< "workgroupBarrier();\n"
<< "var reduce_size : u32 = workgroup_size_x;\n"
<< "for (var curr_size = reduce_size >> 1; curr_size > 0; curr_size = reduce_size >> 1) {\n"
<< " reduce_size = curr_size + (reduce_size & 1);\n"
<< " if (ix < curr_size) {\n"
<< " sum_shared[ix] += sum_shared[ix + reduce_size];\n"
<< " sum_squared_shared[ix] += sum_squared_shared[ix + reduce_size];\n"
<< " }\n"
<< " workgroupBarrier();\n"
<< "}\n"
<< "let sum = sum_shared[0];\n"
<< "let square_sum = sum_squared_shared[0];\n"
<< "let mean = " << SumVector("sum", components) << " / f32(uniforms.hidden_size);\n"
<< "let inv_std_dev = inverseSqrt(" << SumVector("square_sum", components) << " / f32(uniforms.hidden_size) " << simpl1 << "+ uniforms.epsilon);\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " output[offset + i] = (output[offset + i] " << simpl2 << ") * x_element_t(inv_std_dev) * gamma[offset1d + i]" << beta << ";\n"
<< "};\n";
}

return Status::OK();
}
Expand All @@ -100,14 +156,15 @@ Status SkipLayerNorm<simplified>::ComputeInternal(onnxruntime::webgpu::ComputeCo
return Status::OK();
}

const bool is_fp16 = x->GetElementType() == ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16;
const uint32_t hidden_size = onnxruntime::narrow<uint32_t>(x_shape[x_shape.NumDimensions() - 1]);
const int components = GetMaxComponents(hidden_size);
const bool has_input_skip_bias_sum = input_skip_bias_sum != nullptr;
const uint32_t norm_count = onnxruntime::narrow<uint32_t>(x_shape.SizeToDimension(x_shape.NumDimensions() - 1));
const bool split_hidden_dim = hidden_size % 512 == 0 && norm_count == 1;

SkipLayerNormProgram program{beta != nullptr, bias != nullptr, epsilon_, hidden_size, has_input_skip_bias_sum, is_fp16, simplified};
SkipLayerNormProgram program{beta != nullptr, bias != nullptr, epsilon_, hidden_size, has_input_skip_bias_sum, simplified, split_hidden_dim};
program
.CacheHint(simplified, has_input_skip_bias_sum)
.CacheHint(simplified, has_input_skip_bias_sum, split_hidden_dim)
.AddInputs({{x, ProgramTensorMetadataDependency::Type, components}})
.AddInputs({{skip, ProgramTensorMetadataDependency::Type, components}})
.AddInputs({{gamma, ProgramTensorMetadataDependency::Type, components}})
Expand All @@ -123,6 +180,13 @@ Status SkipLayerNorm<simplified>::ComputeInternal(onnxruntime::webgpu::ComputeCo
{static_cast<float>(epsilon_)},
});

if (split_hidden_dim) {
const uint32_t workgroup_size_x = 128;
const uint32_t dispatch_size_x = (has_input_skip_bias_sum ? 2 : 1) * hidden_size / (workgroup_size_x * components);
program.SetDispatchGroupSize(dispatch_size_x, 1, 1)
.SetWorkgroupSize(workgroup_size_x);
}

if (beta != nullptr) {
program.AddInput({beta, ProgramTensorMetadataDependency::Type, components});
}
Expand Down
6 changes: 3 additions & 3 deletions onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,15 +15,15 @@ using onnxruntime::webgpu::ComputeContext;

class SkipLayerNormProgram final : public Program<SkipLayerNormProgram> {
public:
SkipLayerNormProgram(bool hasBeta, bool hasBias, float epsilon, uint32_t hidden_size, bool has_input_skip_bias_sum, bool is_fp16, bool simplified) : Program{"SkipLayerNorm"} {
SkipLayerNormProgram(bool hasBeta, bool hasBias, float epsilon, uint32_t hidden_size, bool has_input_skip_bias_sum, bool simplified, bool split_hidden_dim) : Program{"SkipLayerNorm"} {
epsilon_ = epsilon;
hasBeta_ = hasBeta;
hasBias_ = hasBias;
epsilon_ = epsilon;
hidden_size_ = hidden_size;
has_input_skip_bias_sum_ = has_input_skip_bias_sum;
simplified_ = simplified;
is_fp16_ = is_fp16;
split_hidden_dim_ = split_hidden_dim;
}

Status GenerateShaderCode(ShaderHelper& sh) const override;
Expand All @@ -39,8 +39,8 @@ class SkipLayerNormProgram final : public Program<SkipLayerNormProgram> {
float epsilon_;
uint32_t hidden_size_;
bool has_input_skip_bias_sum_;
bool is_fp16_;
bool simplified_;
bool split_hidden_dim_;
};

template <bool simplified>
Expand Down