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
28 changes: 19 additions & 9 deletions extension/named_data_map/merged_data_map.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ namespace ET_MERGED_DATA_MAP_NAMESPACE {
"No non-empty named data maps provided to merge");

// Check for duplicate keys.
std::unordered_map<std::string, uint32_t> key_to_map_index;
KeyToMapIndex key_to_map_index;
for (const uint32_t i : c10::irange(valid_data_maps.size())) {
const auto cur_map = valid_data_maps[i];
uint32_t num_keys = cur_map->get_num_keys().get();
Expand All @@ -47,7 +47,8 @@ namespace ET_MERGED_DATA_MAP_NAMESPACE {
ET_CHECK_OR_RETURN_ERROR(
inserted,
InvalidArgument,
"Duplicate key %s in named data maps at index %u and %" PRIu32,
"Duplicate key %s in named data maps at index %" PRIu32
" and %" PRIu32,
cur_key,
it->second,
i);
Expand All @@ -58,19 +59,27 @@ namespace ET_MERGED_DATA_MAP_NAMESPACE {

ET_NODISCARD Result<const TensorLayout> MergedDataMap::get_tensor_layout(
string_view key) const {
const auto it = key_to_map_index_.find(key.data());
const auto it = key_to_map_index_.find(key);
if (it == key_to_map_index_.end()) {
ET_LOG(Debug, "Key %s not found in named data maps.", key.data());
ET_LOG(
Debug,
"Key %.*s not found in named data maps.",
static_cast<int>(key.size()),
key.data());
return Error::NotFound;
}
return named_data_maps_.at(it->second)->get_tensor_layout(key);
}

ET_NODISCARD
Result<FreeableBuffer> MergedDataMap::get_data(string_view key) const {
const auto it = key_to_map_index_.find(key.data());
const auto it = key_to_map_index_.find(key);
if (it == key_to_map_index_.end()) {
ET_LOG(Debug, "Key %s not found in named data maps.", key.data());
ET_LOG(
Debug,
"Key %.*s not found in named data maps.",
static_cast<int>(key.size()),
key.data());
return Error::NotFound;
}
return named_data_maps_.at(it->second)->get_data(key);
Expand All @@ -80,11 +89,12 @@ ET_NODISCARD Error MergedDataMap::load_data_into(
string_view key,
void* buffer,
size_t size) const {
const auto it = key_to_map_index_.find(key.data());
const auto it = key_to_map_index_.find(key);
ET_CHECK_OR_RETURN_ERROR(
it != key_to_map_index_.end(),
NotFound,
"Key %s not found in named data maps",
"Key %.*s not found in named data maps",
static_cast<int>(key.size()),
key.data());
return named_data_maps_.at(it->second)->load_data_into(key, buffer, size);
}
Expand All @@ -98,7 +108,7 @@ ET_NODISCARD Result<const char*> MergedDataMap::get_key(uint32_t index) const {
ET_CHECK_OR_RETURN_ERROR(
index < total_num_keys,
InvalidArgument,
"Index %u out of range of size %u",
"Index %" PRIu32 " out of range of size %" PRIu32,
index,
total_num_keys);
for (auto i : c10::irange(named_data_maps_.size())) {
Expand Down
9 changes: 6 additions & 3 deletions extension/named_data_map/merged_data_map.h
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

#include <executorch/runtime/core/named_data_map.h>

#include <string_view>
#include <unordered_map>
#include <vector>

Expand Down Expand Up @@ -92,10 +93,12 @@ class MergedDataMap final
~MergedDataMap() override = default;

private:
using KeyToMapIndex = std::unordered_map<std::string_view, uint32_t>;

MergedDataMap(
std::vector<const executorch::ET_RUNTIME_NAMESPACE::NamedDataMap*>
named_data_maps,
std::unordered_map<std::string, uint32_t> key_to_map_index)
KeyToMapIndex key_to_map_index)
: named_data_maps_(std::move(named_data_maps)),
key_to_map_index_(std::move(key_to_map_index)) {}

Expand All @@ -107,8 +110,8 @@ class MergedDataMap final
std::vector<const executorch::ET_RUNTIME_NAMESPACE::NamedDataMap*>
named_data_maps_;

// Map from key to index in the named_data_maps_ vector.
std::unordered_map<std::string, uint32_t> key_to_map_index_;
// Keys alias stable storage owned by the wrapped, longer-lived data maps.
KeyToMapIndex key_to_map_index_;
};

} // namespace ET_MERGED_DATA_MAP_NAMESPACE
Expand Down
27 changes: 25 additions & 2 deletions extension/named_data_map/test/merged_data_map_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,9 @@

#include <gtest/gtest.h>

#include <cstring>
#include <memory>
#include <string>
#include <unordered_map>
#include <vector>

Expand Down Expand Up @@ -120,8 +122,7 @@ void compare_ndm_api_calls(
EXPECT_EQ(merged_load_into, Error::Ok);
for (auto j : c10::irange(ndm_meta.nbytes())) {
EXPECT_EQ(
((uint8_t*)merged_buffer.get())[j],
((uint8_t*)merged_buffer.get())[j]);
((uint8_t*)ndm_buffer.get())[j], ((uint8_t*)merged_buffer.get())[j]);
}
}
}
Expand Down Expand Up @@ -171,3 +172,25 @@ TEST_F(MergedDataMapTest, CheckDataMapContents) {
compare_ndm_api_calls(data_maps_["addmul"].get(), &merged_map.get());
compare_ndm_api_calls(data_maps_["linear"].get(), &merged_map.get());
}

TEST_F(MergedDataMapTest, LookupUsesStringViewLength) {
std::vector<const NamedDataMap*> ndms = {data_maps_["addmul"].get()};
Result<MergedDataMap> merged_map =
MergedDataMap::load(Span<const NamedDataMap*>(ndms.data(), ndms.size()));
ASSERT_EQ(merged_map.error(), Error::Ok);

const char* key = data_maps_["addmul"]->get_key(0).get();
const size_t key_size = std::strlen(key);
const std::string storage = std::string(key) + ".not-part-of-key";
const std::string_view bounded_key(storage.data(), key_size);

EXPECT_EQ(merged_map->get_data(bounded_key).error(), Error::Ok);
EXPECT_EQ(merged_map->get_tensor_layout(bounded_key).error(), Error::Ok);
const auto layout = merged_map->get_tensor_layout(bounded_key);
ASSERT_EQ(layout.error(), Error::Ok);
const size_t data_size = layout->nbytes();
auto buffer = std::make_unique<uint8_t[]>(data_size);
EXPECT_EQ(
merged_map->load_data_into(bounded_key, buffer.get(), data_size),
Error::Ok);
}
Loading