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
17 changes: 4 additions & 13 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,9 @@ function(add_example_or_test TARGET_NAME ...)
endfunction()


option(HNSWLIB_ENABLE_EXCEPTIONS "Whether to enable exceptions in hnswlib" ON)
# Affects example/test executables in this repo only. The INTERFACE library
# does not export -fno-exceptions / /EHsc; consumers set their own flags.
option(HNSWLIB_ENABLE_EXCEPTIONS "Enable exceptions on hnswlib example/test targets" ON)
if(HNSWLIB_ENABLE_EXCEPTIONS)
message("Exceptions are enabled using HNSWLIB_ENABLE_EXCEPTIONS=ON (default)")
else()
Expand Down Expand Up @@ -145,7 +147,7 @@ endif()

if (CMAKE_CXX_COMPILER_ID STREQUAL "MSVC")
set(ENABLE_EXCEPTIONS_FLAGS /EHsc)
set(DISABLE_EXCEPTIONS_FLAGS /GR- /D_HAS_EXCEPTIONS=0)
set(DISABLE_EXCEPTIONS_FLAGS /EHs-c- /GR- /D_HAS_EXCEPTIONS=0)
else()
set(ENABLE_EXCEPTIONS_FLAGS -fexceptions)
set(DISABLE_EXCEPTIONS_FLAGS -fno-exceptions)
Expand Down Expand Up @@ -194,17 +196,6 @@ if(HNSWLIB_EXAMPLES)
add_cxx_flags(-lrt)
endif()
elseif (CMAKE_CXX_COMPILER_ID STREQUAL "MSVC")
if (NOT HNSWLIB_ENABLE_EXCEPTIONS)
# Do not enable exceptions by default. We will enable them on a
# case by case basis when needed.
foreach(config IN ITEMS Debug Release RelWithDebInfo MinSizeRel)
string(TOUPPER ${config} config_upper)
set(FLAGS_VAR "CMAKE_CXX_FLAGS_${config_upper}")
string(REPLACE "/EHsc" "" ${FLAGS_VAR} "${${FLAGS_VAR}}")
set(${FLAGS_VAR} "${${FLAGS_VAR}}" CACHE STRING
"Flags for ${config} configuration." FORCE)
endforeach()
endif()
add_cxx_flags(/O2 /W1 /openmp)
endif()
add_cxx_flags(-DHAVE_CXX0X)
Expand Down
130 changes: 95 additions & 35 deletions hnswlib/bruteforce.h
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,13 @@ class BruteforceSearch : public AlgorithmInterface<dist_t> {
free(data_);
}

// Labels sit at a packed offset; load via memcpy (nmslib/hnswlib#665).
inline labeltype getExternalLabel(size_t internal_id) const {
labeltype return_label;
memcpy(&return_label, data_ + internal_id * size_per_element_ + data_size_, sizeof(labeltype));
return return_label;
}


Status addPointNoExceptions(const void *datapoint, labeltype label, bool replace_deleted = false) override {
int idx;
Expand Down Expand Up @@ -99,7 +106,7 @@ class BruteforceSearch : public AlgorithmInterface<dist_t> {
dict_external_to_internal.erase(found);

size_t cur_c = found->second;
labeltype label = *((labeltype*)(data_ + size_per_element_ * (cur_element_count-1) + data_size_));
labeltype label = getExternalLabel(cur_element_count - 1);
dict_external_to_internal[label] = cur_c;
memcpy(data_ + size_per_element_ * cur_c,
data_ + size_per_element_ * (cur_element_count-1),
Expand All @@ -117,7 +124,7 @@ class BruteforceSearch : public AlgorithmInterface<dist_t> {
for (int i = 0; i < cur_element_count; i++) {
dist_t dist = fstdistfunc_(query_data, data_ + size_per_element_ * i, dist_func_param_);
if (dist <= lastdist || topResults.size() < k) {
labeltype label = *((labeltype *) (data_ + size_per_element_ * i + data_size_));
labeltype label = getExternalLabel(i);
if ((!isIdAllowed) || (*isIdAllowed)(label)) {
topResults.emplace(dist, label);
if (topResults.size() > k)
Expand All @@ -132,24 +139,23 @@ class BruteforceSearch : public AlgorithmInterface<dist_t> {


Status saveIndexNoExceptions(std::ostream &output) {
StreamExceptionsOff guard(output);
return invokeWithoutStreamThrow([&]() -> Status {
if (!output) {
return Status("Cannot save index: output stream is not open or in a failed state");
}
writeBinaryPOD(output, maxelements_);
writeBinaryPOD(output, size_per_element_);
writeBinaryPOD(output, cur_element_count);
if (!output.good()) {
return Status("Failed writing index metadata");
}
// *NoExceptions I/O checks stream state. Callers must leave the default
// iostream exception mask (goodbit); enabling failbit/badbit can throw.
if (!output) {
return Status("Cannot save index: output stream is not open or in a failed state");
}
writeBinaryPOD(output, maxelements_);
writeBinaryPOD(output, size_per_element_);
writeBinaryPOD(output, cur_element_count);
if (!output.good()) {
return Status("Failed writing index metadata");
}

output.write(data_, maxelements_ * size_per_element_);
if (!output.good()) {
return Status("Failed writing vector data");
}
return OkStatus();
});
output.write(data_, maxelements_ * size_per_element_);
if (!output.good()) {
return Status("Failed writing vector data");
}
return OkStatus();
}


Expand All @@ -162,32 +168,86 @@ class BruteforceSearch : public AlgorithmInterface<dist_t> {
}


void loadIndex(std::istream &input, SpaceInterface<dist_t> *s) {
Status loadIndexNoExceptions(std::istream &input, SpaceInterface<dist_t> *s) {
if (!input) {
return Status("Cannot load index: input stream is not open or not readable");
}
size_t file_maxelements = 0;
size_t file_size_per_element = 0;
size_t file_cur_count = 0;
readBinaryPOD(input, file_maxelements);
readBinaryPOD(input, file_size_per_element);
readBinaryPOD(input, file_cur_count);
if (!input) {
HNSWLIB_THROW_RUNTIME_ERROR("Cannot load index: input stream is not open or not readable");
return Status("Cannot load index: failed to read index header");
}
if (file_cur_count > file_maxelements) {
return Status("Cannot load index: cur_element_count exceeds maxelements");
}
readBinaryPOD(input, maxelements_);
readBinaryPOD(input, size_per_element_);
readBinaryPOD(input, cur_element_count);

data_size_ = s->get_data_size();
size_t data_size = s->get_data_size();
size_t size_per_element = data_size + sizeof(labeltype);
if (file_size_per_element != size_per_element) {
return Status("Cannot load index: size_per_element does not match space");
}
char *new_data = (char *) malloc(file_maxelements * size_per_element);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It would be preferable to use RAII to manage this allocation.

if (new_data == nullptr)
return Status("Not enough memory: loadIndex failed to allocate data");
input.read(new_data, file_maxelements * size_per_element);
if (!input) {
free(new_data);
return Status("Cannot load index: failed to read vector data");
}

std::unordered_map<labeltype, size_t> new_dict;
#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Instead of repeating defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 in multiple places, we could define a HNSWLIB_... macro, e.g.

In hnswlib.h:

#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1
#define HNSWLIB_EXCEPTIONS_ENABLED 1
#else
#undef HNSWLIB_EXCEPTIONS_ENABLED
#endif

Then here

#ifdef HNSWLIB_EXCEPTIONS_ENABLED
try {
#endif

etc.

try {
#endif
for (size_t i = 0; i < file_cur_count; i++) {
labeltype lab;
memcpy(&lab, new_data + i * size_per_element + data_size, sizeof(lab));
new_dict[lab] = i;
}
#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1
} catch (const std::bad_alloc&) {
free(new_data);
return Status("Not enough memory: loadIndex failed to rebuild label map");
}
#endif

free(data_);
data_ = new_data;
maxelements_ = file_maxelements;
cur_element_count = file_cur_count;
data_size_ = data_size;
size_per_element_ = size_per_element;
fstdistfunc_ = s->get_dist_func();
dist_func_param_ = s->get_dist_func_param();
size_per_element_ = data_size_ + sizeof(labeltype);
data_ = (char *) malloc(maxelements_ * size_per_element_);
if (data_ == nullptr)
HNSWLIB_THROW_RUNTIME_ERROR("Not enough memory: loadIndex failed to allocate data");

input.read(data_, maxelements_ * size_per_element_);
dict_external_to_internal.swap(new_dict);
return OkStatus();
}


void loadIndex(const std::string &location, SpaceInterface<dist_t> *s) {
Status loadIndexNoExceptions(const std::string &location, SpaceInterface<dist_t> *s) {
std::ifstream input(location, std::ios::binary);
if (!input.is_open()) {
return Status("Cannot load index: input stream is not open or not readable");
}
return loadIndexNoExceptions(input, s);
}

void loadIndex(std::istream &input, SpaceInterface<dist_t> *s) {
Status status = loadIndexNoExceptions(input, s);
if (!status.ok()) {
HNSWLIB_THROW_RUNTIME_ERROR(status.message());
}
}

loadIndex(input, s);

input.close();
void loadIndex(const std::string &location, SpaceInterface<dist_t> *s) {
Status status = loadIndexNoExceptions(location, s);
if (!status.ok()) {
HNSWLIB_THROW_RUNTIME_ERROR(status.message());
}
}
};
} // namespace hnswlib
Loading