Attention fusion: update uint8 tensor parsing for ONNX upgrade (#9564)

* use UnpackTensor to parse uint8 tensor
* address review feedback
This commit is contained in:
Tianlei Wu 2021-10-28 10:38:10 -07:00 committed by GitHub
parent 17cf39a964
commit a555740708
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -2,6 +2,7 @@
// Licensed under the MIT License.
#include "onnx/defs/shape_inference.h"
#include "onnx/defs/tensor_proto_util.h"
#include "core/framework/tensorprotoutils.h"
#pragma once
@ -322,8 +323,23 @@ bool ValidateUnidirMask(const Graph& graph, const NodeArg& mask, bool& is_unidir
}
if (tensor_proto->data_type() == ONNX_NAMESPACE::TensorProto_DataType_UINT8) {
std::vector<int32_t> int32_data = ONNX_NAMESPACE::ParseData<int32_t>(tensor_proto);
if (!ValidateUnidirMask(int32_data, shape->dim(2).dim_value(), is_unidirectional)) {
size_t bytes;
if (!utils::GetSizeInBytesFromTensorProto<0>(*tensor_proto, &bytes).IsOK()) {
return false;
}
auto data = std::make_unique<uint8_t[]>(bytes);
uint8_t* p = data.get();
if (!utils::UnpackTensor<uint8_t>(
*tensor_proto,
tensor_proto->raw_data().size() ? tensor_proto->raw_data().data() : nullptr,
tensor_proto->raw_data().size(),
p,
bytes)
.IsOK()) {
return false;
}
std::vector<uint8_t> mask_data(p, p + bytes);
if (!ValidateUnidirMask(mask_data, shape->dim(2).dim_value(), is_unidirectional)) {
DEBUG_LOG("Mask is neither unidirectional nor all ones");
return false;
}