dongjoon-hyun commented on code in PR #58747:
URL: https://github.com/apache/spark/pull/58747#discussion_r3998051507


##########
core/src/main/scala/org/apache/spark/memory/ExecutionMemoryPool.scala:
##########
@@ -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

Review Comment:
   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.
   
   ```scala
   // 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.



##########
core/src/test/scala/org/apache/spark/memory/ExecutionMemoryPoolSuite.scala:
##########
@@ -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)) {
+    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)") {

Review Comment:
   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.



##########
core/src/test/scala/org/apache/spark/memory/ExecutionMemoryPoolSuite.scala:
##########
@@ -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)) {

Review Comment:
   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.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to