mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
[WebNN EP] Support rest Reduce* ops (#16824)
Add ReduceL1, ReduceL2, ReduceLogSum, ReduceLogSumExp, ReduceMin, ReduceProd, ReduceSum, ReduceSumSquare.
This commit is contained in:
parent
0c1a5098dc
commit
d0df83e408
3 changed files with 52 additions and 5 deletions
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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>());
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue