aboutsummaryrefslogtreecommitdiff
path: root/src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp
diff options
context:
space:
mode:
Diffstat (limited to 'src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp')
-rw-r--r--src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp3
1 files changed, 2 insertions, 1 deletions
diff --git a/src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp b/src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp
index ca42cacb39..c45ab2cded 100644
--- a/src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp
+++ b/src/armnn/optimizations/ConvertFp32NetworkToBf16.hpp
@@ -31,7 +31,8 @@ inline LayerT* ConvertWeight(Layer* l)
info.GetNumElements(),
newValues.data());
- TensorInfo newInfo(info.GetShape(), DataType::BFloat16);
+ TensorInfo newInfo(info);
+ newInfo.SetDataType(DataType::BFloat16);
ConstTensor newInput(newInfo, newValues);
layer->m_Weight.reset(new ScopedCpuTensorHandle(newInput));
}