Skip to content

Commit fe3edc7

Browse files
authored
Give DataLoader 64-bit source offsets (pytorch#23168)
Summary: DataLoader uses size_t for source offsets, which truncates offsets beyond 4 GiB on 32-bit targets Add 64-bit source-offset APIs and use them for FlatTensor/PTD reads. Buffer sizes remain size_t, so this supports individual resources located beyond 4 GiB, but not a single resource larger than the address space. bypass-github-executorch-ci-checks Reviewed By: rascani, JacobSzwejbka Differential Revision: D120705447 Pull Request resolved: pytorch#23168
1 parent 0f2c4fc commit fe3edc7

4 files changed

Lines changed: 592 additions & 47 deletions

File tree

‎extension/data_loader/test/buffer_data_loader_test.cpp‎

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,8 +8,10 @@
88

99
#include <executorch/extension/data_loader/buffer_data_loader.h>
1010

11+
#include <array>
1112
#include <cstdint>
1213
#include <cstring>
14+
#include <limits>
1315

1416
#include <gtest/gtest.h>
1517

@@ -216,3 +218,62 @@ TEST_F(BufferDataLoaderTest, InBoundsLoadIntoSucceeds) {
216218
EXPECT_EQ(data[0], 1);
217219
}
218220
}
221+
222+
TEST_F(BufferDataLoaderTest, WideInterfaceForwardsToLegacyMethods) {
223+
std::array<uint8_t, 3> data{1, 2, 3};
224+
BufferDataLoader loader(data.data(), data.size());
225+
const DataLoader::SegmentInfo segment_info(
226+
DataLoader::SegmentInfo::Type::Program);
227+
228+
Result<FreeableBuffer> loaded = loader.load_at_offset(1, 2, segment_info);
229+
ASSERT_TRUE(loaded.ok());
230+
const std::array<uint8_t, 2> expected{2, 3};
231+
EXPECT_EQ(std::memcmp(loaded->data(), expected.data(), expected.size()), 0);
232+
233+
std::array<uint8_t, 2> destination{};
234+
EXPECT_EQ(
235+
loader.load_into_at_offset(
236+
1, destination.size(), segment_info, destination.data()),
237+
Error::Ok);
238+
EXPECT_EQ(destination, expected);
239+
240+
Result<uint64_t> source_size = loader.source_size();
241+
ASSERT_TRUE(source_size.ok());
242+
EXPECT_EQ(source_size.get(), sizeof(data));
243+
}
244+
245+
TEST_F(BufferDataLoaderTest, WideInterfaceRejectsUnrepresentableRange) {
246+
std::array<uint8_t, 1> data{};
247+
BufferDataLoader loader(data.data(), data.size());
248+
const DataLoader::SegmentInfo segment_info(
249+
DataLoader::SegmentInfo::Type::Program);
250+
251+
EXPECT_EQ(
252+
loader
253+
.load_at_offset(std::numeric_limits<uint64_t>::max(), 1, segment_info)
254+
.error(),
255+
Error::NotSupported);
256+
EXPECT_EQ(
257+
loader.load_into_at_offset(
258+
std::numeric_limits<uint64_t>::max(), 1, segment_info, data.data()),
259+
Error::NotSupported);
260+
}
261+
262+
#if SIZE_MAX < UINT64_MAX
263+
TEST_F(BufferDataLoaderTest, WideInterfaceRejectsUnrepresentableOffsets) {
264+
std::array<uint8_t, 1> data{};
265+
BufferDataLoader loader(data.data(), data.size());
266+
const DataLoader::SegmentInfo segment_info(
267+
DataLoader::SegmentInfo::Type::Program);
268+
constexpr uint64_t kUnrepresentableOffset =
269+
static_cast<uint64_t>(SIZE_MAX) + 1;
270+
271+
EXPECT_EQ(
272+
loader.load_at_offset(kUnrepresentableOffset, 0, segment_info).error(),
273+
Error::NotSupported);
274+
EXPECT_EQ(
275+
loader.load_into_at_offset(
276+
kUnrepresentableOffset, 0, segment_info, data.data()),
277+
Error::NotSupported);
278+
}
279+
#endif

‎extension/flat_tensor/flat_tensor_data_map.cpp‎

Lines changed: 92 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
#include <executorch/runtime/platform/compiler.h>
2323

2424
#include <cinttypes>
25+
#include <limits>
2526

2627
using executorch::runtime::Error;
2728
using executorch::runtime::FreeableBuffer;
@@ -54,11 +55,14 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
5455
flatbuffers::Offset<flat_tensor_flatbuffer::NamedData>>* named_data,
5556
const flatbuffers::Vector<
5657
flatbuffers::Offset<flat_tensor_flatbuffer::DataSegment>>* segments,
57-
uint64_t segment_end_offset) {
58+
uint64_t segment_data_size) {
5859
// Linear search by name.
5960
if (named_data == nullptr) {
6061
return Error::NotFound;
6162
}
63+
if (segments == nullptr) {
64+
return Error::InvalidExternalData;
65+
}
6266
for (flatbuffers::uoffset_t i = 0; i < named_data->size(); ++i) {
6367
if (key.size() == named_data->Get(i)->key()->size() &&
6468
std::strncmp(
@@ -67,12 +71,12 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
6771
named_data->Get(i)->key()->size()) == 0) {
6872
const auto* found = named_data->Get(i);
6973
// Validate the named_data.
70-
size_t segment_index = found->segment_index();
74+
const flatbuffers::uoffset_t segment_index = found->segment_index();
7175
ET_CHECK_OR_RETURN_ERROR(
72-
segment_index >= 0 && segment_index < segments->size(),
76+
segment_index < segments->size(),
7377
InvalidExternalData,
7478
"Segment index %zu for key %.*s is out of bounds for segment size %d. Malformed PTD file.",
75-
segment_index,
79+
static_cast<size_t>(segment_index),
7680
static_cast<int>(key.size()),
7781
key.data(),
7882
segments->size());
@@ -83,13 +87,13 @@ Result<const flat_tensor_flatbuffer::NamedData*> get_named_data(
8387
static_cast<uint64_t>(segments->Get(segment_index)->offset()),
8488
static_cast<uint64_t>(segments->Get(segment_index)->size()),
8589
&seg_end) &&
86-
seg_end <= segment_end_offset,
90+
seg_end <= segment_data_size,
8791
InvalidExternalData,
88-
"Invalid segment offset %" PRIu64
89-
" is larger than the segment_base_offset + segment_data_size %" PRIu64
90-
"; malformed PTD file.",
92+
"Segment offset %" PRIu64 " + size %" PRIu64
93+
" exceeds segment_data_size %" PRIu64 "; malformed PTD file.",
9194
segments->Get(segment_index)->offset(),
92-
segment_end_offset);
95+
segments->Get(segment_index)->size(),
96+
segment_data_size);
9397
return found;
9498
}
9599
}
@@ -111,6 +115,21 @@ Result<uint64_t> get_segment_end_offset(const FlatTensorHeader& header) {
111115
return segment_end_offset;
112116
}
113117

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+
114133
Result<const TensorLayout> create_tensor_layout(
115134
const flat_tensor_flatbuffer::TensorLayout* tensor_layout) {
116135
ScalarType scalar_type =
@@ -136,7 +155,7 @@ ET_NODISCARD Result<const TensorLayout> FlatTensorDataMap::get_tensor_layout(
136155
key,
137156
flat_tensor_->named_data(),
138157
flat_tensor_->segments(),
139-
segment_end_offset.get());
158+
header_.segment_data_size);
140159
if (!named_data.ok()) {
141160
return named_data.error();
142161
}
@@ -153,7 +172,7 @@ ET_NODISCARD Result<FreeableBuffer> FlatTensorDataMap::get_data(
153172
key,
154173
flat_tensor_->named_data(),
155174
flat_tensor_->segments(),
156-
segment_end_offset.get());
175+
header_.segment_data_size);
157176
if (!named_data.ok()) {
158177
return named_data.error();
159178
}
@@ -163,9 +182,21 @@ ET_NODISCARD Result<FreeableBuffer> FlatTensorDataMap::get_data(
163182
flat_tensor_->segments()->Get(segment_index)->offset();
164183
uint64_t segment_size = flat_tensor_->segments()->Get(segment_index)->size();
165184

166-
return loader_->load(
167-
/*offset=*/header_.segment_base_offset + segment_offset,
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",
168194
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),
169200
DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Constant));
170201
}
171202

@@ -181,14 +212,21 @@ ET_NODISCARD Error FlatTensorDataMap::load_data_into(
181212
key,
182213
flat_tensor_->named_data(),
183214
flat_tensor_->segments(),
184-
segment_end_offset.get());
215+
header_.segment_data_size);
185216
if (!named_data.ok()) {
186217
return named_data.error();
187218
}
188219

189220
uint32_t segment_index = named_data.get()->segment_index();
190221
uint64_t segment_offset =
191222
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+
}
192230

193231
Result<const TensorLayout> tensor_layout =
194232
create_tensor_layout(named_data.get()->tensor_layout());
@@ -200,18 +238,21 @@ ET_NODISCARD Error FlatTensorDataMap::load_data_into(
200238
ET_CHECK_OR_RETURN_ERROR(
201239
size <= tensor_layout.get().nbytes(),
202240
InvalidArgument,
203-
"Buffer size %zu is smaller than tensor size %zu",
241+
"Requested size %zu exceeds tensor size %zu",
204242
size,
205243
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);
206250

207251
// Load mutable data.
208252
DataLoader::SegmentInfo info = DataLoader::SegmentInfo(
209253
DataLoader::SegmentInfo::Type::Mutable, 0, nullptr);
210-
return loader_->load_into(
211-
header_.segment_base_offset + segment_offset,
212-
tensor_layout.get().nbytes(),
213-
info,
214-
buffer);
254+
return loader_->load_into_at_offset(
255+
absolute_offset.get(), size, info, buffer);
215256
}
216257

217258
ET_NODISCARD Result<uint32_t> FlatTensorDataMap::get_num_keys() const {
@@ -233,7 +274,7 @@ ET_NODISCARD Result<const char*> FlatTensorDataMap::get_key(
233274
/* static */ Result<FlatTensorDataMap> FlatTensorDataMap::load(
234275
DataLoader* loader) {
235276
// Check header.
236-
Result<FreeableBuffer> header = loader->load(
277+
Result<FreeableBuffer> header = loader->load_at_offset(
237278
/*offset=*/0,
238279
FlatTensorHeader::kNumHeadBytes,
239280
DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Program));
@@ -250,19 +291,41 @@ ET_NODISCARD Result<const char*> FlatTensorDataMap::get_key(
250291
"Failed to parse FlatTensor header with error code %u. File may be corrupt.",
251292
static_cast<uint32_t>(fh.error()));
252293

253-
size_t expected_size = fh->segment_base_offset + fh->segment_data_size;
254-
size_t actual_size = loader->size().get();
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+
}
255302
ET_CHECK_OR_RETURN_ERROR(
256-
expected_size <= actual_size,
303+
expected_size.get() <= actual_size.get(),
257304
InvalidExternalData,
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);
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());
261324

262325
// Load flatbuffer data as a segment.
263-
Result<FreeableBuffer> flat_tensor_data = loader->load(
326+
Result<FreeableBuffer> flat_tensor_data = loader->load_at_offset(
264327
/*offset=*/0,
265-
fh->flatbuffer_offset + fh->flatbuffer_size,
328+
static_cast<size_t>(flat_tensor_data_size),
266329
DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Program));
267330
if (!flat_tensor_data.ok()) {
268331
ET_LOG(Error, "Failed to load flat_tensor data.");

0 commit comments

Comments
 (0)