Skip to content

Commit 6f73663

Browse files
[core] Support batch vector search
1 parent cf9cdf9 commit 6f73663

10 files changed

Lines changed: 528 additions & 34 deletions

File tree

paimon-common/src/main/java/org/apache/paimon/globalindex/GlobalIndexReader.java

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
import org.apache.paimon.predicate.VectorSearch;
2525

2626
import java.io.Closeable;
27+
import java.util.ArrayList;
2728
import java.util.List;
2829
import java.util.Optional;
2930
import java.util.concurrent.CompletableFuture;
@@ -59,4 +60,23 @@ default CompletableFuture<Optional<ScoredGlobalIndexResult>> visitFullTextSearch
5960
FullTextSearch fullTextSearch) {
6061
throw new UnsupportedOperationException();
6162
}
63+
64+
default CompletableFuture<List<Optional<ScoredGlobalIndexResult>>> visitBatchVectorSearch(
65+
VectorSearch vectorSearch) {
66+
List<CompletableFuture<Optional<ScoredGlobalIndexResult>>> futures = new ArrayList<>();
67+
for (int i = 0; i < vectorSearch.vectorCount(); i++) {
68+
futures.add(visitVectorSearch(vectorSearch.forIndex(i)));
69+
}
70+
return CompletableFuture.allOf(futures.toArray(new CompletableFuture[0]))
71+
.thenApply(
72+
ignored -> {
73+
List<Optional<ScoredGlobalIndexResult>> results =
74+
new ArrayList<>(futures.size());
75+
for (CompletableFuture<Optional<ScoredGlobalIndexResult>> future :
76+
futures) {
77+
results.add(future.join());
78+
}
79+
return results;
80+
});
81+
}
6282
}

paimon-common/src/main/java/org/apache/paimon/globalindex/OffsetGlobalIndexReader.java

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
import org.apache.paimon.predicate.VectorSearch;
2424

2525
import java.io.IOException;
26+
import java.util.ArrayList;
2627
import java.util.List;
2728
import java.util.Optional;
2829
import java.util.concurrent.CompletableFuture;
@@ -145,6 +146,21 @@ public CompletableFuture<Optional<ScoredGlobalIndexResult>> visitFullTextSearch(
145146
.thenApply(opt -> opt.map(r -> r.offset(offset)));
146147
}
147148

149+
@Override
150+
public CompletableFuture<List<Optional<ScoredGlobalIndexResult>>> visitBatchVectorSearch(
151+
VectorSearch vectorSearch) {
152+
return wrapped.visitBatchVectorSearch(vectorSearch.offsetRange(this.offset, this.to))
153+
.thenApply(
154+
results -> {
155+
List<Optional<ScoredGlobalIndexResult>> offsetResults =
156+
new ArrayList<>(results.size());
157+
for (Optional<ScoredGlobalIndexResult> result : results) {
158+
offsetResults.add(result.map(r -> r.offset(offset)));
159+
}
160+
return offsetResults;
161+
});
162+
}
163+
148164
private Optional<GlobalIndexResult> applyOffset(Optional<GlobalIndexResult> result) {
149165
return result.map(r -> r.offset(offset));
150166
}

paimon-common/src/main/java/org/apache/paimon/predicate/VectorSearch.java

Lines changed: 33 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -30,29 +30,54 @@ public class VectorSearch implements Serializable {
3030

3131
private static final long serialVersionUID = 1L;
3232

33-
private final float[] vector;
33+
private final float[][] vectors;
3434
private final String fieldName;
3535
private final int limit;
3636

3737
@Nullable private RoaringNavigableMap64 includeRowIds;
3838

3939
public VectorSearch(float[] vector, int limit, String fieldName) {
40-
if (vector == null) {
41-
throw new IllegalArgumentException("Search cannot be null");
40+
this(new float[][] {vector}, limit, fieldName);
41+
}
42+
43+
public VectorSearch(float[][] vectors, int limit, String fieldName) {
44+
if (vectors == null || vectors.length == 0) {
45+
throw new IllegalArgumentException("Search vectors cannot be null or empty");
46+
}
47+
for (float[] v : vectors) {
48+
if (v == null) {
49+
throw new IllegalArgumentException("Search vector element cannot be null");
50+
}
4251
}
4352
if (limit <= 0) {
4453
throw new IllegalArgumentException("Limit must be positive, got: " + limit);
4554
}
4655
if (fieldName == null || fieldName.isEmpty()) {
4756
throw new IllegalArgumentException("Field name cannot be null or empty");
4857
}
49-
this.vector = vector;
58+
this.vectors = vectors;
5059
this.limit = limit;
5160
this.fieldName = fieldName;
5261
}
5362

5463
public float[] vector() {
55-
return vector;
64+
return vectors[0];
65+
}
66+
67+
public float[][] vectors() {
68+
return vectors;
69+
}
70+
71+
public int vectorCount() {
72+
return vectors.length;
73+
}
74+
75+
public VectorSearch forIndex(int i) {
76+
VectorSearch single = new VectorSearch(vectors[i], limit, fieldName);
77+
if (includeRowIds != null) {
78+
single.withIncludeRowIds(includeRowIds);
79+
}
80+
return single;
5681
}
5782

5883
public int limit() {
@@ -81,7 +106,7 @@ public VectorSearch offsetRange(long from, long to) {
81106
for (long rowId : and64) {
82107
roaringNavigableMap64Offset.add(rowId - from);
83108
}
84-
VectorSearch target = new VectorSearch(vector, limit, fieldName);
109+
VectorSearch target = new VectorSearch(vectors, limit, fieldName);
85110
target.withIncludeRowIds(roaringNavigableMap64Offset);
86111
return target;
87112
}
@@ -90,6 +115,7 @@ public VectorSearch offsetRange(long from, long to) {
90115

91116
@Override
92117
public String toString() {
93-
return String.format("FieldName(%s), Limit(%s)", fieldName, limit);
118+
return String.format(
119+
"FieldName(%s), Limit(%s), VectorCount(%s)", fieldName, limit, vectors.length);
94120
}
95121
}

paimon-core/src/main/java/org/apache/paimon/table/source/VectorRead.java

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020

2121
import org.apache.paimon.globalindex.GlobalIndexResult;
2222

23+
import java.util.ArrayList;
2324
import java.util.List;
2425

2526
/** Vector read to read index files. */
@@ -30,4 +31,14 @@ default GlobalIndexResult read(VectorScan.Plan plan) {
3031
}
3132

3233
GlobalIndexResult read(List<VectorSearchSplit> splits);
34+
35+
default List<GlobalIndexResult> readBatch(VectorScan.Plan plan) {
36+
return readBatch(plan.splits());
37+
}
38+
39+
default List<GlobalIndexResult> readBatch(List<VectorSearchSplit> splits) {
40+
List<GlobalIndexResult> results = new ArrayList<>(1);
41+
results.add(read(splits));
42+
return results;
43+
}
3344
}

paimon-core/src/main/java/org/apache/paimon/table/source/VectorReadImpl.java

Lines changed: 55 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -61,25 +61,39 @@ public class VectorReadImpl implements VectorRead {
6161
private final Predicate filter;
6262
private final int limit;
6363
private final DataField vectorColumn;
64-
private final float[] vector;
64+
private final float[][] vectors;
6565

6666
public VectorReadImpl(
6767
FileStoreTable table,
6868
Predicate filter,
6969
int limit,
7070
DataField vectorColumn,
71-
float[] vector) {
71+
float[][] vectors) {
7272
this.table = table;
7373
this.filter = filter;
7474
this.limit = limit;
7575
this.vectorColumn = vectorColumn;
76-
this.vector = vector;
76+
this.vectors = vectors;
7777
}
7878

7979
@Override
8080
public GlobalIndexResult read(List<VectorSearchSplit> splits) {
81+
if (vectors.length > 1) {
82+
throw new IllegalStateException(
83+
"read() supports single vector only; use readBatch() for multiple vectors");
84+
}
85+
return readBatch(splits).get(0);
86+
}
87+
88+
@Override
89+
public List<GlobalIndexResult> readBatch(List<VectorSearchSplit> splits) {
90+
int n = vectors.length;
8191
if (splits.isEmpty()) {
82-
return GlobalIndexResult.createEmpty();
92+
List<GlobalIndexResult> empty = new ArrayList<>(n);
93+
for (int i = 0; i < n; i++) {
94+
empty.add(GlobalIndexResult.createEmpty());
95+
}
96+
return empty;
8397
}
8498

8599
RoaringNavigableMap64 preFilter = preFilter(splits).orElse(null);
@@ -93,11 +107,11 @@ public GlobalIndexResult read(List<VectorSearchSplit> splits) {
93107
int parallelism = table.coreOptions().toConfiguration().get(GLOBAL_INDEX_THREAD_NUM);
94108
ExecutorService executor = GlobalIndexReadThreadPool.getExecutorService(parallelism);
95109

96-
List<CompletableFuture<Optional<ScoredGlobalIndexResult>>> futures =
110+
List<CompletableFuture<List<Optional<ScoredGlobalIndexResult>>>> futures =
97111
new ArrayList<>(splits.size());
98112
for (VectorSearchSplit split : splits) {
99113
futures.add(
100-
eval(
114+
evalBatch(
101115
globalIndexer,
102116
indexPathFactory,
103117
split.rowRangeStart(),
@@ -109,15 +123,25 @@ public GlobalIndexResult read(List<VectorSearchSplit> splits) {
109123

110124
CompletableFuture.allOf(futures.toArray(new CompletableFuture[0])).join();
111125

112-
ScoredGlobalIndexResult result = ScoredGlobalIndexResult.createEmpty();
113-
for (CompletableFuture<Optional<ScoredGlobalIndexResult>> f : futures) {
114-
Optional<ScoredGlobalIndexResult> next = f.join();
115-
if (next.isPresent()) {
116-
result = result.or(next.get());
126+
ScoredGlobalIndexResult[] merged = new ScoredGlobalIndexResult[n];
127+
for (int i = 0; i < n; i++) {
128+
merged[i] = ScoredGlobalIndexResult.createEmpty();
129+
}
130+
131+
for (CompletableFuture<List<Optional<ScoredGlobalIndexResult>>> future : futures) {
132+
List<Optional<ScoredGlobalIndexResult>> splitResults = future.join();
133+
for (int i = 0; i < n; i++) {
134+
if (splitResults.get(i).isPresent()) {
135+
merged[i] = merged[i].or(splitResults.get(i).get());
136+
}
117137
}
118138
}
119139

120-
return result.topK(limit);
140+
List<GlobalIndexResult> results = new ArrayList<>(n);
141+
for (int i = 0; i < n; i++) {
142+
results.add(merged[i].topK(limit));
143+
}
144+
return results;
121145
}
122146

123147
private Optional<RoaringNavigableMap64> preFilter(List<VectorSearchSplit> splits) {
@@ -139,33 +163,40 @@ private Optional<RoaringNavigableMap64> preFilter(List<VectorSearchSplit> splits
139163
}
140164
}
141165

142-
private CompletableFuture<Optional<ScoredGlobalIndexResult>> eval(
166+
private CompletableFuture<List<Optional<ScoredGlobalIndexResult>>> evalBatch(
143167
GlobalIndexer globalIndexer,
144168
IndexPathFactory indexPathFactory,
145169
long rowRangeStart,
146170
long rowRangeEnd,
147171
List<IndexFileMeta> vectorIndexFiles,
148172
@Nullable RoaringNavigableMap64 includeRowIds,
149173
ExecutorService executor) {
150-
List<GlobalIndexIOMeta> indexIOMetaList = new ArrayList<>();
151-
for (IndexFileMeta indexFile : vectorIndexFiles) {
152-
GlobalIndexMeta meta = checkNotNull(indexFile.globalIndexMeta());
153-
indexIOMetaList.add(
154-
new GlobalIndexIOMeta(
155-
indexPathFactory.toPath(indexFile),
156-
indexFile.fileSize(),
157-
meta.indexMeta()));
158-
}
174+
List<GlobalIndexIOMeta> indexIOMetaList =
175+
buildIOMetaList(indexPathFactory, vectorIndexFiles);
159176
@SuppressWarnings("resource")
160177
FileIO fileIO = table.fileIO();
161178
GlobalIndexFileReader indexFileReader = m -> fileIO.newInputStream(m.filePath());
162179
GlobalIndexReader reader =
163180
globalIndexer.createReader(indexFileReader, indexIOMetaList, executor);
164181
VectorSearch vectorSearch =
165-
new VectorSearch(vector, limit, vectorColumn.name())
182+
new VectorSearch(vectors, limit, vectorColumn.name())
166183
.withIncludeRowIds(includeRowIds);
167184
return new OffsetGlobalIndexReader(reader, rowRangeStart, rowRangeEnd)
168-
.visitVectorSearch(vectorSearch)
185+
.visitBatchVectorSearch(vectorSearch)
169186
.whenComplete((r, t) -> IOUtils.closeQuietly(reader));
170187
}
188+
189+
private List<GlobalIndexIOMeta> buildIOMetaList(
190+
IndexPathFactory indexPathFactory, List<IndexFileMeta> vectorIndexFiles) {
191+
List<GlobalIndexIOMeta> indexIOMetaList = new ArrayList<>();
192+
for (IndexFileMeta indexFile : vectorIndexFiles) {
193+
GlobalIndexMeta meta = checkNotNull(indexFile.globalIndexMeta());
194+
indexIOMetaList.add(
195+
new GlobalIndexIOMeta(
196+
indexPathFactory.toPath(indexFile),
197+
indexFile.fileSize(),
198+
meta.indexMeta()));
199+
}
200+
return indexIOMetaList;
201+
}
171202
}

paimon-core/src/main/java/org/apache/paimon/table/source/VectorSearchBuilder.java

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
import org.apache.paimon.predicate.Predicate;
2424

2525
import java.io.Serializable;
26+
import java.util.List;
2627

2728
/** Builder to build vector search. */
2829
public interface VectorSearchBuilder extends Serializable {
@@ -42,6 +43,9 @@ public interface VectorSearchBuilder extends Serializable {
4243
/** The vector to search. */
4344
VectorSearchBuilder withVector(float[] vector);
4445

46+
/** The vectors to batch search. */
47+
VectorSearchBuilder withVectors(float[][] vectors);
48+
4549
/** Create vector scan to scan index files. */
4650
VectorScan newVectorScan();
4751

@@ -52,4 +56,9 @@ public interface VectorSearchBuilder extends Serializable {
5256
default GlobalIndexResult executeLocal() {
5357
return newVectorRead().read(newVectorScan().scan());
5458
}
59+
60+
/** Execute batch vector index search in local. */
61+
default List<GlobalIndexResult> executeBatchLocal() {
62+
return newVectorRead().readBatch(newVectorScan().scan());
63+
}
5564
}

paimon-core/src/main/java/org/apache/paimon/table/source/VectorSearchBuilderImpl.java

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
import org.apache.paimon.types.DataField;
2727

2828
import static org.apache.paimon.partition.PartitionPredicate.splitPartitionPredicate;
29+
import static org.apache.paimon.utils.Preconditions.checkNotNull;
2930

3031
/** Implementation for {@link VectorSearchBuilder}. */
3132
public class VectorSearchBuilderImpl implements VectorSearchBuilder {
@@ -38,7 +39,7 @@ public class VectorSearchBuilderImpl implements VectorSearchBuilder {
3839
private Predicate filter;
3940
private int limit;
4041
private DataField vectorColumn;
41-
private float[] vector;
42+
private float[][] vectors;
4243

4344
public VectorSearchBuilderImpl(InnerTable table) {
4445
this.table = (FileStoreTable) table;
@@ -76,7 +77,13 @@ public VectorSearchBuilder withVectorColumn(String name) {
7677

7778
@Override
7879
public VectorSearchBuilder withVector(float[] vector) {
79-
this.vector = vector;
80+
this.vectors = new float[][] {vector};
81+
return this;
82+
}
83+
84+
@Override
85+
public VectorSearchBuilder withVectors(float[][] vectors) {
86+
this.vectors = vectors;
8087
return this;
8188
}
8289

@@ -87,6 +94,7 @@ public VectorScan newVectorScan() {
8794

8895
@Override
8996
public VectorRead newVectorRead() {
90-
return new VectorReadImpl(table, filter, limit, vectorColumn, vector);
97+
checkNotNull(vectors, "vectors must be set via withVector() or withVectors()");
98+
return new VectorReadImpl(table, filter, limit, vectorColumn, vectors);
9199
}
92100
}

0 commit comments

Comments
 (0)