mirror of
https://github.com/opencv/opencv.git
synced 2026-07-30 15:53:03 +04:00
Merge pull request #27845 from abhishek-gola:bitwise_layer_add
Added Bitwise layer to new DNN engine #27845 ### Pull Request Readiness Checklist See details at https://github.com/opencv/opencv/wiki/How_to_contribute#making-a-good-pull-request - [x] I agree to contribute to the project under Apache 2 License. - [x] To the best of my knowledge, the proposed patch is not based on a code under GPL or another license that is incompatible with OpenCV - [x] The PR is proposed to the proper branch - [x] There is a reference to the original bug report and related work - [x] There is accuracy test, performance test and test data in opencv_extra repository, if applicable Patch to opencv_extra has the same branch name. - [x] The feature is well documented and sample code can be built with the project CMake
This commit is contained in:
@@ -195,6 +195,9 @@ public:
|
||||
ADD,
|
||||
DIV,
|
||||
WHERE,
|
||||
BITWISE_AND,
|
||||
BITWISE_OR,
|
||||
BITWISE_XOR,
|
||||
} op;
|
||||
|
||||
NaryEltwiseLayerImpl(const LayerParams& params)
|
||||
@@ -245,6 +248,12 @@ public:
|
||||
op = OPERATION::XOR;
|
||||
else if (operation == "where")
|
||||
op = OPERATION::WHERE;
|
||||
else if (operation == "bitwise_and")
|
||||
op = OPERATION::BITWISE_AND;
|
||||
else if (operation == "bitwise_or")
|
||||
op = OPERATION::BITWISE_OR;
|
||||
else if (operation == "bitwise_xor")
|
||||
op = OPERATION::BITWISE_XOR;
|
||||
else
|
||||
CV_Error(cv::Error::StsBadArg, "Unknown operation type \"" + operation + "\"");
|
||||
}
|
||||
@@ -361,6 +370,18 @@ public:
|
||||
return;
|
||||
}
|
||||
|
||||
if (op == OPERATION::BITWISE_AND || op == OPERATION::BITWISE_OR || op == OPERATION::BITWISE_XOR)
|
||||
{
|
||||
CV_Assert(inputs.size());
|
||||
for (auto input : inputs)
|
||||
{
|
||||
CV_CheckTypeEQ(inputs[0], input, "All inputs should have equal types");
|
||||
CV_CheckType(input, input == CV_8S || input == CV_8U || input == CV_16S || input == CV_16U || input == CV_32S || input == CV_32U || input == CV_64S || input == CV_64U, "");
|
||||
}
|
||||
outputs.assign(requiredOutputs, inputs[0]);
|
||||
return;
|
||||
}
|
||||
|
||||
if (op == OPERATION::POW) {
|
||||
CV_Assert(inputs.size() == 2);
|
||||
auto isIntegerType = [](int t) {
|
||||
@@ -807,9 +828,9 @@ public:
|
||||
}
|
||||
|
||||
template<typename T, typename... Args>
|
||||
inline void opDispatch(size_t ninputs, Args&&... args)
|
||||
inline typename std::enable_if<std::is_integral<T>::value && !std::is_same<T, bool>::value, void>::type opDispatch(size_t ninputs, Args&&... args)
|
||||
{
|
||||
if (ninputs == 2) { // Operators that take two operands
|
||||
if (ninputs == 2) {
|
||||
switch (op) {
|
||||
case OPERATION::EQUAL: {
|
||||
auto equal = [](const T &a, const T &b) { return a == b; };
|
||||
@@ -892,6 +913,136 @@ public:
|
||||
binary_forward<T, T>(div, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::BITWISE_AND: {
|
||||
auto band = [](const T &a, const T &b) { return (T)(a & b); };
|
||||
binary_forward<T, T>(band, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::BITWISE_OR: {
|
||||
auto bor = [](const T &a, const T &b) { return (T)(a | b); };
|
||||
binary_forward<T, T>(bor, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::BITWISE_XOR: {
|
||||
auto bxor = [](const T &a, const T &b) { return (T)(a ^ b); };
|
||||
binary_forward<T, T>(bxor, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
default: CV_Error(Error::StsBadArg, "Unsupported operation");
|
||||
}
|
||||
} else if (ninputs == 3 && op == OPERATION::WHERE) {
|
||||
auto where = [](const T &a, const T &b, const T &c) { return a ? b : c; };
|
||||
ternary_forward<bool, T, T, T>(where, std::forward<Args>(args)...);
|
||||
} else {
|
||||
switch (op)
|
||||
{
|
||||
case OPERATION::MAX: {
|
||||
auto max = [](const T &a, const T &b) { return std::max(a, b); };
|
||||
nary_forward<T>(max, T{1}, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MEAN: {
|
||||
auto sum = [](const T &a, const T &b) { return a + b; };
|
||||
nary_forward<T>(sum, T{1} / ninputs, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MIN: {
|
||||
auto min = [](const T &a, const T &b) { return std::min(a, b); };
|
||||
nary_forward<T>(min, T{1}, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::SUM: {
|
||||
auto sum = [](const T &a, const T &b) { return a + b; };
|
||||
nary_forward<T>(sum, T{1}, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
default:
|
||||
CV_Error(Error::StsBadArg, "Unsupported operation.");
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
template<typename T, typename... Args>
|
||||
inline typename std::enable_if<!std::is_integral<T>::value || std::is_same<T, bool>::value, void>::type opDispatch(size_t ninputs, Args&&... args)
|
||||
{
|
||||
if (ninputs == 2) { // Operators that take two operands
|
||||
switch (op) {
|
||||
case OPERATION::EQUAL: {
|
||||
auto equal = [](const T &a, const T &b) { return a == b; };
|
||||
binary_forward<T, bool>(equal, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::GREATER: {
|
||||
auto greater = [](const T &a, const T &b) { return a > b; };
|
||||
binary_forward<T, bool>(greater, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::GREATER_EQUAL: {
|
||||
auto greater_equal = [](const T &a, const T &b) { return a >= b; };
|
||||
binary_forward<T, bool>(greater_equal, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::LESS: {
|
||||
auto less = [](const T &a, const T &b) { return a < b; };
|
||||
binary_forward<T, bool>(less, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::LESS_EQUAL: {
|
||||
auto less_equal = [](const T &a, const T &b) { return a <= b; };
|
||||
binary_forward<T, bool>(less_equal, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::POW: {
|
||||
auto pow = [] (const T& a, const T& b) { return saturate_cast<T>(std::pow((double)a, (double)b)); };
|
||||
binary_forward<T, T>(pow, std::forward<Args>(args)..., 1e5);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MAX: {
|
||||
auto max = [](const T &a, const T &b) { return std::max(a, b); };
|
||||
binary_forward<T, T>(max, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MEAN: {
|
||||
auto mean = [](const T &a, const T &b) { return (a + b) / T{2}; };
|
||||
binary_forward<T, T>(mean, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MIN: {
|
||||
auto min = [](const T &a, const T &b) { return std::min(a, b); };
|
||||
binary_forward<T, T>(min, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::MOD: {
|
||||
auto mod = [] (const T &a, const T &b) { return static_cast<T>(_mod(int(a), int(b))); };
|
||||
binary_forward<T, T>(mod, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::FMOD: {
|
||||
auto fmod = [](const T &a, const T &b) { return std::fmod(a, b); };
|
||||
binary_forward<T, T>(fmod, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::PROD: {
|
||||
auto prod = [](const T &a, const T &b) { return a * b; };
|
||||
binary_forward<T, T>(prod, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::SUB: {
|
||||
auto sub = [](const T &a, const T &b) { return a - b; };
|
||||
binary_forward<T, T>(sub, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::ADD:
|
||||
case OPERATION::SUM: {
|
||||
auto sum = [](const T &a, const T &b) { return a + b; };
|
||||
binary_forward<T, T>(sum, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
case OPERATION::DIV: {
|
||||
auto div = [](const T &a, const T &b) { return a / b; };
|
||||
binary_forward<T, T>(div, std::forward<Args>(args)...);
|
||||
break;
|
||||
}
|
||||
default: CV_Error(Error::StsBadArg, "Unsupported operation");
|
||||
}
|
||||
} else if (ninputs == 3 && op == OPERATION::WHERE) { // Operators that take three operands
|
||||
@@ -977,12 +1128,16 @@ public:
|
||||
break;
|
||||
case CV_32F:
|
||||
CV_Assert(op != OPERATION::BITSHIFT && op != OPERATION::AND &&
|
||||
op != OPERATION::OR && op != OPERATION::XOR);
|
||||
op != OPERATION::OR && op != OPERATION::XOR &&
|
||||
op != OPERATION::BITWISE_AND && op != OPERATION::BITWISE_OR &&
|
||||
op != OPERATION::BITWISE_XOR);
|
||||
opDispatch<float>(std::forward<Args>(args)...);
|
||||
break;
|
||||
case CV_64F:
|
||||
CV_Assert(op != OPERATION::BITSHIFT && op != OPERATION::AND &&
|
||||
op != OPERATION::OR && op != OPERATION::XOR);
|
||||
op != OPERATION::OR && op != OPERATION::XOR &&
|
||||
op != OPERATION::BITWISE_AND && op != OPERATION::BITWISE_OR &&
|
||||
op != OPERATION::BITWISE_XOR);
|
||||
opDispatch<double>(std::forward<Args>(args)...);
|
||||
break;
|
||||
case CV_16S:
|
||||
|
||||
@@ -40,8 +40,10 @@ public:
|
||||
std::vector<MatType>& outputs,
|
||||
std::vector<MatType>& internals) const CV_OVERRIDE
|
||||
{
|
||||
CV_CheckTypeEQ(inputs[0], CV_Bool, "");
|
||||
outputs.assign(1, CV_Bool);
|
||||
int t = inputs[0];
|
||||
bool isInt = (t == CV_8S || t == CV_8U || t == CV_16S || t == CV_16U || t == CV_32S || t == CV_32U || t == CV_64S || t == CV_64U);
|
||||
CV_CheckType(inputs[0], t == CV_Bool || isInt, "Not/BitwiseNot expects bool or integer tensor");
|
||||
outputs.assign(1, t);
|
||||
}
|
||||
|
||||
void forward(InputArrayOfArrays inputs_arr, OutputArrayOfArrays outputs_arr, OutputArrayOfArrays internals_arr) CV_OVERRIDE
|
||||
@@ -52,19 +54,25 @@ public:
|
||||
std::vector<Mat> inputs, outputs;
|
||||
inputs_arr.getMatVector(inputs);
|
||||
outputs_arr.getMatVector(outputs);
|
||||
|
||||
CV_CheckTypeEQ(inputs[0].type(), CV_Bool, "");
|
||||
CV_CheckTypeEQ(outputs[0].type(), CV_Bool, "");
|
||||
|
||||
CV_Assert(inputs[0].isContinuous());
|
||||
CV_Assert(outputs[0].isContinuous());
|
||||
|
||||
bool* input = inputs[0].ptr<bool>();
|
||||
bool* output = outputs[0].ptr<bool>();
|
||||
int t = inputs[0].type();
|
||||
size_t size = inputs[0].total();
|
||||
if (t == CV_Bool)
|
||||
{
|
||||
bool* input = inputs[0].ptr<bool>();
|
||||
bool* output = outputs[0].ptr<bool>();
|
||||
for (size_t i = 0; i < size; ++i)
|
||||
output[i] = !input[i];
|
||||
return;
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < size; ++i)
|
||||
output[i] = !input[i];
|
||||
const unsigned char* in = inputs[0].ptr<unsigned char>();
|
||||
unsigned char* out = outputs[0].ptr<unsigned char>();
|
||||
size_t totalBytes = size * inputs[0].elemSize();
|
||||
for (size_t i = 0; i < totalBytes; ++i)
|
||||
out[i] = static_cast<unsigned char>(~in[i]);
|
||||
}
|
||||
|
||||
#ifdef HAVE_DNN_NGRAPH
|
||||
|
||||
@@ -239,6 +239,8 @@ protected:
|
||||
void parseNonMaxSuprression (LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto);
|
||||
void parseTopK2 (LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto);
|
||||
void parseBitShift (LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto);
|
||||
void parseBitwise (LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto);
|
||||
void parseBitwiseNot (LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto);
|
||||
|
||||
// Domain: com.microsoft
|
||||
// URL: https://github.com/microsoft/onnxruntime/blob/master/docs/ContribOperators.md
|
||||
@@ -1680,6 +1682,27 @@ void ONNXImporter2::parseBitShift(LayerParams& layerParams, const opencv_onnx::N
|
||||
addLayer(layerParams, node_proto);
|
||||
}
|
||||
|
||||
void ONNXImporter2::parseBitwise(LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto)
|
||||
{
|
||||
const std::string& op_type = node_proto.op_type();
|
||||
layerParams.type = "NaryEltwise";
|
||||
if (op_type == "BitwiseAnd")
|
||||
layerParams.set("operation", String("bitwise_and"));
|
||||
else if (op_type == "BitwiseOr")
|
||||
layerParams.set("operation", String("bitwise_or"));
|
||||
else if (op_type == "BitwiseXor")
|
||||
layerParams.set("operation", String("bitwise_xor"));
|
||||
else
|
||||
CV_Error(Error::StsNotImplemented, String("Unsupported bitwise op: ") + op_type);
|
||||
addLayer(layerParams, node_proto);
|
||||
}
|
||||
|
||||
void ONNXImporter2::parseBitwiseNot(LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto)
|
||||
{
|
||||
layerParams.type = "Not";
|
||||
addLayer(layerParams, node_proto);
|
||||
}
|
||||
|
||||
void ONNXImporter2::parseTrilu(LayerParams& layerParams, const opencv_onnx::NodeProto& node_proto)
|
||||
{
|
||||
int ninputs = node_proto.input_size();
|
||||
@@ -2563,6 +2586,10 @@ void ONNXImporter2::buildDispatchMap_ONNX_AI(int opset_version)
|
||||
dispatch["GridSample"] = &ONNXImporter2::parseGridSample;
|
||||
dispatch["Upsample"] = &ONNXImporter2::parseUpsample;
|
||||
dispatch["BitShift"] = &ONNXImporter2::parseBitShift;
|
||||
dispatch["BitwiseAnd"] = &ONNXImporter2::parseBitwise;
|
||||
dispatch["BitwiseOr"] = &ONNXImporter2::parseBitwise;
|
||||
dispatch["BitwiseXor"] = &ONNXImporter2::parseBitwise;
|
||||
dispatch["BitwiseNot"] = &ONNXImporter2::parseBitwiseNot;
|
||||
dispatch["NonMaxSuprression"] = &ONNXImporter2::parseNonMaxSuprression;
|
||||
dispatch["SoftMax"] = dispatch["Softmax"] = dispatch["LogSoftmax"] = &ONNXImporter2::parseSoftMax;
|
||||
dispatch["DetectionOutput"] = &ONNXImporter2::parseDetectionOutput;
|
||||
|
||||
@@ -272,6 +272,36 @@ CASE(test_bernoulli_seed)
|
||||
// no filter
|
||||
CASE(test_bernoulli_seed_expanded)
|
||||
// no filter
|
||||
CASE(test_bitwise_and_i16_3d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_and_i32_2d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_and_ui64_bcast_3v1d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_and_ui8_bcast_4v3d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_not_2d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_not_3d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_not_4d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_or_i16_4d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_or_i32_2d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_or_ui64_bcast_3v1d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_or_ui8_bcast_4v3d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_xor_i16_3d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_xor_i32_2d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_xor_ui64_bcast_3v1d)
|
||||
SKIP;
|
||||
CASE(test_bitwise_xor_ui8_bcast_4v3d)
|
||||
SKIP;
|
||||
CASE(test_bitshift_left_uint16)
|
||||
SKIP;
|
||||
CASE(test_bitshift_left_uint32)
|
||||
|
||||
@@ -307,3 +307,18 @@
|
||||
"test_gelu_default_2_expanded",
|
||||
"test_gelu_tanh_1_expanded",
|
||||
"test_gelu_tanh_2_expanded",
|
||||
"test_bitwise_and_i16_3d",
|
||||
"test_bitwise_and_i32_2d",
|
||||
"test_bitwise_and_ui64_bcast_3v1d",
|
||||
"test_bitwise_and_ui8_bcast_4v3d",
|
||||
"test_bitwise_not_2d",
|
||||
"test_bitwise_not_3d",
|
||||
"test_bitwise_not_4d",
|
||||
"test_bitwise_or_i16_4d",
|
||||
"test_bitwise_or_i32_2d",
|
||||
"test_bitwise_or_ui64_bcast_3v1d",
|
||||
"test_bitwise_or_ui8_bcast_4v3d",
|
||||
"test_bitwise_xor_i16_3d",
|
||||
"test_bitwise_xor_i32_2d",
|
||||
"test_bitwise_xor_ui64_bcast_3v1d",
|
||||
"test_bitwise_xor_ui8_bcast_4v3d",
|
||||
|
||||
@@ -164,21 +164,6 @@
|
||||
"test_bernoulli_expanded", // ---- same as above ---
|
||||
"test_bernoulli_seed", // ---- same as above ---
|
||||
"test_bernoulli_seed_expanded", // ---- same as above ---
|
||||
"test_bitwise_and_i16_3d",
|
||||
"test_bitwise_and_i32_2d",
|
||||
"test_bitwise_and_ui64_bcast_3v1d",
|
||||
"test_bitwise_and_ui8_bcast_4v3d",
|
||||
"test_bitwise_not_2d",
|
||||
"test_bitwise_not_3d",
|
||||
"test_bitwise_not_4d",
|
||||
"test_bitwise_or_i16_4d",
|
||||
"test_bitwise_or_i32_2d",
|
||||
"test_bitwise_or_ui64_bcast_3v1d",
|
||||
"test_bitwise_or_ui8_bcast_4v3d",
|
||||
"test_bitwise_xor_i16_3d",
|
||||
"test_bitwise_xor_i32_2d",
|
||||
"test_bitwise_xor_ui64_bcast_3v1d",
|
||||
"test_bitwise_xor_ui8_bcast_4v3d",
|
||||
"test_blackmanwindow",
|
||||
"test_blackmanwindow_expanded",
|
||||
"test_blackmanwindow_symmetric",
|
||||
|
||||
Reference in New Issue
Block a user