aboutsummaryrefslogtreecommitdiff
path: root/src/backends/neon/workloads/NeonWorkloadUtils.hpp
diff options
context:
space:
mode:
Diffstat (limited to 'src/backends/neon/workloads/NeonWorkloadUtils.hpp')
-rw-r--r--src/backends/neon/workloads/NeonWorkloadUtils.hpp38
1 files changed, 37 insertions, 1 deletions
diff --git a/src/backends/neon/workloads/NeonWorkloadUtils.hpp b/src/backends/neon/workloads/NeonWorkloadUtils.hpp
index f9c3718e14..9f8bb9540e 100644
--- a/src/backends/neon/workloads/NeonWorkloadUtils.hpp
+++ b/src/backends/neon/workloads/NeonWorkloadUtils.hpp
@@ -1,5 +1,5 @@
//
-// Copyright © 2022 Arm Ltd and Contributors. All rights reserved.
+// Copyright © 2017,2022 Arm Ltd and Contributors. All rights reserved.
// SPDX-License-Identifier: MIT
//
#pragma once
@@ -58,6 +58,42 @@ void CopyArmComputeTensorData(arm_compute::Tensor& dstTensor, const T* srcData)
}
inline void InitializeArmComputeTensorData(arm_compute::Tensor& tensor,
+ TensorInfo tensorInfo,
+ const ITensorHandle* handle)
+{
+ ARMNN_ASSERT(handle);
+
+ switch(tensorInfo.GetDataType())
+ {
+ case DataType::Float16:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const armnn::Half*>(handle->Map()));
+ break;
+ case DataType::Float32:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const float*>(handle->Map()));
+ break;
+ case DataType::QAsymmU8:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const uint8_t*>(handle->Map()));
+ break;
+ case DataType::QSymmS8:
+ case DataType::QAsymmS8:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const int8_t*>(handle->Map()));
+ break;
+ case DataType::Signed32:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const int32_t*>(handle->Map()));
+ break;
+ case DataType::QSymmS16:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const int16_t*>(handle->Map()));
+ break;
+ case DataType::BFloat16:
+ CopyArmComputeTensorData(tensor, reinterpret_cast<const armnn::BFloat16*>(handle->Map()));
+ break;
+ default:
+ // Throw exception; assertion not called in release build.
+ throw Exception("Unexpected tensor type during InitializeArmComputeTensorData().");
+ }
+};
+
+inline void InitializeArmComputeTensorData(arm_compute::Tensor& tensor,
const ConstTensorHandle* handle)
{
ARMNN_ASSERT(handle);