@@ -1736,7 +1736,9 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
17361736 // NOTE: this can be null in the driver side.
17371737 protected lazy val env = SparkEnv .get
17381738 protected lazy val blockManager = env.blockManager
1739- protected lazy val shouldFallThroughOnEverything = {
1739+ // Stable reasons to always fall back to SortShuffleManager, evaluated once at
1740+ // first shuffle registration.
1741+ protected lazy val shouldAlwaysFallBack = {
17401742 val fallThroughReasons = new ListBuffer [String ]()
17411743 if (! rapidsConf.isMultiThreadedShuffleManagerMode) {
17421744 if (GpuShuffleEnv .isExternalShuffleEnabled) {
@@ -1749,19 +1751,28 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
17491751 if (rapidsConf.isSqlExplainOnlyEnabled) {
17501752 fallThroughReasons += " Plugin is in explain only mode"
17511753 }
1752- if (GpuShuffleEnv .isRowBasedChecksumEnabled) {
1753- fallThroughReasons += " Detected order-independent checksum enabled " +
1754- " (spark.sql.shuffle.orderIndependentChecksum.enabled or " +
1755- " enableFullRetryOnMismatch). " +
1756- " This Spark 4.1+ feature is not yet supported by Spark-Rapids."
1757- }
17581754 if (fallThroughReasons.nonEmpty) {
17591755 logWarning(s " Rapids Shuffle Plugin is falling back to SortShuffleManager " +
17601756 s " because: ${fallThroughReasons.mkString(" , " )}" )
17611757 }
17621758 fallThroughReasons.nonEmpty
17631759 }
17641760
1761+ private val rowBasedChecksumFallbackLogged = new AtomicBoolean (false )
1762+
1763+ private def shouldFallThroughForShuffle : Boolean = {
1764+ val rowBasedChecksumFallback = GpuShuffleEnv .isRowBasedChecksumEnabled
1765+ if (rowBasedChecksumFallback) {
1766+ if (rowBasedChecksumFallbackLogged.compareAndSet(false , true )) {
1767+ logWarning(" Rapids Shuffle Plugin is falling back to SortShuffleManager because: " +
1768+ " Detected order-independent checksum enabled " +
1769+ " (spark.sql.shuffle.orderIndependentChecksum.enabled or enableFullRetryOnMismatch). " +
1770+ " This Spark 4.1+ feature is not yet supported by Spark-Rapids." )
1771+ }
1772+ }
1773+ shouldAlwaysFallBack || rowBasedChecksumFallback
1774+ }
1775+
17651776 private lazy val localBlockManagerId = blockManager.blockManagerId
17661777
17671778 // Used to prevent stopping multiple times RAPIDS Shuffle Manager internals.
@@ -1776,7 +1787,7 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
17761787 " RapidsShuffleManager is configured" ))
17771788
17781789 protected lazy val resolver =
1779- if (shouldFallThroughOnEverything ) {
1790+ if (shouldAlwaysFallBack ) {
17801791 wrapped.shuffleBlockResolver
17811792 } else if (rapidsConf.isMultiThreadedShuffleManagerMode) {
17821793 // MULTITHREADED mode: use GpuShuffleBlockResolver
@@ -1848,11 +1859,13 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
18481859 val orig = wrapped.registerShuffle(shuffleId, dependency)
18491860
18501861 dependency match {
1851- case _ if shouldFallThroughOnEverything ||
1852- rapidsConf.isMultiThreadedShuffleManagerMode => orig
18531862 case gpuDependency : GpuShuffleDependency [K , V , C ] if gpuDependency.useGPUShuffle =>
1854- new GpuShuffleHandle (orig,
1855- dependency.asInstanceOf [GpuShuffleDependency [K , V , V ]])
1863+ val gpuDep = gpuDependency.asInstanceOf [GpuShuffleDependency [K , V , V ]]
1864+ gpuDep.checksumFallback = shouldFallThroughForShuffle
1865+ if (rapidsConf.isMultiThreadedShuffleManagerMode) orig
1866+ else new GpuShuffleHandle (orig, gpuDep)
1867+ case _ if shouldAlwaysFallBack ||
1868+ rapidsConf.isMultiThreadedShuffleManagerMode => orig
18561869 case _ => orig
18571870 }
18581871 }
@@ -1900,6 +1913,8 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
19001913 context : TaskContext ,
19011914 metricsReporter : ShuffleWriteMetricsReporter ): ShuffleWriter [K , V ] = {
19021915 handle match {
1916+ case gpu : GpuShuffleHandle [_, _] if gpu.dependency.checksumFallback =>
1917+ wrapped.getWriter(gpu.wrapped, mapId, context, metricsReporter)
19031918 case gpu : GpuShuffleHandle [_, _] =>
19041919 registerGpuShuffle(handle.shuffleId)
19051920 new RapidsCachingWriter (
@@ -1914,6 +1929,7 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
19141929 handle.dependency match {
19151930 case gpuDep : GpuShuffleDependency [_, _, _]
19161931 if gpuDep.useMultiThreadedShuffle &&
1932+ ! gpuDep.checksumFallback &&
19171933 rapidsConf.shuffleMultiThreadedWriterThreads > 0 =>
19181934 // use the threaded writer if the number of threads specified is 1 or above,
19191935 // with 0 threads we fallback to the Spark-provided writer.
@@ -1957,6 +1973,9 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
19571973 context : TaskContext ,
19581974 metrics : ShuffleReadMetricsReporter ): ShuffleReader [K , C ] = {
19591975 handle match {
1976+ case gpuHandle : GpuShuffleHandle [_, _] if gpuHandle.dependency.checksumFallback =>
1977+ ShuffleManagerShims .getReader(wrapped, gpuHandle.wrapped, startMapIndex, endMapIndex,
1978+ startPartition, endPartition, context, metrics)
19601979 case gpuHandle : GpuShuffleHandle [_, _] =>
19611980 logInfo(s " Asking map output tracker for dependency ${gpuHandle.dependency}, " +
19621981 s " map output sizes for: ${gpuHandle.shuffleId}, parts= $startPartition- $endPartition" )
@@ -1995,7 +2014,8 @@ class RapidsShuffleInternalManagerBase(conf: SparkConf, val isDriver: Boolean)
19952014 // would need to be made to deal with missing metrics, for example, for a regular
19962015 // Exchange node.
19972016 baseHandle.dependency match {
1998- case gpuDep : GpuShuffleDependency [K , C , C ] if gpuDep.useMultiThreadedShuffle =>
2017+ case gpuDep : GpuShuffleDependency [K , C , C ]
2018+ if gpuDep.useMultiThreadedShuffle && ! gpuDep.checksumFallback =>
19992019 // We want to use batch fetch in the non-push shuffle case. Spark
20002020 // checks for a config to see if batch fetch is enabled (this check), and
20012021 // it also checks when getting (potentially merged) map status from
0 commit comments