[WebNN EP] Support Abs and Neg ops (#16672)

This commit is contained in:
Wanming Lin 2023-07-13 15:44:22 +08:00 committed by GitHub
parent d5b76cff60
commit 7cac114e52
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 11 additions and 1 deletions

View file

@ -92,6 +92,7 @@ std::vector<std::vector<NodeIndex>> GetSupportedNodes(const GraphViewer& graph_v
const WebnnDeviceType device_type,
const logging::Logger& logger);
static const InlinedHashMap<std::string, std::string> op_map = {
{"Abs", "abs"},
{"ArgMax", "argMax"},
{"ArgMin", "argMin"},
{"Add", "add"},
@ -104,6 +105,7 @@ static const InlinedHashMap<std::string, std::string> op_map = {
{"Equal", "equal"},
{"Erf", "erf"},
{"Exp", "exp"},
{"Neg", "neg"},
{"Not", "logicalNot"},
{"Floor", "floor"},
{"Flatten", "flattenTo2d"},

View file

@ -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<emscripten::val>("abs", input);
} else if (op_type == "Ceil") {
output = model_builder.GetBuilder().call<emscripten::val>("ceil", input);
} else if (op_type == "Cos") {
output = model_builder.GetBuilder().call<emscripten::val>("cos", input);
@ -41,6 +43,8 @@ Status UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
output = model_builder.GetBuilder().call<emscripten::val>("floor", input);
} else if (op_type == "Identity") {
output = model_builder.GetBuilder().call<emscripten::val>("identity", input);
} else if (op_type == "Neg") {
output = model_builder.GetBuilder().call<emscripten::val>("neg", input);
} else if (op_type == "Not") {
output = model_builder.GetBuilder().call<emscripten::val>("logicalNot", input);
} else if (op_type == "Reciprocal") {
@ -66,12 +70,14 @@ void CreateUnaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& op
static std::vector<std::string> op_types =
{
"Abs",
"Ceil",
"Cos",
"Erf",
"Exp",
"Floor",
"Identity",
"Neg",
"Not",
"Reciprocal",
"Sin",

View file

@ -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);