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
7 changes: 7 additions & 0 deletions faiss/IndexShards.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -219,6 +219,13 @@ void IndexShardsTemplate<IndexT>::search(
}
}

FAISS_THROW_IF_NOT_MSG(
!(params && params->sel && !translations.empty() &&
translations.back() != 0),
"IDSelector search is not supported when "
"successive_ids shifts shard IDs; use successive_ids=false "
"with globally assigned shard IDs");

auto fn = [n, k, x, params, &all_distances, &all_labels, &translations](
int no, const IndexT* index) {
if (index->verbose) {
Expand Down
4 changes: 4 additions & 0 deletions faiss/IndexShards.h
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,10 @@ struct IndexShardsTemplate : public ThreadedIndex<IndexT> {
void add_with_ids(idx_t n, const component_t* x, const idx_t* xids)
override;

/// Forward search parameters unchanged to each shard. A non-null
/// params->sel is unsupported if successive_ids gives any shard a
/// nonzero ID offset. To filter globally assigned IDs, assign them to
/// the shards first and use successive_ids=false.
void search(
idx_t n,
const component_t* x,
Expand Down
199 changes: 199 additions & 0 deletions tests/test_threaded_index.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,19 @@
* LICENSE file in the root directory of this source tree.
*/

#include <faiss/IndexBinaryFlat.h>
#include <faiss/IndexFlat.h>
#include <faiss/IndexIVF.h>
#include <faiss/IndexReplicas.h>
#include <faiss/IndexShards.h>
#include <faiss/impl/IDSelector.h>
#include <faiss/impl/ThreadedIndex.h>

#include <gtest/gtest.h>
#include <array>
#include <chrono>
#include <memory>
#include <string>
#include <thread>
#include <vector>

Expand All @@ -21,6 +27,25 @@ struct TestException : public std::exception {};

using idx_t = faiss::idx_t;

struct RecordingFlatIndex : faiss::IndexFlatL2 {
using faiss::IndexFlatL2::IndexFlatL2;

mutable int search_calls = 0;
mutable const faiss::SearchParameters* last_params = nullptr;

void search(
idx_t n,
const float* x,
idx_t k,
float* distances,
idx_t* labels,
const faiss::SearchParameters* params = nullptr) const override {
++search_calls;
last_params = params;
faiss::IndexFlatL2::search(n, x, k, distances, labels, params);
}
};

struct MockIndex : public faiss::Index {
explicit MockIndex(idx_t d_in) : faiss::Index(d_in) {
resetMock();
Expand Down Expand Up @@ -258,3 +283,177 @@ TEST(ThreadedIndex, TestShards) {
}
}
}

TEST(ThreadedIndex, ShardsRejectTranslatedIDSelector) {
constexpr idx_t d = 1;
const std::array<float, 4> xb = {0.0F, 10.0F, 20.0F, 30.0F};
const std::array<float, 1> xq = {20.0F};

for (bool threaded : {false, true}) {
for (idx_t selected_id : {2, 0}) {
RecordingFlatIndex first(d);
RecordingFlatIndex second(d);
first.add(2, xb.data());
second.add(2, xb.data() + 2);

faiss::IndexShards shards(d, threaded, true);
shards.add_shard(&first);
shards.add_shard(&second);

faiss::IDSelectorBatch selector(1, &selected_id);
faiss::SearchParametersIVF params;
params.nprobe = 17;
params.sel = &selector;

std::array<float, 2> distances = {-11.0F, -12.0F};
std::array<idx_t, 2> labels = {-21, -22};
try {
shards.search(
1,
xq.data(),
labels.size(),
distances.data(),
labels.data(),
&params);
FAIL() << "expected translated IDSelector search to fail";
} catch (const faiss::FaissException& exception) {
EXPECT_NE(
std::string(exception.what())
.find("IDSelector search is not supported when "
"successive_ids shifts shard IDs"),
std::string::npos);
}

EXPECT_EQ(params.sel, &selector);
EXPECT_EQ(params.nprobe, 17U);
EXPECT_EQ(first.search_calls, 0);
EXPECT_EQ(second.search_calls, 0);
EXPECT_EQ(distances, (std::array<float, 2>{-11.0F, -12.0F}));
EXPECT_EQ(labels, (std::array<idx_t, 2>{-21, -22}));
}
}
}

TEST(ThreadedIndex, ShardsAllowSelectorWithoutIDTranslation) {
constexpr idx_t d = 1;
const std::array<float, 1> xb = {20.0F};
const std::array<float, 1> xq = {20.0F};
idx_t selected_id = 0;
faiss::IDSelectorBatch selector(1, &selected_id);
faiss::SearchParameters params;
params.sel = &selector;

for (bool threaded : {false, true}) {
faiss::IndexFlatL2 single(d);
single.add(1, xb.data());
faiss::IndexShards one_shard(d, threaded, true);
one_shard.add_shard(&single);

std::array<float, 1> distances{};
std::array<idx_t, 1> labels{};
one_shard.search(
1, xq.data(), 1, distances.data(), labels.data(), &params);
EXPECT_EQ(labels[0], 0);
EXPECT_EQ(distances[0], 0.0F);

faiss::IndexFlatL2 empty(d);
faiss::IndexFlatL2 populated(d);
populated.add(1, xb.data());
faiss::IndexShards zero_offsets(d, threaded, true);
zero_offsets.add_shard(&empty);
zero_offsets.add_shard(&populated);

zero_offsets.search(
1, xq.data(), 1, distances.data(), labels.data(), &params);
EXPECT_EQ(labels[0], 0);
EXPECT_EQ(distances[0], 0.0F);
}
}

TEST(ThreadedIndex, ShardsControlsRemainUnchanged) {
constexpr idx_t d = 1;
const std::array<float, 4> xb = {0.0F, 10.0F, 20.0F, 30.0F};
const std::array<float, 1> xq = {20.0F};

for (bool threaded : {false, true}) {
RecordingFlatIndex first(d);
RecordingFlatIndex second(d);
first.add(2, xb.data());
second.add(2, xb.data() + 2);
faiss::IndexShards translated(d, threaded, true);
translated.add_shard(&first);
translated.add_shard(&second);

std::array<float, 1> distances{};
std::array<idx_t, 1> labels{};
translated.search(1, xq.data(), 1, distances.data(), labels.data());
EXPECT_EQ(labels[0], 2);
EXPECT_EQ(distances[0], 0.0F);

faiss::SearchParametersIVF unfiltered_params;
unfiltered_params.nprobe = 17;
translated.search(
1,
xq.data(),
1,
distances.data(),
labels.data(),
&unfiltered_params);
EXPECT_EQ(labels[0], 2);
EXPECT_EQ(first.last_params, &unfiltered_params);
EXPECT_EQ(second.last_params, &unfiltered_params);
EXPECT_EQ(unfiltered_params.nprobe, 17U);
EXPECT_EQ(unfiltered_params.sel, nullptr);

faiss::IndexFlatL2 local_first(d);
faiss::IndexFlatL2 local_second(d);
local_first.add(2, xb.data());
local_second.add(2, xb.data() + 2);
faiss::IndexShards local_ids(d, threaded, false);
local_ids.add_shard(&local_first);
local_ids.add_shard(&local_second);

idx_t selected_id = 0;
faiss::IDSelectorBatch selector(1, &selected_id);
faiss::SearchParameters params;
params.sel = &selector;
local_ids.search(
1, xq.data(), 1, distances.data(), labels.data(), &params);
EXPECT_EQ(labels[0], 0);
EXPECT_EQ(distances[0], 0.0F);
}
}

TEST(ThreadedIndex, BinaryShardsRejectTranslatedIDSelector) {
const uint8_t xb[] = {0, 1};
idx_t selected_id = 1;
faiss::IDSelectorBatch selector(1, &selected_id);
faiss::SearchParameters params;
params.sel = &selector;

for (bool threaded : {false, true}) {
faiss::IndexBinaryFlat first(8);
faiss::IndexBinaryFlat second(8);
first.add(1, xb);
second.add(1, xb + 1);
faiss::IndexBinaryShards shards(8, threaded, true);
shards.add_shard(&first);
shards.add_shard(&second);

int32_t distance = -11;
idx_t label = -21;
try {
shards.search(1, xb, 1, &distance, &label, &params);
FAIL() << "expected translated IDSelector search to fail";
} catch (const faiss::FaissException& exception) {
EXPECT_NE(
std::string(exception.what())
.find("IDSelector search is not supported when "
"successive_ids shifts shard IDs"),
std::string::npos);
}
EXPECT_EQ(distance, -11);
EXPECT_EQ(label, -21);
EXPECT_EQ(params.sel, &selector);
}
}