From e3a21a8762d86327eb79a12ede10dc64cba5b406 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 30 Sep 2026 11:29:15 +0800 Subject: [PATCH] fix: account for and release native shuffle spill ranges --- .../contributor-guide/native_shuffle.md | 9 ++- .../src/partitioners/multi_partition.rs | 26 ++++-- native/shuffle/src/shuffle_writer.rs | 21 ++++- .../writers/local/local_partition_writer.rs | 18 ++++- native/shuffle/src/writers/local/spill.rs | 79 ++++++++++++++++++- .../shuffle/src/writers/partition_writer.rs | 8 ++ .../comet/exec/CometNativeShuffleSuite.scala | 18 +++++ 7 files changed, 165 insertions(+), 14 deletions(-) diff --git a/docs/source/contributor-guide/native_shuffle.md b/docs/source/contributor-guide/native_shuffle.md index e1cabfd6c61..78e8ee99ba0 100644 --- a/docs/source/contributor-guide/native_shuffle.md +++ b/docs/source/contributor-guide/native_shuffle.md @@ -449,7 +449,8 @@ The `MultiPartitionShuffleRepartitioner` holds: - `buffered_batches`, a `Vec` of incoming batches, alongside `partition_indices` recording which rows of those batches belong to each partition. Rows are not copied into per-partition buffers as they arrive. -- `reservation`, a `MemoryReservation` charged for the bytes each buffered batch newly pins. +- `reservation`, an `Arc` shared with the local writer's spill metadata and + charged for the bytes each buffered batch newly pins. `pinned_buffers` tracks backing buffer start addresses so one allocation shared by many sliced batches is charged once rather than once per slice. Charging a per-batch size instead would overstate memory by the slice count and spill spuriously. @@ -457,6 +458,12 @@ The `MultiPartitionShuffleRepartitioner` holds: pressure as the only trigger. - `scratch`, reusable buffers for partition ID computation. +The local writer charges the allocated capacity of both the outer range table and each partition's +range vector to the shared reservation. Spilling releases the buffered input's charge while retaining +the metadata charge. After a partition's spilled blocks have been copied into the output, its range +vector is dropped and its charge released. The outer table remains charged until the writer drops. +The fixed `max_buffer_bytes` threshold counts buffered input only. + The spill file is owned by `PartitionedSpill` in `writers/local/spill.rs`. It holds one DataFusion `SpillFile`, created lazily on the first spill and shared by every output partition, and tracks per partition the byte ranges holding that partition's blocks in write order. A partial write sets a diff --git a/native/shuffle/src/partitioners/multi_partition.rs b/native/shuffle/src/partitioners/multi_partition.rs index 05b14b53243..ea2960b5357 100644 --- a/native/shuffle/src/partitioners/multi_partition.rs +++ b/native/shuffle/src/partitioners/multi_partition.rs @@ -114,8 +114,10 @@ pub(crate) struct MultiPartitionShuffleRepartitioner { /// The configured batch size batch_size: usize, /// Reservation for repartitioning - reservation: MemoryReservation, - /// Spill once the reservation reaches this many bytes, independently of whether the memory + reservation: Arc, + /// The portion of the shared reservation released when buffered input spills. + reserved_input_bytes: usize, + /// Spill once buffered input reaches this many bytes, independently of whether the memory /// pool still has capacity. `None` disables the limit, leaving pool pressure as the only /// spill trigger. max_buffer_bytes: Option, @@ -232,9 +234,13 @@ impl MultiPartitionShuffleRepartitioner { }, }; - let reservation = MemoryConsumer::new(format!("ShuffleRepartitioner[{partition}]")) - .with_can_spill(true) - .register(&runtime.memory_pool); + let reservation = partition_writer.memory_reservation().unwrap_or_else(|| { + Arc::new( + MemoryConsumer::new(format!("ShuffleRepartitioner[{partition}]")) + .with_can_spill(true) + .register(&runtime.memory_pool), + ) + }); Ok(Self { buffered_batches: vec![], @@ -249,6 +255,7 @@ impl MultiPartitionShuffleRepartitioner { scratch, batch_size, reservation, + reserved_input_bytes: 0, max_buffer_bytes, tracing_enabled, pinned_buffers: HashSet::new(), @@ -577,12 +584,15 @@ impl MultiPartitionShuffleRepartitioner { // A rejected reservation does not include this batch's memory, even though the batch // and its partition indices have already been buffered and must be counted as spilled. let reservation_failed = self.reservation.try_grow(mem_growth).is_err(); + if !reservation_failed { + self.reserved_input_bytes += mem_growth; + } // Checking after buffering lets the writer overshoot the limit by at most one batch, // which is how the memory-pressure trigger already behaves. if reservation_failed || self .max_buffer_bytes - .is_some_and(|limit| self.reservation.size() >= limit) + .is_some_and(|limit| self.reserved_input_bytes >= limit) { self.spill(if reservation_failed { mem_growth } else { 0 })?; } @@ -652,7 +662,9 @@ impl MultiPartitionShuffleRepartitioner { // rejected reservation. Shared allocations are charged once within a spill, but // contribute again if buffered for a later spill, regardless of input batching. // Also release and count all buffered inputs when the writer fails partway through. - let memory_spilled_bytes = self.reservation.free().saturating_add(unreserved_bytes); + let reserved_input_bytes = std::mem::take(&mut self.reserved_input_bytes); + self.reservation.shrink(reserved_input_bytes); + let memory_spilled_bytes = reserved_input_bytes.saturating_add(unreserved_bytes); self.metrics.memory_spilled_bytes.add(memory_spilled_bytes); self.pinned_buffers.clear(); self.metrics.spill_count.add(1); diff --git a/native/shuffle/src/shuffle_writer.rs b/native/shuffle/src/shuffle_writer.rs index 0e9085916de..6e55af33f7d 100644 --- a/native/shuffle/src/shuffle_writer.rs +++ b/native/shuffle/src/shuffle_writer.rs @@ -624,6 +624,8 @@ mod test { let num_partitions = 2; let runtime_env = create_runtime(memory_limit); let metrics_set = ExecutionPlanMetricsSet::new(); + let metrics = ShufflePartitionerMetrics::new(&metrics_set, 0); + let memory_spilled_bytes = metrics.memory_spilled_bytes.clone(); let shuffle_block_writer = ShuffleBlockWriter::try_new(batch.schema().as_ref(), CompressionCodec::Lz4Frame) .unwrap(); @@ -641,14 +643,15 @@ mod test { 0, local_partition_writer, CometPartitioning::Hash(vec![Arc::new(Column::new("a", 0))], num_partitions), - ShufflePartitionerMetrics::new(&metrics_set, 0), - runtime_env, + metrics, + Arc::clone(&runtime_env), 1024, false, None, ) .unwrap(); + let base_metadata_bytes = runtime_env.memory_pool.reserved(); repartitioner.insert_batch(batch.clone()).await.unwrap(); assert!(!repartitioner @@ -656,7 +659,14 @@ mod test { .get_spill() .has_spill_file()); + let before_spill = runtime_env.memory_pool.reserved(); repartitioner.spill(0).unwrap(); + let retained_bytes = runtime_env.memory_pool.reserved(); + assert!(retained_bytes > base_metadata_bytes); + assert_eq!( + memory_spilled_bytes.value(), + before_spill - base_metadata_bytes + ); // after spill, both partitions' blocks are in the one spill file { @@ -668,6 +678,13 @@ mod test { // insert another batch after spilling repartitioner.insert_batch(batch.clone()).await.unwrap(); + repartitioner.spill(0).unwrap(); + // The second range fits the existing allocation, which stays charged across spills. + assert_eq!(runtime_env.memory_pool.reserved(), retained_bytes); + repartitioner.shuffle_write().unwrap(); + assert_eq!(runtime_env.memory_pool.reserved(), base_metadata_bytes); + drop(repartitioner); + assert_eq!(runtime_env.memory_pool.reserved(), 0); } /// The zstd context is reused within one encode burst but must not survive past it: a diff --git a/native/shuffle/src/writers/local/local_partition_writer.rs b/native/shuffle/src/writers/local/local_partition_writer.rs index d76ff23eff0..4a61385b76d 100644 --- a/native/shuffle/src/writers/local/local_partition_writer.rs +++ b/native/shuffle/src/writers/local/local_partition_writer.rs @@ -23,6 +23,7 @@ use crate::writers::BufBatchWriter; use crate::{PartitionOffsets, ShuffleBlockWriter}; use arrow::array::RecordBatch; use datafusion::common::DataFusionError; +use datafusion::execution::memory_pool::MemoryReservation; use datafusion::execution::runtime_env::RuntimeEnv; use std::fs::{File, OpenOptions}; use std::io::{BufWriter, ErrorKind, Read, Seek, SeekFrom, Write}; @@ -130,6 +131,7 @@ impl LocalPartitionWriter { write_buffer_size, batch_size, num_output_partitions, + &runtime, ); DataOutput::Multi { output_writer, @@ -172,6 +174,13 @@ impl LocalPartitionWriter { } impl PartitionWriter for LocalPartitionWriter { + fn memory_reservation(&self) -> Option> { + match &self.data_output { + DataOutput::Multi { spill, .. } => Some(spill.memory_reservation()), + DataOutput::Single { .. } => None, + } + } + fn write( &mut self, pid: usize, @@ -292,6 +301,7 @@ impl PartitionWriter for LocalPartitionWriter { } write_timer.stop(); } + spill.release_ranges(pid); // Write in memory batches to output data file. Each partition uses its // own writer so coalescing does not cross partition boundaries, but the @@ -595,6 +605,7 @@ mod tests { writer .finish_partition(pid, &mut vec![Ok(b)].into_iter(), &metrics) .unwrap(); + assert!(writer.get_spill().ranges(pid).unwrap().is_empty()); } writer.finish_all(&metrics).unwrap(); @@ -646,13 +657,15 @@ mod tests { fn finish_partition_fails_when_spill_file_is_truncated() { for write_buffer_size in [1 << 20, 64] { let dir = tempfile::tempdir().unwrap(); + let runtime = Arc::new(RuntimeEnv::default()); let mut writer = partition_writer_with( &test_batch(), 2, write_buffer_size, &dir, - Arc::new(RuntimeEnv::default()), + Arc::clone(&runtime), ); + let base_metadata_bytes = runtime.memory_pool.reserved(); let metrics = ShufflePartitionerMetrics::new(&ExecutionPlanMetricsSet::new(), 0); for pid in 0..2 { writer @@ -679,6 +692,9 @@ mod tests { err.to_string().contains("truncated"), "write buffer {write_buffer_size}: unexpected error: {err}" ); + assert!(runtime.memory_pool.reserved() > base_metadata_bytes); + drop(writer); + assert_eq!(runtime.memory_pool.reserved(), 0); } } } diff --git a/native/shuffle/src/writers/local/spill.rs b/native/shuffle/src/writers/local/spill.rs index ae4b8adfa44..1d166162686 100644 --- a/native/shuffle/src/writers/local/spill.rs +++ b/native/shuffle/src/writers/local/spill.rs @@ -21,6 +21,7 @@ use crate::writers::BufBatchWriter; use crate::ShuffleBlockWriter; use arrow::record_batch::RecordBatch; use datafusion::common::DataFusionError; +use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; use datafusion::execution::runtime_env::RuntimeEnv; use datafusion::execution::SpillFile as DfSpillFile; use datafusion::execution::SpillWriter as DfSpillWriter; @@ -64,6 +65,8 @@ pub(crate) struct PartitionedSpill { len: u64, /// Per partition, the spill file ranges holding its blocks, in write order. ranges: Vec>>, + /// Shared with the repartitioner's input buffers; range charges survive each spill. + reservation: Arc, /// Set when a write fails partway, after which `len` may not match the file. failed: bool, } @@ -74,14 +77,21 @@ impl PartitionedSpill { write_buffer_size: usize, batch_size: usize, num_partitions: usize, + runtime: &RuntimeEnv, ) -> Self { + let ranges = vec![Vec::new(); num_partitions]; + let reservation = MemoryConsumer::new("ShuffleSpill") + .with_can_spill(true) + .register(&runtime.memory_pool); + reservation.grow(ranges.capacity() * size_of::>>()); Self { shuffle_block_writer, write_buffer_size, batch_size, spill_file: None, len: 0, - ranges: vec![Vec::new(); num_partitions], + ranges, + reservation: Arc::new(reservation), failed: false, } } @@ -149,7 +159,7 @@ impl PartitionedSpill { if bytes_written > 0 { let start = self.len; self.len += bytes_written; - self.ranges[pid].push(start..self.len); + self.push_range(pid, start..self.len); } metrics .spilled_bytes @@ -167,6 +177,28 @@ impl PartitionedSpill { Ok(&self.ranges[pid]) } + pub(crate) fn memory_reservation(&self) -> Arc { + Arc::clone(&self.reservation) + } + + fn push_range(&mut self, pid: usize, range: Range) { + let ranges = &mut self.ranges[pid]; + let capacity = ranges.capacity(); + ranges.push(range); + let growth = (ranges.capacity() - capacity) * size_of::>(); + if growth != 0 { + // A spill must be able to record its own bookkeeping under memory pressure. + self.reservation.grow(growth); + } + } + + /// Releases the range allocation after a partition's spilled blocks have been copied. + pub(crate) fn release_ranges(&mut self, pid: usize) { + let bytes = self.ranges[pid].capacity() * size_of::>(); + self.ranges[pid] = Vec::new(); + self.reservation.shrink(bytes); + } + /// Writes buffered spill bytes to the spill file. pub(crate) fn flush(&mut self) -> datafusion::common::Result<()> { self.check_usable()?; @@ -335,13 +367,54 @@ mod tests { ShuffleBlockWriter::try_new(batch.schema_ref().as_ref(), CompressionCodec::None) .unwrap(); // batch_size below the row count so a write serializes into the scratch - PartitionedSpill::new(block_writer, 1 << 20, 10, num_partitions) + PartitionedSpill::new( + block_writer, + 1 << 20, + 10, + num_partitions, + &RuntimeEnv::default(), + ) } fn metrics() -> ShufflePartitionerMetrics { ShufflePartitionerMetrics::new(&ExecutionPlanMetricsSet::new(), 0) } + #[test] + fn merging_releases_range_capacity() { + let num_partitions = 16_000; + let rounds = 88; + let mut spill = partitioned_spill(&test_batch(), num_partitions); + // Model the metadata from many spill rounds without writing the payload to disk. + let base_bytes = spill.reservation.size(); + for pid in 0..num_partitions { + for round in 0..rounds { + spill.push_range(pid, round..round + 1); + } + } + let capacity = spill.ranges[0].capacity(); + let initial_bytes = num_partitions * capacity * size_of::>(); + assert_eq!(spill.reservation.size(), base_bytes + initial_bytes); + for pid in 0..num_partitions { + spill.release_ranges(pid); + assert_eq!(spill.ranges[pid].capacity(), 0); + assert_eq!( + spill.reservation.size(), + base_bytes + (num_partitions - pid - 1) * capacity * size_of::>() + ); + if pid + 1 < num_partitions { + assert_eq!(spill.ranges[pid + 1].len(), rounds as usize); + } + } + let retained_bytes: usize = spill + .ranges + .iter() + .map(|ranges| ranges.capacity() * size_of::>()) + .sum(); + assert_eq!(retained_bytes, 0); + println!("range capacity: {initial_bytes} bytes before merge, {retained_bytes} after"); + } + fn failing_write(spill: &mut PartitionedSpill, recycled: &mut Vec) { let mut iter = vec![ Ok(test_batch()), diff --git a/native/shuffle/src/writers/partition_writer.rs b/native/shuffle/src/writers/partition_writer.rs index 0fcec17c00d..9c054771558 100644 --- a/native/shuffle/src/writers/partition_writer.rs +++ b/native/shuffle/src/writers/partition_writer.rs @@ -17,6 +17,8 @@ use crate::metrics::ShufflePartitionerMetrics; use arrow::record_batch::RecordBatch; +use datafusion::execution::memory_pool::MemoryReservation; +use std::sync::Arc; /// Storage backend abstraction for shuffle partition output. /// @@ -33,6 +35,12 @@ use arrow::record_batch::RecordBatch; /// /// [`LocalPartitionWriter`]: crate::writers::local::local_partition_writer::LocalPartitionWriter pub(crate) trait PartitionWriter: Send { + /// Shares the writer's reservation with the repartitioner, if it reserves memory. + /// Both input buffers and retained metadata must count against the same fair allowance. + fn memory_reservation(&self) -> Option> { + None + } + /// Stages the batches from `iter` for partition `pid` without finalizing it. /// /// Used to stream single-partition output and to stage multi-partition diff --git a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala index 66e5ccbce94..e72ce6f1dad 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala @@ -1302,6 +1302,24 @@ class CometNativeShuffleSuite extends CometTestBase with AdaptiveSparkPlanHelper } } + test("native shuffle: spill metadata can exceed the fair memory limit") { + withParquetTable((0 until 1000).map(i => (i, i.toLong)), "tbl") { + withSQLConf( + CometConf.COMET_OFFHEAP_MEMORY_POOL_TYPE.key -> "fair_unified", + CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.00000001", + CometConf.COMET_BATCH_SIZE.key -> "64") { + // The fair allowance is smaller than the range table itself. Every input batch spills, + // but recording spill metadata must still succeed and preserve every output row. + val shuffled = sql("SELECT * FROM tbl").repartition(10, $"_1") + checkShuffleAnswer(shuffled, 1) + shuffled.collect() + assert(collectFirst(shuffled.queryExecution.executedPlan) { + case e: CometShuffleExchangeExec => e.metrics("spill_count").value + }.exists(_ > 0)) + } + } + } + test("native shuffle: round robin partitioning") { withSQLConf(CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "true") { withParquetTable((0 until 100).map(i => (i, (i + 1).toLong, s"str$i")), "tbl") {