mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Upgrade onehot to OpSet 11 (#2185)
* Upgrade onehot to OpSet 11 * Move Onehot test out of blacklist * Add negative indices support besides negative axis. * PR comments - 1 * PR comments-2
This commit is contained in:
parent
69970d1f2a
commit
aae18a3fe3
5 changed files with 106 additions and 29 deletions
|
|
@ -239,15 +239,16 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sign);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Erf);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_float_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int32_t_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_float_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_int32_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_float, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_float_float, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_int32_t_float, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_float_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int32_t_float_int32_t, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int32_t_float_float, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_float_float, OneHot);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, MaxUnpool);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sinh);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Cosh);
|
||||
|
|
@ -410,6 +411,16 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Ga
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unique);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, TopK);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_int64_t_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_string_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_float_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_int32_t_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_float_int64_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t_float_int32_t, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t_float_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_float_float, OneHot);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, Resize);
|
||||
|
||||
Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
||||
|
|
@ -814,24 +825,26 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sign)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Erf)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_int32_t_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_float_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int32_t_float_int32_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int32_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
MaxUnpool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sinh)>,
|
||||
|
|
@ -1047,6 +1060,26 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unique)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, TopK)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
float_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
float_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
float_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_int32_t_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_float_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int32_t_float_int32_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int32_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, Resize)>,
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -28,10 +28,22 @@ namespace onnxruntime {
|
|||
// spec: https://github.com/onnx/onnx/blob/master/docs/Operators.md#OneHot
|
||||
|
||||
// T1: indices, T2: depth, T3: values
|
||||
#define REG_TYPED_ONE_HOT_OP(types_str, in_type, out_type, depth_type) \
|
||||
#define REG_TYPED_ONE_HOT_OP_V9_10(types_str, in_type, out_type, depth_type) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( \
|
||||
OneHot, \
|
||||
9, 10, \
|
||||
types_str, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<in_type>()) \
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<depth_type>()) \
|
||||
.TypeConstraint("T3", DataTypeImpl::GetTensorType<out_type>()), \
|
||||
OneHotOp<in_type, out_type, depth_type>);
|
||||
|
||||
// T1: indices, T2: depth, T3: values
|
||||
#define REG_TYPED_ONE_HOT_OP_V11(types_str, in_type, out_type, depth_type) \
|
||||
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
|
||||
OneHot, \
|
||||
9, \
|
||||
11, \
|
||||
types_str, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<in_type>()) \
|
||||
|
|
@ -39,8 +51,9 @@ namespace onnxruntime {
|
|||
.TypeConstraint("T3", DataTypeImpl::GetTensorType<out_type>()), \
|
||||
OneHotOp<in_type, out_type, depth_type>);
|
||||
|
||||
#define REG_ONE_HOT_OP(in_type, out_type, depth_type) \
|
||||
REG_TYPED_ONE_HOT_OP(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type)
|
||||
#define REG_ONE_HOT_OP(in_type, out_type, depth_type) \
|
||||
REG_TYPED_ONE_HOT_OP_V9_10(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type); \
|
||||
REG_TYPED_ONE_HOT_OP_V11(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type)
|
||||
|
||||
REG_ONE_HOT_OP(int64_t, int64_t, int64_t);
|
||||
REG_ONE_HOT_OP(float, int64_t, int64_t);
|
||||
|
|
@ -51,6 +64,7 @@ REG_ONE_HOT_OP(int32_t, float, int32_t);
|
|||
REG_ONE_HOT_OP(int32_t, float, float);
|
||||
REG_ONE_HOT_OP(float, float, float); // added this to satisfy onnx model tests
|
||||
REG_ONE_HOT_OP(int64_t, int32_t, float); // added this to satisfy onnx model tests
|
||||
REG_ONE_HOT_OP(int64_t, float, float); // added this to satisfy onnx model tests
|
||||
|
||||
Status ValidateInputs(const Tensor* depth,
|
||||
const Tensor* values) {
|
||||
|
|
@ -129,7 +143,7 @@ Status OneHotOp<in_type, out_type, depth_type>::Compute(OpKernelContext* p_op_ke
|
|||
const auto output_rank = static_cast<int64_t>(indices_num_dims + 1);
|
||||
if (axis_ >= output_rank || axis_ < -output_rank) {
|
||||
std::ostringstream oss;
|
||||
oss << "'axis' attribute must have a value in the range [" << -output_rank
|
||||
oss << "'axis' attribute must have a value in the range [" << -output_rank
|
||||
<< "," << indices_num_dims << "]";
|
||||
return Status(ONNXRUNTIME, INVALID_ARGUMENT, oss.str());
|
||||
}
|
||||
|
|
@ -152,8 +166,23 @@ Status OneHotOp<in_type, out_type, depth_type>::Compute(OpKernelContext* p_op_ke
|
|||
|
||||
// Split indices into matrix of size prefix_dim_size x suffix_dim_size
|
||||
Eigen::array<Eigen::DenseIndex, 2> indices_dims_e = {
|
||||
{static_cast<Eigen::DenseIndex>(prefix_dim_size), static_cast<Eigen::DenseIndex>(suffix_dim_size)}};
|
||||
{static_cast<Eigen::DenseIndex>(prefix_dim_size), static_cast<Eigen::DenseIndex>(suffix_dim_size)}};
|
||||
|
||||
// Handle negative indices. It's faster to create a new indices instead of comparing in generator
|
||||
// since generator has much larger loops.
|
||||
const auto* indices_data = indices->Data<in_type>();
|
||||
const auto indices_size = indices_shape.Size();
|
||||
std::vector<in_type> adjusted_indices;
|
||||
adjusted_indices.reserve(indices_size);
|
||||
for (int64_t i = 0; i < indices_size; ++i)
|
||||
{
|
||||
if (indices_data[i] < 0)
|
||||
adjusted_indices.push_back(indices_data[i] + static_cast<in_type>(depth_val));
|
||||
else
|
||||
adjusted_indices.push_back(indices_data[i]);
|
||||
}
|
||||
indices_data = adjusted_indices.data();
|
||||
|
||||
typename EigenTensorTypes<in_type, 2>::ConstEigenTensorMap indices_tensor_e(indices_data, indices_dims_e);
|
||||
|
||||
// Split output into 3-Tensor of size:
|
||||
|
|
|
|||
|
|
@ -446,10 +446,6 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
|
|||
{"sequence_model2", "SequenceConstruct not implemented yet"},
|
||||
{"sequence_model1", "Sequence* not implemented yet"},
|
||||
{"scatter_elements_with_negative_indices", "ScatterElements(11) not implemented yet"},
|
||||
{"onehot_without_axis", "OneHot(11) not implemented yet"},
|
||||
{"onehot_with_negative_axis", "OneHot(11) not implemented yet"},
|
||||
{"onehot_with_axis", "OneHot(11) not implemented yet"},
|
||||
{"onehot_negative_indices", "OneHot(11) not implemented yet"},
|
||||
{"bitshift_right_uint8", "BitShift(11) uint8 support not enabled currently"},
|
||||
{"bitshift_right_uint16", "BitShift(11) uint16 support not enabled currently"},
|
||||
{"bitshift_left_uint8", "BitShift(11) uint8 support not enabled currently"},
|
||||
|
|
|
|||
|
|
@ -206,5 +206,25 @@ TEST(OneHotOpTest, FloatString) {
|
|||
"off", "off", "off", "off", "off", "off", "on", "off", "off", "off",});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(OneHotOpTest, Axis_Negative_NegIndex_NonDefault) {
|
||||
OpTester test("OneHot", 11);
|
||||
int64_t axis = -3;
|
||||
test.AddAttribute("axis", axis);
|
||||
test.AddInput<int64_t>("indices", {2, 3}, {1, -1, 8, 2, 4, 6});
|
||||
test.AddInput<int64_t>("depth", {1}, {10});
|
||||
test.AddInput<int64_t>("values", {2}, {0, 1});
|
||||
test.AddOutput<int64_t>("output", {10, 2, 3}, { 0, 0, 0, 0, 0, 0,
|
||||
1, 0, 0, 0, 0, 0,
|
||||
0, 0, 0, 1, 0, 0,
|
||||
0, 0, 0, 0, 0, 0,
|
||||
0, 0, 0, 0, 1, 0,
|
||||
0, 0, 0, 0, 0, 0,
|
||||
0, 0, 0, 0, 0, 1,
|
||||
0, 0, 0, 0, 0, 0,
|
||||
0, 0, 1, 0, 0, 0,
|
||||
0, 1, 0, 0, 0, 0,});
|
||||
test.Run();
|
||||
}
|
||||
}
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -151,7 +151,6 @@ def create_backend_test(testname=None):
|
|||
'^test_resize_upsample_sizes_nearest_round_prefer_ceil_asymmetric_cpu',
|
||||
'^test_sequence_*',
|
||||
'^test_scatter_.*',
|
||||
'^test_onehot_.*',
|
||||
'^test_edge_pad_cpu.*', # test data type `int32_t` not supported yet, the `float` equivalent is covered via unit tests
|
||||
'^test_reflect_pad_cpu.*' # test data type `int32_t` not supported yet, the `float` equivalent is covered via unit tests
|
||||
]
|
||||
|
|
|
|||
Loading…
Reference in a new issue