From ab810ad9463b2fb23590befea7a77717073213c3 Mon Sep 17 00:00:00 2001 From: Jihoon Son Date: Tue, 8 Sep 2026 15:55:22 -0700 Subject: [PATCH 1/4] Fix Delta v1 writer detection for spark 4.0+ Signed-off-by: Jihoon Son --- .../GpuCreateDeltaTableCommandBase.scala | 23 ++----- .../sql/delta/rapids/DeltaRuntimeShim.scala | 16 ++++- .../rapids/delta40x/Delta40xRuntimeShim.scala | 10 +++ .../rapids/delta41x/Delta41xRuntimeShim.scala | 9 +++ .../src/main/python/delta_lake_write_test.py | 65 +++++++++++++++++++ 5 files changed, 104 insertions(+), 19 deletions(-) diff --git a/delta-lake/common/src/main/delta-33x-41x/scala/org/apache/spark/sql/delta/rapids/GpuCreateDeltaTableCommandBase.scala b/delta-lake/common/src/main/delta-33x-41x/scala/org/apache/spark/sql/delta/rapids/GpuCreateDeltaTableCommandBase.scala index e02c12b7298..7c713e2f82b 100644 --- a/delta-lake/common/src/main/delta-33x-41x/scala/org/apache/spark/sql/delta/rapids/GpuCreateDeltaTableCommandBase.scala +++ b/delta-lake/common/src/main/delta-33x-41x/scala/org/apache/spark/sql/delta/rapids/GpuCreateDeltaTableCommandBase.scala @@ -28,7 +28,7 @@ import org.apache.hadoop.conf.Configuration import org.apache.hadoop.fs.{FileSystem, Path} import org.apache.spark.SparkContext -import org.apache.spark.sql.{DataFrame, DataFrameWriter, Row, SaveMode, SparkSession} +import org.apache.spark.sql.{DataFrame, Row, SaveMode, SparkSession} import org.apache.spark.sql.catalyst.catalog.{CatalogTable, CatalogTableType} import org.apache.spark.sql.catalyst.expressions.Attribute import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan @@ -297,6 +297,8 @@ abstract class GpuCreateDeltaTableCommandBase( tableWithLocation: CatalogTable): Unit = { val isManagedTable = tableWithLocation.tableType == CatalogTableType.MANAGED val options = new DeltaOptions(table.storage.properties, sparkSession.sessionState.conf) + val isV1WriterSaveAsTableOverwrite = + DeltaRuntimeShim.isV1WriterSaveAsTableOverwrite(options, mode) // Execute write command for `deltaWriter` by // - replacing the metadata new target table for DataFrameWriterV2 writer if it is a @@ -309,7 +311,7 @@ abstract class GpuCreateDeltaTableCommandBase( schema: StructType): (TaggedCommitData[Action], DeltaOperations.Operation) = { // In the V2 Writer, methods like "replace" and "createOrReplace" implicitly mean that // the metadata should be changed. This wasn't the behavior for DataFrameWriterV1. - if (!isV1Writer) { + if (!isV1WriterSaveAsTableOverwrite) { replaceMetadataIfNecessary( txn, tableWithLocation, @@ -327,13 +329,13 @@ abstract class GpuCreateDeltaTableCommandBase( // saveAsTable() command uses this same code path and is marked as a V1 writer. // We do not want saveAsTable() to be treated as a REPLACE command wrt dynamic partition // overwrite. - isTableReplace = isReplace && !isV1Writer + isTableReplace = isReplace && !isV1WriterSaveAsTableOverwrite ) // Metadata updates for creating table (with any writer) and replacing table // (only with V1 writer) will be handled inside WriteIntoDelta. // For createOrReplace operation, metadata updates are handled here if the table already // exists (replacing table), otherwise it is handled inside WriteIntoDelta (creating table). - if (!isV1Writer && isReplace && txn.readVersion > -1L) { + if (!isV1WriterSaveAsTableOverwrite && isReplace && txn.readVersion > -1L) { val newDomainMetadata = Seq.empty[DomainMetadata] ++ ClusteredTableUtils.getDomainMetadataFromTransaction( ClusteredTableUtils.getClusterBySpecOptional(table), txn) @@ -807,19 +809,6 @@ abstract class GpuCreateDeltaTableCommandBase( } } - /** - * Horrible hack to differentiate between DataFrameWriterV1 and V2 so that we can decide - * what to do with table metadata. In DataFrameWriterV1, mode("overwrite").saveAsTable, - * behaves as a CreateOrReplace table, but we have asked for "overwriteSchema" as an - * explicit option to overwrite partitioning or schema information. With DataFrameWriterV2, - * the behavior asked for by the user is clearer: .createOrReplace(), which means that we - * should overwrite schema and/or partitioning. Therefore we have this hack. - */ - private def isV1Writer: Boolean = { - Thread.currentThread().getStackTrace.exists(_.toString.contains( - classOf[DataFrameWriter[_]].getCanonicalName + ".")) - } - /** Returns true if the current operation could be replacing a table. */ private def isReplace: Boolean = { operation == TableCreationModes.CreateOrReplace || diff --git a/delta-lake/common/src/main/delta-io/scala/org/apache/spark/sql/delta/rapids/DeltaRuntimeShim.scala b/delta-lake/common/src/main/delta-io/scala/org/apache/spark/sql/delta/rapids/DeltaRuntimeShim.scala index 039368e7356..f803bb24f9e 100644 --- a/delta-lake/common/src/main/delta-io/scala/org/apache/spark/sql/delta/rapids/DeltaRuntimeShim.scala +++ b/delta-lake/common/src/main/delta-io/scala/org/apache/spark/sql/delta/rapids/DeltaRuntimeShim.scala @@ -21,10 +21,10 @@ import scala.util.Try import com.nvidia.spark.rapids.{RapidsConf, ShimLoader, ShimReflectionUtils, VersionUtils} import com.nvidia.spark.rapids.delta.{DeltaConfigChecker, DeltaProvider} -import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.{DataFrameWriter, SaveMode, SparkSession} import org.apache.spark.sql.catalyst.catalog.CatalogTable import org.apache.spark.sql.connector.catalog.StagingTableCatalog -import org.apache.spark.sql.delta.{DeltaLog, DeltaUDF, Snapshot} +import org.apache.spark.sql.delta.{DeltaLog, DeltaOptions, DeltaUDF, Snapshot} import org.apache.spark.sql.delta.catalog.DeltaCatalog import org.apache.spark.sql.execution.datasources.FileFormat import org.apache.spark.sql.expressions.UserDefinedFunction @@ -48,6 +48,15 @@ trait DeltaRuntimeShim { def getTightBoundColumnOnFileInitDisabled(spark: SparkSession): Boolean def getGpuDeltaCatalog(cpuCatalog: DeltaCatalog, rapidsConf: RapidsConf): StagingTableCatalog + + /** + * Detect a DataFrameWriter V1 mode("overwrite").saveAsTable operation so it retains the + * existing table metadata. Delta versions before 4.1 require stack-trace inspection. + */ + def isV1WriterSaveAsTableOverwrite(options: DeltaOptions, mode: SaveMode): Boolean = { + mode == SaveMode.Overwrite && Thread.currentThread().getStackTrace.exists(_.toString.contains( + classOf[DataFrameWriter[_]].getCanonicalName + ".")) + } } object DeltaRuntimeShim { @@ -113,6 +122,9 @@ object DeltaRuntimeShim { def getTightBoundColumnOnFileInitDisabled(spark: SparkSession): Boolean = shimInstance.getTightBoundColumnOnFileInitDisabled(spark) + def isV1WriterSaveAsTableOverwrite(options: DeltaOptions, mode: SaveMode): Boolean = + shimInstance.isV1WriterSaveAsTableOverwrite(options, mode) + def getGpuDeltaCatalog(cpuCatalog: DeltaCatalog, rapidsConf: RapidsConf): StagingTableCatalog = { shimInstance.getGpuDeltaCatalog(cpuCatalog, rapidsConf) } diff --git a/delta-lake/delta-40x/src/main/scala/org/apache/spark/sql/delta/rapids/delta40x/Delta40xRuntimeShim.scala b/delta-lake/delta-40x/src/main/scala/org/apache/spark/sql/delta/rapids/delta40x/Delta40xRuntimeShim.scala index f34e0e9137f..effcf598ec0 100644 --- a/delta-lake/delta-40x/src/main/scala/org/apache/spark/sql/delta/rapids/delta40x/Delta40xRuntimeShim.scala +++ b/delta-lake/delta-40x/src/main/scala/org/apache/spark/sql/delta/rapids/delta40x/Delta40xRuntimeShim.scala @@ -21,7 +21,10 @@ import com.nvidia.spark.rapids.delta.DeltaProvider import com.nvidia.spark.rapids.delta.delta40x.Delta40xProvider import com.nvidia.spark.rapids.delta.delta40x.GpuDeltaCatalog +import org.apache.spark.sql.SaveMode +import org.apache.spark.sql.classic.DataFrameWriter import org.apache.spark.sql.connector.catalog.StagingTableCatalog +import org.apache.spark.sql.delta.DeltaOptions import org.apache.spark.sql.delta.catalog.DeltaCatalog import org.apache.spark.sql.delta.rapids.{DeltaRuntimeShimBase, GpuOptimisticTransaction, GpuOptimisticTransactionBase, StartTransactionArg} @@ -35,6 +38,13 @@ class Delta40xRuntimeShim extends DeltaRuntimeShimBase { override def getDeltaProvider: DeltaProvider = Delta40xProvider + override def isV1WriterSaveAsTableOverwrite( + options: DeltaOptions, + mode: SaveMode): Boolean = { + mode == SaveMode.Overwrite && Thread.currentThread().getStackTrace.exists(_.toString.contains( + classOf[DataFrameWriter[_]].getCanonicalName + ".")) + } + override def getGpuDeltaCatalog( cpuCatalog: DeltaCatalog, rapidsConf: RapidsConf): StagingTableCatalog = { diff --git a/delta-lake/delta-41x/src/main/scala/org/apache/spark/sql/delta/rapids/delta41x/Delta41xRuntimeShim.scala b/delta-lake/delta-41x/src/main/scala/org/apache/spark/sql/delta/rapids/delta41x/Delta41xRuntimeShim.scala index ba0c00c6b10..247f2e08636 100644 --- a/delta-lake/delta-41x/src/main/scala/org/apache/spark/sql/delta/rapids/delta41x/Delta41xRuntimeShim.scala +++ b/delta-lake/delta-41x/src/main/scala/org/apache/spark/sql/delta/rapids/delta41x/Delta41xRuntimeShim.scala @@ -21,8 +21,11 @@ import com.nvidia.spark.rapids.delta.DeltaProvider import com.nvidia.spark.rapids.delta.delta41x.Delta41xProvider import com.nvidia.spark.rapids.delta.delta41x.GpuDeltaCatalog +import org.apache.spark.sql.SaveMode import org.apache.spark.sql.connector.catalog.StagingTableCatalog +import org.apache.spark.sql.delta.DeltaOptions import org.apache.spark.sql.delta.catalog.DeltaCatalog +import org.apache.spark.sql.delta.commands.CreateDeltaTableLikeShims import org.apache.spark.sql.delta.rapids.{ DeltaRuntimeShimBase, GpuOptimisticTransaction, @@ -34,6 +37,12 @@ class Delta41xRuntimeShim extends DeltaRuntimeShimBase { override def getDeltaProvider: DeltaProvider = Delta41xProvider + override def isV1WriterSaveAsTableOverwrite( + options: DeltaOptions, + mode: SaveMode): Boolean = { + CreateDeltaTableLikeShims.isV1WriterSaveAsTableOverwrite(options, mode) + } + override def getGpuDeltaCatalog( cpuCatalog: DeltaCatalog, rapidsConf: RapidsConf): StagingTableCatalog = { diff --git a/integration_tests/src/main/python/delta_lake_write_test.py b/integration_tests/src/main/python/delta_lake_write_test.py index 78fd9d448cf..48d20388d27 100644 --- a/integration_tests/src/main/python/delta_lake_write_test.py +++ b/integration_tests/src/main/python/delta_lake_write_test.py @@ -868,6 +868,71 @@ def read_ids(spark, table): assert [row.id for row in gpu_rows] == list(range(10, 20)) +@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.x saveAsTable overwrite regression") +@pytest.mark.xfail(is_databricks_runtime(), + reason="https://github.com/NVIDIA/spark-rapids/issues/11169") +def test_delta_replace_where_save_as_table_preserves_partitioning(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): + initial_values = ", ".join( + f"({record_id}L, '{region}', {record_id}.0D)" + for record_id, region in enumerate( + ["NA", "EMEA", "APAC", "LATAM", "NA", "EMEA", "APAC", "LATAM"])) + for table in [cpu_table, gpu_table]: + spark.sql( + f"CREATE TABLE {table} (record_id BIGINT, region STRING, amount DOUBLE) " + f"USING DELTA PARTITIONED BY (region)") + spark.sql(f"INSERT INTO {table} VALUES {initial_values}") + + def replace_na_partition(spark, table): + replacement = spark.sql( + "SELECT 100L AS record_id, 'NA' AS region, 100.0D AS amount") + (replacement.write.format("delta").mode("overwrite") + .option("replaceWhere", "region = 'NA'") + .saveAsTable(table)) + + with_cpu_session(create_initial_tables, conf=confs) + with_cpu_session(lambda spark: replace_na_partition(spark, cpu_table), conf=confs) + + callback = spark_jvm().org.apache.spark.sql.rapids.ExecutionPlanCaptureCallback + callback.startCapture() + try: + with_gpu_session(lambda spark: replace_na_partition(spark, gpu_table), conf=confs) + plans = callback.getResultsWithTimeout(10000) + assert any(callback.contains(plan, "GpuOverwriteByExpressionExecV1") + for plan in plans), "GpuOverwriteByExpressionExecV1 was not executed" + finally: + callback.endCapture() + + def partition_counts(spark, table): + return spark.sql( + f"SELECT region, COUNT(*) AS count FROM {table} " + f"GROUP BY region ORDER BY region").collect() + + cpu_counts = with_cpu_session(lambda spark: partition_counts(spark, cpu_table), conf=confs) + gpu_counts = with_cpu_session(lambda spark: partition_counts(spark, gpu_table), conf=confs) + assert_equal(cpu_counts, gpu_counts) + assert [(row.region, row["count"]) for row in gpu_counts] == [ + ("APAC", 2), ("EMEA", 2), ("LATAM", 2), ("NA", 1)] + + def partition_columns(spark, table): + return spark.sql(f"DESCRIBE DETAIL {table}").select("partitionColumns").head()[0] + + cpu_partition_columns = with_cpu_session( + lambda spark: partition_columns(spark, cpu_table), conf=confs) + gpu_partition_columns = with_cpu_session( + lambda spark: partition_columns(spark, gpu_table), conf=confs) + assert cpu_partition_columns == ["region"] + assert gpu_partition_columns == cpu_partition_columns + + @allow_non_gpu(*delta_meta_allow) @delta_lake @ignore_order(local=True) From cf91f7f0e3949d52a6524682d22193ec03d54228 Mon Sep 17 00:00:00 2001 From: Jihoon Son Date: Tue, 8 Sep 2026 21:11:21 -0700 Subject: [PATCH 2/4] remove the skipif for oss --- integration_tests/src/main/python/delta_lake_write_test.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/integration_tests/src/main/python/delta_lake_write_test.py b/integration_tests/src/main/python/delta_lake_write_test.py index 48d20388d27..26ccd40d10f 100644 --- a/integration_tests/src/main/python/delta_lake_write_test.py +++ b/integration_tests/src/main/python/delta_lake_write_test.py @@ -871,8 +871,6 @@ def read_ids(spark, table): @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.x saveAsTable overwrite regression") @pytest.mark.xfail(is_databricks_runtime(), reason="https://github.com/NVIDIA/spark-rapids/issues/11169") def test_delta_replace_where_save_as_table_preserves_partitioning(spark_tmp_table_factory): @@ -906,8 +904,8 @@ def replace_na_partition(spark, table): try: with_gpu_session(lambda spark: replace_na_partition(spark, gpu_table), conf=confs) plans = callback.getResultsWithTimeout(10000) - assert any(callback.contains(plan, "GpuOverwriteByExpressionExecV1") - for plan in plans), "GpuOverwriteByExpressionExecV1 was not executed" + assert any(callback.contains(plan, "GpuAtomicReplaceTableAsSelectExec") + for plan in plans), "GpuAtomicReplaceTableAsSelectExec was not executed" finally: callback.endCapture() From 9d11ed4317e82eb6e6e936965785ce0338d9e3e2 Mon Sep 17 00:00:00 2001 From: Jihoon Son Date: Tue, 8 Sep 2026 22:24:31 -0700 Subject: [PATCH 3/4] verify v1 path --- integration_tests/src/main/python/delta_lake_write_test.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/integration_tests/src/main/python/delta_lake_write_test.py b/integration_tests/src/main/python/delta_lake_write_test.py index 26ccd40d10f..cdf9a97ee63 100644 --- a/integration_tests/src/main/python/delta_lake_write_test.py +++ b/integration_tests/src/main/python/delta_lake_write_test.py @@ -906,6 +906,12 @@ def replace_na_partition(spark, table): plans = callback.getResultsWithTimeout(10000) assert any(callback.contains(plan, "GpuAtomicReplaceTableAsSelectExec") for plan in plans), "GpuAtomicReplaceTableAsSelectExec was not executed" + # The RTAS data write runs as a nested query execution: Spark 4.0+ issues it as + # OverwriteByExpression, earlier versions as AppendData. + v1_write_node = ("GpuOverwriteByExpressionExecV1" if is_spark_400_or_later() + else "GpuAppendDataExecV1") + assert any(callback.contains(plan, v1_write_node) + for plan in plans), f"{v1_write_node} was not executed" finally: callback.endCapture() From 51bb7a573582c13f04218f029e9b3cb84178b147 Mon Sep 17 00:00:00 2001 From: Jihoon Son Date: Fri, 11 Sep 2026 10:22:37 -0700 Subject: [PATCH 4/4] fix the bug for delta 4.2 --- .../sql/delta/rapids/delta42x/Delta42xRuntimeShim.scala | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/delta-lake/delta-42x/src/main/scala/org/apache/spark/sql/delta/rapids/delta42x/Delta42xRuntimeShim.scala b/delta-lake/delta-42x/src/main/scala/org/apache/spark/sql/delta/rapids/delta42x/Delta42xRuntimeShim.scala index 21213c8f932..2f1f94fcaf7 100644 --- a/delta-lake/delta-42x/src/main/scala/org/apache/spark/sql/delta/rapids/delta42x/Delta42xRuntimeShim.scala +++ b/delta-lake/delta-42x/src/main/scala/org/apache/spark/sql/delta/rapids/delta42x/Delta42xRuntimeShim.scala @@ -29,7 +29,7 @@ import org.apache.spark.sql.connector.catalog.StagingTableCatalog import org.apache.spark.sql.delta.{DeltaOperations, DeltaOptions} import org.apache.spark.sql.delta.actions.Metadata import org.apache.spark.sql.delta.catalog.DeltaCatalog -import org.apache.spark.sql.delta.commands.WriteIntoDelta +import org.apache.spark.sql.delta.commands.{CreateDeltaTableLikeShims, WriteIntoDelta} import org.apache.spark.sql.delta.hooks.GpuAutoCompact42x import org.apache.spark.sql.delta.rapids.{DeltaRuntimeShimBase, GpuDeltaLog, GpuOptimisticTransaction, GpuOptimisticTransactionBase, GpuWriteIntoDeltaLike, StartTransactionArg} @@ -40,6 +40,12 @@ class Delta42xRuntimeShim extends DeltaRuntimeShimBase { override def getDeltaProvider: DeltaProvider = Delta42xProvider + override def isV1WriterSaveAsTableOverwrite( + options: DeltaOptions, + mode: SaveMode): Boolean = { + CreateDeltaTableLikeShims.isV1WriterSaveAsTableOverwrite(options, mode) + } + override def getGpuDeltaCatalog( cpuCatalog: DeltaCatalog, rapidsConf: RapidsConf): StagingTableCatalog = {