diff --git a/extension/named_data_map/merged_data_map.cpp b/extension/named_data_map/merged_data_map.cpp index d76f741fbf4..6778d0beffe 100644 --- a/extension/named_data_map/merged_data_map.cpp +++ b/extension/named_data_map/merged_data_map.cpp @@ -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 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(); @@ -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); @@ -58,9 +59,13 @@ namespace ET_MERGED_DATA_MAP_NAMESPACE { ET_NODISCARD Result 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(key.size()), + key.data()); return Error::NotFound; } return named_data_maps_.at(it->second)->get_tensor_layout(key); @@ -68,9 +73,13 @@ ET_NODISCARD Result MergedDataMap::get_tensor_layout( ET_NODISCARD Result 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(key.size()), + key.data()); return Error::NotFound; } return named_data_maps_.at(it->second)->get_data(key); @@ -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(key.size()), key.data()); return named_data_maps_.at(it->second)->load_data_into(key, buffer, size); } @@ -98,7 +108,7 @@ ET_NODISCARD Result 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())) { diff --git a/extension/named_data_map/merged_data_map.h b/extension/named_data_map/merged_data_map.h index cc291b4d093..69cece3aaf9 100644 --- a/extension/named_data_map/merged_data_map.h +++ b/extension/named_data_map/merged_data_map.h @@ -10,6 +10,7 @@ #include +#include #include #include @@ -92,10 +93,12 @@ class MergedDataMap final ~MergedDataMap() override = default; private: + using KeyToMapIndex = std::unordered_map; + MergedDataMap( std::vector named_data_maps, - std::unordered_map 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)) {} @@ -107,8 +110,8 @@ class MergedDataMap final std::vector named_data_maps_; - // Map from key to index in the named_data_maps_ vector. - std::unordered_map 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 diff --git a/extension/named_data_map/test/merged_data_map_test.cpp b/extension/named_data_map/test/merged_data_map_test.cpp index 581deb74773..b69a488b8dc 100644 --- a/extension/named_data_map/test/merged_data_map_test.cpp +++ b/extension/named_data_map/test/merged_data_map_test.cpp @@ -16,7 +16,9 @@ #include +#include #include +#include #include #include @@ -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]); } } } @@ -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 ndms = {data_maps_["addmul"].get()}; + Result merged_map = + MergedDataMap::load(Span(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(data_size); + EXPECT_EQ( + merged_map->load_data_into(bounded_key, buffer.get(), data_size), + Error::Ok); +}