From 7558ba5e3f96d50c2676f1a693c67cd30c34f955 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Tue, 16 Jun 2026 06:00:08 +0000 Subject: [PATCH] Address feedbacks --- .../core/flatbuffers/flatbuffers_utils.cc | 1 - .../core/flatbuffers/flatbuffers_utils.h | 18 --------- onnxruntime/core/graph/graph.cc | 5 --- onnxruntime/core/graph/model.cc | 1 - .../test/framework/ort_model_only_test.cc | 37 +++++-------------- 5 files changed, 9 insertions(+), 53 deletions(-) diff --git a/onnxruntime/core/flatbuffers/flatbuffers_utils.cc b/onnxruntime/core/flatbuffers/flatbuffers_utils.cc index 876c533073a96..42dff12eaa2db 100644 --- a/onnxruntime/core/flatbuffers/flatbuffers_utils.cc +++ b/onnxruntime/core/flatbuffers/flatbuffers_utils.cc @@ -282,7 +282,6 @@ Status LoadValueInfoOrtFormat(const fbs::ValueInfo& fbs_value_info, Status LoadOpsetImportOrtFormat(const flatbuffers::Vector>* fbs_op_set_ids, std::unordered_map& domain_to_version) { ORT_RETURN_IF(nullptr == fbs_op_set_ids, "Model must have opset imports. Invalid ORT format model."); - ORT_RETURN_IF_ERROR(ValidateRequiredTableOffsets(fbs_op_set_ids, "opset import")); domain_to_version.clear(); domain_to_version.reserve(fbs_op_set_ids->size()); diff --git a/onnxruntime/core/flatbuffers/flatbuffers_utils.h b/onnxruntime/core/flatbuffers/flatbuffers_utils.h index b0fbc34e4fa26..aed0c201a2dd5 100644 --- a/onnxruntime/core/flatbuffers/flatbuffers_utils.h +++ b/onnxruntime/core/flatbuffers/flatbuffers_utils.h @@ -50,24 +50,6 @@ onnxruntime::common::Status LoadOpsetImportOrtFormat( const flatbuffers::Vector>* fbs_op_set_ids, std::unordered_map& domain_to_version); -template -inline onnxruntime::common::Status ValidateRequiredTableOffsets( - const flatbuffers::Vector>* fbs_entries, - const char* entry_description) { - if (fbs_entries == nullptr) { - return onnxruntime::common::Status::OK(); - } - - const auto* raw_offsets = reinterpret_cast(fbs_entries->Data()); - for (flatbuffers::uoffset_t i = 0; i < fbs_entries->size(); ++i) { - const auto entry_offset = - flatbuffers::ReadScalar(raw_offsets + i * sizeof(flatbuffers::uoffset_t)); - ORT_RETURN_IF(entry_offset == 0, "Null ", entry_description, " entry. Invalid ORT format model."); - } - - return onnxruntime::common::Status::OK(); -} - // check if filename ends in .ort bool IsOrtFormatModel(const PathString& filename); diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index f87d530f9fbbf..110878fd8d0b0 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -6659,10 +6659,8 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph // Initializers auto fbs_initializers = fbs_graph.initializers(); - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_initializers, "initializer")); #if !defined(DISABLE_SPARSE_TENSORS) auto fbs_sparse_initializers = fbs_graph.sparse_initializers(); - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_sparse_initializers, "sparse initializer")); flatbuffers::uoffset_t map_size = (fbs_initializers != nullptr ? fbs_initializers->size() : 0U) + (fbs_sparse_initializers != nullptr ? fbs_sparse_initializers->size() : 0U); #else @@ -6732,7 +6730,6 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph // NodeArgs auto fbs_node_args = fbs_graph.node_args(); if (fbs_node_args) { - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_node_args, "node arg")); node_args_.reserve(fbs_node_args->size()); for (const auto* fbs_value_info : *fbs_node_args) { ORT_RETURN_IF(nullptr == fbs_value_info, "NodeArg is missing. Invalid ORT format model."); @@ -6751,9 +6748,7 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph // referenced indices. We compute the required slot count from actual node and edge data rather // than trusting the serialized max_node_index field. auto* fbs_nodes = fbs_graph.nodes(); - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_nodes, "node")); auto* fbs_node_edges = fbs_graph.node_edges(); - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_node_edges, "node edge")); uint32_t max_referenced_node_index = 0; bool has_referenced_node_index = false; diff --git a/onnxruntime/core/graph/model.cc b/onnxruntime/core/graph/model.cc index f4fd754adc596..bfa25a5cb2e9a 100644 --- a/onnxruntime/core/graph/model.cc +++ b/onnxruntime/core/graph/model.cc @@ -998,7 +998,6 @@ common::Status Model::LoadFromOrtFormat(const fbs::Model& fbs_model, // Load the model metadata if (const auto* fbs_metadata_props = fbs_model.metadata_props()) { - ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_metadata_props, "metadata property")); model->model_metadata_.reserve(fbs_metadata_props->size()); for (const auto* prop : *fbs_metadata_props) { ORT_RETURN_IF(nullptr == prop, "Null entry in metadata_props. Invalid ORT format model."); diff --git a/onnxruntime/test/framework/ort_model_only_test.cc b/onnxruntime/test/framework/ort_model_only_test.cc index 8b0a29ae72bae..9f43f8373b871 100644 --- a/onnxruntime/test/framework/ort_model_only_test.cc +++ b/onnxruntime/test/framework/ort_model_only_test.cc @@ -1,8 +1,6 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include - #include "core/flatbuffers/ort_format_version.h" #include "core/flatbuffers/schema/ort.fbs.h" #include "core/framework/data_types.h" @@ -76,9 +74,14 @@ std::vector BuildOrtModelBuffer( return std::vector(builder.GetBufferPointer(), builder.GetBufferPointer() + builder.GetSize()); } -Status LoadOrtBuffer(const std::vector& buffer) { +Status LoadOrtBuffer(const std::vector& buffer, bool use_buffer_for_initializers = false) { SessionOptions so; ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigLoadModelFormat, "ORT")); + if (use_buffer_for_initializers) { + ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigUseORTModelBytesDirectly, "1")); + ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigUseORTModelBytesForInitializers, + "1")); + } InferenceSessionWrapper session_object{so, GetEnvironment()}; return session_object.Load(buffer.data(), static_cast(buffer.size())); @@ -135,40 +138,18 @@ static void RunOrtModel(const OrtModelTestInfo& test_info) { TEST(OrtModelTest, RejectsInitializerRawDataSizeMismatch) { const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) { - std::vector dims{1}; - std::vector raw_data(sizeof(float) * 2, 0); + std::vector dims{32}; + std::vector raw_data(sizeof(float) * 33, 0); std::vector> initializers{ fbs::CreateTensorDirect(builder, "bad_initializer", "", &dims, fbs::TensorDataType::FLOAT, &raw_data)}; return fbs::CreateGraphDirect(builder, &initializers); }); - const auto status = LoadOrtBuffer(buffer); + const auto status = LoadOrtBuffer(buffer, true); ASSERT_FALSE(status.IsOK()); EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr("raw data size mismatch")); } -TEST(OrtModelTest, RejectsNullNodeArgTableEntry) { - auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) { - std::vector> node_args{ - fbs::CreateValueInfoDirect(builder, "X", "", CreateFloatTensorTypeInfo(builder, 1))}; - return fbs::CreateGraphDirect(builder, nullptr, &node_args); - }); - - const auto* fbs_session = fbs::GetInferenceSession(buffer.data()); - ASSERT_NE(fbs_session, nullptr); - const auto* fbs_node_args = fbs_session->model()->graph()->node_args(); - ASSERT_NE(fbs_node_args, nullptr); - - auto* raw_offsets = const_cast(reinterpret_cast(fbs_node_args->Data())); - std::fill_n(raw_offsets, sizeof(flatbuffers::uoffset_t), 0); - - const auto status = LoadOrtBuffer(buffer); - ASSERT_FALSE(status.IsOK()); - EXPECT_THAT(status.ErrorMessage(), - testing::AnyOf(testing::HasSubstr("Null node arg entry"), - testing::HasSubstr("verification failed"))); -} - TEST(OrtModelTest, RejectsDanglingNodeEdge) { const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) { std::vector> node_edges{