Skip to content
Merged
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
5 changes: 5 additions & 0 deletions docs/source/contributor-guide/native_shuffle.md
Original file line number Diff line number Diff line change
Expand Up @@ -222,6 +222,11 @@ decides during plan serialization. When direct read is enabled and the sink's in
exchange, `convertToShuffleScan` emits a `ShuffleScan` operator. When either is false the sink falls
through to the base `CometSink.convert`, which emits the usual `Scan`.

Under AQE, an operator that shares its logical node with a shuffle stage, such as the final aggregate
of a two-phase aggregate, comes back from re-planning as the node already planned, whose input was
serialized as a `Scan` before the stage existed. `CometExecRule` refreshes such a node: once its input
is a sink that emits a `ShuffleScan`, that `ShuffleScan` replaces the stale `Scan` leaf.

The two are not alternatives on failure. If any output type fails `supportedSinkDataType`,
`convertToShuffleScan` records the fallback reason `Unsupported data type for shuffle direct read`
and returns `None`. It does not retry as a regular `Scan`, and retrying would not help, because
Expand Down
96 changes: 93 additions & 3 deletions spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
package org.apache.comet.rules

import scala.collection.mutable.ListBuffer
import scala.jdk.CollectionConverters._

import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder}
Expand Down Expand Up @@ -57,6 +58,7 @@ import org.apache.comet.CometConf.{COMET_SPARK_TO_ARROW_ENABLED, COMET_SPARK_TO_
import org.apache.comet.CometSparkSessionExtensions._
import org.apache.comet.rules.CometExecRule.allExecs
import org.apache.comet.serde._
import org.apache.comet.serde.OperatorOuterClass.Operator
import org.apache.comet.serde.operator._
import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimCometWindowGroupLimit, ShimSubqueryBroadcast}

Expand Down Expand Up @@ -88,6 +90,13 @@ object CometExecRule {
val COMET_UNSAFE_PARTIAL: TreeNodeTag[String] =
TreeNodeTag[String]("comet.unsafePartialAgg")

/**
* Info message for a native operator that keeps a `Scan` over an input that now reads its
* shuffle directly, because its native plan's leaves do not line up with its inputs.
*/
val STALE_SCAN_KEPT: String =
"Shuffle direct read not applied: the native plan's leaves do not match its inputs"

/**
* Fully native operators.
*/
Expand Down Expand Up @@ -605,7 +614,7 @@ case class CometExecRule(session: SparkSession)
}

plan.transformUp { case op =>
val converted = convertNode(op)
val converted = convertNode(refreshStaleShuffleScans(op))
// Replace SubqueryBroadcastExec with CometSubqueryBroadcastExec in DPP expressions
// when the broadcast child has a Comet plan underneath. This enables exchange reuse
// between the DPP subquery and the join's CometBroadcastExchangeExec because both
Expand All @@ -614,6 +623,88 @@ case class CometExecRule(session: SparkSession)
}
}

/**
* AQE re-plans around a materialized stage by reusing the physical node linked to it, so a
* native operator that shares its logical node with a shuffle stage (the final aggregate of a
* two-phase aggregate) keeps the native plan it got while that input was a bare exchange, read
* through a plain `Scan`. Once the input is a sink that reads the shuffle directly, its
* `ShuffleScan` takes the place of the stale leaf. The leaf is patched in place because
* converting the node again from `originalPlan` would drop the stage's logical link that AQE
* relies on and re-run serde on a node that is already planned.
*/
private def refreshStaleShuffleScans(op: SparkPlan): SparkPlan = op match {
case _ if scansChildAsSeparateBlock(op) => op
case native: CometNativeExec if native.children.nonEmpty =>
refreshedNativeOp(native) match {
case Some(newOp) =>
val refreshed = native.withRefreshedNativeOp(newOp)
// An operator that does not hold its native plan as a field cannot take a new one.
if (refreshed.nativeOp eq newOp) refreshed else op
case None => op
}
case _ => op
}

/**
* The native plan of `native` with each `Scan` leaf whose input is now a `ShuffleScan` of the
* same field types replaced by that `ShuffleScan`, or None if there is no such leaf or the plan
* children cannot be matched to the leaves. In the second case `native` gets an info message
* for extended explain, since the block still reads its shuffle through the JVM.
*/
private def refreshedNativeOp(native: CometNativeExec): Option[Operator] = {
val children = native.children.collect { case child: CometNativeExec => child }
// Only a sink that reads a shuffle directly, or a native child that may hold one, can feed
// a `ShuffleScan`.
val mayFeedShuffleScan = children.exists {
case sink: CometSinkPlaceHolder => sink.nativeOp.hasShuffleScan
case _ => true
}
if (children.length != native.children.length || !mayFeedShuffleScan) return None
val leaves = CometExec.nativeLeaves(native.nativeOp)
if (!leaves.exists(_.hasScan)) return None

// Each plan child feeds a run of leaves, in order: a sink feeds one, and a native child
// feeds the leaves of its own native plan.
val current = children.flatMap {
case sink: CometSinkPlaceHolder => Seq(sink.nativeOp)
case child => CometExec.nativeLeaves(child.nativeOp)
}
def isStale(leaf: Operator, input: Operator): Boolean = leaf.hasScan && input.hasShuffleScan
val stale = leaves.zip(current).filter { case (leaf, input) => isStale(leaf, input) }
if (stale.isEmpty) return None
val isRefreshable = current.length == leaves.length &&
stale.forall { case (leaf, input) =>
leaf.getScan.getFieldsList == input.getShuffleScan.getFieldsList
}
if (!isRefreshable) {
withInfo(native, CometExecRule.STALE_SCAN_KEPT)
return None
}
val inputs = current.iterator
Some(mapLeaves(native.nativeOp) { leaf =>
val input = inputs.next()
if (isStale(leaf, input)) input else leaf
})
}

/** `op` with `f` applied to each childless operator, in `CometExec.nativeLeaves` order. */
private def mapLeaves(op: Operator)(f: Operator => Operator): Operator =
if (op.getChildrenCount == 0) {
f(op)
} else {
val children = op.getChildrenList.asScala.map(mapLeaves(_)(f))
op.toBuilder.clearChildren().addAllChildren(children.asJava).build()
}

/**
* Whether `op` is a writer whose native plan builds its own `Scan` over its child, so that the
* child runs as a separate native block instead of inside that plan.
*/
private def scansChildAsSeparateBlock(op: SparkPlan): Boolean = op match {
case _: CometNativeWriteExec | _: CometIcebergWriteExec | _: CometWriteFilesExec => true
case _ => false
}

/**
* Replace SubqueryBroadcastExec with CometSubqueryBroadcastExec in a node's expressions
* (non-AQE DPP), and wrap SubqueryAdaptiveBroadcastExec in CometSubqueryAdaptiveBroadcastExec
Expand Down Expand Up @@ -929,8 +1020,7 @@ case class CometExecRule(session: SparkSession)
// its child (e.g., CometNativeScanExec, or a CometProject over an AQEShuffleRead)
// needs its own serialization. Reset the flag so children can start their own native
// execution blocks.
if (op.isInstanceOf[CometNativeWriteExec] || op.isInstanceOf[CometIcebergWriteExec] ||
op.isInstanceOf[CometWriteFilesExec]) {
if (scansChildAsSeparateBlock(op)) {
firstNativeOp = true
}

Expand Down
58 changes: 36 additions & 22 deletions spark/src/main/scala/org/apache/spark/sql/comet/operators.scala
Original file line number Diff line number Diff line change
Expand Up @@ -728,6 +728,25 @@ object CometExec {
bytes
}

/**
* The childless operators of a native plan, depth first. Its `Scan` and `ShuffleScan` leaves
* read the block's inputs in this order.
*/
def nativeLeaves(op: Operator): Seq[Operator] =
if (op.getChildrenCount == 0) Seq(op)
else op.getChildrenList.asScala.toSeq.flatMap(nativeLeaves)

/**
* The input indices of a native plan that its `ShuffleScan` leaves read. Each `Scan` or
* `ShuffleScan` leaf reads one input, in [[nativeLeaves]] order.
*/
def findShuffleScanIndices(plan: Operator): Set[Int] =
nativeLeaves(plan)
.filter(leaf => leaf.hasScan || leaf.hasShuffleScan)
.zipWithIndex
.collect { case (leaf, index) if leaf.hasShuffleScan => index }
.toSet

def getCometIterator(
inputObjects: Array[Object],
numOutputCols: Int,
Expand Down Expand Up @@ -995,7 +1014,7 @@ abstract class CometNativeExec extends CometExec {
// (`ShuffleQueryStageExec`), so a bare non-AQE `CometShuffleExchangeExec` always serializes
// as a regular Scan regardless of `COMET_SHUFFLE_DIRECT_READ_ENABLED`. Driving the JVM
// dispatch from `shuffleScanIndices` instead of the conf keeps the two aligned.
val shuffleScanIndices = findShuffleScanIndices(nativeOp)
val shuffleScanIndices = CometExec.findShuffleScanIndices(nativeOp)

def isBroadcastInput(plan: SparkPlan): Boolean = plan match {
case _: CometBroadcastExchangeExec => true
Expand Down Expand Up @@ -1176,27 +1195,6 @@ abstract class CometNativeExec extends CometExec {
}
}

/**
* Walk the protobuf operator tree depth-first to find which input indices correspond to
* ShuffleScan vs Scan leaf nodes. Each Scan or ShuffleScan leaf consumes one input in order.
*/
private def findShuffleScanIndices(plan: OperatorOuterClass.Operator): Set[Int] = {
var scanIndex = 0
val indices = mutable.Set.empty[Int]
def walk(op: OperatorOuterClass.Operator): Unit = {
if (op.hasShuffleScan) {
indices += scanIndex
scanIndex += 1
} else if (op.hasScan) {
scanIndex += 1
} else {
op.getChildrenList.asScala.foreach(walk)
}
}
walk(plan)
indices.toSet
}

/**
* Converts this native Comet operator and its children into a native block which can be
* executed as a whole (i.e., in a single JNI call) from the native side.
Expand All @@ -1216,6 +1214,22 @@ abstract class CometNativeExec extends CometExec {
makeCopy(newArgs).asInstanceOf[CometNativeExec]
}

/**
* Copies this operator with `newOp` as its native plan and no serialized plan, so that
* `convertBlock` serializes the block again.
*/
def withRefreshedNativeOp(newOp: Operator): CometNativeExec = {
def transform(arg: Any): AnyRef = arg match {
case op: Operator if op eq nativeOp => newOp
case _: SerializedPlan => SerializedPlan(None)
case other: AnyRef => other
case null => null
}

val newArgs = mapProductIterator(transform)
makeCopy(newArgs).asInstanceOf[CometNativeExec]
}

/**
* Cleans the serialized plan from this native Comet operator. Used to canonicalize the plan.
*/
Expand Down
Loading
Loading