diff --git a/cmake/onnxruntime_webassembly.cmake b/cmake/onnxruntime_webassembly.cmake index 4df1e3c615..7804c31cc2 100644 --- a/cmake/onnxruntime_webassembly.cmake +++ b/cmake/onnxruntime_webassembly.cmake @@ -262,7 +262,7 @@ else() endif() if (onnxruntime_USE_WEBNN) - set_property(TARGET onnxruntime_webassembly APPEND_STRING PROPERTY LINK_FLAGS " --bind") + set_property(TARGET onnxruntime_webassembly APPEND_STRING PROPERTY LINK_FLAGS " --bind -sWASM_BIGINT") endif() # Set link flag to enable exceptions support, this will override default disabling exception throwing behavior when disable exceptions. diff --git a/js/common/lib/tensor-impl.ts b/js/common/lib/tensor-impl.ts index 6535f79d53..42d92e4d2b 100644 --- a/js/common/lib/tensor-impl.ts +++ b/js/common/lib/tensor-impl.ts @@ -17,6 +17,7 @@ const NUMERIC_TENSOR_TYPE_TO_TYPEDARRAY_MAP = new Map { return DataType.int32; case 'uint32': return DataType.uint32; + case 'float16': + return DataType.float16; case 'float32': return DataType.float; case 'float64': @@ -80,6 +82,8 @@ export const tensorDataTypeEnumToString = (typeProto: DataType): Tensor.Type => return 'int32'; case DataType.uint32: return 'uint32'; + case DataType.float16: + return 'float16'; case DataType.float: return 'float32'; case DataType.double: @@ -110,6 +114,8 @@ export const tensorTypeToTypedArrayConstructor = (type: Tensor.Type): Float32Arr Int8ArrayConstructor|Uint16ArrayConstructor|Int16ArrayConstructor|Int32ArrayConstructor|BigInt64ArrayConstructor| Uint8ArrayConstructor|Float64ArrayConstructor|Uint32ArrayConstructor|BigUint64ArrayConstructor => { switch (type) { + case 'float16': + return Uint16Array; case 'float32': return Float32Array; case 'uint8': diff --git a/onnxruntime/core/providers/webnn/builders/helper.cc b/onnxruntime/core/providers/webnn/builders/helper.cc index c52a6b18c6..f758d27ea5 100644 --- a/onnxruntime/core/providers/webnn/builders/helper.cc +++ b/onnxruntime/core/providers/webnn/builders/helper.cc @@ -27,11 +27,12 @@ bool GetShape(const NodeArg& node_arg, std::vector& shape, const loggin return true; } -bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const logging::Logger& logger) { +bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, + const WebnnDeviceType device_type, const logging::Logger& logger) { const auto& op_builders = GetOpBuilders(); if (Contains(op_builders, node.OpType())) { const auto* op_builder = op_builders.at(node.OpType()); - return op_builder->IsOpSupported(graph_viewer.GetAllInitializedTensors(), node, logger); + return op_builder->IsOpSupported(graph_viewer.GetAllInitializedTensors(), node, device_type, logger); } else { return false; } @@ -40,6 +41,10 @@ bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const lo bool IsInputSupported(const NodeArg& input, const std::string& parent_name, const logging::Logger& logger) { const auto& input_name = input.Name(); const auto* shape_proto = input.Shape(); + // Optional tensors can be indicated by an empty name, just ignore it. + if (input_name.empty()) { + return true; + } // We do not support input with no shape. if (!shape_proto) { LOGS(logger, VERBOSE) << "Input [" << input_name << "] of [" << parent_name @@ -59,6 +64,7 @@ bool IsInputSupported(const NodeArg& input, const std::string& parent_name, cons std::vector> GetSupportedNodes(const GraphViewer& graph_viewer, const emscripten::val& wnn_builder_, + const WebnnDeviceType device_type, const logging::Logger& logger) { std::vector> supported_node_groups; @@ -78,7 +84,7 @@ std::vector> GetSupportedNodes(const GraphViewer& graph_v // Firstly check if platform supports the WebNN op. if (CheckSingleOp(node->OpType(), wnn_builder_)) { LOGS(logger, VERBOSE) << "Operator type: [" << node->OpType() << "] is supported by browser"; - supported = IsNodeSupported(*node, graph_viewer, logger); + supported = IsNodeSupported(*node, graph_viewer, device_type, logger); } LOGS(logger, VERBOSE) << "Operator type: [" << node->OpType() @@ -103,5 +109,35 @@ std::vector> GetSupportedNodes(const GraphViewer& graph_v return supported_node_groups; } +bool IsSupportedDataType(const int32_t data_type, const WebnnDeviceType device_type) { + // Current data type implementation status of WebNN is inconsistent along with different backends, + // The XNNPack backend supports only FP32, while the DML backend POC supports more. + if (device_type == WebnnDeviceType::CPU) { + return std::find(supported_cpu_data_types.begin(), supported_cpu_data_types.end(), data_type) != + supported_cpu_data_types.end(); + } else { + return std::find(supported_gpu_data_types.begin(), supported_gpu_data_types.end(), data_type) != + supported_gpu_data_types.end(); + } +} + +bool IsValidMultidirectionalBroadcast(std::vector& shape_a, + std::vector& shape_b, + const logging::Logger& logger) { + int64_t size_a = shape_a.size(); + int64_t size_b = shape_b.size(); + int64_t smaller_size = std::min(size_a, size_b); + for (int64_t i = 0; i < smaller_size; i++) { + // right alignment + int64_t axis_a = size_a - i - 1; + int64_t axis_b = size_b - i - 1; + // Broadcastable tensors must either have each dimension the same size or equal to one. + if (shape_a[axis_a] != shape_b[axis_b] && shape_a[axis_a] != 1 && shape_b[axis_b] != 1) { + return false; + } + } + return true; +} + } // namespace webnn } // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/helper.h b/onnxruntime/core/providers/webnn/builders/helper.h index 9a386e6e34..a8603fbadf 100644 --- a/onnxruntime/core/providers/webnn/builders/helper.h +++ b/onnxruntime/core/providers/webnn/builders/helper.h @@ -7,7 +7,9 @@ #include #include "core/common/inlined_containers.h" #include +#include "core/optimizer/initializer.h" #include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" #include #include @@ -23,38 +25,138 @@ class Logger; namespace webnn { +enum class WebnnDeviceType { + CPU, + GPU, +}; + bool GetShape(const NodeArg& node_arg, std::vector& shape, const logging::Logger& logger); +template +std::string GetShapeString(std::vector& shape) { + std::stringstream shape_info; + shape_info << "["; + for (size_t i = 0; i < shape.size(); i++) { + if (i != 0) { + shape_info << ", "; + } + shape_info << shape[i]; + } + shape_info << "]"; + return shape_info.str(); +} + +template +bool ReadIntArrayFrom1DTensor(const onnx::TensorProto& tensor, std::vector& array, const logging::Logger& logger) { + std::vector unpacked_tensor; + auto status = onnxruntime::utils::UnpackInitializerData(tensor, unpacked_tensor); + if (!status.IsOK()) { + LOGS(logger, ERROR) << "Error while unpacking shape: " << status.ErrorMessage(); + return false; + } + const auto& dims = tensor.dims(); + if (dims.size() != 1) { + LOGS(logger, VERBOSE) << "The tensor must be 1D."; + return false; + } + int64_t rank = dims[0]; + switch (tensor.data_type()) { + case ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64: { + const int64_t* array_data = reinterpret_cast(unpacked_tensor.data()); + if constexpr (std::is_same::value) { + array.assign(array_data, array_data + rank); + } else { + std::transform(array_data, array_data + rank, + std::back_inserter(array), + [](int64_t dim) -> T { return SafeInt(dim); }); + }; + break; + } + + case ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32: { + const int32_t* array_data = reinterpret_cast(unpacked_tensor.data()); + array.assign(array_data, array_data + rank); + break; + } + default: + return false; + } + return true; +} + bool IsInputSupported(const NodeArg& node_arg, const std::string& parent_name, const logging::Logger& logger); // Get a list of groups of supported nodes, each group represents a subgraph supported by WebNN EP. std::vector> GetSupportedNodes(const GraphViewer& graph_viewer, const emscripten::val& wnn_builder_, + const WebnnDeviceType device_type, const logging::Logger& logger); static const InlinedHashMap op_map = { + {"ArgMax", "argMax"}, + {"ArgMin", "argMin"}, {"Add", "add"}, {"Sub", "sub"}, {"Mul", "mul"}, {"Div", "div"}, + {"Pow", "pow"}, + {"Cos", "cos"}, + {"Equal", "equal"}, + {"Erf", "erf"}, + {"Not", "logicalNot"}, + {"Floor", "floor"}, + {"Flatten", "flattenTo2d"}, + {"Sin", "sin"}, + {"Sqrt", "sqrt"}, {"Relu", "relu"}, {"LeakyRelu", "leakyRelu"}, {"Sigmoid", "sigmoid"}, + {"Slice", "slice"}, + {"Softmax", "softmax"}, + {"Cast", "cast"}, {"Clip", "clamp"}, {"Conv", "conv2d"}, {"ConvTranspose", "convTranspose2d"}, {"Concat", "concat"}, + {"Expand", "expand"}, + {"Gather", "gather"}, {"Gemm", "gemm"}, + {"MatMul", "matmul"}, {"GlobalAveragePool", "averagePool2d"}, {"GlobalMaxPool", "maxPool2d"}, {"AveragePool", "averagePool2d"}, + {"LayerNormalization", "meanVarianceNormalization"}, {"MaxPool", "maxPool2d"}, + {"ReduceMax", "reduceMax"}, + {"ReduceMean", "reduceMean"}, {"Reshape", "reshape"}, {"Resize", "resample2d"}, - {"Transpose", "transpose"}}; + {"Split", "split"}, + {"Transpose", "transpose"}, + {"Unsqueeze", "unsqueeze"}, +}; inline bool CheckSingleOp(const std::string& op_type, const emscripten::val& wnn_builder_) { return op_map.find(op_type) != op_map.end() && wnn_builder_[op_map.find(op_type)->second].as(); } +constexpr std::array supported_cpu_data_types = { + ONNX_NAMESPACE::TensorProto_DataType_FLOAT, +}; + +constexpr std::array supported_gpu_data_types = { + ONNX_NAMESPACE::TensorProto_DataType_BOOL, + ONNX_NAMESPACE::TensorProto_DataType_FLOAT16, + ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + ONNX_NAMESPACE::TensorProto_DataType_INT32, + ONNX_NAMESPACE::TensorProto_DataType_INT64, + ONNX_NAMESPACE::TensorProto_DataType_UINT32, + ONNX_NAMESPACE::TensorProto_DataType_UINT64, +}; + +bool IsSupportedDataType(const int32_t data_type, const WebnnDeviceType device_type); + +bool IsValidMultidirectionalBroadcast(std::vector& shape_a, + std::vector& shape_b, + const logging::Logger& logger); } // namespace webnn } // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/argmax_min_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/argmax_min_op_builder.cc new file mode 100644 index 0000000000..47304ff02b --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/argmax_min_op_builder.cc @@ -0,0 +1,95 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class ArgMaxMinOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. + +Status ArgMaxMinOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get shape"); + const auto input_rank = input_shape.size(); + + NodeAttrHelper helper(node); + int64_t axis = helper.Get("axis", 0); + const auto keep_dims = helper.Get("keepdims", 1); + const auto select_last_index = helper.Get("select_last_index", 0); + + axis = HandleNegativeAxis(axis, input_rank); + + emscripten::val options = emscripten::val::object(); + options.set("axis", static_cast(axis)); + options.set("keepDimensions", keep_dims == 1); + options.set("selectLastIndex", select_last_index == 1); + emscripten::val output = emscripten::val::object(); + + const auto& op_type = node.OpType(); + if (op_type == "ArgMax") { + output = model_builder.GetBuilder().call("argMax", input, options); + } else if (op_type == "ArgMin") { + output = model_builder.GetBuilder().call("argMin", input, options); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "ArgMaxMinOpBuilder, unknown op: ", op_type); + } + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. +bool ArgMaxMinOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) + return false; + + return true; +} + +void CreateArgMaxMinOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + if (op_registrations.op_builder_map.find(op_type) != op_registrations.op_builder_map.cend()) + return; + + static std::vector op_types = + { + "ArgMax", + "ArgMin", + }; + + op_registrations.builders.push_back(std::make_unique()); + for (const auto& type : op_types) { + op_registrations.op_builder_map.emplace(type, op_registrations.builders.back().get()); + } +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.cc index 0f0ab712d5..a44562edc7 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.cc @@ -38,7 +38,7 @@ bool HasExternalInitializer(const InitializedTensorSet& initializers, const Node Status BaseOpBuilder::AddToModelBuilder(ModelBuilder& model_builder, const Node& node, const logging::Logger& logger) const { ORT_RETURN_IF_NOT( - IsOpSupported(model_builder.GetInitializerTensors(), node, logger), + IsOpSupported(model_builder.GetInitializerTensors(), node, model_builder.GetWebnnDeviceType(), logger), "Unsupported operator ", node.OpType()); ORT_RETURN_IF_ERROR(AddToModelBuilderImpl(model_builder, node, logger)); @@ -50,8 +50,8 @@ Status BaseOpBuilder::AddToModelBuilder(ModelBuilder& model_builder, const Node& // Operator support related. bool BaseOpBuilder::IsOpSupported(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const { - if (!HasSupportedInputs(node, logger)) + const WebnnDeviceType device_type, const logging::Logger& logger) const { + if (!HasSupportedInputs(node, device_type, logger)) return false; // We do not support external initializers for now. @@ -61,10 +61,11 @@ bool BaseOpBuilder::IsOpSupported(const InitializedTensorSet& initializers, cons if (!HasSupportedOpSet(node, logger)) return false; - return IsOpSupportedImpl(initializers, node, logger); + return IsOpSupportedImpl(initializers, node, device_type, logger); } -bool BaseOpBuilder::HasSupportedInputs(const Node& node, const logging::Logger& logger) const { +bool BaseOpBuilder::HasSupportedInputs(const Node& node, const WebnnDeviceType device_type, + const logging::Logger& logger) const { const auto node_name = MakeString("Node [", node.Name(), "] type [", node.OpType(), "]"); for (const auto* input : node.InputDefs()) { if (!IsInputSupported(*input, node_name, logger)) { @@ -72,10 +73,12 @@ bool BaseOpBuilder::HasSupportedInputs(const Node& node, const logging::Logger& } } - return HasSupportedInputsImpl(node, logger); + return HasSupportedInputsImpl(node, device_type, logger); } -bool BaseOpBuilder::HasSupportedInputsImpl(const Node& node, const logging::Logger& logger) const { +bool BaseOpBuilder::HasSupportedInputsImpl(const Node& node, + const WebnnDeviceType device_type, + const logging::Logger& logger) const { // We only check the type of input 0 by default, specific op builder can override this. const auto& input = *node.InputDefs()[0]; @@ -83,7 +86,7 @@ bool BaseOpBuilder::HasSupportedInputsImpl(const Node& node, const logging::Logg if (!GetType(input, input_type, logger)) return false; - if (input_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + if (!IsSupportedDataType(input_type, device_type)) { LOGS(logger, VERBOSE) << "[" << node.OpType() << "] Input type: [" << input_type << "] is not supported for now"; diff --git a/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.h b/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.h index 0759d93a9e..203e5b773c 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.h +++ b/onnxruntime/core/providers/webnn/builders/impl/base_op_builder.h @@ -28,22 +28,23 @@ class BaseOpBuilder : public IOpBuilder { // Operator support related. public: bool IsOpSupported(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const override; + const WebnnDeviceType device_type, const logging::Logger& logger) const override; protected: virtual bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& /* node */, - const logging::Logger& /* logger */) const { + const WebnnDeviceType /* device_type */, const logging::Logger& /* logger */) const { return true; } - virtual bool HasSupportedInputsImpl(const Node& node, const logging::Logger& logger) const; + virtual bool HasSupportedInputsImpl(const Node& node, const WebnnDeviceType device_type, + const logging::Logger& logger) const; virtual int GetMinSupportedOpSet(const Node& /* node */) const { return 1; } virtual int GetMaxSupportedOpSet(const Node& /* node */) const { return 19; } private: bool HasSupportedOpSet(const Node& node, const logging::Logger& logger) const; - bool HasSupportedInputs(const Node& node, const logging::Logger& logger) const; + bool HasSupportedInputs(const Node& node, const WebnnDeviceType device_type, const logging::Logger& logger) const; }; } // namespace webnn diff --git a/onnxruntime/core/providers/webnn/builders/impl/binary_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/binary_op_builder.cc index 96ee9cf3f1..8b45f0d16e 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/binary_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/binary_op_builder.cc @@ -41,6 +41,8 @@ Status BinaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const output = model_builder.GetBuilder().call("mul", input0, input1); } else if (op_type == "Div") { output = model_builder.GetBuilder().call("div", input0, input1); + } else if (op_type == "Pow") { + output = model_builder.GetBuilder().call("pow", input0, input1); } else { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "BinaryOpBuilder::AddToModelBuilderImpl, unknown op: ", op_type); @@ -53,7 +55,7 @@ Status BinaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const // Operator support related. int BinaryOpBuilder::GetMinSupportedOpSet(const Node& /* node */) const { - // Add/Sub/Mul/Div opset 6- has broadcast attributes we do not support now. + // Add/Sub/Mul/Div/Pow opset 6- has broadcast attributes we do not support now. return 7; } @@ -67,6 +69,7 @@ void CreateBinaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& o "Sub", "Mul", "Div", + "Pow", }; op_registrations.builders.push_back(std::make_unique()); diff --git a/onnxruntime/core/providers/webnn/builders/impl/cast_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/cast_op_builder.cc new file mode 100644 index 0000000000..7c401dafd1 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/cast_op_builder.cc @@ -0,0 +1,105 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" +#include "core/providers/shared/utils/utils.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class CastOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType device_type, const logging::Logger& logger) const override; + + int GetMinSupportedOpSet(const Node& node) const override; +}; + +// Add operator related. + +Status CastOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_name = node.InputDefs()[0]->Name(); + emscripten::val input = model_builder.GetOperand(input_name); + + NodeAttrHelper helper(node); + // We already checked the "to" type in IsOpSupportedImpl. + const auto to_type = helper.Get("to", ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + std::string operand_type; + switch (to_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + operand_type = "uint8"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + operand_type = "float16"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + operand_type = "float32"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + operand_type = "int32"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + operand_type = "int64"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + operand_type = "uint32"; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + operand_type = "uint64"; + break; + default: + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "The Cast node has unsupported 'to' type, name: ", + node.Name(), " type: ", to_type); + } + + emscripten::val output = + model_builder.GetBuilder().call("cast", input, emscripten::val(operand_type)); + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. + +int CastOpBuilder::GetMinSupportedOpSet(const Node& /* node */) const { + // Since opset 6, Cast uses attribute "to" as int type. + return 6; +} + +bool CastOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType device_type, + const logging::Logger& logger) const { + NodeAttrHelper helper(node); + // Check cast output type. + const auto to_type = helper.Get("to", ONNX_NAMESPACE::TensorProto_DataType_UNDEFINED); + if (!IsSupportedDataType(to_type, device_type)) { + LOGS(logger, VERBOSE) << "Invalid cast to type " << to_type << "."; + return false; + } + + return true; +} + +void CreateCastOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/clip_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/clip_op_builder.cc index 1c07439c8e..9de5b88980 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/clip_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/clip_op_builder.cc @@ -24,7 +24,7 @@ class ClipOpBuilder : public BaseOpBuilder { // Operator support related. private: bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const override; + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; }; // Add operator related. @@ -66,7 +66,9 @@ Status ClipOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, // Operator support related. -bool ClipOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, +bool ClipOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { float min, max; return GetClipMinMax(initializers, node, min, max, logger); diff --git a/onnxruntime/core/providers/webnn/builders/impl/conv_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/conv_op_builder.cc index aec845ae3b..cb7d27f86f 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/conv_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/conv_op_builder.cc @@ -27,8 +27,8 @@ class ConvOpBuilder : public BaseOpBuilder { // Operator support related. private: - bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& /* node */, - const logging::Logger& /* logger */) const override; + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; }; void ConvOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { @@ -96,7 +96,7 @@ Status AddInitializerInNewLayout(ModelBuilder& model_builder, bool is_conv) { const auto& tensor = *model_builder.GetInitializerTensors().at(name); auto data_type = tensor.data_type(); - if (data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + if (!IsSupportedDataType(data_type, model_builder.GetWebnnDeviceType())) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "The initializer of graph has unsupported type, name: ", tensor.name(), " type: ", data_type); @@ -123,7 +123,29 @@ Status AddInitializerInNewLayout(ModelBuilder& model_builder, SafeInt num_elements = SafeInt(Product(dest_shape)); - size_t element_size = 4; + size_t element_size{0}; + switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + element_size = sizeof(uint16_t); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + element_size = sizeof(float); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + element_size = sizeof(int32_t); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + element_size = sizeof(int64_t); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + element_size = sizeof(uint32_t); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + element_size = sizeof(uint64_t); + break; + default: + break; + } std::unique_ptr buffer_holder(new uint8_t[element_size * num_elements]); uint8_t* buffer = buffer_holder.get(); @@ -157,7 +179,7 @@ Status AddInitializerInNewLayout(ModelBuilder& model_builder, } } ORT_RETURN_IF_ERROR(model_builder.AddOperandFromPersistMemoryBuffer(name, buffer, num_elements * element_size, - dest_shape, 4)); + dest_shape, data_type)); return Status::OK(); } @@ -175,7 +197,6 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N const auto dilations = helper.Get("dilations", std::vector{1, 1}); auto pads = helper.Get("pads", std::vector{0, 0, 0, 0}); const auto& weight = input_defs[1]->Name(); - if (op_type == "Conv") { emscripten::val options = emscripten::val::object(); ORT_RETURN_IF_ERROR(SetConvBaseOptions(model_builder, node, options, strides, dilations, pads, logger)); @@ -193,6 +214,7 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N } } emscripten::val filter = model_builder.GetOperand(input_defs[1]->Name()); + output = model_builder.GetBuilder().call("conv2d", input, filter, options); } else { emscripten::val options = emscripten::val::object(); @@ -242,7 +264,6 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N options.set("outputPadding", emscripten::val::array(output_padding)); } emscripten::val filter = model_builder.GetOperand(input_defs[1]->Name()); - output = model_builder.GetBuilder().call("convTranspose2d", input, filter, options); } @@ -252,7 +273,9 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N // Operator support related. -bool ConvOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, +bool ConvOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { const auto& name = node.Name(); const auto& op_type = node.OpType(); diff --git a/onnxruntime/core/providers/webnn/builders/impl/expand_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/expand_op_builder.cc new file mode 100644 index 0000000000..9d4de45fda --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/expand_op_builder.cc @@ -0,0 +1,116 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/safeint.h" +#include "core/framework/tensorprotoutils.h" +#include "core/optimizer/initializer.h" +#include "core/providers/common.h" +#include "core/providers/cpu/tensor/reshape_helper.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class ExpandOpBuilder : public BaseOpBuilder { + // Add operator related. + public: + void AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const override; + + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. + +void ExpandOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { + model_builder.AddInitializerToSkip(node.InputDefs()[1]->Name()); +} + +Status ExpandOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + const auto& initializers(model_builder.GetInitializerTensors()); + const auto& shape_tensor = *initializers.at(input_defs[1]->Name()); + std::vector new_shape; + ORT_RETURN_IF_NOT(ReadIntArrayFrom1DTensor(shape_tensor, new_shape, logger), "Cannot get shape."); + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get input's shape."); + if (new_shape.size() < input_shape.size()) { + // Enlarge new shape to input.rank, right aligned with leading ones + new_shape.insert(new_shape.begin(), input_shape.size() - new_shape.size(), 1); + } + emscripten::val output = + model_builder.GetBuilder().call("expand", + input, emscripten::val::array(new_shape)); + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. + +bool ExpandOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + const auto& shape_name = input_defs[1]->Name(); + if (!Contains(initializers, shape_name)) { + LOGS(logger, VERBOSE) << "The shape must be a constant initializer."; + return false; + } + + std::vector new_shape; + const auto& shape_tensor = *initializers.at(shape_name); + if (shape_tensor.data_type() != ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64) { + LOGS(logger, VERBOSE) << "The type of tensor's element data must be INT64."; + return false; + } + if (!ReadIntArrayFrom1DTensor(shape_tensor, new_shape, logger)) { + LOGS(logger, VERBOSE) << "Cannot get shape."; + return false; + } + + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) { + LOGS(logger, VERBOSE) << "Cannot get input's shape."; + return false; + } + + if (input_shape.empty()) { + LOGS(logger, VERBOSE) << "Expand does not support empty input's shape."; + return false; + } + + if (new_shape.size() > input_shape.size()) { + LOGS(logger, VERBOSE) << "The size of shape must be less than or equal to the rank of input."; + } + + if (!IsValidMultidirectionalBroadcast(input_shape, new_shape, logger)) { + LOGS(logger, VERBOSE) << "The input cannot expand to shape " << GetShapeString(new_shape); + return false; + } + + return true; +} + +void CreateExpandOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc new file mode 100644 index 0000000000..6c59ca451f --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc @@ -0,0 +1,58 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/safeint.h" +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class FlattenOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; +}; + +// Add operator related. + +Status FlattenOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + ORT_RETURN_IF(input_defs.size() < 1, "Flatten has no input tensor"); + if (!GetShape(*input_defs[0], input_shape, logger)) { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "FlattenOpBuilder::AddToModelBuilderImpl, cannot get input shape"); + } + int64_t rank = input_shape.size(); + NodeAttrHelper helper(node); + int64_t axis = helper.Get("axis", 1); + ORT_ENFORCE(axis >= -rank && axis <= rank, "axis ", axis, + " is not in valid range [-", rank, ",", rank, "]"); + if (axis < 0) { + axis += rank; + } + emscripten::val inputs = model_builder.GetOperand(input_defs[0]->Name()); + emscripten::val output = model_builder.GetBuilder().call("flattenTo2d", inputs, + static_cast(axis)); + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +void CreateFlattenOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/gather_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/gather_op_builder.cc new file mode 100644 index 0000000000..74a8f74474 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/gather_op_builder.cc @@ -0,0 +1,75 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class GatherOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. + +Status GatherOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get shape"); + const auto rank = input_shape.size(); + NodeAttrHelper helper(node); + const uint32_t axis = static_cast(HandleNegativeAxis(helper.Get("axis", 1), rank)); + + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + emscripten::val indices = model_builder.GetOperand(input_defs[1]->Name()); + emscripten::val options = emscripten::val::object(); + options.set("axis", axis); + emscripten::val output = model_builder.GetBuilder().call("gather", input, indices, options); + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. + +bool GatherOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) + return false; + const auto rank = input_shape.size(); + if (rank < 1) { + LOGS(logger, VERBOSE) << "Gather only supports input shapes >= 1D, but input is " + << rank << "d shape"; + return false; + } + + return true; +} + +void CreateGatherOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/gemm_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/gemm_op_builder.cc index 2330d85f91..b6f4f0a7b8 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/gemm_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/gemm_op_builder.cc @@ -23,106 +23,117 @@ class GemmOpBuilder : public BaseOpBuilder { // Operator support related. private: - bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& /* node */, - const logging::Logger& /* logger */) const override; + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; }; // Add operator related. Status GemmOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, const logging::Logger& /* logger */) const { + const auto& op_type = node.OpType(); const auto& input_defs = node.InputDefs(); const size_t a_idx = 0, b_idx = 1, c_idx = 2; // A*B+C emscripten::val a = model_builder.GetOperand(node.InputDefs()[a_idx]->Name()); emscripten::val b = model_builder.GetOperand(node.InputDefs()[b_idx]->Name()); - emscripten::val options = emscripten::val::object(); - NodeAttrHelper helper(node); - const auto transA = helper.Get("transA", 0); - options.set("aTranspose", emscripten::val(transA == 1)); - const auto transB = helper.Get("transB", 0); - options.set("bTranspose", emscripten::val(transB == 1)); - const auto alpha = helper.Get("alpha", 1.0f); - const auto beta = helper.Get("beta", 1.0f); - options.set("alpha", alpha); - options.set("beta", beta); + emscripten::val output = emscripten::val::object(); + if (op_type == "MatMul") { + output = model_builder.GetBuilder().call("matmul", a, b); + } else { // Gemm + emscripten::val options = emscripten::val::object(); + NodeAttrHelper helper(node); + const auto transA = helper.Get("transA", 0); + options.set("aTranspose", emscripten::val(transA == 1)); + const auto transB = helper.Get("transB", 0); + options.set("bTranspose", emscripten::val(transB == 1)); + const auto alpha = helper.Get("alpha", 1.0f); + const auto beta = helper.Get("beta", 1.0f); + options.set("alpha", alpha); + options.set("beta", beta); - // Add bias if present. - if (input_defs.size() > 2) { - options.set("c", model_builder.GetOperand(node.InputDefs()[c_idx]->Name())); + // Add bias if present. + if (input_defs.size() > 2) { + options.set("c", model_builder.GetOperand(node.InputDefs()[c_idx]->Name())); + } + + output = model_builder.GetBuilder().call("gemm", a, b, options); } - emscripten::val output = model_builder.GetBuilder().call("gemm", a, b, options); - model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); return Status::OK(); } // Operator support related. -bool GemmOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, +bool GemmOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { (void)initializers; + const auto& op_type = node.OpType(); const auto& input_defs(node.InputDefs()); const size_t a_idx = 0, b_idx = 1, c_idx = 2; // A*B+C - std::vector a_shape; - { - if (!GetShape(*input_defs[a_idx], a_shape, logger)) - return false; - - if (a_shape.size() != 2) { - LOGS(logger, VERBOSE) << "A must be 2D"; - return false; - } - - if (Product(a_shape) == 0) { - LOGS(logger, VERBOSE) << "A must be non-empty"; - return false; - } - } - - std::vector b_shape; - { - if (!GetShape(*input_defs[b_idx], b_shape, logger)) - return false; - - if (b_shape.size() != 2) { - LOGS(logger, VERBOSE) << "B must be 2D"; - return false; - } - - if (Product(b_shape) == 0) { - LOGS(logger, VERBOSE) << "B must be non-empty"; - return false; - } - } - - // C of Gemm. - if (input_defs.size() == 3) { - std::vector c_shape; - if (!GetShape(*input_defs[c_idx], c_shape, logger)) - return false; - - size_t c_dim = c_shape.size(); - - if (c_dim > 1) { - // TODO: Supports other shape of C. - // Currently WebNN implementation in Chromium only supports 1-D C. - return false; - } - if (c_dim == 0) { - LOGS(logger, VERBOSE) << "C of Gemm is a scalar"; - } else { - auto c_size = c_shape[c_dim - 1]; - NodeAttrHelper helper(node); - const auto transB = helper.Get("transB", 0); - if (c_size != (transB == 0 ? b_shape[1] : b_shape[0])) { - LOGS(logger, VERBOSE) << "C of Gemm must be a vector of b_shape[" - << (transB == 0 ? "1" : "0") << "]" - << " b_shape: [" << b_shape[0] << ", " << b_shape[1] << "]" - << " c_size: " << c_size; - + if (op_type == "Gemm") { + std::vector a_shape; + { + if (!GetShape(*input_defs[a_idx], a_shape, logger)) return false; + + if (a_shape.size() != 2) { + LOGS(logger, VERBOSE) << "A must be 2D"; + return false; + } + + if (Product(a_shape) == 0) { + LOGS(logger, VERBOSE) << "A must be non-empty"; + return false; + } + } + + std::vector b_shape; + { + if (!GetShape(*input_defs[b_idx], b_shape, logger)) + return false; + + if (b_shape.size() != 2) { + LOGS(logger, VERBOSE) << "B must be 2D"; + return false; + } + + if (Product(b_shape) == 0) { + LOGS(logger, VERBOSE) << "B must be non-empty"; + return false; + } + } + + // C of Gemm. + if (input_defs.size() == 3) { + std::vector c_shape; + if (!GetShape(*input_defs[c_idx], c_shape, logger)) + return false; + + size_t c_dim = c_shape.size(); + + if (c_dim > 1) { + // TODO: Supports other shape of C. + // Currently WebNN implementation in Chromium only supports 1-D C. + return false; + } + if (c_dim == 0) { + LOGS(logger, VERBOSE) << "C of Gemm is a scalar"; + } else { + auto c_size = c_shape[c_dim - 1]; + NodeAttrHelper helper(node); + const auto transB = helper.Get("transB", 0); + if (c_size != (transB == 0 ? b_shape[1] : b_shape[0])) { + LOGS(logger, VERBOSE) << "C of Gemm must be a vector of b_shape[" + << (transB == 0 ? "1" : "0") << "]" + << " b_shape: [" << b_shape[0] << ", " << b_shape[1] << "]" + << " c_size: " << c_size; + + return false; + } } } } @@ -131,8 +142,19 @@ bool GemmOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, } void CreateGemmOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + if (op_registrations.op_builder_map.find(op_type) != op_registrations.op_builder_map.cend()) + return; + + static std::vector op_types = + { + "Gemm", + "MatMul", + }; + op_registrations.builders.push_back(std::make_unique()); - op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); + for (const auto& type : op_types) { + op_registrations.op_builder_map.emplace(type, op_registrations.builders.back().get()); + } } } // namespace webnn } // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/logical_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/logical_op_builder.cc new file mode 100644 index 0000000000..ef3b7a60d1 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/logical_op_builder.cc @@ -0,0 +1,76 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include + +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" +#include "core/providers/webnn/builders/helper.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class LogicalOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + // Operator support related. + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. + +Status LogicalOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& /* logger */) const { + const auto& op_type = node.OpType(); + emscripten::val input0 = model_builder.GetOperand(node.InputDefs()[0]->Name()); + emscripten::val input1 = model_builder.GetOperand(node.InputDefs()[1]->Name()); + emscripten::val output = emscripten::val::object(); + if (op_type == "Equal") { + output = model_builder.GetBuilder().call("equal", input0, input1); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "LogicalOpBuilder::AddToModelBuilderImpl, unknown op: ", op_type); + } + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +void CreateLogicalOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + if (op_registrations.op_builder_map.find(op_type) != op_registrations.op_builder_map.cend()) + return; + + static std::vector op_types = + { + "Equal", + }; + + op_registrations.builders.push_back(std::make_unique()); + for (const auto& type : op_types) { + op_registrations.op_builder_map.emplace(type, op_registrations.builders.back().get()); + } +} + +bool LogicalOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& name = node.Name(); + const auto& op_type = node.OpType(); + const auto& input_defs = node.InputDefs(); + if (input_defs.size() < 2) { + LOGS(logger, VERBOSE) << op_type << " [" << name << "] requires at least 2 inputs, actual: " + << input_defs.size(); + return false; + } + return true; +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc new file mode 100644 index 0000000000..bea4fe72d6 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/normalization_op_builder.cc @@ -0,0 +1,153 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/safeint.h" +#include "core/optimizer/initializer.h" +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class NormalizationOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +Status NormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + ORT_RETURN_IF_NOT(input_defs.size() >= 2, "LayerNormalization requires at least two inputs."); + + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get input shape"); + const auto rank = input_shape.size(); + + emscripten::val options = emscripten::val::object(); + + std::vector scale_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[1], scale_shape, logger), "Cannot get scale shape"); + const auto scale_size = scale_shape.size(); + ORT_RETURN_IF_NOT(scale_size >= 1 && scale_size <= rank, "The scale size should be less than or equal to input size."); + + if (scale_size < rank) { + // Enlarge new shape to input.rank, right aligned with leading ones + scale_shape.insert(scale_shape.begin(), rank - scale_size, 1); + std::vector new_scale_shape; + std::transform(scale_shape.cbegin(), scale_shape.cend(), + std::back_inserter(new_scale_shape), + [](int64_t dim) -> int32_t { return SafeInt(dim); }); + emscripten::val reshape_scale = model_builder.GetOperand(input_defs[1]->Name()); + emscripten::val reshape_output_scale = + model_builder.GetBuilder().call("reshape", reshape_scale, emscripten::val::array(new_scale_shape)); + options.set("scale", reshape_output_scale); + } else { + options.set("scale", model_builder.GetOperand(input_defs[1]->Name())); + } + + if (input_defs.size() == 3) { + // Inputs contain optional bias + std::vector bias_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[2], bias_shape, logger), "Cannot get bias shape"); + const auto bias_size = bias_shape.size(); + ORT_RETURN_IF_NOT(bias_size >= 1 && bias_size <= rank, "The bias size should be less than or equal to input size."); + + if (bias_size < rank) { + // Enlarge new shape to input.rank, right aligned with leading ones + bias_shape.insert(bias_shape.begin(), rank - bias_size, 1); + std::vector new_bias_shape; + std::transform(bias_shape.cbegin(), bias_shape.cend(), + std::back_inserter(new_bias_shape), + [](int64_t dim) -> int32_t { return SafeInt(dim); }); + emscripten::val reshape_bias = model_builder.GetOperand(input_defs[2]->Name()); + emscripten::val reshape_output_bias = + model_builder.GetBuilder().call("reshape", reshape_bias, emscripten::val::array(new_bias_shape)); + options.set("bias", reshape_output_bias); + } else { + options.set("bias", model_builder.GetOperand(input_defs[2]->Name())); + } + } + + NodeAttrHelper helper(node); + options.set("epsilon", helper.Get("epsilon", 1e-05f)); + + int64_t axis = helper.Get("axis", -1); + axis = HandleNegativeAxis(axis, rank); + std::vector axes{static_cast(axis)}; + options.set("axes", emscripten::val::array(axes)); + + emscripten::val output = model_builder.GetBuilder().call("meanVarianceNormalization", input, options); + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + + return Status::OK(); +} + +// Operator support related. + +bool NormalizationOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + if (input_defs.size() < 2) { + LOGS(logger, VERBOSE) << "LayerNormalization requires at least two inputs."; + return false; + } + + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) { + LOGS(logger, VERBOSE) << "Cannot get input shape."; + return false; + } + const auto rank = input_shape.size(); + + NodeAttrHelper helper(node); + int64_t axis = helper.Get("axis", -1); + axis = HandleNegativeAxis(axis, rank); + + const auto& scale_name = input_defs[1]->Name(); + if (!Contains(initializers, scale_name)) { + LOGS(logger, VERBOSE) << "The scale must be a constant initializer."; + return false; + } + + if (input_defs.size() == 3) { + // Inputs contain optional bias + const auto& bias_name = input_defs[2]->Name(); + if (!Contains(initializers, bias_name)) { + LOGS(logger, VERBOSE) << "The bias must be a constant initializer."; + return false; + } + } + + const auto& output_defs = node.OutputDefs(); + if (output_defs.size() != 1) { + LOGS(logger, VERBOSE) << "MeanVarianceNormalization output count must be one."; + return false; + } + + return true; +} + +void CreateNormalizationOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/pool_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/pool_op_builder.cc index 66ca0212a4..eda05f70bf 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/pool_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/pool_op_builder.cc @@ -23,8 +23,8 @@ class PoolOpBuilder : public BaseOpBuilder { // Operator support related. private: - bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const override; + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; }; // Add operator related. @@ -108,7 +108,9 @@ Status PoolOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, } // Operator support related. -bool PoolOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, +bool PoolOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { const auto& op_type = node.OpType(); const auto& input_defs = node.InputDefs(); diff --git a/onnxruntime/core/providers/webnn/builders/impl/reduction_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/reduction_op_builder.cc new file mode 100644 index 0000000000..54352cf8a7 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/reduction_op_builder.cc @@ -0,0 +1,139 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "core/optimizer/initializer.h" + +#include "base_op_builder.h" +#include "builder_utils.h" + +namespace onnxruntime { +namespace webnn { + +class ReductionOpBuilder : public BaseOpBuilder { + // Add operator related. + public: + void AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const override; + + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. +void ReductionOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { + const auto& input_defs = node.InputDefs(); + if (input_defs.size() > 1) { + model_builder.AddInitializerToSkip(input_defs[1]->Name()); // axes + } +} + +// Add operator related. + +Status ReductionOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get shape"); + const auto input_rank = input_shape.size(); + + NodeAttrHelper helper(node); + const auto keep_dims = helper.Get("keepdims", 1); + emscripten::val options = emscripten::val::object(); + options.set("keepDimensions", keep_dims == 1); + std::vector axes_data; + + emscripten::val output = emscripten::val::object(); + + const auto opset = node.SinceVersion(); + if (opset >= 18) { + // Since opset 18, axes is an optional input. + const auto noop_with_empty_axes = helper.Get("noop_with_empty_axes", 0); + if (input_defs.size() > 1) { + // Optional input axes is provided, use axes initializer data. + const auto& initializers(model_builder.GetInitializerTensors()); + const auto& axes_tensor = *initializers.at(input_defs[1]->Name()); + Initializer axes_initializer(axes_tensor); + const auto axes_data_span = axes_initializer.DataAsSpan(); + std::transform( + axes_data_span.begin(), axes_data_span.end(), std::back_inserter(axes_data), + [input_rank](int64_t axis) -> int32_t { return HandleNegativeAxis(axis, input_rank); }); + } else { + if (noop_with_empty_axes) { + // When axes is empty and this attribute is set to true, input tensor will not be reduced. + output = input; + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); + } + } + } else { + if (helper.HasAttr("axes")) { + auto axes = helper.Get("axes", std::vector{}); + std::transform( + axes.begin(), axes.end(), std::back_inserter(axes_data), + [input_rank](int64_t axis) -> int32_t { return HandleNegativeAxis(axis, input_rank); }); + } + } + if (axes_data.size() > 0) { + options.set("axes", emscripten::val::array(axes_data)); + } + + const auto& op_type = node.OpType(); + if (op_type == "ReduceMax") { + output = model_builder.GetBuilder().call("reduceMax", input, options); + } else if (op_type == "ReduceMean") { + output = model_builder.GetBuilder().call("reduceMean", input, options); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "ReductionOpBuilder, unknown op: ", op_type); + } + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. +bool ReductionOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) + return false; + + return true; +} + +void CreateReductionOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + if (op_registrations.op_builder_map.find(op_type) != op_registrations.op_builder_map.cend()) + return; + + static std::vector op_types = + { + "ReduceMax", + "ReduceMean", + }; + + op_registrations.builders.push_back(std::make_unique()); + for (const auto& type : op_types) { + op_registrations.op_builder_map.emplace(type, op_registrations.builders.back().get()); + } +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/reshape_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/reshape_op_builder.cc index 2e4dc3f4ad..5d1db43653 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/reshape_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/reshape_op_builder.cc @@ -29,7 +29,7 @@ class ReshapeOpBuilder : public BaseOpBuilder { // Operator support related. private: bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const override; + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; // Reshape opset 4- uses attributes for new shape which we do not support for now. int GetMinSupportedOpSet(const Node& /* node */) const override { return 5; } @@ -69,7 +69,9 @@ Status ReshapeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, // Operator support related. -bool ReshapeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, +bool ReshapeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { const auto& input_defs = node.InputDefs(); const auto& perm_name = input_defs[1]->Name(); diff --git a/onnxruntime/core/providers/webnn/builders/impl/resize_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/resize_op_builder.cc index 9b2d0f8e70..2afef28b10 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/resize_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/resize_op_builder.cc @@ -30,7 +30,7 @@ class ResizeOpBuilder : public BaseOpBuilder { // Operator support related. private: bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const override; + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; // Resize opset 10- is very different than Resize opset 11+, with many key attributes missing. // We only support Resize opset 11+ here. @@ -159,7 +159,9 @@ Status ResizeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, // Operator support related. -bool ResizeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, +bool ResizeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const { const auto& input_defs = node.InputDefs(); @@ -255,26 +257,6 @@ bool ResizeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers << ". input_size_c, " << input_shape[c_idx] << ", output_size_c, " << output_sizes[c_idx]; return false; } - - // For now we only support upscale, so the output_size_h and output_size_w should be an integer >= 1. - // TODO support ResizeBilinear - auto output_size_h = output_sizes[2]; - auto output_size_w = output_sizes[3]; - auto input_size_h = input_shape[2]; - auto input_size_w = input_shape[3]; - - // Onnx spec requires output sizes to be a positive integer, so we are not checking that here. - if (output_size_h % input_size_h != 0) { - LOGS(logger, VERBOSE) << "Resize: output_size_h: " << output_size_h - << " is not a multiple of input_size_h: " << input_size_h; - return false; - } - - if (output_size_w % input_size_w != 0) { - LOGS(logger, VERBOSE) << "Resize: output_size_w: " << output_size_w - << " is not a multiple of input_size_w: " << input_size_w; - return false; - } } } diff --git a/onnxruntime/core/providers/webnn/builders/impl/slice_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/slice_op_builder.cc new file mode 100644 index 0000000000..8778bb2414 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/slice_op_builder.cc @@ -0,0 +1,164 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/safeint.h" +#include "core/framework/tensorprotoutils.h" +#include "core/optimizer/initializer.h" +#include "core/providers/common.h" +#include "core/providers/cpu/tensor/slice_helper.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class SliceOpBuilder : public BaseOpBuilder { + // Add operator related. + public: + void AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const override; + + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; + // TODO: Support Slice opset < 10, which uses attributes for starts and ends. + int GetMinSupportedOpSet(const Node& /* node */) const override { return 10; } +}; + +// Add operator related. + +void SliceOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { + // Skip all initializer except the first inputs(data). + for (size_t i = 1; i < node.InputDefs().size(); i++) { + model_builder.AddInitializerToSkip(node.InputDefs()[i]->Name()); + } +} + +Status SliceOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get input shape"); + auto rank = input_shape.size(); + NodeAttrHelper helper(node); + + emscripten::val inputs = model_builder.GetOperand(input_defs[0]->Name()); + std::vector starts(rank); + std::vector sizes(rank); + + // Copy the data from the starts/ends/axes/steps initializers. + std::vector input_starts; + std::vector input_ends; + std::vector input_axes; + std::vector input_steps; + SliceOp::PrepareForComputeMetadata compute_metadata(input_shape); + const auto CopyInputData = [&input_defs, &model_builder, &logger](size_t input_idx, std::vector& data, + bool is_required = false) { + data.clear(); + std::string input_name; + // This is an optional input, return empty vector. + if (!is_required) { + if (input_defs.size() <= input_idx) + return Status::OK(); + input_name = input_defs[input_idx]->Name(); + if (input_name.empty()) + return Status::OK(); + } + input_name = input_defs[input_idx]->Name(); + const auto& initializers(model_builder.GetInitializerTensors()); + const auto& tensor = *initializers.at(input_name); + if (!ReadIntArrayFrom1DTensor(tensor, data, logger)) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, + "Data type for starts and ends inputs is not supported in this build."); + } + + return Status::OK(); + }; + ORT_RETURN_IF_ERROR(CopyInputData(1, input_starts, true)); + ORT_RETURN_IF_ERROR(CopyInputData(2, input_ends, true)); + ORT_RETURN_IF_ERROR(CopyInputData(3, input_axes)); + ORT_RETURN_IF_ERROR(CopyInputData(4, input_steps)); + ORT_RETURN_IF_ERROR( + SliceOp::PrepareForComputeHelper(input_starts, input_ends, input_axes, input_steps, compute_metadata)); + + std::transform(compute_metadata.starts_.cbegin(), compute_metadata.starts_.cend(), + starts.begin(), + [](int64_t i) { return SafeInt(i); }); + std::transform(compute_metadata.ends_.cbegin(), compute_metadata.ends_.cend(), compute_metadata.starts_.cbegin(), + sizes.begin(), + [](int64_t i, int64_t j) { return SafeInt(i - j); }); + + emscripten::val output = model_builder.GetBuilder().call("slice", inputs, + emscripten::val::array(starts), + emscripten::val::array(sizes)); + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +bool SliceOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& name = node.Name(); + const auto& op_type = node.OpType(); + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) { + return false; + } + if (input_defs.size() == 5) { // Check steps. + const auto& steps_tensor = *initializers.at(input_defs[4]->Name()); + std::vector unpacked_tensor; + auto status = onnxruntime::utils::UnpackInitializerData(steps_tensor, unpacked_tensor); + if (!status.IsOK()) { + LOGS(logger, ERROR) << "Error while unpacking steps_tensor: " << status.ErrorMessage(); + return false; + } + const auto data_type = steps_tensor.data_type(); + // WebNN doesn't support steps other than 1. + if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT64) { + if (!std::all_of(reinterpret_cast(unpacked_tensor.data()), + reinterpret_cast(unpacked_tensor.data() + unpacked_tensor.size()), + [](int64_t i) { return i == 1; })) { + return false; + } + } else if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT32) { + if (!std::all_of(reinterpret_cast(unpacked_tensor.data()), + reinterpret_cast(unpacked_tensor.data()) + + unpacked_tensor.size() / sizeof(int32_t), + [](int32_t i) { return i == 1; })) { + return false; + } + } + } + + if (input_defs.size() < 3) { + LOGS(logger, VERBOSE) << op_type << " [" << name << "] requires at least 3 inputs (data starts and ends) but got " + << input_defs.size(); + return false; + } + + const auto& starts_name = input_defs[1]->Name(); + const auto& ends_name = input_defs[2]->Name(); + if (!Contains(initializers, starts_name) || !Contains(initializers, ends_name)) { + LOGS(logger, VERBOSE) << op_type << " [" << name << "] need starts and ends as initializer."; + return false; + } + return true; +} + +void CreateSliceOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/softmax_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/softmax_op_builder.cc new file mode 100644 index 0000000000..e3f481db65 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/softmax_op_builder.cc @@ -0,0 +1,94 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class SoftmaxOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +Status SoftmaxOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + emscripten::val input = model_builder.GetOperand(node.InputDefs()[0]->Name()); + emscripten::val output = emscripten::val::object(); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get shape"); + const auto input_size = input_shape.size(); + // WebNN Softmax only support 2d input shape, reshape input to 2d. + if (input_size != 2) { + int32_t new_shape_0 = input_shape.data()[0]; + for (size_t i = 1; i < input_size - 1; i++) { + new_shape_0 *= input_shape.data()[i]; + } + emscripten::val new_shape = emscripten::val::array(); + new_shape.call("push", new_shape_0); + new_shape.call("push", static_cast(input_shape.back())); + input = model_builder.GetBuilder().call("reshape", input, new_shape); + } + output = model_builder.GetBuilder().call("softmax", input); + // Reshape output to the same shape of input. + if (input_size != 2) { + emscripten::val new_shape = emscripten::val::array(); + for (size_t i = 0; i < input_size; i++) { + new_shape.call("push", static_cast(input_shape[i])); + } + output = model_builder.GetBuilder().call("reshape", output, new_shape); + } + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. + +bool SoftmaxOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) + return false; + const auto input_size = input_shape.size(); + if (input_size < 2) { + LOGS(logger, VERBOSE) << "SoftMax only support input size >= 2d shape, input is " + << input_size << "d shape"; + return false; + } + NodeAttrHelper helper(node); + const int32_t axis = helper.Get("axis", 1); + // WebNN softmax only support input axis 1 + if (axis != 1 && axis != -1) { + LOGS(logger, VERBOSE) << "SoftMax only support axis 1 or -1, input axis: " << axis; + return false; + } + + return true; +} + +void CreateSoftmaxOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/split_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/split_op_builder.cc new file mode 100644 index 0000000000..ff94496133 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/split_op_builder.cc @@ -0,0 +1,182 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/safeint.h" +#include "core/optimizer/initializer.h" +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class SplitOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; + + int GetMinSupportedOpSet(const Node& node) const override; +}; + +Status SplitOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + emscripten::val output_array; + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get shape"); + const auto rank = input_shape.size(); + emscripten::val options = emscripten::val::object(); + + NodeAttrHelper helper(node); + int64_t axis = helper.Get("axis", 0); + axis = HandleNegativeAxis(axis, rank); + options.set("axis", static_cast(axis)); + + if (input_defs.size() == 2) { + // Inputs contains optional 'split' input + std::vector splits; + const auto& initializers(model_builder.GetInitializerTensors()); + const auto& split_tensor = *initializers.at(input_defs[1]->Name()); + ORT_RETURN_IF_NOT(ReadIntArrayFrom1DTensor(split_tensor, splits, logger), "Cannot get split."); + output_array = model_builder.GetBuilder().call("split", + input, + emscripten::val::array(splits), + options); + ORT_RETURN_IF_NOT(output_array["length"].as() == static_cast(splits.size()), + "The size of outputs must be equal to the size of 'split' input."); + } else { + if (helper.HasAttr("num_outputs")) { + const int64_t num_outputs = helper.Get("num_outputs", 1); + ORT_RETURN_IF_NOT(num_outputs > 0, "The 'num_outputs' must be a positive integer."); + if (input_shape[axis] % num_outputs == 0) { + // The 'num_outputs' evenly divide the dim value at 'axis' specified. + output_array = model_builder.GetBuilder().call("split", + input, + static_cast(num_outputs), + options); + } else { + std::vector mapping_split; + mapping_split.insert(mapping_split.begin(), num_outputs - 1, input_shape[axis] / num_outputs); + mapping_split.insert(mapping_split.end(), input_shape[axis] % num_outputs); + std::vector converted_splits; + std::transform(mapping_split.cbegin(), mapping_split.cend(), + std::back_inserter(converted_splits), + [](int64_t dim) -> int32_t { return SafeInt(dim); }); + output_array = model_builder.GetBuilder().call("split", + input, + emscripten::val::array(converted_splits), + options); + } + ORT_RETURN_IF_NOT(output_array["length"].as() == static_cast(num_outputs), + "The size of outputs must be equal to 'num_outputs'."); + } else { + // w/o 'split' input for opset 13 + // Refer to https://github.com/microsoft/onnxruntime/blob/a7ad859e3ab60bddfcf2fefa96bfcb550f0fc04c/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp#L984-L989 + // split input stream equally across output streams. + const auto& output_defs = node.OutputDefs(); + auto output_count = output_defs.size(); + output_array = model_builder.GetBuilder().call("split", + input, static_cast(output_count), + options); + ORT_RETURN_IF_NOT(output_array["length"].as() == static_cast(output_count), + "The size of outputs must be equal to the count of output nodes."); + } + } + for (int64_t i = 0, count = output_array["length"].as(); i < count; i++) { + model_builder.AddOperand(node.OutputDefs()[i]->Name(), std::move(output_array[i])); + } + return Status::OK(); +} + +// Operator support related. + +int SplitOpBuilder::GetMinSupportedOpSet(const Node& /* node */) const { + // Since opset 13, Split has optional 'split' input. + return 13; +} + +bool SplitOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) { + LOGS(logger, VERBOSE) << "Cannot get input's shape."; + return false; + } + const auto rank = input_shape.size(); + + NodeAttrHelper helper(node); + int64_t axis = helper.Get("axis", 0); + axis = HandleNegativeAxis(axis, rank); + + if (input_defs.size() == 2) { + // Inputs contains optional 'split' input + const auto& split_name = input_defs[1]->Name(); + if (!Contains(initializers, split_name)) { + LOGS(logger, VERBOSE) << "The split must be a constant initializer."; + return false; + } + // Values should be >= 0. Sum of the values must be equal to the dim value at 'axis' specified. + std::vector split; + const auto& split_tensor = *initializers.at(input_defs[1]->Name()); + if (split_tensor.data_type() != ONNXTensorElementDataType::ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64) { + LOGS(logger, VERBOSE) << "The type of tensor's element data must be INT64."; + return false; + } + if (!ReadIntArrayFrom1DTensor(split_tensor, split, logger)) { + LOGS(logger, VERBOSE) << "Cannot get split."; + return false; + } + int64_t sum = 0; + for (int64_t i = 0; i < split.size(); i++) { + if (split[i] < 0) { + LOGS(logger, VERBOSE) << "Value of split should be greater than or equal to 0."; + return false; + } + sum += split[i]; + } + if (sum != input_shape[axis]) { + LOGS(logger, VERBOSE) << "Sum of the split's values must be equal to the dim value at 'axis' specified."; + return false; + } + } else if (input_defs.size() == 1) { + if (helper.HasAttr("num_outputs")) { + // Split has 'num_outputs' attribute when opset is 18. + const int32_t num_outputs = helper.Get("num_outputs", 1); + if (num_outputs < 1) { + LOGS(logger, VERBOSE) << "The 'num_outputs' must be a positive integer."; + return false; + } + } else { + const auto opset = node.SinceVersion(); + if (opset >= 18) { + LOGS(logger, VERBOSE) << "The 'num_outputs' should be specified when 'split' isn't specified."; + return false; + } + } + } + return true; +} + +void CreateSplitOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc new file mode 100644 index 0000000000..b243a5c79b --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/unary_op_builder.cc @@ -0,0 +1,73 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include + +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" +#include "core/providers/webnn/builders/helper.h" + +#include "base_op_builder.h" + +namespace onnxruntime { +namespace webnn { + +class UnaryOpBuilder : public BaseOpBuilder { + // Add operator related. + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; +}; + +// Add operator related. + +Status UnaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& /* logger */) const { + const auto& op_type(node.OpType()); + + emscripten::val input = model_builder.GetOperand(node.InputDefs()[0]->Name()); + emscripten::val output = emscripten::val::object(); + if (op_type == "Cos") { + output = model_builder.GetBuilder().call("cos", input); + } else if (op_type == "Erf") { + output = model_builder.GetBuilder().call("erf", input); + } else if (op_type == "Floor") { + output = model_builder.GetBuilder().call("floor", input); + } else if (op_type == "Not") { + output = model_builder.GetBuilder().call("logicalNot", input); + } else if (op_type == "Sin") { + output = model_builder.GetBuilder().call("sin", input); + } else if (op_type == "Sqrt") { + output = model_builder.GetBuilder().call("sqrt", input); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "UnaryOpBuilder::AddToModelBuilderImpl, unknown op: ", op_type); + } + + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +void CreateUnaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + if (op_registrations.op_builder_map.find(op_type) != op_registrations.op_builder_map.cend()) + return; + + static std::vector op_types = + { + "Cos", + "Erf", + "Floor", + "Not", + "Sin", + "Sqrt", + }; + + op_registrations.builders.push_back(std::make_unique()); + for (const auto& type : op_types) { + op_registrations.op_builder_map.emplace(type, op_registrations.builders.back().get()); + } +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc new file mode 100644 index 0000000000..18f14382c7 --- /dev/null +++ b/onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc @@ -0,0 +1,126 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Intel Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/common.h" +#include "core/providers/shared/utils/utils.h" +#include "core/providers/webnn/builders/helper.h" +#include "core/providers/webnn/builders/model_builder.h" +#include "core/providers/webnn/builders/op_builder_factory.h" + +#include "core/optimizer/initializer.h" + +#include "base_op_builder.h" +#include "builder_utils.h" + +namespace onnxruntime { +namespace webnn { + +class UnsqueezeOpBuilder : public BaseOpBuilder { + // Add operator related. + public: + void AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const override; + + private: + Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node, + const logging::Logger& logger) const override ORT_MUST_USE_RESULT; + + // Operator support related. + private: + bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node, + const WebnnDeviceType /* device_type */, const logging::Logger& logger) const override; +}; + +// Add operator related. +void UnsqueezeOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { + // Unsqueeze opset 13 uses input 1 as axes, add it to initializer skip list. + const auto& input_defs = node.InputDefs(); + if (node.SinceVersion() > 12 && input_defs.size() > 1) { + model_builder.AddInitializerToSkip(input_defs[1]->Name()); // "axes" + model_builder.AddInputToSkip(input_defs[1]->Name()); + } +} + +// Add operator related. + +Status UnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + emscripten::val input = model_builder.GetOperand(input_defs[0]->Name()); + std::vector input_shape; + ORT_RETURN_IF_NOT(GetShape(*input_defs[0], input_shape, logger), "Cannot get input shape"); + const auto input_rank = input_shape.size(); + + NodeAttrHelper helper(node); + emscripten::val options = emscripten::val::object(); + std::vector axes_data; + + if (node.SinceVersion() >= 13) { + // Input axes is provided, use axes initializer data. + const auto& initializers = model_builder.GetInitializerTensors(); + const auto& axes_tensor = *initializers.at(input_defs[1]->Name()); + Initializer axes_initializer(axes_tensor); + const auto axes_data_span = axes_initializer.DataAsSpan(); + const auto output_rank = input_rank + axes_data_span.size(); + std::transform( + axes_data_span.begin(), axes_data_span.end(), std::back_inserter(axes_data), + [output_rank](int64_t axis) -> int32_t { return HandleNegativeAxis(axis, output_rank); }); + } else { + if (helper.HasAttr("axes")) { + auto axes = helper.Get("axes", std::vector{}); + const auto output_rank = input_rank + axes.size(); + std::transform( + axes.begin(), axes.end(), std::back_inserter(axes_data), + [output_rank](int64_t axis) -> int32_t { return HandleNegativeAxis(axis, output_rank); }); + } + } + + if (axes_data.size() > 0) { + options.set("axes", emscripten::val::array(axes_data)); + } + + emscripten::val output = model_builder.GetBuilder().call("unsqueeze", input, options); + model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); + return Status::OK(); +} + +// Operator support related. + +bool UnsqueezeOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers, + const Node& node, + const WebnnDeviceType /* device_type */, + const logging::Logger& logger) const { + const auto& input_defs = node.InputDefs(); + std::vector input_shape; + if (!GetShape(*input_defs[0], input_shape, logger)) + return false; + + // Unsqueeze opset 13 uses input 1 as axes, it needs to be an initializer. + if (node.SinceVersion() >= 13) { + if (input_defs.size() < 2) { + LOGS(logger, ERROR) << "Input axes of Unsqueeze must be provided"; + return false; + } + const auto& axes_name = input_defs[1]->Name(); + if (!Contains(initializers, axes_name)) { + LOGS(logger, ERROR) << "Input axes of Unsqueeze must be known"; + return false; + } + } else { + if (input_defs.size() < 1) { + LOGS(logger, ERROR) << "Unsqueeze has no input tensor"; + return false; + } + } + + return true; +} + +void CreateUnsqueezeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) { + op_registrations.builders.push_back(std::make_unique()); + op_registrations.op_builder_map.emplace(op_type, op_registrations.builders.back().get()); +} + +} // namespace webnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/builders/model.cc b/onnxruntime/core/providers/webnn/builders/model.cc index bc385dea6c..b25d00d45a 100644 --- a/onnxruntime/core/providers/webnn/builders/model.cc +++ b/onnxruntime/core/providers/webnn/builders/model.cc @@ -29,13 +29,42 @@ Status Model::Predict(const InlinedHashMap& inputs, for (const auto& input : inputs) { const std::string& name = input.first; const struct OnnxTensorData tensor = input.second; - if (tensor.tensor_info.data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "The input of graph has unsupported type, name: ", - name, " type: ", tensor.tensor_info.data_type); - } auto num_elements = SafeInt(Product(tensor.tensor_info.shape)); - emscripten::val view{emscripten::typed_memory_view(num_elements, static_cast(tensor.buffer))}; + emscripten::val view = emscripten::val::undefined(); + switch (tensor.tensor_info.data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + default: + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "The input of graph has unsupported type, name: ", + name, " type: ", tensor.tensor_info.data_type); + } #ifdef ENABLE_WEBASSEMBLY_THREADS // Copy the inputs from Wasm SharedArrayBuffer to the pre-allocated ArrayBuffers. wnn_inputs_[name].call("set", view); @@ -55,13 +84,43 @@ Status Model::Predict(const InlinedHashMap& inputs, for (const auto& output : outputs) { const std::string& name = output.first; const struct OnnxTensorData tensor = output.second; - if (tensor.tensor_info.data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "The input of graph has unsupported type, name: ", - name, " type: ", tensor.tensor_info.data_type); - } auto num_elements = SafeInt(Product(tensor.tensor_info.shape)); - emscripten::val view{emscripten::typed_memory_view(num_elements, static_cast(tensor.buffer))}; + emscripten::val view = emscripten::val::undefined(); + switch (tensor.tensor_info.data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + view = emscripten::val{emscripten::typed_memory_view(num_elements, + static_cast(tensor.buffer))}; + break; + default: + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "The output of graph has unsupported type, name: ", + name, " type: ", tensor.tensor_info.data_type); + } + #ifdef ENABLE_WEBASSEMBLY_THREADS output_views.insert({name, view}); #else @@ -69,7 +128,6 @@ Status Model::Predict(const InlinedHashMap& inputs, #endif } wnn_context_.call("computeSync", wnn_graph_, wnn_inputs_, wnn_outputs_); - #ifdef ENABLE_WEBASSEMBLY_THREADS // Copy the outputs from pre-allocated ArrayBuffers back to the Wasm SharedArrayBuffer. for (const auto& output : outputs) { @@ -102,16 +160,64 @@ void Model::AllocateInputOutputBuffers() { for (const auto& input : inputs_) { const auto& input_info = input_output_info_.at(input); const auto input_shape = input_info.shape; - const auto num_elements = SafeInt(Product(input_shape)); - wnn_inputs_.set(input, - emscripten::val::global("Float32Array").new_(static_cast(num_elements))); + const int32_t num_elements = SafeInt(Product(input_shape)); + const auto data_type = input_info.data_type; + switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + wnn_inputs_.set(input, emscripten::val::global("Uint8Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + wnn_inputs_.set(input, emscripten::val::global("Uint16Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + wnn_inputs_.set(input, emscripten::val::global("Float32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + wnn_inputs_.set(input, emscripten::val::global("Int32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + wnn_inputs_.set(input, emscripten::val::global("BigInt64Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + wnn_inputs_.set(input, emscripten::val::global("Uint32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + wnn_inputs_.set(input, emscripten::val::global("BigUint64Array").new_(num_elements)); + break; + default: + break; + } } for (const auto& output : outputs_) { const auto& output_info = input_output_info_.at(output); const auto output_shape = output_info.shape; - const auto num_elements = SafeInt(Product(output_shape)); - wnn_outputs_.set(output, - emscripten::val::global("Float32Array").new_(static_cast(num_elements))); + const int32_t num_elements = SafeInt(Product(output_shape)); + const auto data_type = output_info.data_type; + switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + wnn_outputs_.set(output, emscripten::val::global("Uint8Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + wnn_outputs_.set(output, emscripten::val::global("Uint16Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + wnn_outputs_.set(output, emscripten::val::global("Float32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + wnn_outputs_.set(output, emscripten::val::global("Int32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + wnn_outputs_.set(output, emscripten::val::global("BigInt64Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + wnn_outputs_.set(output, emscripten::val::global("Uint32Array").new_(num_elements)); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + wnn_outputs_.set(output, emscripten::val::global("BigUint64Array").new_(num_elements)); + break; + default: + break; + } } } diff --git a/onnxruntime/core/providers/webnn/builders/model_builder.cc b/onnxruntime/core/providers/webnn/builders/model_builder.cc index 92b6b29dca..098193bc63 100644 --- a/onnxruntime/core/providers/webnn/builders/model_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/model_builder.cc @@ -19,12 +19,13 @@ namespace webnn { ModelBuilder::ModelBuilder(const GraphViewer& graph_viewer, const logging::Logger& logger, const emscripten::val& context, const emscripten::val& builder, - const DataLayout preferred_layout) + const DataLayout preferred_layout, const WebnnDeviceType wnn_device_type) : graph_viewer_(graph_viewer), logger_(logger), wnn_context_(context), wnn_builder_(builder), - preferred_layout_(preferred_layout) {} + preferred_layout_(preferred_layout), + wnn_device_type_(wnn_device_type) {} Status ModelBuilder::Initialize() { PreprocessInitializers(); @@ -91,32 +92,66 @@ Status ModelBuilder::RegisterInitializers() { for (const auto& pair : GetInitializerTensors()) { const auto& tensor = *pair.second; const auto& name = tensor.name(); - if (Contains(skipped_initializers_, name)) + // Optional tensors can be indicated by an empty name, just ignore it. + if (name.empty() || Contains(skipped_initializers_, name)) continue; const auto& shape = tensor.dims(); std::vector dims; - if (shape.empty()) { - // This is a scalar initializer, WebNN requires a shape, make this a {1} tensor. - dims = {1}; - } else { - std::transform(shape.cbegin(), shape.cend(), - std::back_inserter(dims), - [](int64_t dim) -> int32_t { return SafeInt(dim); }); - } + // When the shape is empty, it is scalar initializer that dims = {}; + std::transform(shape.cbegin(), shape.cend(), + std::back_inserter(dims), + [](int64_t dim) -> int32_t { return SafeInt(dim); }); emscripten::val desc = emscripten::val::object(); desc.set("dimensions", emscripten::val::array(dims)); auto data_type = tensor.data_type(); emscripten::val operand = emscripten::val::object(); - if (data_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + if (IsSupportedDataType(data_type, wnn_device_type_)) { unpacked_tensors_.push_back({}); std::vector& unpacked_tensor = unpacked_tensors_.back(); ORT_RETURN_IF_ERROR(onnxruntime::utils::UnpackInitializerData(tensor, unpacked_tensor)); auto num_elements = SafeInt(Product(tensor.dims())); - desc.set("type", emscripten::val("float32")); - emscripten::val view{emscripten::typed_memory_view(num_elements, - reinterpret_cast(unpacked_tensor.data()))}; + emscripten::val view = emscripten::val::undefined(); + switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + desc.set("type", emscripten::val("uint8")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + desc.set("type", emscripten::val("float16")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + desc.set("type", emscripten::val("float32")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + desc.set("type", emscripten::val("int32")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + desc.set("type", emscripten::val("int64")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + desc.set("type", emscripten::val("uint32")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + desc.set("type", emscripten::val("uint64")); + view = emscripten::val{emscripten::typed_memory_view(num_elements, + reinterpret_cast(unpacked_tensor.data()))}; + break; + default: + break; + } #ifdef ENABLE_WEBASSEMBLY_THREADS // Workaround for WebAssembly multi-threads enabled since WebNN API only accepts non-shared ArrayBufferView. // https://www.w3.org/TR/webnn/#typedefdef-mlnamedarraybufferviews @@ -188,18 +223,36 @@ Status ModelBuilder::RegisterModelInputOutput(const NodeArg& node_arg, bool is_i const auto* type_proto = node_arg.TypeAsProto(); if (!type_proto || !type_proto->tensor_type().has_elem_type()) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "The ", input_output_type, " of graph doesn't have elem_type: ", name); + "The ", input_output_type, " of graph doesn't have elem_type: ", name); } data_type = type_proto->tensor_type().elem_type(); switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + desc.set("type", emscripten::val("uint8")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + desc.set("type", emscripten::val("float16")); + break; case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: desc.set("type", emscripten::val("float32")); break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + desc.set("type", emscripten::val("int32")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + desc.set("type", emscripten::val("int64")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + desc.set("type", emscripten::val("uint32")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + desc.set("type", emscripten::val("uint64")); + break; default: { // TODO: support other type. return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "The ", input_output_type, " of graph doesn't have valid type, name: ", name, + "The ", input_output_type, " of graph doesn't have valid type, name: ", name, " type: ", type_proto->tensor_type().elem_type()); } } @@ -246,14 +299,53 @@ Status ModelBuilder::AddOperations() { Status ModelBuilder::AddOperandFromPersistMemoryBuffer( const std::string& name, const void* buffer, const size_t size, - const std::vector shape, const size_t element_size) { + const std::vector shape, const int32_t data_type) { auto persist_buffer = std::make_unique(size); uint8_t* dest = persist_buffer.get(); memcpy(dest, buffer, size); - emscripten::val view{emscripten::typed_memory_view(size / element_size, reinterpret_cast(dest))}; + emscripten::val view = emscripten::val::undefined(); emscripten::val desc = emscripten::val::object(); + switch (data_type) { + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(uint8_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("uint8")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(uint16_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("float16")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(float), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("float32")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT32: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(int32_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("int32")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_INT64: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(int64_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("int64")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT32: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(uint32_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("uint32")); + break; + case ONNX_NAMESPACE::TensorProto_DataType_UINT64: + view = emscripten::val{emscripten::typed_memory_view(size / sizeof(uint64_t), + reinterpret_cast(dest))}; + desc.set("type", emscripten::val("uint64")); + break; + default: + break; + } + desc.set("dimensions", emscripten::val::array(shape)); - desc.set("type", emscripten::val("float32")); emscripten::val operand = emscripten::val::object(); #ifdef ENABLE_WEBASSEMBLY_THREADS // Workaround for WebAssembly multi-threads enabled since WebNN API only accepts non-shared ArrayBufferView. diff --git a/onnxruntime/core/providers/webnn/builders/model_builder.h b/onnxruntime/core/providers/webnn/builders/model_builder.h index 7c445b97f3..c381eef3f4 100644 --- a/onnxruntime/core/providers/webnn/builders/model_builder.h +++ b/onnxruntime/core/providers/webnn/builders/model_builder.h @@ -9,6 +9,7 @@ #include "model.h" #include "core/framework/execution_provider.h" +#include "core/providers/webnn/builders/helper.h" #include #include @@ -20,8 +21,9 @@ class IOpBuilder; class ModelBuilder { public: - ModelBuilder(const GraphViewer& graph_viewer, const logging::Logger& logger, const emscripten::val& context, - const emscripten::val& builder, const DataLayout preferred_layout); + ModelBuilder(const GraphViewer& graph_viewer, const logging::Logger& logger, + const emscripten::val& context, const emscripten::val& builder, + const DataLayout preferred_layout, const WebnnDeviceType wnn_device_type); ~ModelBuilder() = default; Status Compile(std::unique_ptr& model) ORT_MUST_USE_RESULT; @@ -40,7 +42,7 @@ class ModelBuilder { // Add a constant operand (allocate persist buffer and move the ownership to mem_persist_buffers_). Status AddOperandFromPersistMemoryBuffer( const std::string& name, const void* buffer, - const size_t size, const std::vector shape, const size_t element_size = 4); + const size_t size, const std::vector shape, const int32_t data_type); // Find if an output has a fuseable activation (e.g., Relu). emscripten::val FindActivation(const Node& node, const NodeArg& output, const InlinedHashSet supported_nodes = {}); @@ -50,6 +52,8 @@ class ModelBuilder { DataLayout GetPreferredLayout() const { return preferred_layout_; } + WebnnDeviceType GetWebnnDeviceType() const { return wnn_device_type_; } + // The initializer will be processed separately, skip it as an initializer. void AddInitializerToSkip(const std::string& tensor_name); @@ -66,6 +70,7 @@ class ModelBuilder { emscripten::val wnn_context_ = emscripten::val::object(); emscripten::val wnn_builder_ = emscripten::val::object(); DataLayout preferred_layout_; + WebnnDeviceType wnn_device_type_; std::vector> unpacked_tensors_; InlinedHashMap wnn_operands_; std::vector input_names_; diff --git a/onnxruntime/core/providers/webnn/builders/op_builder.h b/onnxruntime/core/providers/webnn/builders/op_builder.h index efa70ab3d5..6ecc5d1068 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder.h +++ b/onnxruntime/core/providers/webnn/builders/op_builder.h @@ -2,6 +2,8 @@ // Copyright (c) Intel Corporation. All rights reserved. // Licensed under the MIT License. +#include "core/providers/webnn/builders/helper.h" + #pragma once namespace onnxruntime { @@ -27,7 +29,7 @@ class IOpBuilder { public: // Check if an operator is supported. virtual bool IsOpSupported(const InitializedTensorSet& initializers, const Node& node, - const logging::Logger& logger) const = 0; + const WebnnDeviceType device_type, const logging::Logger& logger) const = 0; }; } // namespace webnn diff --git a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc index c13677547f..e4a7aa69f5 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc +++ b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc @@ -15,11 +15,21 @@ namespace webnn { static OpBuilderRegistrations CreateOpBuilderRegistrations() { OpBuilderRegistrations op_registrations; + { // Unary + CreateUnaryOpBuilder("Cos", op_registrations); + CreateUnaryOpBuilder("Erf", op_registrations); + CreateUnaryOpBuilder("Floor", op_registrations); + CreateUnaryOpBuilder("Not", op_registrations); + CreateUnaryOpBuilder("Sin", op_registrations); + CreateUnaryOpBuilder("Sqrt", op_registrations); + } + { // Binary CreateBinaryOpBuilder("Add", op_registrations); CreateBinaryOpBuilder("Sub", op_registrations); CreateBinaryOpBuilder("Mul", op_registrations); CreateBinaryOpBuilder("Div", op_registrations); + CreateBinaryOpBuilder("Pow", op_registrations); } { // Activations @@ -28,6 +38,15 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { CreateActivationOpBuilder("Sigmoid", op_registrations); } + { // ArgMax/ArgMin + CreateArgMaxMinOpBuilder("ArgMax", op_registrations); + CreateArgMaxMinOpBuilder("ArgMin", op_registrations); + } + + { // Cast + CreateCastOpBuilder("Cast", op_registrations); + } + { // Clip CreateClipOpBuilder("Clip", op_registrations); } @@ -41,8 +60,29 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { CreateConcatOpBuilder("Concat", op_registrations); } - { // Gemm + { // Expand + CreateExpandOpBuilder("Expand", op_registrations); + } + + { // Gather + CreateGatherOpBuilder("Gather", op_registrations); + } + + { // Flatten + CreateFlattenOpBuilder("Flatten", op_registrations); + } + + { // Gemm/MatMul CreateGemmOpBuilder("Gemm", op_registrations); + CreateGemmOpBuilder("MatMul", op_registrations); + } + + { // Logical + CreateLogicalOpBuilder("Equal", op_registrations); + } + + { // LayerNormalization + CreateNormalizationOpBuilder("LayerNormalization", op_registrations); } { // Pool @@ -52,6 +92,11 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { CreatePoolOpBuilder("MaxPool", op_registrations); } + { // Reduction + CreateReductionOpBuilder("ReduceMax", op_registrations); + CreateReductionOpBuilder("ReduceMean", op_registrations); + } + { // Reshape CreateReshapeOpBuilder("Reshape", op_registrations); } @@ -60,10 +105,26 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { CreateResizeOpBuilder("Resize", op_registrations); } + { // Slice + CreateSliceOpBuilder("Slice", op_registrations); + } + + { // Softmax + CreateSoftmaxOpBuilder("Softmax", op_registrations); + } + + { // Split + CreateSplitOpBuilder("Split", op_registrations); + } + { // Transpose CreateTransposeOpBuilder("Transpose", op_registrations); } + { // Unsqueeze + CreateUnsqueezeOpBuilder("Unsqueeze", op_registrations); + } + return op_registrations; } diff --git a/onnxruntime/core/providers/webnn/builders/op_builder_factory.h b/onnxruntime/core/providers/webnn/builders/op_builder_factory.h index ffbbf2d92d..40f96b4bfb 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder_factory.h +++ b/onnxruntime/core/providers/webnn/builders/op_builder_factory.h @@ -20,15 +20,28 @@ struct OpBuilderRegistrations { const InlinedHashMap& GetOpBuilders(); void CreateActivationOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateArgMaxMinOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateBinaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateCastOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateClipOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateConvOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateConcatOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateExpandOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateFlattenOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateGatherOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateGemmOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateLogicalOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateNormalizationOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreatePoolOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateReductionOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateReshapeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateResizeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateSliceOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateSoftmaxOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateSplitOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); void CreateTransposeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateUnaryOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); +void CreateUnsqueezeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations); } // namespace webnn } // namespace onnxruntime diff --git a/onnxruntime/core/providers/webnn/webnn_execution_provider.cc b/onnxruntime/core/providers/webnn/webnn_execution_provider.cc index 9e49192428..d98978cc4a 100644 --- a/onnxruntime/core/providers/webnn/webnn_execution_provider.cc +++ b/onnxruntime/core/providers/webnn/webnn_execution_provider.cc @@ -52,8 +52,10 @@ WebNNExecutionProvider::WebNNExecutionProvider( // WebNN EP uses NHWC layout for CPU XNNPACK backend and NCHW for GPU DML backend. if (webnn_device_flags.compare("cpu") == 0) { preferred_layout_ = DataLayout::NHWC; + wnn_device_type_ = webnn::WebnnDeviceType::CPU; } else { preferred_layout_ = DataLayout::NCHW; + wnn_device_type_ = webnn::WebnnDeviceType::GPU; } if (webnn_power_flags.compare("default") != 0) { context_options.set("powerPreference", emscripten::val(webnn_power_flags)); @@ -100,7 +102,7 @@ WebNNExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view const auto& logger = *GetLogger(); - const auto node_groups = webnn::GetSupportedNodes(graph_viewer, wnn_builder_, logger); + const auto node_groups = webnn::GetSupportedNodes(graph_viewer, wnn_builder_, wnn_device_type_, logger); if (node_groups.empty()) { return result; @@ -130,13 +132,18 @@ WebNNExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view InlinedHashSet subgraph_inputs; InlinedHashSet subgraph_outputs; std::vector ordered_subgraph_inputs; - std::vector ordered_subgraph_outputs; + // Output should be unique. It may be produced as graph output and subgraph output. + InlinedHashSet ordered_subgraph_outputs; for (const auto& index : group) { sub_graph->nodes.push_back(index); const auto* node = graph_viewer.GetNode(index); for (const auto* input : node->InputDefs()) { + if (!input->Exists()) { + // skip the placeholder inputs. + continue; + } // if the node input was not produced by this subgraph, add it to the subgraph inputs. if (node_outputs.count(input) == 0) { if (subgraph_inputs.count(input) == 0) { @@ -151,7 +158,7 @@ WebNNExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view node_outputs.insert(output_def); // if output is overall graph output we need to produce it. if (graph_outputs.count(output_def) != 0) { - ordered_subgraph_outputs.push_back(output_def); + ordered_subgraph_outputs.insert(output_def); } } @@ -161,7 +168,7 @@ WebNNExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view const auto* output_def = output_defs[it->GetSrcArgIndex()]; if (subgraph_outputs.count(output_def) == 0) { subgraph_outputs.insert(output_def); - ordered_subgraph_outputs.push_back(output_def); + ordered_subgraph_outputs.insert(output_def); } } } @@ -213,7 +220,8 @@ common::Status WebNNExecutionProvider::Compile(const std::vector model; ORT_RETURN_IF_ERROR(builder.Compile(model)); // Build map from input name to its index in input definitions. @@ -311,7 +319,13 @@ common::Status WebNNExecutionProvider::Compile(const std::vector #include @@ -44,6 +45,7 @@ class WebNNExecutionProvider : public IExecutionProvider { emscripten::val wnn_builder_ = emscripten::val::object(); DataLayout preferred_layout_; + webnn::WebnnDeviceType wnn_device_type_; InlinedHashMap> models_; }; } // namespace onnxruntime