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
18 changes: 2 additions & 16 deletions extension/flat_tensor/flat_tensor_data_map.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@
#include <executorch/runtime/platform/compiler.h>

#include <cinttypes>
#include <limits>

using executorch::runtime::Error;
using executorch::runtime::FreeableBuffer;
Expand Down Expand Up @@ -187,16 +186,9 @@ ET_NODISCARD Result<FreeableBuffer> FlatTensorDataMap::get_data(
if (!absolute_offset.ok()) {
return absolute_offset.error();
}
ET_CHECK_OR_RETURN_ERROR(
segment_size <= std::numeric_limits<size_t>::max(),
NotSupported,
"Segment size %" PRIu64 " exceeds the maximum load size %zu",
segment_size,
std::numeric_limits<size_t>::max());

return loader_->load_at_offset(
absolute_offset.get(),
static_cast<size_t>(segment_size),
segment_size,
DataLoader::SegmentInfo(DataLoader::SegmentInfo::Type::Constant));
}

Expand Down Expand Up @@ -316,16 +308,10 @@ ET_NODISCARD Result<const char*> 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<size_t>::max(),
NotSupported,
"FlatTensor metadata size exceeds the addressable buffer size %zu",
std::numeric_limits<size_t>::max());

// Load flatbuffer data as a segment.
Result<FreeableBuffer> flat_tensor_data = loader->load_at_offset(
/*offset=*/0,
static_cast<size_t>(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.");
Expand Down
43 changes: 40 additions & 3 deletions extension/flat_tensor/test/flat_tensor_data_map_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -300,20 +300,23 @@ class WideOffsetDataLoader final : public DataLoader {

Result<FreeableBuffer> 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<size_t>(offset), size, nullptr);
metadata_ + static_cast<size_t>(offset),
static_cast<size_t>(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_t>(size), nullptr);
}

Error load_into_at_offset(
Expand All @@ -337,13 +340,18 @@ 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_;
const uint8_t* segment_data_;
size_t segment_size_;
uint64_t source_size_;
mutable uint64_t last_offset_{0};
mutable uint64_t last_size_{0};
};

} // namespace
Expand Down Expand Up @@ -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<uint8_t> data = CreateDataWithVersion(
FlatTensorDataMap::kMaxSupportedSchemaVersion, &kSegment);

alignas(std::max_align_t) std::array<uint8_t, 1024> aligned_buffer{};
ASSERT_LE(data.size(), aligned_buffer.size());
std::memcpy(aligned_buffer.data(), data.data(), data.size());

Result<FlatTensorHeader> 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<const uint8_t*>(&segment_data),
sizeof(segment_data),
header->segment_base_offset + kSegmentSize);

Result<FlatTensorDataMap> 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{
Expand Down
9 changes: 5 additions & 4 deletions runtime/core/data_loader.h
Original file line number Diff line number Diff line change
Expand Up @@ -130,10 +130,10 @@ class DataLoader {
*/
ET_NODISCARD virtual Result<size_t> 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<FreeableBuffer> 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<size_t>::max()) {
Expand All @@ -143,14 +143,15 @@ class DataLoader {
return Error::NotSupported;
}
#endif
if (static_cast<uint64_t>(size) >
if (size >
static_cast<uint64_t>(std::numeric_limits<size_t>::max()) - offset) {
ET_LOG(
Error,
"load_at_offset() source range cannot be represented by this data loader.");
return Error::NotSupported;
}
return load(static_cast<size_t>(offset), size, segment_info);
return load(
static_cast<size_t>(offset), static_cast<size_t>(size), segment_info);
}

/** Loads data from a 64-bit source offset into the provided buffer. */
Expand Down
89 changes: 76 additions & 13 deletions runtime/core/freeable_buffer.h
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@

#include <cstddef>
#include <cstdint>
#include <type_traits>
#include <utility>
#include <variant>

#include <executorch/runtime/core/error.h>
Expand All @@ -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_;
Expand All @@ -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:
Expand Down Expand Up @@ -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<Callback, LegacyFreeUInt64Fn> &&
!std::is_convertible_v<Callback, FreeUInt64SizeFn>,
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<Callback>(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.
Expand All @@ -109,8 +142,10 @@ class FreeableBuffer final {
size_(rhs.size_) {
if (std::holds_alternative<PointerData>(rhs.data_)) {
rhs.data_ = PointerData{nullptr, nullptr};
} else {
} else if (std::holds_alternative<UInt64Data>(rhs.data_)) {
rhs.data_ = UInt64Data{0, nullptr};
} else {
rhs.data_ = LegacyUInt64Data{0, nullptr};
}
rhs.free_fn_context_ = nullptr;
rhs.size_ = 0;
Expand All @@ -130,24 +165,48 @@ 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<void*>(ptr_data.data_), size_);
free_fn_context_,
const_cast<void*>(ptr_data.data_),
static_cast<size_t>(size_));
}
ptr_data.data_ = nullptr;
size_ = 0;
} else {
} else if (std::holds_alternative<UInt64Data>(data_)) {
UInt64Data& int64_data = std::get<UInt64Data>(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<uint64_t>(0);
size_ = 0;
} else {
LegacyUInt64Data& int64_data = std::get<LegacyUInt64Data>(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_t>(size_));
}
int64_data.data_ = static_cast<uint64_t>(0);
size_ = 0;
}
}

/**
* 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_t>(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_;
}

Expand Down Expand Up @@ -182,10 +241,14 @@ class FreeableBuffer final {
*/
Result<uint64_t> data_uint64_type() const {
ET_CHECK_OR_RETURN_ERROR(
std::holds_alternative<UInt64Data>(data_),
std::holds_alternative<UInt64Data>(data_) ||
std::holds_alternative<LegacyUInt64Data>(data_),
InvalidType,
"FreeableBuffer is backed by a void*, please use the data() API.");
return std::get<UInt64Data>(data_).data_;
if (std::holds_alternative<UInt64Data>(data_)) {
return std::get<UInt64Data>(data_).data_;
}
return std::get<LegacyUInt64Data>(data_).data_;
}

private:
Expand All @@ -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<PointerData, UInt64Data> data_;
std::variant<PointerData, UInt64Data, LegacyUInt64Data> data_;

void* free_fn_context_;
size_t size_;
uint64_t size_;
};

} // namespace runtime
Expand Down
Loading
Loading