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
63 changes: 63 additions & 0 deletions cpp/test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1144,6 +1144,68 @@ void test_filtered_search() {
}
}


void test_radius_search() {
constexpr std::size_t dataset_count = 256;
constexpr std::size_t dimensions = 32;
metric_punned_t metric(dimensions, metric_kind_t::l2sq_k);

std::random_device seed_source;
std::mt19937 generator(seed_source());
std::uniform_real_distribution<float> distribution(0.0, 1.0);
using vector_of_vectors_t = std::vector<std::vector<float>>;

vector_of_vectors_t vector_of_vectors(dataset_count);
for (auto& vector : vector_of_vectors) {
vector.resize(dimensions);
std::generate(vector.begin(), vector.end(), [&] { return distribution(generator); });
}

index_dense_t index = index_dense_t::make(metric);
index.reserve(dataset_count);
for (std::size_t idx = 0; idx < dataset_count; ++idx)
index.add(idx, vector_of_vectors[idx].data());
expect_eq(index.size(), dataset_count);

// Search without radius (default) - should return up to `wanted` results
auto result_default = index.search(vector_of_vectors[0].data(), 10);
expect(result_default);
expect(result_default.count > 0);
expect(result_default.count <= 10);

// Search with a tight radius - verify all returned distances are within radius
float tight_radius = 0.5f;
auto result_tight = index.search(vector_of_vectors[0].data(), 10,
index_dense_t::any_thread(), false, tight_radius);
expect(result_tight);
expect(result_tight.count <= result_default.count);

std::vector<index_dense_t::vector_key_t> keys(result_tight.count);
std::vector<float> distances(result_tight.count);
result_tight.dump_to(keys.data(), distances.data());
for (std::size_t i = 0; i < result_tight.count; ++i)
expect(distances[i] <= tight_radius);

// Search with radius = 0.0 for a vector that is in the index.
// L2sq self-distance of identical f32 vectors is exactly 0.0.
auto result_zero = index.search(vector_of_vectors[0].data(), 10,
index_dense_t::any_thread(), false, 0.0f);
expect(result_zero);
expect(result_zero.count >= 1);
keys.resize(result_zero.count);
distances.resize(result_zero.count);
result_zero.dump_to(keys.data(), distances.data());
for (std::size_t i = 0; i < result_zero.count; ++i)
expect(distances[i] <= 0.0f);

// Infinite radius (default) should match no-radius search
auto result_inf = index.search(vector_of_vectors[0].data(), 10,
index_dense_t::any_thread(), false,
std::numeric_limits<float>::infinity());
expect(result_inf);
expect_eq(result_inf.count, result_default.count);
}

void test_isolate() {
constexpr std::size_t dataset_count = 16;
constexpr std::size_t dimensions = 32;
Expand Down Expand Up @@ -1256,6 +1318,7 @@ int main(int, char**) {
test_strings<std::int64_t, slot32_t>();

test_filtered_search();
test_radius_search();
test_isolate();
return 0;
}
16 changes: 16 additions & 0 deletions include/usearch/index.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -1437,6 +1437,10 @@ struct index_search_config_t {

/// @brief Brute-forces exhaustive search over all entries in the index.
bool exact = false;

/// @brief Maximum search radius. Results with distance greater than this are excluded.
/// Defaults to infinity, meaning no filtering is applied.
float radius = std::numeric_limits<float>::infinity();
};

struct index_cluster_config_t {
Expand Down Expand Up @@ -3073,6 +3077,18 @@ class index_gt {
top.sort_ascending();
top.shrink(wanted);

if (std::isfinite(config.radius)) {
candidate_t const* data = top.data();
std::size_t within_radius = top.size();
for (std::size_t i = 0; i < top.size(); ++i) {
if (data[i].distance > static_cast<distance_t>(config.radius)) {
within_radius = i;
break;
}
}
top.shrink(within_radius);
}

// Normalize stats
result.computed_distances = context.computed_distances - result.computed_distances;
result.visited_members = context.iteration_cycles - result.visited_members;
Expand Down
29 changes: 15 additions & 14 deletions include/usearch/index_dense.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -764,19 +764,19 @@ class index_dense_gt {
add_result_t add(vector_key_t key, f32_t const* vector, std::size_t thread = any_thread(), bool copy_vector = true) { return add_(key, vector, thread, copy_vector, casts_.from.f32); }
add_result_t add(vector_key_t key, f64_t const* vector, std::size_t thread = any_thread(), bool copy_vector = true) { return add_(key, vector, thread, copy_vector, casts_.from.f64); }

search_result_t search(b1x8_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.b1x8); }
search_result_t search(i8_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.i8); }
search_result_t search(f16_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.f16); }
search_result_t search(bf16_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.bf16); }
search_result_t search(f32_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.f32); }
search_result_t search(f64_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, casts_.from.f64); }

template <typename predicate_at> search_result_t filtered_search(b1x8_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.b1x8); }
template <typename predicate_at> search_result_t filtered_search(i8_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.i8); }
template <typename predicate_at> search_result_t filtered_search(f16_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.f16); }
template <typename predicate_at> search_result_t filtered_search(bf16_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.bf16); }
template <typename predicate_at> search_result_t filtered_search(f32_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.f32); }
template <typename predicate_at> search_result_t filtered_search(f64_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, casts_.from.f64); }
search_result_t search(b1x8_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.b1x8); }
search_result_t search(i8_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.i8); }
search_result_t search(f16_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.f16); }
search_result_t search(bf16_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.bf16); }
search_result_t search(f32_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.f32); }
search_result_t search(f64_t const* vector, std::size_t wanted, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, dummy_predicate_t {}, thread, exact, radius, casts_.from.f64); }

template <typename predicate_at> search_result_t filtered_search(b1x8_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.b1x8); }
template <typename predicate_at> search_result_t filtered_search(i8_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.i8); }
template <typename predicate_at> search_result_t filtered_search(f16_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.f16); }
template <typename predicate_at> search_result_t filtered_search(bf16_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.bf16); }
template <typename predicate_at> search_result_t filtered_search(f32_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.f32); }
template <typename predicate_at> search_result_t filtered_search(f64_t const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread = any_thread(), bool exact = false, float radius = std::numeric_limits<float>::infinity()) const { return search_(vector, wanted, std::forward<predicate_at>(predicate), thread, exact, radius, casts_.from.f64); }

std::size_t get(vector_key_t key, b1x8_t* vector, std::size_t vectors_count = 1) const { return get_(key, vector, vectors_count, casts_.to.b1x8); }
std::size_t get(vector_key_t key, i8_t* vector, std::size_t vectors_count = 1) const { return get_(key, vector, vectors_count, casts_.to.i8); }
Expand Down Expand Up @@ -2051,7 +2051,7 @@ class index_dense_gt {

template <typename scalar_at, typename predicate_at>
search_result_t search_(scalar_at const* vector, std::size_t wanted, predicate_at&& predicate, std::size_t thread,
bool exact, cast_punned_t const& cast) const {
bool exact, float radius, cast_punned_t const& cast) const {

// Cast the vector, if needed for compatibility with `metric_`
thread_lock_t lock = thread_lock_(thread);
Expand All @@ -2067,6 +2067,7 @@ class index_dense_gt {
search_config.thread = lock.thread_id;
search_config.expansion = config_.expansion_search;
search_config.exact = exact;
search_config.radius = radius;

vector_key_t free_key_copy = free_key_;
if (std::is_same<typename std::decay<predicate_at>::type, dummy_predicate_t>::value) {
Expand Down
104 changes: 104 additions & 0 deletions python/scripts/test_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -395,6 +395,110 @@ def test_index_oversubscribed_search(batch_size: int, threads: int):
assert len(match.keys) == batch_size


@pytest.mark.parametrize("ndim", [3, 97, 25, 1024, 4096])
@pytest.mark.parametrize("metric", [MetricKind.Cos, MetricKind.L2sq])
def test_index_search_radius_filters_results(ndim, metric):
"""Verify that radius constrains returned results by distance."""
reset_randomness()

batch_size = 100
index = Index(ndim=ndim, metric=metric, dtype=ScalarKind.F32, multi=False)
vectors = random_vectors(count=batch_size, ndim=ndim, dtype=np.float32)
keys = np.arange(batch_size)
index.add(keys, vectors, threads=threads)

query = vectors[:1]

# Unconstrained search
matches_all: Matches = index.search(query, 10, threads=threads)
assert len(matches_all) > 0

# Pick a radius that is smaller than the farthest distance
if len(matches_all) >= 2:
mid_distance = float(matches_all.distances[len(matches_all) // 2])
matches_radius: Matches = index.search(query, 10, radius=mid_distance, threads=threads)
assert len(matches_radius) <= len(matches_all)
for i in range(len(matches_radius)):
assert matches_radius.distances[i] <= mid_distance


def test_index_search_radius_default_matches_no_radius():
"""Verify that radius=math.inf returns identical results to no-radius search."""
import math

reset_randomness()

ndim = 32
batch_size = 50
index = Index(ndim=ndim, metric=MetricKind.L2sq, dtype=ScalarKind.F32, multi=False)
vectors = random_vectors(count=batch_size, ndim=ndim, dtype=np.float32)
keys = np.arange(batch_size)
index.add(keys, vectors, threads=threads)

query = vectors[:1]
matches_default: Matches = index.search(query, 10, threads=threads)
matches_inf: Matches = index.search(query, 10, radius=math.inf, threads=threads)

assert len(matches_default) == len(matches_inf)
assert np.array_equal(matches_default.keys, matches_inf.keys)
assert np.allclose(matches_default.distances, matches_inf.distances)


def test_index_search_radius_zero():
"""Verify that radius=0 returns only exact matches (distance = 0.0)."""
reset_randomness()

ndim = 32
batch_size = 50
index = Index(ndim=ndim, metric=MetricKind.L2sq, dtype=ScalarKind.F32, multi=False)
vectors = random_vectors(count=batch_size, ndim=ndim, dtype=np.float32)
keys = np.arange(batch_size)
index.add(keys, vectors, threads=threads)

# Search for a vector that is in the index; L2sq self-distance is exactly 0.0
query = vectors[:1]
matches: Matches = index.search(query, 10, radius=0.0, threads=threads)
assert len(matches) >= 1
for i in range(len(matches)):
assert matches.distances[i] <= 0.0


@pytest.mark.parametrize("batch_size", [1, 7, 32])
def test_index_search_radius_single_and_batch(batch_size):
"""Verify radius works for both single (Matches) and batch (BatchMatches) queries."""
reset_randomness()

ndim = 16
dataset_size = 100
index = Index(ndim=ndim, metric=MetricKind.L2sq, dtype=ScalarKind.F32, multi=False)
vectors = random_vectors(count=dataset_size, ndim=ndim, dtype=np.float32)
keys = np.arange(dataset_size)
index.add(keys, vectors, threads=threads)

queries = vectors[:batch_size]

# First do an unconstrained search to find a reasonable radius
if batch_size == 1:
unconstrained: Matches = index.search(queries, 10, threads=threads)
if len(unconstrained) >= 2:
radius = float(unconstrained.distances[len(unconstrained) // 2])
constrained: Matches = index.search(queries, 10, radius=radius, threads=threads)
assert isinstance(constrained, Matches)
assert len(constrained) <= len(unconstrained)
for i in range(len(constrained)):
assert constrained.distances[i] <= radius
else:
unconstrained: BatchMatches = index.search(queries, 10, threads=threads)
# Use a tight radius to ensure filtering happens
radius = 0.0
constrained: BatchMatches = index.search(queries, 10, radius=radius, threads=threads)
assert isinstance(constrained, BatchMatches)
for i in range(len(constrained)):
match = constrained[i]
for j in range(len(match)):
assert match.distances[j] <= radius


@pytest.mark.parametrize("ndim", [3, 97, 256])
@pytest.mark.parametrize("metric", [MetricKind.Cos, MetricKind.L2sq])
@pytest.mark.parametrize("batch_size", [500, 1024])
Expand Down
Loading