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
4 changes: 4 additions & 0 deletions .github/workflows/pr_build_linux.yml
Original file line number Diff line number Diff line change
Expand Up @@ -516,6 +516,7 @@ jobs:
org.apache.comet.exec.CometShuffleEncryptionSuite
org.apache.comet.exec.CometShuffleManagerSuite
org.apache.comet.exec.CometShuffleReadCoalesceSuite
org.apache.spark.sql.comet.execution.shuffle.CometBlockStoreShuffleReaderSuite
org.apache.comet.exec.CometAsyncShuffleSuite
org.apache.comet.exec.DisableAQECometShuffleSuite
org.apache.comet.exec.DisableAQECometAsyncShuffleSuite
Expand All @@ -527,6 +528,8 @@ jobs:
org.apache.comet.exec.CometAggregateSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
org.apache.spark.sql.comet.CometNativeTaskKillSuite
org.apache.comet.exec.CometEmptyRelationExecSuite
org.apache.comet.exec.CometInMemoryCacheSuite
org.apache.comet.exec.CometInMemoryCacheKryoSuite
Expand Down Expand Up @@ -602,6 +605,7 @@ jobs:
org.apache.comet.CometVariantTypeSuite
org.apache.comet.CometHashExpressionSuite
org.apache.comet.CometTemporalExpressionSuite
org.apache.comet.CometTimestampComparisonSuite
org.apache.comet.CometArrayExpressionSuite
org.apache.comet.CometNativeCastSuite
org.apache.comet.CometDateTimeUtilsSuite
Expand Down
4 changes: 4 additions & 0 deletions .github/workflows/pr_build_macos.yml
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,7 @@ jobs:
org.apache.comet.exec.CometShuffleEncryptionSuite
org.apache.comet.exec.CometShuffleManagerSuite
org.apache.comet.exec.CometShuffleReadCoalesceSuite
org.apache.spark.sql.comet.execution.shuffle.CometBlockStoreShuffleReaderSuite
org.apache.comet.exec.CometAsyncShuffleSuite
org.apache.comet.exec.DisableAQECometShuffleSuite
org.apache.comet.exec.DisableAQECometAsyncShuffleSuite
Expand All @@ -175,6 +176,8 @@ jobs:
org.apache.comet.exec.CometAggregateSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
org.apache.spark.sql.comet.CometNativeTaskKillSuite
org.apache.comet.exec.CometEmptyRelationExecSuite
org.apache.comet.exec.CometInMemoryCacheSuite
org.apache.comet.exec.CometInMemoryCacheKryoSuite
Expand Down Expand Up @@ -250,6 +253,7 @@ jobs:
org.apache.comet.CometVariantTypeSuite
org.apache.comet.CometHashExpressionSuite
org.apache.comet.CometTemporalExpressionSuite
org.apache.comet.CometTimestampComparisonSuite
org.apache.comet.CometArrayExpressionSuite
org.apache.comet.CometNativeCastSuite
org.apache.comet.CometDateTimeUtilsSuite
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@

package org.apache.spark.sql.comet

import scala.jdk.CollectionConverters._

import org.apache.spark.rdd.RDD
import org.apache.spark.sql.catalyst.expressions._
import org.apache.spark.sql.catalyst.plans.QueryPlan
Expand All @@ -30,7 +32,7 @@ import org.apache.spark.sql.types.StructType
import org.apache.spark.sql.vectorized.ColumnarBatch

import org.apache.comet.contrib.delta.DeltaSparkScanEnvelope
import org.apache.comet.serde.OperatorOuterClass
import org.apache.comet.serde.{OperatorOuterClass, QueryContextInterner}
import org.apache.comet.serde.OperatorOuterClass.Operator

/**
Expand Down Expand Up @@ -148,12 +150,17 @@ case class CometDeltaNativeScanExec(

override def doExecuteColumnar(): RDD[ColumnarBatch] = {
val nativeMetrics = CometMetricNode.fromCometPlan(this)
val serializedPlan = CometExec.serializeNativePlan(nativeOp)
val plan = PlanDataInjector.internScans(QueryContextInterner.intern(nativeOp))
val sqlTextPool = new QueryContextInterner.Pool(plan.getSqlTextPoolList.asScala.toSeq)
val commonByKey =
PlanDataInjector.internCommons(plan, Map(sourceKey -> commonData), sqlTextPool)
val serializedPlan = CometExec.serializeNativePlan(
plan.toBuilder.addAllSqlTextPool(sqlTextPool.added.asJava).build())

new CometExecRDD(
sparkContext,
Seq.empty,
Map(sourceKey -> commonData),
commonByKey,
Map(sourceKey -> perPartitionData),
serializedPlan,
PlanDataInjector.planFingerprint(serializedPlan),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,16 @@ class DeltaPlanDataInjector extends PlanDataInjector {

op.toBuilder.setContribScan(DeltaSparkScanEnvelope.pack(scanBuilder.build())).build()
}

override def internScan(op: Operator, pool: QueryContextInterner.Pool): Operator = {
val scan = DeltaSparkScanEnvelope.unpack(op)
val common = QueryContextInterner.internScanCommon(scan.getCommon, pool)
if (common eq scan.getCommon) op
else {
val interned = scan.toBuilder.setCommon(common).build()
op.toBuilder.setContribScan(DeltaSparkScanEnvelope.pack(interned)).build()
}
}
}

object DeltaPlanDataInjector {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3881,4 +3881,54 @@ class CometDeltaNativeScanSuite extends CometDeltaTestBase {
}
}
}

test("stage task binaries over delta scans do not repeat the query text per expression") {
withTempDir { dir =>
val paths = (0 until 20).map { i =>
val path = new File(dir, s"t$i").getAbsolutePath
spark
.range(0, 200)
.selectExpr("id", s"id % ${i + 3} AS k", "concat('data/', cast(id AS string)) AS name")
.write
.format("delta")
.save(path)
path
}
val sql = paths
.map(p =>
s"SELECT k, name FROM delta.`$p` WHERE name LIKE 'data/%' AND NOT endswith(name, '/') " +
"AND id % 7 <> 3 AND length(name) > 5")
.mkString("SELECT k, count(*) AS c FROM (\n", "\nUNION ALL\n", "\n) GROUP BY k")
def mapStageBytes(cometEnabled: Boolean): Long = {
var bytes = 0L
withSQLConf(
CometConf.COMET_ENABLED.key -> cometEnabled.toString,
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
val df = spark.sql(sql)
val serializer = org.apache.spark.SparkEnv.get.closureSerializer.newInstance()
val deps = mutable.ArrayBuffer.empty[org.apache.spark.ShuffleDependency[_, _, _]]
def walk(rdd: org.apache.spark.rdd.RDD[_]): Unit = {
rdd.partitions
rdd.dependencies.foreach {
case d: org.apache.spark.ShuffleDependency[_, _, _] => deps += d
case d => walk(d.rdd)
}
}
walk(df.queryExecution.executedPlan.execute())
deps.foreach { d =>
walk(d.rdd)
bytes = math.max(bytes, serializer.serialize((d.rdd, d)).limit().toLong)
}
checkSparkAnswer(df)
}
bytes
}
val vanilla = mapStageBytes(cometEnabled = false)
val comet = mapStageBytes(cometEnabled = true)
// scalastyle:off println
println(s"delta task binary: vanilla=$vanilla comet=$comet sql=${sql.length}")
// scalastyle:on println
assert(comet < vanilla + 20L * sql.length * 3, s"vanilla=$vanilla comet=$comet")
}
}
}
2 changes: 2 additions & 0 deletions docs/source/user-guide/latest/tuning.md
Original file line number Diff line number Diff line change
Expand Up @@ -614,6 +614,8 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n
- `agg` for hash, object hash and sort aggregates, over the leaves of the grouping keys, half for each phase of a
two-phase aggregate, plus the price of the class of each aggregate function (`aggDeclarative`, `aggCollectList`,
`aggCollectSet`, `aggPercentile`, `aggPercentileApprox`, `aggOther`) and `aggObjectHash` for an object hash aggregate.
The leaves of grouping keys holding an array add `aggArrayKey`, also half for each phase, and the leaves of a grouping
key computed through the JVM codegen dispatcher add `codegenDispatch` once, in the phase that computes it.
- `window` over the leaves of its input, plus `windowAggregate`, `windowOffset` or `windowRank` for each window
function, at `L` the number of window functions; `wglPartial` and `wglFinal` for window group limits.
- `expand`, per projection, and `generate`: free with Spark's whole-stage codegen, `expandNoCodegen` and
Expand Down
166 changes: 166 additions & 0 deletions native/common/src/cancellation.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
// 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.

//! Cancellation of a native plan whose Spark task has been killed. The JVM sets it from outside
//! the task thread, and the plan's operators, its blocking waits and its Tokio tasks stop at
//! their next check instead of running the plan to completion.

use datafusion::common::DataFusionError;
use datafusion::execution::TaskContext;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};

pub const CANCELLED_MESSAGE: &str =
"Comet native execution was cancelled because its Spark task was killed";

type Callback = Box<dyn FnOnce() + Send>;

#[derive(Default)]
pub struct PlanCancellation {
cancelled: AtomicBool,
waiters: Mutex<Waiters>,
}

#[derive(Default)]
struct Waiters {
wakers: Vec<Waker>,
callbacks: Vec<Callback>,
}

impl std::fmt::Debug for PlanCancellation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PlanCancellation")
.field("cancelled", &self.is_cancelled())
.finish()
}
}

impl PlanCancellation {
pub fn new() -> Arc<Self> {
Arc::new(Self::default())
}

/// The cancellation of the plan `context` runs for, if it has one.
pub fn of(context: &TaskContext) -> Option<Arc<Self>> {
context.session_config().get_extension::<Self>()
}

pub fn cancel(&self) {
if self.cancelled.swap(true, Ordering::AcqRel) {
return;
}
let Waiters { wakers, callbacks } = std::mem::take(&mut *self.lock());
wakers.into_iter().for_each(Waker::wake);
callbacks.into_iter().for_each(|callback| callback());
}

pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::Acquire)
}

pub fn check(&self) -> Result<(), DataFusionError> {
if self.is_cancelled() {
Err(cancelled_error())
} else {
Ok(())
}
}

/// `Ready` once cancelled; otherwise wakes `cx` when it is.
pub fn poll_cancelled(&self, cx: &mut Context<'_>) -> Poll<()> {
if self.is_cancelled() {
return Poll::Ready(());
}
let mut waiters = self.lock();
if self.is_cancelled() {
return Poll::Ready(());
}
if !waiters.wakers.iter().any(|w| w.will_wake(cx.waker())) {
waiters.wakers.push(cx.waker().clone());
}
Poll::Pending
}

/// Runs `callback` once the plan is cancelled, at once if it already is.
pub fn on_cancel(&self, callback: impl FnOnce() + Send + 'static) {
{
let mut waiters = self.lock();
if !self.is_cancelled() {
waiters.callbacks.push(Box::new(callback));
return;
}
}
callback()
}

fn lock(&self) -> std::sync::MutexGuard<'_, Waiters> {
self.waiters.lock().unwrap_or_else(|e| e.into_inner())
}
}

pub fn cancelled_error() -> DataFusionError {
DataFusionError::Execution(CANCELLED_MESSAGE.to_string())
}

#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
use std::task::Wake;

struct CountingWaker(AtomicUsize);

impl Wake for CountingWaker {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}

#[test]
fn cancel_wakes_each_registered_waker_once_and_runs_callbacks() {
let cancellation = PlanCancellation::new();
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
let waker = Waker::from(Arc::clone(&counter));
let mut cx = Context::from_waker(&waker);
assert!(cancellation.poll_cancelled(&mut cx).is_pending());
assert!(cancellation.poll_cancelled(&mut cx).is_pending());
let calls = Arc::new(AtomicUsize::new(0));
let c = Arc::clone(&calls);
cancellation.on_cancel(move || {
c.fetch_add(1, Ordering::SeqCst);
});
assert!(cancellation.check().is_ok());

cancellation.cancel();
cancellation.cancel();
assert_eq!(counter.0.load(Ordering::SeqCst), 1);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert!(cancellation.poll_cancelled(&mut cx).is_ready());
assert!(cancellation
.check()
.unwrap_err()
.to_string()
.contains(CANCELLED_MESSAGE));

let c = Arc::clone(&calls);
cancellation.on_cancel(move || {
c.fetch_add(1, Ordering::SeqCst);
});
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
}
2 changes: 2 additions & 0 deletions native/common/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@
// specific language governing permissions and limitations
// under the License.

pub mod cancellation;
mod error;
pub mod offset_extents;
mod query_context;
mod schema;
pub mod struct_nulls;
Expand Down
Loading
Loading