ef = Optional.empty();
@@ -175,6 +201,34 @@ public Builder setK(int k) {
return this;
}
+ /**
+ * Sets the inclusive lower bound for distances returned by the nearest-neighbor search.
+ *
+ * This can be set independently of {@link #setUpperBound(float)}. A query with lower bound
+ * {@code lower} retains results whose distance satisfies {@code distance >= lower}.
+ *
+ * @param lowerBound The inclusive lower distance bound.
+ * @return The Builder instance for method chaining.
+ */
+ public Builder setLowerBound(float lowerBound) {
+ this.lowerBound = Optional.of(lowerBound);
+ return this;
+ }
+
+ /**
+ * Sets the exclusive upper bound for distances returned by the nearest-neighbor search.
+ *
+ *
This can be set independently of {@link #setLowerBound(float)}. A query with upper bound
+ * {@code upper} retains results whose distance satisfies {@code distance < upper}.
+ *
+ * @param upperBound The exclusive upper distance bound.
+ * @return The Builder instance for method chaining.
+ */
+ public Builder setUpperBound(float upperBound) {
+ this.upperBound = Optional.of(upperBound);
+ return this;
+ }
+
/**
* Sets the number of probes to load and search.
*
diff --git a/java/src/test/java/org/lance/JNITest.java b/java/src/test/java/org/lance/JNITest.java
index 94db13d6dea..8b9ad4ce1b1 100644
--- a/java/src/test/java/org/lance/JNITest.java
+++ b/java/src/test/java/org/lance/JNITest.java
@@ -54,21 +54,27 @@ public void testQuery() {
Query defaultQuery =
new Query.Builder().setColumn("column").setKey(new float[] {1.0f, 2.0f, 3.0f}).build();
assertEquals(ApproxMode.NORMAL, defaultQuery.getApproxMode());
-
- JniTestHelper.parseQuery(
- Optional.of(
- new Query.Builder()
- .setColumn("column")
- .setKey(new float[] {1.0f, 2.0f, 3.0f})
- .setK(10)
- .setNprobes(20)
- .setEf(30)
- .setRefineFactor(40)
- .setDistanceType(DistanceType.L2)
- .setUseIndex(true)
- .setQueryParallelism(-1)
- .setApproxMode(ApproxMode.ACCURATE)
- .build()));
+ assertEquals(Optional.empty(), defaultQuery.getLowerBound());
+ assertEquals(Optional.empty(), defaultQuery.getUpperBound());
+
+ Query query =
+ new Query.Builder()
+ .setColumn("column")
+ .setKey(new float[] {1.0f, 2.0f, 3.0f})
+ .setK(10)
+ .setLowerBound(1.5f)
+ .setUpperBound(2.5f)
+ .setNprobes(20)
+ .setEf(30)
+ .setRefineFactor(40)
+ .setDistanceType(DistanceType.L2)
+ .setUseIndex(true)
+ .setQueryParallelism(-1)
+ .setApproxMode(ApproxMode.ACCURATE)
+ .build();
+ assertEquals(Optional.of(1.5f), query.getLowerBound());
+ assertEquals(Optional.of(2.5f), query.getUpperBound());
+ JniTestHelper.parseQuery(Optional.of(query));
}
@Test
diff --git a/java/src/test/java/org/lance/VectorSearchTest.java b/java/src/test/java/org/lance/VectorSearchTest.java
index 8a82ecb2849..49dd1cae6bd 100644
--- a/java/src/test/java/org/lance/VectorSearchTest.java
+++ b/java/src/test/java/org/lance/VectorSearchTest.java
@@ -13,6 +13,10 @@
*/
package org.lance;
+import org.lance.index.DistanceType;
+import org.lance.index.IndexParams;
+import org.lance.index.IndexType;
+import org.lance.index.vector.VectorIndexParams;
import org.lance.ipc.Query;
import org.lance.ipc.ScanOptions;
@@ -156,6 +160,70 @@ void test_knn(boolean createVectorIndex) throws Exception {
}
}
+ @ParameterizedTest
+ @ValueSource(booleans = {false, true})
+ void test_knn_with_distance_range(boolean createVectorIndex) throws Exception {
+ try (TestVectorDataset testVectorDataset =
+ new TestVectorDataset(tempDir.resolve("test_knn_with_distance_range"))) {
+ try (Dataset dataset = testVectorDataset.create()) {
+ if (createVectorIndex) {
+ IndexParams params =
+ IndexParams.builder()
+ .setVectorIndexParams(VectorIndexParams.ivfFlat(2, DistanceType.L2))
+ .build();
+ dataset.createIndex(
+ Arrays.asList(TestVectorDataset.vectorColumnName),
+ IndexType.VECTOR,
+ Optional.of(TestVectorDataset.indexName),
+ params,
+ true);
+ }
+
+ float[] key = new float[32];
+ for (int i = 0; i < 32; i++) {
+ key[i] = i;
+ }
+ ScanOptions options =
+ new ScanOptions.Builder()
+ .nearest(
+ new Query.Builder()
+ .setColumn(TestVectorDataset.vectorColumnName)
+ .setKey(key)
+ .setK(400)
+ .setLowerBound(32768.0f)
+ .setUpperBound(131072.0f)
+ .setNprobes(2)
+ .setUseIndex(createVectorIndex)
+ .build())
+ .build();
+
+ try (Scanner scanner = dataset.newScan(options);
+ ArrowReader reader = scanner.scanBatches()) {
+ VectorSchemaRoot root = reader.getVectorSchemaRoot();
+ assertTrue(reader.loadNextBatch(), "Expected distance-range matches");
+
+ IntVector iVector = (IntVector) root.getVector("i");
+ Set actualI = new HashSet<>();
+ for (int i = 0; i < iVector.getValueCount(); i++) {
+ actualI.add(iVector.get(i));
+ }
+ assertEquals(
+ new HashSet<>(Arrays.asList(1, 81, 161, 241, 321)),
+ actualI,
+ "Distance range should include its lower bound and exclude its upper bound");
+
+ Float4Vector distanceVector = (Float4Vector) root.getVector("_distance");
+ for (int i = 0; i < distanceVector.getValueCount(); i++) {
+ float distance = distanceVector.get(i);
+ assertTrue(distance >= 32768.0f);
+ assertTrue(distance < 131072.0f);
+ }
+ assertFalse(reader.loadNextBatch(), "Expected only one batch");
+ }
+ }
+ }
+ }
+
@Test
void test_knn_with_new_data() throws Exception {
try (TestVectorDataset testVectorDataset =