Skip to content
Open
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
111 changes: 75 additions & 36 deletions core/src/main/scala/org/apache/spark/memory/ExecutionMemoryPool.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown
Member

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, because TaskMemoryManager.acquireExecutionMemory is synchronized (this) and holds the TMM monitor across lock.wait(). For the same reason cleanUpAllAllocatedMemory (also synchronized (this)) cannot run while a waiter exists, so releaseAllMemoryForTask never races with a waiter either. So the per-task counter, the lazy allocation / null reset, the multiple-waiter cases and the releaseAll=true variant only cover schedules that require calling this private[memory] pool directly.

A much smaller fix would be to re-register the entry at the top of the loop, e.g.

// The entry may have been removed by a concurrent release of this task's last byte
// while we were waiting. Re-register so the accounting below stays consistent.
val curMem = memoryForTask.getOrElseUpdate(taskAttemptId, 0L)
val numActiveTasks = memoryForTask.keys.size

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.


private def hasWaitingAcquisition(taskAttemptId: Long): Boolean = {
waitingAcquisitions != null && waitingAcquisitions.contains(taskAttemptId)
}

override def memoryUsed: Long = lock.synchronized {
memoryForTask.values.sum
}
Expand Down Expand Up @@ -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
}

/**
Expand All @@ -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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The 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 MemoryManager should use the same wording.

* until that acquisition completes or is interrupted.
* @return the number of bytes freed.
*/
def releaseAllMemoryForTask(taskAttemptId: Long): Long = lock.synchronized {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -143,7 +143,8 @@ private[spark] abstract class MemoryManager(
}

/**
* 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
* until that acquisition completes or is interrupted.
*
* @return the number of bytes freed.
*/
Expand Down
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)) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nit: ExecutionMemoryPool does not branch on MemoryMode at all (it only affects the pool name), so running every pool-level case in both modes doubles the runtime without adding coverage. Only the UnifiedMemoryManager case above benefits from the mode loop.

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)") {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This case (previous=false) and the pool-growth-callback case below assert a behavior change rather than the regression. On master, a task that registers with 0 bytes and is then interrupted stays active until releaseAllMemoryForTask is called; that is by design, and it is still what happens for a non-waiting acquisition that is granted 0 bytes. With the finally block, the same task is removed only if it happened to wait first, so the lifecycle now depends on whether the acquisition waited.

This also affects the "16 cases failed" negative control in the description: some of those failures are these new-behavior assertions, not the NoSuchElementException. If we go with the smaller fix, I would drop these two cases; otherwise the description should distinguish them.

val pool = newPool(mode)
if (previousAllocation) {
assert(pool.acquireMemory(100L, 1L) == 100L)
}
val waiter = acquireAsync(pool, 300L)
if (previousAllocation) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The 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)
}
}
}