diff --git a/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.cc b/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.cc index 61f701f7911a7..2126022f8b547 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.cc @@ -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" @@ -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_) { @@ -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" : (components == 2 ? "vec2" : "f32")) << ";\n" - << "var sum_shared : array;\n" - << "var sum_squared_shared : array;\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 sum_shared : array;\n" + << "var sum_squared_shared : array;\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(0);\n" + << " var sum_squared_vec4 = vec4(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(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" : (components == 2 ? "vec2" : "f32")) << ";\n" + << "var sum_shared : array;\n" + << "var sum_squared_shared : array;\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(); } @@ -100,14 +156,15 @@ Status SkipLayerNorm::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(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(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}}) @@ -123,6 +180,13 @@ Status SkipLayerNorm::ComputeInternal(onnxruntime::webgpu::ComputeCo {static_cast(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}); } diff --git a/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.h b/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.h index 03de1a4b568b9..73f02f0ad8ec0 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.h +++ b/onnxruntime/contrib_ops/webgpu/bert/skip_layer_norm.h @@ -15,7 +15,7 @@ using onnxruntime::webgpu::ComputeContext; class SkipLayerNormProgram final : public Program { 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; @@ -23,7 +23,7 @@ class SkipLayerNormProgram final : public Program { 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; @@ -39,8 +39,8 @@ class SkipLayerNormProgram final : public Program { float epsilon_; uint32_t hidden_size_; bool has_input_skip_bias_sum_; - bool is_fp16_; bool simplified_; + bool split_hidden_dim_; }; template