From 50ab076a05bc1331a1da0c0ffa656791fe79a4a3 Mon Sep 17 00:00:00 2001 From: shaoyijie Date: Wed, 22 Jul 2026 00:55:10 -0700 Subject: [PATCH 1/2] [spark] Repartition limited lateral vector search input --- .../spark_connector_configuration.html | 6 + .../paimon/spark/SparkConnectorOptions.java | 8 + .../RepartitionLateralVectorSearchInput.scala | 91 ++++++ .../PaimonSparkSessionExtensions.scala | 3 +- .../spark/sql/TableValuedFunctionsTest.scala | 261 +++++++++++++++++- 5 files changed, 367 insertions(+), 2 deletions(-) create mode 100644 paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala diff --git a/docs/generated/spark_connector_configuration.html b/docs/generated/spark_connector_configuration.html index 6937d857579f..cd95fd5fd41a 100644 --- a/docs/generated/spark_connector_configuration.html +++ b/docs/generated/spark_connector_configuration.html @@ -86,6 +86,12 @@ Boolean Whether to adjust the target split size based on pruned (projected) columns. If enabled, split size estimation uses only the columns actually being read. + +
vector-search.lateral-join.parallelism
+ 16 + Integer + Parallelism used to repartition a single-partition LIMIT input before executing a lateral vector search. +
write.data-evolution.update-conflict-retry.max-attempts
20 diff --git a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java index 6e0bbf9c5897..2f315b8df0f5 100644 --- a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java +++ b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java @@ -39,6 +39,14 @@ public class SparkConnectorOptions { .withDescription( "If true, map Paimon TIMESTAMP to Spark TIMESTAMP instead of TIMESTAMP_NTZ."); + public static final ConfigOption VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM = + key("vector-search.lateral-join.parallelism") + .intType() + .defaultValue(16) + .withDescription( + "Parallelism used to repartition a single-partition LIMIT input before " + + "executing a lateral vector search."); + public static final ConfigOption MERGE_SCHEMA = key("write.merge-schema") .booleanType() diff --git a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala new file mode 100644 index 000000000000..cf36007128ab --- /dev/null +++ b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala @@ -0,0 +1,91 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.paimon.spark.catalyst.optimizer + +import org.apache.paimon.spark.SparkConnectorOptions +import org.apache.paimon.spark.catalyst.plans.logical.LateralVectorSearch +import org.apache.paimon.spark.util.OptionUtils + +import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, CTERelationRef, GlobalLimit, HintInfo, Join, LogicalPlan, Repartition, RepartitionOperation, ResolvedHint, UnaryNode, WithCTE} +import org.apache.spark.sql.catalyst.rules.Rule + +/** Restores parallelism lost by a global limit before executing a lateral vector search. */ +object RepartitionLateralVectorSearchInput extends Rule[LogicalPlan] { + + override def apply(plan: LogicalPlan): LogicalPlan = { + val cteDefinitions = plan + .collect { case withCTE: WithCTE => withCTE.cteDefs } + .flatten + .map(definition => definition.id -> definition.child) + .toMap + + plan.transformUp { + case lateralVectorSearch: LateralVectorSearch + if hasUnrepartitionedGlobalLimit(lateralVectorSearch.left, cteDefinitions, Set.empty) => + lateralVectorSearch.copy( + left = Repartition(parallelism, shuffle = true, lateralVectorSearch.left)) + } + } + + private def parallelism: Int = { + val value = + OptionUtils + .getOptionString(SparkConnectorOptions.VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM) + .toInt + require( + value > 0, + s"spark.paimon.${SparkConnectorOptions.VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM.key()} " + + s"must be positive, but got $value") + value + } + + private def hasUnrepartitionedGlobalLimit( + plan: LogicalPlan, + cteDefinitions: Map[Long, LogicalPlan], + visitedCTEs: Set[Long]): Boolean = plan match { + case repartition: RepartitionOperation if repartition.shuffle => false + case repartition: RepartitionOperation => + hasUnrepartitionedGlobalLimit(repartition.child, cteDefinitions, visitedCTEs) + case _: GlobalLimit => true + case reference: CTERelationRef if !visitedCTEs.contains(reference.cteId) => + cteDefinitions + .get(reference.cteId) + .exists(hasUnrepartitionedGlobalLimit(_, cteDefinitions, visitedCTEs + reference.cteId)) + case join: Join + if hasBroadcastHint(join.hint.rightHint) || hasResolvedBroadcastHint(join.right) => + hasUnrepartitionedGlobalLimit(join.left, cteDefinitions, visitedCTEs) + case join: Join + if hasBroadcastHint(join.hint.leftHint) || hasResolvedBroadcastHint(join.left) => + hasUnrepartitionedGlobalLimit(join.right, cteDefinitions, visitedCTEs) + case unary: UnaryNode => + hasUnrepartitionedGlobalLimit(unary.child, cteDefinitions, visitedCTEs) + case _ => false + } + + private def hasBroadcastHint(hint: Option[HintInfo]): Boolean = { + hint.flatMap(_.strategy).contains(BROADCAST) + } + + private def hasResolvedBroadcastHint(plan: LogicalPlan): Boolean = { + plan.exists { + case hint: ResolvedHint => hasBroadcastHint(Some(hint.hints)) + case _ => false + } + } +} diff --git a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala index 388889bbeaaa..69b3b5f28b61 100644 --- a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala +++ b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala @@ -19,7 +19,7 @@ package org.apache.paimon.spark.extensions import org.apache.paimon.spark.catalyst.analysis.{PaimonAnalysis, PaimonDeleteTable, PaimonFunctionResolver, PaimonIncompatibleResolutionRules, PaimonMergeInto, PaimonPostHocResolutionRules, PaimonProcedureResolver, PaimonUpdateTable, PaimonViewResolver, ReplacePaimonFunctions, RewriteUpsertTable} -import org.apache.paimon.spark.catalyst.optimizer.{MergePaimonScalarSubqueries, OptimizeMetadataOnlyDeleteFromPaimonTable, PushDownLateralVectorSearchFilter} +import org.apache.paimon.spark.catalyst.optimizer.{MergePaimonScalarSubqueries, OptimizeMetadataOnlyDeleteFromPaimonTable, PushDownLateralVectorSearchFilter, RepartitionLateralVectorSearchInput} import org.apache.paimon.spark.catalyst.plans.logical.PaimonTableValuedFunctions import org.apache.paimon.spark.commands.BucketExpression import org.apache.paimon.spark.execution.{OldCompatibleStrategy, PaimonStrategy} @@ -104,6 +104,7 @@ class PaimonSparkSessionExtensions extends (SparkSessionExtensions => Unit) { // TODO: Enable MAP selected-key pushdown after core reader supports // __PAIMON_MAP_SELECTED_KEYS read type. extensions.injectOptimizerRule(_ => MergePaimonScalarSubqueries) + extensions.injectOptimizerRule(_ => RepartitionLateralVectorSearchInput) extensions.injectOptimizerRule(_ => PushDownLateralVectorSearchFilter) // planner extensions diff --git a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala index a44711d5b911..92b68e3e2eef 100644 --- a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala +++ b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala @@ -25,7 +25,7 @@ import org.apache.paimon.spark.catalyst.plans.logical.{LateralVectorSearch, Paim import org.apache.paimon.utils.DateTimeUtils import org.apache.spark.sql.{DataFrame, Row} -import org.apache.spark.sql.catalyst.plans.logical.Filter +import org.apache.spark.sql.catalyst.plans.logical.{Filter, Repartition} import java.time.LocalDateTime import java.util.Collections @@ -42,6 +42,251 @@ class TableValuedFunctionsTest extends PaimonHiveTestBase { assert(error.getMessage.contains("Limit must be no greater than")) } + test("lateral vector search repartitions global limit input") { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT gid, embs + | FROM vector_search_source + | WHERE dt = '20260629' + | LIMIT 1000 + |) AS q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + val repartitions = lateralVectorSearch.left.collect { + case repartition: Repartition => repartition + } + + assert(repartitions.size == 1, optimizedPlan.toString) + assert(repartitions.head.shuffle, optimizedPlan.toString) + assert(repartitions.head.numPartitions == 16, optimizedPlan.toString) + } + } + + test("lateral vector search uses configured repartition parallelism") { + val parallelismKey = "spark.paimon.vector-search.lateral-join.parallelism" + spark.conf.set(parallelismKey, "4") + try { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT gid, embs + | FROM vector_search_source + | LIMIT 1000 + |) AS q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val repartition = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .flatMap(_.left.collectFirst { case repartition: Repartition => repartition }) + .getOrElse(fail(optimizedPlan.toString)) + + assert(repartition.numPartitions == 4, optimizedPlan.toString) + } + } finally { + spark.conf.unset(parallelismKey) + } + } + + test("lateral vector search repartitions above a limited repartition") { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT /*+ REPARTITION(4) */ gid, embs + | FROM vector_search_source + | LIMIT 1000 + |) AS q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + val repartitions = lateralVectorSearch.left.collect { + case repartition: Repartition => repartition + } + + assert(lateralVectorSearch.left.isInstanceOf[Repartition], optimizedPlan.toString) + assert(repartitions.map(_.numPartitions) == Seq(16, 4), optimizedPlan.toString) + } + } + + test("lateral vector search preserves repartition above limit") { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |WITH q_limit AS ( + | SELECT gid, embs + | FROM vector_search_source + | LIMIT 1000 + |), + |q AS ( + | SELECT /*+ REPARTITION(4) */ gid, embs + | FROM q_limit + |) + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + val repartitions = lateralVectorSearch.left.collect { + case repartition: Repartition => repartition + } + + assert(repartitions.map(_.numPartitions) == Seq(4), optimizedPlan.toString) + } + } + + test("lateral vector search repartitions CTE limited input") { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |WITH q_limit AS ( + | SELECT gid, embs + | FROM vector_search_source + | LIMIT 1000 + |), + |q AS ( + | SELECT gid, embs + | FROM q_limit + |) + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + + assert(lateralVectorSearch.left.isInstanceOf[Repartition], optimizedPlan.toString) + val repartition = lateralVectorSearch.left.asInstanceOf[Repartition] + assert(repartition.shuffle, optimizedPlan.toString) + assert(repartition.numPartitions == 16, optimizedPlan.toString) + } + } + + test("lateral vector search repartitions broadcast join streamed limited input") { + withTable("vector_search_source", "vector_search_dimension") { + createVectorSearchSource() + spark.sql("CREATE TABLE vector_search_dimension (gid BIGINT) USING paimon") + + val optimizedPlan = spark + .sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT /*+ BROADCAST(d) */ s.gid, s.embs + | FROM ( + | SELECT gid, embs + | FROM vector_search_source + | LIMIT 1000 + | ) s + | JOIN vector_search_dimension d + | ON s.gid = d.gid + |) q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + + assert(lateralVectorSearch.left.isInstanceOf[Repartition], optimizedPlan.toString) + val repartition = lateralVectorSearch.left.asInstanceOf[Repartition] + assert(repartition.shuffle, optimizedPlan.toString) + assert(repartition.numPartitions == 16, optimizedPlan.toString) + } + } + + test("lateral vector search repartitions coalesced limited input") { + withTable("vector_search_source") { + createVectorSearchSource() + + val optimizedPlan = spark + .sql(""" + |WITH q_limit AS ( + | SELECT gid, embs + | FROM vector_search_source + | LIMIT 1000 + |) + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT /*+ COALESCE(16) */ gid, embs + | FROM q_limit + |) q, + |LATERAL ( + | SELECT gid + | FROM vector_search('vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + .queryExecution + .optimizedPlan + + val lateralVectorSearch = optimizedPlan + .collectFirst { case lvs: LateralVectorSearch => lvs } + .getOrElse(fail(optimizedPlan.toString)) + + assert(lateralVectorSearch.left.isInstanceOf[Repartition], optimizedPlan.toString) + val repartition = lateralVectorSearch.left.asInstanceOf[Repartition] + assert(repartition.shuffle, optimizedPlan.toString) + assert(repartition.numPartitions == 16, optimizedPlan.toString) + } + } + test("lateral vector search preserves subquery alias qualifiers") { withTable("vector_search_source", "vector_search_result") { spark.sql(""" @@ -542,6 +787,20 @@ class TableValuedFunctionsTest extends PaimonHiveTestBase { .table(tableIdent) } + private def createVectorSearchSource(): Unit = { + spark.sql(""" + |CREATE TABLE vector_search_source (gid BIGINT, embs ARRAY, dt STRING) + |USING paimon + |TBLPROPERTIES ( + | 'vector.file.format' = 'lance', + | 'vector-field' = 'embs', + | 'field.embs.vector-dim' = '3', + | 'row-tracking.enabled' = 'true', + | 'data-evolution.enabled' = 'true') + |PARTITIONED BY (dt) + |""".stripMargin) + } + private def utcMills(timestamp: String) = Timestamp.fromLocalDateTime(LocalDateTime.parse(timestamp)).getMillisecond } From 31bf5d0de18ea5c53bf8c25fd74c611977fa6c59 Mon Sep 17 00:00:00 2001 From: shaoyijie Date: Thu, 23 Jul 2026 05:38:16 -0700 Subject: [PATCH 2/2] [spark] Handle automatically broadcast lateral vector search input --- .../sql/paimon/shims/MinorVersionShim.scala | 4 + .../sql/paimon/shims/MinorVersionShim.scala | 7 ++ .../spark/sql/paimon/shims/Spark4Shim.scala | 9 ++ .../RepartitionLateralVectorSearchInput.scala | 2 +- .../spark/execution/PaimonStrategy.scala | 31 +++++- .../spark/sql/paimon/shims/SparkShim.scala | 3 + .../spark/sql/TableValuedFunctionsTest.scala | 103 +++++++++++++++++- .../sql/paimon/shims/MinorVersionShim.scala | 7 ++ .../spark/sql/paimon/shims/Spark3Shim.scala | 6 + .../spark/sql/paimon/shims/Spark4Shim.scala | 9 ++ 10 files changed, 178 insertions(+), 3 deletions(-) diff --git a/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala b/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala index dffd53beb9c4..aed66477bc05 100644 --- a/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala +++ b/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala @@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression} import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef, LogicalPlan, MergeAction, MergeIntoTable} +import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution, Distribution} object MinorVersionShim { @@ -29,6 +30,9 @@ object MinorVersionShim { output: Seq[Attribute], isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved, output) + def createClusteredDistribution(expressions: Seq[Expression], numPartitions: Int): Distribution = + ClusteredDistribution(expressions, Some(numPartitions)) + def createMergeIntoTable( targetTable: LogicalPlan, sourceTable: LogicalPlan, diff --git a/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala b/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala index ae0d96e8f8e1..41e5ed4c1201 100644 --- a/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala +++ b/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala @@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression} import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef, LogicalPlan, MergeAction, MergeIntoTable} +import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution, Distribution} object MinorVersionShim { @@ -29,6 +30,12 @@ object MinorVersionShim { output: Seq[Attribute], isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved, output) + def createClusteredDistribution(expressions: Seq[Expression], numPartitions: Int): Distribution = + ClusteredDistribution( + expressions, + requireAllClusterKeys = false, + requiredNumPartitions = Some(numPartitions)) + def createMergeIntoTable( targetTable: LogicalPlan, sourceTable: LogicalPlan, diff --git a/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala b/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala index e3046f564f02..854dda9fa2b8 100644 --- a/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala +++ b/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala @@ -39,6 +39,7 @@ import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment, ColumnDefinition, CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable, MergeRows, SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction} import org.apache.spark.sql.catalyst.plans.logical.MergeRows.Keep +import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution, Distribution} import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.util.{ArrayData, GeneratedColumn, IdentityColumn, ResolveDefaultColumns} import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column, Identifier, StagingTableCatalog, Table, TableCatalog} @@ -243,6 +244,14 @@ class Spark4Shim extends SparkShim { CTERelationRef(cteId, resolved, output.toSeq, isStreaming) } + override def createClusteredDistribution( + expressions: Seq[Expression], + numPartitions: Int): Distribution = + ClusteredDistribution( + expressions, + requireAllClusterKeys = false, + requiredNumPartitions = Some(numPartitions)) + override def supportsHashAggregate( aggregateBufferAttributes: Seq[Attribute], groupingExpression: Seq[Expression]): Boolean = { diff --git a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala index cf36007128ab..c44d47c16fab 100644 --- a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala +++ b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala @@ -43,7 +43,7 @@ object RepartitionLateralVectorSearchInput extends Rule[LogicalPlan] { } } - private def parallelism: Int = { + private[spark] def parallelism: Int = { val value = OptionUtils .getOptionString(SparkConnectorOptions.VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM) diff --git a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala index 2686ef7d7666..77c50acf7059 100644 --- a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala +++ b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala @@ -27,6 +27,7 @@ import org.apache.paimon.predicate.{Predicate, PredicateBuilder} import org.apache.paimon.spark.{PaimonRecordReaderIterator, SparkCatalog, SparkGenericCatalog, SparkTable, SparkUtils} import org.apache.paimon.spark.catalog.{SparkBaseCatalog, SupportView} import org.apache.paimon.spark.catalyst.analysis.ResolvedPaimonView +import org.apache.paimon.spark.catalyst.optimizer.RepartitionLateralVectorSearchInput import org.apache.paimon.spark.catalyst.plans.logical.{CopyIntoLocationCommand, CopyIntoLocationSource, CopyIntoTableCommand, CreateOrReplaceTagCommand, CreatePaimonView, DeleteTagCommand, DropPaimonView, LateralVectorSearch, PaimonCallCommand, PaimonDropPartitions, PaimonTableValuedFunctions, RenameTagCommand, ResolvedIdentifier, ShowPaimonViews, ShowTagsCommand, TruncatePaimonTableWithFilter} import org.apache.paimon.spark.data.SparkInternalRow import org.apache.paimon.spark.format.PaimonFormatTable @@ -43,11 +44,15 @@ import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.{ResolvedNamespace, ResolvedTable} import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, GenericInternalRow, JoinedRow, PredicateHelper, UnsafeProjection} +import org.apache.spark.sql.catalyst.optimizer.BuildRight import org.apache.spark.sql.catalyst.plans.logical.{AddPartitions, CreateTableAsSelect, DescribeRelation, DropPartitions, LogicalPlan, RepairTable, ReplaceTable, ReplaceTableAsSelect, ShowCreateTable} +import org.apache.spark.sql.catalyst.plans.physical.{Distribution, UnspecifiedDistribution} import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.connector.catalog.{Identifier, PaimonLookupCatalog, TableCatalog} -import org.apache.spark.sql.execution.{PaimonDescribeTableExec, SparkPlan, SparkStrategy} +import org.apache.spark.sql.execution.{GlobalLimitExec, PaimonDescribeTableExec, SparkPlan, SparkStrategy, UnaryExecNode} import org.apache.spark.sql.execution.datasources.v2.{DataSourceV2Implicits, DataSourceV2Relation} +import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike +import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, BroadcastNestedLoopJoinExec} import org.apache.spark.sql.execution.shim.{PaimonCreateTableAsSelectStrategy, PaimonReplaceTableAsSelectStrategy, PaimonReplaceTableStrategy} import org.apache.spark.sql.paimon.shims.SparkShimLoader @@ -309,6 +314,30 @@ case class LateralVectorSearchExec( override def output: Seq[Attribute] = child.output ++ projectOutput + // Statistics-based broadcast selection is only known after physical planning. Request a + // distribution here so EnsureRequirements can restore the streamed LIMIT side's parallelism. + override def requiredChildDistribution: Seq[Distribution] = { + if (hasUnrepartitionedGlobalLimit(child)) { + Seq( + SparkShimLoader.shim.createClusteredDistribution( + child.output, + RepartitionLateralVectorSearchInput.parallelism)) + } else { + Seq(UnspecifiedDistribution) + } + } + + private def hasUnrepartitionedGlobalLimit(plan: SparkPlan): Boolean = plan match { + case _: ShuffleExchangeLike => false + case _: GlobalLimitExec => true + case join: BroadcastHashJoinExec => + hasUnrepartitionedGlobalLimit(if (join.buildSide == BuildRight) join.left else join.right) + case join: BroadcastNestedLoopJoinExec => + hasUnrepartitionedGlobalLimit(if (join.buildSide == BuildRight) join.left else join.right) + case unary: UnaryExecNode => hasUnrepartitionedGlobalLimit(unary.child) + case _ => false + } + @transient override lazy val producedAttributes: AttributeSet = { AttributeSet(vectorSearchOutput ++ output.filterNot(attr => inputSet.contains(attr))) } diff --git a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala index 85c923325dd3..b65079d06747 100644 --- a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala +++ b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala @@ -33,6 +33,7 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression} import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Assignment, CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable, SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction} +import org.apache.spark.sql.catalyst.plans.physical.Distribution import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.connector.catalog.{Column, Identifier, StagingTableCatalog, Table, TableCatalog} @@ -159,6 +160,8 @@ trait SparkShim { output: Seq[Attribute], isStreaming: Boolean): CTERelationRef + def createClusteredDistribution(expressions: Seq[Expression], numPartitions: Int): Distribution + def supportsHashAggregate( aggregateBufferAttributes: Seq[Attribute], groupingExpression: Seq[Expression]): Boolean diff --git a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala index 92b68e3e2eef..e46a00358d18 100644 --- a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala +++ b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala @@ -22,15 +22,19 @@ import org.apache.paimon.data.{BinaryString, GenericRow, Timestamp} import org.apache.paimon.manifest.ManifestCommittable import org.apache.paimon.spark.PaimonHiveTestBase import org.apache.paimon.spark.catalyst.plans.logical.{LateralVectorSearch, PaimonTableValuedFunctions} +import org.apache.paimon.spark.execution.LateralVectorSearchExec import org.apache.paimon.utils.DateTimeUtils import org.apache.spark.sql.{DataFrame, Row} +import org.apache.spark.sql.catalyst.optimizer.{BuildLeft, BuildRight} import org.apache.spark.sql.catalyst.plans.logical.{Filter, Repartition} +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.joins.BroadcastHashJoinExec import java.time.LocalDateTime import java.util.Collections -class TableValuedFunctionsTest extends PaimonHiveTestBase { +class TableValuedFunctionsTest extends PaimonHiveTestBase with AdaptiveSparkPlanHelper { test("parse positive limit rejects overflowing long") { val longValue: Long = 4294967297L @@ -252,6 +256,103 @@ class TableValuedFunctionsTest extends PaimonHiveTestBase { } } + test("lateral vector search repartitions automatically broadcast join streamed limited input") { + Seq(false, true).foreach { + aqeEnabled => + withSparkSQLConf( + "spark.sql.adaptive.enabled" -> aqeEnabled.toString, + "spark.sql.autoBroadcastJoinThreshold" -> "1024", + "spark.sql.adaptive.autoBroadcastJoinThreshold" -> "1024", + "spark.paimon.vector-search.lateral-join.parallelism" -> "4" + ) { + withTable("vector_search_source") { + createVectorSearchSource() + + val result = spark.sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT s.gid, s.embs + | FROM ( + | SELECT id AS gid, array(1.0F, 2.0F, 3.0F) AS embs + | FROM range(0, 10000, 1, 8) + | LIMIT 1000 + | ) s + | JOIN VALUES (0L) AS d(gid) + | ON s.gid = d.gid + |) q, + |LATERAL ( + | SELECT gid + | FROM vector_search( + | 'vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + val executedPlan = result.queryExecution.executedPlan + val broadcastJoin = collect(executedPlan) { + case join: BroadcastHashJoinExec => join + }.headOption.getOrElse(fail(executedPlan.toString)) + val lateralVectorSearch = collect(executedPlan) { + case exec: LateralVectorSearchExec => exec + }.headOption.getOrElse(fail(executedPlan.toString)) + + withClue(s"AQE enabled: $aqeEnabled\n$executedPlan") { + assert(broadcastJoin.buildSide == BuildRight) + assert(lateralVectorSearch.child.outputPartitioning.numPartitions == 4) + } + } + } + } + } + + test("lateral vector search preserves automatically broadcast join streamed parallelism") { + Seq(false, true).foreach { + aqeEnabled => + withSparkSQLConf( + "spark.sql.adaptive.enabled" -> aqeEnabled.toString, + "spark.sql.autoBroadcastJoinThreshold" -> "1024", + "spark.sql.adaptive.autoBroadcastJoinThreshold" -> "1024", + "spark.paimon.vector-search.lateral-join.parallelism" -> "4" + ) { + withTable("vector_search_source") { + createVectorSearchSource() + + val result = spark.sql(""" + |SELECT q.gid AS query_gid, r.gid AS result_gid + |FROM ( + | SELECT d.gid, d.embs + | FROM ( + | SELECT id AS gid + | FROM range(0, 10, 1, 1) + | LIMIT 10 + | ) s + | JOIN ( + | SELECT id AS gid, array(1.0F, 2.0F, 3.0F) AS embs + | FROM range(0, 10000, 1, 8) + | ) d + | ON s.gid = d.gid + |) q, + |LATERAL ( + | SELECT gid + | FROM vector_search( + | 'vector_search_source', 'embs', q.embs, 3) + |) AS r + |""".stripMargin) + val executedPlan = result.queryExecution.executedPlan + val broadcastJoin = collect(executedPlan) { + case join: BroadcastHashJoinExec => join + }.headOption.getOrElse(fail(executedPlan.toString)) + val lateralVectorSearch = collect(executedPlan) { + case exec: LateralVectorSearchExec => exec + }.headOption.getOrElse(fail(executedPlan.toString)) + + withClue(s"AQE enabled: $aqeEnabled\n$executedPlan") { + assert(broadcastJoin.buildSide == BuildLeft) + assert(lateralVectorSearch.child.outputPartitioning.numPartitions == 8) + } + } + } + } + } + test("lateral vector search repartitions coalesced limited input") { withTable("vector_search_source") { createVectorSearchSource() diff --git a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala index 8877446d7615..078453090977 100644 --- a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala +++ b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala @@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression} import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef, LogicalPlan, MergeAction, MergeIntoTable} +import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution, Distribution} object MinorVersionShim { @@ -29,6 +30,12 @@ object MinorVersionShim { output: Seq[Attribute], isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved, output, isStreaming) + def createClusteredDistribution(expressions: Seq[Expression], numPartitions: Int): Distribution = + ClusteredDistribution( + expressions, + requireAllClusterKeys = false, + requiredNumPartitions = Some(numPartitions)) + def createMergeIntoTable( targetTable: LogicalPlan, sourceTable: LogicalPlan, diff --git a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala index 9bde530b28c7..f14c21259044 100644 --- a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala +++ b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala @@ -38,6 +38,7 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment, CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable, SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction} +import org.apache.spark.sql.catalyst.plans.physical.Distribution // NOTE: `MergeRows` / `MergeRows.Keep` were introduced in Spark 3.4. We access them only via // reflection inside the `mergeRowsKeep*` method bodies so that loading `Spark3Shim` does not fail // on Spark 3.2 / 3.3 runtimes that still ship `paimon-spark3-common` (the module targets 3.5.8 at @@ -228,6 +229,11 @@ class Spark3Shim extends SparkShim { isStreaming: Boolean): CTERelationRef = MinorVersionShim.createCTERelationRef(cteId, resolved, output, isStreaming) + override def createClusteredDistribution( + expressions: Seq[Expression], + numPartitions: Int): Distribution = + MinorVersionShim.createClusteredDistribution(expressions, numPartitions) + override def supportsHashAggregate( aggregateBufferAttributes: Seq[Attribute], groupingExpression: Seq[Expression]): Boolean = diff --git a/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala b/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala index 7a7cdc70f531..91664858ac3b 100644 --- a/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala +++ b/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala @@ -39,6 +39,7 @@ import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment, ColumnDefinition, CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable, MergeRows, SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction} import org.apache.spark.sql.catalyst.plans.logical.MergeRows.{Copy, Insert, Keep, Update} +import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution, Distribution} import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.util.{ArrayData, GeneratedColumn, IdentityColumn, ResolveDefaultColumns} import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column, Identifier, StagingTableCatalog, Table, TableCatalog} @@ -227,6 +228,14 @@ class Spark4Shim extends SparkShim { CTERelationRef(cteId, resolved, output.toSeq, isStreaming) } + override def createClusteredDistribution( + expressions: Seq[Expression], + numPartitions: Int): Distribution = + ClusteredDistribution( + expressions, + requireAllClusterKeys = false, + requiredNumPartitions = Some(numPartitions)) + override def supportsHashAggregate( aggregateBufferAttributes: Seq[Attribute], groupingExpression: Seq[Expression]): Boolean = {