Skip to content
Open
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
118 changes: 9 additions & 109 deletions onnxruntime/core/providers/webgpu/nn/layer_norm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#include "core/providers/webgpu/shader_helper.h"
#include "core/providers/webgpu/webgpu_supported_types.h"
#include "core/providers/webgpu/webgpu_utils.h"
#include "core/providers/webgpu/wgsl_templates/wgsl_gen.h"
#include "core/providers/webgpu/nn/layer_norm.h"

namespace onnxruntime {
Expand Down Expand Up @@ -39,116 +40,15 @@ Status LayerNormProgram::GenerateShaderCode(ShaderHelper& shader) const {
shader.AddOutput("inv_std_dev_output", ShaderUsage::None);
}

std::string simpl1 = (simplified_) ? "" : "- mean * mean ";
std::string simpl2 = (simplified_) ? "" : "- x_element_t(mean) ";
int components = x.NumComponents();

if (split_norm_dim_) {
shader.AdditionalImplementation()
<< "var<workgroup> sum_shared : array<f32, workgroup_size_x>;\n"
<< "var<workgroup> sum_squared_shared : array<f32, workgroup_size_x>;\n";

shader.MainFunctionBody()
<< " var sum_vec4 = vec4<f32>(0);\n"
<< " var sum_squared_vec4 = vec4<f32>(0);\n"
<< " var cur_input = x_value_t(0);\n"
<< " for (var i: u32 = 0; i < uniforms.norm_size / (workgroup_size_x * 4); i++) {\n"
<< " let input_offset = i * workgroup_size_x + local_idx;\n"
<< " let input_value = x[input_offset];\n"
<< " if (i == workgroup_idx) {\n"
<< " cur_input = input_value;\n"
<< " }\n"
<< " let f32_value = vec4<f32>(input_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.norm_size);\n"
<< " let inv_std_dev = inverseSqrt(sum_squared_shared[0] / f32(uniforms.norm_size) " << simpl1 << "+ uniforms.epsilon);\n"
<< " let offset = workgroup_idx * workgroup_size_x + local_idx;\n"
<< " y[offset] = ((cur_input " << simpl2 << ") * x_element_t(inv_std_dev) * scale[offset]" << (has_bias_ ? " + bias[offset] " : "") << ");\n";

if (has_mean_output_) {
shader.MainFunctionBody() << " if (local_idx == 0 && workgroup_idx == 0) {\n"
<< " mean_output[global_idx / uniforms.norm_size] = mean;\n"
<< " }\n";
}
if (has_inv_std_dev_output_) {
shader.MainFunctionBody() << " if (local_idx == 0 && workgroup_idx == 0) {\n"
<< " inv_std_dev_output[global_idx / uniforms.norm_size] = inv_std_dev;\n"
<< " }\n";
}
} else {
int components = x.NumComponents();
std::string bias = (has_bias_) ? " + bias[offset1d + i] " : "";

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 norm_size_vectorized: u32 = uniforms.norm_size / uniforms.components;\n"
<< "var stride = norm_size_vectorized / workgroup_size_x;\n"
<< "let offset = ix * stride + iy * norm_size_vectorized;\n"
<< "let offset1d = stride * ix;\n"
<< "sum_shared[ix] = f32_val_t(0);\n"
<< "sum_squared_shared[ix] = f32_val_t(0);\n"
<< "if (ix == workgroup_size_x - 1) {\n"
<< " stride = norm_size_vectorized - stride * ix;\n"
<< "}\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " let input_value = x[offset + i];\n"
<< " y[offset + i] = input_value;\n"
<< " let f32_value = f32_val_t(input_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.norm_size);\n"
<< "let inv_std_dev = inverseSqrt(" << SumVector("square_sum", components) << " / f32(uniforms.norm_size) " << simpl1 << "+ uniforms.epsilon);\n"
<< "for (var i: u32 = 0; i < stride; i++) {\n"
<< " y[offset + i] = (y[offset + i] " << simpl2 << ") * x_element_t(inv_std_dev) * scale[offset1d + i]" << bias << ";\n"
<< "};\n";

if (has_mean_output_) {
shader.MainFunctionBody() << "if (ix == 0) {\n"
<< " mean_output[iy] = mean;\n"
<< "}\n";
}
if (has_inv_std_dev_output_) {
shader.MainFunctionBody() << "if (ix == 0) {\n"
<< " inv_std_dev_output[iy] = inv_std_dev;\n"
<< "}\n";
}
}

return Status::OK();
return WGSL_TEMPLATE_APPLY(shader, "nn/layer_norm.wgsl.template",
WGSL_TEMPLATE_PARAMETER(components, components),
WGSL_TEMPLATE_PARAMETER(has_bias, has_bias_),
WGSL_TEMPLATE_PARAMETER(has_inv_std_dev_output, has_inv_std_dev_output_),
WGSL_TEMPLATE_PARAMETER(has_mean_output, has_mean_output_),
WGSL_TEMPLATE_PARAMETER(simplified, simplified_),
WGSL_TEMPLATE_PARAMETER(split_norm_dim, split_norm_dim_));
}

template <bool simplified>
Expand Down
158 changes: 158 additions & 0 deletions onnxruntime/core/providers/webgpu/nn/layer_norm.wgsl.template
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#param has_bias
#param simplified
#param has_mean_output
#param has_inv_std_dev_output
#param split_norm_dim
#param components

#if split_norm_dim

var<workgroup> sum_shared : array<f32, workgroup_size_x>;
var<workgroup> sum_squared_shared : array<f32, workgroup_size_x>;

#else

#if components == 4
alias f32_val_t = vec4<f32>;
fn SumVector(v: f32_val_t) -> f32 {
return v.x + v.y + v.z + v.w;
}
#elif components == 2
alias f32_val_t = vec2<f32>;
fn SumVector(v: f32_val_t) -> f32 {
return v.x + v.y;
}
#else
alias f32_val_t = f32;
fn SumVector(v: f32_val_t) -> f32 {
return v;
}
#endif

var<workgroup> sum_shared : array<f32_val_t, workgroup_size_x>;
var<workgroup> sum_squared_shared : array<f32_val_t, workgroup_size_x>;

#endif

$MAIN {
let norm_size : f32 = f32(uniforms.norm_size);
#if split_norm_dim
var sum_vec4 = vec4<f32>(0);
var sum_squared_vec4 = vec4<f32>(0);
var cur_input = x_value_t(0);
for (var i: u32 = 0; i < uniforms.norm_size / (workgroup_size_x * 4); i++) {
let input_offset = i * workgroup_size_x + local_idx;
let input_value = x[input_offset];
if (i == workgroup_idx) {
cur_input = input_value;
}
let f32_value = vec4<f32>(input_value);
sum_vec4 += f32_value;
sum_squared_vec4 += f32_value * f32_value;
}
var sum = (sum_vec4.x + sum_vec4.y + sum_vec4.z + sum_vec4.w);
var sum_squared = (sum_squared_vec4.x + sum_squared_vec4.y + sum_squared_vec4.z + sum_squared_vec4.w);
sum_shared[local_idx] = sum;
sum_squared_shared[local_idx] = sum_squared;
workgroupBarrier();
var reduce_size : u32 = workgroup_size_x;
for (var curr_size = reduce_size >> 1; curr_size > 0; curr_size = reduce_size >> 1) {
reduce_size = curr_size + (reduce_size & 1);
if (local_idx < curr_size) {
sum_shared[local_idx] += sum_shared[local_idx + reduce_size];
sum_squared_shared[local_idx] += sum_squared_shared[local_idx + reduce_size];
}
workgroupBarrier();
}
let mean = sum_shared[0] / norm_size;
let offset = workgroup_idx * workgroup_size_x + local_idx;
#if has_bias
let bias_value = bias[offset];
#else
const bias_value = x_element_t(0.0);
#endif
#if simplified
let inv_std_dev = inverseSqrt(sum_squared_shared[0] / norm_size + uniforms.epsilon);
y[offset] = ((cur_input) * x_element_t(inv_std_dev) * scale[offset] + bias_value);
#else
let inv_std_dev = inverseSqrt(sum_squared_shared[0] / norm_size - mean * mean + uniforms.epsilon);
y[offset] = ((cur_input - x_element_t(mean)) * x_element_t(inv_std_dev) * scale[offset] + bias_value);
#endif
#if has_mean_output
if (local_idx == 0 && workgroup_idx == 0) {
mean_output[global_idx / uniforms.norm_size] = mean;
}
#endif
#if has_inv_std_dev_output
if (local_idx == 0 && workgroup_idx == 0) {
inv_std_dev_output[global_idx / uniforms.norm_size] = inv_std_dev;
}
#endif

#else

let ix = local_idx;
let iy = global_idx / workgroup_size_x;
let norm_size_vectorized: u32 = uniforms.norm_size / components;
var stride = norm_size_vectorized / workgroup_size_x;
let offset = ix * stride + iy * norm_size_vectorized;
let offset1d = stride * ix;
sum_shared[ix] = f32_val_t(0);
sum_squared_shared[ix] = f32_val_t(0);
if (ix == workgroup_size_x - 1) {
stride = norm_size_vectorized - stride * ix;
}
for (var i: u32 = 0; i < stride; i++) {
let input_value = x[offset + i];
y[offset + i] = input_value;
let f32_value = f32_val_t(input_value);
sum_shared[ix] += f32_value;
sum_squared_shared[ix] += f32_value * f32_value;
}
workgroupBarrier();
var reduce_size : u32 = workgroup_size_x;
for (var curr_size = reduce_size >> 1; curr_size > 0; curr_size = reduce_size >> 1) {
reduce_size = curr_size + (reduce_size & 1);
if (ix < curr_size) {
sum_shared[ix] += sum_shared[ix + reduce_size];
sum_squared_shared[ix] += sum_squared_shared[ix + reduce_size];
}
workgroupBarrier();
}
let sum = sum_shared[0];
let square_sum = sum_squared_shared[0];
let mean = SumVector(sum) / norm_size;
#if simplified
let inv_std_dev = inverseSqrt(SumVector(square_sum) / norm_size + uniforms.epsilon);
#else
let inv_std_dev = inverseSqrt(SumVector(square_sum) / norm_size - mean * mean + uniforms.epsilon);
#endif

for (var i: u32 = 0; i < stride; i++) {
#if has_bias
let bias_value = bias[offset1d + i];
#else
const bias_value = x_element_t(0.0);
#endif
#if simplified
y[offset + i] = (y[offset + i]) * x_element_t(inv_std_dev) * scale[offset1d + i] + bias_value;
#else
y[offset + i] = (y[offset + i] - x_element_t(mean)) * x_element_t(inv_std_dev) * scale[offset1d + i] + bias_value;
#endif
};
#if has_mean_output
if (ix == 0) {
mean_output[iy] = mean;
}
#endif
#if has_inv_std_dev_output
if (ix == 0) {
inv_std_dev_output[iy] = inv_std_dev;
}
#endif

#endif
} // MAIN
Loading