From c37e1d20ae09342d744f98de823f7aca14b388ee Mon Sep 17 00:00:00 2001 From: David Mollitor Date: Mon, 21 Sep 2026 19:00:43 +0000 Subject: [PATCH] [SPARK-59697][CORE] Skip the redundant record comparator in the external-sort spill merge when the key prefix is a total order The external-sort spill merge (`UnsafeSorterSpillMerger` / `UnsafeSorterBoundedSpillMerger`, reached via `UnsafeExternalSorter.getSortedIterator()`) orders spill-run heads by the 8-byte key prefix and, on a prefix tie, falls back to the full `RecordComparator`, which decodes and byte-compares the records. When the sort qualifies for radix sort (`canUseRadixSort` -- a single, prefix-sortable key) AND that key is non-null, the prefix is a lossless, order-preserving total order over actual rows, so equal prefixes are equal keys and the tie-break always returns 0 -- dead work on every prefix collision. The non-null requirement matters: a null is encoded in the prefix as an in-range sentinel long (e.g. `Long.MinValue`) that can collide with a real key equal to that sentinel, and the spill-merge iterator carries only the prefix (no isNull), so the record comparator is the only thing that separates a null from an equal- prefix real value. The in-memory radix path is null-aware (RadixSortSupport `nullsFirst()`), but the merge is not -- so the tie-break may only be skipped for a non-null key. This threads `canUseRadixSort` (already computed, previously only forwarded to `UnsafeInMemorySorter`) plus the sort key's nullability into both merge paths, and skips the record comparator only when `canUseRadixSort && !nullable`. The merger now takes a `@Nullable RecordComparator`: null means "prefix is a total order, compare by prefix only", so no separate boolean is needed. It is the merge-side analogue of the null-aware in-memory radix optimization and is behavior-preserving; all other sorts (multi-key, strings, large decimals, or a nullable key) keep the record-comparator tie-break unchanged. Co-authored-by: Isaac --- .../unsafe/sort/UnsafeExternalSorter.java | 38 +++++++++++++--- .../sort/UnsafeSorterBoundedSpillMerger.java | 6 ++- .../unsafe/sort/UnsafeSorterSpillMerger.java | 36 ++++++++++----- .../sort/UnsafeExternalSorterSuite.java | 45 +++++++++++++++++-- .../execution/UnsafeExternalRowSorter.java | 15 ++++--- .../sql/execution/UnsafeKVExternalSorter.java | 3 +- .../ExternalAppendOnlyUnsafeRowArray.scala | 3 +- .../apache/spark/sql/execution/SortExec.scala | 3 +- ...nalAppendOnlyUnsafeRowArrayBenchmark.scala | 3 +- 9 files changed, 121 insertions(+), 31 deletions(-) diff --git a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorter.java b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorter.java index b725919cfb3e1..1060c0e5e048a 100644 --- a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorter.java +++ b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorter.java @@ -57,6 +57,16 @@ public final class UnsafeExternalSorter extends MemoryConsumer { @Nullable private final PrefixComparator prefixComparator; + // Whether the sort key is prefix-sortable (a single key whose prefix is a total order -- the + // condition that enables the in-memory radix sort). + private final boolean canUseRadixSort; + + // Whether that sort key may be null. A null is encoded in the prefix as an in-range sentinel that + // can collide with a real value, so when the key is nullable the record comparator is still + // required to break prefix ties in the spill merge; only a non-null prefix-sortable key makes the + // prefix a total order over actual rows (see getSortedIterator / prepareBoundedMerge). + private final boolean keyNullable; + /** * {@link RecordComparator} may probably keep the reference to the records they compared last * time, so we should not keep a {@link RecordComparator} instance inside @@ -127,7 +137,7 @@ public static UnsafeExternalSorter createWithExistingInMemorySorter( UnsafeExternalSorter sorter = new UnsafeExternalSorter(taskMemoryManager, blockManager, serializerManager, taskContext, recordComparatorSupplier, prefixComparator, initialSize, pageSizeBytes, numElementsForSpillThreshold, sizeInBytesForSpillThreshold, - spillMergeFactor, inMemorySorter, false /* ignored */); + spillMergeFactor, inMemorySorter, false /* ignored */, true /* keyNullable: ignored */); sorter.spill(Long.MAX_VALUE, sorter); taskContext.taskMetrics().incMemoryBytesSpilled(existingMemoryConsumption); sorter.totalSpillBytes += existingMemoryConsumption; @@ -148,11 +158,12 @@ public static UnsafeExternalSorter create( int numElementsForSpillThreshold, long sizeInBytesForSpillThreshold, int spillMergeFactor, - boolean canUseRadixSort) { + boolean canUseRadixSort, + boolean keyNullable) { return new UnsafeExternalSorter(taskMemoryManager, blockManager, serializerManager, taskContext, recordComparatorSupplier, prefixComparator, initialSize, pageSizeBytes, numElementsForSpillThreshold, sizeInBytesForSpillThreshold, spillMergeFactor, - null, canUseRadixSort); + null, canUseRadixSort, keyNullable); } private UnsafeExternalSorter( @@ -168,7 +179,8 @@ private UnsafeExternalSorter( long sizeInBytesForSpillThreshold, int spillMergeFactor, @Nullable UnsafeInMemorySorter existingInMemorySorter, - boolean canUseRadixSort) { + boolean canUseRadixSort, + boolean keyNullable) { super(taskMemoryManager, pageSizeBytes, taskMemoryManager.getTungstenMemoryMode()); this.taskMemoryManager = taskMemoryManager; this.blockManager = blockManager; @@ -176,6 +188,8 @@ private UnsafeExternalSorter( this.taskContext = taskContext; this.recordComparatorSupplier = recordComparatorSupplier; this.prefixComparator = prefixComparator; + this.canUseRadixSort = canUseRadixSort; + this.keyNullable = keyNullable; this.spillMergeFactor = spillMergeFactor; // Use getSizeAsKb (not bytes) to maintain backwards compatibility for units // this.fileBufferSizeBytes = (int) conf.getSizeAsKb("spark.shuffle.file.buffer", "32k") * 1024 @@ -578,6 +592,18 @@ public void merge(UnsafeExternalSorter other) throws IOException { other.cleanupResources(); } + /** + * The record comparator the spill merge uses to break ties between records with equal key + * prefixes, or {@code null} when the prefix is a total order over actual rows -- a single, + * non-null, prefix-sortable key ({@code canUseRadixSort && !keyNullable}). In that case equal + * prefixes are equal keys, so the tie-break is skipped. A nullable key keeps the comparator + * because a null encodes to an in-range sentinel prefix that can collide with a real value. + */ + @Nullable + private RecordComparator mergeTieBreakComparator() { + return (canUseRadixSort && !keyNullable) ? null : recordComparatorSupplier.get(); + } + /** * Returns a sorted iterator. It is the caller's responsibility to call `cleanupResources()` * after consuming this iterator. @@ -601,7 +627,7 @@ public UnsafeSorterIterator getSortedIterator() throws IOException { logger.info("Merging {} spill files in single round", MDC.of(LogKeys.NUM_SPILL_WRITERS, spillWriters.size())); final UnsafeSorterSpillMerger spillMerger = new UnsafeSorterSpillMerger( - recordComparatorSupplier.get(), prefixComparator, spillWriters.size()); + mergeTieBreakComparator(), prefixComparator, spillWriters.size()); for (UnsafeSorterSpillWriter spillWriter : spillWriters) { spillMerger.addSpillIfNotEmpty(spillWriter.getReader(serializerManager)); } @@ -652,7 +678,7 @@ BoundedMergerContext prepareBoundedMerge() { // blocks. final UnsafeSorterBoundedSpillMerger merger = new UnsafeSorterBoundedSpillMerger( spillMergeFactor, - recordComparatorSupplier.get(), + mergeTieBreakComparator(), prefixComparator, blockManager, serializerManager, diff --git a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterBoundedSpillMerger.java b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterBoundedSpillMerger.java index b844f9816bf3c..8e0bf61a9d6b2 100644 --- a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterBoundedSpillMerger.java +++ b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterBoundedSpillMerger.java @@ -59,7 +59,9 @@ final class UnsafeSorterBoundedSpillMerger { SparkLoggerFactory.getLogger(UnsafeSorterBoundedSpillMerger.class); private final int mergeFactor; - private final RecordComparator recordComparator; + // Null when the key prefix is a total order and the per-round mergers can skip the + // record-comparator tie-break on equal prefixes (see UnsafeSorterSpillMerger). + @Nullable private final RecordComparator recordComparator; private final PrefixComparator prefixComparator; private final BlockManager blockManager; private final SerializerManager serializerManager; @@ -71,7 +73,7 @@ final class UnsafeSorterBoundedSpillMerger { UnsafeSorterBoundedSpillMerger( int mergeFactor, - RecordComparator recordComparator, + @Nullable RecordComparator recordComparator, PrefixComparator prefixComparator, BlockManager blockManager, SerializerManager serializerManager, diff --git a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterSpillMerger.java b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterSpillMerger.java index f8603c5799e9b..6cc2549dfd798 100644 --- a/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterSpillMerger.java +++ b/core/src/main/java/org/apache/spark/util/collection/unsafe/sort/UnsafeSorterSpillMerger.java @@ -21,26 +21,40 @@ import java.util.Comparator; import java.util.PriorityQueue; +import javax.annotation.Nullable; + final class UnsafeSorterSpillMerger { private int numRecords = 0; private final PriorityQueue priorityQueue; + /** + * @param recordComparator breaks ties between records whose key prefixes are equal, or + * {@code null} when the key prefix is a total order (a single, non-null, prefix-sortable + * sort key -- the same precondition the in-memory radix sort relies on). When {@code null}, + * equal prefixes are equal keys, so the record-level tie-break is unnecessary and skipped. + */ UnsafeSorterSpillMerger( - RecordComparator recordComparator, + @Nullable RecordComparator recordComparator, PrefixComparator prefixComparator, int numSpills) { - Comparator comparator = (left, right) -> { - int prefixComparisonResult = + Comparator comparator; + if (recordComparator == null) { + comparator = (left, right) -> prefixComparator.compare(left.getKeyPrefix(), right.getKeyPrefix()); - if (prefixComparisonResult == 0) { - return recordComparator.compare( - left.getBaseObject(), left.getBaseOffset(), left.getRecordLength(), - right.getBaseObject(), right.getBaseOffset(), right.getRecordLength()); - } else { - return prefixComparisonResult; - } - }; + } else { + comparator = (left, right) -> { + int prefixComparisonResult = + prefixComparator.compare(left.getKeyPrefix(), right.getKeyPrefix()); + if (prefixComparisonResult == 0) { + return recordComparator.compare( + left.getBaseObject(), left.getBaseOffset(), left.getRecordLength(), + right.getBaseObject(), right.getBaseOffset(), right.getRecordLength()); + } else { + return prefixComparisonResult; + } + }; + } priorityQueue = new PriorityQueue<>(numSpills, comparator); } diff --git a/core/src/test/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorterSuite.java b/core/src/test/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorterSuite.java index 675070819a634..c25402adf7ce3 100644 --- a/core/src/test/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorterSuite.java +++ b/core/src/test/java/org/apache/spark/util/collection/unsafe/sort/UnsafeExternalSorterSuite.java @@ -173,7 +173,8 @@ private UnsafeExternalSorter newSorter() throws IOException { spillElementsThreshold, spillSizeThreshold, /* spillMergeFactor */ -1, - shouldUseRadixSort()); + shouldUseRadixSort(), + /* keyNullable */ false); } @Test @@ -200,6 +201,42 @@ public void testSortingOnlyByPrefix() throws Exception { assertSpillFilesWereCleanedUp(); } + @Test + public void testSortingWithDuplicatePrefixesAcrossSpills() throws Exception { + // The LONG prefix comparator makes the 8-byte prefix a total order for the key, so equal + // prefixes are equal keys. Insert many records with duplicated prefixes spread across several + // spill files to exercise the spill-merge tie-break on equal prefixes -- the path where a + // total-order prefix lets the merge skip the record comparator. Output must stay in + // non-decreasing prefix order with every record preserved, with or without radix sort. + final UnsafeExternalSorter sorter = newSorter(); + final int numDistinct = 8; + final int copiesPerBatch = 48; + final int numBatches = 6; + for (int batch = 0; batch < numBatches; batch++) { + for (int i = 0; i < copiesPerBatch; i++) { + insertNumber(sorter, i % numDistinct); + } + sorter.spill(); + } + + UnsafeSorterIterator iter = sorter.getSortedIterator(); + long previousPrefix = Long.MIN_VALUE; + int count = 0; + while (iter.hasNext()) { + iter.loadNext(); + final long prefix = iter.getKeyPrefix(); + assertTrue(prefix >= previousPrefix, "prefixes must be non-decreasing"); + // insertNumber writes the value as both the prefix and the 4-byte payload. + assertEquals(prefix, Platform.getInt(iter.getBaseObject(), iter.getBaseOffset())); + previousPrefix = prefix; + count++; + } + assertEquals(numBatches * copiesPerBatch, count); + + sorter.cleanupResources(); + assertSpillFilesWereCleanedUp(); + } + @Test public void testSortingEmptyArrays() throws Exception { final UnsafeExternalSorter sorter = newSorter(); @@ -465,7 +502,8 @@ public void forcedSpillingWithoutComparator() throws Exception { spillElementsThreshold, spillSizeThreshold, /* spillMergeFactor */ -1, - shouldUseRadixSort()); + shouldUseRadixSort(), + /* keyNullable */ false); long[] record = new long[100]; int recordSize = record.length * 8; int n = (int) pageSizeBytes / recordSize * 3; @@ -529,7 +567,8 @@ public void testPeakMemoryUsed() throws Exception { spillElementsThreshold, spillSizeThreshold, /* spillMergeFactor */ -1, - shouldUseRadixSort()); + shouldUseRadixSort(), + /* keyNullable */ false); // Peak memory should be monotonically increasing. More specifically, every time // we allocate a new page it should increase by exactly the size of the page. diff --git a/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeExternalRowSorter.java b/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeExternalRowSorter.java index 2f6d1b8c48ce4..6c00321921cf6 100644 --- a/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeExternalRowSorter.java +++ b/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeExternalRowSorter.java @@ -81,8 +81,10 @@ public static UnsafeExternalRowSorter createWithRecordComparator( UnsafeExternalRowSorter.PrefixComputer prefixComputer, long pageSizeBytes, boolean canUseRadixSort) throws IOException { + // Conservative: this path (e.g. ShuffleExchangeExec) does not supply the sort key's + // nullability, so treat the key as nullable and keep the record-comparator tie-break. return new UnsafeExternalRowSorter(schema, recordComparatorSupplier, prefixComparator, - prefixComputer, pageSizeBytes, canUseRadixSort); + prefixComputer, pageSizeBytes, canUseRadixSort, true /* keyNullable */); } public static UnsafeExternalRowSorter create( @@ -91,11 +93,12 @@ public static UnsafeExternalRowSorter create( PrefixComparator prefixComparator, UnsafeExternalRowSorter.PrefixComputer prefixComputer, long pageSizeBytes, - boolean canUseRadixSort) throws IOException { + boolean canUseRadixSort, + boolean keyNullable) throws IOException { Supplier recordComparatorSupplier = () -> new RowComparator(ordering, schema.length()); return new UnsafeExternalRowSorter(schema, recordComparatorSupplier, prefixComparator, - prefixComputer, pageSizeBytes, canUseRadixSort); + prefixComputer, pageSizeBytes, canUseRadixSort, keyNullable); } private UnsafeExternalRowSorter( @@ -104,7 +107,8 @@ private UnsafeExternalRowSorter( PrefixComparator prefixComparator, UnsafeExternalRowSorter.PrefixComputer prefixComputer, long pageSizeBytes, - boolean canUseRadixSort) { + boolean canUseRadixSort, + boolean keyNullable) { this.schema = schema; this.prefixComputer = prefixComputer; final SparkEnv sparkEnv = SparkEnv.get(); @@ -123,7 +127,8 @@ private UnsafeExternalRowSorter( (long) SparkEnv.get().conf().get( package$.MODULE$.SHUFFLE_SPILL_MAX_SIZE_FORCE_SPILL_THRESHOLD()), (int) sparkEnv.conf().get(package$.MODULE$.UNSAFE_SORTER_SPILL_MERGE_FACTOR()), - canUseRadixSort + canUseRadixSort, + keyNullable ); } diff --git a/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeKVExternalSorter.java b/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeKVExternalSorter.java index 73cc8a1e9f932..5a19f84a0a81b 100644 --- a/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeKVExternalSorter.java +++ b/sql/core/src/main/java/org/apache/spark/sql/execution/UnsafeKVExternalSorter.java @@ -102,7 +102,8 @@ public UnsafeKVExternalSorter( numElementsForSpillThreshold, sizeInBytesForSpillThreshold, (int) SparkEnv.get().conf().get(package$.MODULE$.UNSAFE_SORTER_SPILL_MERGE_FACTOR()), - canUseRadixSort); + canUseRadixSort, + keySchema.length() > 0 && keySchema.apply(0).nullable()); } else { // During spilling, the pointer array in `BytesToBytesMap` will not be used, so we can borrow // that and use it as the pointer array for `UnsafeInMemorySorter`. diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArray.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArray.scala index 5b9944772f16c..f0cb7e40edc8c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArray.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArray.scala @@ -155,7 +155,8 @@ class ExternalAppendOnlyUnsafeRowArray( numRowsSpillThreshold, sizeInBytesSpillThreshold, -1, // bounded merge not applicable — this class does not sort - false) + false, + false) // canUseRadixSort / keyNullable: unused, this class does not sort // populate with existing in-memory buffered rows if (inMemoryBuffer != null) { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SortExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SortExec.scala index d0f7a6bc32910..0589cd1024291 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SortExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SortExec.scala @@ -98,7 +98,8 @@ case class SortExec( val pageSize = SparkEnv.get.memoryManager.pageSizeBytes rowSorter = UnsafeExternalRowSorter.create( - schema, ordering, prefixComparator, prefixComputer, pageSize, canUseRadixSort) + schema, ordering, prefixComparator, prefixComputer, pageSize, canUseRadixSort, + sortOrder.head.child.nullable) if (testSpillFrequency > 0) { rowSorter.setTestSpillFrequency(testSpillFrequency) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArrayBenchmark.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArrayBenchmark.scala index fb42f0487c52d..caeda6ad795fc 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArrayBenchmark.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/ExternalAppendOnlyUnsafeRowArrayBenchmark.scala @@ -150,7 +150,8 @@ object ExternalAppendOnlyUnsafeRowArrayBenchmark extends BenchmarkBase { numSpillThreshold, Long.MaxValue, -1, // bounded merge not applicable — benchmark does not sort - false) + false, + false) // canUseRadixSort / keyNullable: unused, benchmark does not sort rows.foreach(x => array.insertRecord(