2222#include < executorch/runtime/platform/compiler.h>
2323
2424#include < cinttypes>
25- #include < limits>
2625
2726using executorch::runtime::Error;
2827using executorch::runtime::FreeableBuffer;
@@ -55,14 +54,11 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
5554 flatbuffers::Offset<flat_tensor_flatbuffer::NamedData>>* named_data,
5655 const flatbuffers::Vector<
5756 flatbuffers::Offset<flat_tensor_flatbuffer::DataSegment>>* segments,
58- uint64_t segment_data_size ) {
57+ uint64_t segment_end_offset ) {
5958 // Linear search by name.
6059 if (named_data == nullptr ) {
6160 return Error::NotFound;
6261 }
63- if (segments == nullptr ) {
64- return Error::InvalidExternalData;
65- }
6662 for (flatbuffers::uoffset_t i = 0 ; i < named_data->size (); ++i) {
6763 if (key.size () == named_data->Get (i)->key ()->size () &&
6864 std::strncmp (
@@ -71,12 +67,12 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
7167 named_data->Get (i)->key ()->size ()) == 0 ) {
7268 const auto * found = named_data->Get (i);
7369 // Validate the named_data.
74- const flatbuffers:: uoffset_t segment_index = found->segment_index ();
70+ size_t segment_index = found->segment_index ();
7571 ET_CHECK_OR_RETURN_ERROR (
76- segment_index < segments->size (),
72+ segment_index >= 0 && segment_index < segments->size (),
7773 InvalidExternalData,
7874 " Segment index %zu for key %.*s is out of bounds for segment size %d. Malformed PTD file." ,
79- static_cast < size_t >( segment_index) ,
75+ segment_index,
8076 static_cast <int >(key.size ()),
8177 key.data (),
8278 segments->size ());
@@ -87,13 +83,13 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
8783 static_cast <uint64_t >(segments->Get (segment_index)->offset ()),
8884 static_cast <uint64_t >(segments->Get (segment_index)->size ()),
8985 &seg_end) &&
90- seg_end <= segment_data_size ,
86+ seg_end <= segment_end_offset ,
9187 InvalidExternalData,
92- " Segment offset %" PRIu64 " + size %" PRIu64
93- " exceeds segment_data_size %" PRIu64 " ; malformed PTD file." ,
88+ " Invalid segment offset %" PRIu64
89+ " is larger than the segment_base_offset + segment_data_size %" PRIu64
90+ " ; malformed PTD file." ,
9491 segments->Get (segment_index)->offset (),
95- segments->Get (segment_index)->size (),
96- segment_data_size);
92+ segment_end_offset);
9793 return found;
9894 }
9995 }
@@ -115,21 +111,6 @@ Result<uint64_t> get_segment_end_offset(const FlatTensorHeader& header) {
115111 return segment_end_offset;
116112}
117113
118- Result<uint64_t > get_absolute_segment_offset (
119- const FlatTensorHeader& header,
120- uint64_t segment_offset) {
121- uint64_t absolute_offset = 0 ;
122- ET_CHECK_OR_RETURN_ERROR (
123- !c10::add_overflows (
124- header.segment_base_offset , segment_offset, &absolute_offset),
125- InvalidExternalData,
126- " segment_base_offset %" PRIu64 " + segment offset %" PRIu64
127- " overflows uint64_t; malformed PTD file." ,
128- header.segment_base_offset ,
129- segment_offset);
130- return absolute_offset;
131- }
132-
133114Result<const TensorLayout> create_tensor_layout (
134115 const flat_tensor_flatbuffer::TensorLayout* tensor_layout) {
135116 ScalarType scalar_type =
@@ -155,7 +136,7 @@ ET_NODISCARD Result<const TensorLayout> FlatTensorDataMap::get_tensor_layout(
155136 key,
156137 flat_tensor_->named_data (),
157138 flat_tensor_->segments (),
158- header_. segment_data_size );
139+ segment_end_offset. get () );
159140 if (!named_data.ok ()) {
160141 return named_data.error ();
161142 }
@@ -172,7 +153,7 @@ ET_NODISCARD Result<FreeableBuffer> FlatTensorDataMap::get_data(
172153 key,
173154 flat_tensor_->named_data (),
174155 flat_tensor_->segments (),
175- header_. segment_data_size );
156+ segment_end_offset. get () );
176157 if (!named_data.ok ()) {
177158 return named_data.error ();
178159 }
@@ -182,21 +163,9 @@ ET_NODISCARD Result<FreeableBuffer> FlatTensorDataMap::get_data(
182163 flat_tensor_->segments ()->Get (segment_index)->offset ();
183164 uint64_t segment_size = flat_tensor_->segments ()->Get (segment_index)->size ();
184165
185- Result<uint64_t > absolute_offset =
186- get_absolute_segment_offset (header_, segment_offset);
187- if (!absolute_offset.ok ()) {
188- return absolute_offset.error ();
189- }
190- ET_CHECK_OR_RETURN_ERROR (
191- segment_size <= std::numeric_limits<size_t >::max (),
192- NotSupported,
193- " Segment size %" PRIu64 " exceeds the maximum load size %zu" ,
166+ return loader_->load (
167+ /* offset=*/ header_.segment_base_offset + segment_offset,
194168 segment_size,
195- std::numeric_limits<size_t >::max ());
196-
197- return loader_->load_at_offset (
198- absolute_offset.get (),
199- static_cast <size_t >(segment_size),
200169 DataLoader::SegmentInfo (DataLoader::SegmentInfo::Type::Constant));
201170}
202171
@@ -212,21 +181,14 @@ ET_NODISCARD Error FlatTensorDataMap::load_data_into(
212181 key,
213182 flat_tensor_->named_data (),
214183 flat_tensor_->segments (),
215- header_. segment_data_size );
184+ segment_end_offset. get () );
216185 if (!named_data.ok ()) {
217186 return named_data.error ();
218187 }
219188
220189 uint32_t segment_index = named_data.get ()->segment_index ();
221190 uint64_t segment_offset =
222191 flat_tensor_->segments ()->Get (segment_index)->offset ();
223- uint64_t segment_size = flat_tensor_->segments ()->Get (segment_index)->size ();
224-
225- Result<uint64_t > absolute_offset =
226- get_absolute_segment_offset (header_, segment_offset);
227- if (!absolute_offset.ok ()) {
228- return absolute_offset.error ();
229- }
230192
231193 Result<const TensorLayout> tensor_layout =
232194 create_tensor_layout (named_data.get ()->tensor_layout ());
@@ -238,21 +200,18 @@ ET_NODISCARD Error FlatTensorDataMap::load_data_into(
238200 ET_CHECK_OR_RETURN_ERROR (
239201 size <= tensor_layout.get ().nbytes (),
240202 InvalidArgument,
241- " Requested size %zu exceeds tensor size %zu" ,
203+ " Buffer size %zu is smaller than tensor size %zu" ,
242204 size,
243205 tensor_layout.get ().nbytes ());
244- ET_CHECK_OR_RETURN_ERROR (
245- static_cast <uint64_t >(size) <= segment_size,
246- InvalidExternalData,
247- " Requested size %zu exceeds segment size %" PRIu64,
248- size,
249- segment_size);
250206
251207 // Load mutable data.
252208 DataLoader::SegmentInfo info = DataLoader::SegmentInfo (
253209 DataLoader::SegmentInfo::Type::Mutable, 0 , nullptr );
254- return loader_->load_into_at_offset (
255- absolute_offset.get (), size, info, buffer);
210+ return loader_->load_into (
211+ header_.segment_base_offset + segment_offset,
212+ tensor_layout.get ().nbytes (),
213+ info,
214+ buffer);
256215}
257216
258217ET_NODISCARD Result<uint32_t > FlatTensorDataMap::get_num_keys () const {
@@ -274,7 +233,7 @@ ET_NODISCARD Result<const char*> FlatTensorDataMap::get_key(
274233/* static */ Result<FlatTensorDataMap> FlatTensorDataMap::load (
275234 DataLoader* loader) {
276235 // Check header.
277- Result<FreeableBuffer> header = loader->load_at_offset (
236+ Result<FreeableBuffer> header = loader->load (
278237 /* offset=*/ 0 ,
279238 FlatTensorHeader::kNumHeadBytes ,
280239 DataLoader::SegmentInfo (DataLoader::SegmentInfo::Type::Program));
@@ -291,41 +250,19 @@ ET_NODISCARD Result<const char*> FlatTensorDataMap::get_key(
291250 " Failed to parse FlatTensor header with error code %u. File may be corrupt." ,
292251 static_cast <uint32_t >(fh.error ()));
293252
294- Result<uint64_t > expected_size = get_segment_end_offset (fh.get ());
295- if (!expected_size.ok ()) {
296- return expected_size.error ();
297- }
298- Result<uint64_t > actual_size = loader->source_size ();
299- if (!actual_size.ok ()) {
300- return actual_size.error ();
301- }
253+ size_t expected_size = fh->segment_base_offset + fh->segment_data_size ;
254+ size_t actual_size = loader->size ().get ();
302255 ET_CHECK_OR_RETURN_ERROR (
303- expected_size. get () <= actual_size. get () ,
256+ expected_size <= actual_size,
304257 InvalidExternalData,
305- " File size is too small; file may be corrupted or truncated. Expected %" PRIu64
306- " from flat_tensor header, received %" PRIu64 " from data loader" ,
307- expected_size.get (),
308- actual_size.get ());
309-
310- uint64_t flat_tensor_data_size = 0 ;
311- ET_CHECK_OR_RETURN_ERROR (
312- !c10::add_overflows (
313- fh->flatbuffer_offset , fh->flatbuffer_size , &flat_tensor_data_size),
314- InvalidExternalData,
315- " flatbuffer_offset %" PRIu64 " + flatbuffer_size %" PRIu64
316- " overflows uint64_t; malformed PTD file." ,
317- fh->flatbuffer_offset ,
318- fh->flatbuffer_size );
319- ET_CHECK_OR_RETURN_ERROR (
320- flat_tensor_data_size <= std::numeric_limits<size_t >::max (),
321- NotSupported,
322- " FlatTensor metadata size exceeds the addressable buffer size %zu" ,
323- std::numeric_limits<size_t >::max ());
258+ " File size is too small; file may be corrupted or truncated. Expected %zu from flat_tensor header, received %zu from data loader" ,
259+ expected_size,
260+ actual_size);
324261
325262 // Load flatbuffer data as a segment.
326- Result<FreeableBuffer> flat_tensor_data = loader->load_at_offset (
263+ Result<FreeableBuffer> flat_tensor_data = loader->load (
327264 /* offset=*/ 0 ,
328- static_cast < size_t >(flat_tensor_data_size) ,
265+ fh-> flatbuffer_offset + fh-> flatbuffer_size ,
329266 DataLoader::SegmentInfo (DataLoader::SegmentInfo::Type::Program));
330267 if (!flat_tensor_data.ok ()) {
331268 ET_LOG (Error, " Failed to load flat_tensor data." );
0 commit comments