mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
[WebNN] Fix bug in SkipSimplifiedLayerNormalization (#23236)
The input should be added by skip and bias (if it exits) firstly.
This commit is contained in:
parent
655b3efee4
commit
519fae019b
1 changed files with 19 additions and 17 deletions
|
|
@ -76,8 +76,6 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder
|
|||
options.set("epsilon", epsilon);
|
||||
|
||||
emscripten::val output = emscripten::val::undefined();
|
||||
// SkipSimplifiedLayerNormalization's output: input_skip_bias_sum.
|
||||
emscripten::val input_skip_bias_sum = emscripten::val::undefined();
|
||||
if (op_type == "BatchNormalization") {
|
||||
ORT_RETURN_IF_NOT(input_defs.size() == 5, "BatchNormalization requires five inputs.");
|
||||
emscripten::val mean = model_builder.GetOperand(input_defs[3]->Name());
|
||||
|
|
@ -107,7 +105,7 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder
|
|||
| | | | | |
|
||||
Y:2 axis B:epsilon A:X A:scale B:bias
|
||||
|
||||
If it is SkipSimplifiedLayerNormalization and its output input_skip_bias_sum exists,
|
||||
If it is SkipSimplifiedLayerNormalization, X should be input_skip_bias_sum:
|
||||
input_skip_bias_sum = X + skip + bias (if it exists)
|
||||
*/
|
||||
|
||||
|
|
@ -115,6 +113,23 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder
|
|||
ORT_RETURN_IF_NOT(GetType(*input_defs[0], input_type, logger), "Cannot get input type");
|
||||
emscripten::val common_options = emscripten::val::object();
|
||||
|
||||
// If it is SkipSimplifiedLayerNormalization, add the skip and bias (if it exists) to the input.
|
||||
if (op_type == "SkipSimplifiedLayerNormalization") {
|
||||
emscripten::val skip = model_builder.GetOperand(input_defs[1]->Name());
|
||||
common_options.set("label", node.Name() + "_add_skip");
|
||||
input = model_builder.GetBuilder().call<emscripten::val>("add", input, skip, common_options);
|
||||
if (!bias.isUndefined()) {
|
||||
common_options.set("label", node.Name() + "_add_skip_bias");
|
||||
input = model_builder.GetBuilder().call<emscripten::val>("add", input, bias, common_options);
|
||||
}
|
||||
|
||||
// Add SkipSimplifiedLayerNormalization's output input_skip_bias_sum if it exists.
|
||||
// Now input equals to input_skip_bias_sum.
|
||||
if (TensorExists(output_defs, 3)) {
|
||||
model_builder.AddOperand(output_defs[3]->Name(), input);
|
||||
}
|
||||
}
|
||||
|
||||
// Pow
|
||||
emscripten::val pow_constant = model_builder.CreateOrGetConstant<float>(input_type, 2);
|
||||
common_options.set("label", node.Name() + "_pow");
|
||||
|
|
@ -146,24 +161,11 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder
|
|||
common_options.set("label", node.Name() + "_mul");
|
||||
output = model_builder.GetBuilder().call<emscripten::val>("mul", scale, div, common_options);
|
||||
|
||||
// Add (if bias exits)
|
||||
// Add (if bias exists)
|
||||
if (!bias.isUndefined()) {
|
||||
common_options.set("label", node.Name() + "_add_bias");
|
||||
output = model_builder.GetBuilder().call<emscripten::val>("add", output, bias, common_options);
|
||||
}
|
||||
|
||||
// SkipSimplifiedLayerNormalization's output input_skip_bias_sum is the sum of input, skip, and bias.
|
||||
if (op_type == "SkipSimplifiedLayerNormalization" && TensorExists(output_defs, 3)) {
|
||||
emscripten::val skip = model_builder.GetOperand(input_defs[1]->Name());
|
||||
common_options.set("label", node.Name() + "_add_skip");
|
||||
input_skip_bias_sum = model_builder.GetBuilder().call<emscripten::val>("add", input, skip, common_options);
|
||||
if (!bias.isUndefined()) {
|
||||
common_options.set("label", node.Name() + "_add_skip_bias");
|
||||
input_skip_bias_sum = model_builder.GetBuilder().call<emscripten::val>(
|
||||
"add", input_skip_bias_sum, bias, common_options);
|
||||
}
|
||||
model_builder.AddOperand(output_defs[3]->Name(), std::move(input_skip_bias_sum));
|
||||
}
|
||||
}
|
||||
} else if (op_type == "InstanceNormalization") {
|
||||
// WebNN spec only supports 4D input for instanceNormalization.
|
||||
|
|
|
|||
Loading…
Reference in a new issue