From 38b640c797613e2396f2975ccd4d8ff0e95a5baa Mon Sep 17 00:00:00 2001 From: Wanming Lin Date: Thu, 30 Nov 2023 00:00:23 +0800 Subject: [PATCH] [WebNN EP] Re-implement Unsqueeze, Squeeze, Flatten with WebNN's reshape (#18585) WebNN will not provide `unsqueeze`, `squeeze`, `flatten2d` ops, as it can be easily implemented by reshape. --- .../core/providers/webnn/builders/helper.h | 6 +-- .../webnn/builders/impl/flatten_op_builder.cc | 20 ++++++--- .../impl/squeeze_unsqueeze_op_builder.cc | 43 ++++++++++++++----- 3 files changed, 49 insertions(+), 20 deletions(-) diff --git a/onnxruntime/core/providers/webnn/builders/helper.h b/onnxruntime/core/providers/webnn/builders/helper.h index 28b54b9c9c..617108c57d 100644 --- a/onnxruntime/core/providers/webnn/builders/helper.h +++ b/onnxruntime/core/providers/webnn/builders/helper.h @@ -153,7 +153,7 @@ static const InlinedHashMap op_map = { {"Erf", {"erf", false}}, {"Exp", {"exp", false}}, {"Expand", {"expand", false}}, - {"Flatten", {"flattenTo2d", false}}, + {"Flatten", {"reshape", true}}, {"Floor", {"floor", true}}, {"Gather", {"gather", false}}, {"Gemm", {"gemm", true}}, @@ -206,12 +206,12 @@ static const InlinedHashMap op_map = { {"Softmax", {"softmax", true}}, {"Split", {"split", true}}, {"Sqrt", {"sqrt", false}}, - {"Squeeze", {"squeeze", false}}, + {"Squeeze", {"reshape", true}}, {"Sub", {"sub", true}}, {"Tan", {"tan", false}}, {"Tanh", {"tanh", true}}, {"Transpose", {"transpose", true}}, - {"Unsqueeze", {"unsqueeze", false}}, + {"Unsqueeze", {"reshape", true}}, {"Where", {"elementwiseIf", false}}, }; diff --git a/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc index 6c59ca451f..f0df27b523 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/flatten_op_builder.cc @@ -36,14 +36,20 @@ Status FlattenOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, 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; - } + axis = HandleNegativeAxis(axis, rank); + + // Use WebNN's reshape to implement Flatten. + int64_t num_pre_axis_elements = std::accumulate( + input_shape.begin(), input_shape.begin() + static_cast(axis), 1, std::multiplies()); + int64_t num_post_axis_elements = std::accumulate( + input_shape.begin() + static_cast(axis), input_shape.end(), 1, std::multiplies()); + + std::vector new_shape = {SafeInt(num_pre_axis_elements), + SafeInt(num_post_axis_elements)}; + emscripten::val inputs = model_builder.GetOperand(input_defs[0]->Name()); - emscripten::val output = model_builder.GetBuilder().call("flattenTo2d", inputs, - static_cast(axis)); + emscripten::val output = model_builder.GetBuilder().call( + "reshape", inputs, emscripten::val::array(new_shape)); model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); return Status::OK(); diff --git a/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc b/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc index 1c0258944d..2a1672c001 100644 --- a/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc +++ b/onnxruntime/core/providers/webnn/builders/impl/squeeze_unsqueeze_op_builder.cc @@ -56,6 +56,7 @@ Status SqueezeUnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_buil emscripten::val options = emscripten::val::object(); std::vector axes_data; + auto rank = input_rank; if (node.SinceVersion() >= 13 && input_defs.size() > 1) { // Input axes is provided, use axes initializer data. @@ -63,35 +64,57 @@ Status SqueezeUnsqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_buil 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(); + if (op_type == "Unsqueeze") { + // Unsqueeze should check the expanded rank. + 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 SafeInt(HandleNegativeAxis(axis, output_rank)); }); + [rank](int64_t axis) -> int32_t { return SafeInt(HandleNegativeAxis(axis, rank)); }); } else { NodeAttrHelper helper(node); if (helper.HasAttr("axes")) { auto axes = helper.Get("axes", std::vector{}); - const auto output_rank = input_rank + axes.size(); + if (op_type == "Unsqueeze") { + // Unsqueeze should check the expanded rank. + rank = input_rank + axes.size(); + } std::transform( axes.begin(), axes.end(), std::back_inserter(axes_data), - [output_rank](int64_t axis) -> int32_t { return SafeInt(HandleNegativeAxis(axis, output_rank)); }); + [rank](int64_t axis) -> int32_t { return SafeInt(HandleNegativeAxis(axis, rank)); }); } } - if (axes_data.size() > 0) { - options.set("axes", emscripten::val::array(axes_data)); - } - emscripten::val output = emscripten::val::undefined(); + // Use WebNN's reshape to implement Squeeze/Unsqueeze. + std::vector new_shape; + std::transform( + input_shape.begin(), input_shape.end(), std::back_inserter(new_shape), + [](int64_t data) -> uint32_t { return SafeInt(data); }); + // Sort axes_data in ascending order. + std::sort(axes_data.begin(), axes_data.end()); if (op_type == "Squeeze") { - output = model_builder.GetBuilder().call("squeeze", input, options); + if (!axes_data.empty()) { + for (auto axis = axes_data.rbegin(); axis != axes_data.rend(); ++axis) { + size_t index = *axis; + new_shape.erase(new_shape.begin() + index); + } + } else { + // Remove all the single dimensions. + new_shape.erase( + std::remove_if(new_shape.begin(), new_shape.end(), [](uint32_t axis) { return axis == 1; }), new_shape.end()); + } } else if (op_type == "Unsqueeze") { - output = model_builder.GetBuilder().call("unsqueeze", input, options); + // Expand new_shape according to axes_data. + for (const int32_t& axis : axes_data) { + new_shape.insert(new_shape.begin() + axis, 1); + } } else { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "SqueezeUnsqueezeOpBuilder::AddToModelBuilderImpl, unknown op: ", op_type); } + output = model_builder.GetBuilder().call("reshape", input, emscripten::val::array(new_shape)); model_builder.AddOperand(node.OutputDefs()[0]->Name(), std::move(output)); return Status::OK(); }