This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new e0960f66d6 [spark] Fix residual filter in distributed vector search
(#8664)
e0960f66d6 is described below
commit e0960f66d6f24d37625ffb4a4a1dd514b9335091
Author: Jingsong Lee <[email protected]>
AuthorDate: Wed Jul 15 23:41:58 2026 +0800
[spark] Fix residual filter in distributed vector search (#8664)
---
.../spark/read/SparkPrimaryKeyVectorRead.java | 15 ++++++-
.../spark/read/SparkVectorSearchBuilderImpl.java | 3 +-
.../spark/sql/PrimaryKeyVectorSearchTest.scala | 48 ++++++++++++++++++++++
3 files changed, 64 insertions(+), 2 deletions(-)
diff --git
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkPrimaryKeyVectorRead.java
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkPrimaryKeyVectorRead.java
index 5fdd42895f..3e4a35f142 100644
---
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkPrimaryKeyVectorRead.java
+++
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkPrimaryKeyVectorRead.java
@@ -19,6 +19,7 @@
package org.apache.paimon.spark.read;
import org.apache.paimon.globalindex.GlobalIndexResult;
+import org.apache.paimon.predicate.Predicate;
import org.apache.paimon.table.FileStoreTable;
import org.apache.paimon.table.source.BucketVectorSearchSplit;
import org.apache.paimon.table.source.PrimaryKeyVectorRead;
@@ -28,6 +29,8 @@ import org.apache.paimon.types.DataField;
import org.apache.paimon.utils.InstantiationUtil;
import org.apache.paimon.utils.SerializableFunction;
+import javax.annotation.Nullable;
+
import java.io.IOException;
import java.util.ArrayList;
import java.util.List;
@@ -46,7 +49,17 @@ public class SparkPrimaryKeyVectorRead extends
PrimaryKeyVectorRead {
float[] query,
int limit,
Map<String, String> searchOptions) {
- super(table, vectorField, query, limit, searchOptions);
+ this(table, vectorField, query, limit, searchOptions, null);
+ }
+
+ public SparkPrimaryKeyVectorRead(
+ FileStoreTable table,
+ DataField vectorField,
+ float[] query,
+ int limit,
+ Map<String, String> searchOptions,
+ @Nullable Predicate filter) {
+ super(table, vectorField, query, limit, searchOptions, filter);
}
@Override
diff --git
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkVectorSearchBuilderImpl.java
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkVectorSearchBuilderImpl.java
index 7638f57269..a1748b6513 100644
---
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkVectorSearchBuilderImpl.java
+++
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/read/SparkVectorSearchBuilderImpl.java
@@ -39,7 +39,8 @@ public class SparkVectorSearchBuilderImpl extends
VectorSearchBuilderImpl {
@Override
public VectorRead newVectorRead() {
if (isPrimaryKeyVectorSearch()) {
- return new SparkPrimaryKeyVectorRead(table, vectorColumn, vector,
limit, options);
+ return new SparkPrimaryKeyVectorRead(
+ table, vectorColumn, vector, limit, options, filter);
}
return new SparkDataEvolutionVectorRead(
table, partitionFilter, filter, limit, vectorColumn, vector,
options);
diff --git
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/PrimaryKeyVectorSearchTest.scala
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/PrimaryKeyVectorSearchTest.scala
index 061fc5238e..b9b71c9045 100644
---
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/PrimaryKeyVectorSearchTest.scala
+++
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/PrimaryKeyVectorSearchTest.scala
@@ -290,6 +290,54 @@ class PrimaryKeyVectorSearchTest extends
PaimonSparkTestBase {
}
}
+ test("distributed primary-key vector search applies residual filter before
top k") {
+ withTable("T") {
+ createVectorTable(
+ columns = "id INT, payload STRING, embedding ARRAY<FLOAT>",
+ bucket = 4,
+ extraOptions = Seq("global-index.thread-num" -> "2"))
+ spark.sql("""
+ |INSERT INTO T VALUES
+ | (1, 'drop', array(1.0f, 0.0f)),
+ | (2, 'keep', array(2.0f, 0.0f)),
+ | (3, 'keep', array(3.0f, 0.0f)),
+ | (4, 'keep', array(4.0f, 0.0f)),
+ | (5, 'keep', array(5.0f, 0.0f)),
+ | (6, 'keep', array(6.0f, 0.0f)),
+ | (7, 'keep', array(7.0f, 0.0f)),
+ | (8, 'keep', array(8.0f, 0.0f)),
+ | (9, 'keep', array(9.0f, 0.0f)),
+ | (10, 'keep', array(10.0f, 0.0f)),
+ | (11, 'keep', array(11.0f, 0.0f)),
+ | (12, 'keep', array(12.0f, 0.0f)),
+ | (13, 'keep', array(13.0f, 0.0f)),
+ | (14, 'keep', array(14.0f, 0.0f)),
+ | (15, 'keep', array(15.0f, 0.0f)),
+ | (16, 'keep', array(16.0f, 0.0f))
+ |""".stripMargin)
+
+ val jobGroup = s"primary-key-vector-residual-filter-${System.nanoTime()}"
+ spark.sparkContext.setJobGroup(jobGroup, jobGroup)
+ try {
+ withSparkSQLConf("spark.paimon.vector-search.distribute.enabled" ->
"true") {
+ val ids = spark
+ .sql("""
+ |SELECT id
+ |FROM vector_search('T', 'embedding', array(0.0f, 0.0f), 2)
+ |WHERE payload = 'keep'
+ |""".stripMargin)
+ .collect()
+ .map(_.getInt(0))
+ .toSet
+ assert(ids == Set(2, 3))
+ }
+ } finally {
+ spark.sparkContext.clearJobGroup()
+ }
+
assert(spark.sparkContext.statusTracker.getJobIdsForGroup(jobGroup).nonEmpty)
+ }
+ }
+
test("deduplicate updates and deletes primary-key vector results") {
withTable("T") {
createVectorTable()