This is an automated email from the ASF dual-hosted git repository.

chengpan pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/incubator-celeborn.git


The following commit(s) were added to refs/heads/main by this push:
     new 4c67325a3 [CELEBORN-720][SPARK] Correct metric peakExecutionMemory of 
SortBasedShuffleWriter
4c67325a3 is described below

commit 4c67325a3d8a48be5791548a5ce45ab358974948
Author: Angerszhuuuu <[email protected]>
AuthorDate: Tue Jun 27 18:40:06 2023 +0800

    [CELEBORN-720][SPARK] Correct metric peakExecutionMemory of 
SortBasedShuffleWriter
    
    ### What changes were proposed in this pull request?
    Currently SortBasedShuffleWriter won't update peakMemoryUsedBytes, this pr 
support this.
    
    ### Why are the changes needed?
    
    ### Does this PR introduce _any_ user-facing change?
    
    ### How was this patch tested?
    
    Closes #1632 from AngersZhuuuu/CELEBORN-720.
    
    Authored-by: Angerszhuuuu <[email protected]>
    Signed-off-by: Cheng Pan <[email protected]>
---
 .../spark/shuffle/celeborn/SortBasedPusher.java    | 28 +++++++++++++++++++-
 .../shuffle/celeborn/SortBasedShuffleWriter.java   | 30 +++++++++++++++++++++-
 .../shuffle/celeborn/SortBasedShuffleWriter.java   | 30 ++++++++++++++++++++--
 3 files changed, 84 insertions(+), 4 deletions(-)

diff --git 
a/client-spark/common/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedPusher.java
 
b/client-spark/common/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedPusher.java
index 573a02b19..f534ba3c7 100644
--- 
a/client-spark/common/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedPusher.java
+++ 
b/client-spark/common/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedPusher.java
@@ -44,6 +44,9 @@ public class SortBasedPusher extends MemoryConsumer {
 
   private static final Logger logger = 
LoggerFactory.getLogger(SortBasedPusher.class);
 
+  /** Peak memory used by this sorter so far, in bytes. * */
+  private long peakMemoryUsedBytes;
+
   private ShuffleInMemorySorter inMemSorter;
   private final LinkedList<MemoryBlock> allocatedPages = new LinkedList<>();
   private MemoryBlock currentPage = null;
@@ -135,7 +138,8 @@ public class SortBasedPusher extends MemoryConsumer {
     this.pushSortMemoryThreshold = pushSortMemoryThreshold;
 
     int initialSize = Math.min((int) pushSortMemoryThreshold / 8, 1024 * 1024);
-    inMemSorter = new ShuffleInMemorySorter(this, initialSize);
+    this.inMemSorter = new ShuffleInMemorySorter(this, initialSize);
+    this.peakMemoryUsedBytes = getMemoryUsage();
     this.sharedPushLock = sharedPushLock;
     this.executorService = executorService;
   }
@@ -395,7 +399,29 @@ public class SortBasedPusher extends MemoryConsumer {
     return 0;
   }
 
+  private long getMemoryUsage() {
+    long totalPageSize = 0;
+    for (MemoryBlock page : allocatedPages) {
+      totalPageSize += page.size();
+    }
+    return ((inMemSorter == null) ? 0 : inMemSorter.getMemoryUsage()) + 
totalPageSize;
+  }
+
+  private void updatePeakMemoryUsed() {
+    long mem = getMemoryUsage();
+    if (mem > peakMemoryUsedBytes) {
+      peakMemoryUsedBytes = mem;
+    }
+  }
+
+  /** Return the peak memory used so far, in bytes. */
+  long getPeakMemoryUsedBytes() {
+    updatePeakMemoryUsed();
+    return peakMemoryUsedBytes;
+  }
+
   private long freeMemory() {
+    updatePeakMemoryUsed();
     long memoryFreed = 0;
     for (MemoryBlock block : allocatedPages) {
       memoryFreed += block.size();
diff --git 
a/client-spark/spark-2/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
 
b/client-spark/spark-2/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
index e4f587cca..75fb3c834 100644
--- 
a/client-spark/spark-2/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
+++ 
b/client-spark/spark-2/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
@@ -251,6 +251,34 @@ public class SortBasedShuffleWriter<K, V, C> extends 
ShuffleWriter<K, V> {
     }
   }
 
+  private void updatePeakMemoryUsed() {
+    // sorter can be null if this writer is closed
+    if (pipelined) {
+      for (int i = 0; i < pushers.length; i++) {
+
+        if (pushers[i] != null) {
+          long mem = pushers[i].getPeakMemoryUsedBytes();
+          if (mem > peakMemoryUsedBytes) {
+            peakMemoryUsedBytes = mem;
+          }
+        }
+      }
+    } else {
+      if (currentPusher != null) {
+        long mem = currentPusher.getPeakMemoryUsedBytes();
+        if (mem > peakMemoryUsedBytes) {
+          peakMemoryUsedBytes = mem;
+        }
+      }
+    }
+  }
+
+  /** Return the peak memory used so far, in bytes. */
+  public long getPeakMemoryUsedBytes() {
+    updatePeakMemoryUsed();
+    return peakMemoryUsedBytes;
+  }
+
   private void write0(scala.collection.Iterator iterator) throws IOException {
     final scala.collection.Iterator<Product2<K, ?>> records = iterator;
 
@@ -355,7 +383,7 @@ public class SortBasedShuffleWriter<K, V, C> extends 
ShuffleWriter<K, V> {
   @Override
   public Option<MapStatus> stop(boolean success) {
     try {
-      taskContext.taskMetrics().incPeakExecutionMemory(peakMemoryUsedBytes);
+      
taskContext.taskMetrics().incPeakExecutionMemory(getPeakMemoryUsedBytes());
       if (stopping) {
         return Option.apply(null);
       } else {
diff --git 
a/client-spark/spark-3/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
 
b/client-spark/spark-3/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
index b4ff196fd..b28ba8a81 100644
--- 
a/client-spark/spark-3/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
+++ 
b/client-spark/spark-3/src/main/java/org/apache/spark/shuffle/celeborn/SortBasedShuffleWriter.java
@@ -74,7 +74,6 @@ public class SortBasedShuffleWriter<K, V, C> extends 
ShuffleWriter<K, V> {
   private final boolean pipelined;
   private final SortBasedPusher[] pushers = new SortBasedPusher[2];
   private SortBasedPusher currentPusher;
-  // TODO it isn't be updated after initialization
   private long peakMemoryUsedBytes = 0;
 
   private final OpenByteArrayOutputStream serBuffer;
@@ -189,6 +188,33 @@ public class SortBasedShuffleWriter<K, V, C> extends 
ShuffleWriter<K, V> {
         executorService);
   }
 
+  private void updatePeakMemoryUsed() {
+    // sorter can be null if this writer is closed
+    if (pipelined) {
+      for (SortBasedPusher pusher : pushers) {
+        if (pusher != null) {
+          long mem = pusher.getPeakMemoryUsedBytes();
+          if (mem > peakMemoryUsedBytes) {
+            peakMemoryUsedBytes = mem;
+          }
+        }
+      }
+    } else {
+      if (currentPusher != null) {
+        long mem = currentPusher.getPeakMemoryUsedBytes();
+        if (mem > peakMemoryUsedBytes) {
+          peakMemoryUsedBytes = mem;
+        }
+      }
+    }
+  }
+
+  /** Return the peak memory used so far, in bytes. */
+  public long getPeakMemoryUsedBytes() {
+    updatePeakMemoryUsed();
+    return peakMemoryUsedBytes;
+  }
+
   @Override
   public void write(scala.collection.Iterator<Product2<K, V>> records) throws 
IOException {
     if (canUseFastWrite()) {
@@ -377,7 +403,7 @@ public class SortBasedShuffleWriter<K, V, C> extends 
ShuffleWriter<K, V> {
   @Override
   public Option<MapStatus> stop(boolean success) {
     try {
-      taskContext.taskMetrics().incPeakExecutionMemory(peakMemoryUsedBytes);
+      
taskContext.taskMetrics().incPeakExecutionMemory(getPeakMemoryUsedBytes());
 
       if (stopping) {
         return Option.empty();

Reply via email to