[WebNN EP] Support rest Reduce* ops (#16824)

Add ReduceL1, ReduceL2, ReduceLogSum, ReduceLogSumExp, ReduceMin,
ReduceProd, ReduceSum, ReduceSumSquare.
This commit is contained in:
Wanming Lin 2023-07-26 08:26:48 +08:00 committed by GitHub
parent 0c1a5098dc
commit d0df83e408
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 52 additions and 5 deletions

View file

@ -171,8 +171,16 @@ static const InlinedHashMap<std::string, std::string> op_map = {
{"Pow", "pow"},
{"PRelu", "prelu"},
{"Reciprocal", "reciprocal"},
{"ReduceL1", "reduceL1"},
{"ReduceL2", "reduceL2"},
{"ReduceLogSum", "reduceLogSum"},
{"ReduceLogSumExp", "reduceLogSumExp"},
{"ReduceMax", "reduceMax"},
{"ReduceMean", "reduceMean"},
{"ReduceMin", "reduceMin"},
{"ReduceProd", "reduceProduct"},
{"ReduceSum", "reduceSum"},
{"ReduceSumSquare", "reduceSumSquare"},
{"Relu", "relu"},
{"Reshape", "reshape"},
{"Resize", "resample2d"},

View file

@ -61,8 +61,9 @@ Status ReductionOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder,
emscripten::val output = emscripten::val::object();
const auto opset = node.SinceVersion();
if (opset >= 18) {
// Since opset 18, axes is an optional input.
const auto& op_type = node.OpType();
if (opset >= 18 || (op_type == "ReduceSum" && opset >= 13)) {
// '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.
@ -93,11 +94,26 @@ Status ReductionOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder,
options.set("axes", emscripten::val::array(axes_data));
}
const auto& op_type = node.OpType();
if (op_type == "ReduceMax") {
if (op_type == "ReduceL1") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceL1", input, options);
} else if (op_type == "ReduceL2") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceL2", input, options);
} else if (op_type == "ReduceLogSum") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceLogSum", input, options);
} else if (op_type == "ReduceLogSumExp") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceLogSumExp", input, options);
} else if (op_type == "ReduceMax") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceMax", input, options);
} else if (op_type == "ReduceMean") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceMean", input, options);
} else if (op_type == "ReduceMin") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceMin", input, options);
} else if (op_type == "ReduceProd") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceProduct", input, options);
} else if (op_type == "ReduceSum") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceSum", input, options);
} else if (op_type == "ReduceSumSquare") {
output = model_builder.GetBuilder().call<emscripten::val>("reduceSumSquare", input, options);
} else {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "ReductionOpBuilder, unknown op: ", op_type);
}
@ -107,7 +123,7 @@ Status ReductionOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder,
}
// Operator support related.
bool ReductionOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */,
bool ReductionOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& initializers,
const Node& node,
const WebnnDeviceType /* device_type */,
const logging::Logger& logger) const {
@ -117,6 +133,13 @@ bool ReductionOpBuilder::IsOpSupportedImpl(const InitializedTensorSet& /* initia
if (!GetShape(*input_defs[0], input_shape, logger))
return false;
const auto& op_type = node.OpType();
// If the optional input 'axes' is provided, it must be an initializer.
if (input_defs.size() > 1 && !Contains(initializers, input_defs[1]->Name())) {
LOGS(logger, VERBOSE) << "Input axes of " << op_type << " must be a constant";
return false;
}
return true;
}
@ -126,8 +149,16 @@ void CreateReductionOpBuilder(const std::string& op_type, OpBuilderRegistrations
static std::vector<std::string> op_types =
{
"ReduceL1",
"ReduceL2",
"ReduceLogSum",
"ReduceLogSumExp",
"ReduceMax",
"ReduceMean",
"ReduceMin",
"ReduceProd",
"ReduceSum",
"ReduceSumSquare",
};
op_registrations.builders.push_back(std::make_unique<ReductionOpBuilder>());

View file

@ -119,8 +119,16 @@ static OpBuilderRegistrations CreateOpBuilderRegistrations() {
}
{ // Reduction
CreateReductionOpBuilder("ReduceL1", op_registrations);
CreateReductionOpBuilder("ReduceL2", op_registrations);
CreateReductionOpBuilder("ReduceLogSum", op_registrations);
CreateReductionOpBuilder("ReduceLogSumExp", op_registrations);
CreateReductionOpBuilder("ReduceMax", op_registrations);
CreateReductionOpBuilder("ReduceMean", op_registrations);
CreateReductionOpBuilder("ReduceMin", op_registrations);
CreateReductionOpBuilder("ReduceProd", op_registrations);
CreateReductionOpBuilder("ReduceSum", op_registrations);
CreateReductionOpBuilder("ReduceSumSquare", op_registrations);
}
{ // Reshape