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
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ import org.apache.spark.sql.catalyst.catalog.{CatalogTable, CatalogTableType, Ca
import org.apache.spark.sql.connector.catalog.{Identifier, StagedTable, StagingTableCatalog, SupportsWrite, Table, TableCapability, TableCatalog, TableChange}
import org.apache.spark.sql.connector.catalog.TableCapability._
import org.apache.spark.sql.connector.expressions.Transform
import org.apache.spark.sql.connector.write.{LogicalWriteInfo, V1Write, WriteBuilder}
import org.apache.spark.sql.connector.write.{LogicalWriteInfo, SupportsTruncate, V1Write, WriteBuilder}
import org.apache.spark.sql.delta.{ColumnWithDefaultExprUtils, DeltaConfigs, DeltaErrors, DeltaLog, DeltaOptions}
import org.apache.spark.sql.delta.catalog.DeltaCatalog
import org.apache.spark.sql.delta.commands.{TableCreationModes, WriteIntoDelta}
Expand Down Expand Up @@ -471,7 +471,7 @@ abstract class GpuDeltaCatalogBase(
override def abortStagedChanges(): Unit = {}

override def capabilities(): util.Set[TableCapability] = {
Set(V1_BATCH_WRITE).asJava
Set(V1_BATCH_WRITE, TRUNCATE).asJava
}

override def newWriteBuilder(info: LogicalWriteInfo): WriteBuilder = {
Expand All @@ -482,7 +482,9 @@ abstract class GpuDeltaCatalogBase(
/*
* WriteBuilder for creating a Delta table.
*/
private class DeltaV1WriteBuilder extends WriteBuilder {
private class DeltaV1WriteBuilder extends WriteBuilder with SupportsTruncate {
override def truncate(): this.type = this

override def build(): V1Write = new V1Write {
override def toInsertableRelation(): InsertableRelation = {
new InsertableRelation {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@ import org.apache.spark.sql.delta.catalog.DeltaTableV2
import org.apache.spark.sql.delta.commands.{DeleteCommand, MergeIntoCommand, OptimizeTableCommand, UpdateCommand}
import org.apache.spark.sql.execution.command.RunnableCommand
import org.apache.spark.sql.execution.datasources.FileFormat
import org.apache.spark.sql.execution.datasources.v2.AppendDataExecV1
import org.apache.spark.sql.execution.datasources.v2.{AppendDataExecV1, OverwriteByExpressionExecV1}

object Delta33xProvider extends DeltaProviderBase with Logging {

Expand All @@ -52,6 +52,21 @@ object Delta33xProvider extends DeltaProviderBase with Logging {
}
}

override def tagForGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): Unit = {
if (!meta.conf.isDeltaWriteEnabled) {
meta.willNotWorkOnGpu("Delta Lake output acceleration has been disabled. To enable set " +
s"${RapidsConf.ENABLE_DELTA_WRITE} to true")
}

cpuExec.table match {
case _: DeltaTableV2 => super.tagForGpu(cpuExec, meta)
case _: GpuDeltaCatalog#GpuStagedDeltaTableV2 =>
case _ => meta.willNotWorkOnGpu(s"${cpuExec.table} table class not supported on GPU")
}
}

override def getRunnableCommandRules: Map[Class[_ <: RunnableCommand],
RunnableCommandRule[_ <: RunnableCommand]] = {
Seq(
Expand Down Expand Up @@ -109,4 +124,17 @@ object Delta33xProvider extends DeltaProviderBase with Logging {
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}

override def convertToGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): GpuExec = {
cpuExec.table match {
case _: DeltaTableV2 =>
super.convertToGpu(cpuExec, meta)
case _: GpuDeltaCatalog#GpuStagedDeltaTableV2 =>
GpuOverwriteByExpressionExecV1(
cpuExec.table, cpuExec.plan, cpuExec.refreshCache, cpuExec.write)
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ import org.apache.spark.sql.delta.commands.{DeleteCommand, MergeIntoCommand, Opt
import org.apache.spark.sql.delta.rapids.GpuDeltaCatalog4x
import org.apache.spark.sql.execution.command.RunnableCommand
import org.apache.spark.sql.execution.datasources.FileFormat
import org.apache.spark.sql.execution.datasources.v2.AppendDataExecV1
import org.apache.spark.sql.execution.datasources.v2.{AppendDataExecV1, OverwriteByExpressionExecV1}

object Delta40xProvider extends DeltaProviderBase with Logging {

Expand All @@ -55,6 +55,21 @@ object Delta40xProvider extends DeltaProviderBase with Logging {
}
}

override def tagForGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): Unit = {
if (!meta.conf.isDeltaWriteEnabled) {
meta.willNotWorkOnGpu("Delta Lake output acceleration has been disabled. To enable set " +
s"${RapidsConf.ENABLE_DELTA_WRITE} to true")
}

cpuExec.table match {
case _: DeltaTableV2 => super.tagForGpu(cpuExec, meta)
case _: GpuDeltaCatalog4x#GpuStagedDeltaTableV2 =>
case _ => meta.willNotWorkOnGpu(s"${cpuExec.table} table class not supported on GPU")
}
}

override def getRunnableCommandRules: Map[Class[_ <: RunnableCommand],
RunnableCommandRule[_ <: RunnableCommand]] = {
Seq(
Expand Down Expand Up @@ -119,4 +134,17 @@ object Delta40xProvider extends DeltaProviderBase with Logging {
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}

override def convertToGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): GpuExec = {
cpuExec.table match {
case _: DeltaTableV2 =>
super.convertToGpu(cpuExec, meta)
case _: GpuDeltaCatalog4x#GpuStagedDeltaTableV2 =>
GpuOverwriteByExpressionExecV1(
cpuExec.table, cpuExec.plan, cpuExec.refreshCache, cpuExec.write)
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ import org.apache.spark.sql.delta.commands.{DeleteCommand, MergeIntoCommand, Opt
import org.apache.spark.sql.delta.rapids.GpuDeltaCatalog4x
import org.apache.spark.sql.execution.command.RunnableCommand
import org.apache.spark.sql.execution.datasources.FileFormat
import org.apache.spark.sql.execution.datasources.v2.AppendDataExecV1
import org.apache.spark.sql.execution.datasources.v2.{AppendDataExecV1, OverwriteByExpressionExecV1}

object Delta41xProvider extends DeltaProviderBase with Logging {

Expand All @@ -55,6 +55,21 @@ object Delta41xProvider extends DeltaProviderBase with Logging {
}
}

override def tagForGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): Unit = {
if (!meta.conf.isDeltaWriteEnabled) {
meta.willNotWorkOnGpu("Delta Lake output acceleration has been disabled. To enable set " +
s"${RapidsConf.ENABLE_DELTA_WRITE} to true")
}

cpuExec.table match {
case _: DeltaTableV2 => super.tagForGpu(cpuExec, meta)
case _: GpuDeltaCatalog4x#GpuStagedDeltaTableV2 =>
case _ => meta.willNotWorkOnGpu(s"${cpuExec.table} table class not supported on GPU")
}
}

override def getRunnableCommandRules: Map[Class[_ <: RunnableCommand],
RunnableCommandRule[_ <: RunnableCommand]] = {
Seq(
Expand Down Expand Up @@ -119,4 +134,17 @@ object Delta41xProvider extends DeltaProviderBase with Logging {
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}

override def convertToGpu(
cpuExec: OverwriteByExpressionExecV1,
meta: OverwriteByExpressionExecV1Meta): GpuExec = {
cpuExec.table match {
case _: DeltaTableV2 =>
super.convertToGpu(cpuExec, meta)
case _: GpuDeltaCatalog4x#GpuStagedDeltaTableV2 =>
GpuOverwriteByExpressionExecV1(
cpuExec.table, cpuExec.plan, cpuExec.refreshCache, cpuExec.write)
case unknown => throw new IllegalStateException(s"$unknown doesn't match any of the known ")
}
}
}
53 changes: 51 additions & 2 deletions integration_tests/src/main/python/delta_lake_write_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -745,7 +745,8 @@ def test_delta_atomic_create_table_as_select(spark_tmp_table_factory, spark_tmp_
@pytest.mark.skipif(is_before_spark_320(), reason="Delta Lake writes are not supported before Spark 3.2.x")
@pytest.mark.parametrize("enable_deletion_vectors", deletion_vector_values_with_xfail_reasons(
enabled_xfail_reason="https://github.com/NVIDIA/spark-rapids/issues/12041"), ids=idfn)
@pytest.mark.xfail(is_spark_356_or_later(), reason="https://github.com/delta-io/delta/issues/4671")
@pytest.mark.xfail(is_spark_356_or_later() and not is_spark_400_or_later(),
reason="https://github.com/delta-io/delta/issues/4671")
@pytest.mark.xfail(is_databricks_runtime(), reason="https://github.com/NVIDIA/spark-rapids/issues/11169")
def test_delta_atomic_replace_table_as_select(spark_tmp_table_factory, spark_tmp_path, enable_deletion_vectors):
_atomic_write_table_as_select(delta_write_gens, spark_tmp_table_factory, spark_tmp_path,
Expand Down Expand Up @@ -816,13 +817,61 @@ def test_delta_ctas_sql(spark_tmp_table_factory, enable_deletion_vectors, use_cd
@pytest.mark.parametrize("enable_deletion_vectors", deletion_vector_values_with_xfail_reasons(
enabled_xfail_reason="https://github.com/NVIDIA/spark-rapids/issues/12041"), ids=idfn)
@pytest.mark.parametrize("use_cdf", [True, False], ids=idfn)
@pytest.mark.xfail(is_spark_356_or_later(), reason="https://github.com/delta-io/delta/issues/4671")
@pytest.mark.xfail(is_spark_356_or_later() and not is_spark_400_or_later(),
reason="https://github.com/delta-io/delta/issues/4671")
@pytest.mark.xfail(is_databricks_runtime(), reason="https://github.com/NVIDIA/spark-rapids/issues/11169")
def test_delta_rtas_sql(spark_tmp_table_factory, enable_deletion_vectors, use_cdf):
_atomic_write_table_as_select_sql(delta_write_gens, spark_tmp_table_factory,
True, enable_deletion_vectors, use_cdf)


@allow_non_gpu_conditional(
is_databricks_runtime(),
'AppendDataExecV1, AtomicCreateTableAsSelectExec, AtomicReplaceTableAsSelectExec')
@allow_non_gpu('DataWritingCommandExec', 'WriteFilesExec', *delta_meta_allow)
@delta_lake
@ignore_order(local=True)
@pytest.mark.skipif(not is_spark_400_or_later(),
reason="Delta Lake 4.0 contains the native truncate capability")
def test_delta_rtas_truncate_capability(spark_tmp_table_factory):
cpu_table = spark_tmp_table_factory.get()
gpu_table = spark_tmp_table_factory.get()
confs = copy_and_update(writer_confs, delta_writes_enabled_conf)

def create_initial_tables(spark):
for table in [cpu_table, gpu_table]:
spark.sql(f"CREATE TABLE {table} USING DELTA AS SELECT id FROM range(10)")

def replace_table(spark, table):
spark.sql(
f"CREATE OR REPLACE TABLE {table} USING DELTA "
f"AS SELECT id FROM range(20) WHERE id >= 10")

with_cpu_session(create_initial_tables, conf=confs)
with_cpu_session(lambda spark: replace_table(spark, cpu_table), conf=confs)

callback = spark_jvm().org.apache.spark.sql.rapids.ExecutionPlanCaptureCallback
callback.startCapture()
try:
with_gpu_session(lambda spark: replace_table(spark, gpu_table), conf=confs)
plans = callback.getResultsWithTimeout(10000)
assert any(callback.contains(plan, "GpuAtomicReplaceTableAsSelectExec")
for plan in plans), "GpuAtomicReplaceTableAsSelectExec was not executed"
if not is_databricks_runtime():
assert any(callback.contains(plan, "GpuOverwriteByExpressionExecV1")
for plan in plans), "GpuOverwriteByExpressionExecV1 was not executed"
finally:
callback.endCapture()

def read_ids(spark, table):
return spark.sql(f"SELECT id FROM {table} ORDER BY id").collect()

cpu_rows = with_cpu_session(lambda spark: read_ids(spark, cpu_table), conf=confs)
gpu_rows = with_cpu_session(lambda spark: read_ids(spark, gpu_table), conf=confs)
assert_equal(cpu_rows, gpu_rows)
assert [row.id for row in gpu_rows] == list(range(10, 20))


@allow_non_gpu(*delta_meta_allow)
@delta_lake
@ignore_order(local=True)
Expand Down
Loading