From 7cac114e522c694d669819869c965ea6f4f89751 Mon Sep 17 00:00:00 2001 From: Wanming Lin Date: Thu, 13 Jul 2023 15:44:22 +0800 Subject: [PATCH] [WebNN EP] Support Abs and Neg ops (#16672) --- onnxruntime/core/providers/webnn/builders/helper.h | 2 ++ .../providers/webnn/builders/impl/unary_op_builder.cc | 8 +++++++- .../core/providers/webnn/builders/op_builder_factory.cc | 2 ++ 3 files changed, 11 insertions(+), 1 deletion(-) diff --git a/onnxruntime/core/providers/webnn/builders/helper.h b/onnxruntime/core/providers/webnn/builders/helper.h index 0b4b6e4687..b4b7cb175a 100644 --- a/onnxruntime/core/providers/webnn/builders/helper.h +++ b/onnxruntime/core/providers/webnn/builders/helper.h @@ -92,6 +92,7 @@ std::vector> GetSupportedNodes(const GraphViewer& graph_v const WebnnDeviceType device_type, const logging::Logger& logger); static const InlinedHashMap op_map = { + {"Abs", "abs"}, {"ArgMax", "argMax"}, {"ArgMin", "argMin"}, {"Add", "add"}, @@ -104,6 +105,7 @@ static const InlinedHashMap op_map = { {"Equal", "equal"}, {"Erf", "erf"}, {"Exp", "exp"}, + {"Neg", "neg"}, {"Not", "logicalNot"}, {"Floor", "floor"}, {"Flatten", "flattenTo2d"}, diff --git a/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc index 20ef6d7ac8..95624a22e9 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc @@ -29,7 +29,9 @@ Status UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const emscripten::val input = model_builder.GetOperand(node.InputDefs()[0]->Name()); emscripten::val output = emscripten::val::object(); - if (op_type == "Ceil") { + if (op_type == "Abs") { + output = model_builder.GetBuilder().call("abs", input); + } else if (op_type == "Ceil") { output = model_builder.GetBuilder().call("ceil", input); } else if (op_type == "Cos") { output = model_builder.GetBuilder().call("cos", input); @@ -41,6 +43,8 @@ Status UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const output = model_builder.GetBuilder().call("floor", input); } else if (op_type == "Identity") { output = model_builder.GetBuilder().call("identity", input); + } else if (op_type == "Neg") { + output = model_builder.GetBuilder().call("neg", input); } else if (op_type == "Not") { output = model_builder.GetBuilder().call("logicalNot", input); } else if (op_type == "Reciprocal") { @@ -66,12 +70,14 @@ void CreateUnaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& op static std::vector op_types = { + "Abs", "Ceil", "Cos", "Erf", "Exp", "Floor", "Identity", + "Neg", "Not", "Reciprocal", "Sin", diff --git a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc index b390853da1..b334fdc638 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc +++ b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc @@ -16,12 +16,14 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { OpBuilderRegistrations op_registrations; { // Unary + CreateUnaryOpBuilder("Abs", op_registrations); CreateUnaryOpBuilder("Ceil", op_registrations); CreateUnaryOpBuilder("Cos", op_registrations); CreateUnaryOpBuilder("Erf", op_registrations); CreateUnaryOpBuilder("Exp", op_registrations); CreateUnaryOpBuilder("Floor", op_registrations); CreateUnaryOpBuilder("Identity", op_registrations); + CreateUnaryOpBuilder("Neg", op_registrations); CreateUnaryOpBuilder("Not", op_registrations); CreateUnaryOpBuilder("Reciprocal", op_registrations); CreateUnaryOpBuilder("Sin", op_registrations);