From 06c40b644bcedb87b0b0adb2d4bbe2790678d073 Mon Sep 17 00:00:00 2001 From: Chris Thompson Date: Thu, 8 Oct 2026 15:05:45 -0700 Subject: [PATCH] Allow uint64-backed FreeableBuffers larger than size_t (#23364) Summary: Pull Request resolved: https://github.com/pytorch/executorch/pull/23364 Follow-on to D120705447, which made DataLoader source offsets 64-bit but kept load sizes size_t, capping a named-data resource at 4 GiB on 32-bit targets. Addresses the review feedback there that weights over 4 GiB should not have to be split into multiple segments. The size is now 64-bit along the named-data path. DataLoader::load_at_offset takes a uint64_t size, FlatTensorDataMap::get_data no longer rejects segments over SIZE_MAX, and FreeableBuffer stores a uint64_t size exposed via size_uint64(). size() aborts if the value doesn't fit, as data() already does for uint64-backed buffers. DualAddressDataLoader and TCEBackend are updated to match. On a 32-bit core: - >4 GiB weight via NamedDataMap::get_data: supported when the loader returns a uint64-backed buffer. The CPU only forwards the 64-bit address and size to a device, e.g. the TCE weight blob. CPU-pointer loaders still return NotSupported. Review follow-up: separate LegacyUInt64Data callback storage and preserve implicit conversion of legacy callbacks, including captureless lambdas. The legacy callback alias and constructor remain deprecated. Added move/free and 32-bit bounds regression coverage. Embedded consumers retain C++17 support. Review follow-up implemented with Codex. Reviewed By: rascani Differential Revision: D122398972 --- .../flat_tensor/flat_tensor_data_map.cpp | 18 +-- .../test/flat_tensor_data_map_test.cpp | 43 +++++- runtime/core/data_loader.h | 9 +- runtime/core/freeable_buffer.h | 89 ++++++++++-- runtime/core/test/freeable_buffer_test.cpp | 136 +++++++++++++++++- 5 files changed, 257 insertions(+), 38 deletions(-) diff --git a/extension/flat_tensor/flat_tensor_data_map.cpp b/extension/flat_tensor/flat_tensor_data_map.cpp index 0ab70f5ab58..fa1cbc4e874 100644 --- a/extension/flat_tensor/flat_tensor_data_map.cpp +++ b/extension/flat_tensor/flat_tensor_data_map.cpp @@ -22,7 +22,6 @@ #include #include -#include using executorch::runtime::Error; using executorch::runtime::FreeableBuffer; @@ -187,16 +186,9 @@ ET_NODISCARD Result FlatTensorDataMap::get_data( if (!absolute_offset.ok()) { return absolute_offset.error(); } - ET_CHECK_OR_RETURN_ERROR( - segment_size <= std::numeric_limits::max(), - NotSupported, - "Segment size %" PRIu64 " exceeds the maximum load size %zu", - segment_size, - std::numeric_limits::max()); - return loader_->load_at_offset( absolute_offset.get(), - static_cast(segment_size), + segment_size, DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Constant)); } @@ -316,16 +308,10 @@ ET_NODISCARD Result FlatTensorDataMap::get_key( " overflows uint64_t; malformed PTD file.", fh->flatbuffer_offset, fh->flatbuffer_size); - ET_CHECK_OR_RETURN_ERROR( - flat_tensor_data_size <= std::numeric_limits::max(), - NotSupported, - "FlatTensor metadata size exceeds the addressable buffer size %zu", - std::numeric_limits::max()); - // Load flatbuffer data as a segment. Result flat_tensor_data = loader->load_at_offset( /*offset=*/0, - static_cast(flat_tensor_data_size), + flat_tensor_data_size, DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Program)); if (!flat_tensor_data.ok()) { ET_LOG(Error, "Failed to load flat_tensor data."); diff --git a/extension/flat_tensor/test/flat_tensor_data_map_test.cpp b/extension/flat_tensor/test/flat_tensor_data_map_test.cpp index 3c768ef3020..d628594b8eb 100644 --- a/extension/flat_tensor/test/flat_tensor_data_map_test.cpp +++ b/extension/flat_tensor/test/flat_tensor_data_map_test.cpp @@ -300,20 +300,23 @@ class WideOffsetDataLoader final : public DataLoader { Result load_at_offset( uint64_t offset, - size_t size, + uint64_t size, const SegmentInfo& segment_info) const override { if (segment_info.segment_type == SegmentInfo::Type::Program) { if (offset > metadata_size_ || size > metadata_size_ - offset) { return Error::InvalidArgument; } return FreeableBuffer( - metadata_ + static_cast(offset), size, nullptr); + metadata_ + static_cast(offset), + static_cast(size), + nullptr); } + last_size_ = size; if (size > segment_size_) { return Error::InvalidArgument; } last_offset_ = offset; - return FreeableBuffer(segment_data_, size, nullptr); + return FreeableBuffer(segment_data_, static_cast(size), nullptr); } Error load_into_at_offset( @@ -337,6 +340,10 @@ class WideOffsetDataLoader final : public DataLoader { return last_offset_; } + uint64_t last_size() const { + return last_size_; + } + private: const uint8_t* metadata_; size_t metadata_size_; @@ -344,6 +351,7 @@ class WideOffsetDataLoader final : public DataLoader { size_t segment_size_; uint64_t source_size_; mutable uint64_t last_offset_{0}; + mutable uint64_t last_size_{0}; }; } // namespace @@ -447,6 +455,35 @@ TEST_F(FlatTensorDataMapTest, PreservesWideOffsetWhenLoadingData) { EXPECT_EQ(loaded->data(), &segment_data); } +TEST_F(FlatTensorDataMapTest, PreservesWideSizeWhenLoadingData) { + constexpr uint64_t kSegmentSize = (uint64_t{1} << 32) + 16; + constexpr SegmentSpec kSegment{0, kSegmentSize, kSegmentSize}; + std::vector data = CreateDataWithVersion( + FlatTensorDataMap::kMaxSupportedSchemaVersion, &kSegment); + + alignas(std::max_align_t) std::array aligned_buffer{}; + ASSERT_LE(data.size(), aligned_buffer.size()); + std::memcpy(aligned_buffer.data(), data.data(), data.size()); + + Result header = + FlatTensorHeader::Parse(aligned_buffer.data(), data.size()); + ASSERT_TRUE(header.ok()); + const float segment_data = 3.0f; + WideOffsetDataLoader loader( + aligned_buffer.data(), + data.size(), + reinterpret_cast(&segment_data), + sizeof(segment_data), + header->segment_base_offset + kSegmentSize); + + Result data_map = FlatTensorDataMap::load(&loader); + ASSERT_TRUE(data_map.ok()); + + // The test loader cannot back this size; it only records the request. + EXPECT_EQ(data_map->get_data(kWideTensorKey).error(), Error::InvalidArgument); + EXPECT_EQ(loader.last_size(), kSegmentSize); +} + TEST_F(FlatTensorDataMapTest, PreservesWideOffsetWhenLoadingDataIntoBuffer) { constexpr uint64_t kSegmentOffset = (uint64_t{1} << 32) + 0x5678; constexpr SegmentSpec kSegment{ diff --git a/runtime/core/data_loader.h b/runtime/core/data_loader.h index 1fd10f2e37c..003f069f5b9 100644 --- a/runtime/core/data_loader.h +++ b/runtime/core/data_loader.h @@ -130,10 +130,10 @@ class DataLoader { */ ET_NODISCARD virtual Result size() const = 0; - /** Loads data from a 64-bit source offset. */ + /** Loads data using a 64-bit source offset and size. */ ET_NODISCARD virtual Result load_at_offset( uint64_t offset, - size_t size, + uint64_t size, const SegmentInfo& segment_info) const { #if SIZE_MAX < UINT64_MAX if (offset > std::numeric_limits::max()) { @@ -143,14 +143,15 @@ class DataLoader { return Error::NotSupported; } #endif - if (static_cast(size) > + if (size > static_cast(std::numeric_limits::max()) - offset) { ET_LOG( Error, "load_at_offset() source range cannot be represented by this data loader."); return Error::NotSupported; } - return load(static_cast(offset), size, segment_info); + return load( + static_cast(offset), static_cast(size), segment_info); } /** Loads data from a 64-bit source offset into the provided buffer. */ diff --git a/runtime/core/freeable_buffer.h b/runtime/core/freeable_buffer.h index c743f32116a..87ed6f1ab11 100644 --- a/runtime/core/freeable_buffer.h +++ b/runtime/core/freeable_buffer.h @@ -10,6 +10,8 @@ #include #include +#include +#include #include #include @@ -26,10 +28,16 @@ class FreeableBuffer final { public: // Callback signature for the function that does the freeing. using FreeFn = void (*)(void* context, void* data, size_t size); - using FreeUInt64Fn = + using FreeUInt64SizeFn = + void (*)(void* context, uint64_t data_uint64, uint64_t size); + /// DEPRECATED: Use FreeUInt64SizeFn instead. + using FreeUInt64Fn ET_DEPRECATED = void (*)(void* context, uint64_t data_uint64, size_t size); private: + using LegacyFreeUInt64Fn = + void (*)(void* context, uint64_t data_uint64, size_t size); + // Forward declare types. struct PointerData { const void* data_; @@ -39,7 +47,12 @@ class FreeableBuffer final { struct UInt64Data { // A pointer value cast to uint64_t. uint64_t data_; - FreeUInt64Fn free_fn_; + FreeUInt64SizeFn free_fn_; + }; + + struct LegacyUInt64Data { + uint64_t data_; + LegacyFreeUInt64Fn free_fn_; }; public: @@ -92,13 +105,33 @@ class FreeableBuffer final { */ explicit FreeableBuffer( const uint64_t data_uint64, - size_t size, - FreeUInt64Fn free_fn, + uint64_t size, + FreeUInt64SizeFn free_fn, void* free_fn_context = nullptr) : data_(UInt64Data{data_uint64, free_fn}), free_fn_context_(free_fn_context), size_(size) {} + /// DEPRECATED: Use the FreeUInt64SizeFn ctor instead. + // Callbacks convertible to the new signature, including nullptr, use the + // non-deprecated ctor. Embedded consumers still require C++17. + // NOLINTBEGIN(facebook-modernize-sfinae-enable-if-to-concepts) + template < + typename Callback, + std::enable_if_t< + std::is_convertible_v && + !std::is_convertible_v, + int> = 0> + ET_DEPRECATED explicit FreeableBuffer( + const uint64_t data_uint64, + size_t size, + Callback&& free_fn, + void* free_fn_context = nullptr) + : data_(LegacyUInt64Data{data_uint64, std::forward(free_fn)}), + free_fn_context_(free_fn_context), + size_(size) {} + // NOLINTEND(facebook-modernize-sfinae-enable-if-to-concepts) + /** * Move ctor. Takes the ownership of the data previously owned by `rhs`, * leaving `rhs` pointing to nullptr. @@ -109,8 +142,10 @@ class FreeableBuffer final { size_(rhs.size_) { if (std::holds_alternative(rhs.data_)) { rhs.data_ = PointerData{nullptr, nullptr}; - } else { + } else if (std::holds_alternative(rhs.data_)) { rhs.data_ = UInt64Data{0, nullptr}; + } else { + rhs.data_ = LegacyUInt64Data{0, nullptr}; } rhs.free_fn_context_ = nullptr; rhs.size_ = 0; @@ -130,17 +165,28 @@ class FreeableBuffer final { // Do not need to check for truncation here, as free_fn_ is only set // using the void* ctor. ptr_data.free_fn_( - free_fn_context_, const_cast(ptr_data.data_), size_); + free_fn_context_, + const_cast(ptr_data.data_), + static_cast(size_)); } ptr_data.data_ = nullptr; size_ = 0; - } else { + } else if (std::holds_alternative(data_)) { UInt64Data& int64_data = std::get(data_); if (int64_data.data_ != 0 && int64_data.free_fn_ != nullptr) { int64_data.free_fn_(free_fn_context_, int64_data.data_, size_); } int64_data.data_ = static_cast(0); size_ = 0; + } else { + LegacyUInt64Data& int64_data = std::get(data_); + if (int64_data.data_ != 0 && int64_data.free_fn_ != nullptr) { + // No truncation: the legacy ctor only accepts a size_t size. + int64_data.free_fn_( + free_fn_context_, int64_data.data_, static_cast(size_)); + } + int64_data.data_ = static_cast(0); + size_ = 0; } } @@ -148,6 +194,19 @@ class FreeableBuffer final { * Size of the data in bytes. Returns 0 if the data has been freed. */ size_t size() const { +#if SIZE_MAX < UINT64_MAX + ET_CHECK_MSG( + size_ <= SIZE_MAX, + "FreeableBuffer size exceeds size_t, please use the size_uint64() API."); +#endif + return static_cast(size_); + } + + /** + * Size of the data in bytes as a uint64_t. Returns 0 if the data has been + * freed. Only needed for uint64_t-backed buffers larger than size_t. + */ + uint64_t size_uint64() const { return size_; } @@ -182,10 +241,14 @@ class FreeableBuffer final { */ Result data_uint64_type() const { ET_CHECK_OR_RETURN_ERROR( - std::holds_alternative(data_), + std::holds_alternative(data_) || + std::holds_alternative(data_), InvalidType, "FreeableBuffer is backed by a void*, please use the data() API."); - return std::get(data_).data_; + if (std::holds_alternative(data_)) { + return std::get(data_).data_; + } + return std::get(data_).data_; } private: @@ -194,16 +257,16 @@ class FreeableBuffer final { FreeableBuffer& operator=(FreeableBuffer&& rhs) noexcept = delete; FreeableBuffer& operator=(const FreeableBuffer& rhs) = delete; - // This stores either a PointerData or a UInt64Data structure. Most users - // should use the PointerData variant and the void* ctor. This creates a + // This stores a PointerData, UInt64Data, or LegacyUInt64Data structure. Most + // users should use the PointerData variant and the void* ctor. This creates a // FreeableBuffer backed by void*, accessed using the void* getter data(). // The UInt64Data variant is only helpful in situations where the // FreeableBuffer points to memory on a different core whose pointer value // is larger than the local core's void*. - std::variant data_; + std::variant data_; void* free_fn_context_; - size_t size_; + uint64_t size_; }; } // namespace runtime diff --git a/runtime/core/test/freeable_buffer_test.cpp b/runtime/core/test/freeable_buffer_test.cpp index 2848a6b049d..667b02e76d8 100644 --- a/runtime/core/test/freeable_buffer_test.cpp +++ b/runtime/core/test/freeable_buffer_test.cpp @@ -21,7 +21,7 @@ using executorch::runtime::FreeableBuffer; struct FreeCallArgs { size_t calls; std::variant data; - size_t size; + uint64_t size; }; void RecordFree(void* context, void* data, size_t size) { @@ -31,7 +31,14 @@ void RecordFree(void* context, void* data, size_t size) { call->size = size; } -void RecordInt64Free(void* context, uint64_t data, size_t size) { +void RecordInt64Free(void* context, uint64_t data, uint64_t size) { + auto* call = reinterpret_cast(context); + call->calls++; + call->data = data; + call->size = size; +} + +void RecordLegacyInt64Free(void* context, uint64_t data, size_t size) { auto* call = reinterpret_cast(context); call->calls++; call->data = data; @@ -250,6 +257,127 @@ TEST(FreeableBufferTest, MoveTest) { EXPECT_EQ(call2.size, sizeof(i64)); } +TEST(FreeableBufferTest, UInt64SizeTest) { + constexpr uint64_t kSize = (uint64_t{1} << 32) + 1; + FreeCallArgs call = {}; + FreeableBuffer fb( + /*data_uint64=*/uint64_t{0x900000000}, + /*size=*/kSize, + /*free_fn=*/RecordInt64Free, + /*free_fn_context=*/&call); + + EXPECT_EQ(fb.size_uint64(), kSize); + + FreeableBuffer moved(std::move(fb)); + EXPECT_EQ(fb.size_uint64(), 0); // NOLINT(bugprone-use-after-move) + EXPECT_EQ(fb.data_uint64_type().get(), 0); // NOLINT(bugprone-use-after-move) + fb.Free(); + EXPECT_EQ(call.calls, 0); + EXPECT_EQ(moved.size_uint64(), kSize); + EXPECT_EQ(moved.data_uint64_type().get(), uint64_t{0x900000000}); + + moved.Free(); + EXPECT_EQ(call.calls, 1); + EXPECT_EQ(call.size, kSize); + EXPECT_EQ(moved.size_uint64(), 0); + moved.Free(); + EXPECT_EQ(call.calls, 1); +} + +TEST(FreeableBufferTest, UInt64SizeWithNullFreeFnTest) { + constexpr uint64_t kAddress = 0x900000000; + constexpr uint64_t kSize = (uint64_t{1} << 32) + 1; + FreeableBuffer fb(kAddress, kSize, nullptr); + + EXPECT_EQ(fb.size_uint64(), kSize); + EXPECT_EQ(fb.data_uint64_type().get(), kAddress); + fb.Free(); + EXPECT_EQ(fb.size_uint64(), 0); + EXPECT_EQ(fb.data_uint64_type().get(), 0); +} + +#ifdef __GNUC__ +#pragma GCC diagnostic push +#pragma GCC diagnostic ignored "-Wdeprecated-declarations" +#endif +TEST(FreeableBufferTest, LegacySizeTFreeFnTest) { + executorch::runtime::pal_init(); + const uint64_t i64 = 0x900000000; + FreeCallArgs call = {}; + FreeableBuffer fb( + /*data_uint64=*/i64, + /*size=*/sizeof(i64), + /*free_fn=*/RecordLegacyInt64Free, + /*free_fn_context=*/&call); + + EXPECT_EQ(fb.size(), sizeof(i64)); + EXPECT_EQ(fb.data_uint64_type().get(), i64); + EXPECT_EQ(fb.data_safe().error(), Error::InvalidType); + fb.Free(); + EXPECT_EQ(call.calls, 1); + EXPECT_EQ(std::get(call.data), i64); + EXPECT_EQ(call.size, sizeof(i64)); + EXPECT_EQ(fb.size_uint64(), 0); + EXPECT_EQ(fb.data_uint64_type().get(), 0); + fb.Free(); + EXPECT_EQ(call.calls, 1); +} + +TEST(FreeableBufferTest, LegacyConvertibleFreeFnTest) { + struct Callback { + using FreeFn = FreeableBuffer::FreeUInt64Fn; + Callback() = default; + Callback(const Callback&) = delete; + Callback& operator=(const Callback&) = delete; + operator FreeFn() & { + return RecordLegacyInt64Free; + } + } callback; + FreeCallArgs call = {}; + constexpr uint64_t kAddress = 0x900000000; + FreeableBuffer fb(kAddress, size_t{16}, callback, &call); + + fb.Free(); + EXPECT_EQ(call.calls, 1); + EXPECT_EQ(std::get(call.data), kAddress); + EXPECT_EQ(call.size, 16); +} + +TEST(FreeableBufferTest, LegacyLambdaFreeFnMoveTest) { + executorch::runtime::pal_init(); + constexpr uint64_t kAddress = 0x900000000; + constexpr size_t kSize = 16; + FreeCallArgs call = {}; + { + FreeableBuffer source( + kAddress, + kSize, + [](void* context, uint64_t data, size_t size) { + RecordLegacyInt64Free(context, data, size); + }, + &call); + FreeableBuffer destination(std::move(source)); + + EXPECT_EQ(source.size(), 0); // NOLINT(bugprone-use-after-move) + EXPECT_EQ( + source.data_uint64_type().get(), 0); // NOLINT(bugprone-use-after-move) + source.Free(); + EXPECT_EQ(call.calls, 0); + EXPECT_EQ(destination.size(), kSize); + EXPECT_EQ(destination.data_uint64_type().get(), kAddress); + EXPECT_EQ(destination.data_safe().error(), Error::InvalidType); + destination.Free(); + destination.Free(); + EXPECT_EQ(call.calls, 1); + } + EXPECT_EQ(call.calls, 1); + EXPECT_EQ(std::get(call.data), kAddress); + EXPECT_EQ(call.size, kSize); +} +#ifdef __GNUC__ +#pragma GCC diagnostic pop +#endif + TEST(FreeableBufferTest, APIMisuseDeathTest) { executorch::runtime::pal_init(); int i; @@ -266,4 +394,8 @@ TEST(FreeableBufferTest, APIMisuseDeathTest) { /*free_fn=*/nullptr); EXPECT_EQ(fb2.data_safe().error(), Error::InvalidType); ET_EXPECT_DEATH(fb2.data(), ".*"); +#if SIZE_MAX < UINT64_MAX + FreeableBuffer wide(i64, uint64_t{SIZE_MAX} + 1, nullptr); + ET_EXPECT_DEATH(wide.size(), "FreeableBuffer size exceeds size_t"); +#endif }