diff --git a/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/BufferSendState.scala b/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/BufferSendState.scala index 0a7942bd581..f4190b0e140 100644 --- a/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/BufferSendState.scala +++ b/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/BufferSendState.scala @@ -1,5 +1,5 @@ /* - * Copyright (c) 2020-2024, NVIDIA CORPORATION. + * Copyright (c) 2020-2026, NVIDIA CORPORATION. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -118,6 +118,35 @@ class BufferSendState( private[this] var acquiredBuffs: Seq[RangeBuffer] = Seq.empty + // A retry window belongs to this transfer request. Its first materialization OOM starts an + // episode, and only a successful preparation by this state resets the episode. Results from + // other BufferSendState instances that happen to share a server batch do not affect it. + private[this] var oomRetryStartNanos: Option[Long] = None + private[this] var oomRetryAttempts: Int = 0 + + private[shuffle] def recordOomAndGetRetryAttempt( + nowNanos: Long, + timeoutNanos: Long): Option[Int] = + synchronized { + val startNanos = oomRetryStartNanos.getOrElse { + oomRetryStartNanos = Some(nowNanos) + nowNanos + } + if (nowNanos - startNanos < timeoutNanos) { + oomRetryAttempts += 1 + Some(oomRetryAttempts) + } else { + None + } + } + + private[shuffle] def resetOomRetryWindow(): Unit = synchronized { + if (oomRetryStartNanos.isDefined) { + oomRetryStartNanos = None + oomRetryAttempts = 0 + } + } + def getRequestTransaction: Transaction = synchronized { transaction } @@ -182,7 +211,15 @@ class BufferSendState( // using `releaseAcquiredToCatalog` //these are closed later, after we synchronize streams val spillable = blockRange.block.bufferHandle.spillable - val buff = spillable.materialize() + val buff = try { + spillable.materialize() + } catch { + case oom: OutOfMemoryError => + throw new RapidsShuffleSendPrepareException( + s"Memory exhausted while materializing a shuffle buffer for executor " + + s"${peerExecutorId} and header " + + s"${TransportUtils.toHex(peerBufferReceiveHeader)}: ${oom.toString}", oom) + } buff match { case _: DeviceMemoryBuffer => deviceBuffs += blockRange.rangeSize() @@ -214,6 +251,8 @@ class BufferSendState( } needsCleanup = false } catch { + case ex: RapidsShuffleSendPrepareException => + throw ex case ex: Exception => throw new RapidsShuffleSendPrepareException( s"Error while copying to bounce buffer for executor ${peerExecutorId} and " + @@ -244,6 +283,8 @@ class BufferSendState( logDebug(s"Sending ${buffsToSend} for transfer request, " + s" [peer_executor_id=${transaction.peerExecutorId()}]") + // Preparing this state's next send ends its continuous materialization-OOM episode. + resetOomRetryWindow() buffsToSend } diff --git a/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServer.scala b/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServer.scala index 8d7817da595..4604a89dd80 100644 --- a/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServer.scala +++ b/sql-plugin/src/main/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServer.scala @@ -1,5 +1,5 @@ /* - * Copyright (c) 2020-2025, NVIDIA CORPORATION. + * Copyright (c) 2020-2026, NVIDIA CORPORATION. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -16,17 +16,19 @@ package com.nvidia.spark.rapids.shuffle -import java.util.concurrent.{ConcurrentLinkedQueue, Executor} +import java.util.concurrent.{ConcurrentLinkedQueue, Executor, TimeUnit} import scala.collection.mutable.ArrayBuffer import ai.rapids.cudf.{Cuda, MemoryBuffer} import com.nvidia.spark.rapids.{NvtxRegistry, RapidsConf, RapidsShuffleHandle, ShuffleMetadata} import com.nvidia.spark.rapids.Arm.{closeOnExcept, withResource} +import com.nvidia.spark.rapids.RapidsPluginImplicits._ import com.nvidia.spark.rapids.format.TableMeta import org.apache.spark.internal.Logging import org.apache.spark.shuffle.rapids.RapidsShuffleSendPrepareException +import org.apache.spark.sql.rapids.GpuShuffleEnv import org.apache.spark.sql.rapids.execution.TrampolineUtil import org.apache.spark.storage.{BlockManagerId, ShuffleBlockBatchId} @@ -91,6 +93,33 @@ class RapidsShuffleServer(transport: RapidsShuffleTransport, */ private[this] var started = true + private[shuffle] def currentTimeNanos(): Long = System.nanoTime() + + private[shuffle] def oomRetryTimeoutNanos: Long = + TimeUnit.SECONDS.toNanos(GpuShuffleEnv.shuffleFetchTimeoutSeconds) + + private[shuffle] def oomRetryBackoffMillis(retryAttempt: Int): Long = { + val shift = math.min(retryAttempt - 1, 4) + math.min(10L << shift, 100L) + } + + private[shuffle] def waitBeforeOomRetry(backoffMillis: Long): Unit = + Thread.sleep(backoffMillis) + + private def stopOomRetries( + states: Seq[BufferSendState], + errors: Seq[Throwable], + reason: String): Unit = { + val failure = new IllegalStateException( + s"Unable to prepare shuffle sends. $reason These sends will not be retried.") + errors.foreach(failure.addSuppressed) + states.foreach(_.safeClose(failure)) + logError(failure.getMessage, failure) + bssExec.synchronized { + bssExec.notifyAll() + } + } + private object ShuffleServerOps { /** * When a transfer request is received during a callback, the handle code is offloaded via this @@ -337,7 +366,7 @@ class RapidsShuffleServer(transport: RapidsShuffleTransport, case ex: RapidsShuffleSendPrepareException => // We failed to prepare the send (copy to bounce buffer), and got an exception. // Put the `bufferSendState` back in the continue queue, so it can be retried. - // If no `BufferSendState` could be handled without error, nothing is retried. + // If no `BufferSendState` could be handled, retry only transient OOM failures. // TODO: we should respond with a failure to the client. // Please see: https://github.com/NVIDIA/spark-rapids/issues/3040 if (toTryAgain == null) { @@ -352,23 +381,73 @@ class RapidsShuffleServer(transport: RapidsShuffleTransport, if (toTryAgain != null) { // we failed at least 1 time to copy to the bounce buffer - if (bssBuffers.isEmpty) { - // we were not able to handle anything, error out. + val failures = toTryAgain.toSeq.zip(supressedErrors.toSeq) + def isMaterializationOom(error: Throwable): Boolean = error match { + case ex: RapidsShuffleSendPrepareException => + ex.getCause.isInstanceOf[OutOfMemoryError] + case _ => false + } + + // Preserve the existing fail-fast behavior when an entire batch fails and at least one + // preparation failure is not a materialization OOM. + if (bssBuffers.isEmpty && !failures.forall(f => isMaterializationOom(f._2))) { val ise = new IllegalStateException("Unable to prepare any sends. " + "This issue can occur when requesting too many shuffle blocks. " + "The sends will not be retried.") supressedErrors.foreach(ise.addSuppressed) throw ise + } + + val oomFailures = failures.filter(f => isMaterializationOom(f._2)) + val oomDecisions = if (oomFailures.nonEmpty) { + val nowNanos = currentTimeNanos() + val timeoutNanos = oomRetryTimeoutNanos + oomFailures.map { case (state, error) => + (state, error, state.recordOomAndGetRetryAttempt(nowNanos, timeoutNanos)) + } } else { - // we at least handled 1 `BufferSendState`, lets continue to retry - logWarning(s"Unable to prepare ${toTryAgain.size} sends. " + + Seq.empty + } + val (retryableOom, expiredOom) = oomDecisions.partition(_._3.isDefined) + + if (expiredOom.nonEmpty) { + stopOomRetries( + expiredOom.map(_._1), + expiredOom.map(_._2), + "The per-send OOM retry window derived from spark.network.timeout expired.") + } + + val retryableOomStates = retryableOom.map(_._1).toSet + val statesToRetry = failures.collect { + case (state, error) + if !isMaterializationOom(error) || retryableOomStates.contains(state) => + state + } + + if (bssBuffers.isEmpty && retryableOom.nonEmpty) { + val retryAttempt = retryableOom.flatMap(_._3).max + val backoffMillis = oomRetryBackoffMillis(retryAttempt) + val message = s"Memory exhausted while preparing ${retryableOom.size} sends. " + + s"Retry attempt $retryAttempt will start after $backoffMillis ms." + if (retryAttempt == 1 || retryAttempt % 100 == 0) { + logWarning(message) + } else { + logDebug(message) + } + // Avoid a hot loop because each failed materialization can repeat spill sweeps and + // device synchronization before it reports the allocation failure. + waitBeforeOomRetry(backoffMillis) + } else if (statesToRetry.nonEmpty) { + logWarning(s"Unable to prepare ${statesToRetry.size} sends. " + "This issue can occur when requesting many shuffle blocks. " + "The sends will be retried.") } - // If we are still able to handle at least one `BufferSendState`, add any - // others that also failed due back to the queue. - addToContinueQueue(toTryAgain.toSeq) + // Each BufferSendState owns its OOM retry window. Incidental co-batching cannot expire a + // fresh state or reset a failing state when an unrelated state makes progress. + if (statesToRetry.nonEmpty) { + addToContinueQueue(statesToRetry) + } } serverStream.sync() diff --git a/tests/src/test/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServerSuite.scala b/tests/src/test/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServerSuite.scala index 8d7415fba04..828eadb680f 100644 --- a/tests/src/test/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServerSuite.scala +++ b/tests/src/test/scala/com/nvidia/spark/rapids/shuffle/RapidsShuffleServerSuite.scala @@ -1,5 +1,5 @@ /* - * Copyright (c) 2020-2024, NVIDIA CORPORATION. + * Copyright (c) 2020-2026, NVIDIA CORPORATION. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -323,6 +323,235 @@ class RapidsShuffleServerSuite extends RapidsShuffleTestHelper { } } + test("when GPU OOM prevents all sends, stop the request at the shuffle fetch timeout") { + val mockSendBuffer = mock[SendBounceBuffers] + val mockDeviceBounceBuffer = mock[BounceBuffer] + val mockDeviceMemoryBuffer = mock[DeviceMemoryBuffer] + val mockServerConnection = mock[ServerConnection] + when(mockDeviceBounceBuffer.buffer).thenReturn(mockDeviceMemoryBuffer) + when(mockSendBuffer.bounceBufferSize).thenReturn(1024) + when(mockSendBuffer.hostBounceBuffer).thenReturn(None) + when(mockSendBuffer.deviceBounceBuffer).thenReturn(mockDeviceBounceBuffer) + + val tr = ShuffleMetadata.buildTransferRequest(0, Seq(1, 2)) + when(mockTransaction.releaseMessage()).thenReturn( + new MetadataTransportBuffer(new RefCountedDirectByteBuffer(tr))) + + val mockRequestHandler = mock[RapidsShuffleRequestHandler] + val bb = ByteBuffer.allocateDirect(123) + withResource(new RefCountedDirectByteBuffer(bb)) { _ => + val tableMeta = MetaUtils.buildTableMeta(1, 456, bb, 100) + val mockHandle = mock[SpillableDeviceBufferHandle] + val mockHandleThatThrows = mock[SpillableDeviceBufferHandle] + val mockMaterialized = mock[DeviceMemoryBuffer] + when(mockHandle.sizeInBytes).thenReturn(tableMeta.bufferMeta().size()) + when(mockHandle.materialize()).thenReturn(mockMaterialized) + when(mockHandleThatThrows.sizeInBytes).thenReturn(tableMeta.bufferMeta().size()) + val oom = new OutOfMemoryError("GPU allocation failed in test") + when(mockHandleThatThrows.materialize()).thenThrow(oom) + val rapidsBuffer = RapidsShuffleHandle(mockHandle, tableMeta) + val rapidsBufferThatThrows = RapidsShuffleHandle(mockHandleThatThrows, tableMeta) + when(mockRequestHandler.getShuffleHandle(ArgumentMatchers.eq(1))) + .thenReturn(rapidsBuffer) + when(mockRequestHandler.getShuffleHandle(ArgumentMatchers.eq(2))) + .thenReturn(rapidsBufferThatThrows) + + var nowNanos = 0L + val server = spy(new RapidsShuffleServer( + mockTransport, + mockServerConnection, + RapidsShuffleTestHelper.makeMockBlockManager("1", "foo"), + mockRequestHandler, + mockExecutor, + mockBssExecutor, + mockConf) { + override private[shuffle] def currentTimeNanos(): Long = nowNanos + override private[shuffle] def oomRetryTimeoutNanos: Long = 1L + override private[shuffle] def waitBeforeOomRetry(backoffMillis: Long): Unit = {} + }) + + val bss = new BufferSendState(mockTransaction, mockSendBuffer, mockRequestHandler, null) + assertResult(10L)(server.oomRetryBackoffMillis(1)) + assertResult(20L)(server.oomRetryBackoffMillis(2)) + assertResult(100L)(server.oomRetryBackoffMillis(5)) + assertResult(100L)(server.oomRetryBackoffMillis(1000)) + server.doHandleTransferRequest(Seq(bss)) + + verify(server, times(1)).addToContinueQueue(Seq(bss)) + verify(server, times(1)).waitBeforeOomRetry(10L) + verify(mockSendBuffer, times(0)).close() + + nowNanos = 1L + server.doHandleTransferRequest(Seq(bss)) + verify(server, times(1)).addToContinueQueue(Seq(bss)) + verify(server, times(1)).waitBeforeOomRetry(10L) + verify(mockHandle, times(2)).materialize() + verify(mockHandleThatThrows, times(2)).materialize() + verify(mockMaterialized, times(2)).close() + verify(mockSendBuffer, times(1)).close() + } + } + + test("successful preparation resets the OOM retry window on the same send state") { + val sendBuffer = mock[SendBounceBuffers] + val deviceBounceBuffer = mock[BounceBuffer] + val deviceMemoryBuffer = mock[DeviceMemoryBuffer] + val bufferSlice = mock[DeviceMemoryBuffer] + when(deviceBounceBuffer.buffer).thenReturn(deviceMemoryBuffer) + when(deviceMemoryBuffer.getLength).thenReturn(1024L) + when(sendBuffer.bounceBufferSize).thenReturn(1024) + when(sendBuffer.hostBounceBuffer).thenReturn(None) + when(sendBuffer.deviceBounceBuffer).thenReturn(deviceBounceBuffer) + when(deviceMemoryBuffer.slice(ArgumentMatchers.anyLong(), ArgumentMatchers.anyLong())) + .thenReturn(bufferSlice) + + val request = ShuffleMetadata.buildTransferRequest(0, Seq(1)) + when(mockTransaction.releaseMessage()).thenReturn( + new MetadataTransportBuffer(new RefCountedDirectByteBuffer(request))) + + val requestHandler = mock[RapidsShuffleRequestHandler] + val metadataBuffer = ByteBuffer.allocateDirect(123) + withResource(new RefCountedDirectByteBuffer(metadataBuffer)) { _ => + val tableMeta = MetaUtils.buildTableMeta(1, 456, metadataBuffer, 100) + val handle = mock[SpillableDeviceBufferHandle] + val materialized = mock[DeviceMemoryBuffer] + when(handle.sizeInBytes).thenReturn(tableMeta.bufferMeta().size()) + when(handle.materialize()).thenReturn(materialized) + when(requestHandler.getShuffleHandle(ArgumentMatchers.eq(1))) + .thenReturn(RapidsShuffleHandle(handle, tableMeta)) + + withResource(new BufferSendState( + mockTransaction, sendBuffer, requestHandler, null)) { bss => + assertResult(Some(1))(bss.recordOomAndGetRetryAttempt(0L, 10L)) + withResource(bss.getBufferToSend()) { _ => } + + // This is a new OOM episode. If resetOomRetryWindow were a no-op or its successful + // preparation call site were removed, the old window would expire at this timestamp. + assertResult(Some(1))(bss.recordOomAndGetRetryAttempt(10L, 10L)) + } + } + } + + test("an expired OOM window does not close a co-batched send with a fresh window") { + def makeSendBuffer(name: String): SendBounceBuffers = { + val sb = mock[SendBounceBuffers](name) + val dbb = mock[BounceBuffer] + when(dbb.buffer).thenReturn(mock[DeviceMemoryBuffer]) + when(sb.bounceBufferSize).thenReturn(1024) + when(sb.hostBounceBuffer).thenReturn(None) + when(sb.deviceBounceBuffer).thenReturn(dbb) + sb + } + val sendBufferA = makeSendBuffer("sendBufferA") + val sendBufferB = makeSendBuffer("sendBufferB") + + val requestHandler = mock[RapidsShuffleRequestHandler] + val metadataBuffer = ByteBuffer.allocateDirect(123) + withResource(new RefCountedDirectByteBuffer(metadataBuffer)) { _ => + val tableMeta = MetaUtils.buildTableMeta(1, 456, metadataBuffer, 100) + val oomHandle = mock[SpillableDeviceBufferHandle] + when(oomHandle.sizeInBytes).thenReturn(tableMeta.bufferMeta().size()) + when(oomHandle.materialize()) + .thenThrow(new OutOfMemoryError("GPU allocation failed in test")) + when(requestHandler.getShuffleHandle(ArgumentMatchers.eq(1))) + .thenReturn(RapidsShuffleHandle(oomHandle, tableMeta)) + + def makeTx(peerExecutorId: Long): Transaction = { + val tx = mock[Transaction] + when(tx.peerExecutorId()).thenReturn(peerExecutorId) + when(tx.releaseMessage()).thenReturn(new MetadataTransportBuffer( + new RefCountedDirectByteBuffer(ShuffleMetadata.buildTransferRequest(0, Seq(1))))) + tx + } + + var nowNanos = 0L + val server = spy(new RapidsShuffleServer( + mockTransport, + mock[ServerConnection], + RapidsShuffleTestHelper.makeMockBlockManager("1", "foo"), + requestHandler, + mockExecutor, + mockBssExecutor, + mockConf) { + override private[shuffle] def currentTimeNanos(): Long = nowNanos + override private[shuffle] def oomRetryTimeoutNanos: Long = 10L + override private[shuffle] def waitBeforeOomRetry(backoffMillis: Long): Unit = {} + }) + + val bssA = new BufferSendState(makeTx(1L), sendBufferA, requestHandler, null) + server.doHandleTransferRequest(Seq(bssA)) + verify(server, times(1)).addToContinueQueue(Seq(bssA)) + + nowNanos = 10L + val bssB = new BufferSendState(makeTx(2L), sendBufferB, requestHandler, null) + server.doHandleTransferRequest(Seq(bssA, bssB)) + + verify(sendBufferA, times(1)).close() + verify(sendBufferB, times(0)).close() + verify(server, times(1)).addToContinueQueue(Seq(bssB)) + } + } + + test("progress by a co-batched send does not reset another send state's OOM window") { + val sendBuffer = mock[SendBounceBuffers] + val deviceBounceBuffer = mock[BounceBuffer] + when(deviceBounceBuffer.buffer).thenReturn(mock[DeviceMemoryBuffer]) + when(sendBuffer.bounceBufferSize).thenReturn(1024) + when(sendBuffer.hostBounceBuffer).thenReturn(None) + when(sendBuffer.deviceBounceBuffer).thenReturn(deviceBounceBuffer) + + val request = ShuffleMetadata.buildTransferRequest(0, Seq(1)) + when(mockTransaction.releaseMessage()).thenReturn( + new MetadataTransportBuffer(new RefCountedDirectByteBuffer(request))) + + val requestHandler = mock[RapidsShuffleRequestHandler] + val metadataBuffer = ByteBuffer.allocateDirect(123) + withResource(new RefCountedDirectByteBuffer(metadataBuffer)) { _ => + val tableMeta = MetaUtils.buildTableMeta(1, 456, metadataBuffer, 100) + val oomHandle = mock[SpillableDeviceBufferHandle] + when(oomHandle.sizeInBytes).thenReturn(tableMeta.bufferMeta().size()) + when(oomHandle.materialize()) + .thenThrow(new OutOfMemoryError("GPU allocation failed in test")) + when(requestHandler.getShuffleHandle(ArgumentMatchers.eq(1))) + .thenReturn(RapidsShuffleHandle(oomHandle, tableMeta)) + + val successfulState = mock[BufferSendState] + val successfulBuffer = mock[MemoryBuffer] + when(successfulState.hasMoreSends).thenReturn(true) + when(successfulState.getBufferToSend()).thenReturn(successfulBuffer) + + val serverConnection = mock[ServerConnection] + when(serverConnection.send( + any(), any(), any(), any[MemoryBuffer](), any[TransactionCallback]())) + .thenReturn(mock[Transaction]) + + var nowNanos = 0L + val server = spy(new RapidsShuffleServer( + mockTransport, + serverConnection, + RapidsShuffleTestHelper.makeMockBlockManager("1", "foo"), + requestHandler, + mockExecutor, + mockBssExecutor, + mockConf) { + override private[shuffle] def currentTimeNanos(): Long = nowNanos + override private[shuffle] def oomRetryTimeoutNanos: Long = 10L + override private[shuffle] def waitBeforeOomRetry(backoffMillis: Long): Unit = {} + }) + + val oomState = new BufferSendState( + mockTransaction, sendBuffer, requestHandler, null) + server.doHandleTransferRequest(Seq(oomState)) + verify(server, times(1)).addToContinueQueue(Seq(oomState)) + + nowNanos = 10L + server.doHandleTransferRequest(Seq(oomState, successfulState)) + + verify(sendBuffer, times(1)).close() + verify(server, times(1)).addToContinueQueue(any()) + } + } + test("when we fail to prepare a send, re-queue the request if anything can be handled") { val mockSendBuffer = mock[SendBounceBuffers] val mockDeviceBounceBuffer = mock[BounceBuffer]