-
Notifications
You must be signed in to change notification settings - Fork 29.4k
[SPARK-59444][CORE] Retain execution-memory task registration while allocations wait #58747
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: master
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -56,6 +56,15 @@ private[memory] class ExecutionMemoryPool( | |
| @GuardedBy("lock") | ||
| private val memoryForTask = new mutable.HashMap[Long, Long]() | ||
|
|
||
| // Only acquisitions that actually wait need a retained zero-byte task entry. | ||
| // Lazily allocated under lock; the non-waiting admission path does no map work. | ||
| @GuardedBy("lock") | ||
| private var waitingAcquisitions: mutable.LongMap[Int] = null | ||
|
|
||
| private def hasWaitingAcquisition(taskAttemptId: Long): Boolean = { | ||
| waitingAcquisitions != null && waitingAcquisitions.contains(taskAttemptId) | ||
| } | ||
|
|
||
| override def memoryUsed: Long = lock.synchronized { | ||
| memoryForTask.values.sum | ||
| } | ||
|
|
@@ -110,42 +119,71 @@ private[memory] class ExecutionMemoryPool( | |
| // task would have more than 1 / numActiveTasks of the memory) or we have enough free | ||
| // memory to give it (we always let each task get at least 1 / (2 * numActiveTasks)). | ||
| // TODO: simplify this to limit each task to its own slot | ||
| while (true) { | ||
| val numActiveTasks = memoryForTask.keys.size | ||
| val curMem = memoryForTask(taskAttemptId) | ||
|
|
||
| // In every iteration of this loop, we should first try to reclaim any borrowed execution | ||
| // space from storage. This is necessary because of the potential race condition where new | ||
| // storage blocks may steal the free execution memory that this task was waiting for. | ||
| maybeGrowPool(numBytes - memoryFree) | ||
|
|
||
| // Maximum size the pool would have after potentially growing the pool. | ||
| // This is used to compute the upper bound of how much memory each task can occupy. This | ||
| // must take into account potential free memory as well as the amount this pool currently | ||
| // occupies. Otherwise, we may run into SPARK-12155 where, in unified memory management, | ||
| // we did not take into account space that could have been freed by evicting cached blocks. | ||
| val maxPoolSize = computeMaxPoolSize() | ||
| val maxMemoryPerTask = maxPoolSize / numActiveTasks | ||
| val minMemoryPerTask = poolSize / (2 * numActiveTasks) | ||
|
|
||
| // How much we can grant this task; keep its share within 0 <= X <= 1 / numActiveTasks | ||
| val maxToGrant = math.min(numBytes, math.max(0, maxMemoryPerTask - curMem)) | ||
| // Only give it as much memory as is free, which might be none if it reached 1 / numTasks | ||
| val toGrant = math.min(maxToGrant, memoryFree) | ||
|
|
||
| // We want to let each task get at least 1 / (2 * numActiveTasks) before blocking; | ||
| // if we can't give it this much now, wait for other tasks to free up memory | ||
| // (this happens if older tasks allocated lots of memory before N grew) | ||
| if (toGrant < numBytes && curMem + toGrant < minMemoryPerTask) { | ||
| logInfo(log"TID ${MDC(TASK_ATTEMPT_ID, taskAttemptId)} waiting for at least 1/2N of" + | ||
| log" ${MDC(POOL_NAME, poolName)} pool to be free") | ||
| lock.wait() | ||
| } else { | ||
| memoryForTask(taskAttemptId) += toGrant | ||
| return toGrant | ||
| var registeredWaiter = false | ||
| try { | ||
| while (true) { | ||
| val numActiveTasks = memoryForTask.keys.size | ||
| val curMem = memoryForTask(taskAttemptId) | ||
|
|
||
| // In every iteration of this loop, we should first try to reclaim any borrowed execution | ||
| // space from storage. This is necessary because of the potential race condition where new | ||
| // storage blocks may steal the free execution memory that this task was waiting for. | ||
| maybeGrowPool(numBytes - memoryFree) | ||
|
|
||
| // Maximum size the pool would have after potentially growing the pool. | ||
| // This is used to compute the upper bound of how much memory each task can occupy. This | ||
| // must take into account potential free memory as well as the amount this pool currently | ||
| // occupies. Otherwise, we may run into SPARK-12155 where, in unified memory management, | ||
| // we did not take into account space that could have been freed by evicting cached blocks. | ||
| val maxPoolSize = computeMaxPoolSize() | ||
| val maxMemoryPerTask = maxPoolSize / numActiveTasks | ||
| val minMemoryPerTask = poolSize / (2 * numActiveTasks) | ||
|
|
||
| // How much we can grant this task; keep its share within 0 <= X <= 1 / numActiveTasks | ||
| val maxToGrant = math.min(numBytes, math.max(0, maxMemoryPerTask - curMem)) | ||
| // Only give it as much memory as is free, which might be none if it reached 1 / numTasks | ||
| val toGrant = math.min(maxToGrant, memoryFree) | ||
|
|
||
| // We want to let each task get at least 1 / (2 * numActiveTasks) before blocking; | ||
| // if we can't give it this much now, wait for other tasks to free up memory | ||
| // (this happens if older tasks allocated lots of memory before N grew) | ||
| if (toGrant < numBytes && curMem + toGrant < minMemoryPerTask) { | ||
| logInfo(log"TID ${MDC(TASK_ATTEMPT_ID, taskAttemptId)} waiting for at least 1/2N of" + | ||
| log" ${MDC(POOL_NAME, poolName)} pool to be free") | ||
| if (!registeredWaiter) { | ||
| if (waitingAcquisitions == null) { | ||
| waitingAcquisitions = mutable.LongMap.empty[Int] | ||
| } | ||
| waitingAcquisitions(taskAttemptId) = | ||
| waitingAcquisitions.getOrElse(taskAttemptId, 0) + 1 | ||
| registeredWaiter = true | ||
| } | ||
| lock.wait() | ||
| } else { | ||
| memoryForTask(taskAttemptId) += toGrant | ||
| return toGrant | ||
| } | ||
| } | ||
| 0L // Never reached | ||
| } finally { | ||
| if (registeredWaiter) { | ||
| val remaining = waitingAcquisitions(taskAttemptId) - 1 | ||
| if (remaining == 0) { | ||
| waitingAcquisitions.remove(taskAttemptId) | ||
| if (waitingAcquisitions.isEmpty) { | ||
| waitingAcquisitions = null | ||
| } | ||
| } else { | ||
| waitingAcquisitions(taskAttemptId) = remaining | ||
| } | ||
| // The last exiting waiter must not leave a zero-byte fairness participant. | ||
| if (!hasWaitingAcquisition(taskAttemptId) && | ||
| memoryForTask.get(taskAttemptId).contains(0L)) { | ||
| memoryForTask.remove(taskAttemptId) | ||
| lock.notifyAll() | ||
| } | ||
| } | ||
| } | ||
| 0L // Never reached | ||
| } | ||
|
|
||
| /** | ||
|
|
@@ -164,15 +202,16 @@ private[memory] class ExecutionMemoryPool( | |
| } | ||
| if (memoryForTask.contains(taskAttemptId)) { | ||
| memoryForTask(taskAttemptId) -= memoryToFree | ||
| if (memoryForTask(taskAttemptId) <= 0) { | ||
| if (memoryForTask(taskAttemptId) <= 0 && !hasWaitingAcquisition(taskAttemptId)) { | ||
| memoryForTask.remove(taskAttemptId) | ||
| } | ||
| } | ||
| lock.notifyAll() // Notify waiters in acquireMemory() that memory has been freed | ||
| } | ||
|
|
||
| /** | ||
| * Release all memory for the given task and mark it as inactive (e.g. when a task ends). | ||
| * Release all memory for the given task. A task with a waiting acquisition remains active | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Could we make the lifecycle conditions more explicit here? A successful acquisition can leave the task active because it now holds memory. Perhaps state that this releases currently reserved memory without canceling pending acquisitions, and that the task entry is removed only when it has neither reserved memory nor waiting acquisitions. The corresponding comment in |
||
| * until that acquisition completes or is interrupted. | ||
| * @return the number of bytes freed. | ||
| */ | ||
| def releaseAllMemoryForTask(taskAttemptId: Long): Long = lock.synchronized { | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,242 @@ | ||
| /* | ||
| * 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. | ||
| */ | ||
|
|
||
| package org.apache.spark.memory | ||
|
|
||
| import java.util.concurrent.{CompletableFuture, ExecutionException, TimeUnit} | ||
| import java.util.concurrent.atomic.AtomicBoolean | ||
|
|
||
| import scala.collection.mutable.ArrayBuffer | ||
| import scala.util.control.NonFatal | ||
|
|
||
| import org.scalatest.concurrent.Eventually | ||
| import org.scalatest.time.SpanSugar._ | ||
|
|
||
| import org.apache.spark.{SparkConf, SparkFunSuite} | ||
|
|
||
| class ExecutionMemoryPoolSuite extends SparkFunSuite with Eventually { | ||
| private class WaitingAcquire(acquire: () => Long) { | ||
| val result = new CompletableFuture[Long]() | ||
| val thread = new Thread("execution-memory-waiter") { | ||
| override def run(): Unit = { | ||
| try { | ||
| result.complete(acquire()) | ||
| } catch { | ||
| case e: InterruptedException => result.completeExceptionally(e) | ||
| case NonFatal(e) => result.completeExceptionally(e) | ||
| } | ||
| } | ||
| } | ||
| thread.setDaemon(true) | ||
|
|
||
| def awaitWaiting(): Unit = eventually(timeout(10.seconds)) { | ||
| assert(!result.isDone, "acquisition completed instead of waiting") | ||
| assert(thread.getState == Thread.State.WAITING) | ||
| assert(thread.getStackTrace.exists { frame => | ||
| frame.getClassName == classOf[ExecutionMemoryPool].getName && | ||
| frame.getMethodName == "acquireMemory" | ||
| }) | ||
| } | ||
|
|
||
| def acquired(): Long = result.get(10, TimeUnit.SECONDS) | ||
|
|
||
| def failure(): Throwable = { | ||
| val error = intercept[ExecutionException] { | ||
| result.get(10, TimeUnit.SECONDS) | ||
| } | ||
| error.getCause | ||
| } | ||
|
|
||
| def interrupt(): Unit = { | ||
| thread.interrupt() | ||
| assert(failure().isInstanceOf[InterruptedException]) | ||
| } | ||
| } | ||
|
|
||
| private val waiters = new ArrayBuffer[WaitingAcquire]() | ||
|
|
||
| override protected def afterEach(): Unit = { | ||
| try { | ||
| waiters.foreach(_.thread.interrupt()) | ||
| waiters.foreach(_.thread.join(10000)) | ||
| assert(waiters.forall(!_.thread.isAlive), "memory acquisition thread did not terminate") | ||
| } finally { | ||
| waiters.clear() | ||
| super.afterEach() | ||
| } | ||
| } | ||
|
|
||
| private def acquireAsync( | ||
| pool: ExecutionMemoryPool, | ||
| bytes: Long, | ||
| maybeGrowPool: Long => Unit = _ => ()): WaitingAcquire = { | ||
| acquireAsync(pool.acquireMemory(bytes, 1L, maybeGrowPool)) | ||
| } | ||
|
|
||
| private def acquireAsync(acquire: => Long): WaitingAcquire = { | ||
| val waiter = new WaitingAcquire(() => acquire) | ||
| waiters += waiter | ||
| waiter.thread.start() | ||
| waiter.awaitWaiting() | ||
| waiter | ||
| } | ||
|
|
||
| private def newPool(mode: MemoryMode): ExecutionMemoryPool = { | ||
| val pool = new ExecutionMemoryPool(new Object, mode) | ||
| pool.incrementPoolSize(1000L) | ||
| assert(pool.acquireMemory(900L, 2L) == 900L) | ||
| pool | ||
| } | ||
|
|
||
| for (mode <- Seq(MemoryMode.ON_HEAP, MemoryMode.OFF_HEAP)) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Nit: |
||
| test(s"retain task registration across two consumers of one TaskMemoryManager ($mode)") { | ||
| val conf = new SparkConf(false) | ||
| .set("spark.memory.offHeap.enabled", "true") | ||
| .set("spark.memory.offHeap.size", "1000") | ||
| val memory = new UnifiedMemoryManager(conf, 1000L, 500L, 1) | ||
| val task = new TaskMemoryManager(memory, 1L) | ||
| val peerTask = new TaskMemoryManager(memory, 2L) | ||
| val owner = new TestMemoryConsumer(task, mode) | ||
| val requester = new TestMemoryConsumer(task, mode) | ||
| val peer = new TestMemoryConsumer(peerTask, mode) | ||
| assert(peer.acquireMemory(900L) == 900L) | ||
| assert(owner.acquireMemory(100L) == 100L) | ||
| val waiter = acquireAsync(requester.acquireMemory(300L)) | ||
|
|
||
| // Release through another consumer of the waiting task, not through the pool directly. | ||
| owner.freeMemory(100L) | ||
| peer.freeMemory(300L) | ||
| assert(waiter.acquired() == 300L) | ||
| assert(owner.getUsed() == 0L) | ||
| assert(requester.getUsed() == 300L) | ||
| assert(task.getMemoryConsumptionForThisTask() == 300L) | ||
| assert(peerTask.getMemoryConsumptionForThisTask() == 600L) | ||
| assert(memory.executionMemoryUsed == 900L) | ||
|
|
||
| requester.freeMemory(300L) | ||
| peer.freeMemory(600L) | ||
| assert(task.cleanUpAllAllocatedMemory() == 0L) | ||
| assert(peerTask.cleanUpAllAllocatedMemory() == 0L) | ||
| assert(memory.executionMemoryUsed == 0L) | ||
| } | ||
|
|
||
| for (releaseAll <- Seq(false, true)) { | ||
| test(s"retain a waiting task after its last release ($mode, releaseAll=$releaseAll)") { | ||
| val pool = newPool(mode) | ||
| assert(pool.acquireMemory(100L, 1L) == 100L) | ||
| val waiter = acquireAsync(pool, 300L) | ||
|
|
||
| if (releaseAll) { | ||
| assert(pool.releaseAllMemoryForTask(1L) == 100L) | ||
| } else { | ||
| pool.releaseMemory(100L, 1L) | ||
| } | ||
| pool.releaseMemory(300L, 2L) | ||
|
|
||
| assert(waiter.acquired() == 300L) | ||
| assert(pool.getMemoryUsageForTask(1L) == 300L) | ||
| assert(pool.memoryUsed == 900L) | ||
| } | ||
| } | ||
|
|
||
| test(s"retain a task until all its waiting acquisitions complete ($mode)") { | ||
| val pool = newPool(mode) | ||
| assert(pool.acquireMemory(100L, 1L) == 100L) | ||
| val first = acquireAsync(pool, 200L) | ||
| val second = acquireAsync(pool, 200L) | ||
| pool.releaseMemory(100L, 1L) | ||
| pool.releaseMemory(100L, 2L) | ||
|
|
||
| eventually(timeout(10.seconds)) { | ||
| assert(first.result.isDone || second.result.isDone) | ||
| } | ||
| val (completed, remaining) = if (first.result.isDone) (first, second) else (second, first) | ||
| assert(completed.acquired() == 200L) | ||
| remaining.awaitWaiting() | ||
| pool.releaseMemory(200L, 1L) | ||
|
|
||
| assert(remaining.acquired() == 200L) | ||
| assert(pool.getMemoryUsageForTask(1L) == 200L) | ||
| assert(pool.memoryUsed == 1000L) | ||
| } | ||
|
|
||
| test(s"preserve a waiting task's remaining allocation after a partial release ($mode)") { | ||
| val pool = newPool(mode) | ||
| assert(pool.acquireMemory(100L, 1L) == 100L) | ||
| val waiter = acquireAsync(pool, 300L) | ||
| pool.releaseMemory(40L, 1L) | ||
| pool.releaseMemory(300L, 2L) | ||
|
|
||
| assert(waiter.acquired() == 300L) | ||
| assert(pool.getMemoryUsageForTask(1L) == 360L) | ||
| assert(pool.memoryUsed == 960L) | ||
| } | ||
|
|
||
| for (previousAllocation <- Seq(false, true)) { | ||
| test(s"remove an interrupted zero-byte task ($mode, previous=$previousAllocation)") { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This case ( This also affects the "16 cases failed" negative control in the description: some of those failures are these new-behavior assertions, not the |
||
| val pool = newPool(mode) | ||
| if (previousAllocation) { | ||
| assert(pool.acquireMemory(100L, 1L) == 100L) | ||
| } | ||
| val waiter = acquireAsync(pool, 300L) | ||
| if (previousAllocation) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Could we also cover interrupting the waiter while the task still holds its original 100-byte reservation? Both variants here have zero reserved bytes by the time the waiter is interrupted. The additional case would verify that exiting the last waiter preserves an existing reservation until it is explicitly released. |
||
| pool.releaseMemory(100L, 1L) | ||
| } | ||
| waiter.interrupt() | ||
|
|
||
| // The interrupted task must no longer reduce the remaining task's fair share. | ||
| assert(pool.acquireMemory(100L, 2L) == 100L) | ||
| assert(pool.getMemoryUsageForTask(1L) == 0L) | ||
| assert(pool.releaseAllMemoryForTask(2L) == 1000L) | ||
| assert(pool.memoryUsed == 0L) | ||
| } | ||
| } | ||
|
|
||
| test(s"remove a zero-byte task if the pool-growth callback fails after waiting ($mode)") { | ||
| val pool = newPool(mode) | ||
| val failAfterWaiting = new AtomicBoolean(false) | ||
| val error = new IllegalStateException("pool-growth failure") | ||
| val waiter = acquireAsync(pool, 300L, _ => { | ||
| if (failAfterWaiting.get()) { | ||
| throw error | ||
| } | ||
| }) | ||
| failAfterWaiting.set(true) | ||
| pool.releaseMemory(0L, 2L) // Wake the waiter to run the callback again. | ||
|
|
||
| assert(waiter.failure() eq error) | ||
| assert(pool.acquireMemory(100L, 2L) == 100L) | ||
| assert(pool.releaseAllMemoryForTask(2L) == 1000L) | ||
| assert(pool.memoryUsed == 0L) | ||
| } | ||
|
|
||
| test(s"interrupting one acquisition must retain the same task's other waiter ($mode)") { | ||
| val pool = newPool(mode) | ||
| assert(pool.acquireMemory(100L, 1L) == 100L) | ||
| val interrupted = acquireAsync(pool, 300L) | ||
| val remaining = acquireAsync(pool, 300L) | ||
| pool.releaseMemory(100L, 1L) | ||
| interrupted.interrupt() | ||
| remaining.awaitWaiting() | ||
| pool.releaseMemory(300L, 2L) | ||
|
|
||
| assert(remaining.acquired() == 300L) | ||
| assert(pool.getMemoryUsageForTask(1L) == 300L) | ||
| assert(pool.memoryUsed == 900L) | ||
| } | ||
| } | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Through
TaskMemoryManager, at most one acquisition per task can be waiting in this pool at any time, becauseTaskMemoryManager.acquireExecutionMemoryissynchronized (this)and holds the TMM monitor acrosslock.wait(). For the same reasoncleanUpAllAllocatedMemory(alsosynchronized (this)) cannot run while a waiter exists, soreleaseAllMemoryForTasknever races with a waiter either. So the per-task counter, the lazy allocation / null reset, the multiple-waiter cases and thereleaseAll=truevariant only cover schedules that require calling thisprivate[memory]pool directly.A much smaller fix would be to re-register the entry at the top of the loop, e.g.
The trade-off is a transient fairness gap: between the removal and the wake-up another task may compute its share with N-1 tasks. That is not a correctness issue, and it keeps the existing lifecycle semantics unchanged. If you prefer to keep the current structure, a
mutable.HashSet[Long]allocated once would already be enough; the counter and the null handling do not buy anything through the public API.