diff --git a/onnxruntime/core/providers/webnn/builders/helper.h b/onnxruntime/core/providers/webnn/builders/helper.h index 85a03050c3..728317f595 100644 --- a/onnxruntime/core/providers/webnn/builders/helper.h +++ b/onnxruntime/core/providers/webnn/builders/helper.h @@ -137,6 +137,7 @@ static const InlinedHashMap op_map = { {"Resize", "resample2d"}, {"Shape", "slice"}, {"Split", "split"}, + {"Squeeze", "squeeze"}, {"Transpose", "transpose"}, {"Unsqueeze", "unsqueeze"}, }; diff --git a/onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc similarity index 55% rename from onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc rename to onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc index 18f14382c7..e9ad02da7f 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/unsqueeze_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc @@ -16,7 +16,7 @@ namespace onnxruntime { namespace webnn { -class UnsqueezeOpBuilder : public BaseOpBuilder { +class SqueezeUnsqueezeOpBuilder : public BaseOpBuilder { // Add operator related. public: void AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const override; @@ -32,10 +32,10 @@ class UnsqueezeOpBuilder : public BaseOpBuilder { }; // 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. +void SqueezeUnsqueezeOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) const { + // Squeeze/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) { + if (node.SinceVersion() >= 13 && input_defs.size() > 1) { model_builder.AddInitializerToSkip(input_defs[1]->Name()); // "axes" model_builder.AddInputToSkip(input_defs[1]->Name()); } @@ -43,20 +43,20 @@ void UnsqueezeOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, cons // Add operator related. -Status UnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, - const Node& node, - const logging::Logger& logger) const { +Status SqueezeUnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, + const Node& node, + const logging::Logger& logger) const { + const auto& op_type = node.OpType(); 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) { + if (node.SinceVersion() >= 13 && input_defs.size() > 1) { // Input axes is provided, use axes initializer data. const auto& initializers = model_builder.GetInitializerTensors(); const auto& axes_tensor = *initializers.at(input_defs[1]->Name()); @@ -67,6 +67,7 @@ Status UnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, 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 { + NodeAttrHelper helper(node); if (helper.HasAttr("axes")) { auto axes = helper.Get("axes", std::vector{}); const auto output_rank = input_rank + axes.size(); @@ -80,46 +81,69 @@ Status UnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, options.set("axes", emscripten::val::array(axes_data)); } - emscripten::val output = model_builder.GetBuilder().call("unsqueeze", input, options); + emscripten::val output = emscripten::val::undefined(); + if (op_type == "Squeeze") { + output = model_builder.GetBuilder().call("squeeze", input, options); + } else if (op_type == "Unsqueeze") { + output = model_builder.GetBuilder().call("unsqueeze", input, options); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "SqueezeUnsqueezeOpBuilder::AddToModelBuilderImpl, unknown op: ", op_type); + } + 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 { +bool SqueezeUnsqueezeOpBuilder::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(); 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 (input_defs.size() < 1) { + LOGS(logger, ERROR) << op_type << " has no input tensor"; + return false; + } + + // Squeeze/Unsqueeze opset 13 uses input 1 as axes, it needs to be an initializer. if (node.SinceVersion() >= 13) { - if (input_defs.size() < 2) { + if (input_defs.size() > 1) { + const auto& axes_name = input_defs[1]->Name(); + if (!Contains(initializers, axes_name)) { + LOGS(logger, ERROR) << "Input axes of " << op_type << " is not present and constant"; + return false; + } + } else if (op_type == "Unsqueeze") { + // The axes are optional for Squeeze, but not Unsqueeze. 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()); +void CreateSqueezeUnsqueezeOpBuilder(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 = + { + "Squeeze", + "Unsqueeze", + }; + + 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 diff --git a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc index 33e6f0b860..77bc21561a 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc +++ b/onnxruntime/core/providers/webnn/builders/op_builder_factory.cc @@ -126,12 +126,13 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() { CreateSplitOpBuilder("Split", op_registrations); } - { // Transpose - CreateTransposeOpBuilder("Transpose", op_registrations); + { // Squeeze/Unsqueeze + CreateSqueezeUnsqueezeOpBuilder("Squeeze", op_registrations); + CreateSqueezeUnsqueezeOpBuilder("Unsqueeze", op_registrations); } - { // Unsqueeze - CreateUnsqueezeOpBuilder("Unsqueeze", op_registrations); + { // Transpose + CreateTransposeOpBuilder("Transpose", 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 b178f22b6c..e51d212e51 100644 --- a/onnxruntime/core/providers/webnn/builders/op_builder_factory.h +++ b/onnxruntime/core/providers/webnn/builders/op_builder_factory.h @@ -40,9 +40,9 @@ void CreateShapeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op 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 CreateSqueezeUnsqueezeOpBuilder(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