From b667cb1f6c3207dcfa54d478edc71e4abc491c6e Mon Sep 17 00:00:00 2001 From: Wanming Lin Date: Thu, 11 Jun 2026 16:09:53 +0800 Subject: [PATCH] [WebNN EP] Remove unnecessary cast around normalization ops MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Remove fp16→fp32→fp16 cast workaround from SimplifiedLayerNormalization and SkipSimplifiedLayerNormalization. The WebNN backend now handles float16 precision correctly, making these intermediate casts unnecessary. --- .../builders/impl/normalization_op_builder.cc | 44 +------------------ 1 file changed, 2 insertions(+), 42 deletions(-) diff --git a/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc index c9cf7bb162870..2d1dd3d15cc7b 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc @@ -96,33 +96,9 @@ 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 (input_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16) { - // Decomposed *SimplifiedLayerNormalization may lose precision if its data type is float16. - // So cast all inputs to float32 to ensure precision. - common_options.set("label", node.Name() + "_cast_input_to_fp32"); - input = model_builder.GetBuilder().call("cast", input, - emscripten::val("float32"), common_options); - - common_options.set("label", node.Name() + "_cast_scale_to_fp32"); - scale = model_builder.GetBuilder().call("cast", scale, - emscripten::val("float32"), common_options); - - if (!bias.isUndefined()) { - common_options.set("label", node.Name() + "_cast_bias_to_fp32"); - bias = model_builder.GetBuilder().call("cast", bias, - emscripten::val("float32"), common_options); - } - } - // 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()); - if (input_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16) { - // Cast skip to float32 - common_options.set("label", node.Name() + "_cast_skip_to_fp32"); - skip = model_builder.GetBuilder().call("cast", skip, - emscripten::val("float32"), common_options); - } common_options.set("label", node.Name() + "_add_skip"); input = model_builder.GetBuilder().call("add", input, skip, common_options); if (!bias.isUndefined()) { @@ -134,20 +110,12 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder // Now input equals to input_skip_bias_sum. if (TensorExists(output_defs, 3)) { emscripten::val input_skip_bias_sum = input; - if (input_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16) { - // Cast input_skip_bias_sum back to float16. - common_options.set("label", node.Name() + "_cast_input_skip_bias_sum_to_fp16"); - input_skip_bias_sum = model_builder.GetBuilder().call("cast", input_skip_bias_sum, - emscripten::val("float16"), - common_options); - } model_builder.AddOperand(output_defs[3]->Name(), input_skip_bias_sum); } } // Pow - emscripten::val pow_constant = - model_builder.CreateOrGetConstant(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, 2); + emscripten::val pow_constant = model_builder.CreateOrGetConstant(input_type, 2); common_options.set("label", node.Name() + "_pow"); emscripten::val pow = model_builder.GetBuilder().call("pow", input, pow_constant, common_options); @@ -160,8 +128,7 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder emscripten::val reduce_mean = model_builder.GetBuilder().call("reduceMean", pow, reduce_options); // Add - emscripten::val add_constant = - model_builder.CreateOrGetConstant(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, epsilon); + emscripten::val add_constant = model_builder.CreateOrGetConstant(input_type, epsilon); common_options.set("label", node.Name() + "_add"); emscripten::val add = model_builder.GetBuilder().call("add", reduce_mean, add_constant, common_options); @@ -183,13 +150,6 @@ Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder common_options.set("label", node.Name() + "_add_bias"); output = model_builder.GetBuilder().call("add", output, bias, common_options); } - - if (input_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16) { - // Cast output back to float16. - common_options.set("label", node.Name() + "_cast_output_to_fp16"); - output = model_builder.GetBuilder().call("cast", output, - emscripten::val("float16"), common_options); - } } } else if (op_type == "InstanceNormalization") { // WebNN spec only supports 4D input for instanceNormalization.