diff --git a/onnxruntime/core/providers/dnnl/dnnl_op_manager.cc b/onnxruntime/core/providers/dnnl/dnnl_op_manager.cc index ae98b04033..734b591149 100644 --- a/onnxruntime/core/providers/dnnl/dnnl_op_manager.cc +++ b/onnxruntime/core/providers/dnnl/dnnl_op_manager.cc @@ -17,6 +17,7 @@ DnnlOpManager::DnnlOpManager() { dnnl_ops_map_.emplace(std::make_pair("Div", std::unique_ptr(new DnnlBinaryNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("DynamicQuantizeLinear", std::unique_ptr(new DnnlDynamicQuantizeLinearNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("Elu", std::unique_ptr(new DnnlElementwiseCapability()))); + dnnl_ops_map_.emplace(std::make_pair("Equal", std::unique_ptr(new DnnlBinaryNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("Erf", std::unique_ptr(new DnnlErfNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("Exp", std::unique_ptr(new DnnlElementwiseCapability()))); dnnl_ops_map_.emplace(std::make_pair("FastGelu", std::unique_ptr(new DnnlDefaultNodeCapability()))); @@ -25,8 +26,12 @@ DnnlOpManager::DnnlOpManager() { dnnl_ops_map_.emplace(std::make_pair("Gemm", std::unique_ptr(new DnnlGemmNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("GlobalAveragePool", std::unique_ptr(new DnnlPoolNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("GlobalMaxPool", std::unique_ptr(new DnnlPoolNodeCapability()))); + dnnl_ops_map_.emplace(std::make_pair("Greater", std::unique_ptr(new DnnlBinaryNodeCapability()))); + dnnl_ops_map_.emplace(std::make_pair("GreaterOrEqual", std::unique_ptr(new DnnlBinaryNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("LayerNormalization", std::unique_ptr(new DnnlLayerNormalizationNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("LeakyRelu", std::unique_ptr(new DnnlElementwiseCapability()))); + dnnl_ops_map_.emplace(std::make_pair("Less", std::unique_ptr(new DnnlBinaryNodeCapability()))); + dnnl_ops_map_.emplace(std::make_pair("LessOrEqual", std::unique_ptr(new DnnlBinaryNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("Log", std::unique_ptr(new DnnlElementwiseCapability()))); dnnl_ops_map_.emplace(std::make_pair("LRN", std::unique_ptr(new DnnlDefaultNodeCapability()))); dnnl_ops_map_.emplace(std::make_pair("MatMul", std::unique_ptr(new DnnlMatMulNodeCapability()))); diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph.cc index 12521905d6..5edc03cd27 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph.cc @@ -104,6 +104,9 @@ dnnl::memory::data_type DnnlTensor::Type() const { return dnnl::memory::data_type::s8; case ONNX_NAMESPACE::TensorProto_DataType_UINT8: return dnnl::memory::data_type::u8; + // Same here, we use u8 as the handler for bool + case ONNX_NAMESPACE::TensorProto_DataType_BOOL: + return dnnl::memory::data_type::u8; default: ORT_THROW("Unsupported data type: ", data_type); } diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph_primitive.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph_primitive.cc index ad3a28c1b3..e27b61f766 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph_primitive.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_subgraph_primitive.cc @@ -163,7 +163,7 @@ int Product(dnnl::memory::dims d) { } void DnnlSubgraphPrimitive::AddKernels() { - std::unordered_set binary_ops = {"Add", "Div", "Mul", "Sub"}; + std::unordered_set binary_ops = {"Add", "Div", "Equal", "Greater", "GreaterOrEqual", "Less", "LessOrEqual", "Mul", "Sub"}; std::unordered_set elementwise_ops = {"Abs", "Elu", "Exp", "LeakyRelu", "Log", "Relu", "Round", "Sigmoid", "Softplus", "Sqrt", "Tanh"}; std::unordered_set pool_ops = {"AveragePool", "GlobalAveragePool", "GlobalMaxPool", "MaxPool"}; std::unordered_set reduce_ops = {"ReduceL1", "ReduceL2", "ReduceLogSum", "ReduceLogSumExp", "ReduceMax", "ReduceMean", "ReduceMin", "ReduceProd", "ReduceSum", "ReduceSumSquare"}; diff --git a/onnxruntime/core/providers/dnnl/subgraph/dnnl_util.cc b/onnxruntime/core/providers/dnnl/subgraph/dnnl_util.cc index 37c0c335db..b9b63b25bb 100644 --- a/onnxruntime/core/providers/dnnl/subgraph/dnnl_util.cc +++ b/onnxruntime/core/providers/dnnl/subgraph/dnnl_util.cc @@ -21,10 +21,15 @@ dnnl::algorithm OrtOperatorToDnnlAlgorithm(std::string op) { {"Abs", dnnl::algorithm::eltwise_abs}, {"BiasGelu", dnnl::algorithm::eltwise_gelu_erf}, {"Elu", dnnl::algorithm::eltwise_elu}, // algorithm requires alpha value + {"Equal", dnnl::algorithm::binary_eq}, {"Exp", dnnl::algorithm::eltwise_exp}, {"FastGelu", dnnl::algorithm::eltwise_gelu_tanh}, {"Gelu", dnnl::algorithm::eltwise_gelu_erf}, + {"Greater", dnnl::algorithm::binary_gt}, + {"GreaterOrEqual", dnnl::algorithm::binary_ge}, {"LeakyRelu", dnnl::algorithm::eltwise_relu}, // algorithm requires alpha value + {"Less", dnnl::algorithm::binary_lt}, + {"LessOrEqual", dnnl::algorithm::binary_le}, {"Log", dnnl::algorithm::eltwise_log}, {"Relu", dnnl::algorithm::eltwise_relu}, {"Round", dnnl::algorithm::eltwise_round},