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
36 changes: 31 additions & 5 deletions onnxruntime/core/graph/graph.cc
Original file line number Diff line number Diff line change
Expand Up @@ -880,7 +880,22 @@ Status Node::LoadFromOrtFormat(const onnxruntime::fbs::Node& fbs_node,
ORT_RETURN_IF(nullptr == fbs_input_arg_counts, "Node::LoadFromOrtFormat, input_arg_counts is missing");
auto& input_arg_count = definitions_.input_arg_count;
input_arg_count.reserve(fbs_input_arg_counts->size());
input_arg_count.insert(input_arg_count.begin(), fbs_input_arg_counts->cbegin(), fbs_input_arg_counts->cend());
size_t total_arg_count = 0;
for (int32_t count : *fbs_input_arg_counts) {
ORT_RETURN_IF(count < 0,
"Node::LoadFromOrtFormat, input_arg_counts contains a negative value for node ", name_,
". Invalid ORT format model.");
const auto count_size_t = static_cast<size_t>(count);
ORT_RETURN_IF(count_size_t > std::numeric_limits<size_t>::max() - total_arg_count,
"Node::LoadFromOrtFormat, input_arg_counts total overflows size_t for node ", name_,
". Invalid ORT format model.");
total_arg_count += count_size_t;
input_arg_count.push_back(count);
}
ORT_RETURN_IF(total_arg_count != definitions_.input_defs.size(),
"Node::LoadFromOrtFormat, input_arg_counts total (", total_arg_count,
") does not match number of explicit inputs (", definitions_.input_defs.size(),
") for node ", name_, ". Invalid ORT format model.");
}

ORT_RETURN_IF_ERROR(LoadNodeArgsFromOrtFormat(fbs_node.outputs(), definitions_.output_defs));
Expand Down Expand Up @@ -1098,16 +1113,27 @@ int Node::PruneRemovableAttributes(gsl::span<const std::string> removable_attrib
Status Node::UpdateInputArgCount() {
// The node refers to a primitive operator.
// Infer and verify node input arg type information.
int total_arg_count = std::accumulate(definitions_.input_arg_count.cbegin(),
definitions_.input_arg_count.cend(), 0);
size_t total_arg_count = 0;
for (int arg_count : definitions_.input_arg_count) {
ORT_RETURN_IF(arg_count < 0,
"This is an invalid model. Node (", name_, ") has a negative input arg count.");

if (total_arg_count < 0 || static_cast<size_t>(total_arg_count) != definitions_.input_defs.size()) {
const auto arg_count_size_t = static_cast<size_t>(arg_count);
ORT_RETURN_IF(arg_count_size_t > std::numeric_limits<size_t>::max() - total_arg_count,
"This is an invalid model. Node (", name_, ") input arg count total overflows size_t.");
total_arg_count += arg_count_size_t;
}

if (total_arg_count != definitions_.input_defs.size()) {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL,
"This is an invalid model. "
"The sum of input arg count is not equal to size of input defs in node (",
name_, ")");
}

ORT_RETURN_IF(total_arg_count > static_cast<size_t>(std::numeric_limits<int>::max()),
"This is an invalid model. Node (", name_, ") input arg count total exceeds int range.");

// op_ is always valid when this is called
const ONNX_NAMESPACE::OpSchema& op = *Op();

Expand Down Expand Up @@ -1135,7 +1161,7 @@ Status Node::UpdateInputArgCount() {
auto& input_arg_count = definitions_.input_arg_count;
input_arg_count.clear();
size_t m = 0;
auto arg_count_left = total_arg_count;
auto arg_count_left = static_cast<int>(total_arg_count);

for (; m < op.inputs().size() - 1; ++m) {
if (arg_count_left > 0) {
Expand Down
81 changes: 81 additions & 0 deletions onnxruntime/test/framework/ort_model_only_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -285,6 +285,87 @@ TEST(OrtModelTest, RejectsGraphInputWithUnknownNodeArg) {
testing::HasSubstr("Graph references unknown NodeArg 'nonexistent'"));
}

TEST(OrtModelTest, RejectsNegativeInputArgCount) {
const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<flatbuffers::Offset<fbs::ValueInfo>> node_args{
fbs::CreateValueInfoDirect(builder, "x", "", CreateFloatTensorTypeInfo(builder, 1)),
fbs::CreateValueInfoDirect(builder, "y", "", CreateFloatTensorTypeInfo(builder, 1))};
std::vector<flatbuffers::Offset<flatbuffers::String>> inputs{builder.CreateSharedString("x")};
std::vector<flatbuffers::Offset<flatbuffers::String>> outputs{builder.CreateSharedString("y")};
std::vector<flatbuffers::Offset<flatbuffers::String>> implicit_inputs;
std::vector<int32_t> input_arg_counts{-1};
std::vector<flatbuffers::Offset<fbs::Node>> nodes{
fbs::CreateNodeDirect(builder, "n0", "", "", 13, 0, "Identity",
fbs::NodeType::Primitive, nullptr,
&inputs, &outputs, nullptr,
&input_arg_counts, &implicit_inputs)};

return fbs::CreateGraphDirect(builder, nullptr, &node_args, &nodes, 1, nullptr, &inputs, &outputs);
});

const auto status = LoadOrtBuffer(buffer);
ASSERT_FALSE(status.IsOK());
EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr("input_arg_counts contains a negative value"));
}

TEST(OrtModelTest, RejectsMismatchedInputArgCountTotal) {
const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<flatbuffers::Offset<fbs::ValueInfo>> node_args{
fbs::CreateValueInfoDirect(builder, "x", "", CreateFloatTensorTypeInfo(builder, 1)),
fbs::CreateValueInfoDirect(builder, "y", "", CreateFloatTensorTypeInfo(builder, 1))};
std::vector<flatbuffers::Offset<flatbuffers::String>> inputs{builder.CreateSharedString("x")};
std::vector<flatbuffers::Offset<flatbuffers::String>> outputs{builder.CreateSharedString("y")};
std::vector<flatbuffers::Offset<flatbuffers::String>> implicit_inputs;
std::vector<int32_t> input_arg_counts{2};
std::vector<flatbuffers::Offset<fbs::Node>> nodes{
fbs::CreateNodeDirect(builder, "n0", "", "", 13, 0, "Identity",
fbs::NodeType::Primitive, nullptr,
&inputs, &outputs, nullptr,
&input_arg_counts, &implicit_inputs)};

return fbs::CreateGraphDirect(builder, nullptr, &node_args, &nodes, 1, nullptr, &inputs, &outputs);
});

const auto status = LoadOrtBuffer(buffer);
ASSERT_FALSE(status.IsOK());
EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr("input_arg_counts total (2) does not match"));
}

#if !defined(ORT_MINIMAL_BUILD)
TEST(OrtModelTest, NormalizesValidVariadicInputArgCounts) {
Comment thread
danielsongmicrosoft marked this conversation as resolved.
const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<flatbuffers::Offset<fbs::ValueInfo>> node_args{
fbs::CreateValueInfoDirect(builder, "x", "", CreateFloatTensorTypeInfo(builder, 1)),
fbs::CreateValueInfoDirect(builder, "y", "", CreateFloatTensorTypeInfo(builder, 1)),
fbs::CreateValueInfoDirect(builder, "z", "", CreateFloatTensorTypeInfo(builder, 1))};
std::vector<flatbuffers::Offset<flatbuffers::String>> inputs{
builder.CreateSharedString("x"), builder.CreateSharedString("y")};
std::vector<flatbuffers::Offset<flatbuffers::String>> outputs{builder.CreateSharedString("z")};
std::vector<flatbuffers::Offset<flatbuffers::String>> implicit_inputs;
std::vector<int32_t> input_arg_counts{1, 1};
std::vector<flatbuffers::Offset<fbs::Node>> nodes{
fbs::CreateNodeDirect(builder, "n0", "", "", 13, 0, "Sum",
fbs::NodeType::Primitive, nullptr,
&inputs, &outputs, nullptr,
&input_arg_counts, &implicit_inputs)};

return fbs::CreateGraphDirect(builder, nullptr, &node_args, &nodes, 1, nullptr, &inputs, &outputs);
});

SessionOptions so;
ASSERT_STATUS_OK(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigLoadModelFormat, "ORT"));

InferenceSessionWrapper session_object{so, GetEnvironment()};
ASSERT_STATUS_OK(session_object.Load(buffer.data(), static_cast<int>(buffer.size())));

const auto& graph = session_object.GetGraph();
const auto* node = graph.GetNode(0);
ASSERT_NE(node, nullptr);

EXPECT_THAT(node->InputArgCount(), testing::ElementsAre(2));
}
#endif // !defined(ORT_MINIMAL_BUILD)

#if !defined(ORT_MINIMAL_BUILD)
// Keep the CompareTypeProtos in case we need debug the difference
/*
Expand Down