mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
Uncomment mod tests and make infra run them (#971)
Add np.uint32, np.float16 to numpy to MLDataType mappings. Fix up the mapping and get rid of hash defines in the map. Uncomment mod tests.
This commit is contained in:
parent
f2999cf4c2
commit
34e8bea487
3 changed files with 28 additions and 14 deletions
|
|
@ -24,12 +24,14 @@ int OnnxRuntimeTensorToNumpyType(const DataTypeImpl* tensor_type) {
|
|||
static std::map<MLDataType, int> type_map{
|
||||
{DataTypeImpl::GetType<bool>(), NPY_BOOL},
|
||||
{DataTypeImpl::GetType<float>(), NPY_FLOAT},
|
||||
{DataTypeImpl::GetType<MLFloat16>(), NPY_FLOAT16},
|
||||
{DataTypeImpl::GetType<double>(), NPY_DOUBLE},
|
||||
{DataTypeImpl::GetType<int32_t>(), NPY_INT},
|
||||
{DataTypeImpl::GetType<int8_t>(), NPY_INT8},
|
||||
{DataTypeImpl::GetType<uint8_t>(), NPY_UINT8},
|
||||
{DataTypeImpl::GetType<int16_t>(), NPY_INT16},
|
||||
{DataTypeImpl::GetType<uint16_t>(), NPY_UINT16},
|
||||
{DataTypeImpl::GetType<int32_t>(), NPY_INT},
|
||||
{DataTypeImpl::GetType<uint32_t>(), NPY_UINT},
|
||||
{DataTypeImpl::GetType<int64_t>(), NPY_LONGLONG},
|
||||
{DataTypeImpl::GetType<uint64_t>(), NPY_ULONGLONG},
|
||||
{DataTypeImpl::GetType<std::string>(), NPY_OBJECT},
|
||||
|
|
@ -47,15 +49,32 @@ const DataTypeImpl* NumpyToOnnxRuntimeTensorType(int numpy_type) {
|
|||
static std::map<int, MLDataType> type_map{
|
||||
{NPY_BOOL, DataTypeImpl::GetType<bool>()},
|
||||
{NPY_FLOAT, DataTypeImpl::GetType<float>()},
|
||||
// Special, not a C type expands to enum value of 16
|
||||
{NPY_FLOAT16, DataTypeImpl::GetType<MLFloat16>()},
|
||||
{NPY_DOUBLE, DataTypeImpl::GetType<double>()},
|
||||
{NPY_INT, DataTypeImpl::GetType<int32_t>()},
|
||||
{NPY_INT8, DataTypeImpl::GetType<int8_t>()},
|
||||
{NPY_UINT8, DataTypeImpl::GetType<uint8_t>()},
|
||||
{NPY_INT16, DataTypeImpl::GetType<int16_t>()},
|
||||
{NPY_UINT16, DataTypeImpl::GetType<uint16_t>()},
|
||||
// We don't want to use size specific types such
|
||||
// as NPY_INT32 bc they are not enums but hash defines
|
||||
// which may map into other enums and may conflict with other entries here
|
||||
// also NPY docs define these sizes as platform specific, thus we
|
||||
// choose to do some rudimentary checks for proper mapping on C++ size
|
||||
{NPY_BYTE, DataTypeImpl::GetType<int8_t>()},
|
||||
{NPY_UBYTE, DataTypeImpl::GetType<uint8_t>()},
|
||||
{NPY_SHORT, sizeof(short) == sizeof(int16_t) ? DataTypeImpl::GetType<int16_t>()
|
||||
: DataTypeImpl::GetType<int32_t>()},
|
||||
{NPY_USHORT, sizeof(unsigned short) == sizeof(uint16_t) ? DataTypeImpl::GetType<uint16_t>()
|
||||
: DataTypeImpl::GetType<uint32_t>()},
|
||||
{NPY_INT,
|
||||
sizeof(int) == sizeof(int32_t) ? DataTypeImpl::GetType<int32_t>()
|
||||
: DataTypeImpl::GetType<int64_t>()},
|
||||
{NPY_UINT, sizeof(int) == sizeof(int32_t) ? DataTypeImpl::GetType<uint32_t>()
|
||||
: DataTypeImpl::GetType<uint64_t>()},
|
||||
|
||||
{NPY_LONG,
|
||||
sizeof(long) == sizeof(int) ? DataTypeImpl::GetType<int32_t>()
|
||||
: DataTypeImpl::GetType<int64_t>()},
|
||||
sizeof(long) == sizeof(int32_t) ? DataTypeImpl::GetType<int32_t>()
|
||||
: DataTypeImpl::GetType<int64_t>()},
|
||||
{NPY_ULONG,
|
||||
sizeof(unsigned long) == sizeof(uint32_t) ? DataTypeImpl::GetType<uint32_t>()
|
||||
: DataTypeImpl::GetType<uint64_t>()},
|
||||
{NPY_LONGLONG, DataTypeImpl::GetType<int64_t>()},
|
||||
{NPY_ULONGLONG, DataTypeImpl::GetType<uint64_t>()},
|
||||
{NPY_UNICODE, DataTypeImpl::GetType<std::string>()},
|
||||
|
|
|
|||
|
|
@ -354,8 +354,7 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
|
|||
{"tf_mobilenet_v2_1.4_224", "result mismatch"},
|
||||
{"tf_mobilenet_v1_1.0_224", "result mismatch"},
|
||||
{"mobilenetv2-1.0", "result mismatch"},
|
||||
{"mxnet_arcface", "result mismatch"},
|
||||
{"mod_float_mixed_sign_example", "faulty test"}
|
||||
{"mxnet_arcface", "result mismatch"}
|
||||
};
|
||||
|
||||
#ifdef USE_NGRAPH
|
||||
|
|
|
|||
|
|
@ -87,13 +87,9 @@ backend_test.exclude(r'('
|
|||
'|^test_mvn_cpu.*'
|
||||
'|^test_qlinearconv_cpu.*'
|
||||
'|^test_quantizelinear_cpu.*'
|
||||
'|^test_mod_float_mixed_sign_example.*'
|
||||
'|^test_reversesequence_batch_cpu.*'
|
||||
'|^test_reversesequence_time_cpu.*'
|
||||
'|^test_roialign_cpu.*'
|
||||
'|^test_mod_mixed_sign_float16_cpu.*'
|
||||
'|^test_mod_uint32_cpu.*'
|
||||
'|^test_mod_uint64_cpu.*'
|
||||
')')
|
||||
|
||||
# import all test cases at global scope to make
|
||||
|
|
|
|||
Loading…
Reference in a new issue