Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 26 additions & 9 deletions onnxruntime/core/framework/tensorprotoutils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1443,7 +1443,9 @@ common::Status GetSizeInBytesFromTensorTypeProto(const ONNX_NAMESPACE::TypeProto

template Status GetSizeInBytesFromTensorTypeProto<0>(const ONNX_NAMESPACE::TypeProto_Tensor& tensor_proto, size_t* out);

common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto) {
common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(
const ONNX_NAMESPACE::TensorProto& tensor_proto,
size_t max_embedded_initializer_size_in_bytes) {
ORT_RETURN_IF(HasExternalData(tensor_proto), "Expected to validate an embedded (non-external) TensorProto");

TensorShape tensor_shape = GetTensorShapeFromTensorProto(tensor_proto);
Expand All @@ -1468,21 +1470,32 @@ common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE:
ORT_RETURN_IF_ERROR(GetSizeInBytesFromTensorElemCountAndType<0>(num_elems_unsigned, tensor_proto.data_type(),
&byte_size_from_shape));
}
ORT_RETURN_IF_NOT(byte_size_from_shape <= kMaxEmbeddedInitializerSizeInBytes,
ORT_RETURN_IF_NOT(byte_size_from_shape <= max_embedded_initializer_size_in_bytes,
"Initializer '", tensor_proto.name(), "' declares a size of ", byte_size_from_shape,
" bytes which exceeds the ", kMaxEmbeddedInitializerSizeInBytes,
" bytes which exceeds the ", max_embedded_initializer_size_in_bytes,
" byte limit for embedded initializer data. Use external data for large initializers.");

if (HasRawData(tensor_proto)) {
ORT_RETURN_IF_NOT(tensor_proto.raw_data().size() == byte_size_from_shape,
"Initializer '", tensor_proto.name(), "': raw_data size (", tensor_proto.raw_data().size(),
" bytes) does not match expected size from shape and data type (",
byte_size_from_shape, " bytes)");
} else if (HasString(tensor_proto)) {
if (HasString(tensor_proto)) {
ORT_RETURN_IF(HasRawData(tensor_proto),
"Initializer '", tensor_proto.name(), "': string tensor can not have raw data");
ORT_RETURN_IF_NOT(tensor_proto.string_data_size() == num_elems_signed,
"Initializer '", tensor_proto.name(), "': string_data count (", tensor_proto.string_data_size(),
") does not match expected count from shape (",
num_elems_signed, ")");

size_t total_string_storage_size = byte_size_from_shape;
for (const auto& string_data : tensor_proto.string_data()) {
ORT_RETURN_IF(string_data.size() > max_embedded_initializer_size_in_bytes - total_string_storage_size,
Comment thread
danielsongmicrosoft marked this conversation as resolved.
"Initializer '", tensor_proto.name(), "': string_data shape bytes + payload exceeds the ",
max_embedded_initializer_size_in_bytes,
" byte limit for embedded initializer data. Use external data for large initializers.");
total_string_storage_size += string_data.size();
}
} else if (HasRawData(tensor_proto)) {
ORT_RETURN_IF_NOT(tensor_proto.raw_data().size() == byte_size_from_shape,
"Initializer '", tensor_proto.name(), "': raw_data size (", tensor_proto.raw_data().size(),
" bytes) does not match expected size from shape and data type (",
byte_size_from_shape, " bytes)");
} else {
// Typed data fields. Each data type maps to a specific repeated field in the proto.
int64_t expected_count = 0;
Expand Down Expand Up @@ -1564,6 +1577,10 @@ common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE:
return Status::OK();
}

common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto) {
return ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kMaxEmbeddedInitializerSizeInBytes);
}

TensorShape GetTensorShapeFromTensorShapeProto(const ONNX_NAMESPACE::TensorShapeProto& tensor_shape_proto) {
const auto& dims = tensor_shape_proto.dim();
TensorShapeVector tensor_shape_vec(static_cast<size_t>(dims.size()));
Expand Down
5 changes: 5 additions & 0 deletions onnxruntime/core/framework/tensorprotoutils.h
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,11 @@ Status GetSizeInBytesFromTensorTypeProto(const ONNX_NAMESPACE::TypeProto_Tensor&
/// declared shape when the actual data is absent or much smaller.
/// The caller must ensure that the TensorProto does not use external data; if it does, this function will
/// return an error status.
/// max_embedded_initializer_size_in_bytes limits the total in-memory initializer size. For STRING tensors,
/// the limit includes both the std::string object storage implied by the shape and all string payload bytes.
common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto,
Comment thread
danielsongmicrosoft marked this conversation as resolved.
size_t max_embedded_initializer_size_in_bytes);

common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto);

/**
Expand Down
63 changes: 63 additions & 0 deletions onnxruntime/test/framework/tensorutils_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,69 @@ TEST(TensorProtoUtilsTest, UnpackTensor) {
EXPECT_FALSE(status.IsOK());
}

namespace {
TensorProto CreateStringTensorProto(std::initializer_list<size_t> payload_sizes) {
TensorProto tensor_proto;
tensor_proto.set_name("string_initializer");
tensor_proto.set_data_type(TensorProto_DataType_STRING);
tensor_proto.add_dims(static_cast<int64_t>(payload_sizes.size()));

for (size_t payload_size : payload_sizes) {
tensor_proto.add_string_data(std::string(payload_size, 'a'));
}

return tensor_proto;
}
} // namespace

TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsOversizedPayload) {
constexpr size_t kPayloadBudgetBytes = 128;
constexpr size_t kTestBudgetBytes = sizeof(std::string) + kPayloadBudgetBytes;
const TensorProto tensor_proto = CreateStringTensorProto({kPayloadBudgetBytes + 1});

const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes);

ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string_data shape bytes + payload exceeds");
}

TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsCombinedShapeAndPayloadOverflow) {
constexpr size_t kFirstPayloadBytes = 64;
constexpr size_t kTestBudgetBytes = 2 * sizeof(std::string) + kFirstPayloadBytes;
const TensorProto tensor_proto = CreateStringTensorProto({kFirstPayloadBytes, 1});

// STRING tensors account for the shape-declared bytes and the aggregate payload bytes.
// The first payload exactly fills the remaining allowance after object storage; the second
// payload proves that validation uses cumulative accounting.
const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes);

ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string_data shape bytes + payload exceeds");
}

TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoAcceptsExactPayloadLimit) {
constexpr size_t kPayloadBytes = 17;
constexpr size_t kTestBudgetBytes = sizeof(std::string) + kPayloadBytes;
const TensorProto tensor_proto = CreateStringTensorProto({kPayloadBytes});

ASSERT_STATUS_OK(ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes));
}

TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoAcceptsNormalPayload) {
constexpr size_t kTestBudgetBytes = 3 * sizeof(std::string) + 32;
const TensorProto tensor_proto = CreateStringTensorProto({7, 11, 13});

ASSERT_STATUS_OK(ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes));
}

TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsRawData) {
TensorProto tensor_proto = CreateStringTensorProto({1});
tensor_proto.set_raw_data(std::string(sizeof(std::string), '\0'));

const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape(
tensor_proto, sizeof(std::string) + 1);

ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string tensor can not have raw data");
}

// A bool initializer supplied through raw_data is copied verbatim, so its bytes are not
// restricted to {0, 1}. UnpackTensor must normalize them so downstream consumers (which assume
// canonical bool values) all observe the same result regardless of how they read the byte.
Expand Down
Loading